From e620c7f9daea9ba8abb83eb85a104f770653a422 Mon Sep 17 00:00:00 2001 From: "Tim (Theemathas Chirananthavat)" Date: Sun, 4 Oct 2026 15:30:51 +0700 Subject: [PATCH] consteval: Rework caller location handling Instead of walking up the call stack each time we need the caller location, we emulate the run time behavior of passing an implicit caller location argument in #[track_caller] functions. We record a representation of that argument in each frame. For inlined functions though, we still do a "stack walk" within that one frame. --- .../src/const_eval/machine.rs | 4 +- .../rustc_const_eval/src/interpret/call.rs | 21 +++++- .../src/interpret/eval_context.rs | 66 ++++++++----------- .../src/interpret/intrinsics.rs | 2 +- .../rustc_const_eval/src/interpret/stack.rs | 9 ++- .../rustc_const_eval/src/interpret/util.rs | 2 +- src/tools/miri/src/helpers.rs | 2 + 7 files changed, 58 insertions(+), 48 deletions(-) diff --git a/compiler/rustc_const_eval/src/const_eval/machine.rs b/compiler/rustc_const_eval/src/const_eval/machine.rs index 1b99f54719fb4..244c72d1847c8 100644 --- a/compiler/rustc_const_eval/src/const_eval/machine.rs +++ b/compiler/rustc_const_eval/src/const_eval/machine.rs @@ -253,7 +253,7 @@ impl<'tcx> CompileTimeInterpCx<'tcx> { } let msg = Symbol::intern(self.read_str(&msg_place)?); - let span = self.find_closest_untracked_caller_location(); + let span = self.caller_location(); let (file, line, col) = self.location_triple_for_span(span); return Err(ConstEvalErrKind::Panic { msg, file, line, col }).into(); } else if self.tcx.is_lang_item(def_id, LangItem::PanicFmt) { @@ -463,7 +463,7 @@ impl<'tcx> interpret::Machine<'tcx> for CompileTimeMachine<'tcx> { fn panic_nounwind(ecx: &mut InterpCx<'tcx, Self>, msg: &str) -> InterpResult<'tcx> { let msg = Symbol::intern(msg); - let span = ecx.find_closest_untracked_caller_location(); + let span = ecx.caller_location(); let (file, line, col) = ecx.location_triple_for_span(span); Err(ConstEvalErrKind::Panic { msg, file, line, col }).into() } diff --git a/compiler/rustc_const_eval/src/interpret/call.rs b/compiler/rustc_const_eval/src/interpret/call.rs index c7615a216312e..f1b551441a675 100644 --- a/compiler/rustc_const_eval/src/interpret/call.rs +++ b/compiler/rustc_const_eval/src/interpret/call.rs @@ -10,7 +10,7 @@ use rustc_attr_ir::find_attr; use rustc_hir::def_id::DefId; use rustc_middle::mir; use rustc_middle::ty::layout::{IntegerExt, TyAndLayout}; -use rustc_middle::ty::{self, AdtDef, FieldDef, Instance, Ty, VariantDef}; +use rustc_middle::ty::{self, AdtDef, FieldDef, Instance, InstanceKind, Ty, VariantDef}; use rustc_span::{bug, span_bug}; use rustc_target::callconv::{ArgAbi, FnAbi}; use tracing::field::Empty; @@ -456,12 +456,16 @@ impl<'tcx, M: Machine<'tcx>> InterpCx<'tcx, M> { caller_fn_abi: &FnAbi<'tcx, Ty<'tcx>>, args: &[FnArg<'tcx, M::Provenance>], with_caller_location: bool, + callee_receives_caller_location: bool, destination: &PlaceTy<'tcx, M::Provenance>, mut cont: ReturnContinuation, ) -> InterpResult<'tcx> { let _trace = enter_trace_span!(M, step::init_stack_frame, %instance, tracing_separate_thread = Empty); let def_id = instance.def_id(); + assert!(!with_caller_location || callee_receives_caller_location); + let caller_location = callee_receives_caller_location.then(|| self.caller_location()); + // The first order of business is to figure out the callee signature. // However, that requires the list of variadic arguments. // We use the *caller* information to determine where to split the list of arguments, @@ -542,6 +546,8 @@ impl<'tcx, M: Machine<'tcx>> InterpCx<'tcx, M> { .collect::>() ); + self.frame_mut().track_caller_arg = caller_location; + // Determine whether there is a special VaList argument. This is always the // last argument, and since arguments start at index 1 that's `arg_count`. let va_list_arg = callee_fn_abi.c_variadic.then(|| mir::Local::from_usize(body.arg_count)); @@ -563,7 +569,8 @@ impl<'tcx, M: Machine<'tcx>> InterpCx<'tcx, M> { // The "where they come from" part is easy, we expect the caller to do any special handling // that might be required here (e.g. for untupling). // If `with_caller_location` is set we pretend there is an extra argument (that - // we will not pass; our `caller_location` intrinsic implementation walks the stack instead). + // we will not pass normally; our `caller_location` intrinsic implementation handles + // that separately). assert_eq!( args.len() + if with_caller_location { 1 } else { 0 }, caller_fn_abi.args.len(), @@ -847,12 +854,22 @@ impl<'tcx, M: Machine<'tcx>> InterpCx<'tcx, M> { Cow::from(args) }; + // Figure out if the callee receives a caller location argument. + // In the case where we are calling the fallback body of + // a `#[rustc_intrinsic] #[track_caller] fn`, we pretend that + // the caller doesn't pass a caller location argument (for checking ABI), + // but the caller receives a caller location argument from thin air anyway. + let callee_receives_caller_location = with_caller_location + || (instance.def.requires_caller_location(*self.tcx) + && matches!(instance.def, InstanceKind::Item(def_id) if self.tcx.intrinsic(def_id).is_some())); + self.init_stack_frame( instance, body, caller_fn_abi, &args, with_caller_location, + callee_receives_caller_location, destination, ReturnContinuation::Goto { ret: target, unwind }, ) diff --git a/compiler/rustc_const_eval/src/interpret/eval_context.rs b/compiler/rustc_const_eval/src/interpret/eval_context.rs index a74b323c395e6..d6d80c3e3506a 100644 --- a/compiler/rustc_const_eval/src/interpret/eval_context.rs +++ b/compiler/rustc_const_eval/src/interpret/eval_context.rs @@ -1,5 +1,6 @@ use std::cell::RefCell; use std::collections::hash_map::Entry; +use std::convert::identity; use either::{Left, Right}; use rustc_abi::{Align, HasDataLayout, Size, TargetDataLayout}; @@ -18,7 +19,7 @@ use rustc_middle::ty::{ use rustc_span::{Span, span_bug}; use rustc_structures::Limit; use rustc_target::callconv::FnAbi; -use tracing::{debug, trace}; +use tracing::trace; use super::{ Frame, FrameInfo, GlobalId, InterpErrorKind, InterpResult, MPlaceTy, Machine, MemPlaceMeta, @@ -376,49 +377,34 @@ impl<'tcx, M: Machine<'tcx>> InterpCx<'tcx, M> { } } - /// Walks up the callstack from the intrinsic's callsite, searching for the first callsite in a - /// frame which is not `#[track_caller]`. This matches the `caller_location` intrinsic, + /// Grabs the implicit caller location argument, if there is one, + /// or falls back to the "current" span. This matches the `caller_location` intrinsic, /// and is primarily intended for the panic machinery. - pub(crate) fn find_closest_untracked_caller_location(&self) -> Span { - for frame in self.stack().iter().rev() { - debug!("find_closest_untracked_caller_location: checking frame {:?}", frame.instance); - - // Assert that the frame we look at is actually executing code currently - // (`loc` is `Right` when we are unwinding and the frame does not require cleanup). - let loc = frame.loc.left().unwrap(); - - // This could be a non-`Call` terminator (such as `Drop`), or not a terminator at all - // (such as `box`). Use the normal span by default. - let mut source_info = *frame.body.source_info(loc); - - // If this is a `Call` terminator, use the `fn_span` instead. - let block = &frame.body.basic_blocks[loc.block]; - if loc.statement_index == block.statements.len() { - debug!( - "find_closest_untracked_caller_location: got terminator {:?} ({:?})", - block.terminator(), - block.terminator().kind, - ); - if let mir::TerminatorKind::Call { fn_span, .. } = block.terminator().kind { - source_info.span = fn_span; - } - } + pub(crate) fn caller_location(&self) -> Span { + let frame = self.frame(); - let caller_location = if frame.instance.def.requires_caller_location(*self.tcx) { - // We use `Err(())` as indication that we should continue up the call stack since - // this is a `#[track_caller]` function. - Some(Err(())) - } else { - None - }; - if let Ok(span) = - frame.body.caller_location_span(source_info, caller_location, *self.tcx, Ok) - { - return span; - } + // Assert that the frame we look at is actually executing code currently + // (`loc` is `Right` when we are unwinding and the frame does not require cleanup). + let loc = frame.loc.left().unwrap(); + + assert_eq!( + frame.track_caller_arg.is_some(), + frame.instance.def.requires_caller_location(*self.tcx) + ); + + // This could be a non-`Call` terminator (such as `Drop`), or not a terminator at all + // (such as `box`). Use the normal span by default. + let mut source_info = *frame.body.source_info(loc); + + // If this is a `Call` terminator, use the `fn_span` instead. + let block = &frame.body.basic_blocks[loc.block]; + if loc.statement_index == block.statements.len() + && let mir::TerminatorKind::Call { fn_span, .. } = block.terminator().kind + { + source_info.span = fn_span; } - span_bug!(self.cur_span(), "no non-`#[track_caller]` frame found") + frame.body.caller_location_span(source_info, frame.track_caller_arg, *self.tcx, identity) } /// Returns the actual dynamic size and alignment of the place at the given type. diff --git a/compiler/rustc_const_eval/src/interpret/intrinsics.rs b/compiler/rustc_const_eval/src/interpret/intrinsics.rs index b8b372ad1dab1..d66f4cf900eca 100644 --- a/compiler/rustc_const_eval/src/interpret/intrinsics.rs +++ b/compiler/rustc_const_eval/src/interpret/intrinsics.rs @@ -293,7 +293,7 @@ impl<'tcx, M: Machine<'tcx>> InterpCx<'tcx, M> { } sym::caller_location => { - let span = self.find_closest_untracked_caller_location(); + let span = self.caller_location(); let val = self.tcx.span_as_caller_location(span); let val = self.const_val_to_op(val, self.tcx.caller_location_ty(), Some(dest.layout))?; diff --git a/compiler/rustc_const_eval/src/interpret/stack.rs b/compiler/rustc_const_eval/src/interpret/stack.rs index 8a731e25d8f4b..6d5c7a78d9f73 100644 --- a/compiler/rustc_const_eval/src/interpret/stack.rs +++ b/compiler/rustc_const_eval/src/interpret/stack.rs @@ -96,6 +96,9 @@ pub struct Frame<'tcx, Prov: Provenance = CtfeProvenance, Extra = ()> { /// frame is popped. pub(super) va_list: Vec>, + /// The implicit caller location argument that is passed by this function's caller, if any. + pub(super) track_caller_arg: Option, + /// The span of the `tracing` crate is stored here. /// When the guard is dropped, the span is exited. This gives us /// a full stack trace on all tracing statements. @@ -255,13 +258,14 @@ impl<'tcx, Prov: Provenance> Frame<'tcx, Prov> { Frame { body: self.body, instance: self.instance, + extra, return_cont: self.return_cont, return_place: self.return_place, locals: self.locals, va_list: self.va_list, - loc: self.loc, - extra, + track_caller_arg: self.track_caller_arg, tracing_span: self.tracing_span, + loc: self.loc, } } } @@ -397,6 +401,7 @@ impl<'tcx, M: Machine<'tcx>> InterpCx<'tcx, M> { return_place: return_place.clone(), locals, va_list: vec![], + track_caller_arg: None, instance, tracing_span: SpanGuard::new(), extra: (), diff --git a/compiler/rustc_const_eval/src/interpret/util.rs b/compiler/rustc_const_eval/src/interpret/util.rs index 4efef90581080..416207371e51c 100644 --- a/compiler/rustc_const_eval/src/interpret/util.rs +++ b/compiler/rustc_const_eval/src/interpret/util.rs @@ -25,7 +25,7 @@ pub(crate) fn type_implements_dyn_trait<'tcx, M: Machine<'tcx>>( let ty::Dynamic(preds, _) = trait_ty.kind() else { span_bug!( - ecx.find_closest_untracked_caller_location(), + ecx.cur_span(), "Invalid type provided to type_implements_predicates. U must be dyn Trait, got {trait_ty}." ); }; diff --git a/src/tools/miri/src/helpers.rs b/src/tools/miri/src/helpers.rs index b0ff200c0614b..fbbf99c7f3d87 100644 --- a/src/tools/miri/src/helpers.rs +++ b/src/tools/miri/src/helpers.rs @@ -429,6 +429,7 @@ pub trait EvalContextExt<'tcx>: crate::MiriInterpCxExt<'tcx> { ); let caller_fn_abi = this.fn_abi_of_fn_ptr(ty::Binder::dummy(sig), ty::List::empty())?; + let callee_receives_caller_location = f.def.requires_caller_location(*this.tcx); // This will also show proper errors if there is any ABI mismatch. this.init_stack_frame( f, @@ -436,6 +437,7 @@ pub trait EvalContextExt<'tcx>: crate::MiriInterpCxExt<'tcx> { caller_fn_abi, &args.iter().map(|a| FnArg::Copy(a.clone().into())).collect::>(), /*with_caller_location*/ false, + callee_receives_caller_location, &dest.into(), cont, )