diff --git a/src/compile/c_gen.cpp b/src/compile/c_gen.cpp index 5d6dfed..83eb38c 100644 --- a/src/compile/c_gen.cpp +++ b/src/compile/c_gen.cpp @@ -488,6 +488,7 @@ void CGen::GenerateGlobal(const SyntaxTreeInterfacePtr &chunk) { InferredType global_type = ir().global_const_vars.at(name); const auto cname = CIdent(name); const auto exp_node = std::dynamic_pointer_cast(exp); + // 文件级数值 local 是常量。再赋值在类型推断阶段就已经报错,这里一律 static const。 if (global_type == T_INT) { if (!exp_node || exp_node->GetExpKind() == ExpKind::kNil) { Out() << "static const int64_t " << cname << " = 0;\n"; diff --git a/src/compile/c_runtime_header.h b/src/compile/c_runtime_header.h index 9651bf9..a10c918 100644 --- a/src/compile/c_runtime_header.h +++ b/src/compile/c_runtime_header.h @@ -82,6 +82,10 @@ struct VarTable { 需要重算——因此把 VarTable 整体清零的分配路径天然落在安全的重算分支上。 */ uint32_t seq_len_valid_; int64_t seq_len_; + /* spec 块字节数。0 = 无 spec。与 src/var/var_table.h 同布局。 */ + uint32_t spec_bytes; + /* >0:spec 是连续 CVar(表特化)。0:不透明块,按 spec_bytes 整体搬迁。 */ + uint32_t spec_cvars; }; typedef struct State State; @@ -335,6 +339,10 @@ extern CVar FakeluaAllocMultiCVar(State *s, int count); extern void FakeluaSetMultiCVarElement(CVar *multi, int idx, CVar val); extern CVar FlEvalLoadClosure(State *state, VarClosure *cl, int arg_num, const CVar *args); extern CVar FakeluaInterpCall(State *state, VarClosure *cl, int arg_num, const CVar *args); +// 当前执行引擎上下文(供 JIT C 代码调用闭包前后维护):push 切到 jit 并返回上一个引擎, +// pop 恢复。native 对象方法等在闭包内执行时据此登记派回本引擎的异步回调。 +extern int FakeluaJitContextPush(State *s, int jit); +extern void FakeluaJitContextPop(State *s, int prev); #ifdef __cplusplus } #endif @@ -519,8 +527,13 @@ static inline CVar FlCallClosure(State *state, CVar cl_var, int arg_num, ...) { #define FCARG_31 FCARG_30, arg_arr[30] #define FCARG_32 FCARG_31, arg_arr[31] -#define FCCASE(N) \ - case N: return ((CVar (*)(VarClosure * FCCVAR_##N))(addr))(cl FCARG_##N); +#define FCCASE(N) \ + case N: { \ + int __fl_prev_jit = FakeluaJitContextPush(state, FAKELUA_JIT_TYPE); \ + CVar __fl_ret = ((CVar (*)(VarClosure * FCCVAR_##N))(addr))(cl FCARG_##N); \ + FakeluaJitContextPop(state, __fl_prev_jit); \ + return __fl_ret; \ + } switch (expected_arg_count) { FCCASE(0) FCCASE(1) FCCASE(2) FCCASE(3) FCCASE(4) FCCASE(5) @@ -647,6 +660,8 @@ static inline uint32_t FlHashString(const char *str, int len) { __t->spec_count = 0; \ __t->seq_len_valid_ = 1; \ __t->seq_len_ = 0; \ + __t->spec_bytes = 0; \ + __t->spec_cvars = 0; \ assert(sizeof(__t->quick_data_) == 8 * sizeof(VarEntry)); \ { int __i; for (__i = 0; __i < 8; ++__i) { \ __t->quick_data_[__i].key.type_ = VAR_NIL; \ @@ -667,6 +682,9 @@ static inline uint32_t FlHashString(const char *str, int len) { (v).data_.t->spec_keys = (CVar *)FakeluaAlloc(_S, sizeof(CVar) * (field_count), !__fakelua_init_flag__); \ (v).data_.t->spec_vals = (CVar *)FakeluaAlloc(_S, sizeof(CVar) * (field_count), !__fakelua_init_flag__); \ (v).data_.t->spec_count = (field_count); \ + (v).data_.t->spec_bytes = (uint32_t)sizeof(SpecType); \ + (v).data_.t->spec_cvars = (uint32_t)(field_count); \ + assert(sizeof(SpecType) == (field_count) * sizeof(CVar)); \ } while(0) #define FL_SPEC(SpecType, v, field) (((SpecType *)(v).data_.t->spec)->field) diff --git a/src/compile/type_inferencer.cpp b/src/compile/type_inferencer.cpp index 84a2b65..b6180c2 100644 --- a/src/compile/type_inferencer.cpp +++ b/src/compile/type_inferencer.cpp @@ -32,6 +32,24 @@ std::string SpecStringFieldCName(const std::string &key) { return out; } +// 预处理生成的文件级初始化函数。它的函数体里对 local x = func() 降级出来的 +// 绑定做一次赋值,不能当成用户再赋值。 +bool IsFakeluaInitFunction(const std::shared_ptr &func) { + if (!func) { + return false; + } + const auto fn = std::dynamic_pointer_cast(func->Funcname()); + if (!fn) { + return false; + } + const auto fnl = std::dynamic_pointer_cast(fn->FuncNameList()); + if (!fnl) { + return false; + } + const auto &names = fnl->Funcnames(); + return names.size() == 1 && names[0] == kInitFunctionName; +} + void UniquifySpecFieldCNames(std::vector &fields) { std::unordered_set used; for (auto &f: fields) { @@ -383,7 +401,9 @@ InferredType TypeInferencer::TypeEnvironment::MergeType(const InferredType old_t InferResult TypeInferencer::InferTypes(const ParseResult &pr, const CompileConfig &cfg) { LOG_DEBUG(s_, "engine", "InferTypes: start for {}", pr.file_name); + file_name_ = pr.file_name; file_level_types_.clear(); + file_level_init_exps_.clear(); InferResult ir; EvalTypeSnapshot current_map; TypeEnvironment env; @@ -467,7 +487,12 @@ InferredType TypeInferencer::InferNode(const SyntaxTreeInterfacePtr &node, Trave } case SyntaxTreeType::Function: { const auto func = std::dynamic_pointer_cast(node); + const bool prev_init = in_init_function_; + if (IsFakeluaInitFunction(func)) { + in_init_function_ = true; + } InferNode(func->Funcbody(), tctx); + in_init_function_ = prev_init; return RecordType(current_map, node.get(), T_UNKNOWN); } case SyntaxTreeType::LocalFunction: { @@ -610,6 +635,11 @@ InferredType TypeInferencer::InferLocalVar(const std::shared_ptr if (const auto *init = tctx.env.LookupInitNode(name)) { tctx.var_define_nodes[var.get()] = init; + // 文件级数值字面量是 static const,用户函数里再赋值直接报编译期错误。 + // __fakelua_init 除外:local x = func() 不能做 C 静态初值,预处理把它改成 + // local x = nil,并只在 __fakelua_init 里写一次 x = func()。那是初始化,不是再赋值。 + if (!in_init_function_ && file_level_init_exps_.contains(init)) { + ThrowFakeluaException(std::format("cannot reassign file-level constant '{}' at {}", name, SyntaxTreeLocationStr(file_name_, assign))); + } } current_map[var.get()] = current; @@ -1092,7 +1128,7 @@ EvalTypeSnapshot TypeInferencer::RunTrialInference(const SyntaxTreeInterfacePtr // 运行函数体类型推断(不新开作用域,参数已在当前作用域中定义)。 // Trial 推断不消费 shadow 信息(shadow 只与 AST 结构相关,主推断已覆盖), - // 但 TraversalContext 需要一个引用,因此用一个丢弃式的 set 占位。 + // 但 TraversalContext 需要一个引用,因此用丢弃式的容器占位。 std::set> dummy_shadowed_decls; TraversalContext tctx{current_map, env, &ctx, var_define_nodes, dummy_shadowed_decls}; InferBlock(std::dynamic_pointer_cast(func_block), false, tctx); diff --git a/src/compile/type_inferencer.h b/src/compile/type_inferencer.h index c190263..6f2a2c4 100644 --- a/src/compile/type_inferencer.h +++ b/src/compile/type_inferencer.h @@ -277,7 +277,13 @@ class TypeInferencer { private: State *s_ = nullptr; + std::string file_name_; + // 正在推断 __fakelua_init。它里面的 x = func() 是文件级复杂初值的唯一赋值点,不是再赋值。 + bool in_init_function_ = false; std::unordered_map file_level_types_; + // 文件级数值字面量 local 的 initializer 节点。InferAssign 用它判断赋值目标是不是 + // 这条常量绑定(对函数内同名遮蔽免疫)。__fakelua_init 里的赋值不查这张表。 + std::unordered_set file_level_init_exps_; // 不动点迭代轮次上限(实际通常 2 轮即可收敛)。 static constexpr int kMaxSpecIterations = 16; diff --git a/src/jit/vm.cpp b/src/jit/vm.cpp index 781e824..a621781 100644 --- a/src/jit/vm.cpp +++ b/src/jit/vm.cpp @@ -40,6 +40,18 @@ extern "C" __attribute__((used)) void FakeluaSetMultiCVarElement(CVar *multi, in inter::SetMultiCVarElement(*multi, idx, val); } +// JIT 生成的 C 代码(TCC 是 C 编译器,用不了 C++ 的 State::JitContextScope)调用闭包 +// 前后经这对接口维护"当前执行引擎"。返回旧值供 Pop 恢复,单线程无需原子操作。 +extern "C" int FakeluaJitContextPush(State *s, int jit) { + const int prev = static_cast(s->CurrentJit()); + s->SetCurrentJit(static_cast(jit)); + return prev; +} + +extern "C" void FakeluaJitContextPop(State *s, int prev) { + s->SetCurrentJit(static_cast(prev)); +} + static CVar CallByNameImpl(State *state, int jit_type, const char *name, int arg_num, const CVar *raw_arg_arr); extern "C" __attribute__((used)) CVar FakeluaCallByName(State *state, int jit_type, const char *name, int arg_num, ...) { @@ -70,6 +82,9 @@ extern "C" __attribute__((used)) CVar FakeluaCallByName(State *state, int jit_ty } static CVar CallByNameImpl(State *state, int jit_type, const char *name, int arg_num, const CVar *raw_arg_arr) { + // 记录"当前 Lua 调用方引擎":native 函数内登记的异步回调据此派回同一引擎。 + // 作用域覆盖整个函数(含嵌套的 Lua 回调),退出时恢复外层引擎。 + State::JitContextScope jit_scope(state, static_cast(jit_type)); // 查找函数:优先 JIT,其次 C++ 原生 // 用 string_view 查表,避免每次调用都堆分配 std::string const std::string_view func_name(name); diff --git a/src/native/README.md b/src/native/README.md index 62b9b59..4fcc883 100644 --- a/src/native/README.md +++ b/src/native/README.md @@ -316,7 +316,7 @@ On Windows, `io.open` / `loadfile` / `dofile` / `os.getenv` / `os.tmpname` / Boo | `net.udp_server(config)` | Bind a UDP socket (`ip`, `port`; `port=0` is ephemeral). `send(data, ip, port)` | | `net.udp_client(config)` | Connected UDP client (`ip`/`port` are the peer). `send(data)` | | `obj:dispatch(func_name)` | Register Lua callback function name | -| `obj:send(connid, data)` | Send data (server: specify connid; client: omit) | +| `obj:send(connid, data)` | Send data (server: specify connid; client: omit). From a callback context the send is queued and pumped by the tick after dispatch finishes; failures are logged as WARN | | `obj:close()` | Close connection/server | | `obj:close_connection(connid)` | Close single connection (server only) | | `obj:get_events()` | Event history | @@ -402,6 +402,30 @@ Timers go first so callbacks that come due can be picked up by the I/O dispatch them. Within MySQL, pools are driven before connections so heartbeat and reconnect have run before a script acquires a connection this round. +**Pump order and callback timing (important):** the order is fixed: timer → net → http → mysql → +redis. Each module dispatches its Lua callbacks synchronously inside the tick, therefore: + +- **IO issued from inside a callback produces results at the earliest in the next tick.** A mysql + query issued from a net callback cannot have its result dispatched until the next frame; do not + write code that issues a query from a callback and expects the result synchronously. +- A second `conn:query` on the same connection while one is in flight is **queued** (no error) and + starts automatically once the connection is Ready again; at most one queued query starts per tick. + +**Callback context (C++ → Lua) — allowed / forbidden operations:** + +net's on_event and the mysql/http/redis result callbacks all run in a restricted dispatch context. +Each entry wraps the call in `IoContext::DispatchScope`, which increments `InDispatch()` +(net `DrainEventsWith`, mysql connect/result dispatch, http `CallNamed`, redis result dispatch). +The flag is an integer depth on the State, so Linux, macOS and Windows behave the same way. +`send` queues while a dispatch is active and the tick pumps the queue after the scope ends. + +| Operation | Supported | Notes | +|------|---------|------| +| Mutating fields of a runtime-created table | ✅ | Attaching a table created inside the callback to a long-lived table is a reliable pattern | +| IO such as `send` / mysql `query` | ✅ | sends are pumped by the tick after dispatch finishes; queries queue asynchronously | +| Assigning **fields** of a file-level `local` table | ❌ | file-level tables get CONST_FLAG after init; mutating content throws "attempt to modify a const table" | +| Rebinding a file-level numeric `local` (`x = ...`) | ❌ | a file-level numeric local with a literal initializer is a constant; assigning it later is a compile error that names the source line. `local x = func()` cannot be a C static initializer: the declaration is lowered to `local x = nil` and `__fakelua_init` assigns it once when the shared library is loaded | + --- ## Event @@ -541,9 +565,18 @@ Boost.Container-backed structures stored on a NativeObject (C++ heap). They surv | Function | Args | Description | |----------|------|-------------| -| `json.encode(value)` | 1 | Lua value → JSON string; consecutive int keys 1..N → array; floats use `%.17g` | +| `json.encode(value)` | 1 | Lua value → JSON string; consecutive int keys 1..N → array; **an empty table encodes as `[]`**; floats use `%.17g` | +| `json.encode_array(value)` | 1 | Strict array encoding: the top level must be an array-like table (empty table → `[]`; non-empty requires consecutive int keys from 1), otherwise throws | | `json.decode(str)` | 1 | JSON string → Lua value; `null` → `nil` | +> **Empty-table ambiguity:** Lua cannot distinguish an "empty array" from an "empty object". +> `json.encode` applies the pure-array heuristic and encodes empty tables as `[]`; use +> `json.encode_array` when the client protocol must be strict (array-shaped or error). +> +> **Breaking change:** `json.encode({})` used to produce `{}` and now produces `[]`. Callers that +> need an empty object should encode a table with a string key, or accept both shapes on the client. +> Use `json.encode_array` when the value must be an array. + --- ## MySQL @@ -556,20 +589,66 @@ Boost.Container-backed structures stored on a NativeObject (C++ heap). They surv `ssl`: omit/`false`/`"disable"` keeps plaintext (default). `true`/`"require"` demands TLS. `"enable"` uses TLS when the server offers it. Optional `ssl_ca` PEM enables certificate verification. +**Callback arguments:** every async `cb` accepts only an in-package global function name +(string, e.g. `"on_result"` or `"DB.on_result"`); **inline closures are not supported**. FakeLua +has no GC: runtime values live on the temporary arena, and a top-level `Call` resets that arena, so +a raw closure pointer cannot be held safely across ticks. To pass context, append **bound +arguments** (pure data) after the function name: `conn:query(sql, "on_result", tag, ctx_table)`. +Bound arguments are serialized into a self-owned byte blob at registration time (same wire format +as `serialize.encode`), survive any number of `Reset`s, and are deserialized in the dispatch frame +and appended after the fixed callback arguments. The blob is released when the query completes or +the connection is destroyed, so memory usage tracks in-flight queries only; the const arena is not +used. + +- Bound arguments may be nil/boolean/number/string and nested tables thereof (no cycles); + non-serializable values (closures, native objects, ...) anywhere in the nesting raise a + "bad argument" error — fields are never silently dropped. +- Bound arguments are a **registration-time snapshot** (copy by value): mutating the original + table after registration is not visible to the callback. Put mutable cross-frame state on the + connection object itself (e.g. `conn.query_done`). +- The callback runs in the **same engine that issued the call** (TCC/GCC/interpreter each close + the loop themselves); dispatch is not pinned to TCC. +- `pool:with(fn)` invokes fn **synchronously in the same frame** and never stores it across + ticks, so inline closures remain supported there. + +**Callback contract:** once `conn:query` is called its callback fires **exactly once** — when the +connection is not ready (handshaking / reconnecting / a previous query still in flight) the query is +queued and starts automatically once the connection is usable; if the connection is closed or in a +terminal error state, the callback receives an error string. All callbacks are driven by +`runtime.tick()`. + | Function/Method | Description | |----------|-------------| -| `mysql.connect(config, cb)` | Async connect; callback `function cb(err, conn)` | +| `mysql.connect(config, cb, ...)` | Async connect; callback `cb(conn, err, success, ...)`; `...` are bound args appended after the fixed args | | `mysql_pool.create(config)` | Create connection pool | -| `conn:query(sql, cb)` | Async query; callback `function cb(err, result)` | -| `conn:stmt_prepare(sql, cb)` | Prepare statement | -| `conn:stmt_execute(id, params, cb)` | Execute prepared statement | +| `conn:query(sql, cb, ...)` | Async query; callback `cb(conn, err, result, ...)`; `...` are bound args | +| `conn:stmt_prepare(sql, cb, ...)` | Prepare statement; callback `cb(conn, err, stmt_id, ...)` | +| `conn:stmt_execute(id, params, cb, ...)` | Execute prepared statement; bound args start at the 4th argument | | `conn:stmt_close(id)` | Close prepared statement | | `conn:close()` | Close connection | -| `pool:acquire()` | Get connection from pool | +| `pool:acquire()` | Get connection from pool (nil when none available) | | `pool:release(conn)` | Return connection to pool | +| `pool:with(fn)` | Lease-style usage: acquires a connection and passes it to `fn(conn)`, returning it automatically when fn returns or throws; returns fn's return value. Prevents the "forgot to release / callback never fired so the connection is never returned" leak pattern | | `pool:close()` | Close pool | | `pool:stats()` | Returns `{total, healthy}` | +**Result table layout (SELECT):** `result[1] = true`; `result[2]` is the column info table (each item +`{[1]=column name, [2]=MySQL type code}`); `result[3]` is the rows table, each row indexed by column +position (1-based). + +**Result table layout (INSERT/UPDATE/DELETE/DDL):** `result[1] = false`; `result[4] = affected_rows`; +`result[5] = last_insert_id`; `result[6] = info`. + +**Row value types:** converted by column type — integer columns (TINY/SHORT/LONG/LONGLONG/INT24/YEAR) +return numbers, float/decimal columns (FLOAT/DOUBLE/DECIMAL/NEWDECIMAL) return numbers, all other +columns (strings/dates/BLOBs) return strings, NULL returns nil. Boolean-style `TINYINT(1)` returns a +number (0/1); convert as needed. + +> **DECIMAL / NEWDECIMAL precision:** these columns become IEEE 754 doubles, about 15–16 significant +> decimal digits. A `DECIMAL(18,6)` (and similar high-precision definitions) can lose the last digit +> or two, so it is a poor sole representation of money. Cast to a string in SQL +> (`CAST(price AS CHAR)`) and keep it as a string in Lua when the value must be exact. + --- ## Redis diff --git a/src/native/README.zh.md b/src/native/README.zh.md index 5cfe369..f288c8b 100644 --- a/src/native/README.zh.md +++ b/src/native/README.zh.md @@ -312,7 +312,7 @@ Windows 上 `io.open` / `loadfile` / `dofile` / `os.getenv` / `os.tmpname` 以 | `net.udp_server(config)` | 绑定 UDP(`ip`、`port`;`port=0` 为临时端口)。`send(data, ip, port)` | | `net.udp_client(config)` | 已连接 UDP 客户端(`ip`/`port` 为对端)。`send(data)` | | `obj:dispatch(func_name)` | 注册 Lua 回调函数名 | -| `obj:send(connid, data)` | 发送数据(服务端需指定 connid;客户端省略) | +| `obj:send(connid, data)` | 发送数据(服务端需指定 connid;客户端省略)。在回调上下文中调用时先入队,由本轮 tick 派发完后统一泵出;发送失败记 WARN 日志 | | `obj:close()` | 关闭连接/服务端 | | `obj:close_connection(connid)` | 关闭单个连接(仅服务端) | | `obj:get_events()` | 事件历史 | @@ -397,6 +397,28 @@ per-State 对象列表,所以 server/client、连接和连接池都不再有 定时器放最前,让本轮到期的回调能赶上后面的 IO 派发。mysql 内部先池后连接,这样脚本这一轮 取连接之前,心跳和重连已经推进过了。 +**泵序与回调时序(重要):** 顺序固定为 timer → net → http → mysql → redis。各模块的 Lua 回调 +都在 tick 内同步执行,因此: + +- **回调里发起的 IO 最早在下一轮 tick 才有结果**。net 回调里发起的 mysql query,最早要等下一帧 + 才可能派发结果;不要写"回调里发查询并同步等结果"的代码。 +- mysql 同一连接上飞行中再次 `conn:query` 会**排队**(而不是报错),连接回到 Ready 后由后续 + tick 自动启动;每轮 tick 至多启动一条排队的 query。 + +**回调上下文(C++ → Lua)允许/禁止的操作:** + +net 的 on_event、mysql/http/redis 的结果回调都运行在受限的派发上下文中。派发入口用 +`IoContext::DispatchScope` 把 `InDispatch()` 加一(net 的 DrainEventsWith、mysql 的结果/连接回调、 +http 的 `CallNamed`、redis 的结果回调)。这是 State 上的一个整数深度,Linux / macOS / Windows +行为相同。`send` 看到仍在派发中就入队,本轮 tick 在 scope 结束后统一泵出。规则如下: + +| 操作 | 是否支持 | 说明 | +|------|---------|------| +| 改"运行时创建的 table"的字段 | ✅ | 回调里新建的 table 挂到长生命周期 table 上是可靠模式 | +| `send` / mysql `query` 等 IO | ✅ | send 会在派发结束后由 tick 统一泵出;query 异步排队 | +| 给文件级 `local` 表的**字段**赋值 | ❌ | 文件级表初始化后打 CONST_FLAG,改内容会抛 "attempt to modify a const table" | +| 重绑定文件级数值 `local`(`x = ...`) | ❌ | 字面量初值的文件级数值 local 是常量,函数里再赋值是带行号的编译期错误。`local x = func()` 不能做 C 静态初值:声明先写成 `local x = nil`,加载 so 时由 `__fakelua_init` 赋值一次 | + --- ## Event(事件系统) @@ -536,9 +558,16 @@ PCG-32 算法:64-bit 状态,32-bit 输出,周期 2^64。每个 `random.new | 函数 | 参数 | 说明 | |------|------|------| -| `json.encode(value)` | 1 | Lua 值 → JSON 字符串;连续整数键 1..N → 数组;浮点数用 `%.17g` | +| `json.encode(value)` | 1 | Lua 值 → JSON 字符串;连续整数键 1..N → 数组,**空 table 编码为 `[]`**;浮点数用 `%.17g` | +| `json.encode_array(value)` | 1 | 严格数组编码:顶层必须是数组形 table(空 table → `[]`;非空要求 key 为从 1 起的连续整数),否则抛错 | | `json.decode(str)` | 1 | JSON 字符串 → Lua 值;`null` → `nil` | +> **空 table 的歧义:** Lua 无法区分"空数组"与"空对象"。`json.encode` 按纯数组启发式把空 +> table 编码为 `[]`;需要与客户端协议严格对齐时用 `json.encode_array`(要求数组形,否则报错)。 +> +> **破坏性变更:** `json.encode({})` 从产出 `{}` 改为产出 `[]`。依赖空对象的调用方要改成带字符串键的表, +> 或在客户端同时接受空数组和空对象。明确要数组时用 `json.encode_array`。 + --- ## YAML @@ -612,20 +641,69 @@ PCG-32 算法:64-bit 状态,32-bit 输出,周期 2^64。每个 `random.new `ssl`:省略/`false`/`"disable"` 保持明文(默认)。`true`/`"require"` 强制 TLS。`"enable"` 在服务器支持时使用 TLS。可选 `ssl_ca` PEM 会校验证书。 +**回调参数:** 所有异步回调 `cb` 只接受**包内全局函数名**(字符串,如 `"on_result"` 或 +`"DB.on_result"`),**不支持内联闭包**。fakelua 没有 GC:运行期值在临时 arena 上,顶层 +`Call` 会 `Reset` 这块内存,无法安全地跨 tick 持有闭包裸指针。需要上下文时在函数名后面追加 +**绑定参数**(纯数据):`conn:query(sql, "on_result", tag, ctx_table)`。绑定参数在登记当下被 +序列化成独立字节串暂存(与 `serialize.encode` 同一 wire 格式),跨任意次 `Reset` 不丢; +结果回来时在派发帧内反序列化,追加在固定参数之后调用。字节串随 query 完成/连接销毁释放, +内存只与在途 query 数相关,不占用 const arena。 + +- 绑定参数支持 nil/boolean/number/string 及其嵌套表(不可有环);闭包、native 对象等 + 不可序列化类型出现在任意嵌套位置都会直接报 "bad argument",不会静默丢字段。 +- 绑定参数是**登记时刻的值快照**(按值复制),登记之后调用方再改原表,回调看不到; + 需要跨帧共享的可变状态请挂在连接对象字段上(如 `conn.query_done`)。 +- 回调在**发起调用的同一引擎**里执行(TCC/GCC/解释器各自闭环),不固定派发到 TCC。 +- `pool:with(fn)` 的 fn 是**同帧同步调用**、不跨 tick 暂存,因此仍然支持内联闭包。 + +**回调契约:** `conn:query` 一旦被调用,其回调**恰好被调用一次**——连接未就绪(握手中/重连中/ +上一条 query 仍在飞行)时 query 排队,连接可用后自动执行;连接已关闭或进入错误终态时,回调会 +收到错误字符串。所有回调由 `runtime.tick()` 驱动。 + | 函数/方法 | 说明 | |------|------| -| `mysql.connect(config, cb)` | 异步连接;回调 `function cb(err, conn)` | +| `mysql.connect(config, cb, ...)` | 异步连接;回调 `cb(conn, err, success, ...)`;`...` 为绑定参数,追加在固定参数后 | | `mysql_pool.create(config)` | 创建连接池 | -| `conn:query(sql, cb)` | 异步查询;回调 `function cb(err, result)` | -| `conn:stmt_prepare(sql, cb)` | 预处理语句 | -| `conn:stmt_execute(id, params, cb)` | 执行预处理语句 | +| `conn:query(sql, cb, ...)` | 异步查询;回调 `cb(conn, err, result, ...)`;`...` 为绑定参数 | +| `conn:stmt_prepare(sql, cb, ...)` | 预处理语句;回调 `cb(conn, err, stmt_id, ...)` | +| `conn:stmt_execute(id, params, cb, ...)` | 执行预处理语句;绑定参数从第 4 个参数起 | | `conn:stmt_close(id)` | 关闭预处理语句 | | `conn:close()` | 关闭连接 | -| `pool:acquire()` | 从池获取连接 | +| `pool:acquire()` | 从池获取连接(无可用连接返回 nil) | | `pool:release(conn)` | 归还连接到池 | +| `pool:with(fn)` | 租约式用法:取一条连接传给 `fn(conn)`,fn 返回(或抛错)后自动归还,返回 fn 的返回值。避免"忘 release / 回调不触发导致连接永不归还" | | `pool:close()` | 关闭连接池 | | `pool:stats()` | 返回 `{total, healthy}` | +**结果表布局(SELECT):** `result[1] = true`;`result[2]` 为列信息表(每项 `{[1]=列名, [2]=MySQL 类型码}`); +`result[3]` 为行表,每行按列位置(1 起)索引。 + +**结果表布局(INSERT/UPDATE/DELETE/DDL):** `result[1] = false`;`result[4] = affected_rows`; +`result[5] = last_insert_id`;`result[6] = info`。 + +**行值类型:** 按列类型转换——整数列(TINY/SHORT/LONG/LONGLONG/INT24/YEAR)返回 number(整数), +浮点/小数列(FLOAT/DOUBLE/DECIMAL/NEWDECIMAL)返回 number(浮点),其余列(字符串/日期/BLOB 等) +返回 string,NULL 返回 nil。布尔语义的 `TINYINT(1)` 返回数字(0/1),需要自行换算。 + +> **DECIMAL / NEWDECIMAL 精度:** 这两类列转成 IEEE 754 `double`,大约 15–16 位有效十进制数字。 +> `DECIMAL(18,6)` 这类高精度定义可能丢掉末尾 1–2 位,不适合作为金额的唯一表示。 +> 需要精确小数时在 SQL 里 `CAST(price AS CHAR)`,在 Lua 里按字符串处理。 + +```lua +-- 租约式用法:fn 是同帧同步调用,仍可用内联闭包;fn 返回后连接自动归还,抛错也不泄漏 +local ok = pool:with(function(c) + -- query 是异步的:结果在之后的 runtime.tick() 里通过【具名回调】到达。 + -- 绑定参数(这里的 "kick")登记时序列化暂存,派发时追加在 result 之后。 + c:query("UPDATE user SET online = 0 WHERE last_login < 100", "on_user_offline", "kick") + return true +end) + +function on_user_offline(conn, err, result, tag) + -- tag == "kick",是登记 query 时传入的绑定参数快照 + if err then print("query failed:", err, tag) end +end +``` + --- ## Redis diff --git a/src/native/json/native_json.cpp b/src/native/json/native_json.cpp index cf3040b..0b00ee7 100644 --- a/src/native/json/native_json.cpp +++ b/src/native/json/native_json.cpp @@ -112,7 +112,9 @@ static bj::value LuaToJsonValue(CVar v, int depth, std::unordered_set(VarType::Int)) { @@ -126,7 +128,7 @@ static bj::value LuaToJsonValue(CVar v, int depth, std::unordered_set max_idx) max_idx = key; } - if (is_array && (max_idx <= 0 || static_cast(max_idx) != kvs.size())) { + if (is_array && !kvs.empty() && static_cast(max_idx) != kvs.size()) { is_array = false; } @@ -181,10 +183,40 @@ static CVar JsonEncode(State *s, CVar *args, int n) { return inter::NativeToFakeluaString(s, out); } +// 判断 table 是否为数组形(空 table 视为数组;非空要求 key 为从 1 起的连续整数)。 +static bool IsArrayLikeTable(CVar v) { + if (v.type_ != static_cast(VarType::Table) || !v.data_.t) return false; + auto kvs = table::TableHelper::CollectKVPairs(v); + int64_t max_idx = 0; + for (auto &kv: kvs) { + if (kv.key.type_ != static_cast(VarType::Int)) return false; + int64_t key = kv.key.data_.i; + if (key < 1 || key > 1000000) return false; + if (key > max_idx) max_idx = key; + } + return kvs.empty() || static_cast(max_idx) == kvs.size(); +} + +// json.encode_array(value) → JSON string +// 严格数组编码通道:顶层必须是数组形 table(空 table → []),否则报错。 +// 用于协议字段必须为数组的场景,避免 encode 的对象/数组启发式歧义。 +static CVar JsonEncodeArray(State *s, CVar *args, int n) { + if (n < 1) ThrowBadArgument(1, "json.encode_array", "value expected"); + CVar a0 = inter::GetNativeArg(s, args, n, 0); + if (!IsArrayLikeTable(a0)) { + ThrowFakeluaException("json.encode_array: value is not an array-like table"); + } + std::unordered_set visited; + bj::value jv = LuaToJsonValue(a0, 0, visited); + std::string out = bj::serialize(jv); + return inter::NativeToFakeluaString(s, out); +} + void RegisterJsonLibraryApi(State *s) { if (!s) return; RegisterNativeFunction(s, "json.decode", 1, false, JsonDecode); RegisterNativeFunction(s, "json.encode", 1, false, JsonEncode); + RegisterNativeFunction(s, "json.encode_array", 1, false, JsonEncodeArray); } }// namespace fakelua::json diff --git a/src/native/mysql/mysql_connection.cpp b/src/native/mysql/mysql_connection.cpp index 42bf81d..0357005 100644 --- a/src/native/mysql/mysql_connection.cpp +++ b/src/native/mysql/mysql_connection.cpp @@ -1,5 +1,6 @@ #include "native/mysql/mysql_connection.h" #include "native/native_common.h" +#include "native/serialize/wire_codec.h" #include "native/table/native_table.h" #include "util/logging.h" #include "var/var.h" @@ -39,6 +40,42 @@ int WaitReady(short what) { }// namespace +// 按列类型把行值转成 Lua 值:整数列 → int、浮点/小数列 → float,其余(字符串/日期/BLOB 等) +// 保持 string,NULL 一律 nil。解析失败(理论上不该发生)回退为字符串,不丢数据。 +// 暴露给测试,声明在 mysql_connection.h。 +CVar FieldCellToCVar(::fakelua::State *s, const FieldCell &fv, int mysql_type) { + if (fv.is_null) return inter::NativeToFakeluaNil(s); + switch (mysql_type) { + case MYSQL_TYPE_TINY: + case MYSQL_TYPE_SHORT: + case MYSQL_TYPE_LONG: + case MYSQL_TYPE_LONGLONG: + case MYSQL_TYPE_INT24: + case MYSQL_TYPE_YEAR: { + char *end = nullptr; + long long v = std::strtoll(fv.value.c_str(), &end, 10); + if (end && !fv.value.empty() && *end == '\0') { + return inter::NativeToFakeluaLonglong(s, static_cast(v)); + } + break; + } + case MYSQL_TYPE_DECIMAL: + case MYSQL_TYPE_FLOAT: + case MYSQL_TYPE_DOUBLE: + case MYSQL_TYPE_NEWDECIMAL: { + char *end = nullptr; + double v = std::strtod(fv.value.c_str(), &end); + if (end && !fv.value.empty() && *end == '\0') { + return inter::NativeToFakeluaDouble(s, v); + } + break; + } + default: + break; + } + return inter::NativeToFakeluaString(s, fv.value); +} + MysqlConnection::MysqlConnection(::fakelua::State *state) : io_(state->GetIoContext()) { lua_state_ = state; } @@ -174,12 +211,23 @@ void MysqlConnection::FinishConnect(MYSQL *ret) { } void MysqlConnection::Query(const std::string &sql) { - if (close_pending_ || !mysql_) return; - if (state_ != ConnState::Ready || !ready_) { - pending_result_err_ = "connection not ready"; + // 契约:conn:query 一旦被调用,回调恰好被调用一次。 + // - 连接已死(mysql_ 释放 / 已请求关闭):下一轮 tick 给回调派发错误; + // - 连接未就绪或上一条 query 仍在飞行:排队,等 Ready 后自动启动。 + if (!mysql_) { + pending_result_err_ = "connection closed"; pending_result_ = true; return; } + if (close_pending_) { + pending_result_err_ = "connection closing"; + pending_result_ = true; + return; + } + if (state_ != ConnState::Ready || !ready_) { + queued_queries_.push_back(QueuedQuery{sql, result_cb_}); + return; + } last_sql_ = sql; state_ = ConnState::Querying; query_type_ = QueryType::Query; @@ -283,7 +331,11 @@ ResultsetData MysqlConnection::ConsumeResult(MYSQL_RES *res, bool is_resultset) } void MysqlConnection::StmtPrepare(const std::string &sql) { - if (close_pending_ || !mysql_) return; + if (!mysql_ || close_pending_) { + pending_result_err_ = !mysql_ ? "connection closed" : "connection closing"; + pending_result_ = true; + return; + } if (state_ != ConnState::Ready || !ready_) { pending_result_err_ = "connection not ready for prepare"; pending_result_ = true; @@ -327,7 +379,11 @@ void MysqlConnection::StmtPrepare(const std::string &sql) { } void MysqlConnection::StmtExecute(uint32_t stmt_id, const std::vector ¶ms) { - if (close_pending_ || !mysql_) return; + if (!mysql_ || close_pending_) { + pending_result_err_ = !mysql_ ? "connection closed" : "connection closing"; + pending_result_ = true; + return; + } if (state_ != ConnState::Ready || !ready_) { pending_result_err_ = "connection not ready for execute"; pending_result_ = true; @@ -669,6 +725,31 @@ void MysqlConnection::Tick() { pending_result_err_.clear(); pending_results_.clear(); } + + if (!queued_queries_.empty()) { + DrainQueryQueue(); + } +} + +void MysqlConnection::DrainQueryQueue() { + // 连接进入终态(Error/Idle):再也不会 Ready,逐条给排队的 query 派发错误回调, + // 保证"每条 query 恰好一次回调",而不是静默吞掉。 + if (state_ == ConnState::Error || state_ == ConnState::Idle) { + while (!queued_queries_.empty()) { + QueuedQuery q = std::move(queued_queries_.front()); + queued_queries_.erase(queued_queries_.begin()); + DispatchCallbackWithResult(q.cb, "connection not available"); + } + return; + } + // 连接可用:启动下一条排队的 query。每轮 tick 至多启动一条—— + // 若它同步出错(pending_result_ 被占用),剩下的留到下一轮,避免覆盖未派发的结果。 + if (state_ == ConnState::Ready && ready_) { + QueuedQuery q = std::move(queued_queries_.front()); + queued_queries_.erase(queued_queries_.begin()); + result_cb_ = q.cb; + Query(q.sql); + } } MysqlError MysqlConnection::LastError() const { @@ -685,11 +766,16 @@ bool MysqlConnection::IsRetryable(MysqlErrorType type) { } } -void MysqlConnection::SetConnectCallback(const std::string &name) { connect_cb_ = name; } -void MysqlConnection::SetResultCallback(const std::string &name) { result_cb_ = name; } +void MysqlConnection::SetConnectCallback(const ResultCallback &cb) { connect_cb_ = cb; } +void MysqlConnection::SetResultCallback(const ResultCallback &cb) { result_cb_ = cb; } void MysqlConnection::SetState(::fakelua::State *state) { lua_state_ = state; } void MysqlConnection::SetNativeObject(::fakelua::NativeObject *obj) { native_obj_ = obj; } bool MysqlConnection::Connected() const { return ready_; } +void MysqlConnection::MarkConnectedForTest() { + state_ = ConnState::Ready; + ready_ = true; + close_pending_ = false; +} bool MysqlConnection::Connecting() const { return state_ == ConnState::Connecting || state_ == ConnState::Handshaking; } int MysqlConnection::TickDepth() const { return tick_depth_; } bool MysqlConnection::ClosePending() const { return close_pending_; } @@ -702,80 +788,105 @@ void MysqlConnection::SetError(MysqlErrorType type, uint16_t code, const std::st last_error_.sql_state = sql_state; } -void MysqlConnection::DispatchConnect(const char *err_msg) { - TickDepthGuard guard(tick_depth_); - native::IoContext::DispatchScope dispatch_scope(io_); - if (close_pending_) return; - if (!lua_state_ || connect_cb_.empty()) return; - auto func = lua_state_->GetVM().GetFunction(connect_cb_); +void MysqlConnection::InvokeCallback(const ResultCallback &cb, CVar *args, int n) { + // 派回【登记回调时的引擎】(与旧内联闭包自带 func_ptr 的语义一致)。 + auto func = lua_state_->GetVM().GetFunction(cb.name); if (func.Empty()) return; - void *addr = func.GetAddr(JIT_TCC); - JITType jit_type = JIT_TCC; + JITType jit_type = cb.jit; + void *addr = func.GetAddr(jit_type); if (!addr) { - addr = func.GetAddr(JIT_GCC); - jit_type = JIT_GCC; + // 兜底:登记引擎没有该函数的编译产物(正常不应发生——同一份脚本三引擎都会编译)。 + // 按其余引擎依次找,避免回调静默丢失。 + static constexpr JITType kFallbacks[] = {JIT_TCC, JIT_GCC, JIT_INTERP}; + for (JITType t: kFallbacks) { + if (t == jit_type) continue; + addr = func.GetAddr(t); + if (addr) { + jit_type = t; + break; + } + } } if (!addr) return; - CVar args[3]; - args[0] = native_obj_ ? inter::NativeToFakeluaNativeObject(lua_state_, native_obj_) : inter::NativeToFakeluaNil(lua_state_); + inter::DispatchCall(lua_state_, addr, args, n, jit_type); +} + +namespace { + +// 把 3 个固定参数和登记时暂存的绑定参数拼成完整回调参数: +// 固定参数在前,绑定参数(登记时刻的值快照)追加在后。 +// DecodedArgs 持有解码字符串的后备内存,必须活到回调调用结束。 +std::vector BuildCallbackArgs(CVar a0, CVar a1, CVar a2, serialize::DecodedArgs &decoded) { + std::vector all; + all.reserve(3 + decoded.vars.size()); + all.push_back(a0); + all.push_back(a1); + all.push_back(a2); + for (CVar v: decoded.vars) all.push_back(v); + return all; +} + +}// namespace + +void MysqlConnection::DispatchConnect(const char *err_msg) { + TickDepthGuard guard(tick_depth_); + native::IoContext::DispatchScope dispatch_scope(io_); + if (!lua_state_ || connect_cb_.Empty()) return; + + // 在派发帧内 decode 绑定参数到临时 arena;decoded 生命周期覆盖整个函数。 + serialize::DecodedArgs decoded = serialize::WireDecodeCallbackArgs(lua_state_, connect_cb_.bound); + + CVar a0 = native_obj_ ? inter::NativeToFakeluaNativeObject(lua_state_, native_obj_) : inter::NativeToFakeluaNil(lua_state_); + CVar a1; + CVar a2; if (err_msg && err_msg[0]) { - args[1] = inter::NativeToFakeluaString(lua_state_, err_msg); - args[2] = inter::NativeToFakeluaInt(lua_state_, 0); + a1 = inter::NativeToFakeluaString(lua_state_, err_msg); + a2 = inter::NativeToFakeluaInt(lua_state_, 0); } else { - args[1] = inter::NativeToFakeluaNil(lua_state_); - args[2] = inter::NativeToFakeluaInt(lua_state_, 1); + a1 = inter::NativeToFakeluaNil(lua_state_); + a2 = inter::NativeToFakeluaInt(lua_state_, 1); } - inter::DispatchCall(lua_state_, addr, args, 3, jit_type); + std::vector args = BuildCallbackArgs(a0, a1, a2, decoded); + InvokeCallback(connect_cb_, args.data(), static_cast(args.size())); } -void MysqlConnection::DispatchResult(const char *err_msg) { +void MysqlConnection::DispatchCallbackWithResult(const ResultCallback &cb, const char *err_msg) { TickDepthGuard guard(tick_depth_); native::IoContext::DispatchScope dispatch_scope(io_); - if (close_pending_) return; - if (!lua_state_ || result_cb_.empty()) return; - auto func = lua_state_->GetVM().GetFunction(result_cb_); - if (func.Empty()) return; - void *addr = func.GetAddr(JIT_TCC); - JITType jit_type = JIT_TCC; - if (!addr) { - addr = func.GetAddr(JIT_GCC); - jit_type = JIT_GCC; - } - if (!addr) return; + if (!lua_state_ || cb.Empty()) return; + + // 多结果集可能在本函数内多次调用同一回调;绑定参数只 decode 一次,各次调用复用。 + // decoded 必须活到最后一次 InvokeCallback 返回。 + serialize::DecodedArgs decoded = serialize::WireDecodeCallbackArgs(lua_state_, cb.bound); const char *msg = err_msg && err_msg[0] ? err_msg : "query failed"; - CVar args[3]; - args[0] = native_obj_ ? inter::NativeToFakeluaNativeObject(lua_state_, native_obj_) : inter::NativeToFakeluaNil(lua_state_); + CVar a0 = native_obj_ ? inter::NativeToFakeluaNativeObject(lua_state_, native_obj_) : inter::NativeToFakeluaNil(lua_state_); + CVar nil{}; + nil.type_ = static_cast(VarType::Nil); + if (err_msg) { - args[1] = inter::NativeToFakeluaString(lua_state_, msg); - CVar nil{}; - nil.type_ = static_cast(VarType::Nil); - args[2] = nil; - inter::DispatchCall(lua_state_, addr, args, 3, jit_type); + std::vector args = BuildCallbackArgs(a0, inter::NativeToFakeluaString(lua_state_, msg), nil, decoded); + InvokeCallback(cb, args.data(), static_cast(args.size())); } else if (dispatch_stmt_id_ != 0) { - CVar nil{}; - nil.type_ = static_cast(VarType::Nil); - args[1] = nil; - args[2] = inter::NativeToFakeluaInt(lua_state_, static_cast(dispatch_stmt_id_)); - inter::DispatchCall(lua_state_, addr, args, 3, jit_type); + std::vector args = BuildCallbackArgs(a0, nil, + inter::NativeToFakeluaInt(lua_state_, static_cast(dispatch_stmt_id_)), decoded); + InvokeCallback(cb, args.data(), static_cast(args.size())); } else if (pending_results_.empty()) { - CVar nil{}; - nil.type_ = static_cast(VarType::Nil); - args[1] = nil; - args[2] = table::TableHelper::CreateTable(lua_state_); - inter::DispatchCall(lua_state_, addr, args, 3, jit_type); + std::vector args = BuildCallbackArgs(a0, nil, table::TableHelper::CreateTable(lua_state_), decoded); + InvokeCallback(cb, args.data(), static_cast(args.size())); } else { for (const auto &rs: pending_results_) { - CVar nil{}; - nil.type_ = static_cast(VarType::Nil); - args[1] = nil; - args[2] = ResultsetToLua(lua_state_, rs); - inter::DispatchCall(lua_state_, addr, args, 3, jit_type); + std::vector args = BuildCallbackArgs(a0, nil, ResultsetToLua(lua_state_, rs), decoded); + InvokeCallback(cb, args.data(), static_cast(args.size())); if (close_pending_) break; } } } +void MysqlConnection::DispatchResult(const char *err_msg) { + DispatchCallbackWithResult(result_cb_, err_msg); +} + CVar MysqlConnection::ResultsetToLua(::fakelua::State *s, const ResultsetData &result) { CVar tbl = table::TableHelper::CreateTable(s); if (result.is_resultset) { @@ -795,8 +906,8 @@ CVar MysqlConnection::ResultsetToLua(::fakelua::State *s, const ResultsetData &r CVar row_tbl = table::TableHelper::CreateTable(s); int64_t col_pos = 1; for (const auto &fv: row) { - if (fv.is_null) table::TableHelper::SetTableInt(s, row_tbl, col_pos, inter::NativeToFakeluaNil(s)); - else table::TableHelper::SetTableInt(s, row_tbl, col_pos, inter::NativeToFakeluaString(s, fv.value)); + const int col_type = col_pos <= static_cast(result.columns.size()) ? result.columns[static_cast(col_pos - 1)].second : 0; + table::TableHelper::SetTableInt(s, row_tbl, col_pos, FieldCellToCVar(s, fv, col_type)); ++col_pos; } table::TableHelper::SetTableInt(s, rows_tbl, row_idx++, row_tbl); diff --git a/src/native/mysql/mysql_connection.h b/src/native/mysql/mysql_connection.h index 7e3261b..9f2e526 100644 --- a/src/native/mysql/mysql_connection.h +++ b/src/native/mysql/mysql_connection.h @@ -3,6 +3,7 @@ // mysql_connection.h — async MySQL client using libmysqlclient (MariaDB Connector/C) // MYSQL_OPT_NONBLOCK + libevent wait on mysql_get_socket(). +#include "fakelua.h" #include "native/native_io_context.h" #ifdef __cplusplus @@ -28,6 +29,10 @@ class State; class NativeObject; }// namespace fakelua +// FieldCell 和 FieldCellToCVar 声明在独立的轻量头文件中(不依赖 mysql.h/libevent), +// 可直接在测试代码里 include。 +#include "native/mysql/mysql_field_convert.h" + namespace fakelua::mysql { struct StmtParam { @@ -35,6 +40,25 @@ struct StmtParam { std::string value; }; + +// C++ → Lua 异步回调:全局函数名 + 绑定参数。 +// fakelua 没有 GC,运行期闭包活在临时 arena 上,顶层 Call 的 State::Reset() 会回收, +// 无法安全地跨 tick 持有闭包裸指针。因此异步回调只接受【全局函数名】;脚本需要上下文时 +// 用绑定参数(mysql.connect / conn:query 等名字后面的多余实参): +// 登记当下把纯数据参数序列化成自有的字节串 bound(nil/bool/number/string/table), +// 跨任意次 Reset 暂存;结果派发时在当前 tick 帧内 decode 到临时 arena,追加在固定参数 +// (conn, err, result, ...)之后调用函数。bound 随 query 完成/连接销毁释放, +// 内存只与在途 query 数相关,不占用 const arena。 +struct ResultCallback { + std::string name; + std::string bound;// wire 编码的绑定参数元组;空串表示无绑定参数 + // 登记回调时所在的脚本引擎。结果回来后派回【同一引擎】:在哪个引擎发起的 IO, + // 回调就执行哪个引擎的编译产物,与内联闭包自带 func_ptr 的语义一致。 + JITType jit = JIT_TCC; + + [[nodiscard]] bool Empty() const { return name.empty(); } +}; + enum class SslMode { Disable = 0, Enable, @@ -59,12 +83,8 @@ struct MysqlError { std::string sql_state; }; -struct FieldCell { - bool is_null = false; - std::string value; -}; - struct ResultsetData { + bool is_resultset = false; std::vector> columns; std::vector> rows; @@ -94,14 +114,17 @@ class MysqlConnection { MysqlError LastError() const; static bool IsRetryable(MysqlErrorType type); - void SetConnectCallback(const std::string &name); - void SetResultCallback(const std::string &name); + void SetConnectCallback(const ResultCallback &cb); + void SetResultCallback(const ResultCallback &cb); void SetState(::fakelua::State *state); void SetNativeObject(::fakelua::NativeObject *obj); bool Connected() const; bool Connecting() const; + // 单测:跳过握手,把连接标成可被池 Acquire。生产路径不会调用。 + void MarkConnectedForTest(); + int TickDepth() const; bool ClosePending() const; void RequestClose(); @@ -116,10 +139,18 @@ class MysqlConnection { ::fakelua::State *lua_state_ = nullptr; ::fakelua::NativeObject *native_obj_ = nullptr; - std::string connect_cb_; - std::string result_cb_; + ResultCallback connect_cb_; + ResultCallback result_cb_; std::string last_sql_; + // 同连接飞行中再次发起的 query:排队而不是报 "connection not ready", + // 由 Tick() 在上一条结果派发完毕、连接回到 Ready 后依次启动。 + struct QueuedQuery { + std::string sql; + ResultCallback cb; + }; + std::vector queued_queries_; + enum class QueryType { None, Query, StmtPrepare, StmtExecute, Ping }; QueryType query_type_ = QueryType::None; @@ -164,8 +195,17 @@ class MysqlConnection { void DispatchConnect(const char *err_msg); void DispatchResult(const char *err_msg); + void DispatchCallbackWithResult(const ResultCallback &cb, const char *err_msg); void SetError(MysqlErrorType type, uint16_t code, const std::string &msg, const std::string &sql_state); + // 统一的回调调用入口:按函数名查 VM 注册表并派发。 + // 调用方需已把固定参数和 decode 后的绑定参数拼进 args。 + // 回调缺失或不可调用时返回空 CVar(调用方无需关心返回值)。 + void InvokeCallback(const ResultCallback &cb, CVar *args, int n); + // 连接可用时启动下一条排队 query;连接进入终态(Error/Idle)时 + // 逐条给排队的 query 回调派发错误,保证"每条 query 恰好一次回调"。 + void DrainQueryQueue(); + void Teardown(); void ApplySsl(); void ArmWait(int status); diff --git a/src/native/mysql/mysql_connection_pool.cpp b/src/native/mysql/mysql_connection_pool.cpp index bd651b8..33488f0 100644 --- a/src/native/mysql/mysql_connection_pool.cpp +++ b/src/native/mysql/mysql_connection_pool.cpp @@ -183,6 +183,16 @@ size_t MysqlConnectionPool::HealthyCount() const { return count; } +void MysqlConnectionPool::MarkConnectedForTest() { + std::lock_guard lock(mutex_); + for (auto &entry: pool_) { + if (!entry.conn) continue; + entry.conn->MarkConnectedForTest(); + entry.healthy = true; + entry.in_use = false; + } +} + // Private helpers void MysqlConnectionPool::SendHeartbeat(PoolEntry &entry) { if (!entry.conn || !entry.healthy) return; diff --git a/src/native/mysql/mysql_connection_pool.h b/src/native/mysql/mysql_connection_pool.h index e87122c..9e2453a 100644 --- a/src/native/mysql/mysql_connection_pool.h +++ b/src/native/mysql/mysql_connection_pool.h @@ -68,6 +68,9 @@ class MysqlConnectionPool { size_t TotalCount() const; size_t HealthyCount() const; + // 单测:把已创建的连接标成 healthy + Connected,使 Acquire 不必等真实握手。 + void MarkConnectedForTest(); + private: PoolConfig config_; ::fakelua::State *state_ = nullptr; diff --git a/src/native/mysql/mysql_field_convert.h b/src/native/mysql/mysql_field_convert.h new file mode 100644 index 0000000..866c688 --- /dev/null +++ b/src/native/mysql/mysql_field_convert.h @@ -0,0 +1,31 @@ +#pragma once + +// mysql_field_convert.h — 行值类型转换:FieldCell → Lua CVar。 +// 此头文件不依赖 mysql.h 或 libevent,可直接在测试代码中 include。 + +#include +#include + +namespace fakelua { +struct CVar; +class State; +}// namespace fakelua + +namespace fakelua::mysql { + +// MySQL 结果行中的一个字段值。 +struct FieldCell { + bool is_null = false; + std::string value; +}; + +// 按列类型把行值转成 Lua CVar。 +// mysql_type 为 MariaDB/MySQL 的 enum_field_types 整数值(mysql.h MYSQL_TYPE_*)。 +// 整数列(TINY/SHORT/LONG/LONGLONG/INT24/YEAR)→ VarType::Int +// 浮点/小数列(FLOAT/DOUBLE/DECIMAL/NEWDECIMAL)→ VarType::Float +// 其他列(字符串/日期/BLOB 等)→ VarType::String +// NULL 字段(is_null == true)→ VarType::Nil +// 解析失败(如整数列收到非数字值)→ 回退为 VarType::String,不丢数据。 +CVar FieldCellToCVar(::fakelua::State *s, const FieldCell &fv, int mysql_type); + +}// namespace fakelua::mysql diff --git a/src/native/mysql/native_mysql.cpp b/src/native/mysql/native_mysql.cpp index a4f9257..92e93f2 100644 --- a/src/native/mysql/native_mysql.cpp +++ b/src/native/mysql/native_mysql.cpp @@ -3,6 +3,7 @@ #include "native/mysql/mysql_connection_pool.h" #include "native/native_common.h" #include "native/object/native_object.h" +#include "native/serialize/wire_codec.h" #include "native/table/native_table.h" #include "var/var.h" @@ -36,6 +37,28 @@ static std::string CVarToString(CVar v) { return {}; } +// 解析异步回调:只接受全局函数名(字符串)。 +// fakelua 没有 GC,内联闭包活在临时 arena 上,无法安全地跨多次 tick/Reset 保存, +// 因此异步回调不支持闭包;需要上下文时在函数名后追加纯数据绑定参数: +// conn:query(sql, "on_result", tag, ctx_table) +// 绑定参数在登记当下序列化成自有字节串随回调暂存,派发时 decode 追加在固定参数后。 +// 名称参数非法(闭包/数字/空串等)直接抛错——历史上坏回调会被静默转成空串、永不触发。 +static ResultCallback BuildCallback(State *s, CVar name_v, int name_argno, CVar *args, int bound_first, int n, + const char *fname) { + if (name_v.type_ == static_cast(VarType::Closure)) { + ThrowBadArgument(name_argno, fname, + "global callback function name expected; inline closures are not supported for async callbacks, pass a name string plus data args"); + } + std::string name = CVarToString(name_v); + if (name.empty()) { + ThrowBadArgument(name_argno, fname, "callback function name expected"); + } + // 绑定参数严格校验 + 序列化(含闭包/native 对象等直接抛 bad argument)。 + std::string bound = serialize::WireEncodeCallbackArgs(s, args, bound_first, n, fname, bound_first + 1); + // 记住登记时的引擎:回调派发时派回同一引擎,而不是固定 TCC。 + return ResultCallback{std::move(name), std::move(bound), s->CurrentJit()}; +} + // Retrieve MysqlConnection* from NativeObject MysqlConnection *UnwrapConnNative(NativeObject *self) { if (!self) return nullptr; @@ -193,9 +216,8 @@ static CVar MysqlConnect(State *s, CVar *args, int n) { if (user.empty()) ThrowBadArgument(1, "mysql.connect", "user required"); - // Read callback function name - std::string cb_name = CVarToString(a1); - if (cb_name.empty()) ThrowBadArgument(1, "mysql.connect", "callback function expected"); + // Read callback: global function name; args after it are serialized bound params + ResultCallback cb = BuildCallback(s, a1, 2, args, 2, n, "mysql.connect"); // Create NativeObject wrapper first (so callbacks can dispatch) int64_t gid = s->GetNativeObjectManager().CreateGroup(); @@ -223,7 +245,7 @@ static CVar MysqlConnect(State *s, CVar *args, int n) { // Create connection (async) auto *conn = new MysqlConnection(s); - conn->SetConnectCallback(cb_name); + conn->SetConnectCallback(cb); conn->SetNativeObject(nat); nat->SetInt("__mysql_conn__", reinterpret_cast(conn)); @@ -248,13 +270,15 @@ CVar ConnQuery(NativeObject *self, State *s, CVar *args, int n) { CVar a0 = inter::GetNativeArg(s, args, n, 0); CVar a1 = inter::GetNativeArg(s, args, n, 1); std::string sql = CVarToString(a0); - std::string cb_name = CVarToString(a1); + ResultCallback cb = BuildCallback(s, a1, 2, args, 2, n, "conn:query"); auto *conn = UnwrapConnNative(self); - if (!conn || !conn->Connected()) error("conn:query: connection is closed"); + if (!conn) error("conn:query: connection is closed"); + // 连接未就绪(握手中/重连中/上一条 query 仍在飞行)不再抛错: + // 回调随 query 入队,由连接的 tick 在可用时自动启动(见 MysqlConnection::Query)。 conn->SetState(s); - conn->SetResultCallback(cb_name); + conn->SetResultCallback(cb); conn->Query(sql); MaybeReleaseOwnedConn(self); @@ -269,13 +293,13 @@ CVar ConnStmtPrepare(NativeObject *self, State *s, CVar *args, int n) { CVar a0 = inter::GetNativeArg(s, args, n, 0); CVar a1 = inter::GetNativeArg(s, args, n, 1); std::string sql = CVarToString(a0); - std::string cb_name = CVarToString(a1); + ResultCallback cb = BuildCallback(s, a1, 2, args, 2, n, "conn:stmt_prepare"); auto *conn = UnwrapConnNative(self); if (!conn || !conn->Connected()) error("conn:stmt_prepare: connection is closed"); conn->SetState(s); - conn->SetResultCallback(cb_name); + conn->SetResultCallback(cb); conn->StmtPrepare(sql); MaybeReleaseOwnedConn(self); @@ -292,7 +316,7 @@ CVar ConnStmtExecute(NativeObject *self, State *s, CVar *args, int n) { CVar a2 = inter::GetNativeArg(s, args, n, 2); uint32_t stmt_id = static_cast(inter::CVarToInteger(a0, 0)); - std::string cb_name = CVarToString(a2); + ResultCallback cb = BuildCallback(s, a2, 3, args, 3, n, "conn:stmt_execute"); std::vector params; if (a1.type_ == static_cast(VarType::Table) && a1.data_.t) { @@ -333,7 +357,7 @@ CVar ConnStmtExecute(NativeObject *self, State *s, CVar *args, int n) { if (!conn || !conn->Connected()) error("conn:stmt_execute: connection is closed"); conn->SetState(s); - conn->SetResultCallback(cb_name); + conn->SetResultCallback(cb); conn->StmtExecute(stmt_id, params); MaybeReleaseOwnedConn(self); @@ -396,7 +420,8 @@ CVar ConnPing(NativeObject *self, State *s, CVar * /*args*/, int /*n*/) { void RegisterMysqlLibraryApi(State *s) { if (!s) return; - RegisterNativeFunction(s, "mysql.connect", 2, false, MysqlConnect); + // 可变参数:config、回调函数名之后可跟任意个纯数据绑定参数。 + RegisterNativeFunction(s, "mysql.connect", 2, true, MysqlConnect); } }// namespace fakelua::mysql diff --git a/src/native/mysql/native_mysql.h b/src/native/mysql/native_mysql.h index 551651f..2fdd513 100644 --- a/src/native/mysql/native_mysql.h +++ b/src/native/mysql/native_mysql.h @@ -36,4 +36,8 @@ void OnStateDeleted(State *s); // 驱动本 State 上所有连接池和连接。由 runtime.tick() 调用。 void TickAll(State *s); +// 单测:pool:with 的 fn 抛错后连接必须回到可再次 Acquire 的状态。 +// 归还成功返回 1;fn 没有抛错返回 -1;归还失败(仍被占用)返回 0。 +int TestPoolWithFnThrowReturnsConnection(State *s, CVar fn); + }// namespace fakelua::mysql diff --git a/src/native/mysql/native_mysql_pool.cpp b/src/native/mysql/native_mysql_pool.cpp index fc21a6e..e59f6b6 100644 --- a/src/native/mysql/native_mysql_pool.cpp +++ b/src/native/mysql/native_mysql_pool.cpp @@ -23,6 +23,7 @@ static CVar PoolAcquire(NativeObject *self, State *s, CVar *args, int n); static CVar PoolRelease(NativeObject *self, State *s, CVar *args, int n); static CVar PoolClose(NativeObject *self, State *s, CVar *args, int n); static CVar PoolStats(NativeObject *self, State *s, CVar *args, int n); +static CVar PoolWith(NativeObject *self, State *s, CVar *args, int n); static CVar ConnPoolRelease(NativeObject *self, State *s, CVar *args, int n); static CVar ConnErrorInfo(NativeObject *self, State *s, CVar *args, int n); static MysqlConnection *UnwrapConn(CVar v); @@ -199,20 +200,18 @@ static CVar PoolCreate(State *s, CVar *args, int n) { }); nat->RegisterMethod("acquire", PoolAcquire); nat->RegisterMethod("release", PoolRelease); + nat->RegisterMethod("with", PoolWith); nat->RegisterMethod("close", PoolClose); nat->RegisterMethod("stats", PoolStats); return inter::NativeToFakeluaNativeObject(s, nat); } -// pool:acquire() → connection - -static CVar PoolAcquire(NativeObject *self, State *s, CVar * /*args*/, int /*n*/) { - auto *pool_obj = UnwrapPool(self); - if (!pool_obj || !pool_obj->pool) return inter::NativeToFakeluaNil(s); - +// 从池里取一条连接并包成 NativeObject(pool:acquire / pool:with 共用)。 +// 池满/无健康连接时返回 nil。 +static NativeObject *AcquireConnObject(State *s, PoolObject *pool_obj) { auto *conn = pool_obj->pool->Acquire(); - if (!conn) return inter::NativeToFakeluaNil(s); + if (!conn) return nullptr; // Wrap connection in NativeObject for Lua (use a new group for each connection) int64_t conn_gid = s->GetNativeObjectManager().CreateGroup(); @@ -236,10 +235,95 @@ static CVar PoolAcquire(NativeObject *self, State *s, CVar * /*args*/, int /*n*/ conn->SetState(s); conn->SetNativeObject(nat); + return nat; +} + +// pool:acquire() → connection +static CVar PoolAcquire(NativeObject *self, State *s, CVar * /*args*/, int /*n*/) { + auto *pool_obj = UnwrapPool(self); + if (!pool_obj || !pool_obj->pool) return inter::NativeToFakeluaNil(s); + + auto *nat = AcquireConnObject(s, pool_obj); + if (!nat) return inter::NativeToFakeluaNil(s); return inter::NativeToFakeluaNativeObject(s, nat); } +// pool:with(fn) → fn 的返回值 +// 租约式用法:从池里取一条连接传给 fn,fn 返回(或抛错)后自动归还连接。 +// 彻底避免"忘了 release / 回调不触发导致连接永不归还"的泄漏模式。 + +static CVar PoolWith(NativeObject *self, State *s, CVar *args, int n) { + auto *pool_obj = UnwrapPool(self); + if (!pool_obj || !pool_obj->pool) return inter::NativeToFakeluaNil(s); + if (n < 1) ThrowBadArgument(1, "pool:with", "function expected"); + CVar a0 = inter::GetNativeArg(s, args, n, 0); + if (a0.type_ != static_cast(VarType::Closure) || !a0.data_.cl) { + ThrowBadArgument(1, "pool:with", "function expected"); + } + + auto *nat = AcquireConnObject(s, pool_obj); + if (!nat) return inter::NativeToFakeluaNil(s); + + CVar conn_arg = inter::NativeToFakeluaNativeObject(s, nat); + // 析构时归还:fn 正常返回和抛错都走这里,避免 catch 漏掉某条路径。 + struct LeaseGuard { + NativeObject *nat; + ~LeaseGuard() { DetachAcquiredWrapper(nat); } + } lease{nat}; + // fn 是同帧同步调用,直接沿用当前 Lua 调用方的引擎(异常边界随之正确), + // 不固定 TCC 标记。 + return inter::DispatchCallClosure(s, a0.data_.cl, &conn_arg, 1, s->CurrentJit()); +} + +int TestPoolWithFnThrowReturnsConnection(State *s, CVar fn) { + PoolConfig config; + config.host = "127.0.0.1"; + config.port = 1; + config.user = "root"; + config.password = "x"; + config.database = "test"; + config.pool_size = 1; + config.connect_timeout_ms = 200; + config.read_timeout_ms = 200; + config.heartbeat_interval_ms = 0; + config.max_retries = 0; + + auto *pool_obj = new PoolObject(); + pool_obj->config = config; + pool_obj->pool = std::make_unique(config, s); + pool_obj->pool->Initialize(); + pool_obj->pool->MarkConnectedForTest(); + + int64_t gid = s->GetNativeObjectManager().CreateGroup(); + auto *nat = s->GetNativeObjectManager().Create(gid, "mysql_pool"); + nat->SetInt("__mysql_pool__", reinterpret_cast(pool_obj)); + RegisterMysqlNativeWrapper(s, nat, true); + nat->SetFinalizer([](NativeObject *self) { + UnregisterMysqlNativeWrapper(self); + auto *p = UnwrapPool(self); + if (p) { + InvalidateAcquiredWrappers(p); + delete p; + self->SetInt("__mysql_pool__", 0); + } + }); + + bool threw = false; + try { + CVar args[1] = {fn}; + PoolWith(nat, s, args, 1); + } catch (...) { + threw = true; + } + if (!threw) return -1; + // 归还成功则能再次取到同一条连接;仍被占用则 Acquire 返回空。 + MysqlConnection *again = pool_obj->pool->Acquire(); + if (!again) return 0; + pool_obj->pool->Release(again); + return 1; +} + // pool:release(conn) static CVar PoolRelease(NativeObject *self, State *s, CVar *args, int n) { diff --git a/src/native/native_common.h b/src/native/native_common.h index bd276ac..e5806fe 100644 --- a/src/native/native_common.h +++ b/src/native/native_common.h @@ -10,6 +10,13 @@ namespace fakelua { +// C++ 调用 Lua 闭包时传给 DispatchCall / DispatchCallClosure 的后端标记。 +// 这个值不选择函数地址:闭包的机器码在 VarClosure::func_ptr 里,具名函数由调用方 +// 按 VM 注册表取地址(TCC 没有再回退 GCC)。它只决定异常能否穿过代码页—— +// TCC 没有展开表,必须走错误边界;GCC 多包一层边界同样安全。解释器闭包由 +// IsInterpClosure 识别,与这里的标记无关,因此不会在 GCC/解释器后端静默失败。 +inline constexpr JITType kNativeCallbackJit = JIT_TCC; + // Shared helpers for native library argument validation and error reporting. // Used across native_math, native_table, native_utf8, native_string, native_io. diff --git a/src/native/native_io_context.h b/src/native/native_io_context.h index 2145ad0..4c3228a 100644 --- a/src/native/native_io_context.h +++ b/src/native/native_io_context.h @@ -38,6 +38,14 @@ class IoContext { std::size_t Poll(); + // C++→Lua 回调嵌套深度。DispatchScope 在各派发入口 +1,析构时 -1。 + // 纯计数器,不依赖平台:Linux / macOS / Windows 行为一致。 + // 当前包住回调的位置: + // net TcpServer / TcpClient / UdpSocket::DrainEventsWith + // mysql MysqlConnection::DispatchConnect / DispatchCallbackWithResult + // http CallNamed + // redis CallNamed,以及连接/命令结果派发 + // net 的 send 看到 InDispatch() 为真时改为入队,由本轮 tick 在 scope 结束后泵出。 bool InDispatch() const { return dispatch_depth_ > 0; } diff --git a/src/native/net/native_net.cpp b/src/native/net/native_net.cpp index 1bb9999..284df44 100644 --- a/src/native/net/native_net.cpp +++ b/src/native/net/native_net.cpp @@ -3,6 +3,7 @@ #include "native/net/net_event.h" #include "native/object/native_object.h" #include "native/table/native_table.h" +#include "state/state.h" #include "util/logging.h" #include "var/var.h" @@ -45,7 +46,19 @@ struct NetObject { int tick_depth = 0; bool close_pending = false; + // 回调上下文(native dispatch 期间)里的 send 先入队,由本轮 tick 派发完后统一 + // 泵出——消除"回调里 send 是否生效"的平台相关行为(事件驱动改轮询驱动)。 + struct PendingSend { + enum class Kind { Server, Client, UdpPeer, UdpTo } kind; + int connid = -1; // Server: 目标连接 + std::string data; + std::string ip; // UdpTo: 目标地址 + uint16_t port = 0; // UdpTo: 目标端口 + }; + std::vector pending_sends; + static constexpr size_t kMaxEvents = 1024; + static constexpr size_t kMaxPendingSends = 4096; }; // 辅助:从 CVar 提取字符串 @@ -126,6 +139,8 @@ static CVar CallLuaEvent(State *state, const std::string &func_name, const char // "echo", data → 将 data 发回来源连接 // "close" → 延后关闭本对象(tick 回调内安全) +static bool SendOrQueue(NetObject *obj, NetObject::PendingSend ps); + static void HandleCallbackReturn(NetObject *obj, const CVar &ret, int connid) { if (!obj || obj->close_pending) return; // 检查是否为 Multi(多返回值) @@ -143,14 +158,66 @@ static void HandleCallbackReturn(NetObject *obj, const CVar &ret, int connid) { std::string echo_data = CVarToString(data_var); if (obj->udp) { if (obj->udp->Connected()) { - obj->udp->Send(echo_data.data(), echo_data.size()); + SendOrQueue(obj, {NetObject::PendingSend::Kind::UdpPeer, -1, std::move(echo_data)}); } else { - obj->udp->SendTo(echo_data.data(), echo_data.size(), obj->last_peer_ip, obj->last_peer_port); + SendOrQueue(obj, {NetObject::PendingSend::Kind::UdpTo, -1, std::move(echo_data), obj->last_peer_ip, obj->last_peer_port}); } } else if (obj->is_server) { - obj->server->Send(connid, echo_data.data(), echo_data.size()); + SendOrQueue(obj, {NetObject::PendingSend::Kind::Server, connid, std::move(echo_data)}); } else { - obj->client->Send(echo_data.data(), echo_data.size()); + SendOrQueue(obj, {NetObject::PendingSend::Kind::Client, -1, std::move(echo_data)}); + } + } +} + +// send 的统一实现:按 PendingSend 的 kind 写 socket。返回 false 表示引擎侧 +// 拒绝(连接未就绪/已关闭),调用方记日志,不再静默。 +static bool DoSend(NetObject *obj, const NetObject::PendingSend &ps) { + using Kind = NetObject::PendingSend::Kind; + switch (ps.kind) { + case Kind::Server: + if (obj->server) return obj->server->Send(ps.connid, ps.data.data(), ps.data.size()); + return false; + case Kind::Client: + if (obj->client) return obj->client->Send(ps.data.data(), ps.data.size()); + return false; + case Kind::UdpPeer: + if (obj->udp && obj->udp->Connected()) return obj->udp->Send(ps.data.data(), ps.data.size()); + return false; + case Kind::UdpTo: + if (obj->udp) return obj->udp->SendTo(ps.data.data(), ps.data.size(), ps.ip, ps.port); + return false; + } + return false; +} + +// 判断是否处于 C++→Lua 回调派发上下文(on_event / mysql/http 回调等)。 +static bool InDispatchContext(NetObject *obj) { + return obj && obj->state && obj->state->GetIoContext().InDispatch(); +} + +// send 的统一入口:回调上下文里不直接写 socket,而是入队,由本轮 tick 派发完后 +// 统一泵出——消除"回调里 send 偶尔静默丢失"的平台相关行为;非派发上下文立即发送。 +static bool SendOrQueue(NetObject *obj, NetObject::PendingSend ps) { + if (InDispatchContext(obj)) { + if (obj->pending_sends.size() >= NetObject::kMaxPendingSends) { + LOG_WARN(obj->state, "net", "pending send queue full ({}), dropping {} bytes", obj->pending_sends.size(), ps.data.size()); + return false; + } + obj->pending_sends.push_back(std::move(ps)); + return true; + } + return DoSend(obj, ps); +} + +// 派发完成后统一泵出积压的 send。此时已脱离 DispatchScope,直接写 socket 是安全的。 +static void FlushPendingSends(NetObject *obj) { + if (!obj || obj->pending_sends.empty()) return; + auto sends = std::move(obj->pending_sends); + obj->pending_sends.clear(); + for (const auto &ps: sends) { + if (!DoSend(obj, ps)) { + LOG_WARN(obj->state, "net", "deferred send failed (kind={} connid={} len={})", static_cast(ps.kind), ps.connid, ps.data.size()); } } } @@ -294,10 +361,13 @@ static void TickNetObject(NativeObject *self) { } } catch (...) { obj->tick_depth--; + FlushPendingSends(obj); finish_tick(); throw; } obj->tick_depth--; + // 回调期间积压的 send 在这里统一泵出(socket 写在派发上下文之外执行)。 + FlushPendingSends(obj); finish_tick(); } @@ -314,6 +384,7 @@ void TickAll(State *s) { } // server:send(connid, data) / client:send(data) +// 返回 true 表示已被引擎接受(直接写出,或回调上下文中入队待泵出)。 static CVar NetSend(NativeObject *self, State *s, CVar *args, int n) { auto *obj = Unwrap(self); if (!obj) return inter::NativeToFakeluaBool(s, false); @@ -322,7 +393,7 @@ static CVar NetSend(NativeObject *self, State *s, CVar *args, int n) { if (obj->udp->Connected()) { if (n < 1) ThrowBadArgument(1, "send", "data expected"); std::string data = CVarToString(inter::GetNativeArg(s, args, n, 0)); - bool ok = obj->udp->Send(data.data(), data.size()); + bool ok = SendOrQueue(obj, {NetObject::PendingSend::Kind::UdpPeer, -1, std::move(data)}); return inter::NativeToFakeluaBool(s, ok); } if (n < 3) ThrowBadArgument(1, "send", "data, ip and port expected"); @@ -332,7 +403,7 @@ static CVar NetSend(NativeObject *self, State *s, CVar *args, int n) { if (port_val <= 0 || port_val > 65535) { ThrowFakeluaException(std::format("net: port {} out of range (1-65535)", port_val)); } - bool ok = obj->udp->SendTo(data.data(), data.size(), ip, static_cast(port_val)); + bool ok = SendOrQueue(obj, {NetObject::PendingSend::Kind::UdpTo, -1, std::move(data), std::move(ip), static_cast(port_val)}); return inter::NativeToFakeluaBool(s, ok); } @@ -351,7 +422,7 @@ static CVar NetSend(NativeObject *self, State *s, CVar *args, int n) { } return std::string(); }(); - bool ok = obj->server->Send(connid, data.data(), data.size()); + bool ok = SendOrQueue(obj, {NetObject::PendingSend::Kind::Server, connid, std::move(data)}); return inter::NativeToFakeluaBool(s, ok); } else if (!obj->is_server && obj->client) { // client: send(data) @@ -367,7 +438,7 @@ static CVar NetSend(NativeObject *self, State *s, CVar *args, int n) { } return std::string(); }(); - bool ok = obj->client->Send(data.data(), data.size()); + bool ok = SendOrQueue(obj, {NetObject::PendingSend::Kind::Client, -1, std::move(data)}); return inter::NativeToFakeluaBool(s, ok); } diff --git a/src/native/object/native_object.cpp b/src/native/object/native_object.cpp index 7e5e44a..1e570ce 100644 --- a/src/native/object/native_object.cpp +++ b/src/native/object/native_object.cpp @@ -492,6 +492,8 @@ CVar NativeObject::Wrap(State *s) const { vtbl->spec = spec; vtbl->spec_get = reinterpret_cast(NativeSpecGet); vtbl->spec_set = reinterpret_cast(NativeSpecSet); + vtbl->spec_bytes = static_cast(sizeof(NativeObjectSpec)); + vtbl->spec_cvars = 0; // 填充 spec_keys / spec_vals(供 pairs() 迭代) RefreshSpecKeys(vtbl, this, s); diff --git a/src/native/serialize/native_serialize.cpp b/src/native/serialize/native_serialize.cpp index 50c3c03..9887acc 100644 --- a/src/native/serialize/native_serialize.cpp +++ b/src/native/serialize/native_serialize.cpp @@ -1,4 +1,5 @@ #include "native/serialize/native_serialize.h" +#include "native/serialize/wire_codec.h" #include "native/native_common.h" #include "native/table/native_table.h" #include "util/logging.h" @@ -15,45 +16,21 @@ #include #include #include -#include #include #include #include -#include -#include #include #include namespace fakelua::serialize { -// Wire format(类 protobuf 编码) -// 每个值 = [type_tag(1 byte)] [payload] -// 0x00 nil -// 0x01 false -// 0x02 true -// 0x03 + varint 整数(zigzag 编码:小绝对值 → 小编码) -// 0x04 + 8 bytes double(小端 memcpy) -// 0x05 + varint(len) + bytes 新字符串,加入字典 -// 0x06 + varint(id) 字典中的字符串引用 -// 0x07 + varint(count) + N*(key,value) 表 -// 不支持的类型(闭包等)在表中跳过,顶层编码则抛错。 +// 紧凑 wire 格式(nil/bool/int/float/string/table)的编解码在 wire_codec 中实现, +// Lua serialize.encode/decode 与内部异步回调绑定参数共用同一格式,这里只做 Lua 层封装。 -enum Tag : uint8_t { - TAG_NIL = 0x00, - TAG_FALSE = 0x01, - TAG_TRUE = 0x02, - TAG_INT = 0x03, - TAG_DOUBLE = 0x04, - TAG_STR_NEW = 0x05, - TAG_STR_REF = 0x06, - TAG_TABLE = 0x07, -}; - -// 辅助:从 CVar 提取字符串(二进制安全) +// 辅助:从 CVar 提取字符串(二进制安全),供 Boost.Serialization 树使用。 static std::string CVarToString(CVar v) { if (v.type_ == static_cast(VarType::String) && v.data_.s) { - auto sv = v.data_.s->Str(); - return std::string(sv.data(), sv.size()); + return std::string(v.data_.s->Str()); } if (v.type_ == static_cast(VarType::StringId) && v.data_.i) { const char *ptr = reinterpret_cast(v.data_.i); @@ -63,270 +40,21 @@ static std::string CVarToString(CVar v) { return {}; } -static std::string_view CVarToStringView(CVar v) { - if (v.type_ == static_cast(VarType::String) && v.data_.s) { - return v.data_.s->Str(); - } - if (v.type_ == static_cast(VarType::StringId) && v.data_.i) { - const char *ptr = reinterpret_cast(v.data_.i); - int sz = *reinterpret_cast(ptr); - return std::string_view(ptr + 8, sz); - } - return {}; -} - -// 类型判断 -static bool IsSupported(CVar v) { - switch (v.type_) { - case static_cast(VarType::Nil): - case static_cast(VarType::Bool): - case static_cast(VarType::Int): - case static_cast(VarType::Float): - case static_cast(VarType::String): - case static_cast(VarType::StringId): - case static_cast(VarType::Table): - return true; - default: - return false; - } -} - -// Varint(LEB128 无符号) -static void WriteVarint(std::string &out, uint64_t v) { - while (v >= 0x80) { - out.push_back(static_cast((v & 0x7f) | 0x80)); - v >>= 7; - } - out.push_back(static_cast(v)); -} - -static uint64_t ReadVarint(const std::string &in, size_t &pos) { - uint64_t result = 0; - int shift = 0; - bool terminated = false; - while (pos < in.size()) { - uint8_t b = static_cast(in[pos++]); - result |= static_cast(b & 0x7f) << shift; - if ((b & 0x80) == 0) { - terminated = true; - break; - } - shift += 7; - if (shift >= 64) { - ThrowFakeluaException("serialize.decode: varint too long"); - } - } - if (!terminated) { - ThrowFakeluaException("serialize.decode: truncated varint"); - } - return result; -} - -// Zigzag(有符号整数 ↔ 无符号) -static uint64_t ZigzagEncode(int64_t n) { - return (static_cast(n) << 1) ^ static_cast(n >> 63); -} - -static int64_t ZigzagDecode(uint64_t u) { - return static_cast((u >> 1) ^ (-(u & 1))); -} - -// Double(小端 memcpy) -static void WriteDouble(std::string &out, double v) { - uint8_t buf[8]; - std::memcpy(buf, &v, 8); - out.append(reinterpret_cast(buf), 8); -} - -static double ReadDouble(const std::string &in, size_t &pos) { - if (pos + 8 > in.size()) { - ThrowFakeluaException("serialize.decode: truncated double"); - } - double v; - std::memcpy(&v, in.data() + pos, 8); - pos += 8; - return v; -} - -// 编码 -struct EncodeState { - std::unordered_map dict;// 字符串 → 字典 id - std::unordered_set visited; - int depth = 0; -}; - -static void EncodeValue(std::string &out, CVar v, EncodeState &state) { - switch (v.type_) { - case static_cast(VarType::Nil): - out.push_back(TAG_NIL); - return; - case static_cast(VarType::Bool): - out.push_back(AsVar(v).GetBool() ? TAG_TRUE : TAG_FALSE); - return; - case static_cast(VarType::Int): - out.push_back(TAG_INT); - WriteVarint(out, ZigzagEncode(v.data_.i)); - return; - case static_cast(VarType::Float): - out.push_back(TAG_DOUBLE); - WriteDouble(out, v.data_.f); - return; - case static_cast(VarType::String): - case static_cast(VarType::StringId): { - auto sv = CVarToStringView(v); - auto it = state.dict.find(sv); - if (it != state.dict.end()) { - out.push_back(TAG_STR_REF); - WriteVarint(out, it->second); - } else { - uint32_t id = static_cast(state.dict.size()); - state.dict.emplace(sv, id); - out.push_back(TAG_STR_NEW); - WriteVarint(out, static_cast(sv.size())); - if (!sv.empty()) { - out.append(sv.data(), sv.size()); - } - } - return; - } - case static_cast(VarType::Table): { - VarTable *t = v.data_.t; - if (!t) { - out.push_back(TAG_TABLE); - WriteVarint(out, 0); - return; - } - if (state.depth >= 64) { - ThrowFakeluaException("serialize.encode: nesting too deep"); - } - if (!state.visited.insert(t).second) { - ThrowFakeluaException("serialize.encode: cyclic table"); - } - state.depth++; - std::vector keys; - std::vector vals; - try { - table::TableHelper::ForEachKV(v, [&](CVar k, CVar val) { - if (IsSupported(k) && IsSupported(val)) { - keys.push_back(k); - vals.push_back(val); - } - }); - out.push_back(TAG_TABLE); - WriteVarint(out, static_cast(keys.size())); - for (size_t i = 0; i < keys.size(); ++i) { - EncodeValue(out, keys[i], state); - EncodeValue(out, vals[i], state); - } - } catch (...) { - state.depth--; - state.visited.erase(t); - throw; - } - state.depth--; - state.visited.erase(t); - return; - } - default: - ThrowFakeluaException("serialize.encode: unsupported type: " + VarTypeToString(static_cast(v.type_))); - } -} - -// 解码 -struct DecodeState { - std::vector dict;// id → 字符串 - int depth = 0; -}; - -static CVar DecodeValue(const std::string &in, size_t &pos, DecodeState &state, State *s) { - if (pos >= in.size()) { - ThrowFakeluaException("serialize.decode: unexpected end of input"); - } - uint8_t tag = static_cast(in[pos++]); - switch (tag) { - case TAG_NIL: - return inter::NativeToFakeluaNil(s); - case TAG_FALSE: - return inter::NativeToFakeluaBool(s, false); - case TAG_TRUE: - return inter::NativeToFakeluaBool(s, true); - case TAG_INT: { - uint64_t u = ReadVarint(in, pos); - return inter::NativeToFakeluaLonglong(s, ZigzagDecode(u)); - } - case TAG_DOUBLE: - return inter::NativeToFakeluaDouble(s, ReadDouble(in, pos)); - case TAG_STR_NEW: { - uint64_t len = ReadVarint(in, pos); - if (len > in.size() - pos) { - ThrowFakeluaException("serialize.decode: truncated string"); - } - std::string str(in, pos, static_cast(len)); - pos += static_cast(len); - uint32_t id = static_cast(state.dict.size()); - state.dict.push_back(str); - return inter::NativeToFakeluaString(s, state.dict[id]); - } - case TAG_STR_REF: { - uint64_t id = ReadVarint(in, pos); - if (id >= state.dict.size()) { - ThrowFakeluaException("serialize.decode: bad string ref id"); - } - return inter::NativeToFakeluaString(s, state.dict[id]); - } - case TAG_TABLE: { - if (state.depth >= 64) { - ThrowFakeluaException("serialize.decode: nesting too deep"); - } - uint64_t count = ReadVarint(in, pos); - // Each key/value pair is at least two tags. - if (count > (in.size() - pos) / 2) { - ThrowFakeluaException("serialize.decode: table too large"); - } - CVar tbl = table::TableHelper::CreateTable(s); - state.depth++; - try { - for (uint64_t i = 0; i < count; ++i) { - CVar key = DecodeValue(in, pos, state, s); - CVar val = DecodeValue(in, pos, state, s); - table::TableHelper::SetTable(s, tbl, key, val); - } - } catch (...) { - state.depth--; - throw; - } - state.depth--; - return tbl; - } - default: - ThrowFakeluaException("serialize.decode: unknown tag: " + std::to_string(tag)); - } -} - // 原生函数 static CVar SerializeEncode(State *s, CVar *args, int n) { CVar v = inter::GetNativeArg(s, args, n, 0); - if (!IsSupported(v)) { + if (!WireIsSerializable(v)) { LOG_ERROR(s, "serialize", "serialize.encode: unsupported type: {}", VarTypeToString(static_cast(v.type_))); ThrowFakeluaException("serialize.encode: unsupported type: " + VarTypeToString(static_cast(v.type_))); } - std::string out; - out.reserve(64); - EncodeState state; - EncodeValue(out, v, state); + std::string out = WireEncode(v); LOG_DEBUG(s, "serialize", "serialize.encode: bytes={}", out.size()); return inter::NativeToFakeluaString(s, out); } static CVar SerializeDecode(State *s, CVar *args, int n) { std::string in = CVarToString(inter::GetNativeArg(s, args, n, 0)); - size_t pos = 0; - DecodeState state; - CVar result = DecodeValue(in, pos, state, s); - if (pos != in.size()) { - LOG_ERROR(s, "serialize", "serialize.decode: trailing bytes (read={} total={})", pos, in.size()); - ThrowFakeluaException("serialize.decode: trailing bytes"); - } + CVar result = WireDecode(s, in); LOG_DEBUG(s, "serialize", "serialize.decode: bytes={}", in.size()); return result; } @@ -392,7 +120,7 @@ static SerNode CVarToSerNode(CVar v, std::unordered_set &visited, in n.kind = 5; if (t) { table::TableHelper::ForEachKV(v, [&](CVar k, CVar val) { - if (IsSupported(k) && IsSupported(val)) { + if (WireIsSerializable(k) && WireIsSerializable(val)) { n.keys.push_back(CVarToSerNode(k, visited, depth + 1, op)); n.vals.push_back(CVarToSerNode(val, visited, depth + 1, op)); } @@ -436,7 +164,7 @@ static CVar SerNodeToCVar(const SerNode &n, State *s, int depth) { static SerNode RequireSerNode(State *s, CVar *args, int n, const char *op) { CVar v = inter::GetNativeArg(s, args, n, 0); - if (!IsSupported(v)) { + if (!WireIsSerializable(v)) { ThrowFakeluaException(std::string(op) + ": unsupported type: " + VarTypeToString(static_cast(v.type_))); } std::unordered_set visited; diff --git a/src/native/serialize/wire_codec.cpp b/src/native/serialize/wire_codec.cpp new file mode 100644 index 0000000..6a8d36c --- /dev/null +++ b/src/native/serialize/wire_codec.cpp @@ -0,0 +1,381 @@ +#include "native/serialize/wire_codec.h" + +#include "native/native_common.h" +#include "native/table/native_table.h" +#include "var/var.h" +#include "var/var_string.h" +#include "var/var_table.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace fakelua::serialize { + +// Wire format(类 protobuf 编码)——与 native_serialize.cpp 文档保持一致: +// 每个值 = [type_tag(1 byte)] [payload] +// 0x00 nil +// 0x01 false +// 0x02 true +// 0x03 + varint 整数(zigzag 编码:小绝对值 → 小编码) +// 0x04 + 8 bytes double(小端 memcpy) +// 0x05 + varint(len) + bytes 新字符串,加入字典 +// 0x06 + varint(id) 字典中的字符串引用 +// 0x07 + varint(count) + N*(key,value) 表 + +namespace { + +enum Tag : uint8_t { + TAG_NIL = 0x00, + TAG_FALSE = 0x01, + TAG_TRUE = 0x02, + TAG_INT = 0x03, + TAG_DOUBLE = 0x04, + TAG_STR_NEW = 0x05, + TAG_STR_REF = 0x06, + TAG_TABLE = 0x07, +}; + +constexpr int kWireMaxDepth = 64; + +// 辅助:从 CVar 提取字符串(二进制安全) +std::string CVarToString(CVar v) { + if (v.type_ == static_cast(VarType::String) && v.data_.s) { + auto sv = v.data_.s->Str(); + return std::string(sv.data(), sv.size()); + } + if (v.type_ == static_cast(VarType::StringId) && v.data_.i) { + const char *ptr = reinterpret_cast(v.data_.i); + int sz = *reinterpret_cast(ptr); + return std::string(ptr + 8, sz); + } + return {}; +} + +std::string_view CVarToStringView(CVar v) { + if (v.type_ == static_cast(VarType::String) && v.data_.s) { + return v.data_.s->Str(); + } + if (v.type_ == static_cast(VarType::StringId) && v.data_.i) { + const char *ptr = reinterpret_cast(v.data_.i); + int sz = *reinterpret_cast(ptr); + return std::string_view(ptr + 8, sz); + } + return {}; +} + +bool IsSupported(CVar v) { + switch (v.type_) { + case static_cast(VarType::Nil): + case static_cast(VarType::Bool): + case static_cast(VarType::Int): + case static_cast(VarType::Float): + case static_cast(VarType::String): + case static_cast(VarType::StringId): + case static_cast(VarType::Table): + return true; + default: + return false; + } +} + +// Varint(LEB128 无符号) +void WriteVarint(std::string &out, uint64_t v) { + while (v >= 0x80) { + out.push_back(static_cast((v & 0x7f) | 0x80)); + v >>= 7; + } + out.push_back(static_cast(v)); +} + +uint64_t ReadVarint(std::string_view in, size_t &pos) { + uint64_t result = 0; + int shift = 0; + bool terminated = false; + while (pos < in.size()) { + uint8_t b = static_cast(in[pos++]); + result |= static_cast(b & 0x7f) << shift; + if ((b & 0x80) == 0) { + terminated = true; + break; + } + shift += 7; + if (shift >= 64) { + ThrowFakeluaException("wire.decode: varint too long"); + } + } + if (!terminated) { + ThrowFakeluaException("wire.decode: truncated varint"); + } + return result; +} + +// Zigzag(有符号整数 ↔ 无符号) +uint64_t ZigzagEncode(int64_t n) { + return (static_cast(n) << 1) ^ static_cast(n >> 63); +} + +int64_t ZigzagDecode(uint64_t u) { + return static_cast((u >> 1) ^ (-(u & 1))); +} + +void WriteDouble(std::string &out, double v) { + uint8_t buf[8]; + std::memcpy(buf, &v, 8); + out.append(reinterpret_cast(buf), 8); +} + +double ReadDouble(std::string_view in, size_t &pos) { + if (pos + 8 > in.size()) { + ThrowFakeluaException("wire.decode: truncated double"); + } + double v; + std::memcpy(&v, in.data() + pos, 8); + pos += 8; + return v; +} + +// 编码 +struct EncodeState { + std::unordered_map dict;// 字符串 → 字典 id + std::unordered_set visited; + int depth = 0; +}; + +void EncodeValue(std::string &out, CVar v, EncodeState &state) { + switch (v.type_) { + case static_cast(VarType::Nil): + out.push_back(TAG_NIL); + return; + case static_cast(VarType::Bool): + out.push_back(AsVar(v).GetBool() ? TAG_TRUE : TAG_FALSE); + return; + case static_cast(VarType::Int): + out.push_back(TAG_INT); + WriteVarint(out, ZigzagEncode(v.data_.i)); + return; + case static_cast(VarType::Float): + out.push_back(TAG_DOUBLE); + WriteDouble(out, v.data_.f); + return; + case static_cast(VarType::String): + case static_cast(VarType::StringId): { + auto sv = CVarToStringView(v); + auto it = state.dict.find(sv); + if (it != state.dict.end()) { + out.push_back(TAG_STR_REF); + WriteVarint(out, it->second); + } else { + uint32_t id = static_cast(state.dict.size()); + state.dict.emplace(sv, id); + out.push_back(TAG_STR_NEW); + WriteVarint(out, static_cast(sv.size())); + if (!sv.empty()) { + out.append(sv.data(), sv.size()); + } + } + return; + } + case static_cast(VarType::Table): { + VarTable *t = v.data_.t; + if (!t) { + out.push_back(TAG_TABLE); + WriteVarint(out, 0); + return; + } + if (state.depth >= kWireMaxDepth) { + ThrowFakeluaException("wire.encode: nesting too deep"); + } + if (!state.visited.insert(t).second) { + ThrowFakeluaException("wire.encode: cyclic table"); + } + state.depth++; + std::vector keys; + std::vector vals; + try { + table::TableHelper::ForEachKV(v, [&](CVar k, CVar val) { + if (IsSupported(k) && IsSupported(val)) { + keys.push_back(k); + vals.push_back(val); + } + }); + out.push_back(TAG_TABLE); + WriteVarint(out, static_cast(keys.size())); + for (size_t i = 0; i < keys.size(); ++i) { + EncodeValue(out, keys[i], state); + EncodeValue(out, vals[i], state); + } + } catch (...) { + state.depth--; + state.visited.erase(t); + throw; + } + state.depth--; + state.visited.erase(t); + return; + } + default: + ThrowFakeluaException("wire.encode: unsupported type: " + VarTypeToString(static_cast(v.type_))); + } +} + +// 解码 +struct DecodeState { + std::vector dict;// id → 字符串 + int depth = 0; +}; + +CVar DecodeValue(std::string_view in, size_t &pos, DecodeState &state, State *s) { + if (pos >= in.size()) { + ThrowFakeluaException("wire.decode: unexpected end of input"); + } + uint8_t tag = static_cast(in[pos++]); + switch (tag) { + case TAG_NIL: + return inter::NativeToFakeluaNil(s); + case TAG_FALSE: + return inter::NativeToFakeluaBool(s, false); + case TAG_TRUE: + return inter::NativeToFakeluaBool(s, true); + case TAG_INT: { + uint64_t u = ReadVarint(in, pos); + return inter::NativeToFakeluaLonglong(s, ZigzagDecode(u)); + } + case TAG_DOUBLE: + return inter::NativeToFakeluaDouble(s, ReadDouble(in, pos)); + case TAG_STR_NEW: { + uint64_t len = ReadVarint(in, pos); + if (len > in.size() - pos) { + ThrowFakeluaException("wire.decode: truncated string"); + } + std::string str(in.substr(pos, static_cast(len))); + pos += static_cast(len); + uint32_t id = static_cast(state.dict.size()); + state.dict.push_back(std::move(str)); + return inter::NativeToFakeluaString(s, state.dict[id]); + } + case TAG_STR_REF: { + uint64_t id = ReadVarint(in, pos); + if (id >= state.dict.size()) { + ThrowFakeluaException("wire.decode: bad string ref id"); + } + return inter::NativeToFakeluaString(s, state.dict[id]); + } + case TAG_TABLE: { + if (state.depth >= kWireMaxDepth) { + ThrowFakeluaException("wire.decode: nesting too deep"); + } + uint64_t count = ReadVarint(in, pos); + // Each key/value pair is at least two tags. + if (count > (in.size() - pos) / 2) { + ThrowFakeluaException("wire.decode: table too large"); + } + CVar tbl = table::TableHelper::CreateTable(s); + state.depth++; + try { + for (uint64_t i = 0; i < count; ++i) { + CVar key = DecodeValue(in, pos, state, s); + CVar val = DecodeValue(in, pos, state, s); + table::TableHelper::SetTable(s, tbl, key, val); + } + } catch (...) { + state.depth--; + throw; + } + state.depth--; + return tbl; + } + default: + ThrowFakeluaException("wire.decode: unknown tag: " + std::to_string(tag)); + } +} + +// 严格校验:回调绑定参数不允许静默丢字段。任意位置出现不可序列化类型都响亮报错。 +void ValidateStrict(CVar v, int depth, std::unordered_set &visited, const char *fname, int argno) { + if (!IsSupported(v)) { + ThrowBadArgument(argno, fname, + "callback args must be nil, boolean, number, string or a table thereof (closures and native objects are not allowed)"); + } + if (v.type_ != static_cast(VarType::Table) || !v.data_.t) return; + VarTable *t = v.data_.t; + if (depth >= kWireMaxDepth) { + ThrowFakeluaException("callback args: nesting too deep"); + } + if (!visited.insert(t).second) { + ThrowFakeluaException("callback args: cyclic table"); + } + table::TableHelper::ForEachKV(v, [&](CVar k, CVar val) { + ValidateStrict(k, depth + 1, visited, fname, argno); + ValidateStrict(val, depth + 1, visited, fname, argno); + }); + visited.erase(t); +} + +}// namespace + +bool WireIsSerializable(CVar v) { + return IsSupported(v); +} + +std::string WireEncode(CVar v) { + std::string out; + out.reserve(64); + EncodeState state; + EncodeValue(out, v, state); + return out; +} + +CVar WireDecode(State *s, std::string_view in) { + size_t pos = 0; + DecodeState state; + CVar result = DecodeValue(in, pos, state, s); + if (pos != in.size()) { + ThrowFakeluaException("wire.decode: trailing bytes"); + } + return result; +} + +std::string WireEncodeCallbackArgs(State * /*s*/, CVar *args, int first, int n, const char *fname, int argno_base) { + std::string out; + out.reserve(32); + const int count = n > first ? n - first : 0; + WriteVarint(out, static_cast(count)); + if (count == 0) return out; + + // 先整体严格校验,再编码:不允许走到编码器的"跳过不支持字段"静默分支。 + std::unordered_set visited; + for (int i = 0; i < count; ++i) { + ValidateStrict(args[first + i], 0, visited, fname, argno_base + i); + } + + EncodeState state; + for (int i = 0; i < count; ++i) { + EncodeValue(out, args[first + i], state); + } + return out; +} + +DecodedArgs WireDecodeCallbackArgs(State *s, std::string_view blob) { + DecodedArgs out; + if (blob.empty()) return out; + size_t pos = 0; + uint64_t count = ReadVarint(blob, pos); + out.vars.reserve(static_cast(count)); + DecodeState state; + for (uint64_t i = 0; i < count; ++i) { + out.vars.push_back(DecodeValue(blob, pos, state, s)); + } + if (pos != blob.size()) { + ThrowFakeluaException("wire.decode: trailing bytes in callback args"); + } + // 解码出的字符串可能引用 state.dict 里 std::string 的内容,把后备内存随 DecodedArgs + // 一起带走,调用方持有 DecodedArgs 到回调返回即可保证有效。 + out.str_dict = std::move(state.dict); + return out; +} + +}// namespace fakelua::serialize diff --git a/src/native/serialize/wire_codec.h b/src/native/serialize/wire_codec.h new file mode 100644 index 0000000..c9a690e --- /dev/null +++ b/src/native/serialize/wire_codec.h @@ -0,0 +1,46 @@ +#pragma once + +// wire_codec.h — CVar 紧凑 wire 编解码(内部模块,非 Lua API)。 +// +// 单值格式与 Lua serialize.encode/decode 完全一致(见 native_serialize.cpp), +// 供两类使用者复用: +// 1. serialize 原生库(Lua API 层); +// 2. 异步原生回调(mysql):把"函数名 + 绑定参数"中的数据参数序列化成自有的 +// 字节串暂存,跨多次顶层 Call / State::Reset() 保存,结果回来时在派发帧内 +// decode 回临时 arena。字节串是普通堆内存,回调结束即释放,不占用 const arena。 + +#include "fakelua.h" + +#include +#include +#include + +namespace fakelua::serialize { + +// 值是否可被 wire 编码:nil/bool/int/float/string/stringid/table,其余(闭包/native 对象等)不行。 +bool WireIsSerializable(CVar v); + +// 单个 CVar → wire 字节。表内遇到不可序列化字段会跳过(与 serialize.encode 历史行为一致); +// 顶层值不可序列化时抛错。 +std::string WireEncode(CVar v); + +// wire 字节 → 单个 CVar(分配在 State 临时 arena 上,仅当前帧有效)。 +// 字节必须被完整消费,尾随字节抛错。 +CVar WireDecode(State *s, std::string_view in); + +// 回调绑定参数元组:[varint count][value]... +// 与 WireEncode 不同,这里对参数做【严格】校验:闭包/native 对象等不可序列化类型 +// 出现在任意嵌套位置(含表的键值)都抛 bad argument,不允许静默丢字段。 +// fname/argno_base 用于错误信息(bad argument #argno_base+i to 'fname')。 +std::string WireEncodeCallbackArgs(State *s, CVar *args, int first, int n, const char *fname, int argno_base); + +// 绑定参数元组解码。产物在临时 arena 上;DecodedArgs 同时持有字符串字典等后备内存, +// 调用方必须让它活到回调函数返回之后。 +struct DecodedArgs { + std::vector vars; + std::vector str_dict; +}; + +DecodedArgs WireDecodeCallbackArgs(State *s, std::string_view blob); + +}// namespace fakelua::serialize diff --git a/src/state/fakelua.cpp b/src/state/fakelua.cpp index 4584725..83ea662 100644 --- a/src/state/fakelua.cpp +++ b/src/state/fakelua.cpp @@ -503,6 +503,9 @@ void ThrowIfMultiCVar(const CVar &v) { static CVar DispatchCallRaw(void *addr, const CVar *arg_arr, int arg_count, VarClosure *cl); CVar DispatchCall(State *s, void *addr, const CVar *arg_arr, int arg_count, JITType type, VarClosure *cl) { + // 进入目标引擎代码:切换当前引擎上下文,使本次执行过程中调起的 native 函数 + // (含 native 对象方法桥接)能登记派回本引擎的异步回调。嵌套调用逐层保存恢复。 + State::JitContextScope jit_scope(s, type); // 解释器闭包用 code_str magic 识别,不能看函数指针最低位:GCC 的函数指针可以是奇数。 // type==JIT_INTERP 且 cl 为空:Call()/CallByName 直接调注册的解释器原型。 // cl 是 TCC/GCC 闭包时即使 type 是 JIT_INTERP 也必须走 C 函数指针(解释器回调 JIT 闭包)。 diff --git a/src/state/heap.cpp b/src/state/heap.cpp index 2bcf95f..375651b 100644 --- a/src/state/heap.cpp +++ b/src/state/heap.cpp @@ -70,4 +70,14 @@ size_t HeapAllocator::Size() const { return current_block_index_ * BLOCK_SIZE + current_block_offset_; } +bool HeapAllocator::Contains(const void *p) const { + if (!p) return false; + const auto *addr = static_cast(p); + for (const void *block: blocks_) { + const auto *base = static_cast(block); + if (addr >= base && addr < base + BLOCK_SIZE) return true; + } + return false; +} + }// namespace fakelua diff --git a/src/state/heap.h b/src/state/heap.h index 8515d16..76b7b9b 100644 --- a/src/state/heap.h +++ b/src/state/heap.h @@ -29,6 +29,9 @@ class HeapAllocator { // 当前临时内存使用 [[nodiscard]] size_t Size() const; + // p 是否落在本分配器已经切出的块里。const arena 不 Reset,用来识别已经钉住的对象。 + [[nodiscard]] bool Contains(const void *p) const; + private: struct DestructorInfo { void (*destroyer)(void *); @@ -58,6 +61,10 @@ class Heap { // const_allocator_ 不重置,常量内存一直保留 } + [[nodiscard]] bool OwnsConst(const void *p) const { + return const_allocator_.Contains(p); + } + private: HeapAllocator temp_allocator_; // 临时内存分配器,编译过程中使用,编译结束后重置 HeapAllocator const_allocator_;// 常量内存分配器,编译过程中使用,编译结束后不重置 diff --git a/src/state/state.h b/src/state/state.h index 0ddf487..f2c409c 100644 --- a/src/state/state.h +++ b/src/state/state.h @@ -129,6 +129,37 @@ class State { bool prev_; }; + // 当前正在执行的脚本引擎(TCC/GCC/解释器)。 + // 每次 C++ 原生函数被 Lua 调起(FakeluaCallByName/CallByNameArgs 的汇聚点)都会压入 + // 调用方引擎,嵌套调用(JIT → native → Lua 回调 → native)逐层保存恢复。 + // 用途:异步回调(mysql 等)在登记时记住"是从哪个引擎发起的",结果回来后派回同一引擎, + // 与旧的内联闭包自带 func_ptr 的语义一致;不能固定派发到 TCC。 + // 栈空(C++ 宿主不经 Lua 直接调原生内部 API)时返回 JIT_TCC,与历史行为一致。 + [[nodiscard]] JITType CurrentJit() const { + return current_jit_; + } + + // C 编译单元(TCC 生成的 C 代码)不能用 C++ 的 JitContextScope,由 extern "C" + // 的 FakeluaJitContextPush/Pop 经这对 getter/setter 维护当前引擎。 + void SetCurrentJit(JITType jit) { + current_jit_ = jit; + } + + struct JitContextScope { + explicit JitContextScope(State *s, JITType jit) : s_(s), prev_(s->current_jit_) { + s_->current_jit_ = jit; + } + ~JitContextScope() { + s_->current_jit_ = prev_; + } + JitContextScope(const JitContextScope &) = delete; + JitContextScope &operator=(const JitContextScope &) = delete; + + private: + State *s_; + JITType prev_; + }; + // 本 State 的日志输出目标。为 nullptr 表示没指定日志文件,只打控制台。 LogSink *GetLogSink() const { return log_sink_.get(); @@ -184,6 +215,8 @@ class State { std::function var_interface_new_func_; JitErrorBoundary *jit_error_boundary_ = nullptr; bool interp_const_alloc_ = false; + // 当前执行引擎,初值 TCC(无 Lua 调用栈的边界场景,与 kNativeCallbackJit 一致)。 + JITType current_jit_ = JIT_TCC; int reentrant_count_ = 0; StateConfig config_; Compiler compiler_; diff --git a/src/var/var_table.h b/src/var/var_table.h index fb9e8c6..0a23d18 100644 --- a/src/var/var_table.h +++ b/src/var/var_table.h @@ -40,6 +40,11 @@ struct VarTable { // 需要重算——因此把 VarTable 整体清零的分配路径天然落在安全的重算分支上。 uint32_t seq_len_valid_; int64_t seq_len_; + // spec 块字节数。0 表示没有需要搬迁的 spec(spec 指针为空)。 + // 与 c_runtime_header.h 的 VarTable 保持同布局。 + uint32_t spec_bytes; + // >0 时 spec 块是连续 spec_cvars 个 CVar(JIT 表特化)。0 表示不透明块(如 NativeObjectSpec)。 + uint32_t spec_cvars; }; }// namespace fakelua diff --git a/test/lua/infer/test_global_const_reassign_dynamic.lua b/test/lua/infer/test_global_const_reassign_dynamic.lua new file mode 100644 index 0000000..73f0ea9 --- /dev/null +++ b/test/lua/infer/test_global_const_reassign_dynamic.lua @@ -0,0 +1,7 @@ +-- 右值是运行时值也同样禁止:常量与否不取决于赋值能不能在编译期算出来。 +local map_width = 2000 + +function init(w) + map_width = w + return map_width +end diff --git a/test/lua/infer/test_global_const_reassigned.lua b/test/lua/infer/test_global_const_reassigned.lua new file mode 100644 index 0000000..a84f2f2 --- /dev/null +++ b/test/lua/infer/test_global_const_reassigned.lua @@ -0,0 +1,9 @@ +-- 有初值的文件级数值 local 是常量。函数里再赋值必须在编译期报错, +-- 并带上 Lua 文件位置,而不是等生成的 C 代码因为 const 赋值才失败。 +-- local x = func() 不在此列:声明会降成 nil,只由 __fakelua_init 赋值一次。 +local next_bot_id = 0 + +function init() + next_bot_id = 800000 + return next_bot_id +end diff --git a/test/lua/json/test_json_edge.lua b/test/lua/json/test_json_edge.lua index b3eef44..807d1c0 100644 --- a/test/lua/json/test_json_edge.lua +++ b/test/lua/json/test_json_edge.lua @@ -249,11 +249,38 @@ function test_encode_control_chars() return 1 end --- 测试 JSON 编码:空对象 +-- 测试 JSON 编码:空 table function test_encode_empty_object() local s = json.encode({}) - -- 空表应编码为空对象 - if s ~= "{}" then return 0 end + -- 空 table 按纯数组启发式编码为 [](客户端按数组解析协议字段) + if s ~= "[]" then return 0 end + return 1 +end + +-- 测试 JSON 编码:空数组字段嵌在对象里 +function test_encode_empty_array_field() + local s = json.encode({ type = "rank", list = {} }) + -- 空 list 字段应为 "list":[] 而不是 "list":{}(string.find 为 ECMAScript 正则,[ 需转义) + if not string.find(s, '"list":\\[\\]') then return 0 end + return 1 +end + +-- 测试 json.encode_array:数组形 table 正常编码 +function test_encode_array_ok() + if json.encode_array({}) ~= "[]" then return 0 end + if json.encode_array({ 1, 2, 3 }) ~= "[1,2,3]" then return 0 end + if json.encode_array({ [1] = "a", [2] = "b" }) ~= '["a","b"]' then return 0 end + return 1 +end + +-- 测试 json.encode_array:非数组形 table 报错 +function test_encode_array_reject() + local ok1 = pcall(function() json.encode_array({ a = 1 }) end) + if ok1 then return 0 end + local ok2 = pcall(function() json.encode_array({ [1] = "a", [3] = "c" }) end) + if ok2 then return 0 end + local ok3 = pcall(function() json.encode_array("not a table") end) + if ok3 then return 0 end return 1 end diff --git a/test/lua/mysql/test_mysql_contract.lua b/test/lua/mysql/test_mysql_contract.lua new file mode 100644 index 0000000..82302fe --- /dev/null +++ b/test/lua/mysql/test_mysql_contract.lua @@ -0,0 +1,172 @@ +package "MysqlContractTest" + +-- MySQL 回调契约测试(无需真实 MySQL 服务器:连到死端口验证回调行为) +-- 异步回调模型:函数名 + 绑定参数(纯数据),不支持内联闭包: +-- conn:query(sql, "on_result", arg1, arg2, ...) +-- 绑定参数在登记当下被序列化成独立字节串暂存,跨任意次顶层 Call 的 arena Reset +-- 都不丢;结果派发时在 tick 帧内反序列化,追加在固定参数 (conn, err, result) 之后。 +-- 回调结束字节串即释放,不占用 const arena。 +-- 跨帧可变状态挂在 conn 对象字段上(连接是 C++ 持有的 native 对象,Reset 不回收)。 +-- 回归: +-- P0-1 非法回调(闭包/数字/空串)曾被静默转成空名、回调永不触发,现在必须响亮报错 +-- P0-2 连接未就绪时 query 的回调曾被静默吞掉(Query 提前 return 不设 pending_result_) +-- P1-4 同连接飞行中再 query 曾直接报 "connection not ready";现在排队,回调恰好一次 + +-- 死端口配置:连接必然失败,正好驱动错误终态路径 +local dead_cfg = { + host = "127.0.0.1", + port = 1, + user = "root", + password = "x", + db = "test", + timeout_ms = 500 +} + +-- 观测状态全部挂在 conn 字段上;绑定参数用 tag 区分是哪一次登记。 +function contract_on_connect(c, err, success, tag) + if tag == "connect_tag" then + c.cc = (c.cc or 0) + 1 + end +end + +function contract_on_query(c, err, result, tag, n, ctx) + if tag == "query_tag" then + c.qc = (c.qc or 0) + 1 + c.qerr = err + c.qtag = tag + c.qn = n + c.qctx = ctx + end +end + +-- 场景 1:连接未就绪时发起 query —— 连接失败后回调恰好收到一次错误; +-- 绑定的字符串/数字/表参数跨帧存活并按快照原样送达。 +function test_query_closure_exactly_once() + local conn = mysql.connect(dead_cfg, "contract_on_connect", "connect_tag") + + -- 连接尚未就绪:query 排队,连接进入错误终态后回调收到错误 + conn:query("SELECT 1", "contract_on_query", "query_tag", 7, { k = "v" }) + + for i = 1, 2000 do + runtime.tick() + if conn.qc then break end + os.sleep(1) + end + + if conn.cc ~= 1 then + print("connect callback count = ", conn.cc) + return 0 + end + if conn.qc ~= 1 then + print("query callback count = ", conn.qc) + return 0 + end + if type(conn.qerr) ~= "string" or #conn.qerr == 0 then + print("query callback missing error, got:", conn.qerr) + return 0 + end + if conn.qtag ~= "query_tag" then + print("bound string lost, got:", conn.qtag) + return 0 + end + if conn.qn ~= 7 then + print("bound number lost, got:", conn.qn) + return 0 + end + if type(conn.qctx) ~= "table" or conn.qctx.k ~= "v" then + print("bound table lost, got:", type(conn.qctx)) + return 0 + end + return 1 +end + +-- 场景 2:回调名传非字符串(含内联闭包)、绑定参数含闭包 —— 都必须响亮报错 +function test_bad_callback_type() + local conn = mysql.connect(dead_cfg, "contract_on_connect", "connect_tag") + + local ok1 = pcall(function() conn:query("SELECT 1", 123) end) + if ok1 then return 0 end + local ok2 = pcall(function() conn:query("SELECT 1", nil) end) + if ok2 then return 0 end + local ok3 = pcall(function() conn:query("SELECT 1", {}) end) + if ok3 then return 0 end + -- 内联闭包不再支持:异步回调无法安全持有临时 arena 上的闭包 + local ok4 = pcall(function() conn:query("SELECT 1", function() end) end) + if ok4 then return 0 end + -- 回调名合法,但绑定参数里夹带闭包:同样响亮报错,不能静默丢字段 + local ok5 = pcall(function() conn:query("SELECT 1", "contract_on_query", "x", function() end) end) + if ok5 then return 0 end + -- 连接回调同理 + local ok6 = pcall(function() mysql.connect(dead_cfg, function() end) end) + if ok6 then return 0 end + + conn:close() + return 1 +end + +-- 场景 3:pool:with —— 无健康连接时返回 nil(fn 不执行);非函数参数报错。 +-- pool:with 的函数是【同帧同步调用】、不跨 tick 暂存,因此仍支持内联闭包。 +function test_pool_with() + local pool = mysql_pool.create({ + host = "127.0.0.1", + port = 1, + user = "root", + password = "x", + db = "test", + pool_size = 1, + timeout_ms = 300 + }) + + -- 没有可用连接:with 返回 nil + local r = pool:with(function(c) + return 42 + end) + if r ~= nil then return 0 end + + -- 非函数参数必须报错 + local ok = pcall(function() pool:with(123) end) + if ok then return 0 end + + pool:close() + return 1 +end + +-- 给 C++ 单测用:返回一个会 error() 的闭包,用来验证 pool:with 抛错后仍归还连接。 +function make_thrower() + return function(c) + error("intentional failure") + end +end + +-- 跨帧绑定参数(P0-1 的新形态):绑定参数序列化后必须活过下一次顶层 Call 的 +-- arena Reset。arm 在一次 Call 里登记回调后返回;宿主显式 Reset 临时 arena; +-- pump 是下一次 Call,再 tick 到回调。若字节串暂存失效,ctx 字段会坏掉或回调跑不起来。 +-- 回调派回登记时的引擎:三引擎各自 arm/pump 闭环,文件级 local 各自独立也没问题。 +-- cross_hits 只承载动态值(表字段 / nil),不会被推断成文件级数值常量。 +local cross_hits = nil + +function contract_on_cross(c, err, result, ctx) + if ctx.tag == "ok" and ctx.n == 7 and type(err) == "string" then + cross_hits = ctx.n + else + cross_hits = -1 + end +end + +function arm_cross_frame() + cross_hits = nil + local conn = mysql.connect(dead_cfg, "contract_on_connect", "connect_tag") + + conn:query("SELECT 1", "contract_on_cross", { tag = "ok", n = 7 }) + return 1 +end + +function pump_cross_frame() + for i = 1, 2000 do + runtime.tick() + if cross_hits ~= nil then break end + os.sleep(1) + end + if cross_hits == nil then return 0 end + return cross_hits +end diff --git a/test/lua/mysql/test_mysql_datatypes.lua b/test/lua/mysql/test_mysql_datatypes.lua index 7892fc8..30ca1ea 100644 --- a/test/lua/mysql/test_mysql_datatypes.lua +++ b/test/lua/mysql/test_mysql_datatypes.lua @@ -113,9 +113,9 @@ function test_datatypes() local row = result[3][1] - -- 验证 id 为 "1" - if row[1] ~= "1" then - print("id mismatch, expected '1', got:", tostring(row[1])) + -- 验证 id 为整数 1(INT 列按列类型返回 number) + if row[1] ~= 1 then + print("id mismatch, expected 1, got:", tostring(row[1])) conn:close() return 0 end diff --git a/test/lua/mysql/test_mysql_lifecycle.lua b/test/lua/mysql/test_mysql_lifecycle.lua index 35209e5..d8ce794 100644 --- a/test/lua/mysql/test_mysql_lifecycle.lua +++ b/test/lua/mysql/test_mysql_lifecycle.lua @@ -67,18 +67,32 @@ function test_lifecycle() conn:close() conn:close() - -- 3. 验证对已关闭连接调用 query 能够受控捕获错误 - local ok, err_msg = pcall(function() + -- 3. 验证对已关闭连接调用 query:不抛错,而是在后续 tick 由回调收到 "closed" 错误 + -- (契约:conn:query 的回调恰好触发一次,连接已关闭时回调收到错误字符串) + conn.query_done = false + conn.query_err = nil + local ok, call_err = pcall(function() conn:query("SELECT 1", "on_result") end) - if ok then - print("expected error when querying closed connection, but pcall succeeded") + if not ok then + print("query on closed connection should notify via callback, not throw:", tostring(call_err)) return 0 end - if not string.find(tostring(err_msg), "closed") then - print("error message should mention 'closed', got:", tostring(err_msg)) + for i = 1, 100 do + runtime.tick() + if conn.query_done then break end + os.sleep(1) + end + + if not conn.query_done then + print("closed-connection query callback never fired") + return 0 + end + + if conn.query_err == nil or not string.find(tostring(conn.query_err), "closed") then + print("error message should mention 'closed', got:", tostring(conn.query_err)) return 0 end diff --git a/test/lua/mysql/test_mysql_multi_result.lua b/test/lua/mysql/test_mysql_multi_result.lua index ac832ab..54db96b 100644 --- a/test/lua/mysql/test_mysql_multi_result.lua +++ b/test/lua/mysql/test_mysql_multi_result.lua @@ -90,9 +90,9 @@ function test_multi_result() conn:close() return 0 end - local expected = tostring(i) - if result[3][1][1] ~= expected then - print("result", i, "value mismatch:", result[3][1][1], "expected:", expected) + -- SELECT 字面整数按列类型返回 number(整数) + if result[3][1][1] ~= i then + print("result", i, "value mismatch:", result[3][1][1], "expected:", i) conn:close() return 0 end diff --git a/test/lua/mysql/test_mysql_pool.lua b/test/lua/mysql/test_mysql_pool.lua index 7df1dd4..f665cc7 100644 --- a/test/lua/mysql/test_mysql_pool.lua +++ b/test/lua/mysql/test_mysql_pool.lua @@ -65,7 +65,7 @@ function test_pool() pool:close() return 0 end - if result[3][1][1] ~= "1" then + if result[3][1][1] ~= 1 then print("pool query result mismatch:", result[3][1][1]) pool:release(conn) pool:close() diff --git a/test/lua/mysql/test_mysql_query_error.lua b/test/lua/mysql/test_mysql_query_error.lua index 31742bd..e7c1f38 100644 --- a/test/lua/mysql/test_mysql_query_error.lua +++ b/test/lua/mysql/test_mysql_query_error.lua @@ -108,7 +108,7 @@ function test_query_error() end local res = conn.query_result - if not res or res[1] ~= true or res[3][1][1] ~= "42" then + if not res or res[1] ~= true or res[3][1][1] ~= 42 then print("recovery query result mismatch") conn:close() return 0 diff --git a/test/lua/mysql/test_mysql_stmt.lua b/test/lua/mysql/test_mysql_stmt.lua index ca433b0..9fc0fa1 100644 --- a/test/lua/mysql/test_mysql_stmt.lua +++ b/test/lua/mysql/test_mysql_stmt.lua @@ -186,7 +186,8 @@ function test_stmt() end local row = result[3][1] - if row[1] ~= "1" or row[2] ~= "alice" then + -- INT 列按列类型返回 number(整数),VARCHAR 列返回 string + if row[1] ~= 1 or row[2] ~= "alice" then print("row mismatch:", row[1], row[2]) conn:stmt_close(conn.stmt_id) conn:close() diff --git a/test/lua/mysql/test_mysql_stmt_params.lua b/test/lua/mysql/test_mysql_stmt_params.lua index 67d3cb6..d5003a5 100644 --- a/test/lua/mysql/test_mysql_stmt_params.lua +++ b/test/lua/mysql/test_mysql_stmt_params.lua @@ -167,8 +167,8 @@ function test_stmt_params() end local r1 = res1[3][1] - if r1[1] ~= "1" then - print("r1 id expected '1', got:", tostring(r1[1])) + if r1[1] ~= 1 then + print("r1 id expected 1, got:", tostring(r1[1])) conn:stmt_close(select_stmt_id) conn:close() return 0 @@ -205,7 +205,7 @@ function test_stmt_params() end local r2 = res2[3][1] - if r2[1] ~= "2" or r2[2] ~= "bob" then + if r2[1] ~= 2 or r2[2] ~= "bob" then print("r2 values mismatch:", tostring(r2[1]), tostring(r2[2])) conn:stmt_close(select_stmt_id) conn:close() diff --git a/test/lua/net/test_net_server_client.lua b/test/lua/net/test_net_server_client.lua index fe7bd6a..b2709ac 100644 --- a/test/lua/net/test_net_server_client.lua +++ b/test/lua/net/test_net_server_client.lua @@ -140,3 +140,45 @@ function test_send_buffer_full() cli:close() return ok and 0 or 1 end + +-- 回归(P1-6):在 on_event 回调里直接调用 obj:send()(不走返回 "echo" 的路径)。 +-- 派发上下文中的 send 会先入队,由本轮 tick 派发完后统一泵出;客户端随后收到数据。 +local send_in_cb_server = nil + +function on_send_in_cb(type, connid, data, len, reason) + if type == "recv" then + if send_in_cb_server then + send_in_cb_server:send(connid, "direct:" .. data) + end + end +end + +function test_send_in_callback() + local server = net.server({ port = 19987, maxconn = 4 }) + server:dispatch("NetTest.on_send_in_cb") + send_in_cb_server = server + + local client = net.client({ port = 19987 }) + client:dispatch("NetTest.on_client_event") + + for i = 1, 50 do + runtime.tick() + if server:get_conn_count() >= 1 then break end + os.sleep(1) + end + + client:send("ping") + + -- 回调里直接 send 的数据应经入队泵出后到达客户端 + for i = 1, 50 do + runtime.tick() + if client:get_last_data() == "direct:ping" then break end + os.sleep(1) + end + + local ok = client:get_last_data() == "direct:ping" + send_in_cb_server = nil + server:close() + client:close() + return ok and 1 or 0 +end diff --git a/test/test_exception.cpp b/test/test_exception.cpp index e65d6a0..307e62d 100644 --- a/test/test_exception.cpp +++ b/test/test_exception.cpp @@ -1125,12 +1125,20 @@ TEST(exception, const_no_init) { EXPECT_THROW(CompileFile(s, "./exception/test_const_no_init.lua", {}), std::exception); } +// 文件级数值字面量是常量,函数里再赋值是编译期错误。 TEST(exception, const_reassign) { FakeluaStateGuard sg; auto s = sg.GetState(); ASSERT_NE(s, nullptr); SetDebugLogLevel(s, 0); - EXPECT_THROW(CompileFile(s, "./exception/test_const_reassign.lua", {}), std::exception); + try { + CompileFile(s, "./exception/test_const_reassign.lua", {}); + FAIL() << "reassigning file-level constant a should fail"; + } catch (const std::exception &e) { + const std::string msg = e.what(); + EXPECT_NE(msg.find("cannot reassign file-level constant 'a'"), std::string::npos); + EXPECT_NE(msg.find("test_const_reassign.lua"), std::string::npos); + } } TEST(exception, top_level_bare_local) { diff --git a/test/test_infer.cpp b/test/test_infer.cpp index c99920b..6725194 100644 --- a/test/test_infer.cpp +++ b/test/test_infer.cpp @@ -3023,6 +3023,39 @@ TEST(infer, test_global_const_float) { }); } +// 文件级数值字面量是常量。函数里再赋值必须在编译期报错,并带上 Lua 位置。 +// local x = func() 不走这条路径:C 静态初值不能调用函数,预处理改成 local x = nil, +// 只在 __fakelua_init 里赋值一次。见 test_global_init_multi_names_funcall。 +TEST(infer, test_global_const_reassigned_not_const) { + const auto s = FakeluaNewState(); + ASSERT_NE(s, nullptr); + try { + CompileFile(s, "./infer/test_global_const_reassigned.lua", {}); + FakeluaDeleteState(s); + FAIL() << "reassigning a file-level numeric constant should fail at compile time"; + } catch (const std::exception &e) { + const std::string msg = e.what(); + EXPECT_NE(msg.find("cannot reassign file-level constant 'next_bot_id'"), std::string::npos); + EXPECT_NE(msg.find("test_global_const_reassigned.lua"), std::string::npos); + FakeluaDeleteState(s); + } +} + +TEST(infer, test_global_const_reassign_dynamic) { + const auto s = FakeluaNewState(); + ASSERT_NE(s, nullptr); + try { + CompileFile(s, "./infer/test_global_const_reassign_dynamic.lua", {}); + FakeluaDeleteState(s); + FAIL() << "reassigning a file-level numeric constant with a runtime value should fail"; + } catch (const std::exception &e) { + const std::string msg = e.what(); + EXPECT_NE(msg.find("cannot reassign file-level constant 'map_width'"), std::string::npos); + EXPECT_NE(msg.find("test_global_const_reassign_dynamic.lua"), std::string::npos); + FakeluaDeleteState(s); + } +} + // --------------------------------------------------------------------------- // func() + func() and func(func()) — return expression contains function-call // results used in arithmetic or as arguments to another call. These patterns diff --git a/test/test_json.cpp b/test/test_json.cpp index 58ea29a..c0cbaf0 100644 --- a/test/test_json.cpp +++ b/test/test_json.cpp @@ -628,6 +628,32 @@ TEST(test_json, encode_empty_object) { FakeluaDeleteState(s); } +// 空 list 字段在对象里应编码为 "list":[](回归:曾产出 "list":{} 导致客户端按数组解析崩溃) +TEST(test_json, encode_empty_array_field) { + State *s = FakeluaNewState(); + ASSERT_NE(s, nullptr); + CompileConfig config; + CompileFile(s, "./json/test_json_edge.lua", config); + int64_t ret = 0; + CallAll(s, "JsonTest.test_encode_empty_array_field", ret); + EXPECT_EQ(ret, 1); + FakeluaDeleteState(s); +} + +// json.encode_array:数组形 table 正常编码,非数组形报错 +TEST(test_json, encode_array_strict) { + State *s = FakeluaNewState(); + ASSERT_NE(s, nullptr); + CompileConfig config; + CompileFile(s, "./json/test_json_edge.lua", config); + int64_t ret = 0; + CallAll(s, "JsonTest.test_encode_array_ok", ret); + EXPECT_EQ(ret, 1); + CallAll(s, "JsonTest.test_encode_array_reject", ret); + EXPECT_EQ(ret, 1); + FakeluaDeleteState(s); +} + TEST(test_json, encode_nested_too_deep) { State *s = FakeluaNewState(); ASSERT_NE(s, nullptr); diff --git a/test/test_mysql.cpp b/test/test_mysql.cpp index 53c8693..674b881 100644 --- a/test/test_mysql.cpp +++ b/test/test_mysql.cpp @@ -1,8 +1,13 @@ #include "fakelua.h" +#include "native/mysql/mysql_field_convert.h" +#include "native/mysql/native_mysql.h" +#include "native/native_common.h" #include "test_jit.h" #include "gtest/gtest.h" using namespace fakelua; +using namespace fakelua::mysql; + // MySQL 模块测试 @@ -182,4 +187,182 @@ TEST(test_mysql, bad_port) { CallAll(s, "MysqlTest.test_bad_port", ret); EXPECT_EQ(ret, 1); FakeluaDeleteState(s); -} \ No newline at end of file +} +// 回归(P0-1/P0-2/P1-4):连接未就绪时用闭包发起 query, +// 连接失败后回调应恰好收到一次错误(曾被静默吞掉或报错丢失)。 +TEST(test_mysql, query_closure_callback_exactly_once) { + State *s = FakeluaNewState(); + ASSERT_NE(s, nullptr); + CompileConfig config; + CompileFile(s, "./mysql/test_mysql_contract.lua", config); + int64_t ret = 0; + CallAll(s, "MysqlContractTest.test_query_closure_exactly_once", ret); + EXPECT_EQ(ret, 1); + FakeluaDeleteState(s); +} + +// 回归(P0-1):回调参数传非法类型必须响亮报错,而不是静默丢弃回调。 +TEST(test_mysql, bad_callback_type_throws) { + State *s = FakeluaNewState(); + ASSERT_NE(s, nullptr); + CompileConfig config; + CompileFile(s, "./mysql/test_mysql_contract.lua", config); + int64_t ret = 0; + CallAll(s, "MysqlContractTest.test_bad_callback_type", ret); + EXPECT_EQ(ret, 1); + FakeluaDeleteState(s); +} + +// 回归(P1-4):pool:with 基础行为 —— 无可用连接返回 nil,非函数参数报错。 +TEST(test_mysql, pool_with_lease_api) { + State *s = FakeluaNewState(); + ASSERT_NE(s, nullptr); + CompileConfig config; + CompileFile(s, "./mysql/test_mysql_contract.lua", config); + int64_t ret = 0; + CallAll(s, "MysqlContractTest.test_pool_with", ret); + EXPECT_EQ(ret, 1); + FakeluaDeleteState(s); +} + +// 回归(P1-4):pool:with 的 fn 抛错后连接必须归还,否则下一次 Acquire 拿不到。 +// 不依赖真实 MySQL:测试入口把池里的连接标成已连接再走真正的 PoolWith。 +TEST(test_mysql, pool_with_fn_throw_returns_connection) { + State *s = FakeluaNewState(); + ASSERT_NE(s, nullptr); + CompileConfig config; + CompileFile(s, "./mysql/test_mysql_contract.lua", config); + + int arg_count = 0; + bool is_vararg = false; + void *addr = inter::GetFuncAddr(s, JIT_TCC, "MysqlContractTest.make_thrower", arg_count, is_vararg); + ASSERT_NE(addr, nullptr); + CVar fn = inter::DispatchCall(s, addr, nullptr, arg_count, JIT_TCC); + ASSERT_EQ(fn.type_, static_cast(VarType::Closure)); + + const int released = mysql::TestPoolWithFnThrowReturnsConnection(s, fn); + EXPECT_EQ(released, 1); + FakeluaDeleteState(s); +} + +// 回归(P0-1):异步回调的绑定参数(序列化字节串)跨过下一次顶层 Call 的临时 +// arena Reset 后仍然能被反序列化并原样送达回调。 +// 回调派回【登记时的引擎】:CallAll 三引擎各自 arm 出自己的连接,pump 帧 tick +// 时该连接的具名回调执行同引擎编译产物、写同引擎的文件级 local,三引擎各自闭环, +// 与旧内联闭包模型行为一致。 +TEST(test_mysql, closure_callback_survives_frame_reset) { + State *s = FakeluaNewState(); + ASSERT_NE(s, nullptr); + CompileConfig config; + CompileFile(s, "./mysql/test_mysql_contract.lua", config); + + int64_t armed = 0; + CallAll(s, "MysqlContractTest.arm_cross_frame", armed); + EXPECT_EQ(armed, 1); + // 顶层 Call 返回后显式再 Reset 一次,模拟宿主在帧末回收临时 arena。 + inter::Reset(s); + + int64_t hits = 0; + CallAll(s, "MysqlContractTest.pump_cross_frame", hits); + EXPECT_EQ(hits, 7); + FakeluaDeleteState(s); +} + +// 单元测试(P2-7):FieldCellToCVar 行值类型转换,无需真实 MySQL 连接。 +// 验证整数列返回 Int、浮点/DECIMAL 列返回 Float、其他列返回 String、NULL 返回 nil。 +// MYSQL_TYPE_* 常量取自 MariaDB/MySQL 协议规范(固定值,不随版本变化): +// LONG=3, LONGLONG=8, DOUBLE=5, NEWDECIMAL=246, VAR_STRING=253 +TEST(test_mysql, field_cell_to_cvar_type_conversion) { + State *s = FakeluaNewState(); + ASSERT_NE(s, nullptr); + + // 整数列 TINY(MYSQL_TYPE_TINY = 1) + { + FieldCell fv; + fv.is_null = false; + fv.value = "1"; + CVar v = FieldCellToCVar(s, fv, 1 /*MYSQL_TYPE_TINY*/); + EXPECT_EQ(v.type_, static_cast(VarType::Int)); + EXPECT_EQ(v.data_.i, 1); + } + // 整数列(MYSQL_TYPE_LONG = 3) + { + FieldCell fv; + fv.is_null = false; + fv.value = "42"; + CVar v = FieldCellToCVar(s, fv, 3 /*MYSQL_TYPE_LONG*/); + EXPECT_EQ(v.type_, static_cast(VarType::Int)); + EXPECT_EQ(v.data_.i, 42); + } + // 负整数(MYSQL_TYPE_LONGLONG = 8) + { + FieldCell fv; + fv.is_null = false; + fv.value = "-99"; + CVar v = FieldCellToCVar(s, fv, 8 /*MYSQL_TYPE_LONGLONG*/); + EXPECT_EQ(v.type_, static_cast(VarType::Int)); + EXPECT_EQ(v.data_.i, -99); + } + // 浮点列(MYSQL_TYPE_FLOAT = 4) + { + FieldCell fv; + fv.is_null = false; + fv.value = "1.5"; + CVar v = FieldCellToCVar(s, fv, 4 /*MYSQL_TYPE_FLOAT*/); + EXPECT_EQ(v.type_, static_cast(VarType::Float)); + EXPECT_NEAR(v.data_.f, 1.5, 1e-6); + } + // 浮点列(MYSQL_TYPE_DOUBLE = 5) + { + FieldCell fv; + fv.is_null = false; + fv.value = "3.14"; + CVar v = FieldCellToCVar(s, fv, 5 /*MYSQL_TYPE_DOUBLE*/); + EXPECT_EQ(v.type_, static_cast(VarType::Float)); + EXPECT_NEAR(v.data_.f, 3.14, 1e-9); + } + // DECIMAL 列(MYSQL_TYPE_DECIMAL = 0)与 NEWDECIMAL 都落成 double,精度见 README。 + { + FieldCell fv; + fv.is_null = false; + fv.value = "9.5"; + CVar v = FieldCellToCVar(s, fv, 0 /*MYSQL_TYPE_DECIMAL*/); + EXPECT_EQ(v.type_, static_cast(VarType::Float)); + EXPECT_NEAR(v.data_.f, 9.5, 1e-9); + } + // DECIMAL 列(MYSQL_TYPE_NEWDECIMAL = 246) + { + FieldCell fv; + fv.is_null = false; + fv.value = "123.456"; + CVar v = FieldCellToCVar(s, fv, 246 /*MYSQL_TYPE_NEWDECIMAL*/); + EXPECT_EQ(v.type_, static_cast(VarType::Float)); + EXPECT_NEAR(v.data_.f, 123.456, 1e-9); + } + // 字符串列(MYSQL_TYPE_VAR_STRING = 253) + { + FieldCell fv; + fv.is_null = false; + fv.value = "hello"; + CVar v = FieldCellToCVar(s, fv, 253 /*MYSQL_TYPE_VAR_STRING*/); + EXPECT_EQ(v.type_, static_cast(VarType::String)); + } + // NULL 字段(任意列类型) + { + FieldCell fv; + fv.is_null = true; + fv.value = ""; + CVar v = FieldCellToCVar(s, fv, 3 /*MYSQL_TYPE_LONG*/); + EXPECT_EQ(v.type_, static_cast(VarType::Nil)); + } + // 整数列但值为非数字 → 回退为 string(不崩溃) + { + FieldCell fv; + fv.is_null = false; + fv.value = "not_a_number"; + CVar v = FieldCellToCVar(s, fv, 3 /*MYSQL_TYPE_LONG*/); + EXPECT_EQ(v.type_, static_cast(VarType::String)); + } + + FakeluaDeleteState(s); +} diff --git a/test/test_net.cpp b/test/test_net.cpp index 6164838..a9bea51 100644 --- a/test/test_net.cpp +++ b/test/test_net.cpp @@ -384,3 +384,21 @@ TEST(test_net, test_udp_echo_lua) { FakeluaDeleteState(s); } + +// 回归(P1-6):在 on_event 回调里直接 obj:send() 必须可靠送达。 +// 派发上下文中的 send 入队,由本轮 tick 派发完后统一泵出。 +// 注意:C++ → Lua 回调按 VM 查表固定派发 TCC 编译的函数,文件级可变状态在各后端 +// 有独立副本,驱动端必须与回调同后端(TCC),因此这里不走 CallAll。 +TEST(test_net, test_send_in_callback) { + State *s = FakeluaNewState(); + ASSERT_NE(s, nullptr); + + CompileConfig config; + CompileFile(s, "./net/test_net_server_client.lua", config); + + int64_t ret = 0; + Call(s, JIT_TCC, "NetTest.test_send_in_callback", ret); + EXPECT_EQ(ret, 1); + + FakeluaDeleteState(s); +}