All callers perform the same obj_to_index() calculation to pass the index. Simplify by passing object pointer instead and determining the index by slab_obj_ext(). Reviewed-by: Suren Baghdasaryan Signed-off-by: Vlastimil Babka (SUSE) --- mm/memcontrol.c | 12 +++--------- mm/slab.h | 19 +++++++++++-------- mm/slub.c | 22 +++++++--------------- 3 files changed, 21 insertions(+), 32 deletions(-) diff --git a/mm/memcontrol.c b/mm/memcontrol.c index 6dc4888a90f3..4e427286a88a 100644 --- a/mm/memcontrol.c +++ b/mm/memcontrol.c @@ -2865,15 +2865,13 @@ struct mem_cgroup *mem_cgroup_from_obj_slab(struct slab *slab, void *p) */ unsigned long obj_exts; struct slabobj_ext *obj_ext; - unsigned int off; obj_exts = slab_obj_exts(slab); if (!obj_exts) return NULL; get_slab_obj_exts(obj_exts); - off = obj_to_index(slab->slab_cache, slab, p); - obj_ext = slab_obj_ext(slab, obj_exts, off); + obj_ext = slab_obj_ext(slab->slab_cache, slab, obj_exts, p); if (obj_ext->objcg) { struct obj_cgroup *objcg = obj_ext->objcg; @@ -3541,7 +3539,6 @@ bool __memcg_slab_post_alloc_hook(struct kmem_cache *s, struct list_lru *lru, size_t obj_size = obj_full_size(s); struct obj_cgroup *objcg; struct slab *slab; - unsigned long off; size_t i; /* @@ -3616,8 +3613,7 @@ bool __memcg_slab_post_alloc_hook(struct kmem_cache *s, struct list_lru *lru, obj_exts = slab_obj_exts(slab); get_slab_obj_exts(obj_exts); - off = obj_to_index(s, slab, p[i]); - obj_ext = slab_obj_ext(slab, obj_exts, off); + obj_ext = slab_obj_ext(s, slab, obj_exts, p[i]); obj_cgroup_get(objcg); obj_ext->objcg = objcg; put_slab_obj_exts(obj_exts); @@ -3635,10 +3631,8 @@ void __memcg_slab_free_hook(struct kmem_cache *s, struct slab *slab, struct obj_cgroup *objcg; struct slabobj_ext *obj_ext; struct obj_stock_pcp *stock; - unsigned int off; - off = obj_to_index(s, slab, p[i]); - obj_ext = slab_obj_ext(slab, obj_exts, off); + obj_ext = slab_obj_ext(s, slab, obj_exts, p[i]); objcg = obj_ext->objcg; if (!objcg) continue; diff --git a/mm/slab.h b/mm/slab.h index 7bd361447c54..64cec02b5016 100644 --- a/mm/slab.h +++ b/mm/slab.h @@ -579,7 +579,7 @@ struct slabobj_ext { * obj_exts = slab_obj_exts(slab); * if (obj_exts) { * get_slab_obj_exts(obj_exts); - * obj_ext = slab_obj_ext(slab, obj_exts, obj_to_index(s, slab, obj)); + * obj_ext = slab_obj_ext(s, slab, obj_exts, obj); * // do something with obj_ext * put_slab_obj_exts(obj_exts); * } @@ -639,21 +639,24 @@ static inline unsigned int slab_get_stride(struct slab *slab) /* * slab_obj_ext - get the pointer to the slab object extension metadata * associated with an object in a slab. + * @s: cache that the slab blongs to * @slab: a pointer to the slab struct * @obj_exts: a pointer to the object extension vector - * @index: an index of the object + * @obj: a pointer to the object * * Returns a pointer to the object extension associated with the object. * Must be called within a section covered by get/put_slab_obj_exts(). */ -static inline struct slabobj_ext *slab_obj_ext(struct slab *slab, - unsigned long obj_exts, - unsigned int index) +static inline struct slabobj_ext * +slab_obj_ext(struct kmem_cache *s, struct slab *slab, unsigned long obj_exts, + const void *obj) { struct slabobj_ext *obj_ext; + unsigned int index; VM_WARN_ON_ONCE(obj_exts != slab_obj_exts(slab)); + index = obj_to_index(s, slab, obj); obj_ext = (struct slabobj_ext *)(obj_exts + slab_get_stride(slab) * index); return kasan_reset_tag(obj_ext); @@ -669,9 +672,9 @@ static inline unsigned long slab_obj_exts(struct slab *slab) return 0; } -static inline struct slabobj_ext *slab_obj_ext(struct slab *slab, - unsigned long obj_exts, - unsigned int index) +static inline struct slabobj_ext * +slab_obj_ext(struct kmem_cache *s, struct slab *slab, unsigned long obj_exts, + const void *obj) { return NULL; } diff --git a/mm/slub.c b/mm/slub.c index 8c1031989e41..aa99d7eb6a4d 100644 --- a/mm/slub.c +++ b/mm/slub.c @@ -2070,11 +2070,10 @@ static inline void mark_obj_codetag_empty(const void *obj) obj_slab = virt_to_slab(obj); slab_exts = slab_obj_exts(obj_slab); if (slab_exts) { + struct slabobj_ext *ext; + get_slab_obj_exts(slab_exts); - unsigned int offs = obj_to_index(obj_slab->slab_cache, - obj_slab, obj); - struct slabobj_ext *ext = slab_obj_ext(obj_slab, - slab_exts, offs); + ext = slab_obj_ext(obj_slab->slab_cache, obj_slab, slab_exts, obj); if (is_kfence_address(obj)) { put_slab_obj_exts(slab_exts); @@ -2368,10 +2367,8 @@ __alloc_tagging_slab_alloc_hook(struct kmem_cache *s, void *object, gfp_t flags, * check should be added before alloc_tag_add(). */ if (obj_exts) { - unsigned int obj_idx = obj_to_index(s, slab, object); - get_slab_obj_exts(obj_exts); - obj_ext = slab_obj_ext(slab, obj_exts, obj_idx); + obj_ext = slab_obj_ext(s, slab, obj_exts, object); alloc_tag_add(&obj_ext->ref, current->alloc_tag, s->size); put_slab_obj_exts(obj_exts); } else { @@ -2392,7 +2389,6 @@ static noinline void __alloc_tagging_slab_free_hook(struct kmem_cache *s, struct slab *slab, void **p, int objects) { - int i; unsigned long obj_exts; /* slab->obj_exts might not be NULL if it was created for MEMCG accounting. */ @@ -2404,13 +2400,11 @@ __alloc_tagging_slab_free_hook(struct kmem_cache *s, struct slab *slab, void **p return; get_slab_obj_exts(obj_exts); - for (i = 0; i < objects; i++) { - unsigned int off = obj_to_index(s, slab, p[i]); - + for (int i = 0; i < objects; i++) { if (is_kfence_address(p[i])) continue; - alloc_tag_sub(&slab_obj_ext(slab, obj_exts, off)->ref, s->size); + alloc_tag_sub(&slab_obj_ext(s, slab, obj_exts, p[i])->ref, s->size); } put_slab_obj_exts(obj_exts); } @@ -2495,7 +2489,6 @@ bool memcg_slab_post_charge(void *p, gfp_t flags) struct kmem_cache *s; struct page *page; struct slab *slab; - unsigned long off; page = virt_to_page(p); if (PageLargeKmalloc(page)) { @@ -2535,8 +2528,7 @@ bool memcg_slab_post_charge(void *p, gfp_t flags) obj_exts = slab_obj_exts(slab); if (obj_exts) { get_slab_obj_exts(obj_exts); - off = obj_to_index(s, slab, p); - obj_ext = slab_obj_ext(slab, obj_exts, off); + obj_ext = slab_obj_ext(s, slab, obj_exts, p); if (unlikely(obj_ext->objcg)) { put_slab_obj_exts(obj_exts); return true; -- 2.55.0