Convert the scrutinee of a binding pattern to the right category. (#5662)

When the binding pattern appears within a `var` pattern, convert to a
reference. Otherwise, convert to a value.

This gets the advent of code examples to produce the right answers again
:)

---------

Co-authored-by: Geoff Romer <gromer@google.com>
This commit is contained in:
Richard Smith
2025-06-17 16:45:19 +00:00
committed by GitHub
co-authored by Geoff Romer
parent 1c09de9b87
commit 80529aaef9
29 changed files with 317 additions and 260 deletions
+55 -22
View File
@@ -48,10 +48,13 @@ class MatchContext {
SemIR::InstId pattern_id;
// `None` when processing the callee side.
SemIR::InstId scrutinee_id;
// Whether we are in a context where plain bindings are reference bindings.
// This happens in var patterns.
bool ref_binding_context;
auto Print(llvm::raw_ostream& out) const -> void {
out << "{pattern_id: " << pattern_id << ", scrutinee_id: " << scrutinee_id
<< "}";
<< ", ref_binding_context: " << ref_binding_context << "}";
}
};
@@ -224,10 +227,18 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
InsertHere(context, type_expr_region_id);
auto value_id = SemIR::InstId::None;
if (kind_ == MatchKind::Local) {
auto conversion_kind = entry.ref_binding_context
? ConversionTarget::DurableRef
: ConversionTarget::Value;
if (!bind_name_id.has_value()) {
// TODO: Is this appropriate, or should we perform a conversion based on
// whether the `_` binding is a value or ref binding first, and then
// separately discard the initializer for a `_` binding?
conversion_kind = ConversionTarget::Discarded;
}
value_id =
Convert(context, SemIR::LocId(entry.scrutinee_id), entry.scrutinee_id,
{.kind = bind_name_id.has_value() ? ConversionTarget::ValueOrRef
: ConversionTarget::Discarded,
{.kind = conversion_kind,
.type_id = context.insts().Get(bind_name_id).type_id()});
} else {
// In a function call, conversion is handled while matching the enclosing
@@ -253,7 +264,8 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
// the caller side of the pattern, so we traverse without emitting any
// insts.
AddWork({.pattern_id = addr_pattern.inner_id,
.scrutinee_id = SemIR::InstId::None});
.scrutinee_id = SemIR::InstId::None,
.ref_binding_context = false});
return;
}
CARBON_CHECK(entry.scrutinee_id.has_value());
@@ -279,13 +291,16 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
context, SemIR::LocId(scrutinee_ref_id),
{.type_id = GetPointerType(context, scrutinee_ref_type_inst_id),
.lvalue_id = scrutinee_ref_id});
AddWork({.pattern_id = addr_pattern.inner_id, .scrutinee_id = new_scrutinee});
AddWork({.pattern_id = addr_pattern.inner_id,
.scrutinee_id = new_scrutinee,
.ref_binding_context = false});
}
auto MatchContext::DoEmitPatternMatch(Context& context,
SemIR::ValueParamPattern param_pattern,
SemIR::InstId pattern_inst_id,
WorkItem entry) -> void {
CARBON_CHECK(!entry.ref_binding_context);
switch (kind_) {
case MatchKind::Caller: {
CARBON_CHECK(
@@ -320,7 +335,8 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
.pretty_name_id = SemIR::GetPrettyNameFromPatternId(
context.sem_ir(), entry.pattern_id)});
AddWork({.pattern_id = param_pattern.subpattern_id,
.scrutinee_id = param_id});
.scrutinee_id = param_id,
.ref_binding_context = entry.ref_binding_context});
results_.push_back(param_id);
break;
}
@@ -334,6 +350,7 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
SemIR::RefParamPattern param_pattern,
SemIR::InstId pattern_inst_id,
WorkItem entry) -> void {
CARBON_CHECK(entry.ref_binding_context);
switch (kind_) {
case MatchKind::Caller: {
CARBON_CHECK(
@@ -362,7 +379,8 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
.pretty_name_id = SemIR::GetPrettyNameFromPatternId(
context.sem_ir(), entry.pattern_id)});
AddWork({.pattern_id = param_pattern.subpattern_id,
.scrutinee_id = param_id});
.scrutinee_id = param_id,
.ref_binding_context = entry.ref_binding_context});
results_.push_back(param_id);
break;
}
@@ -376,6 +394,7 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
SemIR::OutParamPattern param_pattern,
SemIR::InstId pattern_inst_id,
WorkItem entry) -> void {
CARBON_CHECK(!entry.ref_binding_context);
switch (kind_) {
case MatchKind::Caller: {
CARBON_CHECK(
@@ -408,7 +427,8 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
.pretty_name_id = SemIR::GetPrettyNameFromPatternId(
context.sem_ir(), entry.pattern_id)});
AddWork({.pattern_id = param_pattern.subpattern_id,
.scrutinee_id = param_id});
.scrutinee_id = param_id,
.ref_binding_context = entry.ref_binding_context});
results_.push_back(param_id);
break;
}
@@ -447,7 +467,8 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
// the caller side of the pattern, so we traverse without emitting any
// insts.
AddWork({.pattern_id = var_pattern.subpattern_id,
.scrutinee_id = SemIR::InstId::None});
.scrutinee_id = SemIR::InstId::None,
.ref_binding_context = true});
return;
}
case MatchKind::Local: {
@@ -481,8 +502,9 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
AddInst<SemIR::Assign>(context, SemIR::LocId(pattern_inst_id),
{.lhs_id = storage_id, .rhs_id = init_id});
}
AddWork(
{.pattern_id = var_pattern.subpattern_id, .scrutinee_id = storage_id});
AddWork({.pattern_id = var_pattern.subpattern_id,
.scrutinee_id = storage_id,
.ref_binding_context = true});
if (context.scope_stack().PeekIndex() == ScopeIndex::Package) {
context.global_init().Suspend();
}
@@ -500,8 +522,9 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
[&](llvm::ArrayRef<SemIR::InstId> subscrutinee_ids) {
for (auto [subpattern_id, subscrutinee_id] :
llvm::reverse(llvm::zip(subpattern_ids, subscrutinee_ids))) {
AddWork(
{.pattern_id = subpattern_id, .scrutinee_id = subscrutinee_id});
AddWork({.pattern_id = subpattern_id,
.scrutinee_id = subscrutinee_id,
.ref_binding_context = entry.ref_binding_context});
}
};
if (!entry.scrutinee_id.has_value()) {
@@ -627,22 +650,25 @@ auto CalleePatternMatch(Context& context,
// in the original order.
if (return_slot_pattern_id.has_value()) {
match.AddWork({.pattern_id = return_slot_pattern_id,
.scrutinee_id = SemIR::InstId::None});
.scrutinee_id = SemIR::InstId::None,
.ref_binding_context = false});
}
if (param_patterns_id.has_value()) {
for (SemIR::InstId inst_id :
llvm::reverse(context.inst_blocks().Get(param_patterns_id))) {
match.AddWork(
{.pattern_id = inst_id, .scrutinee_id = SemIR::InstId::None});
match.AddWork({.pattern_id = inst_id,
.scrutinee_id = SemIR::InstId::None,
.ref_binding_context = false});
}
}
if (implicit_param_patterns_id.has_value()) {
for (SemIR::InstId inst_id :
llvm::reverse(context.inst_blocks().Get(implicit_param_patterns_id))) {
match.AddWork(
{.pattern_id = inst_id, .scrutinee_id = SemIR::InstId::None});
match.AddWork({.pattern_id = inst_id,
.scrutinee_id = SemIR::InstId::None,
.ref_binding_context = false});
}
}
@@ -663,17 +689,22 @@ auto CallerPatternMatch(Context& context, SemIR::SpecificId specific_id,
if (return_slot_arg_id.has_value()) {
CARBON_CHECK(return_slot_pattern_id.has_value());
match.AddWork({.pattern_id = return_slot_pattern_id,
.scrutinee_id = return_slot_arg_id});
.scrutinee_id = return_slot_arg_id,
.ref_binding_context = false});
}
// Check type conversions per-element.
for (auto [arg_id, param_pattern_id] : llvm::reverse(llvm::zip_equal(
arg_refs, context.inst_blocks().GetOrEmpty(param_patterns_id)))) {
match.AddWork({.pattern_id = param_pattern_id, .scrutinee_id = arg_id});
match.AddWork({.pattern_id = param_pattern_id,
.scrutinee_id = arg_id,
.ref_binding_context = false});
}
if (self_pattern_id.has_value()) {
match.AddWork({.pattern_id = self_pattern_id, .scrutinee_id = self_arg_id});
match.AddWork({.pattern_id = self_pattern_id,
.scrutinee_id = self_arg_id,
.ref_binding_context = false});
}
return match.DoWork(context);
@@ -682,7 +713,9 @@ auto CallerPatternMatch(Context& context, SemIR::SpecificId specific_id,
auto LocalPatternMatch(Context& context, SemIR::InstId pattern_id,
SemIR::InstId scrutinee_id) -> void {
MatchContext match(MatchKind::Local);
match.AddWork({.pattern_id = pattern_id, .scrutinee_id = scrutinee_id});
match.AddWork({.pattern_id = pattern_id,
.scrutinee_id = scrutinee_id,
.ref_binding_context = false});
match.DoWork(context);
}