Add Call param patterns to Function (#6586)

This commit is contained in:
Geoff Romer
2026-01-13 19:30:15 +00:00
committed by GitHub
parent 4a47f1ebeb
commit a2737a3189
17 changed files with 409 additions and 298 deletions
+95 -36
View File
@@ -68,11 +68,20 @@ class MatchContext {
// Adds a work item to the stack.
auto AddWork(WorkItem work_item) -> void { stack_.push_back(work_item); }
// Processes all work items on the stack. When performing caller pattern
// matching, returns an inst block with one inst reference for each
// calling-convention argument. When performing callee pattern matching,
// returns an inst block with references to all the emitted BindName insts.
auto DoWork(Context& context) -> SemIR::InstBlockId;
// Processes all work items on the stack.
auto DoWork(Context& context) -> void;
// Returns an inst block of references to all the emitted `Call` arguments.
// Can only be called once, at the end of Caller pattern matching.
auto CallerResults(Context& context) && -> SemIR::InstBlockId;
// Returns an inst block of references to all the emitted `Call` params,
// and an inst block of references to the `Call` param patterns they were
// emitted to match. Can only be called once, at the end of Callee pattern
// matching.
auto CalleeResults(Context& context) && -> CalleePatternMatchResults;
~MatchContext();
private:
// Emits the pattern-match insts necessary to match the pattern inst
@@ -115,12 +124,17 @@ class MatchContext {
// The stack of work to be processed.
llvm::SmallVector<WorkItem> stack_;
// The pending results that will be returned by the current `DoWork` call.
// It represents the contents of the `Call` arguments block when kind_
// is Caller, or the `Call` parameters block when kind_ is Callee
// (it is empty when kind_ is Local). Consequently, it is populated
// only by DoEmitPatternMatch for *ParamPattern insts.
llvm::SmallVector<SemIR::InstId> results_;
// The in-progress contents of the `Call` arguments block. This is populated
// only when kind_ is Caller.
llvm::SmallVector<SemIR::InstId> call_args_;
// The in-progress contents of the `Call` parameters block. This is populated
// only when kind_ is Callee.
llvm::SmallVector<SemIR::InstId> call_params_;
// The in-progress contents of the `Call` parameter patterns block. This is
// populated only when kind_ is Callee.
llvm::SmallVector<SemIR::InstId> call_param_patterns_;
// The kind of pattern match being performed.
MatchKind kind_;
@@ -131,16 +145,55 @@ class MatchContext {
} // namespace
auto MatchContext::DoWork(Context& context) -> SemIR::InstBlockId {
results_.reserve(stack_.size());
auto MatchContext::DoWork(Context& context) -> void {
CARBON_CHECK(call_args_.empty() && call_params_.empty() &&
call_param_patterns_.empty());
switch (kind_) {
case MatchKind::Caller: {
call_args_.reserve(stack_.size());
break;
}
case MatchKind::Callee: {
call_param_patterns_.reserve(stack_.size());
call_params_.reserve(stack_.size());
break;
}
case MatchKind::Local:
break;
}
while (!stack_.empty()) {
EmitPatternMatch(context, stack_.pop_back_val());
}
auto block_id = context.inst_blocks().Add(results_);
results_.clear();
}
auto MatchContext::CallerResults(Context& context) && -> SemIR::InstBlockId {
CARBON_CHECK(kind_ == MatchKind::Caller);
auto block_id = context.inst_blocks().Add(call_args_);
call_args_.clear();
return block_id;
}
auto MatchContext::CalleeResults(
Context& context) && -> CalleePatternMatchResults {
CARBON_CHECK(kind_ == MatchKind::Callee);
CARBON_CHECK(call_params_.size() == call_param_patterns_.size());
auto call_param_patterns_id = context.inst_blocks().Add(call_param_patterns_);
call_param_patterns_.clear();
auto call_params_id = context.inst_blocks().Add(call_params_);
call_params_.clear();
return {.call_param_patterns_id = call_param_patterns_id,
.call_params_id = call_params_id};
}
MatchContext::~MatchContext() {
CARBON_CHECK(call_args_.empty() && call_params_.empty() &&
call_param_patterns_.empty(),
"Unhandled pattern matching outputs. call_args_.size(): {0}, "
"call_params_.size(): {1}, call_param_patterns_.size(): {2}",
call_args_.size(), call_params_.size(),
call_param_patterns_.size());
}
// Inserts the given region into the current code block. If the region
// consists of a single block, this will be implemented as a `splice_block`
// inst. Otherwise, this will end the current block with a branch to the entry
@@ -250,14 +303,14 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
switch (kind_) {
case MatchKind::Caller: {
CARBON_CHECK(
static_cast<size_t>(param_pattern.index.index) == results_.size(),
"Parameters out of order; expecting {0} but got {1}", results_.size(),
param_pattern.index.index);
static_cast<size_t>(param_pattern.index.index) == call_args_.size(),
"Parameters out of order; expecting {0} but got {1}",
call_args_.size(), param_pattern.index.index);
CARBON_CHECK(entry.scrutinee_id.has_value());
if (entry.scrutinee_id == SemIR::ErrorInst::InstId) {
results_.push_back(SemIR::ErrorInst::InstId);
call_args_.push_back(SemIR::ErrorInst::InstId);
} else {
results_.push_back(ConvertToValueOfType(
call_args_.push_back(ConvertToValueOfType(
context, SemIR::LocId(entry.scrutinee_id), entry.scrutinee_id,
ExtractScrutineeType(
context.sem_ir(),
@@ -278,7 +331,8 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
context.sem_ir(), entry.pattern_id)});
AddWork({.pattern_id = param_pattern.subpattern_id,
.scrutinee_id = param_id});
results_.push_back(param_id);
call_params_.push_back(param_id);
call_param_patterns_.push_back(entry.pattern_id);
break;
}
case MatchKind::Local: {
@@ -296,20 +350,20 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
switch (kind_) {
case MatchKind::Caller: {
CARBON_CHECK(
static_cast<size_t>(param_pattern.index.index) == results_.size(),
"Parameters out of order; expecting {0} but got {1}", results_.size(),
param_pattern.index.index);
static_cast<size_t>(param_pattern.index.index) == call_args_.size(),
"Parameters out of order; expecting {0} but got {1}",
call_args_.size(), param_pattern.index.index);
CARBON_CHECK(entry.scrutinee_id.has_value());
if (std::is_same_v<RefParamPatternT, SemIR::VarParamPattern>) {
results_.push_back(entry.scrutinee_id);
call_args_.push_back(entry.scrutinee_id);
break;
}
auto scrutinee_type_id = ExtractScrutineeType(
context.sem_ir(),
SemIR::GetTypeOfInstInSpecific(context.sem_ir(), callee_specific_id_,
entry.pattern_id));
results_.push_back(Convert(
call_args_.push_back(Convert(
context, SemIR::LocId(entry.scrutinee_id), entry.scrutinee_id,
{.kind = entry.allow_unmarked_ref ? ConversionTarget::UnmarkedRefParam
: ConversionTarget::RefParam,
@@ -328,7 +382,8 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
context.sem_ir(), entry.pattern_id)});
AddWork({.pattern_id = param_pattern.subpattern_id,
.scrutinee_id = param_id});
results_.push_back(param_id);
call_params_.push_back(param_id);
call_param_patterns_.push_back(entry.pattern_id);
break;
}
case MatchKind::Local: {
@@ -343,9 +398,9 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
switch (kind_) {
case MatchKind::Caller: {
CARBON_CHECK(
static_cast<size_t>(param_pattern.index.index) == results_.size(),
"Parameters out of order; expecting {0} but got {1}", results_.size(),
param_pattern.index.index);
static_cast<size_t>(param_pattern.index.index) == call_args_.size(),
"Parameters out of order; expecting {0} but got {1}",
call_args_.size(), param_pattern.index.index);
CARBON_CHECK(entry.scrutinee_id.has_value());
CARBON_CHECK(
context.insts().Get(entry.scrutinee_id).type_id() ==
@@ -353,7 +408,7 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
context.sem_ir(),
SemIR::GetTypeOfInstInSpecific(
context.sem_ir(), callee_specific_id_, entry.pattern_id)));
results_.push_back(entry.scrutinee_id);
call_args_.push_back(entry.scrutinee_id);
// Do not traverse farther, because the caller side of the pattern
// ends here.
break;
@@ -370,7 +425,8 @@ auto MatchContext::DoEmitPatternMatch(Context& context,
context.sem_ir(), entry.pattern_id)});
AddWork({.pattern_id = param_pattern.subpattern_id,
.scrutinee_id = param_id});
results_.push_back(param_id);
call_param_patterns_.push_back(entry.pattern_id);
call_params_.push_back(param_id);
break;
}
case MatchKind::Local: {
@@ -588,10 +644,11 @@ auto CalleePatternMatch(Context& context,
SemIR::InstBlockId implicit_param_patterns_id,
SemIR::InstBlockId param_patterns_id,
SemIR::InstBlockId return_patterns_id)
-> SemIR::InstBlockId {
-> CalleePatternMatchResults {
if (!return_patterns_id.has_value() && !param_patterns_id.has_value() &&
!implicit_param_patterns_id.has_value()) {
return SemIR::InstBlockId::None;
return {.call_param_patterns_id = SemIR::InstBlockId::None,
.call_params_id = SemIR::InstBlockId::None};
}
MatchContext match(MatchKind::Callee);
@@ -620,7 +677,8 @@ auto CalleePatternMatch(Context& context,
}
}
return match.DoWork(context);
match.DoWork(context);
return std::move(match).CalleeResults(context);
}
auto CallerPatternMatch(Context& context, SemIR::SpecificId specific_id,
@@ -661,7 +719,8 @@ auto CallerPatternMatch(Context& context, SemIR::SpecificId specific_id,
.allow_unmarked_ref = true});
}
return match.DoWork(context);
match.DoWork(context);
return std::move(match).CallerResults(context);
}
auto LocalPatternMatch(Context& context, SemIR::InstId pattern_id,