#ifndef T
#error "Define T (bit width: 32, 64) before including cnum_defs.h"
#endif
#include <linux/cnum.h>
#include <linux/kernel.h>
#include <linux/limits.h>
#include <linux/minmax.h>
#include <linux/compiler_types.h>
#define cnum_t __PASTE(cnum, T)
#define ut __PASTE(u, T)
#define st __PASTE(s, T)
#define UT_MAX __PASTE(__PASTE(U, T), _MAX)
#define ST_MAX __PASTE(__PASTE(S, T), _MAX)
#define ST_MIN __PASTE(__PASTE(S, T), _MIN)
#define EMPTY __PASTE(__PASTE(CNUM, T), _EMPTY)
#define FN(name) __PASTE(__PASTE(cnum, T), __PASTE(_, name))
struct cnum_t FN(from_urange)(ut min, ut max)
{
return (struct cnum_t){ .base = min, .size = (ut)max - min };
}
struct cnum_t FN(from_srange)(st min, st max)
{
ut size = (ut)max - (ut)min;
ut base = size == UT_MAX ? 0 : (ut)min;
return (struct cnum_t){ .base = base, .size = size };
}
static inline bool FN(urange_overflow)(struct cnum_t cnum)
{
return cnum.size > UT_MAX - (ut)cnum.base;
}
ut FN(umin)(struct cnum_t cnum)
{
return FN(urange_overflow)(cnum) ? 0 : cnum.base;
}
EXPORT_SYMBOL_GPL(FN(umin));
ut FN(umax)(struct cnum_t cnum)
{
return FN(urange_overflow)(cnum) ? UT_MAX : cnum.base + cnum.size;
}
EXPORT_SYMBOL_GPL(FN(umax));
static inline bool FN(srange_overflow)(struct cnum_t cnum)
{
return FN(contains)(cnum, (ut)ST_MAX) && FN(contains)(cnum, (ut)ST_MIN);
}
st FN(smin)(struct cnum_t cnum)
{
return FN(srange_overflow)(cnum)
? ST_MIN
: min((st)cnum.base, (st)(cnum.base + cnum.size));
}
st FN(smax)(struct cnum_t cnum)
{
return FN(srange_overflow)(cnum)
? ST_MAX
: max((st)cnum.base, (st)(cnum.base + cnum.size));
}
struct cnum_t FN(intersect)(struct cnum_t a, struct cnum_t b)
{
struct cnum_t b1;
ut dbase;
if (FN(is_empty)(a) || FN(is_empty)(b))
return EMPTY;
if (a.base > b.base)
swap(a, b);
dbase = b.base - a.base;
b1 = (struct cnum_t){ dbase, b.size };
if (FN(urange_overflow)(b1)) {
if (b1.base <= a.size) {
return a.size <= b.size ? a : b;
} else {
return (struct cnum_t) {
.base = a.base,
.size = min(a.size, (ut)(b1.base + b1.size)),
};
}
} else if (a.size >= b1.base) {
return (struct cnum_t) {
.base = b.base,
.size = min((ut)(a.size - dbase), b.size),
};
} else {
return EMPTY;
}
}
void FN(intersect_with)(struct cnum_t *dst, struct cnum_t src)
{
*dst = FN(intersect)(*dst, src);
}
void FN(intersect_with_urange)(struct cnum_t *dst, ut min, ut max)
{
FN(intersect_with)(dst, FN(from_urange)(min, max));
}
void FN(intersect_with_srange)(struct cnum_t *dst, st min, st max)
{
FN(intersect_with)(dst, FN(from_srange)(min, max));
}
static inline struct cnum_t FN(normalize)(struct cnum_t cnum)
{
if (cnum.size == UT_MAX && cnum.base != 0 && cnum.base != (ut)ST_MAX)
cnum.base = 0;
return cnum;
}
struct cnum_t FN(add)(struct cnum_t a, struct cnum_t b)
{
if (FN(is_empty)(a) || FN(is_empty)(b))
return EMPTY;
if (a.size > UT_MAX - b.size)
return (struct cnum_t){ 0, (ut)UT_MAX };
else
return FN(normalize)((struct cnum_t){ a.base + b.base, a.size + b.size });
}
struct cnum_t FN(negate)(struct cnum_t a)
{
if (FN(is_empty)(a))
return EMPTY;
return FN(normalize)((struct cnum_t){ -((ut)a.base + a.size), a.size });
}
bool FN(is_empty)(struct cnum_t cnum)
{
return cnum.base == EMPTY.base && cnum.size == EMPTY.size;
}
bool FN(contains)(struct cnum_t cnum, ut v)
{
if (FN(is_empty)(cnum))
return false;
if (FN(urange_overflow)(cnum))
return v >= cnum.base || v <= (ut)cnum.base + cnum.size;
else
return v >= cnum.base && v <= (ut)cnum.base + cnum.size;
}
bool FN(is_const)(struct cnum_t cnum)
{
return cnum.size == 0;
}
bool FN(is_subset)(struct cnum_t bigger, struct cnum_t smaller)
{
if (FN(is_empty(smaller)))
return true;
if (FN(is_empty(bigger)))
return false;
smaller.base -= bigger.base;
bigger.base = 0;
if (FN(urange_overflow)(smaller) && bigger.size < UT_MAX)
return false;
return smaller.base + smaller.size <= bigger.size;
}
#undef EMPTY
#undef cnum_t
#undef ut
#undef st
#undef UT_MAX
#undef ST_MAX
#undef ST_MIN
#undef FN