Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 41 additions & 0 deletions compiler/rustc_ast_lowering/src/expr.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ use std::sync::Arc;
use rustc_ast::node_id::NodeMap;
use rustc_ast::visit::{Visitor, walk_expr};
use rustc_ast::*;
use rustc_attr_ir::find_attr;
use rustc_attr_ir::lang_items::LangItem;
use rustc_attr_ir::target::Target;
use rustc_errors::msg;
Expand Down Expand Up @@ -411,6 +412,7 @@ impl<'hir> LoweringContext<'_, 'hir> {
|this| {
this.with_new_scopes(e.span, |this| this.lower_block_expr(block))
},
None,
)
});
let Some(move_expr_state) = move_expr_state else {
Expand Down Expand Up @@ -866,6 +868,7 @@ impl<'hir> LoweringContext<'_, 'hir> {
desugaring_kind: hir::CoroutineDesugaring,
coroutine_source: hir::CoroutineSource,
body: impl FnOnce(&mut Self) -> hir::Expr<'hir>,
captured_caller_location: Option<HirId>,
) -> hir::ExprKind<'hir> {
let closure_def_id = self.local_def_id(closure_node_id);
let coroutine_kind = hir::CoroutineKind::Desugared(desugaring_kind, coroutine_source);
Expand Down Expand Up @@ -957,11 +960,49 @@ impl<'hir> LoweringContext<'_, 'hir> {
kind: hir::ClosureKind::Coroutine(coroutine_kind),
constness: hir::Constness::NotConst,
explicit_captures,
captured_caller_location,
}))
}

/// Checks whether `#[track_caller]` annotation on a coroutine function
/// or coroutine closure exists, and should affect the generated coroutine inside.
/// Currently used only for coroutine fns, not coroutine closures.
///
/// FIXME(closure_track_caller): Change coroutine closures to use this.
pub(super) fn should_track_caller_in_coroutine(
&self,
hir_id: HirId,
node_id: NodeId,
span: Span,
is_in_trait_impl: bool,
) -> bool {
if !self.tcx.features().async_fn_track_caller() {
return false;
}
if let Some(attrs) = self.curr_owner.attrs.get(&hir_id.local_id)
&& find_attr!(*attrs, TrackCaller(_))
{
// The coroutine function itself is annotated with #[track_caller]
return true;
}
if is_in_trait_impl {
// Check if we need to "inherit" #[track_caller] from the trait definition.
let Some(trait_item_def_id) =
self.get_partial_res(node_id).and_then(|r| r.expect_full_res().opt_def_id())
else {
self.dcx().span_delayed_bug(span, "could not resolve trait item being implemented");
return false;
};
return find_attr!(self.tcx, trait_item_def_id, TrackCaller(_));
}
Comment on lines +988 to +997

@theemathas theemathas Sep 29, 2026 •

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I don't understand what this is doing, but I'm doing the same thing as this existing code:

let (effective_ident, impl_kind) = if is_in_trait_impl {
let trait_item_def_id = self
.get_partial_res(i.id)
.and_then(|r| r.expect_full_res().opt_def_id())
.ok_or_else(|| {
self.dcx()
.span_delayed_bug(span, "could not resolve trait item being implemented")
});

View changes since the review

false
}

/// Forwards a possible `#[track_caller]` annotation from `outer_hir_id` to
/// `inner_hir_id` in case the `async_fn_track_caller` feature is enabled.
/// Currently only used for coroutine closures, not coroutine fns.
///
/// FIXME(closure_track_caller): Remove this function.
pub(super) fn maybe_forward_track_caller(&mut self, outer_hir_id: HirId, inner_hir_id: HirId) {
if self.tcx.features().async_fn_track_caller()
&& let Some(attrs) = self.curr_owner.attrs.get(&outer_hir_id.local_id)
Expand Down
6 changes: 6 additions & 0 deletions compiler/rustc_ast_lowering/src/expr/closure.rs
Original file line number Diff line number Diff line change
Expand Up @@ -219,6 +219,7 @@ impl<'hir> LoweringContext<'_, 'hir> {
kind: closure_kind,
constness: self.lower_constness(attrs, constness),
explicit_captures,
captured_caller_location: None,
});

(hir::ExprKind::Closure(c), move_expr_state)
Expand Down Expand Up @@ -315,8 +316,12 @@ impl<'hir> LoweringContext<'_, 'hir> {
body.span,
coroutine_marker,
hir::CoroutineSource::Closure,
false,
);

// FIXME(closure_track_caller): Currently, coroutine closures,
// unlike coroutine fns, have #[track_caller] track the poller instead of
// the caller of the closure.
this.maybe_forward_track_caller(closure_hir_id, expr.hir_id);

(parameters, expr)
Expand Down Expand Up @@ -358,6 +363,7 @@ impl<'hir> LoweringContext<'_, 'hir> {
kind: hir::ClosureKind::CoroutineClosure(coroutine_desugaring),
constness: self.lower_constness(attrs, constness),
explicit_captures,
captured_caller_location: None,
});
hir::ExprKind::Closure(c)
}
Expand Down
81 changes: 67 additions & 14 deletions compiler/rustc_ast_lowering/src/item.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@ use rustc_abi::ExternAbi;
use rustc_ast::visit::AssocCtxt;
use rustc_ast::*;
use rustc_attr_ir::target::Target;
use rustc_attr_ir::{AttributeKind, EiiImplResolution, find_attr};
use rustc_attr_ir::{AttributeKind, EiiImplResolution, LangItem, find_attr};
use rustc_errors::{E0570, ErrorGuaranteed, struct_span_code_err};
use rustc_hir::def::{DefKind, Res};
use rustc_hir::{self as hir, HirId, ImplItemImplKind, LifetimeSource, PredicateOrigin};
Expand Down Expand Up @@ -268,12 +268,14 @@ impl<'hir> LoweringContext<'_, 'hir> {
let body_id = this.lower_maybe_coroutine_body(
*fn_sig_span,
span,
id,
hir_id,
decl,
coroutine_marker,
body.as_deref(),
attrs,
contract.as_deref(),
false,
);

let itctx = ImplTraitContext::Universal;
Expand Down Expand Up @@ -867,12 +869,14 @@ impl<'hir> LoweringContext<'_, 'hir> {
let body_id = self.lower_maybe_coroutine_body(
sig.span,
i.span,
i.id,
hir_id,
&sig.decl,
sig.header.coroutine_marker,
Some(body),
attrs,
contract.as_deref(),
false,
);
let (generics, sig) = self.lower_method_sig(
generics,
Expand Down Expand Up @@ -1077,12 +1081,14 @@ impl<'hir> LoweringContext<'_, 'hir> {
let body_id = self.lower_maybe_coroutine_body(
sig.span,
i.span,
i.id,
hir_id,
&sig.decl,
sig.header.coroutine_marker,
body.as_deref(),
attrs,
contract.as_deref(),
is_in_trait_impl,
);
let (generics, sig) = self.lower_method_sig(
generics,
Expand Down Expand Up @@ -1271,12 +1277,14 @@ impl<'hir> LoweringContext<'_, 'hir> {
&mut self,
fn_decl_span: Span,
span: Span,
node_id: NodeId,
fn_id: hir::HirId,
decl: &FnDecl,
coroutine_marker: Option<CoroutineMarker>,
body: Option<&Block>,
attrs: &'hir [rustc_attr_ir::Attribute],
contract: Option<&FnContract>,
is_in_trait_impl: bool,
) -> hir::BodyId {
let Some(body) = body else {
// Functions without a body are an error, except if this is an intrinsic. For those we
Expand Down Expand Up @@ -1311,20 +1319,15 @@ impl<'hir> LoweringContext<'_, 'hir> {
};
// FIXME(contracts): Support contracts on async fn.
self.lower_body(|this| {
let (parameters, expr) = this.lower_coroutine_body_with_moved_arguments(
this.lower_coroutine_body_with_moved_arguments(
decl,
|this| this.lower_block_expr(body),
fn_decl_span,
body.span,
coroutine_marker,
hir::CoroutineSource::Fn,
);

// FIXME(async_fn_track_caller): Can this be moved above?
let hir_id = expr.hir_id;
this.maybe_forward_track_caller(fn_id, hir_id);

(parameters, expr)
this.should_track_caller_in_coroutine(fn_id, node_id, span, is_in_trait_impl),
)
})
}

Expand All @@ -1340,6 +1343,7 @@ impl<'hir> LoweringContext<'_, 'hir> {
body_span: Span,
coroutine_marker: CoroutineMarker,
coroutine_source: hir::CoroutineSource,
track_caller: bool,
) -> (&'hir [hir::Param<'hir>], hir::Expr<'hir>) {
let mut parameters: Vec<hir::Param<'_>> = Vec::new();
let mut statements: Vec<hir::Stmt<'_>> = Vec::new();
Expand Down Expand Up @@ -1472,6 +1476,42 @@ impl<'hir> LoweringContext<'_, 'hir> {
parameters.push(new_parameter);
}

let (caller_location_init_stmt, caller_location_hir_id) = track_caller
.then(|| {
let ident = Ident::with_dummy_span(sym::__captured_caller_location);
let span = self.mark_span_with_reason(
DesugaringKind::CoroutineFnTrackCaller,
DUMMY_SP,
Some([sym::core_intrinsics].into()),
);

// Get the caller location inside the function/closure body, but outside the coroutine.
let (outer_pat, outer_pat_hir_id) = self.pat_ident(span, ident);
let outer_expr = self.expr_call_lang_item_fn(span, LangItem::CallerLocation, &[]);
let outer_let_stmt = self.stmt_let_pat(
None,
span,
Some(outer_expr),
outer_pat,
hir::LocalSource::AsyncFn,
);

// Capture the stored caller location in the coroutine.
let (inner_pat, _inner_pat_hir_id) = self.pat_ident(span, ident);
let inner_expr = self.expr_ident(span, ident, outer_pat_hir_id);
let inner_let_stmt = self.stmt_let_pat(
None,
span,
Some(inner_expr),
inner_pat,
hir::LocalSource::AsyncFn,
);
statements.push(inner_let_stmt);

(outer_let_stmt, outer_pat_hir_id)
})
.unzip();

let mkbody = |this: &mut LoweringContext<'_, 'hir>| {
// Create a block from the user's function body:
let user_body = lower_body(this);
Expand Down Expand Up @@ -1505,7 +1545,7 @@ impl<'hir> LoweringContext<'_, 'hir> {
};
let closure_id = coroutine_marker.closure_id;

let coroutine_expr = self.make_desugared_coroutine_expr(
let coroutine_expr_kind = self.make_desugared_coroutine_expr(
// The default capture mode here is by-ref. Later on during upvar analysis,
// we will force the captured arguments to by-move, but for async closures,
// we want to make sure that we avoid unnecessarily moving captures, or else
Expand All @@ -1518,15 +1558,28 @@ impl<'hir> LoweringContext<'_, 'hir> {
desugaring_kind,
coroutine_source,
mkbody,
caller_location_hir_id,
);

let expr = hir::Expr {
let coroutine_expr = hir::Expr {
hir_id: self.lower_node_id(closure_id),
kind: coroutine_expr,
kind: coroutine_expr_kind,
span: self.lower_span(body_span),
};

(self.arena.alloc_from_iter(parameters), expr)
let body_expr = match caller_location_init_stmt {
Some(init_stmt) => {
let body_block = self.block_all(
DUMMY_SP,
self.arena.alloc_from_iter([init_stmt]),
Some(self.arena.alloc(coroutine_expr)),
);
let body_expr_kind = hir::ExprKind::Block(body_block, None);
hir::Expr { hir_id: self.next_id(), kind: body_expr_kind, span: DUMMY_SP }
}
None => coroutine_expr,
};

(self.arena.alloc_from_iter(parameters), body_expr)
}

fn lower_method_sig(
Expand Down
4 changes: 4 additions & 0 deletions compiler/rustc_attr_ir/src/lang_items.rs
Original file line number Diff line number Diff line change
Expand Up @@ -469,6 +469,10 @@ language_item_table! {

// Experimental lang item for `Reflection and comptime`(https://goals.rust-lang.org/2025h2/reflection-and-comptime.html)
FnPtr, sym::FnPtr, fn_ptr, Target::Struct, GenericRequirement::None;

// Used in the desugaring of #[track_caller] on coroutine functions
// FIXME(closure_track_caller): Also use this in coroutine closures.
CallerLocation, sym::caller_location, caller_location, Target::Fn, GenericRequirement::Exact(0);
}

/// The requirement imposed on the generics of a lang item
Expand Down
16 changes: 12 additions & 4 deletions compiler/rustc_codegen_cranelift/src/abi/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -325,15 +325,23 @@ pub(crate) fn codegen_fn_prelude<'tcx>(fx: &mut FunctionCx<'_, '_, 'tcx>, start_
.collect::<Vec<(Local, ArgKind<'tcx>, Ty<'tcx>)>>();

assert!(fx.caller_location.is_none());
if fx.instance.def.requires_caller_location(fx.tcx) {
if let Some(coro_info) = fx.mir.coroutine.as_deref()
&& let Some(captured_caller_location_idx) = coro_info.captured_caller_location
{
assert!(
!fx.instance.def.requires_caller_location(fx.tcx),
"should have only one source of truth for caller_location"
);
fx.caller_location = Some(CallerLocation::Captured(captured_caller_location_idx));
} else if fx.instance.def.requires_caller_location(fx.tcx) {
// Store caller location for `#[track_caller]`.
let arg_abi = arg_abis_iter.next().unwrap();
let param = cvalue_for_param(fx, None, None, arg_abi, &mut block_params_iter).unwrap();
assert!(
!param.is_underaligned_pointee,
"caller location argument should not be underaligned",
);
fx.caller_location = Some(param.value);
fx.caller_location = Some(CallerLocation::Direct(param.value));
}

assert_eq!(arg_abis_iter.next(), None, "ArgAbi left behind for {:?}", fx.fn_abi);
Expand Down Expand Up @@ -557,7 +565,7 @@ pub(crate) fn codegen_terminator_call<'tcx>(

// Pass the caller location for `#[track_caller]`.
if instance.is_some_and(|inst| inst.def.requires_caller_location(fx.tcx)) {
let caller_location = fx.get_caller_location(source_info);
let caller_location = fx.codegen_caller_location(source_info);
args.push(CallArgument { value: caller_location, is_owned: false });
}

Expand Down Expand Up @@ -811,7 +819,7 @@ pub(crate) fn codegen_drop<'tcx>(

if drop_instance.def.requires_caller_location(fx.tcx) {
// Pass the caller location for `#[track_caller]`.
let caller_location = fx.get_caller_location(source_info);
let caller_location = fx.codegen_caller_location(source_info);
call_args.extend(adjust_arg_for_abi(
fx,
caller_location,
Expand Down
12 changes: 6 additions & 6 deletions compiler/rustc_codegen_cranelift/src/base.rs
Original file line number Diff line number Diff line change
Expand Up @@ -391,7 +391,7 @@ fn codegen_fn_body(fx: &mut FunctionCx<'_, '_, '_>, start_block: Block) {
AssertKind::BoundsCheck { len, index } => {
let len = codegen_operand(fx, len).load_scalar(fx);
let index = codegen_operand(fx, index).load_scalar(fx);
let location = fx.get_caller_location(source_info).load_scalar(fx);
let location = fx.codegen_caller_location(source_info).load_scalar(fx);

codegen_panic_inner(
fx,
Expand All @@ -404,7 +404,7 @@ fn codegen_fn_body(fx: &mut FunctionCx<'_, '_, '_>, start_block: Block) {
AssertKind::MisalignedPointerDereference { required, found } => {
let required = codegen_operand(fx, required).load_scalar(fx);
let found = codegen_operand(fx, found).load_scalar(fx);
let location = fx.get_caller_location(source_info).load_scalar(fx);
let location = fx.codegen_caller_location(source_info).load_scalar(fx);

codegen_panic_inner(
fx,
Expand All @@ -415,7 +415,7 @@ fn codegen_fn_body(fx: &mut FunctionCx<'_, '_, '_>, start_block: Block) {
);
}
AssertKind::NullPointerDereference => {
let location = fx.get_caller_location(source_info).load_scalar(fx);
let location = fx.codegen_caller_location(source_info).load_scalar(fx);

codegen_panic_inner(
fx,
Expand All @@ -426,7 +426,7 @@ fn codegen_fn_body(fx: &mut FunctionCx<'_, '_, '_>, start_block: Block) {
)
}
AssertKind::NullReferenceConstructed => {
let location = fx.get_caller_location(source_info).load_scalar(fx);
let location = fx.codegen_caller_location(source_info).load_scalar(fx);

codegen_panic_inner(
fx,
Expand All @@ -438,7 +438,7 @@ fn codegen_fn_body(fx: &mut FunctionCx<'_, '_, '_>, start_block: Block) {
}
AssertKind::InvalidEnumConstruction(source) => {
let source = codegen_operand(fx, source).load_scalar(fx);
let location = fx.get_caller_location(source_info).load_scalar(fx);
let location = fx.codegen_caller_location(source_info).load_scalar(fx);

codegen_panic_inner(
fx,
Expand All @@ -449,7 +449,7 @@ fn codegen_fn_body(fx: &mut FunctionCx<'_, '_, '_>, start_block: Block) {
)
}
_ => {
let location = fx.get_caller_location(source_info).load_scalar(fx);
let location = fx.codegen_caller_location(source_info).load_scalar(fx);

codegen_panic_inner(
fx,
Expand Down
Loading
Loading