Implement thunking in terms of constant evaluation (#7332)

The bulk of this change is changing most pattern insts to be `Always`
rather than `AlwaysUnique` constants, so that they can be wrapped in
`SpecificConstant`s to perform substitution. That then lets thunking
rely much more on `SpecificConstant` wrappers instead of deep-copying
the inst tree with modified types.

This approach to thunking should scale better, particularly as things
like form generics make function signatures more complex, because we can
leverage the existing support for constant evaluation and substitution.
Unfortunately, applying this approach to binding patterns will require
more work; see the TODO near the top of `thunk.cpp` for details.

---------

Co-authored-by: Richard Smith <richard@metafoo.co.uk>
This commit is contained in:
Geoff Romer
2026-06-13 00:22:00 +00:00
committed by GitHub
co-authored by Richard Smith
parent 40d2d8f68c
commit e023f75254
589 changed files with 22583 additions and 11517 deletions
+191 -96
View File
@@ -33,9 +33,6 @@ namespace {
struct CallerState {
// The in-progress contents of the `Call` arguments block.
llvm::SmallVector<SemIR::InstId> call_args;
// The SpecificId of the function being called (if any).
SemIR::SpecificId callee_specific_id;
};
// Manages the allocation of call parameter indices.
@@ -62,6 +59,15 @@ class IndexSource {
// State for callee-side pattern matching.
struct CalleeState {
// Pushes an inst onto the list of `Call` parameter patterns whose constant
// value is the substituted constant value of `pattern_id` in `specific_id`.
auto PushCallParamPattern(Context& context, SemIR::LocId loc_id,
SemIR::InstId pattern_id,
SemIR::SpecificId specific_id) -> void {
call_param_patterns.push_back(
WrapInstForSpecific(context, loc_id, pattern_id, specific_id).inst_id);
}
IndexSource index;
// The in-progress contents of the `Call` parameters block.
@@ -140,8 +146,11 @@ class MatchContext {
}
};
// Constructs a MatchContext.
explicit MatchContext(Context& context) : context_(context) {}
// Constructs a MatchContext that uses the substituted constant values of
// patterns in the given specific.
explicit MatchContext(Context& context,
SemIR::SpecificId specific_id = SemIR::SpecificId::None)
: specific_id_stack_({specific_id}), context_(context) {}
// Performs pattern matching for the given work item.
auto Match(State state, WorkItem entry) -> void;
@@ -189,6 +198,10 @@ class MatchContext {
SemIR::InstId scrutinee_id, WorkItem entry) -> void;
auto DoPreWork(State state, SemIR::SpliceInst splice,
SemIR::InstId scrutinee_id, WorkItem entry) -> void;
auto DoPreWork(State state, SemIR::SpecificConstant specific_constant,
SemIR::InstId scrutinee_id, WorkItem entry) -> void;
auto DoPreWork(State state, SemIR::ImportRefLoaded import_ref,
SemIR::InstId scrutinee_id, WorkItem entry) -> void;
// Do the post-work for `entry`. `entry.work` must be a `PostWork`, and
// the pattern argument must be the value of `entry.pattern_id` in `context_`.
@@ -204,15 +217,10 @@ class MatchContext {
WorkItem entry) -> void;
auto DoPostWork(State state, SemIR::TuplePattern tuple_pattern,
WorkItem entry) -> void;
// Asserts that there is a single inst in the top array in `results_stack_`,
// pops that array, and returns the inst.
auto PopResult() -> SemIR::InstId {
CARBON_CHECK(results_stack_.PeekArray().size() == 1);
auto value_id = results_stack_.PeekArray()[0];
results_stack_.PopArray();
return value_id;
}
auto DoPostWork(State state, SemIR::SpecificConstant specific_constant,
WorkItem entry) -> void;
auto DoPostWork(State state, SemIR::ImportRefLoaded import_ref,
WorkItem entry) -> void;
// Performs the core logic of matching a variable pattern whose scrutinee
// type is `scrutinee_type_id`, but returns the scrutinee that its subpattern
@@ -223,6 +231,22 @@ class MatchContext {
SemIR::InstId scrutinee_id, WorkItem entry) const
-> SemIR::InstId;
// Performs the core logic of matching `entry.scrutinee_id` against the
// constant value of `entry.pattern_id`, rather than against the pattern inst
// itself. Factored out so that it can be reused across all inst kinds that
// can indirectly represent patterns, like `ImportRefLoaded` and
// `SpecificConstant`.
auto DoIndirectPreWorkImpl(State state, WorkItem entry) -> void;
// Asserts that there is a single inst in the top array in `results_stack_`,
// pops that array, and returns the inst.
auto PopResult() -> SemIR::InstId {
CARBON_CHECK(results_stack_.PeekArray().size() == 1);
auto value_id = results_stack_.PeekArray()[0];
results_stack_.PopArray();
return value_id;
}
// The stack of work to be processed.
llvm::SmallVector<WorkItem> stack_;
@@ -230,12 +254,26 @@ class MatchContext {
// a single result, which may have multiple sub-results.
ArrayStack<SemIR::InstId> results_stack_;
// The top of this stack represents the specific that should currently be
// substituted into pattern insts before matching them (if any). This is
// never empty, because the constructor pushes an initial specific that
// should never be popped.
llvm::SmallVector<SemIR::SpecificId> specific_id_stack_;
Context& context_;
};
} // namespace
auto MatchContext::Match(State state, WorkItem entry) -> void {
Diagnostics::AnnotationScope annotate_diagnostics(
&context_.emitter(), [&](auto& builder) {
if (std::holds_alternative<CallerState*>(state)) {
CARBON_DIAGNOSTIC(InCallToFunctionParam, Note,
"initializing function parameter");
builder.Note(entry.pattern_id, InCallToFunctionParam);
}
});
CARBON_CHECK(stack_.empty());
stack_.push_back(entry);
while (!stack_.empty()) {
@@ -351,13 +389,54 @@ auto MatchContext::DoPreWork(State state,
}
}
// Returns the substituted scrutinee type of `pattern_id` in `specific_id`. As
// with `GetTypeOfInstInSpecific`, this does not perform substitution, and it
// accepts `SpecificId::None`, treating it as a request for the value to use
// within the generic itself.
static auto GetScrutineeTypeInSpecific(const Context& context,
SemIR::InstId pattern_id,
SemIR::SpecificId specific_id)
-> SemIR::TypeId {
const auto& sem_ir = context.sem_ir();
return ExtractScrutineeType(
sem_ir, SemIR::GetTypeOfInstInSpecific(sem_ir, specific_id, pattern_id));
}
auto MatchContext::DoPostWork(State state,
SemIR::AnyBindingPattern binding_pattern,
WorkItem entry) -> void {
CARBON_CHECK(!std::holds_alternative<CallerState*>(state));
if (std::holds_alternative<ThunkState*>(state)) {
// Pass through the result from the subpattern.
return;
}
auto scrutinee_type_id = GetScrutineeTypeInSpecific(
context_, entry.pattern_id, specific_id_stack_.back());
auto value_id = PopResult();
if (value_id.has_value()) {
value_id = Convert(context_, SemIR::LocId(value_id), value_id,
{.kind = ConversionKindFor(binding_pattern, entry),
.type_id = scrutinee_type_id});
} else {
CARBON_CHECK(binding_pattern.kind == SemIR::SymbolicBindingPattern::Kind);
}
if (need_subpattern_results()) {
results_stack_.AppendToTop(value_id);
}
if (specific_id_stack_.size() > 1) {
// If anything has been pushed onto the specific stack, that means
// that `binding_pattern` is a constant inst. The bindings for constant
// insts are not precomputed (see the documentation for bind_name_map), so
// we have to emit a new one here.
context_.inst_block_stack().AddInstId(
AddBindingForPattern(context_, SemIR::LocId(entry.pattern_id),
binding_pattern, scrutinee_type_id, value_id));
return;
}
// We're logically consuming this map entry, so we invalidate it in order
// to avoid accidentally consuming it twice.
auto [bind_name_id, type_expr_region_id] =
@@ -367,34 +446,13 @@ auto MatchContext::DoPostWork(State state,
if (type_expr_region_id.has_value()) {
InsertHere(context_, type_expr_region_id);
}
auto value_id = PopResult();
CARBON_CHECK(bind_name_id.has_value(), "bind_name_map entry used twice");
if (value_id.has_value()) {
auto conversion_kind = ConversionKindFor(binding_pattern, entry);
if (!bind_name_id.has_value()) {
// TODO: Is this appropriate, or should we perform a conversion based on
// the category of the `_` binding first, and then separately discard the
// initializer for a `_` binding?
conversion_kind = ConversionTarget::Discarded;
}
value_id =
Convert(context_, SemIR::LocId(value_id), value_id,
{.kind = conversion_kind,
.type_id = context_.insts().Get(bind_name_id).type_id()});
} else {
CARBON_CHECK(binding_pattern.kind == SemIR::SymbolicBindingPattern::Kind);
}
if (bind_name_id.has_value()) {
auto bind_name = context_.insts().GetAs<SemIR::AnyBinding>(bind_name_id);
CARBON_CHECK(!bind_name.value_id.has_value());
bind_name.value_id = value_id;
ReplaceInstBeforeConstantUse(context_, bind_name_id, bind_name);
context_.inst_block_stack().AddInstId(bind_name_id);
}
if (need_subpattern_results()) {
results_stack_.AppendToTop(value_id);
}
auto bind_name = context_.insts().GetAs<SemIR::AnyBinding>(bind_name_id);
CARBON_CHECK(!bind_name.value_id.has_value());
bind_name.value_id = value_id;
ReplaceInstBeforeConstantUse(context_, bind_name_id, bind_name);
context_.inst_block_stack().AddInstId(bind_name_id);
}
// Returns the inst kind to use for the parameter corresponding to the given
@@ -413,35 +471,12 @@ static auto ParamKindFor(SemIR::Inst param_pattern) -> SemIR::InstKind {
}
}
// Returns the applicable specific for patterns handled in the given state.
static auto GetSpecific(State state) -> SemIR::SpecificId {
CARBON_KIND_SWITCH(state) {
case CARBON_KIND(CallerState* caller_state): {
return caller_state->callee_specific_id;
}
default: {
return SemIR::SpecificId::None;
}
}
}
// Returns the scrutinee type of `pattern_id` in the specific determined by
// `state`.
static auto GetSpecificScrutineeTypeId(const Context& context, State state,
SemIR::InstId pattern_id)
-> SemIR::TypeId {
const auto& sem_ir = context.sem_ir();
return ExtractScrutineeType(
sem_ir,
SemIR::GetTypeOfInstInSpecific(sem_ir, GetSpecific(state), pattern_id));
}
auto MatchContext::DoPreWork(State state, SemIR::AnyParamPattern param_pattern,
SemIR::InstId scrutinee_id, WorkItem entry)
-> void {
AddAsPostWork(entry);
auto scrutinee_type_id =
GetSpecificScrutineeTypeId(context_, state, entry.pattern_id);
auto scrutinee_type_id = GetScrutineeTypeInSpecific(
context_, entry.pattern_id, specific_id_stack_.back());
// If `param_pattern` is a `VarParamPattern`, match it as a `VarPattern` here,
// and then as a `RefParamPattern` below.
@@ -469,8 +504,7 @@ auto MatchContext::DoPreWork(State state, SemIR::AnyParamPattern param_pattern,
case CARBON_KIND(CalleeState* callee_state): {
SemIR::Inst param =
SemIR::AnyParam{.kind = ParamKindFor(param_pattern),
.type_id = ExtractScrutineeType(
context_.sem_ir(), param_pattern.type_id),
.type_id = scrutinee_type_id,
.index = callee_state->index.Allocate(),
.pretty_name_id = SemIR::GetPrettyNameFromPatternId(
context_.sem_ir(), entry.pattern_id)};
@@ -499,8 +533,10 @@ auto MatchContext::DoPreWork(State state, SemIR::AnyParamPattern param_pattern,
} else {
results_stack_.AppendToTop(param_id);
}
callee_state->PushCallParamPattern(context_, loc_id, entry.pattern_id,
specific_id_stack_.back());
callee_state->call_params.push_back(param_id);
callee_state->call_param_patterns.push_back(entry.pattern_id);
break;
}
case CARBON_KIND(ThunkState* thunk_state): {
@@ -579,11 +615,11 @@ auto MatchContext::DoPreWork(State state,
}
auto MatchContext::DoPostWork(State state,
SemIR::ReturnSlotPattern return_slot_pattern,
SemIR::ReturnSlotPattern /*return_slot_pattern*/,
WorkItem entry) -> void {
CARBON_CHECK(std::holds_alternative<CalleeState*>(state));
auto type_id =
ExtractScrutineeType(context_.sem_ir(), return_slot_pattern.type_id);
auto type_id = GetScrutineeTypeInSpecific(context_, entry.pattern_id,
specific_id_stack_.back());
auto return_slot_id = AddInst<SemIR::ReturnSlot>(
context_, SemIR::LocId(entry.pattern_id),
{.type_id = type_id,
@@ -602,8 +638,8 @@ auto MatchContext::DoPostWork(State state,
auto MatchContext::DoPreWork(State state, SemIR::VarPattern var_pattern,
SemIR::InstId scrutinee_id, WorkItem entry)
-> void {
auto scrutinee_type_id =
GetSpecificScrutineeTypeId(context_, state, entry.pattern_id);
auto scrutinee_type_id = GetScrutineeTypeInSpecific(
context_, entry.pattern_id, specific_id_stack_.back());
auto new_scrutinee_id =
DoVarPreWorkImpl(state, scrutinee_type_id, scrutinee_id, entry);
if (need_subpattern_results()) {
@@ -737,8 +773,8 @@ auto MatchContext::DoPreWork(State state, SemIR::TuplePattern tuple_pattern,
return;
}
auto tuple_type_id =
ExtractScrutineeType(context_.sem_ir(), tuple_pattern.type_id);
auto tuple_type_id = GetScrutineeTypeInSpecific(context_, entry.pattern_id,
specific_id_stack_.back());
auto converted_scrutinee_id = ConvertToValueOrRefOfType(
context_, SemIR::LocId(entry.pattern_id), scrutinee_id, tuple_type_id);
if (auto scrutinee_value = context_.insts().TryGetAs<SemIR::TupleValue>(
@@ -765,15 +801,15 @@ auto MatchContext::DoPreWork(State state, SemIR::TuplePattern tuple_pattern,
}
auto MatchContext::DoPostWork(State /*state*/,
SemIR::TuplePattern tuple_pattern, WorkItem entry)
-> void {
SemIR::TuplePattern /*tuple_pattern*/,
WorkItem entry) -> void {
auto elements_id = context_.inst_blocks().Add(results_stack_.PeekArray());
results_stack_.PopArray();
auto tuple_value_id =
AddInst<SemIR::TupleValue>(context_, SemIR::LocId(entry.pattern_id),
{.type_id = SemIR::ExtractScrutineeType(
context_.sem_ir(), tuple_pattern.type_id),
.elements_id = elements_id});
auto tuple_value_id = AddInst<SemIR::TupleValue>(
context_, SemIR::LocId(entry.pattern_id),
{.type_id = GetScrutineeTypeInSpecific(context_, entry.pattern_id,
specific_id_stack_.back()),
.elements_id = elements_id});
results_stack_.AppendToTop(tuple_value_id);
}
@@ -800,7 +836,9 @@ auto MatchContext::DoPreWork(State state, SemIR::SpliceInst splice,
context_.bundles().Add<SemIR::CalleePatternMatchAction::Args>(
{.pattern_id = entry.pattern_id,
.parent_index = callee_state->index.Allocate()})});
callee_state->call_param_patterns.push_back(entry.pattern_id);
callee_state->PushCallParamPattern(
context_, SemIR::LocId(entry.pattern_id), entry.pattern_id,
specific_id_stack_.back());
callee_state->call_params.push_back(result_id);
results_stack_.AppendToTop(result_id);
}
@@ -813,8 +851,8 @@ auto MatchContext::DoPreWork(State state,
-> void {
CARBON_CHECK(std::holds_alternative<CalleeState*>(state));
if (need_subpattern_results()) {
auto type_id =
ExtractScrutineeType(context_.sem_ir(), return_pattern.type_id);
auto type_id = GetScrutineeTypeInSpecific(context_, entry.pattern_id,
specific_id_stack_.back());
SemIR::InstKind result_kind =
return_pattern.kind == SemIR::RefReturnPattern::Kind
? SemIR::RefReturn::Kind
@@ -827,6 +865,55 @@ auto MatchContext::DoPreWork(State state,
}
}
auto MatchContext::DoIndirectPreWorkImpl(State state, WorkItem entry) -> void {
CARBON_CHECK(!std::holds_alternative<CalleeState*>(state));
AddAsPostWork(entry);
auto constant_id = SemIR::GetConstantValueInSpecific(
context_.sem_ir(), specific_id_stack_.back(), entry.pattern_id);
entry.pattern_id = context_.constant_values().GetInstId(constant_id);
// Now that we've substituted the current specific into the underlying
// pattern, we shouldn't try to redundantly reapply that specific while
// matching it.
specific_id_stack_.push_back(SemIR::SpecificId::None);
AddWork(entry);
}
auto MatchContext::DoPreWork(State state,
SemIR::SpecificConstant specific_constant,
SemIR::InstId /*scrutinee_id*/, WorkItem entry)
-> void {
if (std::holds_alternative<CalleeState*>(state)) {
// In callee pattern matching, we need to produce call parameter patterns
// with the SpecificConstant's specific substituted into them, so we need
// to defer substitution until that point. Instead, we treat that specific
// as the current specific while traversing the underlying inst.
AddAsPostWork(entry);
entry.pattern_id = specific_constant.inst_id;
specific_id_stack_.push_back(specific_constant.specific_id);
AddWork(entry);
} else {
DoIndirectPreWorkImpl(state, entry);
}
}
auto MatchContext::DoPostWork(State /*state*/,
SemIR::SpecificConstant /*specific_constant*/,
WorkItem /*entry*/) -> void {
specific_id_stack_.pop_back();
}
auto MatchContext::DoPreWork(State state, SemIR::ImportRefLoaded /*import_ref*/,
SemIR::InstId /*scrutinee_id*/, WorkItem entry)
-> void {
DoIndirectPreWorkImpl(state, entry);
}
auto MatchContext::DoPostWork(State /*state*/,
SemIR::ImportRefLoaded /*import_ref*/,
WorkItem /*entry*/) -> void {
specific_id_stack_.pop_back();
}
auto MatchContext::Dispatch(State state, WorkItem entry) -> void {
if (entry.pattern_id == SemIR::ErrorInst::InstId) {
if (need_subpattern_results()) {
@@ -834,14 +921,6 @@ auto MatchContext::Dispatch(State state, WorkItem entry) -> void {
}
return;
}
Diagnostics::AnnotationScope annotate_diagnostics(
&context_.emitter(), [&](auto& builder) {
if (std::holds_alternative<CallerState*>(state)) {
CARBON_DIAGNOSTIC(InCallToFunctionParam, Note,
"initializing function parameter");
builder.Note(entry.pattern_id, InCallToFunctionParam);
}
});
auto pattern = context_.insts().Get(entry.pattern_id);
CARBON_KIND_SWITCH(entry.work) {
case CARBON_KIND(PreWork work): {
@@ -885,6 +964,14 @@ auto MatchContext::Dispatch(State state, WorkItem entry) -> void {
DoPreWork(state, splice_inst, work.scrutinee_id, entry);
break;
}
case CARBON_KIND(SemIR::SpecificConstant specific_constant): {
DoPreWork(state, specific_constant, work.scrutinee_id, entry);
break;
}
case CARBON_KIND(SemIR::ImportRefLoaded import_ref): {
DoPreWork(state, import_ref, work.scrutinee_id, entry);
break;
}
default: {
CARBON_FATAL("Inst kind not handled: {0}", pattern.kind());
}
@@ -917,6 +1004,14 @@ auto MatchContext::Dispatch(State state, WorkItem entry) -> void {
DoPostWork(state, tuple_pattern, entry);
break;
}
case CARBON_KIND(SemIR::SpecificConstant specific_constant): {
DoPostWork(state, specific_constant, entry);
break;
}
case CARBON_KIND(SemIR::ImportRefLoaded import_ref): {
DoPostWork(state, import_ref, entry);
break;
}
default: {
CARBON_FATAL("Inst kind not handled: {0}", pattern.kind());
}
@@ -1025,8 +1120,8 @@ auto CallerPatternMatch(Context& context, SemIR::SpecificId specific_id,
llvm::ArrayRef<SemIR::InstId> arg_refs,
SemIR::InstId return_arg_id, bool is_desugared)
-> SemIR::InstBlockId {
CallerState state = {.callee_specific_id = specific_id};
MatchContext match(context);
CallerState state;
MatchContext match(context, specific_id);
// When we have a separate `self_arg_id`, we concatenate that onto the front
// of the arg_refs to match against the first parameter.