diff --git a/toolchain/check/convert.cpp b/toolchain/check/convert.cpp index b1aed909f3b4..c647361cf3c4 100644 --- a/toolchain/check/convert.cpp +++ b/toolchain/check/convert.cpp @@ -1207,18 +1207,18 @@ auto ConvertCallArgs(Context& context, SemIR::LocId call_loc_id, for (auto implicit_param_id : implicit_param_refs) { auto addr_pattern = context.insts().TryGetAs(implicit_param_id); - auto [param_id, param] = SemIR::Function::GetParamFromParamRefId( + auto param_info = SemIR::Function::GetParamFromParamRefId( context.sem_ir(), implicit_param_id); - if (param.name_id == SemIR::NameId::SelfValue) { + if (param_info.GetNameId(context.sem_ir()) == SemIR::NameId::SelfValue) { auto converted_self_id = ConvertSelf( context, call_loc_id, callee.callee_loc, callee_specific_id, - addr_pattern, param_id, param, self_id); + addr_pattern, param_info.inst_id, param_info.inst, self_id); if (converted_self_id == SemIR::InstId::BuiltinError) { return SemIR::InstBlockId::Invalid; } args.push_back(converted_self_id); } else { - CARBON_CHECK(!param.runtime_index.is_valid(), + CARBON_CHECK(!param_info.inst.runtime_index.is_valid(), "Unexpected implicit parameter passed at runtime"); } } @@ -1239,16 +1239,16 @@ auto ConvertCallArgs(Context& context, SemIR::LocId call_loc_id, // TODO: In general we need to perform pattern matching here to find the // argument corresponding to each parameter. - auto [param_id, param] = + auto param_info = SemIR::Function::GetParamFromParamRefId(context.sem_ir(), param_ref_id); - if (!param.runtime_index.is_valid()) { + if (!param_info.inst.runtime_index.is_valid()) { // Not a runtime parameter: we don't pass an argument. continue; } - auto param_type_id = - SemIR::GetTypeInSpecific(context.sem_ir(), callee_specific_id, - context.insts().Get(param_id).type_id()); + auto param_type_id = SemIR::GetTypeInSpecific( + context.sem_ir(), callee_specific_id, + context.insts().Get(param_info.inst_id).type_id()); // TODO: Convert to the proper expression category. For now, we assume // parameters are all `let` bindings. auto converted_arg_id = @@ -1257,7 +1257,8 @@ auto ConvertCallArgs(Context& context, SemIR::LocId call_loc_id, return SemIR::InstBlockId::Invalid; } - CARBON_CHECK(static_cast(args.size()) == param.runtime_index.index, + CARBON_CHECK(static_cast(args.size()) == + param_info.inst.runtime_index.index, "Parameters not numbered in order."); args.push_back(converted_arg_id); } diff --git a/toolchain/check/generic.cpp b/toolchain/check/generic.cpp index 184179cdbe53..8787c9562544 100644 --- a/toolchain/check/generic.cpp +++ b/toolchain/check/generic.cpp @@ -440,11 +440,9 @@ auto RequireGenericParamsOnType(Context& context, SemIR::InstBlockId block_id) return; } for (auto& inst_id : context.inst_blocks().Get(block_id)) { - // TODO: Change GetParamFromParamRefId to return the name instead of - // inspecting param.name_id. - auto [param_id, param] = + auto param_info = SemIR::Function::GetParamFromParamRefId(context.sem_ir(), inst_id); - if (param.name_id == SemIR::NameId::SelfValue) { + if (param_info.GetNameId(context.sem_ir()) == SemIR::NameId::SelfValue) { CARBON_DIAGNOSTIC(SelfParameterNotAllowed, Error, "`self` parameter only allowed on functions"); context.emitter().Emit(inst_id, SelfParameterNotAllowed); @@ -467,11 +465,9 @@ auto RequireGenericOrSelfImplicitFunctionParams(Context& context, return; } for (auto& inst_id : context.inst_blocks().Get(block_id)) { - // TODO: Change GetParamFromParamRefId to return the name instead of - // inspecting param.name_id. - auto [param_id, param] = + auto param_info = SemIR::Function::GetParamFromParamRefId(context.sem_ir(), inst_id); - if (param.name_id != SemIR::NameId::SelfValue && + if (param_info.GetNameId(context.sem_ir()) != SemIR::NameId::SelfValue && !context.constant_values().Get(inst_id).is_constant()) { CARBON_DIAGNOSTIC( ImplictParamMustBeConstant, Error, diff --git a/toolchain/check/handle_function.cpp b/toolchain/check/handle_function.cpp index f34becf11c47..d0e42f0f3b9b 100644 --- a/toolchain/check/handle_function.cpp +++ b/toolchain/check/handle_function.cpp @@ -78,35 +78,15 @@ static auto CheckFunctionSignature(Context& context, for (auto param_id : llvm::concat( context.inst_blocks().GetOrEmpty(name_and_params.implicit_params_id), context.inst_blocks().GetOrEmpty(name_and_params.params_id))) { - auto param = context.insts().Get(param_id); - // Find the parameter in the pattern. - // TODO: This duplicates work done by Function::GetParamFromParamRefId. - if (auto addr_pattern = param.TryAs()) { - param_id = addr_pattern->inner_id; - param = context.insts().Get(param_id); - } - - auto bind_name = param.TryAs(); - if (bind_name) { - param_id = bind_name->value_id; - param = context.insts().Get(param_id); - } - - auto param_inst = param.TryAs(); - if (!param_inst) { - // Once we support more generalized patterns we will need to diagnose - // parameters with unsupported patterns. - context.TODO(param_id, "unexpected syntax for parameter"); - // TODO: Also repair the param ID so downstream code doesn't need to deal - // with this. - continue; - } + auto param_info = + SemIR::Function::GetParamFromParamRefId(context.sem_ir(), param_id); // If this is a runtime parameter, number it. - if (bind_name && bind_name->kind == SemIR::BindName::Kind) { - param_inst->runtime_index = next_index; - context.ReplaceInstBeforeConstantUse(param_id, *param_inst); + if (param_info.bind_name && + param_info.bind_name->kind == SemIR::BindName::Kind) { + param_info.inst.runtime_index = next_index; + context.ReplaceInstBeforeConstantUse(param_info.inst_id, param_info.inst); ++next_index.index; } } @@ -380,17 +360,18 @@ static auto HandleFunctionDefinitionAfterSignature( for (auto param_ref_id : llvm::concat( context.inst_blocks().GetOrEmpty(function.implicit_param_refs_id), context.inst_blocks().GetOrEmpty(function.param_refs_id))) { - auto [param_id, param] = + auto param_info = SemIR::Function::GetParamFromParamRefId(context.sem_ir(), param_ref_id); // The parameter types need to be complete. - context.TryToCompleteType(param.type_id, [&] { + context.TryToCompleteType(param_info.inst.type_id, [&] { CARBON_DIAGNOSTIC( IncompleteTypeInFunctionParam, Error, "parameter has incomplete type `{0}` in function definition", SemIR::TypeId); - return context.emitter().Build(param_id, IncompleteTypeInFunctionParam, - param.type_id); + return context.emitter().Build(param_info.inst_id, + IncompleteTypeInFunctionParam, + param_info.inst.type_id); }); } diff --git a/toolchain/check/import_ref.cpp b/toolchain/check/import_ref.cpp index 00631d274a14..4dadc17170e2 100644 --- a/toolchain/check/import_ref.cpp +++ b/toolchain/check/import_ref.cpp @@ -714,7 +714,8 @@ class ImportRefResolver { llvm::SmallVector new_param_refs; for (auto ref_id : param_refs) { // Figure out the param structure. This echoes - // Function::GetParamFromParamRefId. + // Function::GetParamFromParamRefId, and could use that function if we + // added `bool addr` and `InstId bind_inst_id` to its return `ParamInfo`. // TODO: Consider a different parameter handling to simplify import logic. auto inst = import_ir_.insts().Get(ref_id); auto addr_inst = inst.TryAs(); diff --git a/toolchain/check/member_access.cpp b/toolchain/check/member_access.cpp index aa8cf51cddaa..25772c32a286 100644 --- a/toolchain/check/member_access.cpp +++ b/toolchain/check/member_access.cpp @@ -91,9 +91,10 @@ static auto IsInstanceMethod(const SemIR::File& sem_ir, const auto& function = sem_ir.functions().Get(function_id); for (auto param_id : sem_ir.inst_blocks().GetOrEmpty(function.implicit_param_refs_id)) { - auto param = - SemIR::Function::GetParamFromParamRefId(sem_ir, param_id).second; - if (param.name_id == SemIR::NameId::SelfValue) { + auto param_name_id = + SemIR::Function::GetParamFromParamRefId(sem_ir, param_id) + .GetNameId(sem_ir); + if (param_name_id == SemIR::NameId::SelfValue) { return true; } } diff --git a/toolchain/lower/file_context.cpp b/toolchain/lower/file_context.cpp index c95820c3133a..79af16bef963 100644 --- a/toolchain/lower/file_context.cpp +++ b/toolchain/lower/file_context.cpp @@ -213,12 +213,13 @@ auto FileContext::BuildFunctionDecl(SemIR::FunctionId function_id) } for (auto param_ref_id : llvm::concat(implicit_param_refs, param_refs)) { - auto [param_id, param] = + auto param_info = SemIR::Function::GetParamFromParamRefId(sem_ir(), param_ref_id); - if (!param.runtime_index.is_valid()) { + if (!param_info.inst.runtime_index.is_valid()) { continue; } - switch (auto value_rep = SemIR::ValueRepr::ForType(sem_ir(), param.type_id); + switch (auto value_rep = + SemIR::ValueRepr::ForType(sem_ir(), param_info.inst.type_id); value_rep.kind) { case SemIR::ValueRepr::Unknown: CARBON_FATAL("Incomplete parameter type lowering function declaration"); @@ -264,7 +265,7 @@ auto FileContext::BuildFunctionDecl(SemIR::FunctionId function_id) llvm::Attribute::getWithStructRetType(llvm_context(), return_type)); } else { name_id = SemIR::Function::GetParamFromParamRefId(sem_ir(), inst_id) - .second.name_id; + .GetNameId(sem_ir()); } arg.setName(sem_ir().names().GetIRBaseName(name_id)); } @@ -311,14 +312,14 @@ auto FileContext::BuildFunctionDefinition(SemIR::FunctionId function_id) } for (auto param_ref_id : llvm::concat(implicit_param_refs, param_refs)) { - auto [param_id, param] = + auto param_info = SemIR::Function::GetParamFromParamRefId(sem_ir(), param_ref_id); - if (!param.runtime_index.is_valid()) { + if (!param_info.inst.runtime_index.is_valid()) { continue; } // Get the value of the parameter from the function argument. - auto param_type_id = param.type_id; + auto param_type_id = param_info.inst.type_id; llvm::Value* param_value = llvm::PoisonValue::get(GetType(param_type_id)); if (SemIR::ValueRepr::ForType(sem_ir(), param_type_id).kind != SemIR::ValueRepr::None) { @@ -327,7 +328,7 @@ auto FileContext::BuildFunctionDefinition(SemIR::FunctionId function_id) } // The value of the parameter is the value of the argument. - function_lowering.SetLocal(param_id, param_value); + function_lowering.SetLocal(param_info.inst_id, param_value); // Match the portion of the pattern corresponding to the parameter against // the parameter value. For now this is always a single name binding, diff --git a/toolchain/sem_ir/function.cpp b/toolchain/sem_ir/function.cpp index 2b667b9ba35d..e99a93824f08 100644 --- a/toolchain/sem_ir/function.cpp +++ b/toolchain/sem_ir/function.cpp @@ -43,20 +43,28 @@ auto GetCalleeFunction(const File& sem_ir, InstId callee_id) -> CalleeFunction { } auto Function::GetParamFromParamRefId(const File& sem_ir, InstId param_ref_id) - -> std::pair { + -> ParamInfo { auto ref = sem_ir.insts().Get(param_ref_id); - if (auto addr_pattern = ref.TryAs()) { + if (auto addr_pattern = ref.TryAs()) { param_ref_id = addr_pattern->inner_id; ref = sem_ir.insts().Get(param_ref_id); } - if (auto bind_name = ref.TryAs()) { + auto bind_name = ref.TryAs(); + if (bind_name) { param_ref_id = bind_name->value_id; ref = sem_ir.insts().Get(param_ref_id); } + return {param_ref_id, ref.As(), bind_name}; +} - return {param_ref_id, ref.As()}; +auto Function::ParamInfo::GetNameId(const File& sem_ir) -> NameId { + if (bind_name) { + return sem_ir.entity_names().Get(bind_name->entity_name_id).name_id; + } else { + return NameId::Invalid; + } } auto Function::GetDeclaredReturnType(const File& file, diff --git a/toolchain/sem_ir/function.h b/toolchain/sem_ir/function.h index 4fd512a32363..5b1f3bb92bb5 100644 --- a/toolchain/sem_ir/function.h +++ b/toolchain/sem_ir/function.h @@ -65,10 +65,18 @@ struct Function : public EntityWithParamsBase, } // Given a parameter reference instruction from `param_refs_id` or - // `implicit_param_refs_id`, returns the corresponding `Param` instruction - // and its ID. + // `implicit_param_refs_id`, returns a `ParamInfo` value with the + // corresponding instruction, its ID, and the name binding, if present. + struct ParamInfo { + InstId inst_id; + Param inst; + std::optional bind_name; + + // Gets the name from `bind_name`. Returns invalid if that is not present. + auto GetNameId(const File& sem_ir) -> NameId; + }; static auto GetParamFromParamRefId(const File& sem_ir, InstId param_ref_id) - -> std::pair; + -> ParamInfo; // Gets the declared return type for a specific version of this function, or // the canonical return type for the original declaration no specific is