diff --git a/src/compile/c_gen.cpp b/src/compile/c_gen.cpp index 4509622..cccc39a 100644 --- a/src/compile/c_gen.cpp +++ b/src/compile/c_gen.cpp @@ -558,19 +558,22 @@ void CGen::GenerateDecls(const SyntaxTreeInterfacePtr &chunk, GenResult &gr) { gr.function_names[pkg_func_name] = JitFunctionInfo{static_cast(func->params.size()), is_vararg, name}; } - // 如果原始函数含有数学参数,声明其特化变体 - if (const auto math_it = ir().math_param_positions.find(func->name); math_it != ir().math_param_positions.end()) { - const auto &math_params = math_it->second; - const int num_specs = 1 << static_cast(math_params.size()); - for (int bitmask = 0; bitmask < num_specs; ++bitmask) { - const auto spec_name = SpecFuncName(name, math_params, bitmask); - const auto spec_ret = GetSpecReturnType(func->name, bitmask); - Out() << SpecReturnCTypeName(spec_ret) << " " << spec_name << "("; - EmitSpecParamList(cparams, math_params, bitmask); - Out() << ");\n"; - // 注册特化函数名,使 CompileFunctioncall 能将其识别为 - // 本地调用(同文件直接调用)。 - gr.function_names[spec_name] = JitFunctionInfo{static_cast(func->params.size()), is_vararg, spec_name}; + // 数学特化只挂在文件级函数上。嵌套函数与文件级函数同名时不能套用那份特化, + // 否则多返回值的嵌套函数会生成 int64_t __fl_func_N_0() { return FlMakeMulti(...); }。 + if (func->parent == nullptr) { + if (const auto math_it = ir().math_param_positions.find(func->name); math_it != ir().math_param_positions.end()) { + const auto &math_params = math_it->second; + const int num_specs = 1 << static_cast(math_params.size()); + for (int bitmask = 0; bitmask < num_specs; ++bitmask) { + const auto spec_name = SpecFuncName(name, math_params, bitmask); + const auto spec_ret = GetSpecReturnType(func->name, bitmask); + Out() << SpecReturnCTypeName(spec_ret) << " " << spec_name << "("; + EmitSpecParamList(cparams, math_params, bitmask); + Out() << ");\n"; + // 注册特化函数名,使 CompileFunctioncall 能将其识别为 + // 本地调用(同文件直接调用)。 + gr.function_names[spec_name] = JitFunctionInfo{static_cast(func->params.size()), is_vararg, spec_name}; + } } } } @@ -661,7 +664,9 @@ void CGen::GenerateImpl(const SyntaxTreeInterfacePtr &chunk, GenResult &gr) { const auto func_block = funcbody_ptr->Block(); const auto c_func_params = CIdents(func_params); - if (const auto math_it = ir().math_param_positions.find(func->name); math_it != ir().math_param_positions.end()) { + // 数学特化只挂在文件级函数上。嵌套函数即使与文件级函数同名也不套用那份特化。 + const auto math_it = func->parent == nullptr ? ir().math_param_positions.find(func->name) : ir().math_param_positions.end(); + if (math_it != ir().math_param_positions.end()) { const auto &math_params = math_it->second; const int num_specs = 1 << static_cast(math_params.size()); for (int bitmask = 0; bitmask < num_specs; ++bitmask) { @@ -2748,26 +2753,26 @@ std::string CGen::CompileVar(const SyntaxTreeInterfacePtr &v) { const auto &name = v_ptr->GetName(); DEBUG_ASSERT(cur_section_ != Section::Globals); - // 1. Check if captured (local variable, parameter, or loop variable) + // 1. 被闭包捕获的局部变量 / 形参 / 嵌套 local function。 if (const auto it = var_to_def_map_.find(v_ptr.get()); it != var_to_def_map_.end()) { VarDef *def = it->second; if (def->is_captured) { if (def->defining_func == cur_func_info_) { return "(*__box_" + CIdent(name) + ")"; - } else { - if (cur_func_info_) { - const auto vit = std::ranges::find(cur_func_info_->captured_vars, def); - if (vit != cur_func_info_->captured_vars.end()) { - int idx = static_cast(vit - cur_func_info_->captured_vars.begin()); - return std::format("(*_CL->upvalues[{}])", idx); - } + } + if (cur_func_info_) { + const auto vit = std::ranges::find(cur_func_info_->captured_vars, def); + if (vit != cur_func_info_->captured_vars.end()) { + int idx = static_cast(vit - cur_func_info_->captured_vars.begin()); + return std::format("(*_CL->upvalues[{}])", idx); } } } } - // 2. Check if function referenced as value (non-direct call) - if (local_func_names_.contains(name)) { + // 2. 没有遮蔽它的局部值时,裸函数名引用文件级函数。 + // 形参、局部变量、嵌套 local function 即使与文件级函数同名,也走下面的局部变量。 + if (!BindsLocalValue(v_ptr.get()) && local_func_names_.contains(name)) { const auto &info = local_func_names_.at(name); const std::string &csym = info.c_symbol_name.empty() ? name : info.c_symbol_name; return std::format("FlMakeClosure(_S, (void*){}, 0, {}, {})", csym, info.params_count, info.is_vararg ? "true" : "false"); @@ -3059,6 +3064,9 @@ std::string CGen::TryCompileNativeSpecCallExpr(const SyntaxTreeInterfacePtr &fun DEBUG_ASSERT(callee_pe && callee_pe->GetPrefixKind() == PrefixExpKind::kVar); const auto callee_var = std::dynamic_pointer_cast(callee_pe->GetValue()); DEBUG_ASSERT(callee_var && callee_var->GetVarKind() == VarKind::kSimple); + if (BindsLocalValue(callee_var.get())) { + return {}; + } const auto &callee_name = callee_var->GetName(); const auto explist_ptr = std::dynamic_pointer_cast(args_ptr->Explist()); DEBUG_ASSERT(explist_ptr); @@ -3189,14 +3197,9 @@ std::string CGen::CompileFunctioncall(const SyntaxTreeInterfacePtr &functioncall } std::string call_expr; - bool is_local_callee = false; - if (var_ptr) { - if (const auto it = var_to_def_map_.find(var_ptr); it != var_to_def_map_.end()) { - is_local_callee = true; - } - } + bool is_local_callee = BindsLocalValue(var_ptr); - if (local_func_names_.contains(func_name)) { + if (!is_local_callee && local_func_names_.contains(func_name)) { const auto &info = local_func_names_.at(func_name); if (!info.is_vararg && !has_expansion) { if (static_cast(compiled_args.size()) != info.params_count) { @@ -3940,6 +3943,10 @@ std::string CGen::TryCompileSpecDirectCall(const std::shared_ptrGetPrefixKind() == PrefixExpKind::kVar && args_ptr->GetArgsKind() == ArgsKind::kExpList) { if (const auto callee_var = std::dynamic_pointer_cast(pe_pre_ptr->GetValue()); callee_var && callee_var->GetVarKind() == VarKind::kSimple) { const auto &callee_name = callee_var->GetName(); + // 形参、局部变量、嵌套 local function 不能按文件级函数名做特化直调。 + if (BindsLocalValue(callee_var.get())) { + return ""; + } if (const auto math_it = ir().math_param_positions.find(callee_name); math_it != ir().math_param_positions.end()) { const auto &math_params = math_it->second; const auto explist_arg = args_ptr->Explist(); @@ -4937,6 +4944,20 @@ std::string CGen::CompileUpvaluePointer(VarDef *def) { return "NULL"; } +bool CGen::BindsLocalValue(const SyntaxTreeVar *var) const { + if (!var) { + return false; + } + const auto it = var_to_def_map_.find(var); + if (it == var_to_def_map_.end()) { + return false; + } + const VarDef *def = it->second; + // 文件级 local function 登记在 chunk 作用域,但对应的 CVar 只在 init 里。 + // 其它函数引用它时必须走文件级符号,不能当成局部变量名。 + return !(def->defining_func == nullptr && def->def_node && def->def_node->Type() == SyntaxTreeType::LocalFunction); +} + bool CGen::IsCapturedInStmt(const SyntaxTreeInterface *stmt_ptr, const std::string &name) const { if (const auto it = stmt_var_to_def_.find({stmt_ptr, name}); it != stmt_var_to_def_.end()) { return it->second->is_captured; diff --git a/src/compile/c_gen.h b/src/compile/c_gen.h index 6a221a8..15664b1 100644 --- a/src/compile/c_gen.h +++ b/src/compile/c_gen.h @@ -367,6 +367,10 @@ class CGen { // 查询 name 在 stmt_ptr 对应语句节点上是否被标记为 captured [[nodiscard]] bool IsCapturedInStmt(const SyntaxTreeInterface *stmt_ptr, const std::string &name) const; + // 变量引用的是形参、局部变量或嵌套 local function。 + // 文件级 local function 不算:它的 C 变量只存在于 init,其它函数仍按文件级符号调用。 + [[nodiscard]] bool BindsLocalValue(const SyntaxTreeVar *var) const; + // 发射 heap-boxed captured var 声明: // CVar *__box_ = (CVar *)FakeluaAlloc(_S, sizeof(CVar), false); // *__box_ = ; diff --git a/src/compile/compile_common.h b/src/compile/compile_common.h index e5529d2..cf21dcb 100644 --- a/src/compile/compile_common.h +++ b/src/compile/compile_common.h @@ -292,13 +292,18 @@ struct ParseResult { // SemanticAnalysis 的输出。 // 由 SemanticAnalysis::Analyze 填充,供 CGen 使用。 struct AnalysisResult { - // 函数名 -> 最大返回值数量(-1 代表动态,例如以函数调用结尾) + // 函数名 -> 最大返回值数量(-1 代表动态:以函数调用或 ... 结尾)。 + // return ... 与 return x, ... 都是 -1,调用点会展开多返回值。 std::unordered_map function_max_returns; - // 函数名 -> 有效返回值数量:在 function_max_returns 基础上把「return f()」尾调用 - // 沿被调函数链做定点求解(含递归自调用锚点),例如 f()=return g()、g()=return 1 - // 会解析为 1。-1 表示无法静态确定(未知被调/vararg/无锚点递归环)。 - // 仅供数学参数特化的资格判定使用,不影响既有代码生成路径。 + // 函数名 -> 有效返回值数量。所有路径返回数一致时为该精确值,例如 + // return a, f() 计为 1 + f 的返回数(不是 max(表达式个数, f 的返回数)); + // return ... / return x, ...、路径之间数量不一致、未知被调、无锚点递归环为 -1。 + // 仅供数学参数特化的资格判定使用。 std::unordered_map function_effective_returns; + // 「return <单个调用>」表达式节点 -> 按词法作用域解析后的被调有效返回数。 + // 形参、局部变量、以及与文件级函数同名的嵌套 local function 记为 -1, + // 避免用文件级简单名把多返回值调用放进标量特化。 + std::unordered_map return_call_effective_returns; // 语法分析出的所有函数调用表达式节点集合,供 CGen 直接查询 std::unordered_set function_call_exps; // 语法分析出的所有函数调用到其被调用者名字的映射,供 CGen 直接查询 diff --git a/src/compile/semantic_analysis.cpp b/src/compile/semantic_analysis.cpp index 6d36637..a7a4d1c 100644 --- a/src/compile/semantic_analysis.cpp +++ b/src/compile/semantic_analysis.cpp @@ -3,6 +3,7 @@ #include "util/common.h" #include "util/exception.h" #include +#include namespace fakelua { @@ -86,161 +87,422 @@ void SemanticAnalysis::AnalyzeGlobalConstNames(const SyntaxTreeInterfacePtr &chu } } -void SemanticAnalysis::AnalyzeFunctionReturnCounts(const SyntaxTreeInterfacePtr &chunk, AnalysisResult &ar) { - DEBUG_ASSERT(chunk->Type() == SyntaxTreeType::Block); - const auto block = std::dynamic_pointer_cast(chunk); +namespace { - // 每个文件级函数的返回值明细,用于在 function_max_returns 之上做 - // 「尾调用链有效返回数」定点求解(仅供数学参数特化资格判定)。 - struct FuncReturnInfo { - // 是否存在不以尾调用/vararg 结尾的 return(含裸 return),其最大显式返回数 - bool has_concrete = false; - int concrete_max = 0; - // 存在 return ...(vararg 尾展开) - bool has_vararg_tail = false; - // 存在尾调用指向无法解析的被调(原生函数/方法/外部符号):返回数不可知 - bool has_unknown_tail = false; - // 尾调用指向的本文件级函数名 - std::vector tail_callees; +// 按词法作用域收集每个函数的返回路径,再定点求出「所有路径都相同的精确返回数」。 +// return a, f() 的返回数是 1 + f 的返回数;任一路径是 ...、未知被调,或路径之间数量不一致,则为不可知。 +struct ReturnCountGraph { + struct Binding { + bool is_func = false; + bool file_level = false; + const SyntaxTreeInterface *func = nullptr; }; - std::unordered_map infos; - - for (const auto &stmt: block->Stmts()) { + struct Tail { + int prefix = 0; + const SyntaxTreeInterface *callee = nullptr; + bool unknown = false; + }; + struct Info { + bool has_exact = false; + int exact = 0; + bool conflict = false; + bool force_taint = false; + bool dynamic_max = false; + int max_concrete = 0; + std::vector tails; + std::vector> sole_calls; + }; + struct FileFunc { std::string name; - SyntaxTreeInterfacePtr funcbody; + const SyntaxTreeInterface *node = nullptr; + int max_returns = 0; + }; + + std::function callee_name; + std::unordered_map funcs; + std::unordered_map globals; + std::unordered_map file_level_count; + std::unordered_set file_level_names; + std::vector> scopes; + std::vector file_funcs; + const SyntaxTreeInterface *current = nullptr; + size_t func_base = 0; + + static bool ReadDecl(const SyntaxTreeInterfacePtr &stmt, std::string &name, SyntaxTreeInterfacePtr &body, bool &is_local) { if (stmt->Type() == SyntaxTreeType::Function) { const auto func = std::dynamic_pointer_cast(stmt); const auto funcname_ptr = std::dynamic_pointer_cast(func->Funcname()); const auto funcnamelist = std::dynamic_pointer_cast(funcname_ptr->FuncNameList()); + if (!funcnamelist || funcnamelist->Funcnames().empty()) { + return false; + } name = funcnamelist->Funcnames()[0]; - funcbody = func->Funcbody(); - } else if (stmt->Type() == SyntaxTreeType::LocalFunction) { + body = func->Funcbody(); + is_local = false; + return true; + } + if (stmt->Type() == SyntaxTreeType::LocalFunction) { const auto func = std::dynamic_pointer_cast(stmt); name = func->Name(); - funcbody = func->Funcbody(); + body = func->Funcbody(); + is_local = true; + return true; } - if (!funcbody) { - continue; + return false; + } + + void Prescan(const std::shared_ptr &chunk) { + for (const auto &stmt: chunk->Stmts()) { + std::string name; + SyntaxTreeInterfacePtr body; + bool is_local = false; + if (!ReadDecl(stmt, name, body, is_local) || name.empty()) { + continue; + } + file_level_names.insert(name); + file_level_count[name] += 1; + if (!is_local) { + globals[name] = stmt.get(); + } } + } - const auto funcbody_ptr = std::dynamic_pointer_cast(funcbody); - const auto func_block = funcbody_ptr->Block(); + void AddExact(Info &info, int count) { + if (!info.has_exact) { + info.has_exact = true; + info.exact = count; + } else if (info.exact != count) { + info.conflict = true; + } + if (!info.dynamic_max) { + info.max_concrete = std::max(info.max_concrete, count); + } + } - std::vector returns; - CollectReturnsForBlock(func_block, returns); + [[nodiscard]] int FileCount(const std::string &name) const { + const auto it = file_level_count.find(name); + return it == file_level_count.end() ? 0 : it->second; + } - int max_returns = 0; - auto &info = infos[name]; - for (const auto &ret_node: returns) { - const auto ret = std::dynamic_pointer_cast(ret_node); - const auto el = std::dynamic_pointer_cast(ret->Explist()); - if (!el || el->Exps().empty()) { - // return count is 0 - info.has_concrete = true; + Tail ResolveCall(const SyntaxTreeInterfacePtr &call_exp, int prefix) const { + Tail tail; + tail.prefix = prefix; + const std::string name = callee_name(call_exp); + if (name.empty() || FileCount(name) > 1) { + tail.unknown = true; + return tail; + } + for (int i = static_cast(scopes.size()) - 1; i >= 0; --i) { + const auto it = scopes[static_cast(i)].find(name); + if (it == scopes[static_cast(i)].end()) { continue; } - const auto &ret_exps = el->Exps(); - int count = static_cast(ret_exps.size()); - if (IsFunctionCallExp(ret_exps.back())) { - max_returns = -1;// dynamic - if (count == 1) { - // 唯一返回表达式是尾调用:记录被调函数名用于定点求解 - const std::string callee = GetCalleeName(ret_exps.back()); - if (callee.empty()) { - info.has_unknown_tail = true; - } else { - info.tail_callees.push_back(callee); - } - } else { - // 多返回值 return 且末位是调用:显式部分贡献确定数量, - // 尾调用部分另按被调链传播 - info.has_concrete = true; - info.concrete_max = std::max(info.concrete_max, count); - const std::string callee = GetCalleeName(ret_exps.back()); - if (callee.empty()) { - info.has_unknown_tail = true; - } else { - info.tail_callees.push_back(callee); - } - } + const bool inside = current != nullptr && static_cast(i) >= func_base; + if (!it->second.is_func || (inside && !it->second.file_level && file_level_names.contains(name))) { + // 形参/局部变量,或与文件级函数同名的嵌套 local function:不能用文件级简单名。 + tail.unknown = true; + return tail; + } + tail.callee = it->second.func; + return tail; + } + const auto git = globals.find(name); + if (git == globals.end()) { + tail.unknown = true; + return tail; + } + tail.callee = git->second; + return tail; + } + + void NoteReturn(const SyntaxTreeInterfacePtr &stmt) { + if (!current || !funcs.contains(current)) { + return; + } + auto &info = funcs.at(current); + const auto ret = std::dynamic_pointer_cast(stmt); + const auto el = ret->Explist() ? std::dynamic_pointer_cast(ret->Explist()) : nullptr; + if (!el || el->Exps().empty()) { + AddExact(info, 0); + return; + } + const auto &ret_exps = el->Exps(); + const int count = static_cast(ret_exps.size()); + if (IsVarargExp(ret_exps.back())) { + // return ... 与 return x, ... 的返回数都随调用者变化。 + info.dynamic_max = true; + info.force_taint = true; + return; + } + if (IsFunctionCallExp(ret_exps.back())) { + info.dynamic_max = true; + Tail tail = ResolveCall(ret_exps.back(), count - 1); + if (tail.unknown) { + info.force_taint = true; } else { - if (IsVarargExp(ret_exps.back()) && count == 1) { - // return ...:尾 vararg 展开,返回数随调用者变化 - info.has_vararg_tail = true; - max_returns = -1; - } else if (max_returns >= 0) { - max_returns = std::max(max_returns, count); - } - info.has_concrete = true; - info.concrete_max = std::max(info.concrete_max, count); + info.tails.push_back(tail); + } + if (count == 1) { + info.sole_calls.emplace_back(ret_exps[0].get(), tail); } + return; } - ar.function_max_returns[name] = max_returns; + AddExact(info, count); } - // 尾调用被调若不是本文件级函数(外部 chunk 符号),返回数同样不可知,按 tainted 处理。 - for (auto &[name, info]: infos) { - for (const auto &callee: info.tail_callees) { - if (!infos.contains(callee)) { - info.has_unknown_tail = true; - break; + void AddParams(const SyntaxTreeInterfacePtr &funcbody) { + const auto fb = std::dynamic_pointer_cast(funcbody); + if (!fb || !fb->Parlist()) { + return; + } + const auto parlist = std::dynamic_pointer_cast(fb->Parlist()); + if (!parlist || !parlist->Namelist()) { + return; + } + const auto namelist = std::dynamic_pointer_cast(parlist->Namelist()); + if (!namelist) { + return; + } + for (const auto &pname: namelist->Names()) { + scopes.back()[pname] = Binding{}; + } + } + + void AddFunc(const SyntaxTreeInterface *node, const std::string &name, const SyntaxTreeInterfacePtr &funcbody, bool file_level) { + funcs.emplace(node, Info{}); + const SyntaxTreeInterface *saved = current; + const size_t saved_base = func_base; + current = node; + scopes.emplace_back(); + func_base = scopes.size() - 1; + AddParams(funcbody); + if (funcbody) { + const auto fb = std::dynamic_pointer_cast(funcbody); + if (fb && fb->Block()) { + WalkBlock(fb->Block()); } } + scopes.pop_back(); + current = saved; + func_base = saved_base; + + if (!file_level || name.empty()) { + return; + } + const auto &info = funcs.at(node); + file_funcs.push_back(FileFunc{name, node, info.dynamic_max ? -1 : info.max_concrete}); } - // 定点求解有效返回数,内部使用双状态: - // kTainted(-2):永久不可知(vararg 尾展开 / 未知或外部被调 / 被调链污染); - // kUnknown(-1):暂未解析(被调链尚未收敛,无 concrete 锚点的递归环停留于此); - // >=0:当前确定的最大返回数。 - // 初值:直接 tainted 标记者 → kTainted;存在确定 return(含裸 return)→ concrete_max; - // 仅有尾调用 return → kUnknown 等待被调链解析。 - // 迭代:边到 kTainted → 自身 kTainted(动态返回路径会污染所有上游); - // 边到 >=0 → 取 max;边到 kUnknown → 本轮等待。 - // 递归自调用(如 fact 有 return 1 基线)由 concrete 锚点稳定为 1。 - constexpr int kUnknown = -1; - constexpr int kTainted = -2; - for (const auto &[name, info]: infos) { - if (info.has_vararg_tail || info.has_unknown_tail) { - ar.function_effective_returns[name] = kTainted; - } else if (info.has_concrete) { - ar.function_effective_returns[name] = info.concrete_max; - } else { - ar.function_effective_returns[name] = kUnknown; - } - } - for (size_t round = 0; round < infos.size() + 1; ++round) { - bool changed = false; - for (const auto &[name, info]: infos) { - int &cur = ar.function_effective_returns[name]; - if (cur == kTainted) { - continue; + void WalkBlock(const SyntaxTreeInterfacePtr &node) { + if (!node) { + return; + } + const auto block = std::dynamic_pointer_cast(node); + DEBUG_ASSERT(block); + scopes.emplace_back(); + for (const auto &stmt: block->Stmts()) { + WalkStmt(stmt); + } + scopes.pop_back(); + } + + void BindLocalFunc(const SyntaxTreeInterfacePtr &stmt, const std::string &name) { + Binding binding; + binding.is_func = true; + binding.file_level = current == nullptr; + binding.func = stmt.get(); + scopes.back()[name] = binding; + } + + void WalkStmt(const SyntaxTreeInterfacePtr &stmt) { + switch (stmt->Type()) { + case SyntaxTreeType::Return: + NoteReturn(stmt); + break; + case SyntaxTreeType::Block: + WalkBlock(stmt); + break; + case SyntaxTreeType::If: { + const auto if_node = std::dynamic_pointer_cast(stmt); + WalkBlock(if_node->Block()); + if (const auto elseifs = if_node->ElseIfs()) { + const auto el = std::dynamic_pointer_cast(elseifs); + for (const auto &blk: el->ElseifBlocks()) { + WalkBlock(blk); + } + } + WalkBlock(if_node->ElseBlock()); + break; + } + case SyntaxTreeType::While: { + const auto while_node = std::dynamic_pointer_cast(stmt); + WalkBlock(while_node->Block()); + break; + } + case SyntaxTreeType::Repeat: { + const auto rep = std::dynamic_pointer_cast(stmt); + WalkBlock(rep->Block()); + break; + } + case SyntaxTreeType::ForLoop: { + const auto for_loop = std::dynamic_pointer_cast(stmt); + scopes.emplace_back(); + scopes.back()[for_loop->Name()] = Binding{}; + WalkBlock(for_loop->Block()); + scopes.pop_back(); + break; + } + case SyntaxTreeType::ForIn: { + const auto for_in = std::dynamic_pointer_cast(stmt); + scopes.emplace_back(); + if (const auto nl = std::dynamic_pointer_cast(for_in->Namelist())) { + for (const auto &name: nl->Names()) { + scopes.back()[name] = Binding{}; + } + } + WalkBlock(for_in->Block()); + scopes.pop_back(); + break; + } + case SyntaxTreeType::LocalVar: { + const auto lv = std::dynamic_pointer_cast(stmt); + if (const auto nl = std::dynamic_pointer_cast(lv->Namelist())) { + for (const auto &name: nl->Names()) { + scopes.back()[name] = Binding{}; + } + } + break; } - for (const auto &callee: info.tail_callees) { - const int cv = ar.function_effective_returns.at(callee); - if (cv == kTainted) { - cur = kTainted; - changed = true; + case SyntaxTreeType::LocalFunction: { + std::string name; + SyntaxTreeInterfacePtr body; + bool is_local = false; + if (!ReadDecl(stmt, name, body, is_local)) { break; } - if (cv == kUnknown) { + // 先绑定再进函数体,递归 local function 能解析到自己。 + BindLocalFunc(stmt, name); + AddFunc(stmt.get(), name, body, current == nullptr); + break; + } + case SyntaxTreeType::Function: { + std::string name; + SyntaxTreeInterfacePtr body; + bool is_local = false; + if (!ReadDecl(stmt, name, body, is_local)) { + break; + } + if (current != nullptr && !name.empty()) { + file_level_names.insert(name); + file_level_count[name] += 1; + globals[name] = stmt.get(); + } + AddFunc(stmt.get(), name, body, current == nullptr); + break; + } + case SyntaxTreeType::Assign: + case SyntaxTreeType::FunctionCall: + case SyntaxTreeType::Break: + case SyntaxTreeType::Continue: + case SyntaxTreeType::Goto: + case SyntaxTreeType::Label: + case SyntaxTreeType::Empty: + break; + default: + ThrowFakeluaException(std::format("ReturnCountGraph: unexpected statement type {}", SyntaxTreeTypeToString(stmt->Type()))); + } + } + + void Solve(AnalysisResult &ar) const { + constexpr int kUnknown = -1; + constexpr int kTainted = -2; + std::unordered_map eff; + for (const auto &[node, info]: funcs) { + if (info.force_taint || info.conflict) { + eff[node] = kTainted; + } else if (info.has_exact) { + eff[node] = info.exact; + } else if (info.tails.empty()) { + eff[node] = 0; + } else { + eff[node] = kUnknown; + } + } + for (size_t round = 0; round < funcs.size() + 1; ++round) { + bool changed = false; + for (const auto &[node, info]: funcs) { + int &cur = eff.at(node); + if (cur == kTainted) { continue; } - if (cur == kUnknown || cv > cur) { - cur = cv; - changed = true; + for (const auto &tail: info.tails) { + if (!tail.callee || !eff.contains(tail.callee)) { + cur = kTainted; + changed = true; + break; + } + const int cv = eff.at(tail.callee); + if (cv == kTainted) { + cur = kTainted; + changed = true; + break; + } + if (cv == kUnknown) { + continue; + } + const int total = tail.prefix + cv; + if (cur == kUnknown || (cur >= 0 && cur != total)) { + cur = (cur == kUnknown) ? total : kTainted; + changed = true; + if (cur == kTainted) { + break; + } + } } } + if (!changed) { + break; + } } - if (!changed) { - break; + + for (const auto &ff: file_funcs) { + int value = eff.contains(ff.node) ? eff.at(ff.node) : kTainted; + if (value < 0) { + value = -1; + } + ar.function_effective_returns[ff.name] = value; + ar.function_max_returns[ff.name] = ff.max_returns; } - } - // 对外统一输出:kTainted 与仍未解析的 kUnknown 都记为 -1(不可特化)。 - for (auto &[name, v]: ar.function_effective_returns) { - if (v < 0) { - v = -1; + for (const auto &[node, info]: funcs) { + (void) node; + for (const auto &[exp, tail]: info.sole_calls) { + int value = -1; + if (!tail.unknown && tail.callee && eff.contains(tail.callee)) { + value = eff.at(tail.callee); + if (value < 0) { + value = -1; + } + } + ar.return_call_effective_returns[exp] = value; + } } } + + void Run(const SyntaxTreeInterfacePtr &chunk, AnalysisResult &ar) { + DEBUG_ASSERT(chunk->Type() == SyntaxTreeType::Block); + const auto block = std::dynamic_pointer_cast(chunk); + Prescan(block); + WalkBlock(chunk); + Solve(ar); + } +}; + +}// namespace + +void SemanticAnalysis::AnalyzeFunctionReturnCounts(const SyntaxTreeInterfacePtr &chunk, AnalysisResult &ar) { + ReturnCountGraph graph; + graph.callee_name = [this](const SyntaxTreeInterfacePtr &exp) { return GetCalleeName(exp); }; + graph.Run(chunk, ar); } void SemanticAnalysis::CollectReturnsForBlock(const SyntaxTreeInterfacePtr &node, std::vector &returns) { diff --git a/src/compile/type_inferencer.cpp b/src/compile/type_inferencer.cpp index a08ca24..35355fc 100644 --- a/src/compile/type_inferencer.cpp +++ b/src/compile/type_inferencer.cpp @@ -1035,20 +1035,15 @@ bool TypeInferencer::IsEligibleForMathSpec(const SyntaxTreeInterfacePtr &block_n } // 唯一返回值是尾位置 vararg 展开(return ...):实际元数随调用者变化,排除。 // 唯一返回值是尾位置函数调用(return f()):特化体经 CompileNumericExp - // 直发被调的标量特化,故仅当 f 的有效返回数静态可知恰好为 1 时才安全 - // (见 SemanticAnalysis 的 function_effective_returns 定点求解, - // 可穿透多层尾调用与带单值基线的递归)。 + // 直发被调的标量特化,故仅当 f 的有效返回数静态可知恰好为 1 时才安全。 + // 被调按词法作用域解析:形参/局部变量,或与文件级函数同名的嵌套函数,不算 1。 const auto &only = ret_exps[0]; if (IsVarargExp(only)) { return false; } if (IsFunctionCallExp(only)) { - const auto callee_it = ar.callee_names.find(only.get()); - const std::string callee = (callee_it != ar.callee_names.end()) ? callee_it->second : ""; - // 用跨函数定点求解后的有效返回数:return f() 仅当 f 链静态可知恰好 - // 返回 1 个值时才安全(递归自调用带单值基线也算 1)。 - const auto eff_it = ar.function_effective_returns.find(callee); - if (eff_it == ar.function_effective_returns.end() || eff_it->second != 1) { + const auto eff_it = ar.return_call_effective_returns.find(only.get()); + if (eff_it == ar.return_call_effective_returns.end() || eff_it->second != 1) { return false; } } diff --git a/src/compile/type_inferencer.h b/src/compile/type_inferencer.h index fbd0d53..d535884 100644 --- a/src/compile/type_inferencer.h +++ b/src/compile/type_inferencer.h @@ -136,8 +136,8 @@ class TypeInferencer { // 数学参数特化资格检查:函数的每条 return 都必须与标量/CVar 特化调用约定兼容。 // 多值返回(return a, b, ...)、尾位置 vararg 展开(return ...)、以及尾位置 - // 调用无法静态确定为单返回值函数(return f() 且 f 可能返回多个值)的函数不参与 - // 特化——这些形态只有通用 CVar 变体才能正确处理。 + // 调用无法静态确定为恰好 1 个返回值(return f(),含被局部绑定遮蔽的同名函数) + // 的函数不参与特化——这些形态只有通用 CVar 变体才能正确处理。 // 不递归进入嵌套函数定义;裸 return(0 值)会使特化返回类型退化为 CVar,是安全的。 [[nodiscard]] bool IsEligibleForMathSpec(const SyntaxTreeInterfacePtr &block_node, const AnalysisResult &ar) const; diff --git a/src/interp/codegen.cpp b/src/interp/codegen.cpp index b67d645..086ed6f 100644 --- a/src/interp/codegen.cpp +++ b/src/interp/codegen.cpp @@ -1602,7 +1602,9 @@ int InterpCodegen::CompileFunctioncall(const SyntaxTreeInterfacePtr &functioncal } } - if (!func_name.empty() && (!is_local_callee || file_level_func || named_protos_.contains(func_name))) { + // 形参、局部变量、嵌套 local function 即使与文件级函数同名,也必须调用这个绑定。 + // named_protos_ 只登记文件级函数,不能拿来覆盖局部值。 + if (!func_name.empty() && (!is_local_callee || file_level_func)) { return place_callname(func_name, arg_regs); } diff --git a/test/lua/infer/test_jitbug_return_exact.lua b/test/lua/infer/test_jitbug_return_exact.lua new file mode 100644 index 0000000..5b40612 --- /dev/null +++ b/test/lua/infer/test_jitbug_return_exact.lua @@ -0,0 +1,32 @@ +-- 回归:有效返回数是每条路径的精确值,不是 max()。 +-- exactly_one 的 return x + 1, zero() 在 zero 返回 0 个值时实际只有 1 个值,caller_one 仍可特化。 +-- mixed 一条路径返回 1 个值、另一条返回 0 个值,caller_mixed 不得特化。 + +local function zero() + return +end + +local function exactly_one(x) + return x + 1, zero() +end + +local function mixed(x) + if x > 0 then + return x + end + return zero() +end + +function caller_one(n) + local y = n + 1 - 1 + return exactly_one(y) +end + +function caller_mixed(n) + local y = n + 1 - 1 + return mixed(y) +end + +function test(n) + return caller_one(n) + caller_mixed(n) +end diff --git a/test/lua/infer/test_jitbug_return_scope.lua b/test/lua/infer/test_jitbug_return_scope.lua new file mode 100644 index 0000000..f239c88 --- /dev/null +++ b/test/lua/infer/test_jitbug_return_scope.lua @@ -0,0 +1,39 @@ +-- 回归:return g() 按词法作用域解析被调,不能用文件级同名函数的返回数。 +-- shadowed 里的 local g 返回两个值,即使文件级 g 只返回一个值,也不得特化。 +-- 形参 g 同样遮蔽文件级 g。 +-- 没有同名文件级函数的嵌套 inner 精确返回 1 个值时,外层仍可特化。 + +function g(x) + return x + 1 +end + +function shadowed(n) + local function g(x) + local y = x + 0 + return y, y + 1 + end + local z = n + 1 - 1 + return g(z) +end + +function use_nested(n) + local function inner(x) + return x + 1 + end + local y = n + 1 - 1 + return inner(y) +end + +function apply(n, g) + local y = n + 1 - 1 + return g(y) +end + +function test(n) + local a, b = shadowed(n) + local c = use_nested(n) + local d, e = apply(n, function(x) + return x, x + 10 + end) + return a + b + c + d + e +end diff --git a/test/lua/infer/test_jitbug_vararg_passthrough.lua b/test/lua/infer/test_jitbug_vararg_passthrough.lua new file mode 100644 index 0000000..e1a12f0 --- /dev/null +++ b/test/lua/infer/test_jitbug_vararg_passthrough.lua @@ -0,0 +1,11 @@ +-- 回归:函数体只有 return ... 时,调用点必须展开全部返回值。 +-- 旧逻辑把这种函数的 max_returns 记成 1,local a, b = passthrough(10, 20) 只会留下第一个值。 + +local function passthrough(...) + return ... +end + +function test(n) + local a, b = passthrough(n, n + 1) + return a + b +end diff --git a/test/test_infer.cpp b/test/test_infer.cpp index 1b436a1..9394625 100644 --- a/test/test_infer.cpp +++ b/test/test_infer.cpp @@ -4407,7 +4407,7 @@ TEST(infer, test_infer_cvar_to_int) { // 多返回值 / vararg 展开 / 尾位置透传多返回值的函数即使命中数学参数,也不得生成 // 标量特化(特化函数 return FlMakeMulti 会产生非法 C)。 -// test(10)=115, test(10.5)=120。 +// test(10)=235,test(10.5)=240,其中 factacc(5,1)=120。 TEST(infer, test_jitbug_multi_return_no_spec) { const auto code = InferGetCCode("./infer/test_jitbug_multi_return_spec.lua"); // 通用 CVar 变体必须存在。 @@ -4440,8 +4440,11 @@ TEST(infer, test_jitbug_multi_return_no_spec) { // 表达式位置;generic-for 超过 3 个表达式时丢弃位同理。test()=60。 TEST(infer, test_jitbug_tail_expand_stmt) { const auto code = InferGetCCode("./infer/test_jitbug_tail_expand_stmt.lua"); - // generic-for 第 4 个(丢弃)表达式以独立语句形式求值。 + // 语句宏必须先写成独立语句,不能拼进赋值或 (void)() 的表达式位置。 + ASSERT_NE(code.find("flua_call_res_"), std::string::npos); ASSERT_NE(code.find("(void)("), std::string::npos); + ASSERT_EQ(code.find("= do"), std::string::npos); + ASSERT_EQ(code.find("(void)(do"), std::string::npos); InferRunHelper([](State *s, JITType type, bool debug_mode) { CompileFile(s, "./infer/test_jitbug_tail_expand_stmt.lua", {.debug_mode = debug_mode}); @@ -4451,6 +4454,52 @@ TEST(infer, test_jitbug_tail_expand_stmt) { }); } +// 只有 return ... 的函数,调用点要展开全部返回值。test(10)=21,test(10.5)=22。 +TEST(infer, test_jitbug_vararg_passthrough) { + InferRunHelper([](State *s, JITType type, bool debug_mode) { + CompileFile(s, "./infer/test_jitbug_vararg_passthrough.lua", {.debug_mode = debug_mode}); + int64_t ri = 0; + Call(s, type, "test", ri, 10); + ASSERT_EQ(ri, 21); + double rf = 0; + Call(s, type, "test", rf, 10.5); + ASSERT_NEAR(rf, 22.0, 0.001); + }); +} + +// return a, f() 在 f 返回 0 个值时精确返回数是 1;路径 1 与路径 0 不一致时不可特化。 +// test(10)=21(caller_one=11,caller_mixed=10)。 +TEST(infer, test_jitbug_return_exact) { + const auto code = InferGetCCode("./infer/test_jitbug_return_exact.lua"); + ASSERT_NE(code.find("flua_fn_caller_one_0("), std::string::npos); + ASSERT_EQ(code.find("flua_fn_caller_mixed_0("), std::string::npos); + ASSERT_EQ(code.find("flua_fn_exactly_one_0("), std::string::npos); + + InferRunHelper([](State *s, JITType type, bool debug_mode) { + CompileFile(s, "./infer/test_jitbug_return_exact.lua", {.debug_mode = debug_mode}); + int64_t ret = 0; + Call(s, type, "test", ret, 10); + ASSERT_EQ(ret, 21); + }); +} + +// 嵌套 local function / 形参遮蔽文件级同名函数时,不得按文件级返回数做标量特化。 +// 无同名文件级函数、且嵌套函数精确返回 1 个值时,外层仍特化。test(10)=62。 +TEST(infer, test_jitbug_return_scope) { + const auto code = InferGetCCode("./infer/test_jitbug_return_scope.lua"); + ASSERT_EQ(code.find("flua_fn_shadowed_0("), std::string::npos); + ASSERT_EQ(code.find("flua_fn_apply_0("), std::string::npos); + ASSERT_NE(code.find("flua_fn_use_nested_0("), std::string::npos); + ASSERT_NE(code.find("flua_fn_g_0("), std::string::npos); + + InferRunHelper([](State *s, JITType type, bool debug_mode) { + CompileFile(s, "./infer/test_jitbug_return_scope.lua", {.debug_mode = debug_mode}); + int64_t ret = 0; + Call(s, type, "test", ret, 10); + ASSERT_EQ(ret, 62); + }); +} + // 顶层函数名为 C 保留入口名(main)或 libc 符号(sin)时,C 符号统一加 flua_fn_ // 前缀;Lua 侧仍按原名调用。 TEST(infer, test_jitbug_reserved_func_name) {