Add cnum32_union() and cnum64_union() to compute a smallest circular range containing both inputs. Do so by enumerating the following configurations in a rotated frame (a.base at the origin): 0 UT_MAX |---------------------------------------------------| [= a ==============================] | [= b tail =] [= b main ===================> [= union ===========================================] 0 UT_MAX |---------------------------------------------------| [= a =====================] | [= b tail ======] [= b main =====> [= union tail ============] [= union main => 0 UT_MAX |---------------------------------------------------| [= a =======] | [= b tail =============] [= b main =====> [= union tail =========] [= union main => 0 UT_MAX |---------------------------------------------------| [= a =========================] | | [= b ====================] | [= union ==================================] | 0 UT_MAX |---------------------------------------------------| [= a =====================================] | | [= b ========] | [= union =================================] | 0 UT_MAX |---------------------------------------------------| [= a ============] | | [= b =======] | Two possible covering arcs: [= ab ======================================] | [= ba tail ======] [= ba main =========> Pick the smaller one. Signed-off-by: Eduard Zingerman --- include/linux/cnum.h | 2 ++ kernel/bpf/cnum_defs.h | 44 ++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 46 insertions(+) diff --git a/include/linux/cnum.h b/include/linux/cnum.h index 49b7d0c7645d..ddf4726841e6 100644 --- a/include/linux/cnum.h +++ b/include/linux/cnum.h @@ -39,6 +39,7 @@ u32 cnum32_umin(struct cnum32 cnum); u32 cnum32_umax(struct cnum32 cnum); s32 cnum32_smin(struct cnum32 cnum); s32 cnum32_smax(struct cnum32 cnum); +struct cnum32 cnum32_union(struct cnum32 a, struct cnum32 b); struct cnum32 cnum32_intersect(struct cnum32 a, struct cnum32 b); void cnum32_intersect_with(struct cnum32 *dst, struct cnum32 src); void cnum32_intersect_with_urange(struct cnum32 *dst, u32 min, u32 max); @@ -65,6 +66,7 @@ u64 cnum64_umin(struct cnum64 cnum); u64 cnum64_umax(struct cnum64 cnum); s64 cnum64_smin(struct cnum64 cnum); s64 cnum64_smax(struct cnum64 cnum); +struct cnum64 cnum64_union(struct cnum64 a, struct cnum64 b); struct cnum64 cnum64_intersect(struct cnum64 a, struct cnum64 b); void cnum64_intersect_with(struct cnum64 *dst, struct cnum64 src); void cnum64_intersect_with_urange(struct cnum64 *dst, u64 min, u64 max); diff --git a/kernel/bpf/cnum_defs.h b/kernel/bpf/cnum_defs.h index 30685e43de04..05b17a8b4814 100644 --- a/kernel/bpf/cnum_defs.h +++ b/kernel/bpf/cnum_defs.h @@ -198,6 +198,50 @@ static inline struct cnum_t FN(normalize)(struct cnum_t cnum) return cnum; } +/* + * Return a smallest arc containing both 'a' and 'b'. + * Break equal-size ties by choosing the smaller base. + */ +struct cnum_t FN(union)(struct cnum_t a, struct cnum_t b) +{ + struct cnum_t b1, ab, ba; + ut end; + + if (FN(is_empty)(a)) + return b; + if (FN(is_empty)(b)) + return a; + + /* + * Rotate so that a1.base == 0 and a1.end == a.size. + * Normalize b1 to preserve the full-circle representation. + */ + b1 = FN(normalize)((struct cnum_t){ b.base - a.base, b.size }); + end = max(a.size, (ut)(b1.base + b1.size)); + + if (FN(urange_overflow)(b1)) { + /* a1 reaches b1's main arc: together they cover the circle. */ + if (b1.base <= a.size) + return (struct cnum_t){ 0, UT_MAX }; + + /* Extend b1's tail through a1's end, then rotate back. */ + return FN(normalize)((struct cnum_t){ b.base, end - b1.base }); + } + + /* ab, rotated back, covers both nonwrapping arcs. */ + ab = (struct cnum_t){ a.base, end }; + if (b1.base <= a.size) + return FN(normalize)(ab); + + /* The arcs are disjoint; ba is the other possible covering arc. */ + ba = (struct cnum_t){ b.base, a.size - b1.base }; + if (ba.size < ab.size || + (ba.size == ab.size && ba.base < ab.base)) + ab = ba; + + return FN(normalize)(ab); +} + struct cnum_t FN(add)(struct cnum_t a, struct cnum_t b) { if (FN(is_empty)(a) || FN(is_empty)(b)) -- 2.53.0