[PATCH v5 sched_ext/for-7.3 31/33] tools/sched_ext: Add three-mask cmask intersection iterator

Tejun Heo <[email protected]>
Newsgroups dev.linux.lists.sched-ext,org.kernel.vger.linux-kernel
Message-ID <[email protected]>
Add cmask_next_and2_set() and its round-robin wrap, extending
cmask_next_and_set() to a three-mask intersection: the next cid set in all
three masks at or after @start. A caller iterating the intersection of
three cmasks can then scan it in one pass, folding the third mask into the
word-level AND rather than skipping non-members one candidate at a time.

Signed-off-by: Tejun Heo <[email protected]>
---
 tools/sched_ext/include/scx/cid.bpf.h | 77 +++++++++++++++++++++++++++
 1 file changed, 77 insertions(+)

diff --git a/tools/sched_ext/include/scx/cid.bpf.h b/tools/sched_ext/include/scx/cid.bpf.h
index 6b0b4e41b288..69fb4e97bc77 100644
--- a/tools/sched_ext/include/scx/cid.bpf.h
+++ b/tools/sched_ext/include/scx/cid.bpf.h
@@ -668,6 +668,83 @@ static __always_inline u32 cmask_next_and_set_wrap(const struct scx_cmask __aren
 	return found < start ? found : a_end;
 }
 
+/*
+ * Like cmask_next_and_set() but over the intersection of THREE masks. Return
+ * a->base + a->nr_cids if no cid is set in all three at or after @start.
+ */
+static __always_inline u32 cmask_next_and2_set(const struct scx_cmask __arena *a,
+					       const struct scx_cmask __arena *b,
+					       const struct scx_cmask __arena *c,
+					       u32 start)
+{
+	u32 a_end = a->base + a->nr_cids;
+	u32 b_end = b->base + b->nr_cids;
+	u32 c_end = c->base + c->nr_cids;
+	u32 a_wbase = a->base / 64;
+	u32 b_wbase = b->base / 64;
+	u32 c_wbase = c->base / 64;
+	u32 lo = a->base > b->base ? a->base : b->base;
+	u32 hi = a_end < b_end ? a_end : b_end;
+	u32 last_wi, start_wi, start_bit, i;
+
+	lo = lo > c->base ? lo : c->base;
+	hi = hi < c_end ? hi : c_end;
+
+	if (lo >= hi)
+		return a_end;
+	if (start < lo)
+		start = lo;
+	if (start >= hi)
+		return a_end;
+
+	last_wi = (hi - 1) / 64;
+	start_wi = start / 64;
+	start_bit = start & 63;
+
+	bpf_for(i, 0, CMASK_MAX_WORDS) {
+		u32 abs_wi = start_wi + i;
+		u64 word;
+		u32 found;
+
+		if (abs_wi > last_wi)
+			break;
+
+		word = a->bits[abs_wi - a_wbase] & b->bits[abs_wi - b_wbase] &
+		       c->bits[abs_wi - c_wbase];
+		if (i == 0)
+			word &= GENMASK_U64(63, start_bit);
+		if (!word)
+			continue;
+
+		found = abs_wi * 64 + ctzll(word);
+		if (found >= hi)
+			return a_end;
+		return found;
+	}
+	return a_end;
+}
+
+/*
+ * Round-robin variant of cmask_next_and2_set(): wrap to @a->base if the
+ * three-way intersection has no cid in the forward half. Return a->base +
+ * a->nr_cids if empty.
+ */
+static __always_inline u32 cmask_next_and2_set_wrap(const struct scx_cmask __arena *a,
+						    const struct scx_cmask __arena *b,
+						    const struct scx_cmask __arena *c,
+						    u32 start)
+{
+	u32 a_end = a->base + a->nr_cids;
+	u32 found;
+
+	found = cmask_next_and2_set(a, b, c, start);
+	if (found < a_end || start <= a->base)
+		return found;
+
+	found = cmask_next_and2_set(a, b, c, a->base);
+	return found < start ? found : a_end;
+}
+
 /**
  * cmask_from_cpumask - translate a kernel cpumask to a cid-space cmask
  * @m: cmask to fill. Zeroed first; only bits within [@m->base, @m->base +
-- 
2.55.0
lmpx.com only provides a reader for public news (NNTP) servers. It is not affiliated with the servers or forums shown here and is not responsible for the content of articles, which is written by their respective authors.