Re: [PATCH 2/3] rbtree: update augmented data on the way down in rb_add_augmented_cached()

From: Peter Zijlstra

Date: Mon Sep 28 2026 - 10:39:20 EST


On Mon, Sep 28, 2026 at 03:37:33PM +0200, Peter Zijlstra wrote:
> On Mon, Sep 28, 2026 at 08:26:10PM +0800, Yiwei Lin wrote:
>
> > diff --git a/kernel/sched/fair.c b/kernel/sched/fair.c
> > index 7455a83a6a990..60db4624b9897 100644
> > --- a/kernel/sched/fair.c
> > +++ b/kernel/sched/fair.c
> > @@ -1066,9 +1066,20 @@ static inline bool min_vruntime_update(struct sched_entity *se, bool exit)
> > se->max_slice == old_max_slice;
> > }
> >
> > +/*
> > + * Fold @new's subtree data into @se, for each @se on @new's insertion path.
> > + */
> > +static inline void
> > +min_vruntime_merge(struct sched_entity *se, struct sched_entity *new)
> > +{
> > + __min_vruntime_update(se, &new->run_node);
> > + __min_slice_update(se, &new->run_node);
> > + __max_slice_update(se, &new->run_node);
> > +}
> >
>
> So min_vruntime_update() can be written in terms of this helper like:
>
> static inline bool min_vruntime_update(struct sched_entity *se, bool exit)
> {
> u64 old_min_vruntime = se->min_vruntime;
> u64 old_min_slice = se->min_slice;
> u64 old_max_slice = se->max_slice;
> struct rb_node *node = &se->run_node;
>
> se->min_vruntime = se->vruntime;
> se->min_slice = se->slice;
> se->max_slice = se->slice;
>
> min_vruntime_merge(se, node->rb_right);
> min_vruntime_merge(se, node->rb_left);
>
> return se->min_vruntime == old_min_vruntime &&
> se->min_slice == old_min_slice &&
> se->max_slice == old_max_slice;
> }
>
> And that is *very* close to being generalizable, obviating the need for
> RBCOMPUTE. Does your LLM see a way to make that happen?

Perhaps by doing something like so?

diff --git a/include/linux/rbtree_augmented.h b/include/linux/rbtree_augmented.h
index d2fa1c41bfd2..c4cf5160aa17 100644
--- a/include/linux/rbtree_augmented.h
+++ b/include/linux/rbtree_augmented.h
@@ -15,6 +15,7 @@
#include <linux/compiler.h>
#include <linux/rbtree.h>
#include <linux/rcupdate.h>
+#include <linux/args.h>

/*
* Please note - only struct rb_augment_callbacks and the prototypes for
@@ -86,6 +87,20 @@ rb_add_augmented_cached(struct rb_node *node, struct rb_root_cached *tree,
return leftmost ? node : NULL;
}

+#define FOR_EACH_1(what, x) what(x)
+#define FOR_EACH_2(what, x, ...) what(x) FOR_EACH_1(what, __VA_ARGS__)
+#define FOR_EACH_3(what, x, ...) what(x) FOR_EACH_2(what, __VA_ARGS__)
+#define FOR_EACH_4(what, x, ...) what(x) FOR_EACH_3(what, __VA_ARGS__)
+#define FOR_EACH_5(what, x, ...) what(x) FOR_EACH_4(what, __VA_ARGS__)
+#define FOR_EACH_6(what, x, ...) what(x) FOR_EACH_5(what, __VA_ARGS__)
+#define FOR_EACH_7(what, x, ...) what(x) FOR_EACH_6(what, __VA_ARGS__)
+#define FOR_EACH_8(what, x, ...) what(x) FOR_EACH_7(what, __VA_ARGS__)
+
+#define FOR_EACH(action, ...) \
+ CONCATENATE(FOR_EACH_, COUNT_ARGS(__VA_ARGS__))(action, __VA_ARGS__)
+
+#define COPY_VAL(x) new->x = old->x;
+
/*
* Template for declaring augmented rbtree callbacks (generic multi fields)
*
@@ -93,12 +108,12 @@ rb_add_augmented_cached(struct rb_node *node, struct rb_root_cached *tree,
* RBNAME: name of the rb_augment_callbacks structure
* RBSTRUCT: struct type of the tree nodes
* RBFIELD: name of struct rb_node field within RBSTRUCT
- * RBCOPY: name of function that copies the RBAUGMENTED datas
* RBCOMPUTE: name of function that recomputes the RBAUGMENTED datas
+ * RBAUG...: field names within RBSTRUCT holding data for the subtree
*/

#define RB_DECLARE_CALLBACKS_MULTI(RBSTATIC, RBNAME, \
- RBSTRUCT, RBFIELD, RBCOPY, RBCOMPUTE) \
+ RBSTRUCT, RBFIELD, RBCOMPUTE, RBAUG...) \
static inline void \
RBNAME ## _propagate(struct rb_node *rb, struct rb_node *stop) \
{ \
@@ -110,18 +125,23 @@ RBNAME ## _propagate(struct rb_node *rb, struct rb_node *stop) \
} \
} \
static inline void \
+RBNAME ## __copy(RBSTRUCT *new, RBSTRUCT *old) \
+{ \
+ FOR_EACH(COPY_VAL, RBAUG); \
+} \
+static inline void \
RBNAME ## _copy(struct rb_node *rb_old, struct rb_node *rb_new) \
{ \
RBSTRUCT *old = rb_entry(rb_old, RBSTRUCT, RBFIELD); \
RBSTRUCT *new = rb_entry(rb_new, RBSTRUCT, RBFIELD); \
- RBCOPY(new, old); \
+ RBNAME ## __copy(new, old); \
} \
static void \
RBNAME ## _rotate(struct rb_node *rb_old, struct rb_node *rb_new) \
{ \
RBSTRUCT *old = rb_entry(rb_old, RBSTRUCT, RBFIELD); \
RBSTRUCT *new = rb_entry(rb_new, RBSTRUCT, RBFIELD); \
- RBCOPY(new, old); \
+ RBNAME ## __copy(new, old); \
RBCOMPUTE(old, false); \
} \
RBSTATIC const struct rb_augment_callbacks RBNAME = { \
diff --git a/kernel/sched/fair.c b/kernel/sched/fair.c
index e707da7177df..cc4cd3565268 100644
--- a/kernel/sched/fair.c
+++ b/kernel/sched/fair.c
@@ -1061,9 +1061,8 @@ static inline bool min_vruntime_update(struct sched_entity *se, bool exit)
se->max_slice == old_max_slice;
}

-
RB_DECLARE_CALLBACKS_MULTI(static, min_vruntime_cb, struct sched_entity,
- run_node, min_vruntime_copy, min_vruntime_update);
+ run_node, min_vruntime_update, min_vruntime, min_slice, max_slice)

/*
* Enqueue an entity into the rb-tree: