Compare commits

..

71 Commits

Author SHA1 Message Date
Alex Orlenko 0245d4ce6b v0.11.5 2025-11-22 16:13:08 +00:00
Alex Orlenko a7f105c698 Update Lua::set_memory_category doc 2025-11-22 14:27:02 +00:00
Alex Orlenko d2a8670bef (CI) Update wasi/wasmtime 2025-11-22 13:58:15 +00:00
Alex Orlenko e9de70a030 (CI) Move from x86_64-apple-darwin to aarch64-apple-darwin 2025-11-22 13:57:04 +00:00
Alex Orlenko aee647c6c0 Update CHANGELOG 2025-11-22 13:55:07 +00:00
Alex Orlenko 2e4184e7e4 Update spelling 2025-11-22 13:49:23 +00:00
Alex Orlenko 12b24b6c5b mlua-sys: v0.9.0 2025-11-22 13:47:17 +00:00
Alex Orlenko 1b500b7d47 Remove generic from internal definition of RawLua::create_string 2025-11-22 13:36:50 +00:00
Alex Orlenko ce4fc80e18 Bump luau-src to 0.17.0 (Luau 0.701) 2025-11-22 12:58:59 +00:00
Alex Orlenko 121971f54e Add Lua::set_memory_category and Lua::heap_dump functions to profile Luau memory usage.
This functionality uses Luau private api to dump heap mempory in JSON format for inspection.
The new type `HeapDump` represents memory snapshot with some basic API to calculate stats.
2025-11-22 00:01:34 +00:00
Alex Orlenko 6835537e3b Switch to released verson of luau0-src 2025-11-19 11:25:59 +00:00
Alex Orlenko a2728928cf Temporary disable some send tests on nightly
Aparently there is a regression in the compiler and sync detection does not work correctly
2025-11-18 00:51:24 +00:00
Alex Orlenko 676f3a6983 Fix tests 2025-11-17 13:44:29 +00:00
Alex Orlenko 0beaac228c Update Luau require tests 2025-11-16 23:51:51 +00:00
Alex Orlenko 9a7f75ad6b Update require implementation to satisfy Luau 0.700 2025-11-16 23:13:11 +00:00
Alex Orlenko cd56f92a7f Update Luau definitions in mlua-sys to 0.700 2025-11-16 23:12:06 +00:00
Alex Orlenko 1bd1359f43 Exclude the first arg when checking for yield_with call.
This is part of the previous commit
2025-11-06 23:28:37 +00:00
Alex Orlenko feec72bcbd Reduce number of allocations when calling async function
Instead of creating a uniq poller with upvalue on each async call, return future directly
and pass it to the poller
This also gives about 3-5% perf improvements
2025-11-06 23:09:26 +00:00
Alex Orlenko 0611906c6a Add Lua::type_metatable helper to get metatable of a primitive type.
The accompany function `Lua::set_type_metatable` already exists.
2025-11-05 22:07:14 +00:00
Alex Orlenko 72ac247dca Fix MaybeSend doc 2025-11-04 23:20:53 +00:00
Alex Orlenko f2fd010c5f Add missing lua_remove when discovering function name 2025-10-28 16:04:24 +00:00
Alex Orlenko 0619f264de Add Lua::traceback function to generate stack traces at different levels
This is similar to `debug.traceback`, through does not require debug module.
Close #652
2025-10-28 14:49:22 +00:00
Alex Orlenko ddd44bdd36 Add LUA_LOADED_TABLE constant (Luau) 2025-10-27 20:57:56 +00:00
Alex Orlenko 1152519074 Add add_method_once and add_async_method_once UserData methods (experimental).
They will allow implementing userdata methods that can be called only once, destructing
userdata instance during the call.
2025-10-26 20:14:46 +00:00
Alex Orlenko 3a2fd1ec59 Make AnyUserData::type_name public 2025-10-17 20:00:53 +01:00
Alex Orlenko a4c8b20697 impl IntoLuaMulti for &MultiValue 2025-10-13 12:36:42 +01:00
Alex Crichton 6e353d6c9f Build/test wasm32-wasip2 in CI (#649)
This is a follow-up from mlua-rs/lua-src-rs#13 which verifies/tests that
mlua/lua all work when compiled for a WASI target. While this doesn't
have formal documentation yet it also codifies in CI configuration how
to build for WASI and get tests passing (notably C compiler
configuration and some misc Rust flags).

This moves some `dev-dependencies` that don't compile for
`wasm32-wasip2` to a different section of the manifest. This
additionally annotates panicking tests with `#[cfg(not(panic =
"abort"))]` to skip those tests on WASI.

This does not test either the `send` or `async` feature at this time.
Testing `send` requires threads which WASI does not yet support, and
testing `async` requires more support in Tokio which is not currently
there yet.
2025-10-04 09:52:23 +01:00
Alex Orlenko 247208edb1 v0.11.4 2025-09-28 23:46:55 +01:00
Alex Orlenko e08768cc5e Derive Default for Value (clippy) 2025-09-28 23:42:12 +01:00
Alex Orlenko 5b38af9746 AsyncCallFuture is Unpin 2025-09-19 10:00:28 +01:00
Alex Orlenko 54907f80c5 Add SerializableValue to lib and prelude exports 2025-09-12 12:40:43 +01:00
Alex Orlenko ae512f2b49 Remove const from SerializableValue (it's not really useful) 2025-09-12 12:40:00 +01:00
Alex Orlenko 53c159b6cb Unhide Value::to_serializable 2025-09-12 11:49:43 +01:00
Alex Orlenko 2beca6ebe1 Add test for Table::for_each_value 2025-09-12 11:49:37 +01:00
Alex Orlenko 09da7a41e5 Add new serde option "detect_mixed_tables"
This option would allow detecting mixed tables (with array-like and map-like entries or several borders)
to encoding them chosing the best method (as a map or as a table).
2025-09-12 11:11:18 +01:00
Alex Orlenko bad20374ad Simplify Table::clear method
There is no need to traverse array part, lua_next will cover everything
2025-09-08 23:37:28 +01:00
Alex Orlenko 40b507c3ec Add ObjectLike::get_path helper 2025-09-04 19:12:44 +01:00
Andrew Dunbar 537cc995f6 Copyedit English in README.md (#639) 2025-09-04 14:59:24 +01:00
Alex Orlenko 5d27cb91b2 Add optional __namecall optimization for Luau
Add `UserDataRegistry::enable_namecall()` hint to set `__namecall` metamethod to enable Luau-specific method resolution optimization.
2025-09-02 00:53:12 +01:00
Alex Orlenko c70a636ca9 Remove newlines from yield_with examples 2025-08-30 12:51:53 +01:00
Alex Orlenko 13ff0ca798 v0.11.3 2025-08-29 23:11:21 +01:00
Alex Orlenko 44f49e35d6 Update CHANGELOG 2025-08-29 00:18:06 +01:00
Alex Orlenko e1ee4058a6 Add new benchmark to measure complex userdata method calls 2025-08-28 23:56:03 +01:00
Alex Orlenko f06d0020ea Add test to emulate method through field 2025-08-28 23:50:18 +01:00
Alex Orlenko d399559d30 Add Lua::yield_with to allow yielding Rust async functions and exchange values between Lua coroutine and Rust.
This functionality is similar to `coroutine.yield` and `coroutine.resume` without C restrictions.
2025-08-28 18:41:24 +01:00
Alex Orlenko 30735d5ff1 Fix thread recovery when pushing a bad arg
We should not erase thread stack if a bad argument is pushed before resuming the thread.
2025-08-25 23:07:37 +01:00
Alex Orlenko 75c23e5853 Add lua_cpcall to Luau ffi (0.688+) 2025-08-25 12:54:17 +01:00
Alex Orlenko 347856b806 Do not try to yield at non-yielable points in Luau interrupt
In particular we cannot yeild across metamethod/C-call boundaries.
This behaviour matches with Lua 5.3+ yielding from hooks only at safe points.
Closes #632
2025-08-25 12:19:50 +01:00
Alex Orlenko 774a63bece Add Buffer::cursor() method
This can be useful for providing access to buffers through core IO traits.
2025-08-24 11:29:01 +01:00
Alex Orlenko c481c87eac Add Lua::create_buffer_with_capacity method
This allow creating a preallocated buffer with specified size initialized to zero.
2025-08-23 22:38:55 +01:00
Alex Orlenko 85b280a9d6 Update nightly Rust error message matching 2025-08-23 09:40:13 +01:00
Alex Orlenko db7b782d3c Remove lifetimes from short type names 2025-08-23 09:13:31 +01:00
Alex Orlenko 5f38445558 Fix warnings 2025-08-20 16:25:06 +01:00
Alex Orlenko df0a44d405 Make Lua reference values cheap to clone
Instead of locking the VM and making a copy on auxiliary thread, track number of references using Rust ref counter.
This should also help reducing number of used references (they are limited to to 1M usually) on auxiliary thread.
2025-08-20 12:05:37 +01:00
Alex Orlenko f0806a6d62 Lower fastpath table creation limit to 1 << 26
When Lua is configured without memory restrictions, we use fastpath for table creation (unprotected mode).
In generally it's safe as long as we `abort()` on allocation failure.
However some Lua versions have additional restrictions on table size that we need to adhere in mlua too.
Probably Luau has the lowest limits.
Fixes #627
2025-08-13 22:49:40 +01:00
Alex Orlenko 3516f4c6ca v0.11.2 2025-08-10 00:53:45 +01:00
Alex Orlenko ca73583714 Update CHANGELOG 2025-08-10 00:53:01 +01:00
Alex Orlenko 36560435f7 Add push_into_stack_multi fastpath to Variadic 2025-08-10 00:35:51 +01:00
Alex Orlenko 763c2b2564 Update repl example: don't print newline if no values returned 2025-08-10 00:20:20 +01:00
Alex Orlenko bafdb6138c Update dependencies 2025-08-10 00:19:54 +01:00
Alex Orlenko c9d6a610e1 mlua-sys: v0.8.3 2025-08-10 00:11:05 +01:00
Alex Orlenko bd63f63bc9 Use ascii lowercase for module aliases
This matches with Luau 0.686 changes
2025-08-09 19:14:31 +01:00
piz-ewing c035c23a15 fix: normalize_chunk_name handles Windows paths with drive letter (#623)
Co-authored-by: ewing <ewing@MacBook-Pro.local>
2025-08-04 22:34:36 +01:00
Alex Orlenko cb153a52b2 Make Luau registered aliases case-insensitive
Executing `require("@my_module")` or `require("@My_Module")` should give the same result and use case-insensitive name.
See #620 for details
2025-07-26 22:23:16 +01:00
Alex Orlenko b1c69d3005 Use to_bits comparison to check if a float value can be represented as an integer losslessly.
This allows to simplify the code while still maintaining "negative zeros" edge case.
Thanks @JasonHise for the suggestion.
2025-07-25 21:25:08 +01:00
Alex Orlenko 841bd332e4 Fix LuaJIT negative zero tests 2025-07-25 15:24:04 +01:00
Alex Orlenko 815d1bd7c9 Better handling negative zeros to match Lua 5.3+ behavior
In Lua 5.3+ the function `lua_isinteger` returns "false" for -0.0 numbers.
In earlier Lua versions we should follow the same behavior to avoid losing the sign when converting to Integer.
Close #618
2025-07-25 14:32:47 +01:00
Alex Orlenko 78331ceebe v0.11.1 2025-07-15 22:43:18 +01:00
Alex Orlenko f945a35cbd Execute metatable destructor in Table::set_metatable at the end of invocation
Before this change, destructor was executed shortly after pushing metatable to ref_thread.
2025-07-15 19:14:46 +01:00
Alex Orlenko 459edb6816 Always grow aux ref stack considering the reserve 2025-07-15 16:32:22 +01:00
Alex Orlenko 00328b0b64 Protect Lua::push_c_function for Lua <5.2 2025-07-15 16:11:31 +01:00
80 changed files with 3118 additions and 671 deletions
+38 -3
View File
@@ -14,7 +14,7 @@ jobs:
- os: ubuntu-latest
target: x86_64-unknown-linux-gnu
- os: macos-latest
target: x86_64-apple-darwin
target: aarch64-apple-darwin
- os: windows-latest
target: x86_64-pc-windows-msvc
steps:
@@ -110,7 +110,7 @@ jobs:
- os: ubuntu-latest
target: x86_64-unknown-linux-gnu
- os: macos-latest
target: x86_64-apple-darwin
target: aarch64-apple-darwin
- os: windows-latest
target: x86_64-pc-windows-msvc
steps:
@@ -199,7 +199,7 @@ jobs:
- os: ubuntu-latest
target: x86_64-unknown-linux-gnu
- os: macos-latest
target: x86_64-apple-darwin
target: aarch64-apple-darwin
steps:
- uses: actions/checkout@main
- uses: dtolnay/rust-toolchain@stable
@@ -256,6 +256,41 @@ jobs:
cargo test --tests --features "${{ matrix.lua }},vendored"
cargo test --tests --features "${{ matrix.lua }},vendored,async,serde,macros,anyhow,userdata-wrappers"
test_wasm32_wasip2:
name: Test on wasm32-wasip2
runs-on: ubuntu-latest
needs: build
strategy:
matrix:
lua: [lua54, lua53, lua52, lua51]
steps:
- uses: actions/checkout@main
- uses: dtolnay/rust-toolchain@stable
with:
toolchain: nightly-2025-10-02
target: wasm32-wasip2
- name: Install wasi-sdk/Wasmtime
working-directory: ${{ runner.tool_cache }}
run: |
wasi_sdk=29
wasmtime=v39.0.0
curl -LO https://github.com/WebAssembly/wasi-sdk/releases/download/wasi-sdk-$wasi_sdk/wasi-sdk-$wasi_sdk.0-x86_64-linux.tar.gz
tar xf wasi-sdk-$wasi_sdk.0-x86_64-linux.tar.gz
WASI_SDK_PATH=`pwd`/wasi-sdk-$wasi_sdk.0-x86_64-linux
echo "WASI_SDK_PATH=$WASI_SDK_PATH" >> $GITHUB_ENV
echo "CC_wasm32_wasip2=$WASI_SDK_PATH/bin/clang" >> $GITHUB_ENV
echo "CARGO_TARGET_WASM32_WASIP2_LINKER=$WASI_SDK_PATH/bin/clang" >> $GITHUB_ENV
echo "CARGO_TARGET_WASM32_WASIP2_RUSTFLAGS=-Clink-arg=-Wl,--export=cabi_realloc" >> $GITHUB_ENV
curl -LO https://github.com/bytecodealliance/wasmtime/releases/download/$wasmtime/wasmtime-$wasmtime-x86_64-linux.tar.xz
tar xf wasmtime-$wasmtime-x86_64-linux.tar.xz
echo "CARGO_TARGET_WASM32_WASIP2_RUNNER=`pwd`/wasmtime-$wasmtime-x86_64-linux/wasmtime -W exceptions" >> $GITHUB_ENV
- name: Run ${{ matrix.lua }} tests
run: |
cargo test --target wasm32-wasip2 --tests --features "${{ matrix.lua }},vendored"
cargo test --target wasm32-wasip2 --tests --features "${{ matrix.lua }},vendored,serde,macros,anyhow,userdata-wrappers"
rustfmt:
name: Rustfmt
runs-on: ubuntu-latest
+38
View File
@@ -1,3 +1,41 @@
## v0.11.5 (Nov 22, 2025)
- Luau updated to 0.701
- Added `Lua::set_memory_category` and `Lua::heap_dump` functions to profile (Luau) memory
- Added `Lua::type_metatable` helper to get metatable of a primitive type
- Added `Lua::traceback` function to generate stack traces at different levels
- Added `add_method_once` /`add_async_method_once` UserData methods (experimental)
- Make `AnyUserData::type_name` public
- impl `IntoLuaMulti` for `&MultiValue`
- Bugfixes and async perf improvements
## v0.11.4 (Sep 29, 2025)
- Make `Value::to_serializable` public
- Add new serde option `detect_mixed_tables` (to encode mixed array+map tables)
- Add `ObjectLike::get_path` helper (for tables and userdata)
## v0.11.3 (Aug 30, 2025)
- Add `Lua::yield_with` to use as `coroutine.yield` functional replacement in async functions for any Lua
- Do not try to yield at non-yielable points in Luau interrupt (#632)
- Add `Buffer::cursor` method (Luau)
- Add `Lua::create_buffer_with_capacity` method (Luau)
- Make Lua reference values cheap to clone (only increments ref count)
- Fix panic on large (>67M entries) table creation
## v0.11.2 (Aug 10, 2025)
- Faster stack push for `Variadic<T>`
- Fix handling Windows paths with drive letter in Luau require (#623)
- Make Luau registered aliases ascii case-insensitive (#620)
- Fix deserializing negative zeros `-0.0` (#618)
## v0.11.1 (Jul 15, 2025)
- Fixed bug exhausting Lua auxiliary stack and leaving it without reserve (#615)
- `Lua::push_c_function` now correctly handles OOM for Lua 5.1 and Luau
## v0.11.0 (Jul 14, 2025)
Changes since v0.11.0-beta.3
+12 -11
View File
@@ -1,8 +1,8 @@
[package]
name = "mlua"
version = "0.11.0" # remember to update mlua_derive
version = "0.11.5" # remember to update mlua_derive
authors = ["Aleksandr Orlenko <zxteam@pm.me>", "kyren <catherine@kyju.org>"]
rust-version = "1.79.0"
rust-version = "1.80.0"
edition = "2021"
repository = "https://github.com/mlua-rs/mlua"
documentation = "https://docs.rs/mlua"
@@ -61,25 +61,26 @@ serde-value = { version = "0.7", optional = true }
parking_lot = { version = "0.12", features = ["arc_lock"] }
anyhow = { version = "1.0", optional = true }
rustversion = "1.0"
libc = "0.2"
ffi = { package = "mlua-sys", version = "0.8.0", path = "mlua-sys" }
ffi = { package = "mlua-sys", version = "0.9.0", path = "mlua-sys" }
[dev-dependencies]
trybuild = "1.0"
hyper = { version = "1.2", features = ["full"] }
hyper-util = { version = "0.1.3", features = ["full"] }
http-body-util = "0.1.1"
reqwest = { version = "0.12", features = ["json"] }
tokio = { version = "1.0", features = ["macros", "rt", "time"] }
serde = { version = "1.0", features = ["derive"] }
serde_json = { version = "1.0", features = ["arbitrary_precision"] }
maplit = "1.0"
tempfile = "3"
static_assertions = "1.0"
[target.'cfg(not(target_arch = "wasm32"))'.dev-dependencies]
criterion = { version = "0.6", features = ["async_tokio"] }
rustyline = "16.0"
[target.'cfg(not(target_family = "wasm"))'.dev-dependencies]
hyper = { version = "1.2", features = ["full"] }
hyper-util = { version = "0.1.3", features = ["full"] }
http-body-util = "0.1.1"
reqwest = { version = "0.12", features = ["json"] }
tempfile = "3"
criterion = { version = "0.7", features = ["async_tokio"] }
rustyline = "17.0"
tokio = { version = "1.0", features = ["full"] }
[lints.rust]
+4 -4
View File
@@ -17,14 +17,14 @@
[Benchmarks]: https://github.com/khvzak/script-bench-rs
[FAQ]: FAQ.md
`mlua` is a set of bindings to the [Lua](https://www.lua.org) programming language for Rust with a goal to provide a
`mlua` is a set of bindings to the [Lua](https://www.lua.org) programming language for Rust with a goal of providing a
_safe_ (as much as possible), high level, easy to use, practical and flexible API.
Started as an `rlua` fork, `mlua` supports Lua 5.4, 5.3, 5.2, 5.1 (including LuaJIT) and [Luau] and allows writing native Lua modules in Rust as well as using Lua in a standalone mode.
`mlua` is tested on Windows/macOS/Linux including module mode in [GitHub Actions] on `x86_64` platforms and cross-compilation to `aarch64` (other targets are also supported).
WebAssembly (WASM) is supported through `wasm32-unknown-emscripten` target for all Lua/Luau versions excluding JIT.
WebAssembly (WASM) is supported through the `wasm32-unknown-emscripten` target for all Lua/Luau versions excluding JIT.
[GitHub Actions]: https://github.com/mlua-rs/mlua/actions
[Luau]: https://luau.org
@@ -33,7 +33,7 @@ WebAssembly (WASM) is supported through `wasm32-unknown-emscripten` target for a
### Feature flags
`mlua` uses feature flags to reduce the amount of dependencies and compiled code, and allow to choose only required set of features.
`mlua` uses feature flags to reduce the number of dependencies and compiled code, and allow choosing only the required set of features.
Below is a list of the available feature flags. By default `mlua` does not enable any features.
* `lua54`: enable Lua [5.4] support
@@ -270,7 +270,7 @@ remain usable after a user generated panic, and such panics should not break int
leak Lua stack space. This is mostly important to safely use `mlua` types in Drop impls, as you should not be
using panics for general error handling.
Below is a list of `mlua` behaviors that should be considered a bug.
Below is a list of `mlua` behaviors that should be considered bugs.
If you encounter them, a bug report would be very welcome:
+ If you can cause UB with `mlua` without typing the word "unsafe", this is a bug.
+54
View File
@@ -128,6 +128,22 @@ fn table_traversal_sequence(c: &mut Criterion) {
});
}
fn table_ref_clone(c: &mut Criterion) {
let lua = Lua::new();
let t = lua.create_table().unwrap();
c.bench_function("table [ref clone]", |b| {
b.iter_batched(
|| collect_gc_twice(&lua),
|_| {
let _t2 = t.clone();
},
BatchSize::SmallInput,
);
});
}
fn function_create(c: &mut Criterion) {
let lua = Lua::new();
@@ -350,6 +366,42 @@ fn userdata_call_method(c: &mut Criterion) {
});
}
// A userdata method call that goes through an implicit `__index` function
fn userdata_call_method_complex(c: &mut Criterion) {
struct UserData(u64);
impl LuaUserData for UserData {
fn register(registry: &mut LuaUserDataRegistry<Self>) {
registry.add_field_method_get("val", |_, this| Ok(this.0));
registry.add_method_mut("inc_by", |_, this, by: u64| {
this.0 += by;
Ok(this.0)
});
#[cfg(feature = "luau")]
registry.enable_namecall();
}
}
let lua = Lua::new();
let ud = lua.create_userdata(UserData(0)).unwrap();
let inc_by = lua
.load("function(ud, s) return ud:inc_by(s) end")
.eval::<LuaFunction>()
.unwrap();
c.bench_function("userdata [call method complex]", |b| {
b.iter_batched(
|| {
collect_gc_twice(&lua);
},
|_| {
inc_by.call::<()>((&ud, 1)).unwrap();
},
BatchSize::SmallInput,
);
});
}
fn userdata_async_call_method(c: &mut Criterion) {
struct UserData(i64);
impl LuaUserData for UserData {
@@ -399,6 +451,7 @@ criterion_group! {
table_traversal_pairs,
table_traversal_for_each,
table_traversal_sequence,
table_ref_clone,
function_create,
function_call_sum,
@@ -413,6 +466,7 @@ criterion_group! {
userdata_create,
userdata_call_index,
userdata_call_method,
userdata_call_method_complex,
userdata_async_call_method,
}
+10 -8
View File
@@ -20,14 +20,16 @@ fn main() {
match lua.load(&line).eval::<MultiValue>() {
Ok(values) => {
editor.add_history_entry(line).unwrap();
println!(
"{}",
values
.iter()
.map(|value| format!("{:#?}", value))
.collect::<Vec<_>>()
.join("\t")
);
if values.len() > 0 {
println!(
"{}",
values
.iter()
.map(|value| format!("{:#?}", value))
.collect::<Vec<_>>()
.join("\t")
);
}
break;
}
Err(Error::SyntaxError {
+2 -2
View File
@@ -1,6 +1,6 @@
[package]
name = "mlua-sys"
version = "0.8.2"
version = "0.9.0"
authors = ["Aleksandr Orlenko <zxteam@pm.me>"]
rust-version = "1.71"
edition = "2021"
@@ -41,7 +41,7 @@ cfg-if = "1.0"
pkg-config = "0.3.17"
lua-src = { version = ">= 548.1.0, < 548.2.0", optional = true }
luajit-src = { version = ">= 210.6.0, < 210.7.0", optional = true }
luau0-src = { version = "0.15.4", optional = true }
luau0-src = { version = "0.17.0", optional = true }
[lints.rust]
unexpected_cfgs = { level = "allow", check-cfg = ['cfg(raw_dylib)'] }
+52 -44
View File
@@ -2,6 +2,7 @@
//!
//! Based on github.com/keplerproject/lua-compat-5.3
use std::ffi::CStr;
use std::os::raw::{c_char, c_int, c_void};
use std::{mem, ptr};
@@ -20,8 +21,8 @@ unsafe fn compat53_reverse(L: *mut lua_State, mut a: c_int, mut b: c_int) {
}
}
const COMPAT53_LEVELS1: c_int = 12; // size of the first part of the stack
const COMPAT53_LEVELS2: c_int = 10; // size of the second part of the stack
const COMPAT53_LEVELS1: c_int = 10; // size of the first part of the stack
const COMPAT53_LEVELS2: c_int = 11; // size of the second part of the stack
unsafe fn compat53_countlevels(L: *mut lua_State) -> c_int {
let mut ar: lua_Debug = mem::zeroed();
@@ -88,11 +89,10 @@ unsafe fn compat53_findfield(L: *mut lua_State, objidx: c_int, level: c_int) ->
lua_pop(L, 1); // remove value (but keep name)
return 1;
} else if compat53_findfield(L, objidx, level - 1) != 0 {
// try recursively
lua_remove(L, -2); // remove table (but keep name)
lua_pushliteral(L, c".");
lua_insert(L, -2); // place '.' between the two names
lua_concat(L, 3);
// stack: lib_name, lib_table, field_name (top)
lua_pushliteral(L, c"."); // place '.' between the two names
lua_replace(L, -3); // (in the slot occupied by table)
lua_concat(L, 3); // lib_name.field_name
return 1;
}
}
@@ -101,13 +101,20 @@ unsafe fn compat53_findfield(L: *mut lua_State, objidx: c_int, level: c_int) ->
0 // not found
}
unsafe fn compat53_pushglobalfuncname(L: *mut lua_State, ar: *mut lua_Debug) -> c_int {
unsafe fn compat53_pushglobalfuncname(L: *mut lua_State, L1: *mut lua_State, ar: *mut lua_Debug) -> c_int {
let top = lua_gettop(L);
lua_getinfo(L, cstr!("f"), ar); // push function
lua_getinfo(L1, cstr!("f"), ar); // push function
lua_xmove(L1, L, 1); // and move onto L
lua_pushvalue(L, LUA_GLOBALSINDEX);
luaL_checkstack(L, 6, cstr!("not enough stack")); // slots for 'findfield'
if compat53_findfield(L, top + 1, 2) != 0 {
let name = lua_tostring(L, -1);
if CStr::from_ptr(name).to_bytes().starts_with(b"_G.") {
lua_pushstring(L, name.add(3)); // push name without prefix
lua_remove(L, -2); // remove original name
}
lua_copy(L, -1, top + 1); // move name to proper place
lua_pop(L, 2); // remove pushed values
lua_settop(L, top + 1); // remove pushed values
1
} else {
lua_settop(L, top); // remove function and global table
@@ -115,27 +122,23 @@ unsafe fn compat53_pushglobalfuncname(L: *mut lua_State, ar: *mut lua_Debug) ->
}
}
unsafe fn compat53_pushfuncname(L: *mut lua_State, ar: *mut lua_Debug) {
if *(*ar).namewhat != b'\0' as c_char {
// is there a name?
lua_pushfstring(L, cstr!("function '%s'"), (*ar).name);
unsafe fn compat53_pushfuncname(L: *mut lua_State, L1: *mut lua_State, ar: *mut lua_Debug) {
// try first a global name
if compat53_pushglobalfuncname(L, L1, ar) != 0 {
lua_pushfstring(L, cstr!("function '%s'"), lua_tostring(L, -1));
lua_remove(L, -2); // remove name
} else if *(*ar).namewhat != b'\0' as c_char {
// use name from code
lua_pushfstring(L, cstr!("%s '%s'"), (*ar).namewhat, (*ar).name);
} else if *(*ar).what == b'm' as c_char {
// main?
lua_pushliteral(L, c"main chunk");
} else if *(*ar).what == b'C' as c_char {
if compat53_pushglobalfuncname(L, ar) != 0 {
lua_pushfstring(L, cstr!("function '%s'"), lua_tostring(L, -1));
lua_remove(L, -2); // remove name
} else {
lua_pushliteral(L, c"?");
}
} else if *(*ar).what != b'C' as c_char {
// for Lua functions, use <file:line>
let short_src = (*ar).short_src.as_ptr();
lua_pushfstring(L, cstr!("function <%s:%d>"), short_src, (*ar).linedefined);
} else {
lua_pushfstring(
L,
cstr!("function <%s:%d>"),
(*ar).short_src.as_ptr(),
(*ar).linedefined,
);
lua_pushliteral(L, c"?");
}
}
@@ -186,7 +189,8 @@ pub unsafe fn lua_isinteger(L: *mut lua_State, idx: c_int) -> c_int {
if lua_type(L, idx) == LUA_TNUMBER {
let n = lua_tonumber(L, idx);
let i = lua_tointeger(L, idx);
if (n - i as lua_Number).abs() < lua_Number::EPSILON {
// Lua 5.3+ returns "false" for `-0.0`
if n.to_bits() == (i as lua_Number).to_bits() {
return 1;
}
}
@@ -458,32 +462,36 @@ pub unsafe fn luaL_traceback(L: *mut lua_State, L1: *mut lua_State, msg: *const
let mut ar: lua_Debug = mem::zeroed();
let top = lua_gettop(L);
let numlevels = compat53_countlevels(L1);
let mark = if numlevels > COMPAT53_LEVELS1 + COMPAT53_LEVELS2 {
COMPAT53_LEVELS1
} else {
0
};
#[rustfmt::skip]
let mut limit = if numlevels - level > COMPAT53_LEVELS1 + COMPAT53_LEVELS2 { COMPAT53_LEVELS1 } else { -1 };
if !msg.is_null() {
lua_pushfstring(L, cstr!("%s\n"), msg);
}
lua_pushliteral(L, c"stack traceback:");
while lua_getstack(L1, level, &mut ar) != 0 {
level += 1;
if level == mark {
if limit == 0 {
// too many levels?
lua_pushliteral(L, c"\n\t..."); // add a '...'
level = numlevels - COMPAT53_LEVELS2; // and skip to last ones
let n = numlevels - level - COMPAT53_LEVELS2;
// add warning about skip ("n + 1" because we skip current level too)
lua_pushfstring(L, cstr!("\n\t...\t(skipping %d levels)"), n + 1); // add warning about skip
level += n; // and skip to last levels
} else {
lua_getinfo(L1, cstr!("Slnt"), &mut ar);
lua_pushfstring(L, cstr!("\n\t%s:"), ar.short_src.as_ptr());
if ar.currentline > 0 {
lua_pushfstring(L, cstr!("%d:"), ar.currentline);
lua_getinfo(L1, cstr!("Sln"), &mut ar);
if *ar.what != b't' as c_char {
if ar.currentline <= 0 {
lua_pushfstring(L, cstr!("\n\t%s: in "), ar.short_src.as_ptr());
} else {
lua_pushfstring(L, cstr!("\n\t%s:%d: in "), ar.short_src.as_ptr(), ar.currentline);
}
compat53_pushfuncname(L, L1, &mut ar);
lua_concat(L, lua_gettop(L) - top);
} else {
lua_pushstring(L, cstr!("\n\t(...tail calls...)"));
}
lua_pushliteral(L, c" in ");
compat53_pushfuncname(L, &mut ar);
lua_concat(L, lua_gettop(L) - top);
}
level += 1;
limit -= 1;
}
lua_concat(L, lua_gettop(L) - top);
}
+2 -1
View File
@@ -51,7 +51,8 @@ pub unsafe fn lua_isinteger(L: *mut lua_State, idx: c_int) -> c_int {
if lua_type(L, idx) == LUA_TNUMBER {
let n = lua_tonumber(L, idx);
let i = lua_tointeger(L, idx);
if (n - i as lua_Number).abs() < lua_Number::EPSILON {
// Lua 5.3+ returns "false" for `-0.0`
if n.to_bits() == (i as lua_Number).to_bits() {
return 1;
}
}
+56 -42
View File
@@ -23,8 +23,8 @@ unsafe fn compat53_reverse(L: *mut lua_State, mut a: c_int, mut b: c_int) {
}
}
const COMPAT53_LEVELS1: c_int = 12; // size of the first part of the stack
const COMPAT53_LEVELS2: c_int = 10; // size of the second part of the stack
const COMPAT53_LEVELS1: c_int = 10; // size of the first part of the stack
const COMPAT53_LEVELS2: c_int = 11; // size of the second part of the stack
unsafe fn compat53_findfield(L: *mut lua_State, objidx: c_int, level: c_int) -> c_int {
if level == 0 || lua_istable(L, -1) == 0 {
@@ -41,11 +41,10 @@ unsafe fn compat53_findfield(L: *mut lua_State, objidx: c_int, level: c_int) ->
lua_pop(L, 1); // remove value (but keep name)
return 1;
} else if compat53_findfield(L, objidx, level - 1) != 0 {
// try recursively
lua_remove(L, -2); // remove table (but keep name)
lua_pushliteral(L, c".");
lua_insert(L, -2); // place '.' between the two names
lua_concat(L, 3);
// stack: lib_name, lib_table, field_name (top)
lua_pushliteral(L, c"."); // place '.' between the two names
lua_replace(L, -3); // (in the slot occupied by table)
lua_concat(L, 3); // lib_name.field_name
return 1;
}
}
@@ -54,14 +53,25 @@ unsafe fn compat53_findfield(L: *mut lua_State, objidx: c_int, level: c_int) ->
0 // not found
}
unsafe fn compat53_pushglobalfuncname(L: *mut lua_State, level: c_int, ar: *mut lua_Debug) -> c_int {
unsafe fn compat53_pushglobalfuncname(
L: *mut lua_State,
L1: *mut lua_State,
level: c_int,
ar: *mut lua_Debug,
) -> c_int {
let top = lua_gettop(L);
// push function
lua_getinfo(L, level, cstr!("f"), ar);
lua_getinfo(L1, level, cstr!("f"), ar); // push function
lua_xmove(L1, L, 1); // and move onto L
lua_pushvalue(L, LUA_GLOBALSINDEX);
luaL_checkstack(L, 6, cstr!("not enough stack")); // slots for 'findfield'
if compat53_findfield(L, top + 1, 2) != 0 {
let name = lua_tostring(L, -1);
if CStr::from_ptr(name).to_bytes().starts_with(b"_G.") {
lua_pushstring(L, name.add(3)); // push name without prefix
lua_remove(L, -2); // remove original name
}
lua_copy(L, -1, top + 1); // move name to proper place
lua_pop(L, 2); // remove pushed values
lua_settop(L, top + 1); // remove pushed values
1
} else {
lua_settop(L, top); // remove function and global table
@@ -69,13 +79,16 @@ unsafe fn compat53_pushglobalfuncname(L: *mut lua_State, level: c_int, ar: *mut
}
}
unsafe fn compat53_pushfuncname(L: *mut lua_State, level: c_int, ar: *mut lua_Debug) {
unsafe fn compat53_pushfuncname(L: *mut lua_State, L1: *mut lua_State, level: c_int, ar: *mut lua_Debug) {
if !(*ar).name.is_null() {
// is there a name?
lua_pushfstring(L, cstr!("function '%s'"), (*ar).name);
} else if compat53_pushglobalfuncname(L, level, ar) != 0 {
} else if compat53_pushglobalfuncname(L, L1, level, ar) != 0 {
lua_pushfstring(L, cstr!("function '%s'"), lua_tostring(L, -1));
lua_remove(L, -2); // remove name
} else if *(*ar).what != b'C' as c_char {
// for Lua functions, use <file:line>
lua_pushfstring(L, cstr!("function <%s:%d>"), (*ar).short_src, (*ar).linedefined);
} else {
lua_pushliteral(L, c"?");
}
@@ -120,7 +133,8 @@ pub unsafe fn lua_isinteger(L: *mut lua_State, idx: c_int) -> c_int {
if lua_type(L, idx) == LUA_TNUMBER {
let n = lua_tonumber(L, idx);
let i = lua_tointeger(L, idx);
if (n - i as lua_Number).abs() < lua_Number::EPSILON {
// Lua 5.3+ returns "false" for `-0.0`
if n.to_bits() == (i as lua_Number).to_bits() {
return 1;
}
}
@@ -189,9 +203,7 @@ pub unsafe fn lua_rawgeti(L: *mut lua_State, idx: c_int, n: lua_Integer) -> c_in
#[inline(always)]
pub unsafe fn lua_rawgetp(L: *mut lua_State, idx: c_int, p: *const c_void) -> c_int {
let abs_i = lua_absindex(L, idx);
lua_pushlightuserdata(L, p as *mut c_void);
lua_rawget(L, abs_i)
lua_rawgetptagged(L, idx, p, 0)
}
#[inline(always)]
@@ -225,11 +237,7 @@ pub unsafe fn lua_rawseti(L: *mut lua_State, idx: c_int, n: lua_Integer) {
#[inline(always)]
pub unsafe fn lua_rawsetp(L: *mut lua_State, idx: c_int, p: *const c_void) {
let abs_i = lua_absindex(L, idx);
luaL_checkstack(L, 1, cstr!("not enough stack slots available"));
lua_pushlightuserdata(L, p as *mut c_void);
lua_insert(L, -2);
lua_rawset(L, abs_i);
lua_rawsetptagged(L, idx, p, 0)
}
#[inline(always)]
@@ -451,36 +459,42 @@ pub unsafe fn luaL_len(L: *mut lua_State, idx: c_int) -> lua_Integer {
pub unsafe fn luaL_traceback(L: *mut lua_State, L1: *mut lua_State, msg: *const c_char, mut level: c_int) {
let mut ar: lua_Debug = mem::zeroed();
let top = lua_gettop(L);
let numlevels = lua_stackdepth(L);
let mark = if numlevels > COMPAT53_LEVELS1 + COMPAT53_LEVELS2 {
COMPAT53_LEVELS1
} else {
0
};
#[rustfmt::skip]
let mut limit = if numlevels - level > COMPAT53_LEVELS1 + COMPAT53_LEVELS2 { COMPAT53_LEVELS1 } else { -1 };
let mut buf: luaL_Strbuf = mem::zeroed();
luaL_buffinit(L, &mut buf);
if !msg.is_null() {
lua_pushfstring(L, cstr!("%s\n"), msg);
luaL_addstring(&mut buf, msg);
luaL_addstring(&mut buf, cstr!("\n"));
}
lua_pushliteral(L, c"stack traceback:");
while lua_getinfo(L1, level, cstr!(""), &mut ar) != 0 {
if level + 1 == mark {
luaL_addstring(&mut buf, cstr!("stack traceback:"));
while lua_getinfo(L1, level, cstr!("sln"), &mut ar) != 0 {
if limit == 0 {
// too many levels?
lua_pushliteral(L, c"\n\t..."); // add a '...'
level = numlevels - COMPAT53_LEVELS2; // and skip to last ones
let n = numlevels - level - COMPAT53_LEVELS2;
// add warning about skip ("n + 1" because we skip current level too)
lua_pushfstring(L, cstr!("\n\t...\t(skipping %d levels)"), n + 1);
luaL_addvalue(&mut buf);
level += n; // and skip to last levels
} else {
lua_getinfo(L1, level, cstr!("sln"), &mut ar);
lua_pushfstring(L, cstr!("\n\t%s:"), ar.short_src);
luaL_addstring(&mut buf, cstr!("\n\t"));
luaL_addstring(&mut buf, ar.short_src);
luaL_addstring(&mut buf, cstr!(":"));
if ar.currentline > 0 {
lua_pushfstring(L, cstr!("%d:"), ar.currentline);
luaL_addunsigned(&mut buf, ar.currentline as _);
luaL_addstring(&mut buf, cstr!(":"));
}
lua_pushliteral(L, c" in ");
compat53_pushfuncname(L, level, &mut ar);
lua_concat(L, lua_gettop(L) - top);
luaL_addstring(&mut buf, cstr!(" in "));
compat53_pushfuncname(L, L1, level, &mut ar);
luaL_addvalue(&mut buf);
}
level += 1;
limit -= 1;
}
lua_concat(L, lua_gettop(L) - top);
luaL_pushresult(&mut buf);
}
pub unsafe fn luaL_tolstring(L: *mut lua_State, mut idx: c_int, len: *mut usize) -> *const c_char {
@@ -543,7 +557,7 @@ pub unsafe fn luaL_getsubtable(L: *mut lua_State, idx: c_int, fname: *const c_ch
pub unsafe fn luaL_requiref(L: *mut lua_State, modname: *const c_char, openf: lua_CFunction, glb: c_int) {
luaL_checkstack(L, 3, cstr!("not enough stack slots available"));
luaL_getsubtable(L, LUA_REGISTRYINDEX, cstr!("_LOADED"));
luaL_getsubtable(L, LUA_REGISTRYINDEX, LUA_LOADED_TABLE);
if lua_getfield(L, -1, modname) == LUA_TNIL {
lua_pop(L, 1);
lua_pushcfunction(L, openf);
+18
View File
@@ -5,6 +5,9 @@ use std::ptr;
use super::lua::{self, lua_CFunction, lua_Number, lua_State, lua_Unsigned, LUA_REGISTRYINDEX};
// Key, in the registry, for table of loaded modules
pub const LUA_LOADED_TABLE: *const c_char = cstr!("_LOADED");
#[repr(C)]
pub struct luaL_Reg {
pub name: *const c_char,
@@ -209,3 +212,18 @@ pub unsafe fn luaL_addstring(B: *mut luaL_Strbuf, s: *const c_char) {
}
luaL_addlstring(B, s, len);
}
pub unsafe fn luaL_addunsigned(B: *mut luaL_Strbuf, mut n: lua_Unsigned) {
let mut buf: [c_char; 32] = [0; 32];
let mut i = 32;
loop {
i -= 1;
let digit = (n % 10) as u8;
buf[i] = (b'0' + digit) as c_char;
n /= 10;
if n == 0 {
break;
}
}
luaL_addlstring(B, buf.as_ptr().add(i), 32 - i);
}
+8
View File
@@ -203,6 +203,7 @@ unsafe extern "C-unwind" {
pub fn lua_rawget(L: *mut lua_State, idx: c_int) -> c_int;
#[link_name = "lua_rawgeti"]
pub fn lua_rawgeti_(L: *mut lua_State, idx: c_int, n: c_int) -> c_int;
pub fn lua_rawgetptagged(L: *mut lua_State, idx: c_int, p: *const c_void, tag: c_int) -> c_int;
pub fn lua_createtable(L: *mut lua_State, narr: c_int, nrec: c_int);
pub fn lua_setreadonly(L: *mut lua_State, idx: c_int, enabled: c_int);
@@ -220,6 +221,7 @@ unsafe extern "C-unwind" {
pub fn lua_rawset(L: *mut lua_State, idx: c_int);
#[link_name = "lua_rawseti"]
pub fn lua_rawseti_(L: *mut lua_State, idx: c_int, n: c_int);
pub fn lua_rawsetptagged(L: *mut lua_State, idx: c_int, p: *const c_void, tag: c_int);
pub fn lua_setmetatable(L: *mut lua_State, objindex: c_int) -> c_int;
pub fn lua_setfenv(L: *mut lua_State, idx: c_int) -> c_int;
@@ -235,6 +237,7 @@ unsafe extern "C-unwind" {
) -> c_int;
pub fn lua_call(L: *mut lua_State, nargs: c_int, nresults: c_int);
pub fn lua_pcall(L: *mut lua_State, nargs: c_int, nresults: c_int, errfunc: c_int) -> c_int;
pub fn lua_cpcall(L: *mut lua_State, f: lua_CFunction, ud: *mut c_void) -> c_int;
//
// Coroutine functions
@@ -544,4 +547,9 @@ unsafe extern "C" {
unsafe extern "C" {
pub fn luau_setfflag(name: *const c_char, value: c_int) -> c_int;
pub fn lua_getmetatablepointer(L: *mut lua_State, idx: c_int) -> *const c_void;
pub fn lua_gcdump(
L: *mut lua_State,
file: *mut c_void,
category_name: Option<unsafe extern "C" fn(L: *mut lua_State, memcat: u8) -> *const c_char>,
);
}
+37 -7
View File
@@ -23,6 +23,16 @@ pub enum luarequire_WriteResult {
Failure,
}
/// Represents whether a configuration file is present, and if so, its syntax.
#[repr(C)]
pub enum luarequire_ConfigStatus {
Absent,
// Signals the presence of multiple configuration files
Ambiguous,
PresentJson,
PresentLuau,
}
#[repr(C)]
pub struct luarequire_Configuration {
// Returns whether requires are permitted from the given chunkname.
@@ -48,6 +58,17 @@ pub struct luarequire_Configuration {
path: *const c_char,
) -> luarequire_NavigateResult,
// Provides a final override opportunity if an alias cannot be found in configuration files. If
// NAVIGATE_SUCCESS is returned, this must update the internal state to point at the aliased module.
// Can be left undefined.
pub to_alias_fallback: Option<
unsafe extern "C-unwind" fn(
L: *mut lua_State,
ctx: *mut c_void,
alias_unprefixed: *const c_char,
) -> luarequire_NavigateResult,
>,
// Navigates through the context by making mutations to the internal state.
pub to_parent:
unsafe extern "C-unwind" fn(L: *mut lua_State, ctx: *mut c_void) -> luarequire_NavigateResult,
@@ -90,13 +111,14 @@ pub struct luarequire_Configuration {
size_out: *mut usize,
) -> luarequire_WriteResult,
// Returns whether a configuration file is present in the current context.
// If not, require-by-string will call to_parent until either a configuration file is present or
// Returns whether a configuration file is present in the current context, and if so, its syntax.
// If not present, require-by-string will call to_parent until either a configuration file is present or
// NAVIGATE_FAILURE is returned (at root).
pub is_config_present: unsafe extern "C-unwind" fn(L: *mut lua_State, ctx: *mut c_void) -> bool,
pub get_config_status:
unsafe extern "C-unwind" fn(L: *mut lua_State, ctx: *mut c_void) -> luarequire_ConfigStatus,
// Parses the configuration file in the current context for the given alias and returns its
// value or WRITE_FAILURE if not found. This function is only called if is_config_present
// value or WRITE_FAILURE if not found. This function is only called if get_config_status
// returns true. If this function pointer is set, get_config must not be set. Opting in to this
// function pointer disables parsing configuration files internally and can be used for finer
// control over the configuration file parsing process.
@@ -111,9 +133,10 @@ pub struct luarequire_Configuration {
) -> luarequire_WriteResult,
>,
// Provides the contents of the configuration file in the current context. This function is only called
// if is_config_present returns true. If this function pointer is set, get_alias must not be set. Opting
// in to this function pointer enables parsing configuration files internally.
// Provides the contents of the configuration file in the current context.
// This function is only called if get_config_status does not return CONFIG_ABSENT. If this function
// pointer is set, get_alias must not be set. Opting in to this function pointer enables parsing
// configuration files internally.
pub get_config: Option<
unsafe extern "C-unwind" fn(
L: *mut lua_State,
@@ -124,6 +147,13 @@ pub struct luarequire_Configuration {
) -> luarequire_WriteResult,
>,
// Returns the maximum number of milliseconds to allow for executing a given Luau-syntax configuration
// file. This function is only called if get_config_status returns CONFIG_PRESENT_LUAU and can be left
// undefined if support for Luau-syntax configuration files is not needed. A default value of 2000ms is
// used. Negative values are treated as infinite.
pub get_luau_config_timeout:
Option<unsafe extern "C-unwind" fn(L: *mut lua_State, ctx: *mut c_void) -> c_int>,
// Executes the module and places the result on the stack. Returns the number of results placed on the
// stack.
// Returning -1 directs the requiring thread to yield. In this case, this thread should be resumed with
+79 -4
View File
@@ -1,3 +1,5 @@
use std::io;
#[cfg(feature = "serde")]
use serde::ser::{Serialize, Serializer};
@@ -50,13 +52,18 @@ impl Buffer {
#[track_caller]
pub fn write_bytes(&self, offset: usize, bytes: &[u8]) {
let lua = self.0.lua.lock();
let data = unsafe {
let (buf, size) = self.as_raw_parts(&lua);
std::slice::from_raw_parts_mut(buf, size)
};
let data = self.as_slice_mut(&lua);
data[offset..offset + bytes.len()].copy_from_slice(bytes);
}
/// Returns an adaptor implementing [`io::Read`], [`io::Write`] and [`io::Seek`] over the
/// buffer.
///
/// Buffer operations are infallible, none of the read/write functions will return a Err.
pub fn cursor(self) -> impl io::Read + io::Write + io::Seek {
BufferCursor(self, 0)
}
pub(crate) fn as_slice(&self, lua: &RawLua) -> &[u8] {
unsafe {
let (buf, size) = self.as_raw_parts(lua);
@@ -64,6 +71,14 @@ impl Buffer {
}
}
#[allow(clippy::mut_from_ref)]
fn as_slice_mut(&self, lua: &RawLua) -> &mut [u8] {
unsafe {
let (buf, size) = self.as_raw_parts(lua);
std::slice::from_raw_parts_mut(buf, size)
}
}
#[cfg(feature = "luau")]
unsafe fn as_raw_parts(&self, lua: &RawLua) -> (*mut u8, usize) {
let mut size = 0usize;
@@ -78,6 +93,66 @@ impl Buffer {
}
}
struct BufferCursor(Buffer, usize);
impl io::Read for BufferCursor {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
let lua = self.0 .0.lua.lock();
let data = self.0.as_slice(&lua);
if self.1 == data.len() {
return Ok(0);
}
let len = buf.len().min(data.len() - self.1);
buf[..len].copy_from_slice(&data[self.1..self.1 + len]);
self.1 += len;
Ok(len)
}
}
impl io::Write for BufferCursor {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
let lua = self.0 .0.lua.lock();
let data = self.0.as_slice_mut(&lua);
if self.1 == data.len() {
return Ok(0);
}
let len = buf.len().min(data.len() - self.1);
data[self.1..self.1 + len].copy_from_slice(&buf[..len]);
self.1 += len;
Ok(len)
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
impl io::Seek for BufferCursor {
fn seek(&mut self, pos: io::SeekFrom) -> io::Result<u64> {
let lua = self.0 .0.lua.lock();
let data = self.0.as_slice(&lua);
let new_offset = match pos {
io::SeekFrom::Start(offset) => offset as i64,
io::SeekFrom::End(offset) => data.len() as i64 + offset,
io::SeekFrom::Current(offset) => self.1 as i64 + offset,
};
if new_offset < 0 {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"invalid seek to a negative position",
));
}
if new_offset as usize > data.len() {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
"invalid seek to a position beyond the end of the buffer",
));
}
self.1 = new_offset as usize;
Ok(self.1 as u64)
}
}
#[cfg(feature = "serde")]
impl Serialize for Buffer {
fn serialize<S: Serializer>(&self, serializer: S) -> std::result::Result<S::Ok, S::Error> {
+3 -7
View File
@@ -18,7 +18,7 @@ use {
crate::traits::LuaNativeAsyncFn,
crate::types::AsyncCallback,
std::future::{self, Future},
std::pin::Pin,
std::pin::{pin, Pin},
std::task::{Context, Poll},
};
@@ -669,13 +669,9 @@ impl<R: FromLuaMulti> Future for AsyncCallFuture<R> {
type Output = Result<R>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
// Safety: We're not moving any pinned data
let this = unsafe { self.get_unchecked_mut() };
let this = self.get_mut();
match &mut this.0 {
Ok(thread) => {
let pinned_thread = unsafe { Pin::new_unchecked(thread) };
pinned_thread.poll(cx)
}
Ok(thread) => pin!(thread).poll(cx),
Err(err) => Poll::Ready(Err(err.clone())),
}
}
+5 -2
View File
@@ -132,7 +132,7 @@ pub use crate::{
buffer::Buffer,
chunk::{CompileConstant, Compiler},
function::CoverageInfo,
luau::{NavigateError, Require, TextRequirer},
luau::{HeapDump, NavigateError, Require, TextRequirer},
vector::Vector,
};
@@ -142,7 +142,10 @@ pub use crate::{thread::AsyncThread, traits::LuaNativeAsyncFn};
#[cfg(feature = "serde")]
#[doc(inline)]
pub use crate::serde::{de::Options as DeserializeOptions, ser::Options as SerializeOptions, LuaSerdeExt};
pub use crate::{
serde::{de::Options as DeserializeOptions, ser::Options as SerializeOptions, LuaSerdeExt},
value::SerializableValue,
};
#[cfg(feature = "serde")]
#[cfg_attr(docsrs, doc(cfg(feature = "serde")))]
+178
View File
@@ -0,0 +1,178 @@
use std::collections::HashMap;
use std::hash::Hash;
use std::mem;
use std::os::raw::c_char;
use crate::state::ExtraData;
use super::json::{self, Json};
/// Represents a heap dump of a Luau memory state.
#[cfg(any(feature = "luau", doc))]
#[cfg_attr(docsrs, doc(cfg(feature = "luau")))]
pub struct HeapDump {
data: Json<'static>, // refers to the contents of `buf`
buf: Box<str>,
}
impl HeapDump {
/// Dumps the current Lua heap state.
pub(crate) unsafe fn new(state: *mut ffi::lua_State) -> Option<Self> {
unsafe extern "C" fn category_name(state: *mut ffi::lua_State, cat: u8) -> *const c_char {
(&*ExtraData::get(state))
.mem_categories
.get(cat as usize)
.map(|s| s.as_ptr())
.unwrap_or(cstr!("unknown"))
}
let mut buf = Vec::new();
unsafe {
let file = libc::tmpfile();
if file.is_null() {
return None;
}
ffi::lua_gcdump(state, file as *mut _, Some(category_name));
libc::fseek(file, 0, libc::SEEK_END);
let len = libc::ftell(file) as usize;
libc::rewind(file);
if len > 0 {
buf.reserve(len);
libc::fread(buf.as_mut_ptr() as *mut _, 1, len, file);
buf.set_len(len);
}
libc::fclose(file);
}
let buf = String::from_utf8(buf).ok()?.into_boxed_str();
let data = json::parse(unsafe { mem::transmute::<&str, &'static str>(&buf) }).ok()?;
Some(HeapDump { data, buf })
}
/// Returns the raw JSON representation of the heap dump.
///
/// The JSON structure is an internal detail and may change in future versions.
#[doc(hidden)]
pub fn to_json(&self) -> &str {
&self.buf
}
/// Returns the total size of the Lua heap in bytes.
pub fn size(&self) -> u64 {
self.data["stats"]["size"].as_u64().unwrap_or_default()
}
/// Returns a mapping from object type to (count, total size in bytes).
///
/// If `category` is provided, only objects in that category are considered.
pub fn size_by_type<'a>(&'a self, category: Option<&str>) -> HashMap<&'a str, (usize, u64)> {
self.size_by_type_inner(category).unwrap_or_default()
}
fn size_by_type_inner<'a>(&'a self, category: Option<&str>) -> Option<HashMap<&'a str, (usize, u64)>> {
let category_id = match category {
// If we cannot find the category, return empty result
Some(cat) => Some(self.find_category_id(cat)?),
None => None,
};
let mut size_by_type = HashMap::new();
let objects = self.data["objects"].as_object()?;
for obj in objects.values() {
if let Some(cat_id) = category_id {
if obj["cat"].as_i64()? != cat_id {
continue;
}
}
update_size(&mut size_by_type, obj["type"].as_str()?, obj["size"].as_u64()?);
}
Some(size_by_type)
}
/// Returns a mapping from category name to total size in bytes.
pub fn size_by_category(&self) -> HashMap<&str, u64> {
let mut size_by_category = HashMap::new();
if let Some(categories) = self.data["stats"]["categories"].as_object() {
for cat in categories.values() {
if let Some(cat_name) = cat["name"].as_str() {
size_by_category.insert(cat_name, cat["size"].as_u64().unwrap_or_default());
}
}
}
size_by_category
}
/// Returns a mapping from userdata type to (count, total size in bytes).
pub fn size_by_userdata<'a>(&'a self, category: Option<&str>) -> HashMap<&'a str, (usize, u64)> {
self.size_by_userdata_inner(category).unwrap_or_default()
}
fn size_by_userdata_inner<'a>(
&'a self,
category: Option<&str>,
) -> Option<HashMap<&'a str, (usize, u64)>> {
let category_id = match category {
// If we cannot find the category, return empty result
Some(cat) => Some(self.find_category_id(cat)?),
None => None,
};
let mut size_by_userdata = HashMap::new();
let objects = self.data["objects"].as_object()?;
for obj in objects.values() {
if obj["type"] != "userdata" {
continue;
}
if let Some(cat_id) = category_id {
if obj["cat"].as_i64()? != cat_id {
continue;
}
}
// Determine userdata type from metatable
let mut ud_type = "unknown";
if let Some(metatable_addr) = obj["metatable"].as_str() {
if let Some(t) = get_key(objects, &objects[metatable_addr], "__type") {
ud_type = t;
}
}
update_size(&mut size_by_userdata, ud_type, obj["size"].as_u64()?);
}
Some(size_by_userdata)
}
/// Finds the category ID for a given category name.
fn find_category_id(&self, category: &str) -> Option<i64> {
let categories = self.data["stats"]["categories"].as_object()?;
for (cat_id, cat) in categories {
if cat["name"].as_str() == Some(category) {
return cat_id.parse().ok();
}
}
None
}
}
/// Updates the size mapping for a given key.
fn update_size<K: Eq + Hash>(size_type: &mut HashMap<K, (usize, u64)>, key: K, size: u64) {
let (ref mut count, ref mut total_size) = size_type.entry(key).or_insert((0, 0));
*count += 1;
*total_size += size;
}
/// Retrieves the value associated with a given `key` from a Lua table `tbl`.
fn get_key<'a>(objects: &'a HashMap<&'a str, Json>, tbl: &Json, key: &str) -> Option<&'a str> {
let pairs = tbl["pairs"].as_array()?;
for kv in pairs.chunks_exact(2) {
#[rustfmt::skip]
let (Some(key_addr), Some(val_addr)) = (kv[0].as_str(), kv[1].as_str()) else { continue; };
if objects[key_addr]["type"] == "string" && objects[key_addr]["data"].as_str() == Some(key) {
if objects[val_addr]["type"] == "string" {
return objects[val_addr]["data"].as_str();
} else {
break;
}
}
}
None
}
+327
View File
@@ -0,0 +1,327 @@
use std::array;
use std::collections::HashMap;
use std::iter::Peekable;
use std::ops::Index;
use std::str::CharIndices;
// A simple JSON parser and representation.
// This parser supports only a subset of JSON specification and is intended for Luau's use cases.
#[derive(Debug, PartialEq)]
pub(crate) enum Json<'a> {
Null,
Bool(bool),
Integer(i64),
Number(f64),
String(&'a str),
Array(Vec<Json<'a>>),
Object(HashMap<&'a str, Json<'a>>),
}
impl<'a> Index<&str> for Json<'a> {
type Output = Json<'a>;
fn index(&self, key: &str) -> &Self::Output {
match self {
Json::Object(map) => map.get(key).unwrap_or(&Json::Null),
_ => &Json::Null,
}
}
}
impl PartialEq<&str> for Json<'_> {
fn eq(&self, other: &&str) -> bool {
matches!(self, Json::String(s) if s == other)
}
}
impl<'a> Json<'a> {
pub(crate) fn as_str(&self) -> Option<&'a str> {
match self {
Json::String(s) => Some(s),
_ => None,
}
}
pub(crate) fn as_i64(&self) -> Option<i64> {
match self {
Json::Integer(i) => Some(*i),
Json::Number(n) if n.fract() == 0.0 => Some(*n as i64),
_ => None,
}
}
pub(crate) fn as_u64(&self) -> Option<u64> {
self.as_i64()
.and_then(|i| if i >= 0 { Some(i as u64) } else { None })
}
pub(crate) fn as_array(&self) -> Option<&[Json<'a>]> {
match self {
Json::Array(arr) => Some(arr),
_ => None,
}
}
pub(crate) fn as_object(&self) -> Option<&HashMap<&'a str, Json<'a>>> {
match self {
Json::Object(map) => Some(map),
_ => None,
}
}
}
pub(crate) fn parse<'a>(s: &'a str) -> Result<Json<'a>, &'static str> {
let s = s.trim_ascii();
let mut chars = s.char_indices().peekable();
let value = parse_value(s, &mut chars)?;
Ok(value)
}
fn parse_value<'a>(s: &'a str, chars: &mut Peekable<CharIndices>) -> Result<Json<'a>, &'static str> {
skip_whitespace(chars);
match chars.peek() {
Some((_, '{')) => parse_object(s, chars),
Some((_, '[')) => parse_array(s, chars),
Some((_, '"')) => parse_string(s, chars).map(Json::String),
Some((_, 't' | 'f')) => parse_bool(chars),
Some((_, 'n')) => parse_null(chars),
Some((_, '-' | '0'..='9')) => parse_number(chars),
Some(_) => Err("unexpected character"),
None => Err("unexpected end of input"),
}
}
fn parse_object<'a>(s: &'a str, chars: &mut Peekable<CharIndices>) -> Result<Json<'a>, &'static str> {
chars.next(); // consume '{'
let mut map = HashMap::new();
skip_whitespace(chars);
if matches!(chars.peek(), Some((_, '}'))) {
chars.next();
return Ok(Json::Object(map));
}
loop {
skip_whitespace(chars);
let key = parse_string(s, chars)?;
skip_whitespace(chars);
if !matches!(chars.next(), Some((_, ':'))) {
return Err("expected ':'");
}
let value = parse_value(s, chars)?;
map.insert(key, value);
skip_whitespace(chars);
match chars.next() {
Some((_, ',')) => continue,
Some((_, '}')) => break,
_ => return Err("expected ',' or '}'"),
}
}
Ok(Json::Object(map))
}
fn parse_array<'a>(s: &'a str, chars: &mut Peekable<CharIndices>) -> Result<Json<'a>, &'static str> {
chars.next(); // consume '['
let mut arr = Vec::new();
skip_whitespace(chars);
if matches!(chars.peek(), Some((_, ']'))) {
chars.next();
return Ok(Json::Array(arr));
}
loop {
skip_whitespace(chars);
arr.push(parse_value(s, chars)?);
skip_whitespace(chars);
match chars.next() {
Some((_, ',')) => continue,
Some((_, ']')) => return Ok(Json::Array(arr)),
_ => return Err("expected ',' or ']'"),
}
}
}
fn parse_string<'a>(s: &'a str, chars: &mut Peekable<CharIndices>) -> Result<&'a str, &'static str> {
if !matches!(chars.next(), Some((_, '"'))) {
return Err("expected string starting with '\"'");
}
let start = chars.peek().map(|(i, _)| *i).unwrap_or(0);
for (i, c) in chars {
if c == '"' {
return Ok(&s[start..i]);
}
}
Err("unterminated string")
}
fn parse_number(chars: &mut Peekable<CharIndices>) -> Result<Json<'static>, &'static str> {
let mut is_float = false;
let mut num = String::new();
while let Some((_, c @ ('0'..='9' | '-' | '.' | 'e' | 'E' | '+'))) = chars.peek() {
num.push(*c);
is_float = is_float || matches!(c, '.' | 'e' | 'E');
chars.next();
}
if !is_float {
let i = num.parse::<i64>().map_err(|_| "invalid integer")?;
return Ok(Json::Integer(i));
}
let n = num.parse::<f64>().map_err(|_| "invalid number")?;
Ok(Json::Number(n))
}
fn parse_bool(chars: &mut Peekable<CharIndices>) -> Result<Json<'static>, &'static str> {
let bool = next_chars(chars);
if bool == [Some('t'), Some('r'), Some('u'), Some('e')] {
return Ok(Json::Bool(true));
}
if bool == [Some('f'), Some('a'), Some('l'), Some('s')] && matches!(chars.next(), Some((_, 'e'))) {
return Ok(Json::Bool(false));
}
Err("invalid boolean literal")
}
fn parse_null(chars: &mut Peekable<CharIndices>) -> Result<Json<'static>, &'static str> {
if next_chars(chars) == [Some('n'), Some('u'), Some('l'), Some('l')] {
return Ok(Json::Null);
}
Err("invalid \"null\" literal")
}
fn skip_whitespace(chars: &mut Peekable<CharIndices>) {
while let Some((_, ' ' | '\n' | '\r' | '\t')) = chars.peek() {
chars.next();
}
}
fn next_chars<const N: usize>(chars: &mut Peekable<CharIndices>) -> [Option<char>; N] {
array::from_fn(|_| chars.next().map(|(_, c)| c))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse() {
assert_eq!(parse("null").unwrap(), Json::Null);
assert_eq!(parse("true").unwrap(), Json::Bool(true));
assert_eq!(parse("false").unwrap(), Json::Bool(false));
assert_eq!(parse("42").unwrap(), Json::Integer(42));
assert_eq!(parse("42.0").unwrap(), Json::Number(42.0));
assert_eq!(parse(r#""hello""#).unwrap(), Json::String("hello"));
assert_eq!(
parse("[1,2.0,3]").unwrap(),
Json::Array(vec![Json::Integer(1), Json::Number(2.0), Json::Integer(3)])
);
let mut obj = HashMap::new();
obj.insert("key", Json::String("value"));
assert_eq!(parse(r#"{"key":"value"}"#).unwrap(), Json::Object(obj));
}
#[test]
fn test_whitespace_handling() {
assert_eq!(parse(" null ").unwrap(), Json::Null);
assert_eq!(parse(" true ").unwrap(), Json::Bool(true));
assert_eq!(
parse(" [ 1 , 2.0 , 3 ] ").unwrap(),
Json::Array(vec![Json::Integer(1), Json::Number(2.0), Json::Integer(3)])
);
let mut obj = HashMap::new();
obj.insert("key", Json::String("value"));
assert_eq!(parse(r#" { "key" : "value" } "#).unwrap(), Json::Object(obj));
}
#[test]
fn test_empty_collections() {
assert_eq!(parse("[]").unwrap(), Json::Array(vec![]));
assert_eq!(parse("{}").unwrap(), Json::Object(HashMap::new()));
assert_eq!(parse("[ ]").unwrap(), Json::Array(vec![]));
assert_eq!(parse("{ }").unwrap(), Json::Object(HashMap::new()));
}
#[test]
fn test_nested_structures() {
assert_eq!(
parse(r#"{"nested":{"inner":"value"}}"#).unwrap(),
Json::Object({
let mut outer = HashMap::new();
let mut inner = HashMap::new();
inner.insert("inner", Json::String("value"));
outer.insert("nested", Json::Object(inner));
outer
})
);
assert_eq!(
parse("[[1,2],[3,4]]").unwrap(),
Json::Array(vec![
Json::Array(vec![Json::Integer(1), Json::Integer(2)]),
Json::Array(vec![Json::Integer(3), Json::Integer(4)])
])
);
}
#[test]
fn test_numbers() {
assert_eq!(parse("0").unwrap(), Json::Integer(0));
assert_eq!(parse("-42").unwrap(), Json::Integer(-42));
assert_eq!(parse("3.14").unwrap(), Json::Number(3.14));
assert_eq!(parse("-3.14").unwrap(), Json::Number(-3.14));
assert_eq!(parse("1e10").unwrap(), Json::Number(1e10));
assert_eq!(parse("1E10").unwrap(), Json::Number(1E10));
assert_eq!(parse("1e-10").unwrap(), Json::Number(1e-10));
assert_eq!(parse("1.5e+10").unwrap(), Json::Number(1.5e+10));
}
#[test]
fn test_strings() {
assert_eq!(parse(r#""""#).unwrap(), Json::String(""));
assert_eq!(parse(r#""hello world""#).unwrap(), Json::String("hello world"));
assert_eq!(
parse(r#""with spaces and 123""#).unwrap(),
Json::String("with spaces and 123")
);
}
#[test]
fn test_mixed_array() {
assert_eq!(
parse(r#"[null, true, false, 35.1, 42, "text", [], {}]"#).unwrap(),
Json::Array(vec![
Json::Null,
Json::Bool(true),
Json::Bool(false),
Json::Number(35.1),
Json::Integer(42),
Json::String("text"),
Json::Array(vec![]),
Json::Object(HashMap::new())
])
);
}
#[test]
fn test_object_multiple_keys() {
let mut obj = HashMap::new();
obj.insert("a", Json::Integer(1));
obj.insert("b", Json::Bool(true));
obj.insert("c", Json::Null);
assert_eq!(parse(r#"{"a":1,"b":true,"c":null}"#).unwrap(), Json::Object(obj));
}
#[test]
fn test_error_cases() {
assert!(parse("").is_err());
assert!(parse("nul").is_err());
assert!(parse("tru").is_err());
assert!(parse("fals").is_err());
assert!(parse(r#""unterminated"#).is_err());
assert!(parse("[1,2,]").is_err());
assert!(parse(r#"{"key""#).is_err());
assert!(parse(r#"{"key":"value""#).is_err());
assert!(parse(r#"{"key":"value",}"#).is_err());
assert!(parse("invalid").is_err());
assert!(parse("[1 2]").is_err());
assert!(parse(r#"{"key":"value" "key2":"value2"}"#).is_err());
}
}
+55 -2
View File
@@ -1,14 +1,15 @@
use std::ffi::CStr;
use std::ffi::{CStr, CString};
use std::os::raw::c_int;
use std::ptr;
use crate::chunk::ChunkMode;
use crate::error::Result;
use crate::error::{Error, Result};
use crate::function::Function;
use crate::state::{callback_error_ext, ExtraData, Lua};
use crate::traits::{FromLuaMulti, IntoLua};
use crate::types::MaybeSend;
pub use heap_dump::HeapDump;
pub use require::{NavigateError, Require, TextRequirer};
// Since Luau has some missing standard functions, we re-implement them here
@@ -22,6 +23,56 @@ impl Lua {
require::create_require_function(self, require)
}
/// Set the memory category for subsequent allocations from this Lua state.
///
/// The category "main" is reserved for the default memory category.
/// Maximum of 255 categories can be registered.
/// The category is set per Lua thread (state) and affects all allocations made from that
/// thread.
///
/// Return error if too many categories are registered or if the category name is invalid.
///
/// See [`Lua::heap_dump`] for tracking memory usage by category.
#[cfg(any(feature = "luau", doc))]
#[cfg_attr(docsrs, doc(cfg(feature = "luau")))]
pub fn set_memory_category(&self, category: &str) -> Result<()> {
let lua = self.lock();
if category.contains(|c| !matches!(c, 'a'..='z' | 'A'..='Z' | '0'..='9' | '-' | '_')) {
return Err(Error::runtime("invalid memory category name"));
}
let cat_id = unsafe {
let extra = ExtraData::get(lua.state());
match ((*extra).mem_categories.iter().enumerate())
.find(|&(_, name)| name.as_bytes() == category.as_bytes())
{
Some((id, _)) => id as u8,
None => {
let new_id = (*extra).mem_categories.len() as u8;
if new_id == 255 {
return Err(Error::runtime("too many memory categories registered"));
}
(*extra).mem_categories.push(CString::new(category).unwrap());
new_id
}
}
};
unsafe { ffi::lua_setmemcat(lua.state(), cat_id as i32) };
Ok(())
}
/// Dumps the current Lua VM heap state.
///
/// The returned `HeapDump` can be used to analyze memory usage.
/// It's recommended to call [`Lua::gc_collect`] before dumping the heap.
#[cfg(any(feature = "luau", doc))]
#[cfg_attr(docsrs, doc(cfg(feature = "luau")))]
pub fn heap_dump(&self) -> Result<HeapDump> {
let lua = self.lock();
unsafe { heap_dump::HeapDump::new(lua.state()).ok_or_else(|| Error::runtime("unable to dump heap")) }
}
pub(crate) unsafe fn configure_luau(&self) -> Result<()> {
let globals = self.globals();
@@ -96,4 +147,6 @@ unsafe extern "C-unwind" fn lua_loadstring(state: *mut ffi::lua_State) -> c_int
})
}
mod heap_dump;
mod json;
mod require;
+75 -276
View File
@@ -1,12 +1,10 @@
use std::cell::RefCell;
use std::collections::VecDeque;
use std::ffi::CStr;
use std::io::Result as IoResult;
use std::ops::{Deref, DerefMut};
use std::os::raw::{c_char, c_int, c_void};
use std::path::{Component, Path, PathBuf};
use std::result::Result as StdResult;
use std::{env, fmt, fs, mem, ptr};
use std::{fmt, mem, ptr};
use crate::error::{Error, Result};
use crate::function::Function;
@@ -14,9 +12,10 @@ use crate::state::{callback_error_ext, Lua};
use crate::table::Table;
use crate::types::MaybeSend;
// TODO: Rename to FsRequirer
pub use fs::TextRequirer;
/// An error that can occur during navigation in the Luau `require-by-string` system.
#[cfg(any(feature = "luau", doc))]
#[cfg_attr(docsrs, doc(cfg(feature = "luau")))]
#[derive(Debug, Clone)]
pub enum NavigateError {
Ambiguous,
@@ -50,9 +49,10 @@ impl From<Error> for NavigateError {
#[cfg(feature = "luau")]
type WriteResult = ffi::luarequire_WriteResult;
#[cfg(feature = "luau")]
type ConfigStatus = ffi::luarequire_ConfigStatus;
/// A trait for handling modules loading and navigation in the Luau `require-by-string` system.
#[cfg(any(feature = "luau", doc))]
#[cfg_attr(docsrs, doc(cfg(feature = "luau")))]
pub trait Require {
/// Returns `true` if "require" is permitted for the given chunk name.
fn is_require_allowed(&self, chunk_name: &str) -> bool;
@@ -73,7 +73,7 @@ pub trait Require {
/// Navigate to the given child directory.
fn to_child(&mut self, name: &str) -> StdResult<(), NavigateError>;
/// Returns whether the context is currently pointing at a module
/// Returns whether the context is currently pointing at a module.
fn has_module(&self) -> bool;
/// Provides a cache key representing the current module.
@@ -103,226 +103,31 @@ impl fmt::Debug for dyn Require {
}
}
/// The standard implementation of Luau `require-by-string` navigation.
#[derive(Default, Debug)]
pub struct TextRequirer {
/// An absolute path to the current Luau module (not mapped to a physical file)
abs_path: PathBuf,
/// A relative path to the current Luau module (not mapped to a physical file)
rel_path: PathBuf,
/// A physical path to the current Luau module, which is a file or a directory with an
/// `init.lua(u)` file
resolved_path: Option<PathBuf>,
struct Context {
require: Box<dyn Require>,
config_cache: Option<IoResult<Vec<u8>>>,
}
impl TextRequirer {
/// The prefix used for chunk names in the require system.
/// Only chunk names starting with this prefix are allowed to be used in `require`.
const CHUNK_PREFIX: &str = "@";
/// The file extensions that are considered valid for Luau modules.
const FILE_EXTENSIONS: &[&str] = &["luau", "lua"];
/// Creates a new `TextRequirer` instance.
pub fn new() -> Self {
Self::default()
}
fn normalize_chunk_name(chunk_name: &str) -> &str {
if let Some((path, line)) = chunk_name.split_once(':') {
if line.parse::<u32>().is_ok() {
return path;
}
}
chunk_name
}
// Normalizes the path by removing unnecessary components
fn normalize_path(path: &Path) -> PathBuf {
let mut components = VecDeque::new();
for comp in path.components() {
match comp {
Component::Prefix(..) | Component::RootDir => {
components.push_back(comp);
}
Component::CurDir => {}
Component::ParentDir => {
if matches!(components.back(), None | Some(Component::ParentDir)) {
components.push_back(Component::ParentDir);
} else if matches!(components.back(), Some(Component::Normal(..))) {
components.pop_back();
}
}
Component::Normal(..) => components.push_back(comp),
}
}
if matches!(components.front(), None | Some(Component::Normal(..))) {
components.push_front(Component::CurDir);
}
// Join the components back together
components.into_iter().collect()
}
/// Resolve a Luau module path to a physical file or directory.
///
/// Empty directories without init files are considered valid as "intermediate" directories.
fn resolve_module(path: &Path) -> StdResult<Option<PathBuf>, NavigateError> {
let mut found_path = None;
if path.components().next_back() != Some(Component::Normal("init".as_ref())) {
let current_ext = (path.extension().and_then(|s| s.to_str()))
.map(|s| format!("{s}."))
.unwrap_or_default();
for ext in Self::FILE_EXTENSIONS {
let candidate = path.with_extension(format!("{current_ext}{ext}"));
if candidate.is_file() && found_path.replace(candidate).is_some() {
return Err(NavigateError::Ambiguous);
}
}
}
if path.is_dir() {
for component in Self::FILE_EXTENSIONS.iter().map(|ext| format!("init.{ext}")) {
let candidate = path.join(component);
if candidate.is_file() && found_path.replace(candidate).is_some() {
return Err(NavigateError::Ambiguous);
}
}
if found_path.is_none() {
// Directories without init files are considered valid "intermediate" path
return Ok(None);
}
}
Ok(Some(found_path.ok_or(NavigateError::NotFound)?))
}
}
impl Require for TextRequirer {
fn is_require_allowed(&self, chunk_name: &str) -> bool {
chunk_name.starts_with(Self::CHUNK_PREFIX)
}
fn reset(&mut self, chunk_name: &str) -> StdResult<(), NavigateError> {
if !chunk_name.starts_with(Self::CHUNK_PREFIX) {
return Err(NavigateError::NotFound);
}
let chunk_name = Self::normalize_chunk_name(&chunk_name[1..]);
let chunk_path = Self::normalize_path(chunk_name.as_ref());
if chunk_path.extension() == Some("rs".as_ref()) {
// Special case for Rust source files, reset to the current directory
let chunk_filename = chunk_path.file_name().unwrap();
let cwd = env::current_dir().map_err(|_| NavigateError::NotFound)?;
self.abs_path = Self::normalize_path(&cwd.join(chunk_filename));
self.rel_path = ([Component::CurDir, Component::Normal(chunk_filename)].into_iter()).collect();
self.resolved_path = None;
return Ok(());
}
if chunk_path.is_absolute() {
let resolved_path = Self::resolve_module(&chunk_path)?;
self.abs_path = chunk_path.clone();
self.rel_path = chunk_path;
self.resolved_path = resolved_path;
} else {
// Relative path
let cwd = env::current_dir().map_err(|_| NavigateError::NotFound)?;
let abs_path = Self::normalize_path(&cwd.join(&chunk_path));
let resolved_path = Self::resolve_module(&abs_path)?;
self.abs_path = abs_path;
self.rel_path = chunk_path;
self.resolved_path = resolved_path;
}
Ok(())
}
fn jump_to_alias(&mut self, path: &str) -> StdResult<(), NavigateError> {
let path = Self::normalize_path(path.as_ref());
let resolved_path = Self::resolve_module(&path)?;
self.abs_path = path.clone();
self.rel_path = path;
self.resolved_path = resolved_path;
Ok(())
}
fn to_parent(&mut self) -> StdResult<(), NavigateError> {
let mut abs_path = self.abs_path.clone();
if !abs_path.pop() {
// It's important to return `NotFound` if we reached the root, as it's a "recoverable" error if we
// cannot go beyond the root directory.
// Luau "require-by-string` has a special logic to search for config file to resolve aliases.
return Err(NavigateError::NotFound);
}
let mut rel_parent = self.rel_path.clone();
rel_parent.pop();
let resolved_path = Self::resolve_module(&abs_path)?;
self.abs_path = abs_path;
self.rel_path = Self::normalize_path(&rel_parent);
self.resolved_path = resolved_path;
Ok(())
}
fn to_child(&mut self, name: &str) -> StdResult<(), NavigateError> {
let abs_path = self.abs_path.join(name);
let rel_path = self.rel_path.join(name);
let resolved_path = Self::resolve_module(&abs_path)?;
self.abs_path = abs_path;
self.rel_path = rel_path;
self.resolved_path = resolved_path;
Ok(())
}
fn has_module(&self) -> bool {
(self.resolved_path.as_deref())
.map(Path::is_file)
.unwrap_or(false)
}
fn cache_key(&self) -> String {
self.resolved_path.as_deref().unwrap().display().to_string()
}
fn has_config(&self) -> bool {
self.abs_path.is_dir() && self.abs_path.join(".luaurc").is_file()
}
fn config(&self) -> IoResult<Vec<u8>> {
fs::read(self.abs_path.join(".luaurc"))
}
fn loader(&self, lua: &Lua) -> Result<Function> {
let name = format!("@{}", self.rel_path.display());
lua.load(self.resolved_path.as_deref().unwrap())
.set_name(name)
.into_function()
}
}
struct Context(Box<dyn Require>);
impl Deref for Context {
type Target = dyn Require;
fn deref(&self) -> &Self::Target {
&*self.0
&*self.require
}
}
impl DerefMut for Context {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut *self.0
&mut *self.require
}
}
impl Context {
fn new(require: impl Require + MaybeSend + 'static) -> Self {
Context {
require: Box::new(require),
config_cache: None,
}
}
}
@@ -447,9 +252,18 @@ pub(super) unsafe extern "C-unwind" fn init_config(config: *mut ffi::luarequire_
write_to_buffer(buffer, buffer_size, size_out, cache_key.as_bytes())
}
unsafe extern "C-unwind" fn is_config_present(state: *mut ffi::lua_State, ctx: *mut c_void) -> bool {
let this = try_borrow!(state, ctx);
this.has_config()
unsafe extern "C-unwind" fn get_config_status(
state: *mut ffi::lua_State,
ctx: *mut c_void,
) -> ConfigStatus {
let mut this = try_borrow_mut!(state, ctx);
if this.has_config() {
this.config_cache = Some(this.config());
if let Some(Ok(data)) = &this.config_cache {
return detect_config_format(data);
}
}
ConfigStatus::Absent
}
unsafe extern "C-unwind" fn get_config(
@@ -459,8 +273,10 @@ pub(super) unsafe extern "C-unwind" fn init_config(config: *mut ffi::luarequire_
buffer_size: usize,
size_out: *mut usize,
) -> WriteResult {
let this = try_borrow!(state, ctx);
let config = callback_error_ext(state, ptr::null_mut(), true, move |_, _| Ok(this.config()?));
let mut this = try_borrow_mut!(state, ctx);
let config = callback_error_ext(state, ptr::null_mut(), true, move |_, _| {
Ok(this.config_cache.take().unwrap_or_else(|| this.config())?)
});
write_to_buffer(buffer, buffer_size, size_out, &config)
}
@@ -483,18 +299,32 @@ pub(super) unsafe extern "C-unwind" fn init_config(config: *mut ffi::luarequire_
(*config).is_require_allowed = is_require_allowed;
(*config).reset = reset;
(*config).jump_to_alias = jump_to_alias;
(*config).to_alias_fallback = None;
(*config).to_parent = to_parent;
(*config).to_child = to_child;
(*config).is_module_present = is_module_present;
(*config).get_chunkname = get_chunkname;
(*config).get_loadname = get_loadname;
(*config).get_cache_key = get_cache_key;
(*config).is_config_present = is_config_present;
(*config).get_config_status = get_config_status;
(*config).get_alias = None;
(*config).get_config = Some(get_config);
(*config).load = load;
}
/// Detect configuration file format (JSON or Luau)
#[cfg(feature = "luau")]
fn detect_config_format(data: &[u8]) -> ConfigStatus {
let data = data.trim_ascii();
if data.starts_with(b"{") {
let data = &data[1..].trim_ascii_start();
if data.starts_with(b"\"") || data == b"}" {
return ConfigStatus::PresentJson;
}
}
ConfigStatus::PresentLuau
}
/// Helper function to write data to a buffer
#[cfg(feature = "luau")]
unsafe fn write_to_buffer(
@@ -545,7 +375,7 @@ pub(super) fn create_require_function<R: Require + MaybeSend + 'static>(
let (get_cache_key, find_current_file, proxyrequire, registered_modules, loader_cache) = unsafe {
lua.exec_raw::<(Function, Function, Function, Table, Table)>((), move |state| {
let context = Context(Box::new(require));
let context = Context::new(require);
let context_ptr = ffi::lua_newuserdata_t(state, RefCell::new(context));
ffi::lua_pushcclosured(state, get_cache_key, cstr!("get_cache_key"), 1);
ffi::lua_pushcfunctiond(state, find_current_file, cstr!("find_current_file"));
@@ -567,10 +397,26 @@ pub(super) fn create_require_function<R: Require + MaybeSend + 'static>(
1
}
let (error, r#type) = unsafe {
lua.exec_raw::<(Function, Function)>((), move |state| {
unsafe extern "C-unwind" fn to_lowercase(state: *mut ffi::lua_State) -> c_int {
let s = ffi::luaL_checkstring(state, 1);
let s = CStr::from_ptr(s);
if !s.to_bytes().iter().any(|&c| c.is_ascii_uppercase()) {
// If the string does not contain any uppercase ASCII letters, return it as is
return 1;
}
callback_error_ext(state, ptr::null_mut(), true, |extra, _| {
let s = (s.to_bytes().iter())
.map(|&c| c.to_ascii_lowercase())
.collect::<bstr::BString>();
(*extra).raw_lua().push(s).map(|_| 1)
})
}
let (error, r#type, to_lowercase) = unsafe {
lua.exec_raw::<(Function, Function, Function)>((), move |state| {
ffi::lua_pushcfunctiond(state, error, cstr!("error"));
ffi::lua_pushcfunctiond(state, r#type, cstr!("type"));
ffi::lua_pushcfunctiond(state, to_lowercase, cstr!("to_lowercase"));
})
}?;
@@ -583,6 +429,7 @@ pub(super) fn create_require_function<R: Require + MaybeSend + 'static>(
env.raw_set("LOADER_CACHE", loader_cache)?;
env.raw_set("error", error)?;
env.raw_set("type", r#type)?;
env.raw_set("to_lowercase", to_lowercase)?;
lua.load(
r#"
@@ -592,7 +439,7 @@ pub(super) fn create_require_function<R: Require + MaybeSend + 'static>(
end
-- Check if the module (path) is explicitly registered
local maybe_result = REGISTERED_MODULES[path]
local maybe_result = REGISTERED_MODULES[to_lowercase(path)]
if maybe_result ~= nil then
return maybe_result
end
@@ -620,52 +467,4 @@ pub(super) fn create_require_function<R: Require + MaybeSend + 'static>(
.into_function()
}
#[cfg(test)]
mod tests {
use std::path::Path;
use super::TextRequirer;
#[test]
fn test_path_normalize() {
for (input, expected) in [
// Basic formatting checks
("", "./"),
(".", "./"),
("a/relative/path", "./a/relative/path"),
// Paths containing extraneous '.' and '/' symbols
("./remove/extraneous/symbols/", "./remove/extraneous/symbols"),
("./remove/extraneous//symbols", "./remove/extraneous/symbols"),
("./remove/extraneous/symbols/.", "./remove/extraneous/symbols"),
("./remove/extraneous/./symbols", "./remove/extraneous/symbols"),
("../remove/extraneous/symbols/", "../remove/extraneous/symbols"),
("../remove/extraneous//symbols", "../remove/extraneous/symbols"),
("../remove/extraneous/symbols/.", "../remove/extraneous/symbols"),
("../remove/extraneous/./symbols", "../remove/extraneous/symbols"),
("/remove/extraneous/symbols/", "/remove/extraneous/symbols"),
("/remove/extraneous//symbols", "/remove/extraneous/symbols"),
("/remove/extraneous/symbols/.", "/remove/extraneous/symbols"),
("/remove/extraneous/./symbols", "/remove/extraneous/symbols"),
// Paths containing '..'
("./remove/me/..", "./remove"),
("./remove/me/../", "./remove"),
("../remove/me/..", "../remove"),
("../remove/me/../", "../remove"),
("/remove/me/..", "/remove"),
("/remove/me/../", "/remove"),
("./..", "../"),
("./../", "../"),
("../..", "../../"),
("../../", "../../"),
// '..' disappears if path is absolute and component is non-erasable
("/../", "/"),
] {
let path = TextRequirer::normalize_path(input.as_ref());
assert_eq!(
&path,
expected.as_ref() as &Path,
"wrong normalization for {input}"
);
}
}
}
mod fs;
+278
View File
@@ -0,0 +1,278 @@
use std::collections::VecDeque;
use std::io::Result as IoResult;
use std::path::{Component, Path, PathBuf};
use std::result::Result as StdResult;
use std::{env, fs};
use crate::error::Result;
use crate::function::Function;
use crate::state::Lua;
use super::{NavigateError, Require};
/// The standard implementation of Luau `require-by-string` navigation.
#[derive(Default, Debug)]
pub struct TextRequirer {
/// An absolute path to the current Luau module (not mapped to a physical file)
abs_path: PathBuf,
/// A relative path to the current Luau module (not mapped to a physical file)
rel_path: PathBuf,
/// A physical path to the current Luau module, which is a file or a directory with an
/// `init.lua(u)` file
resolved_path: Option<PathBuf>,
}
impl TextRequirer {
/// The prefix used for chunk names in the require system.
/// Only chunk names starting with this prefix are allowed to be used in `require`.
const CHUNK_PREFIX: &str = "@";
/// The file extensions that are considered valid for Luau modules.
const FILE_EXTENSIONS: &[&str] = &["luau", "lua"];
/// The filename for the JSON configuration file.
const LUAURC_CONFIG_FILENAME: &str = ".luaurc";
/// The filename for the Luau configuration file.
const LUAU_CONFIG_FILENAME: &str = ".config.luau";
/// Creates a new `TextRequirer` instance.
pub fn new() -> Self {
Self::default()
}
fn normalize_chunk_name(chunk_name: &str) -> &str {
if let Some((path, line)) = chunk_name.rsplit_once(':') {
if line.parse::<u32>().is_ok() {
return path;
}
}
chunk_name
}
// Normalizes the path by removing unnecessary components
fn normalize_path(path: &Path) -> PathBuf {
let mut components = VecDeque::new();
for comp in path.components() {
match comp {
Component::Prefix(..) | Component::RootDir => {
components.push_back(comp);
}
Component::CurDir => {}
Component::ParentDir => {
if matches!(components.back(), None | Some(Component::ParentDir)) {
components.push_back(Component::ParentDir);
} else if matches!(components.back(), Some(Component::Normal(..))) {
components.pop_back();
}
}
Component::Normal(..) => components.push_back(comp),
}
}
if matches!(components.front(), None | Some(Component::Normal(..))) {
components.push_front(Component::CurDir);
}
// Join the components back together
components.into_iter().collect()
}
/// Resolve a Luau module path to a physical file or directory.
///
/// Empty directories without init files are considered valid as "intermediate" directories.
fn resolve_module(path: &Path) -> StdResult<Option<PathBuf>, NavigateError> {
let mut found_path = None;
if path.components().next_back() != Some(Component::Normal("init".as_ref())) {
let current_ext = (path.extension().and_then(|s| s.to_str()))
.map(|s| format!("{s}."))
.unwrap_or_default();
for ext in Self::FILE_EXTENSIONS {
let candidate = path.with_extension(format!("{current_ext}{ext}"));
if candidate.is_file() && found_path.replace(candidate).is_some() {
return Err(NavigateError::Ambiguous);
}
}
}
if path.is_dir() {
for component in Self::FILE_EXTENSIONS.iter().map(|ext| format!("init.{ext}")) {
let candidate = path.join(component);
if candidate.is_file() && found_path.replace(candidate).is_some() {
return Err(NavigateError::Ambiguous);
}
}
if found_path.is_none() {
// Directories without init files are considered valid "intermediate" path
return Ok(None);
}
}
Ok(Some(found_path.ok_or(NavigateError::NotFound)?))
}
}
impl Require for TextRequirer {
fn is_require_allowed(&self, chunk_name: &str) -> bool {
chunk_name.starts_with(Self::CHUNK_PREFIX)
}
fn reset(&mut self, chunk_name: &str) -> StdResult<(), NavigateError> {
if !chunk_name.starts_with(Self::CHUNK_PREFIX) {
return Err(NavigateError::NotFound);
}
let chunk_name = Self::normalize_chunk_name(&chunk_name[1..]);
let chunk_path = Self::normalize_path(chunk_name.as_ref());
if chunk_path.extension() == Some("rs".as_ref()) {
// Special case for Rust source files, reset to the current directory
let chunk_filename = chunk_path.file_name().unwrap();
let cwd = env::current_dir().map_err(|_| NavigateError::NotFound)?;
self.abs_path = Self::normalize_path(&cwd.join(chunk_filename));
self.rel_path = ([Component::CurDir, Component::Normal(chunk_filename)].into_iter()).collect();
self.resolved_path = None;
return Ok(());
}
if chunk_path.is_absolute() {
let resolved_path = Self::resolve_module(&chunk_path)?;
self.abs_path = chunk_path.clone();
self.rel_path = chunk_path;
self.resolved_path = resolved_path;
} else {
// Relative path
let cwd = env::current_dir().map_err(|_| NavigateError::NotFound)?;
let abs_path = Self::normalize_path(&cwd.join(&chunk_path));
let resolved_path = Self::resolve_module(&abs_path)?;
self.abs_path = abs_path;
self.rel_path = chunk_path;
self.resolved_path = resolved_path;
}
Ok(())
}
fn jump_to_alias(&mut self, path: &str) -> StdResult<(), NavigateError> {
let path = Self::normalize_path(path.as_ref());
let resolved_path = Self::resolve_module(&path)?;
self.abs_path = path.clone();
self.rel_path = path;
self.resolved_path = resolved_path;
Ok(())
}
fn to_parent(&mut self) -> StdResult<(), NavigateError> {
let mut abs_path = self.abs_path.clone();
if !abs_path.pop() {
// It's important to return `NotFound` if we reached the root, as it's a "recoverable" error if we
// cannot go beyond the root directory.
// Luau "require-by-string` has a special logic to search for config file to resolve aliases.
return Err(NavigateError::NotFound);
}
let mut rel_parent = self.rel_path.clone();
rel_parent.pop();
let resolved_path = Self::resolve_module(&abs_path)?;
self.abs_path = abs_path;
self.rel_path = Self::normalize_path(&rel_parent);
self.resolved_path = resolved_path;
Ok(())
}
fn to_child(&mut self, name: &str) -> StdResult<(), NavigateError> {
let abs_path = self.abs_path.join(name);
let rel_path = self.rel_path.join(name);
let resolved_path = Self::resolve_module(&abs_path)?;
self.abs_path = abs_path;
self.rel_path = rel_path;
self.resolved_path = resolved_path;
Ok(())
}
fn has_module(&self) -> bool {
(self.resolved_path.as_deref())
.map(Path::is_file)
.unwrap_or(false)
}
fn cache_key(&self) -> String {
self.resolved_path.as_deref().unwrap().display().to_string()
}
fn has_config(&self) -> bool {
self.abs_path.is_dir() && self.abs_path.join(Self::LUAURC_CONFIG_FILENAME).is_file()
|| self.abs_path.is_dir() && self.abs_path.join(Self::LUAU_CONFIG_FILENAME).is_file()
}
fn config(&self) -> IoResult<Vec<u8>> {
if self.abs_path.join(Self::LUAURC_CONFIG_FILENAME).is_file() {
return fs::read(self.abs_path.join(Self::LUAURC_CONFIG_FILENAME));
}
fs::read(self.abs_path.join(Self::LUAU_CONFIG_FILENAME))
}
fn loader(&self, lua: &Lua) -> Result<Function> {
let name = format!("@{}", self.rel_path.display());
lua.load(self.resolved_path.as_deref().unwrap())
.set_name(name)
.into_function()
}
}
#[cfg(test)]
mod tests {
use std::path::Path;
use super::TextRequirer;
#[test]
fn test_path_normalize() {
for (input, expected) in [
// Basic formatting checks
("", "./"),
(".", "./"),
("a/relative/path", "./a/relative/path"),
// Paths containing extraneous '.' and '/' symbols
("./remove/extraneous/symbols/", "./remove/extraneous/symbols"),
("./remove/extraneous//symbols", "./remove/extraneous/symbols"),
("./remove/extraneous/symbols/.", "./remove/extraneous/symbols"),
("./remove/extraneous/./symbols", "./remove/extraneous/symbols"),
("../remove/extraneous/symbols/", "../remove/extraneous/symbols"),
("../remove/extraneous//symbols", "../remove/extraneous/symbols"),
("../remove/extraneous/symbols/.", "../remove/extraneous/symbols"),
("../remove/extraneous/./symbols", "../remove/extraneous/symbols"),
("/remove/extraneous/symbols/", "/remove/extraneous/symbols"),
("/remove/extraneous//symbols", "/remove/extraneous/symbols"),
("/remove/extraneous/symbols/.", "/remove/extraneous/symbols"),
("/remove/extraneous/./symbols", "/remove/extraneous/symbols"),
// Paths containing '..'
("./remove/me/..", "./remove"),
("./remove/me/../", "./remove"),
("../remove/me/..", "../remove"),
("../remove/me/../", "../remove"),
("/remove/me/..", "/remove"),
("/remove/me/../", "/remove"),
("./..", "../"),
("./../", "../"),
("../..", "../../"),
("../../", "../../"),
// '..' disappears if path is absolute and component is non-erasable
("/../", "/"),
] {
let path = TextRequirer::normalize_path(input.as_ref());
assert_eq!(
&path,
expected.as_ref() as &Path,
"wrong normalization for {input}"
);
}
}
}
+26
View File
@@ -204,6 +204,23 @@ impl IntoLuaMulti for MultiValue {
}
}
impl IntoLuaMulti for &MultiValue {
#[inline]
fn into_lua_multi(self, _: &Lua) -> Result<MultiValue> {
Ok(self.clone())
}
#[inline]
unsafe fn push_into_stack_multi(self, lua: &RawLua) -> Result<c_int> {
let nresults = self.len() as i32;
check_stack(lua.state(), nresults + 1)?;
for value in &self.0 {
lua.push_value(value)?;
}
Ok(nresults)
}
}
impl FromLuaMulti for MultiValue {
#[inline]
fn from_lua_multi(values: MultiValue, _: &Lua) -> Result<Self> {
@@ -297,6 +314,15 @@ impl<T: IntoLua> IntoLuaMulti for Variadic<T> {
fn into_lua_multi(self, lua: &Lua) -> Result<MultiValue> {
MultiValue::from_lua_iter(lua, self)
}
unsafe fn push_into_stack_multi(self, lua: &RawLua) -> Result<c_int> {
let nresults = self.len() as i32;
check_stack(lua.state(), nresults + 1)?;
for value in self.0 {
value.push_into_stack(lua)?;
}
Ok(nresults)
}
}
impl<T: FromLua> FromLuaMulti for Variadic<T> {
+2 -1
View File
@@ -36,5 +36,6 @@ pub use crate::{AsyncThread as LuaAsyncThread, LuaNativeAsyncFn};
#[cfg(feature = "serde")]
#[doc(no_inline)]
pub use crate::{
DeserializeOptions as LuaDeserializeOptions, LuaSerdeExt, SerializeOptions as LuaSerializeOptions,
DeserializeOptions as LuaDeserializeOptions, LuaSerdeExt, SerializableValue as LuaSerializableValue,
SerializeOptions as LuaSerializeOptions,
};
+41 -9
View File
@@ -15,11 +15,12 @@ use crate::userdata::AnyUserData;
use crate::value::Value;
/// A struct for deserializing Lua values into Rust values.
#[derive(Debug)]
#[derive(Debug, Default)]
pub struct Deserializer {
value: Value,
options: Options,
visited: Rc<RefCell<FxHashSet<*const c_void>>>,
len: Option<usize>, // A length hint for sequences
}
/// A struct with options to change default deserializer behavior.
@@ -54,6 +55,19 @@ pub struct Options {
///
/// Default: **false**
pub encode_empty_tables_as_array: bool,
/// If true, enable detection of mixed tables.
///
/// A mixed table is a table that has both array-like and map-like entries or several borders.
/// See [`The Length Operator`] documentation for details about borders.
///
/// When this option is disabled, a table with a non-zero length (with one or more borders) will
/// be always encoded as an array.
///
/// Default: **false**
///
/// [`The Length Operator`]: https://www.lua.org/manual/5.4/manual.html#3.4.7
pub detect_mixed_tables: bool,
}
impl Default for Options {
@@ -70,6 +84,7 @@ impl Options {
deny_recursive_tables: true,
sort_keys: false,
encode_empty_tables_as_array: false,
detect_mixed_tables: false,
}
}
@@ -108,6 +123,15 @@ impl Options {
self.encode_empty_tables_as_array = enabled;
self
}
/// Sets [`detect_mixed_tables`] option.
///
/// [`detect_mixed_tables`]: #structfield.detect_mixed_tables
#[must_use]
pub const fn detect_mixed_tables(mut self, enable: bool) -> Self {
self.detect_mixed_tables = enable;
self
}
}
impl Deserializer {
@@ -121,7 +145,7 @@ impl Deserializer {
Deserializer {
value,
options,
visited: Rc::new(RefCell::new(FxHashSet::default())),
..Default::default()
}
}
@@ -130,8 +154,14 @@ impl Deserializer {
value,
options,
visited,
..Default::default()
}
}
fn with_len(mut self, len: usize) -> Self {
self.len = Some(len);
self
}
}
impl<'de> serde::Deserializer<'de> for Deserializer {
@@ -155,11 +185,13 @@ impl<'de> serde::Deserializer<'de> for Deserializer {
Ok(s) => visitor.visit_str(&s),
Err(_) => visitor.visit_bytes(&s.as_bytes()),
},
Value::Table(ref t) if t.raw_len() > 0 || t.is_array() => self.deserialize_seq(visitor),
Value::Table(ref t) if self.options.encode_empty_tables_as_array && t.is_empty() => {
self.deserialize_seq(visitor)
Value::Table(ref t) => {
if let Some(len) = t.encode_as_array(self.options) {
self.with_len(len).deserialize_seq(visitor)
} else {
self.deserialize_map(visitor)
}
}
Value::Table(_) => self.deserialize_map(visitor),
Value::LightUserData(ud) if ud.0.is_null() => visitor.visit_none(),
Value::UserData(ud) if ud.is_serializable() => {
serde_userdata(ud, |value| value.deserialize_any(visitor))
@@ -270,14 +302,14 @@ impl<'de> serde::Deserializer<'de> for Deserializer {
Value::Table(t) => {
let _guard = RecursionGuard::new(&t, &self.visited);
let len = t.raw_len();
let len = self.len.unwrap_or_else(|| t.raw_len());
let mut deserializer = SeqDeserializer {
seq: t.sequence_values(),
seq: t.sequence_values().with_len(len),
options: self.options,
visited: self.visited,
};
let seq = visitor.visit_seq(&mut deserializer)?;
if deserializer.seq.count() == 0 {
if deserializer.seq.next().is_none() {
Ok(seq)
} else {
Err(de::Error::invalid_length(len, &"fewer elements in the table"))
+193 -58
View File
@@ -37,6 +37,7 @@ use crate::{buffer::Buffer, chunk::Compiler};
use {
crate::types::LightUserData,
std::future::{self, Future},
std::task::Poll,
};
#[cfg(feature = "serde")]
@@ -358,6 +359,8 @@ impl Lua {
if cfg!(feature = "luau") && !modname.starts_with('@') {
return Err(Error::runtime("module name must begin with '@'"));
}
#[cfg(feature = "luau")]
let modname = modname.to_ascii_lowercase();
unsafe {
self.exec_raw::<()>(value, |state| {
ffi::luaL_getsubtable(state, ffi::LUA_REGISTRYINDEX, LOADED_MODULES_KEY);
@@ -629,13 +632,13 @@ impl Lua {
///
/// Any Luau code is guaranteed to call this handler "eventually"
/// (in practice this can happen at any function call or at any loop iteration).
/// This is similar to `Lua::set_hook` but in more simplified form.
///
/// The provided interrupt function can error, and this error will be propagated through
/// the Luau code that was executing at the time the interrupt was triggered.
/// Also this can be used to implement continuous execution limits by instructing Luau VM to
/// yield by returning [`VmState::Yield`].
///
/// This is similar to `Lua::set_hook` but in more simplified form.
/// yield by returning [`VmState::Yield`]. The yield will happen only at yieldable points
/// of execution (not across metamethod/C-call boundaries).
///
/// # Example
///
@@ -693,7 +696,10 @@ impl Lua {
match result {
VmState::Continue => {}
VmState::Yield => {
ffi::lua_yield(state, 0);
// We can yield only at yieldable points, otherwise ignore and continue
if ffi::lua_isyieldable(state) != 0 {
ffi::lua_yield(state, 0);
}
}
}
}
@@ -869,7 +875,7 @@ impl Lua {
}
}
/// Gets information about the interpreter runtime stack at a given level.
/// Gets information about the interpreter runtime stack at the given level.
///
/// This function calls callback `f`, passing the [`Debug`] structure that can be used to get
/// information about the function executing at a given level.
@@ -893,6 +899,26 @@ impl Lua {
}
}
/// Creates a traceback of the call stack at the given level.
///
/// The `msg` parameter, if provided, is added at the beginning of the traceback.
/// The `level` parameter works the same way as in [`Lua::inspect_stack`].
pub fn traceback(&self, msg: Option<&str>, level: usize) -> Result<String> {
let lua = self.lock();
unsafe {
check_stack(lua.state(), 3)?;
protect_lua!(lua.state(), 0, 1, |state| {
let msg = match msg {
Some(s) => ffi::lua_pushlstring(state, s.as_ptr() as *const c_char, s.len()),
None => ptr::null(),
};
// `protect_lua` adds it's own call frame, so we need to increase level by 1
ffi::luaL_traceback(state, state, msg, (level + 1) as c_int);
})?;
Ok(String(lua.pop_ref()))
}
}
/// Returns the amount of memory (in bytes) currently used inside this Lua state.
pub fn used_memory(&self) -> usize {
let lua = self.lock();
@@ -1148,36 +1174,41 @@ impl Lua {
}
}
/// Create and return an interned Lua string.
/// Creates and returns an interned Lua string.
///
/// Lua strings can be arbitrary `[u8]` data including embedded nulls, so in addition to `&str`
/// and `&String`, you can also pass plain `&[u8]` here.
#[inline]
pub fn create_string(&self, s: impl AsRef<[u8]>) -> Result<String> {
unsafe { self.lock().create_string(s) }
unsafe { self.lock().create_string(s.as_ref()) }
}
/// Create and return a Luau [buffer] object from a byte slice of data.
/// Creates and returns a Luau [buffer] object from a byte slice of data.
///
/// [buffer]: https://luau.org/library#buffer-library
#[cfg(any(feature = "luau", doc))]
#[cfg_attr(docsrs, doc(cfg(feature = "luau")))]
pub fn create_buffer(&self, buf: impl AsRef<[u8]>) -> Result<Buffer> {
pub fn create_buffer(&self, data: impl AsRef<[u8]>) -> Result<Buffer> {
let lua = self.lock();
let state = lua.state();
let data = data.as_ref();
unsafe {
if lua.unlikely_memory_error() {
crate::util::push_buffer(state, buf.as_ref(), false)?;
return Ok(Buffer(lua.pop_ref()));
}
let _sg = StackGuard::new(state);
check_stack(state, 3)?;
crate::util::push_buffer(state, buf.as_ref(), true)?;
Ok(Buffer(lua.pop_ref()))
let (ptr, buffer) = lua.create_buffer_with_capacity(data.len())?;
ptr.copy_from_nonoverlapping(data.as_ptr(), data.len());
Ok(buffer)
}
}
/// Creates and returns a Luau [buffer] object with the specified size.
///
/// Size limit is 1GB. All bytes will be initialized to zero.
///
/// [buffer]: https://luau.org/library#buffer-library
#[cfg(any(feature = "luau", doc))]
#[cfg_attr(docsrs, doc(cfg(feature = "luau")))]
pub fn create_buffer_with_capacity(&self, size: usize) -> Result<Buffer> {
unsafe { Ok(self.lock().create_buffer_with_capacity(size)?.1) }
}
/// Creates and returns a new empty table.
#[inline]
pub fn create_table(&self) -> Result<Table> {
@@ -1286,8 +1317,24 @@ impl Lua {
/// This function is unsafe because provides a way to execute unsafe C function.
pub unsafe fn create_c_function(&self, func: ffi::lua_CFunction) -> Result<Function> {
let lua = self.lock();
ffi::lua_pushcfunction(lua.ref_thread(), func);
Ok(Function(lua.pop_ref_thread()))
if cfg!(any(feature = "lua54", feature = "lua53", feature = "lua52")) {
ffi::lua_pushcfunction(lua.ref_thread(), func);
return Ok(Function(lua.pop_ref_thread()));
}
// Lua <5.2 requires memory allocation to push a C function
let state = lua.state();
{
let _sg = StackGuard::new(state);
check_stack(state, 3)?;
if lua.unlikely_memory_error() {
ffi::lua_pushcfunction(state, func);
} else {
protect_lua!(state, 0, 1, |state| ffi::lua_pushcfunction(state, func))?;
}
Ok(Function(lua.pop_ref()))
}
}
/// Wraps a Rust async function or closure, creating a callable Lua function handle to it.
@@ -1467,7 +1514,27 @@ impl Lua {
unsafe { self.lock().make_userdata(UserDataStorage::new(ud)) }
}
/// Sets the metatable for a Lua builtin type.
/// Gets the metatable of a Lua built-in (primitive) type.
///
/// The metatable is shared by all values of the given type.
///
/// See [`Lua::set_type_metatable`] for examples.
#[allow(private_bounds)]
pub fn type_metatable<T: LuaType>(&self) -> Option<Table> {
let lua = self.lock();
let state = lua.state();
unsafe {
let _sg = StackGuard::new(state);
assert_stack(state, 2);
if lua.push_primitive_type::<T>() && ffi::lua_getmetatable(state, -1) != 0 {
return Some(Table(lua.pop_ref()));
}
}
None
}
/// Sets the metatable for a Lua built-in (primitive) type.
///
/// The metatable will be shared by all values of the given type.
///
@@ -1494,44 +1561,13 @@ impl Lua {
let _sg = StackGuard::new(state);
assert_stack(state, 2);
match T::TYPE_ID {
ffi::LUA_TBOOLEAN => {
ffi::lua_pushboolean(state, 0);
if lua.push_primitive_type::<T>() {
match metatable {
Some(metatable) => lua.push_ref(&metatable.0),
None => ffi::lua_pushnil(state),
}
ffi::LUA_TLIGHTUSERDATA => {
ffi::lua_pushlightuserdata(state, ptr::null_mut());
}
ffi::LUA_TNUMBER => {
ffi::lua_pushnumber(state, 0.);
}
#[cfg(feature = "luau")]
ffi::LUA_TVECTOR => {
#[cfg(not(feature = "luau-vector4"))]
ffi::lua_pushvector(state, 0., 0., 0.);
#[cfg(feature = "luau-vector4")]
ffi::lua_pushvector(state, 0., 0., 0., 0.);
}
ffi::LUA_TSTRING => {
ffi::lua_pushstring(state, b"\0" as *const u8 as *const _);
}
ffi::LUA_TFUNCTION => match self.load("function() end").eval::<Function>() {
Ok(func) => lua.push_ref(&func.0),
Err(_) => return,
},
ffi::LUA_TTHREAD => {
ffi::lua_pushthread(state);
}
#[cfg(feature = "luau")]
ffi::LUA_TBUFFER => {
ffi::lua_newbuffer(state, 0);
}
_ => return,
ffi::lua_setmetatable(state, -2);
}
match metatable {
Some(metatable) => lua.push_ref(&metatable.0),
None => ffi::lua_pushnil(state),
}
ffi::lua_setmetatable(state, -2);
}
}
@@ -2053,6 +2089,105 @@ impl Lua {
LightUserData(&ASYNC_POLL_TERMINATE as *const u8 as *mut std::os::raw::c_void)
}
#[cfg(feature = "async")]
#[inline(always)]
pub(crate) fn poll_yield() -> LightUserData {
static ASYNC_POLL_YIELD: u8 = 0;
LightUserData(&ASYNC_POLL_YIELD as *const u8 as *mut std::os::raw::c_void)
}
/// Suspends the current async function, returning the provided arguments to caller.
///
/// This function is similar to [`coroutine.yield`] but allow yeilding Rust functions
/// and passing values to the caller.
/// Please note that you cannot cross [`Thread`] boundaries (e.g. calling `yield_with` on one
/// thread and resuming on another).
///
/// # Examples
///
/// Async iterator:
///
/// ```
/// # use mlua::{Lua, Result};
/// #
/// async fn generator(lua: Lua, _: ()) -> Result<()> {
/// for i in 0..10 {
/// lua.yield_with::<()>(i).await?;
/// }
/// Ok(())
/// }
///
/// fn main() -> Result<()> {
/// let lua = Lua::new();
/// lua.globals().set("generator", lua.create_async_function(generator)?)?;
///
/// lua.load(r#"
/// local n = 0
/// for i in coroutine.wrap(generator) do
/// n = n + i
/// end
/// assert(n == 45)
/// "#)
/// .exec()
/// }
/// ```
///
/// Exchange values on yield:
///
/// ```
/// # use mlua::{Lua, Result, Value};
/// #
/// async fn pingpong(lua: Lua, mut val: i32) -> Result<()> {
/// loop {
/// val = lua.yield_with::<i32>(val).await? + 1;
/// }
/// Ok(())
/// }
///
/// # fn main() -> Result<()> {
/// let lua = Lua::new();
///
/// let co = lua.create_thread(lua.create_async_function(pingpong)?)?;
/// assert_eq!(co.resume::<i32>(1)?, 1);
/// assert_eq!(co.resume::<i32>(2)?, 3);
/// assert_eq!(co.resume::<i32>(3)?, 4);
///
/// # Ok(())
/// # }
/// ```
///
/// [`coroutine.yield`]: https://www.lua.org/manual/5.4/manual.html#pdf-coroutine.yield
#[cfg(feature = "async")]
#[cfg_attr(docsrs, doc(cfg(feature = "async")))]
pub async fn yield_with<R: FromLuaMulti>(&self, args: impl IntoLuaMulti) -> Result<R> {
let mut args = Some(args.into_lua_multi(self)?);
future::poll_fn(move |_cx| match args.take() {
Some(args) => unsafe {
let lua = self.lock();
lua.push(Self::poll_yield())?; // yield marker
if args.len() <= 1 {
lua.push(args.front())?;
} else {
lua.push(lua.create_sequence_from(&args)?)?;
}
lua.push(args.len())?;
Poll::Pending
},
None => unsafe {
let lua = self.lock();
let state = lua.state();
let top = ffi::lua_gettop(state);
if top == 0 || ffi::lua_type(state, 1) != ffi::LUA_TUSERDATA {
// This must be impossible scenario if used correctly
return Poll::Ready(R::from_stack_multi(0, &lua));
}
let _sg = StackGuard::with_top(state, 1);
Poll::Ready(R::from_stack_multi(top - 1, &lua))
},
})
.await
}
/// Returns a weak reference to the Lua instance.
///
/// This is useful for creating a reference to the Lua instance that does not prevent it from
+6 -2
View File
@@ -64,7 +64,7 @@ pub(crate) struct ExtraData {
pub(super) wrapped_failure_top: usize,
// Pool of `Thread`s (coroutines) for async execution
#[cfg(feature = "async")]
pub(super) thread_pool: Vec<c_int>,
pub(super) thread_pool: Vec<crate::types::ValueRefIndex>,
// Address of `WrappedFailure` metatable
pub(super) wrapped_failure_mt_ptr: *const c_void,
@@ -94,6 +94,8 @@ pub(crate) struct ExtraData {
pub(super) compiler: Option<Compiler>,
#[cfg(feature = "luau-jit")]
pub(super) enable_jit: bool,
#[cfg(feature = "luau")]
pub(crate) mem_categories: Vec<std::ffi::CString>,
}
impl Drop for ExtraData {
@@ -196,6 +198,8 @@ impl ExtraData {
enable_jit: true,
#[cfg(feature = "luau")]
running_gc: false,
#[cfg(feature = "luau")]
mem_categories: vec![std::ffi::CString::new("main").unwrap()],
}));
// Store it in the registry
@@ -270,7 +274,7 @@ impl ExtraData {
// Try to grow max stack size
if self.ref_stack_top >= self.ref_stack_size {
let mut inc = self.ref_stack_size; // Try to double stack size
while inc > 0 && ffi::lua_checkstack(self.ref_thread, inc) == 0 {
while inc > 0 && ffi::lua_checkstack(self.ref_thread, inc + REF_STACK_RESERVE) == 0 {
inc /= 2;
}
if inc == 0 {
+121 -46
View File
@@ -19,7 +19,7 @@ use crate::thread::Thread;
use crate::traits::IntoLua;
use crate::types::{
AppDataRef, AppDataRefMut, Callback, CallbackUpvalue, DestructedUserdata, Integer, LightUserData,
MaybeSend, ReentrantMutex, RegistryKey, ValueRef, XRc,
LuaType, MaybeSend, ReentrantMutex, RegistryKey, ValueRef, XRc,
};
use crate::userdata::{
init_userdata_metatable, AnyUserData, MetaMethod, RawUserDataRegistry, UserData, UserDataRegistry,
@@ -28,8 +28,8 @@ use crate::userdata::{
use crate::util::{
assert_stack, check_stack, get_destructed_userdata_metatable, get_internal_userdata, get_main_state,
get_metatable_ptr, get_userdata, init_error_registry, init_internal_metatable, pop_error,
push_internal_userdata, push_string, push_table, rawset_field, safe_pcall, safe_xpcall, short_type_name,
StackGuard, WrappedFailure,
push_internal_userdata, push_string, push_table, push_userdata, rawset_field, safe_pcall, safe_xpcall,
short_type_name, StackGuard, WrappedFailure,
};
use crate::value::{Nil, Value};
@@ -510,19 +510,33 @@ impl RawLua {
}
/// See [`Lua::create_string`]
pub(crate) unsafe fn create_string(&self, s: impl AsRef<[u8]>) -> Result<String> {
pub(crate) unsafe fn create_string(&self, s: &[u8]) -> Result<String> {
let state = self.state();
if self.unlikely_memory_error() {
push_string(state, s.as_ref(), false)?;
push_string(state, s, false)?;
return Ok(String(self.pop_ref()));
}
let _sg = StackGuard::new(state);
check_stack(state, 3)?;
push_string(state, s.as_ref(), true)?;
push_string(state, s, true)?;
Ok(String(self.pop_ref()))
}
#[cfg(feature = "luau")]
pub(crate) unsafe fn create_buffer_with_capacity(&self, size: usize) -> Result<(*mut u8, crate::Buffer)> {
let state = self.state();
if self.unlikely_memory_error() {
let ptr = crate::util::push_buffer(state, size, false)?;
return Ok((ptr, crate::Buffer(self.pop_ref())));
}
let _sg = StackGuard::new(state);
check_stack(state, 3)?;
let ptr = crate::util::push_buffer(state, size, true)?;
Ok((ptr, crate::Buffer(self.pop_ref())))
}
/// See [`Lua::create_table_with_capacity`]
pub(crate) unsafe fn create_table_with_capacity(&self, narr: usize, nrec: usize) -> Result<Table> {
let state = self.state();
@@ -624,7 +638,7 @@ impl RawLua {
#[cfg(feature = "async")]
pub(crate) unsafe fn create_recycled_thread(&self, func: &Function) -> Result<Thread> {
if let Some(index) = (*self.extra.get()).thread_pool.pop() {
let thread_state = ffi::lua_tothread(self.ref_thread(), index);
let thread_state = ffi::lua_tothread(self.ref_thread(), *index.0);
ffi::lua_xpush(self.ref_thread(), thread_state, func.0.index);
#[cfg(feature = "luau")]
@@ -645,11 +659,52 @@ impl RawLua {
pub(crate) unsafe fn recycle_thread(&self, thread: &mut Thread) {
let extra = &mut *self.extra.get();
if extra.thread_pool.len() < extra.thread_pool.capacity() {
extra.thread_pool.push(thread.0.index);
thread.0.drop = false; // Prevent thread from being garbage collected
if let Some(index) = thread.0.index_count.take() {
extra.thread_pool.push(index);
}
}
}
/// Pushes a primitive type value onto the Lua stack.
pub(crate) unsafe fn push_primitive_type<T: LuaType>(&self) -> bool {
match T::TYPE_ID {
ffi::LUA_TBOOLEAN => {
ffi::lua_pushboolean(self.state(), 0);
}
ffi::LUA_TLIGHTUSERDATA => {
ffi::lua_pushlightuserdata(self.state(), ptr::null_mut());
}
ffi::LUA_TNUMBER => {
ffi::lua_pushnumber(self.state(), 0.);
}
#[cfg(feature = "luau")]
ffi::LUA_TVECTOR => {
#[cfg(not(feature = "luau-vector4"))]
ffi::lua_pushvector(self.state(), 0., 0., 0.);
#[cfg(feature = "luau-vector4")]
ffi::lua_pushvector(self.state(), 0., 0., 0., 0.);
}
ffi::LUA_TSTRING => {
ffi::lua_pushstring(self.state(), b"\0" as *const u8 as *const _);
}
ffi::LUA_TFUNCTION => {
unsafe extern "C-unwind" fn func(_state: *mut ffi::lua_State) -> c_int {
0
}
ffi::lua_pushcfunction(self.state(), func);
}
ffi::LUA_TTHREAD => {
ffi::lua_pushthread(self.state());
}
#[cfg(feature = "luau")]
ffi::LUA_TBUFFER => {
ffi::lua_newbuffer(self.state(), 0);
}
_ => return false,
}
true
}
/// Pushes a value that implements `IntoLua` onto the Lua stack.
///
/// Uses up to 2 stack spaces to push a single value, does not call `checkstack`.
@@ -728,7 +783,7 @@ impl RawLua {
let n = ffi::lua_tonumber(state, idx);
match num_traits::cast(n) {
Some(i) if (n - (i as Number)).abs() < Number::EPSILON => Value::Integer(i),
Some(i) if n.to_bits() == (i as Number).to_bits() => Value::Integer(i),
_ => Value::Number(n),
}
}
@@ -827,13 +882,6 @@ impl RawLua {
ValueRef::new(self, index)
}
#[inline]
pub(crate) unsafe fn clone_ref(&self, vref: &ValueRef) -> ValueRef {
ffi::lua_pushvalue(self.ref_thread(), vref.index);
let index = (*self.extra.get()).ref_stack_pop();
ValueRef::new(self, index)
}
pub(crate) unsafe fn drop_ref(&self, vref: &ValueRef) {
let ref_thread = self.ref_thread();
mlua_debug_assert!(
@@ -920,7 +968,7 @@ impl RawLua {
// We generate metatable first to make sure it *always* available when userdata pushed
let mt_id = get_metatable_id()?;
let protect = !self.unlikely_memory_error();
crate::util::push_userdata(state, data, protect)?;
push_userdata(state, data, protect)?;
ffi::lua_rawgeti(state, ffi::LUA_REGISTRYINDEX, mt_id);
ffi::lua_setmetatable(state, -2);
@@ -1048,6 +1096,18 @@ impl RawLua {
field_setters_index = Some(ffi::lua_absindex(state, -1));
}
// Create methods namecall table
#[cfg_attr(not(feature = "luau"), allow(unused_mut))]
let mut methods_map = None;
#[cfg(feature = "luau")]
if registry.enable_namecall {
let map: &mut rustc_hash::FxHashMap<_, crate::types::CallbackPtr> =
methods_map.get_or_insert_with(Default::default);
for (k, m) in &registry.methods {
map.insert(k.as_bytes().to_vec(), &**m);
}
}
let mut methods_index = None;
let methods_nrec = registry.methods.len();
#[cfg(feature = "async")]
@@ -1095,6 +1155,7 @@ impl RawLua {
field_getters_index,
field_setters_index,
methods_index,
methods_map,
)?;
// Update stack guard to keep metatable after return
@@ -1226,7 +1287,7 @@ impl RawLua {
}
}
unsafe extern "C-unwind" fn call_callback(state: *mut ffi::lua_State) -> c_int {
unsafe extern "C-unwind" fn get_future_callback(state: *mut ffi::lua_State) -> c_int {
// Async functions cannot be scoped and therefore destroyed,
// so the first upvalue is always valid
let upvalue = get_userdata::<AsyncCallbackUpvalue>(state, ffi::lua_upvalueindex(1));
@@ -1240,36 +1301,37 @@ impl RawLua {
let extra = XRc::clone(&(*upvalue).extra);
let protect = !rawlua.unlikely_memory_error();
push_internal_userdata(state, AsyncPollUpvalue { data: fut, extra }, protect)?;
if protect {
protect_lua!(state, 1, 1, fn(state) {
ffi::lua_pushcclosure(state, poll_future, 1);
})?;
} else {
ffi::lua_pushcclosure(state, poll_future, 1);
}
Ok(1)
})
}
unsafe extern "C-unwind" fn poll_future(state: *mut ffi::lua_State) -> c_int {
let upvalue = get_userdata::<AsyncPollUpvalue>(state, ffi::lua_upvalueindex(1));
callback_error_ext(state, (*upvalue).extra.get(), true, |extra, nargs| {
// Future is always passed in the first argument
let future = get_userdata::<AsyncPollUpvalue>(state, 1);
callback_error_ext(state, (*future).extra.get(), true, |extra, nargs| {
// Lua ensures that `LUA_MINSTACK` stack spaces are available (after pushing arguments)
// The lock must be already held as the future is polled
let rawlua = (*extra).raw_lua();
if nargs == 1 && ffi::lua_tolightuserdata(state, -1) == Lua::poll_terminate().0 {
if nargs == 2 && ffi::lua_tolightuserdata(state, -1) == Lua::poll_terminate().0 {
// Destroy the future and terminate the Lua thread
(*upvalue).data.take();
(*future).data.take();
ffi::lua_pushinteger(state, -1);
return Ok(1);
}
let fut = &mut (*upvalue).data;
let fut = &mut (*future).data;
let mut ctx = Context::from_waker(rawlua.waker());
match fut.as_mut().map(|fut| fut.as_mut().poll(&mut ctx)) {
Some(Poll::Pending) => {
let fut_nvals = ffi::lua_gettop(state) - 1; // Exclude the future itself
if fut_nvals >= 3 && ffi::lua_tolightuserdata(state, -3) == Lua::poll_yield().0 {
// We have some values to yield
ffi::lua_pushnil(state);
ffi::lua_replace(state, -4);
return Ok(3);
}
ffi::lua_pushnil(state);
ffi::lua_pushlightuserdata(state, Lua::poll_pending().0);
Ok(2)
@@ -1298,7 +1360,7 @@ impl RawLua {
}
let state = self.state();
let get_poll = unsafe {
let get_future = unsafe {
let _sg = StackGuard::new(state);
check_stack(state, 4)?;
@@ -1308,10 +1370,10 @@ impl RawLua {
push_internal_userdata(state, upvalue, protect)?;
if protect {
protect_lua!(state, 1, 1, fn(state) {
ffi::lua_pushcclosure(state, call_callback, 1);
ffi::lua_pushcclosure(state, get_future_callback, 1);
})?;
} else {
ffi::lua_pushcclosure(state, call_callback, 1);
ffi::lua_pushcclosure(state, get_future_callback, 1);
}
Function(self.pop_ref())
@@ -1330,16 +1392,19 @@ impl RawLua {
let coroutine = lua.globals().get::<Table>("coroutine")?;
// Prepare environment for the async poller
let env = lua.create_table_with_capacity(0, 3)?;
env.set("get_poll", get_poll)?;
let env = lua.create_table_with_capacity(0, 4)?;
env.set("get_future", get_future)?;
env.set("poll", unsafe { lua.create_c_function(poll_future)? })?;
env.set("yield", coroutine.get::<Function>("yield")?)?;
env.set("unpack", unsafe { lua.create_c_function(unpack)? })?;
lua.load(
r#"
local poll = get_poll(...)
local nres, res, res2 = poll()
local poll, yield = poll, yield
local future = get_future(...)
local nres, res, res2 = poll(future)
while true do
-- Poll::Ready branch, `nres` is the number of results
if nres ~= nil then
if nres == 0 then
return
@@ -1355,10 +1420,20 @@ impl RawLua {
return unpack(res, nres)
end
end
-- `res` is a "pending" value
-- `yield` can return a signal to drop the future that we should propagate
-- to the poller
nres, res, res2 = poll(yield(res))
-- Poll::Pending branch
if res2 == nil then
-- `res` is a "pending" value
-- `yield` can return a signal to drop the future that we should propagate
-- to the poller
nres, res, res2 = poll(future, yield(res))
elseif res2 == 0 then
nres, res, res2 = poll(future, yield())
elseif res2 == 1 then
nres, res, res2 = poll(future, yield(res))
else
nres, res, res2 = poll(future, yield(unpack(res, res2)))
end
end
"#,
)
@@ -1370,14 +1445,14 @@ impl RawLua {
#[cfg(feature = "async")]
#[inline]
pub(crate) unsafe fn waker(&self) -> &Waker {
(*self.extra.get()).waker.as_ref()
pub(crate) fn waker(&self) -> &Waker {
unsafe { (*self.extra.get()).waker.as_ref() }
}
#[cfg(feature = "async")]
#[inline]
pub(crate) unsafe fn set_waker(&self, waker: NonNull<Waker>) -> NonNull<Waker> {
mem::replace(&mut (*self.extra.get()).waker, waker)
pub(crate) fn set_waker(&self, waker: NonNull<Waker>) -> NonNull<Waker> {
unsafe { mem::replace(&mut (*self.extra.get()).waker, waker) }
}
}
+119 -38
View File
@@ -1,14 +1,14 @@
use std::collections::HashSet;
use std::fmt;
use std::marker::PhantomData;
use std::os::raw::{c_int, c_void};
use std::os::raw::c_void;
use std::string::String as StdString;
use crate::error::{Error, Result};
use crate::function::Function;
use crate::state::{LuaGuard, RawLua};
use crate::state::{LuaGuard, RawLua, WeakLua};
use crate::traits::{FromLua, FromLuaMulti, IntoLua, IntoLuaMulti, ObjectLike};
use crate::types::{Integer, LuaType, ValueRef};
use crate::types::{Integer, ValueRef};
use crate::util::{assert_stack, check_stack, get_metatable_ptr, StackGuard};
use crate::value::{Nil, Value};
@@ -416,14 +416,7 @@ impl Table {
lua.push_ref(&self.0);
// Clear array part
for i in 1..=ffi::lua_rawlen(state, -1) {
ffi::lua_pushnil(state);
ffi::lua_rawseti(state, -2, i as Integer);
}
// Clear hash part
// It must be safe as long as we don't use invalid keys
// This is safe as long as we don't assign new keys
ffi::lua_pushnil(state);
while ffi::lua_next(state, -2) != 0 {
ffi::lua_pop(state, 1); // pop value
@@ -510,7 +503,7 @@ impl Table {
let lua = self.0.lua.lock();
let ref_thread = lua.ref_thread();
unsafe {
if let Some(metatable) = metatable {
if let Some(metatable) = &metatable {
ffi::lua_pushvalue(ref_thread, metatable.0.index);
} else {
ffi::lua_pushnil(ref_thread);
@@ -675,16 +668,25 @@ impl Table {
guard: self.0.lua.lock(),
table: self,
index: 1,
len: None,
_phantom: PhantomData,
}
}
/// Iterates over the sequence part of the table, invoking the given closure on each value.
///
/// This methods is similar to [`Table::sequence_values`], but optimized for performance.
#[doc(hidden)]
pub fn for_each_value<V>(&self, mut f: impl FnMut(V) -> Result<()>) -> Result<()>
where
V: FromLua,
{
pub fn for_each_value<V: FromLua>(&self, f: impl FnMut(V) -> Result<()>) -> Result<()> {
self.for_each_value_by_len(None, f)
}
fn for_each_value_by_len<V: FromLua>(
&self,
len: impl Into<Option<usize>>,
mut f: impl FnMut(V) -> Result<()>,
) -> Result<()> {
let len = len.into();
let lua = self.0.lua.lock();
let state = lua.state();
unsafe {
@@ -692,9 +694,14 @@ impl Table {
check_stack(state, 4)?;
lua.push_ref(&self.0);
let len = ffi::lua_rawlen(state, -1);
for i in 1..=len {
ffi::lua_rawgeti(state, -1, i as _);
for i in 1.. {
if len.map(|len| i > len).unwrap_or(false) {
break;
}
let t = ffi::lua_rawgeti(state, -1, i as _);
if len.is_none() && t == ffi::LUA_TNIL {
break;
}
f(V::from_stack(-1, &lua)?)?;
ffi::lua_pop(state, 1);
}
@@ -727,8 +734,9 @@ impl Table {
Ok(())
}
/// Checks if the table has the array metatable attached.
#[cfg(feature = "serde")]
pub(crate) fn is_array(&self) -> bool {
fn has_array_metatable(&self) -> bool {
let lua = self.0.lua.lock();
let state = lua.state();
unsafe {
@@ -744,6 +752,70 @@ impl Table {
}
}
/// If the table is an array, returns the number of non-nil elements and max index.
///
/// Returns `None` if the table is not an array.
///
/// This operation has O(n) complexity.
#[cfg(feature = "serde")]
fn find_array_len(&self) -> Option<(usize, usize)> {
let lua = self.0.lua.lock();
let ref_thread = lua.ref_thread();
unsafe {
let _sg = StackGuard::new(ref_thread);
let (mut count, mut max_index) = (0, 0);
ffi::lua_pushnil(ref_thread);
while ffi::lua_next(ref_thread, self.0.index) != 0 {
if ffi::lua_type(ref_thread, -2) != ffi::LUA_TNUMBER {
return None;
}
let k = ffi::lua_tonumber(ref_thread, -2);
if k.trunc() != k || k < 1.0 {
return None;
}
max_index = std::cmp::max(max_index, k as usize);
count += 1;
ffi::lua_pop(ref_thread, 1);
}
Some((count, max_index))
}
}
/// Determines if the table should be encoded as an array or a map.
///
/// The algorithm is the following:
/// 1. If `detect_mixed_tables` is enabled, iterate over all keys in the table checking is they
/// all are positive integers. If non-array key is found, return `None` (encode as map).
/// Otherwise check the sparsity of the array. Too sparse arrays are encoded as maps.
///
/// 2. If `detect_mixed_tables` is disabled, check if the table has a positive length or has the
/// array metatable. If so, encode as array. If the table is empty and
/// `encode_empty_tables_as_array` is enabled, encode as array.
///
/// Returns the length of the array if it should be encoded as an array.
#[cfg(feature = "serde")]
pub(crate) fn encode_as_array(&self, options: crate::serde::de::Options) -> Option<usize> {
if options.detect_mixed_tables {
if let Some((len, max_idx)) = self.find_array_len() {
// If the array is too sparse, serialize it as a map instead
if len < 10 || len * 2 >= max_idx {
return Some(max_idx);
}
}
} else {
let len = self.raw_len();
if len > 0 || self.has_array_metatable() {
return Some(len);
}
if options.encode_empty_tables_as_array && self.is_empty() {
return Some(0);
}
}
None
}
#[cfg(feature = "luau")]
#[inline(always)]
fn check_readonly_write(&self, lua: &RawLua) -> Result<()> {
@@ -863,10 +935,6 @@ where
}
}
impl LuaType for Table {
const TYPE_ID: c_int = ffi::LUA_TTABLE;
}
impl ObjectLike for Table {
#[inline]
fn get<V: FromLua>(&self, key: impl IntoLua) -> Result<V> {
@@ -884,7 +952,7 @@ impl ObjectLike for Table {
R: FromLuaMulti,
{
// Convert table to a function and call via pcall that respects the `__call` metamethod.
Function(self.0.copy()).call(args)
Function(self.0.clone()).call(args)
}
#[cfg(feature = "async")]
@@ -893,7 +961,7 @@ impl ObjectLike for Table {
where
R: FromLuaMulti,
{
Function(self.0.copy()).call_async(args)
Function(self.0.clone()).call_async(args)
}
#[inline]
@@ -941,7 +1009,17 @@ impl ObjectLike for Table {
#[inline]
fn to_string(&self) -> Result<StdString> {
Value::Table(Table(self.0.copy())).to_string()
Value::Table(Table(self.0.clone())).to_string()
}
#[inline]
fn to_value(&self) -> Value {
Value::Table(self.clone())
}
#[inline]
fn weak_lua(&self) -> &WeakLua {
&self.0.lua
}
}
@@ -977,6 +1055,15 @@ impl<'a> SerializableTable<'a> {
}
}
impl<V> TableSequence<'_, V> {
/// Sets the length (hint) of the sequence.
#[cfg(feature = "serde")]
pub(crate) fn with_len(mut self, len: usize) -> Self {
self.len = Some(len);
self
}
}
#[cfg(feature = "serde")]
impl Serialize for SerializableTable<'_> {
fn serialize<S>(&self, serializer: S) -> StdResult<S::Ok, S::Error>
@@ -998,14 +1085,10 @@ impl Serialize for SerializableTable<'_> {
let _guard = RecursionGuard::new(self.table, visited);
// Array
let len = self.table.raw_len();
if len > 0
|| self.table.is_array()
|| (self.options.encode_empty_tables_as_array && self.table.is_empty())
{
if let Some(len) = self.table.encode_as_array(self.options) {
let mut seq = serializer.serialize_seq(Some(len))?;
let mut serialize_err = None;
let res = self.table.for_each_value::<Value>(|value| {
let res = self.table.for_each_value_by_len::<Value>(len, |value| {
let skip = check_value_for_skip(&value, self.options, visited)
.map_err(|err| Error::SerializeError(err.to_string()))?;
if skip {
@@ -1129,13 +1212,11 @@ pub struct TableSequence<'a, V> {
guard: LuaGuard,
table: &'a Table,
index: Integer,
len: Option<usize>,
_phantom: PhantomData<V>,
}
impl<V> Iterator for TableSequence<'_, V>
where
V: FromLua,
{
impl<V: FromLua> Iterator for TableSequence<'_, V> {
type Item = Result<V>;
fn next(&mut self) -> Option<Self::Item> {
@@ -1149,7 +1230,7 @@ where
lua.push_ref(&self.table.0);
match ffi::lua_rawgeti(state, -1, self.index) {
ffi::LUA_TNIL => None,
ffi::LUA_TNIL if self.index as usize > self.len.unwrap_or(0) => None,
_ => {
self.index += 1;
Some(V::from_stack(-1, lua))
+5 -5
View File
@@ -156,7 +156,6 @@ impl Thread {
let thread_state = self.state();
unsafe {
let _sg = StackGuard::new(state);
let _thread_sg = StackGuard::with_top(thread_state, 0);
let nargs = args.push_into_stack_multi(&lua)?;
if nargs > 0 {
@@ -165,6 +164,7 @@ impl Thread {
pushed_nargs += nargs;
}
let _thread_sg = StackGuard::with_top(thread_state, 0);
let (_, nresults) = self.resume_inner(&lua, pushed_nargs)?;
check_stack(state, nresults + 1)?;
ffi::lua_xmove(thread_state, state, nresults);
@@ -192,12 +192,12 @@ impl Thread {
let thread_state = self.state();
unsafe {
let _sg = StackGuard::new(state);
let _thread_sg = StackGuard::with_top(thread_state, 0);
check_stack(state, 1)?;
error.push_into_stack(&lua)?;
ffi::lua_xmove(state, thread_state, 1);
let _thread_sg = StackGuard::with_top(thread_state, 0);
let (_, nresults) = self.resume_inner(&lua, ffi::LUA_RESUMEERROR)?;
check_stack(state, nresults + 1)?;
ffi::lua_xmove(thread_state, state, nresults);
@@ -604,7 +604,7 @@ impl<R: FromLuaMulti> Future for AsyncThread<R> {
if status.is_yielded() {
if !(nresults == 1 && is_poll_pending(thread_state)) {
// Ignore value returned via yield()
// Ignore values returned via yield()
cx.waker().wake_by_ref();
}
return Poll::Pending;
@@ -635,7 +635,7 @@ struct WakerGuard<'lua, 'a> {
impl<'lua, 'a> WakerGuard<'lua, 'a> {
#[inline]
pub fn new(lua: &'lua RawLua, waker: &'a Waker) -> Result<WakerGuard<'lua, 'a>> {
let prev = unsafe { lua.set_waker(NonNull::from(waker)) };
let prev = lua.set_waker(NonNull::from(waker));
Ok(WakerGuard {
lua,
prev,
@@ -647,7 +647,7 @@ impl<'lua, 'a> WakerGuard<'lua, 'a> {
#[cfg(feature = "async")]
impl Drop for WakerGuard<'_, '_> {
fn drop(&mut self) {
unsafe { self.lua.set_waker(self.prev) };
self.lua.set_waker(self.prev);
}
}
+42 -2
View File
@@ -5,9 +5,9 @@ use std::sync::Arc;
use crate::error::{Error, Result};
use crate::multi::MultiValue;
use crate::private::Sealed;
use crate::state::{Lua, RawLua};
use crate::state::{Lua, RawLua, WeakLua};
use crate::types::MaybeSend;
use crate::util::{check_stack, short_type_name};
use crate::util::{check_stack, parse_lookup_path, short_type_name};
use crate::value::Value;
#[cfg(feature = "async")]
@@ -200,10 +200,50 @@ pub trait ObjectLike: Sealed {
where
R: FromLuaMulti;
/// Look up a value by a path of keys.
///
/// The syntax is similar to accessing nested tables in Lua, with additional support for
/// `?` operator to perform safe navigation.
///
/// For example, the path `a[1].c` is equivalent to `table.a[1].c` in Lua.
/// With `?` operator, `a[1]?.c` is equivalent to `table.a[1] and table.a[1].c or nil` in Lua.
///
/// Bracket notation rules:
/// - `[123]` - integer keys
/// - `["string key"]` or `['string key']` - string keys (must be quoted)
/// - String keys support escape sequences: `\"`, `\'`, `\\`
fn get_path<V: FromLua>(&self, path: &str) -> Result<V> {
let mut current = self.to_value();
for (key, safe_nil) in parse_lookup_path(path)? {
current = match current {
Value::Table(table) => table.get::<Value>(key),
Value::UserData(ud) => ud.get::<Value>(key),
_ => {
let type_name = current.type_name();
let err = format!("attempt to index a {type_name} value with key '{key}'");
Err(Error::runtime(err))
}
}?;
if safe_nil && (current == Value::Nil || current == Value::NULL) {
break;
}
}
let lua = self.weak_lua().lock();
V::from_lua(current, lua.lua())
}
/// Converts the object to a string in a human-readable format.
///
/// This might invoke the `__tostring` metamethod.
fn to_string(&self) -> Result<StdString>;
/// Converts the object to a Lua value.
fn to_value(&self) -> Value;
/// Gets a reference to the associated Lua state.
#[doc(hidden)]
fn weak_lua(&self) -> &WeakLua;
}
/// A trait for types that can be used as Lua functions.
+9 -2
View File
@@ -20,6 +20,9 @@ pub use either::Either;
pub use registry_key::RegistryKey;
pub(crate) use value_ref::ValueRef;
#[cfg(feature = "async")]
pub(crate) use value_ref::ValueRefIndex;
/// Type of Lua integer numbers.
pub type Integer = ffi::lua_Integer;
/// Type of Lua floating point numbers.
@@ -35,10 +38,13 @@ unsafe impl Send for LightUserData {}
unsafe impl Sync for LightUserData {}
#[cfg(feature = "send")]
pub(crate) type Callback = Box<dyn Fn(&RawLua, c_int) -> Result<c_int> + Send + 'static>;
type CallbackFn<'a> = dyn Fn(&RawLua, c_int) -> Result<c_int> + Send + 'a;
#[cfg(not(feature = "send"))]
pub(crate) type Callback = Box<dyn Fn(&RawLua, c_int) -> Result<c_int> + 'static>;
type CallbackFn<'a> = dyn Fn(&RawLua, c_int) -> Result<c_int> + 'a;
pub(crate) type Callback = Box<CallbackFn<'static>>;
pub(crate) type CallbackPtr = *const CallbackFn<'static>;
pub(crate) type ScopedCallback<'s> = Box<dyn Fn(&RawLua, c_int) -> Result<c_int> + 's>;
@@ -114,6 +120,7 @@ pub trait MaybeSend: Send {}
#[cfg(feature = "send")]
impl<T: Send> MaybeSend for T {}
/// A trait that adds `Send` requirement if `send` feature is enabled.
#[cfg(not(feature = "send"))]
pub trait MaybeSend {}
#[cfg(not(feature = "send"))]
+28 -23
View File
@@ -1,22 +1,39 @@
use std::fmt;
use std::os::raw::{c_int, c_void};
use super::XRc;
use crate::state::{RawLua, WeakLua};
/// A reference to a Lua (complex) value stored in the Lua auxiliary thread.
#[derive(Clone)]
pub struct ValueRef {
pub(crate) lua: WeakLua,
// Keep index separate to avoid additional indirection when accessing it.
pub(crate) index: c_int,
pub(crate) drop: bool,
// If `index_count` is `None`, the value does not need to be destroyed.
pub(crate) index_count: Option<ValueRefIndex>,
}
/// A reference to a Lua value index in the auxiliary thread.
/// It's cheap to clone and can be used to track the number of references to a value.
#[derive(Clone)]
pub(crate) struct ValueRefIndex(pub(crate) XRc<c_int>);
impl From<c_int> for ValueRefIndex {
#[inline]
fn from(index: c_int) -> Self {
ValueRefIndex(XRc::new(index))
}
}
impl ValueRef {
#[inline]
pub(crate) fn new(lua: &RawLua, index: c_int) -> Self {
pub(crate) fn new(lua: &RawLua, index: impl Into<ValueRefIndex>) -> Self {
let index = index.into();
ValueRef {
lua: lua.weak().clone(),
index,
drop: true,
index: *index.0,
index_count: Some(index),
}
}
@@ -25,16 +42,6 @@ impl ValueRef {
let lua = self.lua.lock();
unsafe { ffi::lua_topointer(lua.ref_thread(), self.index) }
}
/// Returns a copy of the value, which is valid as long as the original value is held.
#[inline]
pub(crate) fn copy(&self) -> Self {
ValueRef {
lua: self.lua.clone(),
index: self.index,
drop: false,
}
}
}
impl fmt::Debug for ValueRef {
@@ -43,17 +50,15 @@ impl fmt::Debug for ValueRef {
}
}
impl Clone for ValueRef {
fn clone(&self) -> Self {
unsafe { self.lua.lock().clone_ref(self) }
}
}
impl Drop for ValueRef {
fn drop(&mut self) {
if self.drop {
if let Some(lua) = self.lua.try_lock() {
unsafe { lua.drop_ref(self) };
if let Some(ValueRefIndex(index)) = self.index_count.take() {
// It's guaranteed that the inner value returns exactly once.
// This means in particular that the value is not dropped.
if XRc::into_inner(index).is_some() {
if let Some(lua) = self.lua.try_lock() {
unsafe { lua.drop_ref(self) };
}
}
}
}
+56 -3
View File
@@ -12,7 +12,7 @@ use crate::string::String;
use crate::table::{Table, TablePairs};
use crate::traits::{FromLua, FromLuaMulti, IntoLua, IntoLuaMulti};
use crate::types::{MaybeSend, ValueRef};
use crate::util::{check_stack, get_userdata, push_string, take_userdata, StackGuard};
use crate::util::{check_stack, get_userdata, push_string, short_type_name, take_userdata, StackGuard};
use crate::value::Value;
#[cfg(feature = "async")]
@@ -273,6 +273,29 @@ pub trait UserDataMethods<T> {
A: FromLuaMulti,
R: IntoLuaMulti;
/// Add a method which accepts `T` as the first parameter.
///
/// The userdata `T` will be moved out of the userdata container. This is useful for
/// methods that need to consume the userdata.
///
/// The method can be called only once per userdata instance, subsequent calls will result in a
/// [`Error::UserDataDestructed`] error.
#[doc(hidden)]
fn add_method_once<M, A, R>(&mut self, name: impl Into<StdString>, method: M)
where
T: 'static,
M: Fn(&Lua, T, A) -> Result<R> + MaybeSend + 'static,
A: FromLuaMulti,
R: IntoLuaMulti,
{
let name = name.into();
let method_name = format!("{}.{name}", short_type_name::<T>());
self.add_function(name, move |lua, (ud, args): (AnyUserData, A)| {
let this = (ud.take()).map_err(|err| Error::bad_self_argument(&method_name, err))?;
method(lua, this, args)
});
}
/// Add an async method which accepts a `&T` as the first parameter and returns [`Future`].
///
/// Refer to [`add_method`] for more information about the implementation.
@@ -303,6 +326,34 @@ pub trait UserDataMethods<T> {
MR: Future<Output = Result<R>> + MaybeSend + 'static,
R: IntoLuaMulti;
/// Add an async method which accepts a `T` as the first parameter and returns [`Future`].
///
/// The userdata `T` will be moved out of the userdata container. This is useful for
/// methods that need to consume the userdata.
///
/// The method can be called only once per userdata instance, subsequent calls will result in a
/// [`Error::UserDataDestructed`] error.
#[cfg(feature = "async")]
#[cfg_attr(docsrs, doc(cfg(feature = "async")))]
#[doc(hidden)]
fn add_async_method_once<M, A, MR, R>(&mut self, name: impl Into<StdString>, method: M)
where
T: 'static,
M: Fn(Lua, T, A) -> MR + MaybeSend + 'static,
A: FromLuaMulti,
MR: Future<Output = Result<R>> + MaybeSend + 'static,
R: IntoLuaMulti,
{
let name = name.into();
let method_name = format!("{}.{name}", short_type_name::<T>());
self.add_async_function(name, move |lua, (ud, args): (AnyUserData, A)| {
match (ud.take()).map_err(|err| Error::bad_self_argument(&method_name, err)) {
Ok(this) => either::Either::Left(method(lua, this, args)),
Err(err) => either::Either::Right(async move { Err(err) }),
}
});
}
/// Add a regular method as a function which accepts generic arguments.
///
/// The first argument will be a [`AnyUserData`] of type `T` if the method is called with Lua
@@ -921,8 +972,10 @@ impl AnyUserData {
lua.get_userdata_ref_type_id(&self.0).ok().flatten()
}
/// Returns a type name of this `UserData` (from a metatable field).
pub(crate) fn type_name(&self) -> Result<Option<StdString>> {
/// Returns a type name of this userdata (from a metatable field).
///
/// If no type name is set, returns `None`.
pub fn type_name(&self) -> Result<Option<StdString>> {
let lua = self.0.lua.lock();
let state = lua.state();
unsafe {
+16 -5
View File
@@ -1,6 +1,7 @@
use std::string::String as StdString;
use crate::error::{Error, Result};
use crate::state::WeakLua;
use crate::table::Table;
use crate::traits::{FromLua, FromLuaMulti, IntoLua, IntoLuaMulti, ObjectLike};
use crate::userdata::AnyUserData;
@@ -15,14 +16,14 @@ impl ObjectLike for AnyUserData {
fn get<V: FromLua>(&self, key: impl IntoLua) -> Result<V> {
// `lua_gettable` method used under the hood can work with any Lua value
// that has `__index` metamethod
Table(self.0.copy()).get_protected(key)
Table(self.0.clone()).get_protected(key)
}
#[inline]
fn set(&self, key: impl IntoLua, value: impl IntoLua) -> Result<()> {
// `lua_settable` method used under the hood can work with any Lua value
// that has `__newindex` metamethod
Table(self.0.copy()).set_protected(key, value)
Table(self.0.clone()).set_protected(key, value)
}
#[inline]
@@ -30,7 +31,7 @@ impl ObjectLike for AnyUserData {
where
R: FromLuaMulti,
{
Function(self.0.copy()).call(args)
Function(self.0.clone()).call(args)
}
#[cfg(feature = "async")]
@@ -39,7 +40,7 @@ impl ObjectLike for AnyUserData {
where
R: FromLuaMulti,
{
Function(self.0.copy()).call_async(args)
Function(self.0.clone()).call_async(args)
}
#[inline]
@@ -88,6 +89,16 @@ impl ObjectLike for AnyUserData {
#[inline]
fn to_string(&self) -> Result<StdString> {
Value::UserData(AnyUserData(self.0.copy())).to_string()
Value::UserData(self.clone()).to_string()
}
#[inline]
fn to_value(&self) -> Value {
Value::UserData(self.clone())
}
#[inline]
fn weak_lua(&self) -> &WeakLua {
&self.0.lua
}
}
+22
View File
@@ -56,6 +56,9 @@ pub(crate) struct RawUserDataRegistry {
pub(crate) destructor: ffi::lua_CFunction,
pub(crate) type_id: Option<TypeId>,
pub(crate) type_name: StdString,
#[cfg(feature = "luau")]
pub(crate) enable_namecall: bool,
}
impl UserDataType {
@@ -100,6 +103,8 @@ impl<T> UserDataRegistry<T> {
destructor: super::util::destroy_userdata_storage::<T>,
type_id: r#type.type_id(),
type_name: short_type_name::<T>(),
#[cfg(feature = "luau")]
enable_namecall: false,
};
UserDataRegistry {
@@ -110,6 +115,23 @@ impl<T> UserDataRegistry<T> {
}
}
/// Enables support for the namecall optimization in Luau.
///
/// This enables methods resolution optimization in Luau for complex userdata types with methods
/// and field getters. When enabled, Luau will use a faster lookup path for method calls when a
/// specific syntax is used (e.g. `obj:method()`.
///
/// This optimization does not play well with async methods, custom `__index` metamethod and
/// field getters as functions. So, it is disabled by default.
///
/// Use with caution.
#[doc(hidden)]
#[cfg(feature = "luau")]
#[cfg_attr(docsrs, doc(cfg(feature = "luau")))]
pub fn enable_namecall(&mut self) {
self.raw.enable_namecall = true;
}
fn box_method<M, A, R>(&self, name: &str, method: M) -> Callback
where
M: Fn(&Lua, &T, A) -> Result<R> + MaybeSend + 'static,
+41
View File
@@ -4,8 +4,11 @@ use std::marker::PhantomData;
use std::os::raw::c_int;
use std::ptr;
use rustc_hash::FxHashMap;
use super::UserDataStorage;
use crate::error::{Error, Result};
use crate::types::CallbackPtr;
use crate::util::{get_userdata, rawget_field, rawset_field, take_userdata};
// This is a trick to check if a type is `Sync` or not.
@@ -244,6 +247,7 @@ pub(crate) unsafe fn init_userdata_metatable(
field_getters: Option<c_int>,
field_setters: Option<c_int>,
methods: Option<c_int>,
_methods_map: Option<FxHashMap<Vec<u8>, CallbackPtr>>, // Used only in Luau for `__namecall`
) -> Result<()> {
if field_getters.is_some() || methods.is_some() {
// Push `__index` generator function
@@ -267,6 +271,13 @@ pub(crate) unsafe fn init_userdata_metatable(
}
rawset_field(state, metatable, "__index")?;
#[cfg(feature = "luau")]
if let Some(methods_map) = _methods_map {
// In Luau we can speedup method calls by providing a dedicated `__namecall` metamethod
push_userdata_metatable_namecall(state, methods_map)?;
rawset_field(state, metatable, "__namecall")?;
}
}
if let Some(field_setters) = field_setters {
@@ -425,6 +436,36 @@ unsafe fn init_userdata_metatable_newindex(state: *mut ffi::lua_State) -> Result
})
}
#[cfg(feature = "luau")]
unsafe fn push_userdata_metatable_namecall(
state: *mut ffi::lua_State,
methods_map: FxHashMap<Vec<u8>, CallbackPtr>,
) -> Result<()> {
unsafe extern "C-unwind" fn namecall(state: *mut ffi::lua_State) -> c_int {
let name = ffi::lua_namecallatom(state, ptr::null_mut());
if name.is_null() {
ffi::luaL_error(state, cstr!("attempt to call an unknown method"));
}
let name_cs = std::ffi::CStr::from_ptr(name);
let methods_map = get_userdata::<FxHashMap<Vec<u8>, CallbackPtr>>(state, ffi::lua_upvalueindex(1));
let callback_ptr = match (*methods_map).get(name_cs.to_bytes()) {
Some(ptr) => *ptr,
#[rustfmt::skip]
None => ffi::luaL_error(state, cstr!("attempt to call an unknown method '%s'"), name),
};
crate::state::callback_error_ext(state, ptr::null_mut(), true, |extra, nargs| {
let rawlua = (*extra).raw_lua();
(*callback_ptr)(rawlua, nargs)
})
}
// Automatic destructor is provided for any Luau userdata
crate::util::push_userdata(state, methods_map, true)?;
protect_lua!(state, 1, 1, |state| {
ffi::lua_pushcclosured(state, namecall, cstr!("__namecall"), 1);
})
}
// This method is called by Lua GC when it's time to collect the userdata.
//
// This method is usually used to collect internal userdata.
+2
View File
@@ -402,6 +402,8 @@ pub(crate) unsafe fn init_error_registry(state: *mut ffi::lua_State) -> Result<(
"__ipairs",
#[cfg(feature = "luau")]
"__iter",
#[cfg(feature = "luau")]
"__namecall",
#[cfg(feature = "lua54")]
"__close",
] {
+8 -8
View File
@@ -9,6 +9,7 @@ pub(crate) use error::{
error_traceback, error_traceback_thread, init_error_registry, pop_error, protect_lua_call,
protect_lua_closure, WrappedFailure,
};
pub(crate) use path::parse_path as parse_lookup_path;
pub(crate) use short_names::short_type_name;
pub(crate) use types::TypeKey;
pub(crate) use userdata::{
@@ -101,15 +102,13 @@ pub(crate) unsafe fn push_string(state: *mut ffi::lua_State, s: &[u8], protect:
// Uses 3 stack spaces (when protect), does not call checkstack.
#[cfg(feature = "luau")]
#[inline(always)]
pub(crate) unsafe fn push_buffer(state: *mut ffi::lua_State, b: &[u8], protect: bool) -> Result<()> {
let data = if protect {
protect_lua!(state, 0, 1, |state| ffi::lua_newbuffer(state, b.len()))?
pub(crate) unsafe fn push_buffer(state: *mut ffi::lua_State, size: usize, protect: bool) -> Result<*mut u8> {
let data = if protect || size > const { 1024 * 1024 * 1024 } {
protect_lua!(state, 0, 1, |state| ffi::lua_newbuffer(state, size))?
} else {
ffi::lua_newbuffer(state, b.len())
ffi::lua_newbuffer(state, size)
};
let buf = slice::from_raw_parts_mut(data as *mut u8, b.len());
buf.copy_from_slice(b);
Ok(())
Ok(data as *mut u8)
}
// Uses 3 stack spaces, does not call checkstack.
@@ -122,7 +121,7 @@ pub(crate) unsafe fn push_table(
) -> Result<()> {
let narr: c_int = narr.try_into().unwrap_or(c_int::MAX);
let nrec: c_int = nrec.try_into().unwrap_or(c_int::MAX);
if protect || narr >= const { 1 << 30 } || nrec >= const { 1 << 27 } {
if protect || narr >= const { 1 << 26 } || nrec >= const { 1 << 26 } {
protect_lua!(state, 0, 1, |state| ffi::lua_createtable(state, narr, nrec))
} else {
ffi::lua_createtable(state, narr, nrec);
@@ -329,6 +328,7 @@ pub(crate) fn linenumber_to_usize(n: c_int) -> Option<usize> {
}
mod error;
mod path;
mod short_names;
mod types;
mod userdata;
+255
View File
@@ -0,0 +1,255 @@
use std::borrow::Cow;
use std::fmt;
use std::iter::Peekable;
use std::str::CharIndices;
use crate::error::{Error, Result};
use crate::state::Lua;
use crate::traits::IntoLua;
use crate::types::Integer;
use crate::value::Value;
#[derive(Debug)]
pub(crate) enum PathKey<'a> {
Str(Cow<'a, str>),
Int(Integer),
}
impl fmt::Display for PathKey<'_> {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
match self {
PathKey::Str(s) => write!(f, "{}", s),
PathKey::Int(i) => write!(f, "{}", i),
}
}
}
impl IntoLua for PathKey<'_> {
fn into_lua(self, lua: &Lua) -> Result<Value> {
match self {
PathKey::Str(s) => Ok(Value::String(lua.create_string(s.as_ref())?)),
PathKey::Int(i) => Ok(Value::Integer(i)),
}
}
}
// Parses a path like `a.b[3]?.c["d"]` into segments of `(key, safe_nil)`.
pub(crate) fn parse_path<'a>(path: &'a str) -> Result<Vec<(PathKey<'a>, bool)>> {
fn read_ident<'a>(path: &'a str, chars: &mut Peekable<CharIndices<'a>>) -> (Cow<'a, str>, bool) {
let mut safe_nil = false;
let start = chars.peek().map(|&(i, _)| i).unwrap_or(path.len());
let mut end = start;
while let Some(&(pos, c)) = chars.peek() {
if c == '.' || c == '?' || c.is_ascii_whitespace() || c == '[' {
if c == '?' {
safe_nil = true;
chars.next(); // consume '?'
}
break;
}
end = pos + c.len_utf8();
chars.next();
}
(Cow::Borrowed(&path[start..end]), safe_nil)
}
let mut segments = Vec::new();
let mut chars = path.char_indices().peekable();
while let Some(&(pos, next)) = chars.peek() {
match next {
'.' => {
// Dot notation: identifier
chars.next();
let (key, safe_nil) = read_ident(path, &mut chars);
if key.is_empty() {
return Err(Error::runtime(format!("empty key in path at position {pos}")));
}
segments.push((PathKey::Str(key), safe_nil));
}
'[' => {
// Bracket notation: either integer or quoted string
chars.next();
let key = match chars.peek() {
Some(&(pos, c @ '0'..='9' | c @ '-')) => {
// Integer key
let negative = c == '-';
if negative {
chars.next(); // consume '-'
}
let mut num: Option<Integer> = None;
while let Some(&(_, c @ '0'..='9')) = chars.peek() {
let new_num = num
.unwrap_or(0)
.checked_mul(10)
.and_then(|n| n.checked_add((c as u8 - b'0') as Integer))
.ok_or_else(|| {
Error::runtime(format!("integer overflow in path at position {pos}"))
})?;
num = Some(new_num);
chars.next(); // consume digit
}
match num {
Some(n) if negative => PathKey::Int(-n),
Some(n) => PathKey::Int(n),
None => {
let err = format!("invalid integer in path at position {pos}");
return Err(Error::runtime(err));
}
}
}
Some((_, '\'' | '"')) => {
// Quoted string
PathKey::Str(unquote_string(path, &mut chars)?)
}
Some((_, ']')) => {
return Err(Error::runtime(format!("empty key in path at position {pos}")));
}
Some((pos, c)) => {
let err = format!("unexpected character '{c}' in path at position {pos}");
return Err(Error::runtime(err));
}
None => {
return Err(Error::runtime("unexpected end of path"));
}
};
// Expect closing bracket
let mut safe_nil = false;
match chars.next() {
Some((_, ']')) => {
// Check for optional safe-nil operator
if let Some(&(_, '?')) = chars.peek() {
safe_nil = true;
chars.next(); // consume '?'
}
}
Some((pos, c)) => {
let err = format!("expected ']' in path at position {pos}, found '{c}'");
return Err(Error::runtime(err));
}
None => {
return Err(Error::runtime("unexpected end of path"));
}
}
segments.push((key, safe_nil));
}
c if c.is_ascii_whitespace() => {
chars.next(); // Skip whitespace
}
_ if segments.is_empty() => {
// First segment without dot/bracket notation
let (key_cow, safe_nil) = read_ident(path, &mut chars);
if key_cow.is_empty() {
return Err(Error::runtime(format!("empty key in path at position {pos}")));
}
segments.push((PathKey::Str(key_cow), safe_nil));
}
c => {
let err = format!("unexpected character '{c}' in path at position {pos}");
return Err(Error::runtime(err));
}
}
}
Ok(segments)
}
fn unquote_string<'a>(path: &'a str, chars: &mut Peekable<CharIndices<'a>>) -> Result<Cow<'a, str>> {
let (start_pos, first_quote) = chars.next().unwrap();
let mut result = String::new();
loop {
match chars.next() {
Some((pos, '\\')) => {
if result.is_empty() {
// First escape found, copy everything up to this point
result.push_str(&path[start_pos + 1..pos]);
}
match chars.next() {
Some((_, '\\')) => result.push('\\'),
Some((_, '"')) => result.push('"'),
Some((_, '\'')) => result.push('\''),
Some((_, other)) => {
result.push('\\');
result.push(other);
}
None => continue, // will be handled by outer loop
}
}
Some((pos, c)) if c == first_quote => {
if !result.is_empty() {
return Ok(Cow::Owned(result));
}
// No escapes, return borrowed slice
return Ok(Cow::Borrowed(&path[start_pos + 1..pos]));
}
Some((_, c)) => {
if !result.is_empty() {
result.push(c);
}
// If no escapes yet, continue tracking for potential borrowed slice
}
None => {
let err = format!("unexpected end of string at position {start_pos}");
return Err(Error::runtime(err));
}
}
}
}
#[cfg(test)]
mod tests {
use super::{parse_path, PathKey};
#[test]
fn test_parse_path() {
// Test valid paths
let path = parse_path("a.b[3]?.c['d']").unwrap();
assert_eq!(path.len(), 5);
assert!(matches!(path[0], (PathKey::Str(ref s), false) if s == "a"));
assert!(matches!(path[1], (PathKey::Str(ref s), false) if s == "b"));
assert!(matches!(path[2], (PathKey::Int(3), true)));
assert!(matches!(path[3], (PathKey::Str(ref s), false) if s == "c"));
assert!(matches!(path[4], (PathKey::Str(ref s), false) if s == "d"));
// Test empty path
let path = parse_path("").unwrap();
assert_eq!(path.len(), 0);
let path = parse_path(" ").unwrap();
assert_eq!(path.len(), 0);
// Test invalid dot syntax
let err = parse_path("a..b").unwrap_err().to_string();
assert_eq!(err, "runtime error: empty key in path at position 1");
let err = parse_path("a.b.").unwrap_err().to_string();
assert_eq!(err, "runtime error: empty key in path at position 3");
// Test invalid bracket syntax
let err = parse_path("a[unclosed").unwrap_err().to_string();
assert_eq!(
err,
"runtime error: unexpected character 'u' in path at position 2"
);
let err = parse_path("a[]").unwrap_err().to_string();
assert_eq!(err, "runtime error: empty key in path at position 1");
let err = parse_path(r#"a["unclosed"#).unwrap_err().to_string();
assert_eq!(err, "runtime error: unexpected end of string at position 2");
let err = parse_path(r#"a["#).unwrap_err().to_string();
assert_eq!(err, "runtime error: unexpected end of path");
let err = parse_path(r#"a[123"#).unwrap_err().to_string();
assert_eq!(err, "runtime error: unexpected end of path");
let err = parse_path(r#"a['bla'123"#).unwrap_err().to_string();
assert_eq!(
err,
"runtime error: expected ']' in path at position 7, found '1'"
);
let err = parse_path(r#"a["bla"]x"#).unwrap_err().to_string();
assert_eq!(
err,
"runtime error: unexpected character 'x' in path at position 8"
);
// Test bad integers
let err = parse_path("a[99999999999999999999]").unwrap_err().to_string();
assert_eq!(err, "runtime error: integer overflow in path at position 2");
let err = parse_path("a[-]").unwrap_err().to_string();
assert_eq!(err, "runtime error: invalid integer in path at position 2");
}
}
+18 -8
View File
@@ -1,6 +1,6 @@
//! Mostly copied from [bevy_utils]
//! Inspired by bevy's [disqualified]
//!
//! [bevy_utils]: https://github.com/bevyengine/bevy/blob/main/crates/bevy_utils/src/short_names.rs
//! [disqualified]: https://github.com/bevyengine/disqualified/blob/main/src/short_name.rs
use std::any::type_name;
@@ -23,8 +23,7 @@ pub(crate) fn short_type_name<T: ?Sized>() -> String {
while index < end_of_string {
let rest_of_string = full_name.get(index..end_of_string).unwrap_or_default();
// Collapse everything up to the next special character,
// then skip over it
// Collapse everything up to the next special character, then skip over it
if let Some(special_character_index) =
rest_of_string.find(|c: char| [' ', '<', '>', '(', ')', '[', ']', ',', ';'].contains(&c))
{
@@ -32,11 +31,16 @@ pub(crate) fn short_type_name<T: ?Sized>() -> String {
parsed_name += collapse_type_name(segment_to_collapse);
// Insert the special character
let special_character = &rest_of_string[special_character_index..=special_character_index];
parsed_name.push_str(special_character);
parsed_name += special_character;
// Remove lifetimes like <'_> or <'_, '_, ...>
if parsed_name.ends_with("<'_>") || parsed_name.ends_with("<'_, ") {
_ = parsed_name.split_off(parsed_name.len() - 4);
}
match special_character {
">" | ")" | "]" if rest_of_string[special_character_index + 1..].starts_with("::") => {
parsed_name.push_str("::");
parsed_name += "::";
// Move the index past the "::"
index += special_character_index + 3;
}
@@ -53,14 +57,18 @@ pub(crate) fn short_type_name<T: ?Sized>() -> String {
}
#[inline(always)]
fn collapse_type_name(string: &str) -> &str {
string.rsplit("::").next().unwrap()
fn collapse_type_name(segment: &str) -> &str {
segment.rsplit("::").next().unwrap()
}
#[cfg(test)]
mod tests {
use super::short_type_name;
use std::collections::HashMap;
use std::marker::PhantomData;
struct MyData<'a, 'b>(PhantomData<&'a &'b ()>);
struct MyDataT<'a, T>(PhantomData<&'a T>);
#[test]
fn tests() {
@@ -73,5 +81,7 @@ mod tests {
"HashMap<String, Option<[i32; 3]>>"
);
assert_eq!(short_type_name::<dyn Fn(i32) -> i32>(), "dyn Fn(i32) -> i32");
assert_eq!(short_type_name::<MyDataT<&str>>(), "MyDataT<&str>");
assert_eq!(short_type_name::<(&MyData, [MyData])>(), "(MyData, [MyData])");
}
}
+17 -12
View File
@@ -28,9 +28,10 @@ use {
/// The non-primitive variants (eg. string/table/function/thread/userdata) contain handle types
/// into the internal Lua state. It is a logic error to mix handle types between separate
/// `Lua` instances, and doing so will result in a panic.
#[derive(Clone)]
#[derive(Clone, Default)]
pub enum Value {
/// The Lua value `nil`.
#[default]
Nil,
/// The Lua value `true` or `false`.
Boolean(bool),
@@ -491,7 +492,6 @@ impl Value {
/// This allows customizing serialization behavior using serde.
#[cfg(feature = "serde")]
#[cfg_attr(docsrs, doc(cfg(feature = "serde")))]
#[doc(hidden)]
pub fn to_serializable(&self) -> SerializableValue<'_> {
SerializableValue::new(self, Default::default(), None)
}
@@ -580,12 +580,6 @@ impl Value {
}
}
impl Default for Value {
fn default() -> Self {
Self::Nil
}
}
impl fmt::Debug for Value {
fn fmt(&self, fmt: &mut fmt::Formatter) -> fmt::Result {
if fmt.alternate() {
@@ -684,7 +678,7 @@ impl<'a> SerializableValue<'a> {
///
/// Default: **true**
#[must_use]
pub const fn deny_unsupported_types(mut self, enabled: bool) -> Self {
pub fn deny_unsupported_types(mut self, enabled: bool) -> Self {
self.options.deny_unsupported_types = enabled;
self
}
@@ -695,7 +689,7 @@ impl<'a> SerializableValue<'a> {
///
/// Default: **true**
#[must_use]
pub const fn deny_recursive_tables(mut self, enabled: bool) -> Self {
pub fn deny_recursive_tables(mut self, enabled: bool) -> Self {
self.options.deny_recursive_tables = enabled;
self
}
@@ -704,7 +698,7 @@ impl<'a> SerializableValue<'a> {
///
/// Default: **false**
#[must_use]
pub const fn sort_keys(mut self, enabled: bool) -> Self {
pub fn sort_keys(mut self, enabled: bool) -> Self {
self.options.sort_keys = enabled;
self
}
@@ -713,10 +707,21 @@ impl<'a> SerializableValue<'a> {
///
/// Default: **false**
#[must_use]
pub const fn encode_empty_tables_as_array(mut self, enabled: bool) -> Self {
pub fn encode_empty_tables_as_array(mut self, enabled: bool) -> Self {
self.options.encode_empty_tables_as_array = enabled;
self
}
/// If true, enable detection of mixed tables.
///
/// A mixed table is a table that has both array-like and map-like entries or several borders.
///
/// Default: **false**
#[must_use]
pub fn detect_mixed_tables(mut self, enabled: bool) -> Self {
self.options.detect_mixed_tables = enabled;
self
}
}
#[cfg(feature = "serde")]
+56 -4
View File
@@ -8,7 +8,7 @@ use futures_util::stream::TryStreamExt;
use tokio::sync::Mutex;
use mlua::{
Error, Function, Lua, LuaOptions, MultiValue, ObjectLike, Result, StdLib, Table, UserData,
Error, Function, Lua, LuaOptions, MultiValue, ObjectLike, Result, StdLib, Table, ThreadStatus, UserData,
UserDataMethods, UserDataRef, Value,
};
@@ -423,9 +423,9 @@ async fn test_async_thread_pool() -> Result<()> {
#[tokio::test]
async fn test_async_userdata() -> Result<()> {
struct MyUserData(u64);
struct MyUserdata(u64);
impl UserData for MyUserData {
impl UserData for MyUserdata {
fn add_methods<M: UserDataMethods<Self>>(methods: &mut M) {
methods.add_async_method("get_value", |_, data, ()| async move {
sleep_ms(10).await;
@@ -438,6 +438,11 @@ async fn test_async_userdata() -> Result<()> {
Ok(())
});
methods.add_async_method_once("take_value", |_, data, ()| async move {
sleep_ms(10).await;
Ok(data.0)
});
methods.add_async_function("sleep", |_, n| async move {
sleep_ms(n).await;
Ok(format!("elapsed:{}ms", n))
@@ -479,7 +484,7 @@ async fn test_async_userdata() -> Result<()> {
let lua = Lua::new();
let globals = lua.globals();
let userdata = lua.create_userdata(MyUserData(11))?;
let userdata = lua.create_userdata(MyUserdata(11))?;
globals.set("userdata", &userdata)?;
lua.load(
@@ -518,6 +523,21 @@ async fn test_async_userdata() -> Result<()> {
#[cfg(not(any(feature = "lua51", feature = "luau")))]
assert_eq!(userdata.call_async::<String>(()).await?, "elapsed:24ms");
// Take value
let userdata2 = lua.create_userdata(MyUserdata(0))?;
globals.set("userdata2", userdata2)?;
lua.load("assert(userdata:take_value() == 24)")
.exec_async()
.await?;
match lua.load("userdata2.take_value(userdata)").exec_async().await {
Err(Error::CallbackError { cause, .. }) => {
let err = cause.to_string();
assert!(err.contains("bad argument `self` to `MyUserdata.take_value`"));
assert!(err.contains("userdata has been destructed"));
}
r => panic!("expected Err(CallbackError), got {r:?}"),
}
Ok(())
}
@@ -667,3 +687,35 @@ async fn test_async_hook() -> Result<()> {
Ok(())
}
#[test]
fn test_async_yield_with() -> Result<()> {
let lua = Lua::new();
let func = lua.create_async_function(|lua, (mut a, mut b): (i32, i32)| async move {
let zero = lua.yield_with::<MultiValue>(()).await?;
assert!(zero.is_empty());
let one = lua.yield_with::<MultiValue>(a + b).await?;
assert_eq!(one.len(), 1);
for _ in 0..3 {
(a, b) = lua.yield_with((a + b, a * b)).await?;
}
Ok((0, 0))
})?;
let thread = lua.create_thread(func)?;
let zero = thread.resume::<MultiValue>((2, 3))?; // function arguments
assert!(zero.is_empty());
let one = thread.resume::<i32>(())?; // value of "zero" is passed here
assert_eq!(one, 5);
assert_eq!(thread.resume::<(i32, i32)>(1)?, (5, 6)); // value of "one" is passed here
assert_eq!(thread.resume::<(i32, i32)>((10, 11))?, (21, 110));
assert_eq!(thread.resume::<(i32, i32)>((11, 12))?, (23, 132));
assert_eq!(thread.resume::<(i32, i32)>((12, 13))?, (0, 0));
assert_eq!(thread.status(), ThreadStatus::Finished);
Ok(())
}
+68 -2
View File
@@ -1,5 +1,7 @@
#![cfg(feature = "luau")]
use std::io::{Read, Seek, SeekFrom, Write};
use mlua::{Lua, Result, Value};
#[test]
@@ -41,7 +43,7 @@ fn test_buffer() -> Result<()> {
}
#[test]
#[should_panic(expected = "range end index 14 out of range for slice of length 13")]
#[should_panic(expected = "out of range for slice of length 13")]
fn test_buffer_out_of_bounds_read() {
let lua = Lua::new();
let buf = lua.create_buffer(b"hello, world!").unwrap();
@@ -49,9 +51,73 @@ fn test_buffer_out_of_bounds_read() {
}
#[test]
#[should_panic(expected = "range end index 16 out of range for slice of length 13")]
#[should_panic(expected = "out of range for slice of length 13")]
fn test_buffer_out_of_bounds_write() {
let lua = Lua::new();
let buf = lua.create_buffer(b"hello, world!").unwrap();
buf.write_bytes(14, b"!!");
}
#[test]
fn create_large_buffer() {
let lua = Lua::new();
let err = lua.create_buffer_with_capacity(1_073_741_824 + 1).unwrap_err(); // 1GB
assert!(err.to_string().contains("memory allocation error"));
// Normal buffer is okay
let buf = lua.create_buffer_with_capacity(1024 * 1024).unwrap();
assert_eq!(buf.len(), 1024 * 1024);
}
#[test]
fn test_buffer_cursor() -> Result<()> {
let lua = Lua::new();
let mut cursor = lua.create_buffer(b"hello, world")?.cursor();
let mut data = Vec::new();
cursor.read_to_end(&mut data)?;
assert_eq!(data, b"hello, world");
// No more data to read
let mut one = [0u8; 1];
assert_eq!(cursor.read(&mut one)?, 0);
// Seek to start
cursor.seek(SeekFrom::Start(0))?;
cursor.read_exact(&mut one)?;
assert_eq!(one, [b'h']);
// Seek to end -5
cursor.seek(SeekFrom::End(-5))?;
let mut five = [0u8; 5];
cursor.read_exact(&mut five)?;
assert_eq!(&five, b"world");
// Seek to current -1
cursor.seek(SeekFrom::Current(-1))?;
cursor.read_exact(&mut one)?;
assert_eq!(one, [b'd']);
// Invalid seek
assert!(cursor.seek(SeekFrom::Current(-100)).is_err());
assert!(cursor.seek(SeekFrom::End(1)).is_err());
// Write data
let buf = lua.create_buffer_with_capacity(100)?;
cursor = buf.clone().cursor();
cursor.write_all(b"hello, ...")?;
cursor.seek(SeekFrom::Current(-3))?;
cursor.write_all(b"Rust!")?;
assert_eq!(&buf.read_bytes::<12>(0), b"hello, Rust!");
// Writing beyond the end of the buffer does nothing
cursor.seek(SeekFrom::End(0))?;
assert_eq!(cursor.write(b".")?, 0);
// Flush is no-op
cursor.flush()?;
Ok(())
}
+1
View File
@@ -21,6 +21,7 @@ fn test_chunk_methods() -> Result<()> {
}
#[test]
#[cfg(not(target_os = "wasi"))]
fn test_chunk_path() -> Result<()> {
let lua = Lua::new();
+74
View File
@@ -330,6 +330,15 @@ fn test_interrupts() -> Result<()> {
assert_eq!(yield_count.load(Ordering::Relaxed), 7);
assert_eq!(co.status(), ThreadStatus::Finished);
// Test no yielding at non-yieldable points
yield_count.store(0, Ordering::Relaxed);
let co = lua.create_thread(lua.create_function(|lua, arg: Value| {
(lua.load("return (function(x) return x end)(...)")).call::<Value>(arg)
})?)?;
let res = co.resume::<String>("abc")?;
assert_eq!(res, "abc".to_string());
assert_eq!(yield_count.load(Ordering::Relaxed), 3);
//
// Test errors in interrupts
//
@@ -461,5 +470,70 @@ fn test_typeof_error() -> Result<()> {
Ok(())
}
#[test]
fn test_memory_category() -> Result<()> {
let lua = Lua::new();
lua.set_memory_category("main").unwrap();
// Invalid category names should be rejected
let err = lua.set_memory_category("invalid$");
assert!(err.is_err());
for i in 0..254 {
let name = format!("category_{}", i);
lua.set_memory_category(&name).unwrap();
}
// 255th category should fail
let err = lua.set_memory_category("category_254");
assert!(err.is_err());
Ok(())
}
#[test]
fn test_heap_dump() -> Result<()> {
let lua = Lua::new();
// Assign a new memory category and create few objects
lua.set_memory_category("test_category")?;
let _t = lua.create_table()?;
let _ud = lua.create_any_userdata("hello, world")?;
let dump = lua.heap_dump()?;
assert!(dump.size() > 0);
let size_by_category = dump.size_by_category();
assert_eq!(size_by_category.len(), 2);
assert!(size_by_category.contains_key("test_category"));
assert!(size_by_category["main"] < dump.size());
// Check size by type within the category
let size_by_type = dump.size_by_type(Some("test_category"));
assert!(!size_by_type.is_empty());
assert!(size_by_type.contains_key("table"));
assert!(size_by_type.contains_key("userdata"));
// Try non-existent category
let size_by_type2 = dump.size_by_type(Some("non_existent_category"));
assert!(size_by_type2.is_empty());
// Remove category filter
let size_by_type_all = dump.size_by_type(None);
assert!(size_by_type.len() < size_by_type_all.len());
// Check size by userdata type within the category
let size_by_udtype = dump.size_by_userdata(Some("test_category"));
assert_eq!(size_by_udtype.len(), 1);
assert!(size_by_udtype.contains_key("&str"));
assert_eq!(size_by_udtype["&str"].0, 1);
// Try non-existent category
let size_by_udtype2 = dump.size_by_userdata(Some("non_existent_category"));
assert!(size_by_udtype2.is_empty());
// Remove category filter
let size_by_udtype_all = dump.size_by_userdata(None);
assert!(size_by_udtype.len() < size_by_udtype_all.len());
Ok(())
}
#[path = "luau/require.rs"]
mod require;
+62 -18
View File
@@ -1,7 +1,7 @@
use std::io::Result as IoResult;
use std::result::Result as StdResult;
use mlua::{Error, IntoLua, Lua, MultiValue, NavigateError, Require, Result, TextRequirer, Value};
use mlua::{Error, FromLua, IntoLua, Lua, MultiValue, NavigateError, Require, Result, TextRequirer, Value};
fn run_require(lua: &Lua, path: impl IntoLua) -> Result<Value> {
lua.load(r#"return require(...)"#).call(path)
@@ -11,9 +11,14 @@ fn run_require_pcall(lua: &Lua, path: impl IntoLua) -> Result<MultiValue> {
lua.load(r#"return pcall(require, ...)"#).call(path)
}
#[track_caller]
fn get_value<V: FromLua>(value: &Value, key: impl IntoLua) -> V {
value.as_table().unwrap().get(key).unwrap()
}
#[track_caller]
fn get_str(value: &Value, key: impl IntoLua) -> String {
value.as_table().unwrap().get::<String>(key).unwrap()
get_value(value, key)
}
#[test]
@@ -47,6 +52,16 @@ fn test_require_errors() {
assert!(res.is_err());
assert!((res.unwrap_err().to_string()).contains("require is not supported in this context"));
// RequireAliasThatDoesNotExist
let res = run_require(&lua, "@this.alias.does.not.exist");
assert!(res.is_err());
assert!((res.unwrap_err().to_string()).contains("@this.alias.does.not.exist is not a valid alias"));
// IllegalAlias
let res = run_require(&lua, "@");
assert!(res.is_err());
assert!((res.unwrap_err().to_string()).contains("@ is not a valid alias"));
// Test throwing mlua::Error
struct MyRequire(TextRequirer);
@@ -171,35 +186,64 @@ fn test_require_without_config() {
assert!(res.is_table());
}
#[test]
fn test_require_with_config() {
fn test_require_with_config_inner(r#type: &str) {
let lua = Lua::new();
let base_path = format!("./tests/luau/require/{type}");
// RequirePathWithAlias
let res = run_require(&lua, "./tests/luau/require/with_config/src/alias_requirer").unwrap();
let res = run_require(&lua, format!("{base_path}/src/alias_requirer")).unwrap();
assert_eq!("result from dependency", get_str(&res, 1));
// RequirePathWithAlias (case-insensitive)
let res2 = run_require(&lua, format!("{base_path}/src/alias_requirer_uc")).unwrap();
assert_eq!("result from dependency", get_str(&res2, 1));
assert_eq!(res.to_pointer(), res2.to_pointer());
// RequirePathWithParentAlias
let res = run_require(&lua, "./tests/luau/require/with_config/src/parent_alias_requirer").unwrap();
let res = run_require(&lua, format!("{base_path}/src/parent_alias_requirer")).unwrap();
assert_eq!("result from other_dependency", get_str(&res, 1));
// RequirePathWithAliasPointingToDirectory
let res = run_require(
&lua,
"./tests/luau/require/with_config/src/directory_alias_requirer",
)
.unwrap();
let res = run_require(&lua, format!("{base_path}/src/directory_alias_requirer")).unwrap();
assert_eq!("result from subdirectory_dependency", get_str(&res, 1));
// RequireAliasThatDoesNotExist
let res = run_require(&lua, "@this.alias.does.not.exist");
assert!(res.is_err());
assert!((res.unwrap_err().to_string()).contains("@this.alias.does.not.exist is not a valid alias"));
// RequireChainedAliasesSuccess
let res = run_require(
&lua,
format!("{base_path}/chained_aliases/subdirectory/successful_requirer"),
)
.unwrap();
assert_eq!("result from inner_dependency", get_str(&get_value(&res, 1), 1));
assert_eq!("result from outer_dependency", get_str(&get_value(&res, 2), 1));
// IllegalAlias
let res = run_require(&lua, "@");
// RequireChainedAliasesFailureCyclic
let res = run_require(
&lua,
format!("{base_path}/chained_aliases/subdirectory/failing_requirer_cyclic"),
);
assert!(res.is_err());
assert!((res.unwrap_err().to_string()).contains("@ is not a valid alias"));
let err_msg = "error requiring module \"@cyclicentry\": detected alias cycle (@cyclic1 -> @cyclic2 -> @cyclic3 -> @cyclic1)";
assert!(res.unwrap_err().to_string().contains(err_msg));
// RequireChainedAliasesFailureMissing
let res = run_require(
&lua,
format!("{base_path}/chained_aliases/subdirectory/failing_requirer_missing"),
);
assert!(res.is_err());
let err_msg = "error requiring module \"@brokenchain\": @missing is not a valid alias";
assert!(res.unwrap_err().to_string().contains(err_msg));
}
#[test]
fn test_require_with_config() {
test_require_with_config_inner("with_config");
}
#[test]
fn test_require_with_config_luau() {
test_require_with_config_inner("with_config_luau");
}
#[cfg(all(feature = "async", not(windows)))]
@@ -0,0 +1,9 @@
{
"aliases":{
"outer": "./",
"cyclicentry": "@cyclic1",
"cyclic1": "@cyclic2",
"cyclic2": "@cyclic3",
"cyclic3": "@cyclic1"
}
}
@@ -0,0 +1 @@
return {"result from outer_dependency"}
@@ -0,0 +1,10 @@
{
"aliases":{
"passthroughinner": "./inner_dependency",
"passthroughouter": "@outer",
"dep": "@passthroughinner",
"outerdep": "@outer/outer_dependency",
"outerdir": "@passthroughouter",
"brokenchain": "@missing"
}
}
@@ -0,0 +1 @@
return require("@cyclicentry")
@@ -0,0 +1 @@
return require("@brokenchain")
@@ -0,0 +1 @@
return {"result from inner_dependency"}
@@ -0,0 +1,7 @@
local result = {}
table.insert(result, require("@dep"))
table.insert(result, require("@outerdep"))
table.insert(result, require("@outerdir/outer_dependency"))
return result
@@ -0,0 +1 @@
return require("@DeP")
@@ -0,0 +1,8 @@
return {
luau = {
aliases = {
dep = "./this_should_be_overwritten_by_child_luaurc",
otherdep = "./src/other_dependency"
}
}
}
@@ -0,0 +1,11 @@
return {
luau = {
aliases = {
outer = "./",
cyclicentry = "@cyclic1",
cyclic1 = "@cyclic2",
cyclic2 = "@cyclic3",
cyclic3 = "@cyclic1"
}
}
}
@@ -0,0 +1 @@
return {"result from outer_dependency"}
@@ -0,0 +1,12 @@
return {
luau = {
aliases = {
passthroughinner = "./inner_dependency",
passthroughouter = "@outer",
dep = "@passthroughinner",
outerdep = "@outer/outer_dependency",
outerdir = "@passthroughouter",
brokenchain = "@missing"
}
}
}
@@ -0,0 +1 @@
return require("@cyclicentry")
@@ -0,0 +1 @@
return require("@brokenchain")
@@ -0,0 +1 @@
return {"result from inner_dependency"}
@@ -0,0 +1,7 @@
local result = {}
table.insert(result, require("@dep"))
table.insert(result, require("@outerdep"))
table.insert(result, require("@outerdir/outer_dependency"))
return result
@@ -0,0 +1,8 @@
return {
luau = {
aliases = {
dep = "./dependency",
subdir = "./subdirectory"
}
}
}
@@ -0,0 +1 @@
return require("@dep")
@@ -0,0 +1 @@
return require("@DeP")
@@ -0,0 +1 @@
return {"result from dependency"}
@@ -0,0 +1 @@
return(require("@subdir/subdirectory_dependency"))
@@ -0,0 +1 @@
return {"result from other_dependency"}
@@ -0,0 +1 @@
return require("@otherdep")
@@ -0,0 +1 @@
return {"result from subdirectory_dependency"}
+20
View File
@@ -72,6 +72,26 @@ fn test_multivalue() {
let _multi2 = MultiValue::from_vec(vec);
}
#[test]
fn test_multivalue_by_ref() -> Result<()> {
let lua = Lua::new();
let multi = MultiValue::from_vec(vec![
Value::Integer(3),
Value::String(lua.create_string("hello")?),
Value::Boolean(true),
]);
let f = lua.create_function(|_, (i, s, b): (i32, String, bool)| {
assert_eq!(i, 3);
assert_eq!(s.to_str()?, "hello");
assert_eq!(b, true);
Ok(())
})?;
f.call::<()>(&multi)?;
Ok(())
}
#[test]
fn test_variadic() {
let mut var = Variadic::with_capacity(3);
+4 -2
View File
@@ -382,7 +382,8 @@ fn test_scope_userdata_ref() -> Result<()> {
modify_userdata(&lua, &ud)?;
// We can only borrow userdata scoped
assert!((matches!(ud.borrow::<MyUserData>(), Err(Error::UserDataTypeMismatch))));
#[rustfmt::skip]
assert!(matches!(ud.borrow::<MyUserData>(), Err(Error::UserDataTypeMismatch)));
ud.borrow_scoped::<MyUserData, ()>(|ud_inst| {
assert_eq!(ud_inst.0.get(), 2);
})?;
@@ -419,7 +420,8 @@ fn test_scope_userdata_ref_mut() -> Result<()> {
let ud = scope.create_userdata_ref_mut(&mut data)?;
modify_userdata(&lua, &ud)?;
assert!((matches!(ud.borrow_mut::<MyUserData>(), Err(Error::UserDataTypeMismatch))));
#[rustfmt::skip]
assert!(matches!(ud.borrow_mut::<MyUserData>(), Err(Error::UserDataTypeMismatch)));
ud.borrow_mut_scoped::<MyUserData, ()>(|ud_inst| {
ud_inst.0 += 10;
})?;
+1
View File
@@ -47,6 +47,7 @@ fn test_userdata_multithread_access_send_only() -> Result<()> {
Ok(())
}
#[rustversion::stable]
#[test]
fn test_userdata_multithread_access_sync() -> Result<()> {
let lua = Lua::new();
+39
View File
@@ -269,6 +269,45 @@ fn test_serialize_empty_table() -> LuaResult<()> {
Ok(())
}
#[test]
fn test_serialize_mixed_table() -> LuaResult<()> {
let lua = Lua::new();
// Check that sparse array is serialized similarly when using direct serialization
// and via `Lua::from_value`
let table = lua.load("{1,2,3,nil,5}").eval::<Value>()?;
let json1 = serde_json::to_string(&table).unwrap();
let json2 = lua.from_value::<serde_json::Value>(table)?;
assert_eq!(json1, json2.to_string());
// A table with several borders should be correctly encoded when `detect_mixed_tables` is enabled
let table = lua
.load(
r#"
local t = {1,2,3,nil,5,6}
t[10] = 10
return t
"#,
)
.eval::<Value>()?;
let json = serde_json::to_string(&table.to_serializable().detect_mixed_tables(true)).unwrap();
assert_eq!(json, r#"[1,2,3,null,5,6,null,null,null,10]"#);
// A mixed table with both array-like and map-like entries
let table = lua.load(r#"{1,2,3, key="value"}"#).eval::<Value>()?;
let json = serde_json::to_string(&table).unwrap();
assert_eq!(json, r#"[1,2,3]"#);
let json = serde_json::to_string(&table.to_serializable().detect_mixed_tables(true)).unwrap();
assert_eq!(json, r#"{"1":1,"2":2,"3":3,"key":"value"}"#);
// A mixed table with duplicate keys of different types
let table = lua.load(r#"{1,2,3, ["1"]="value"}"#).eval::<Value>()?;
let json = serde_json::to_string(&table.to_serializable().detect_mixed_tables(true)).unwrap();
assert_eq!(json, r#"{"1":1,"2":2,"3":3,"1":"value"}"#);
Ok(())
}
#[test]
fn test_to_value_struct() -> LuaResult<()> {
let lua = Lua::new();
+106
View File
@@ -61,6 +61,15 @@ fn test_table() -> Result<()> {
Ok(())
}
#[test]
#[cfg(target_os = "linux")] // Linux allow overcommiting the memory (relevant for CI)
fn test_table_with_large_capacity() {
let lua = Lua::new();
let t = lua.create_table_with_capacity(1 << 26, 1 << 26);
assert!(t.is_ok());
}
#[test]
fn test_table_push_pop() -> Result<()> {
let lua = Lua::new();
@@ -263,6 +272,22 @@ fn test_table_for_each() -> Result<()> {
Ok(())
}
#[test]
fn test_table_for_each_value() -> Result<()> {
let lua = Lua::new();
let table = lua.load("{1, 2, 3, 4, 5, nil, 7}").eval::<Table>()?;
let mut sum = 0;
table.for_each_value::<i32>(|v| {
sum += v;
Ok(())
})?;
// Iterations stops at the first nil
assert_eq!(sum, 1 + 2 + 3 + 4 + 5);
Ok(())
}
#[test]
fn test_table_scope() -> Result<()> {
let lua = Lua::new();
@@ -473,3 +498,84 @@ fn test_table_object_like() -> Result<()> {
Ok(())
}
#[test]
fn test_table_get_path() -> Result<()> {
let lua = Lua::new();
// Create a nested table structure
let table = lua
.load(
r#"
{
a = {
b = {
c = "hello",
d = 42
},
[1] = "first",
["special key"] = "special value"
},
abc = "top level",
x = {},
["🚀"] = "rocket",
[1] = {
["nested-key"] = {
[42] = {
final = "hello!",
},
},
["key\"with\"quotes"] = "value1",
["key'with'quotes"] = "value2",
["key\\with\\backslashes"] = "value3",
[-2] = "negative index",
},
}
"#,
)
.eval::<Table>()?;
// Test basic dot notation
assert_eq!(table.get_path::<String>(".a.b.c")?, "hello");
assert_eq!(table.get_path::<String>("a.b.c")?, "hello");
assert_eq!(table.get_path::<i32>("a.b.d")?, 42);
assert_eq!(table.get_path::<String>("abc")?, "top level");
// Test bracket notation with integer keys
assert_eq!(table.get_path::<String>("a[1]")?, "first");
assert_eq!(table.get_path::<String>("[1][-2]")?, "negative index");
// Test bracket notation with string keys
assert_eq!(table.get_path::<String>("a[\"special key\"]")?, "special value");
assert_eq!(table.get_path::<String>("a['special key']")?, "special value");
assert_eq!(table.get_path::<String>(r#"[1]["key\"with\"quotes"]"#)?, "value1");
assert_eq!(table.get_path::<String>(r#"[1]['key"with"quotes']"#)?, "value1");
assert_eq!(table.get_path::<String>(r#"[1]['key\'with\'quotes']"#)?, "value2");
assert_eq!(
table.get_path::<String>(r#"[1]["key\\with\\backslashes"]"#)?,
"value3"
);
// Test mixed notation
assert_eq!(table.get_path::<String>("[1].nested-key[42].final")?, "hello!");
// Test unicode keys
assert_eq!(table.get_path::<String>("🚀")?, "rocket");
// Test empty path returns the table itself
assert_eq!(table.get_path::<Table>("")?, table);
// Test safe navigation
assert_eq!(table.get_path::<String>("a?.b.c")?, "hello");
assert_eq!(table.get_path::<Value>("x.y?.z")?, Value::Nil);
assert_eq!(table.get_path::<Value>("[1].nested-key[43]?.final")?, Value::Nil);
// Test path with whitespace
assert_eq!(table.get_path::<String>(" .a [\"b\"] .c ")?, "hello");
// Test indexing non-indexable value
let err = table.get_path::<String>("abc.c").unwrap_err().to_string();
assert_eq!(err, "runtime error: attempt to index a string value with key 'c'");
Ok(())
}
+103 -1
View File
@@ -394,6 +394,7 @@ fn test_error() -> Result<()> {
}
#[test]
#[cfg(not(panic = "abort"))]
fn test_panic() -> Result<()> {
fn make_lua(options: LuaOptions) -> Result<Lua> {
let lua = Lua::new_with(StdLib::ALL_SAFE, options)?;
@@ -437,7 +438,7 @@ fn test_panic() -> Result<()> {
{
let lua = make_lua(LuaOptions::default())?;
match catch_unwind(AssertUnwindSafe(|| -> Result<()> {
let _catched_panic = lua
let _caught_panic = lua
.load(
r#"
-- Set global
@@ -602,6 +603,21 @@ fn test_num_conversion() -> Result<()> {
assert_eq!(lua.unpack::<i128>(lua.pack(1i128 << 64)?)?, 1i128 << 64);
// Negative zero
let negative_zero = lua.load("-0.0").eval::<f64>()?;
assert_eq!(negative_zero, 0.0);
// LuaJIT treats -0.0 as a positive zero
#[cfg(not(feature = "luajit"))]
assert!(negative_zero.is_sign_negative());
// In Lua <5.3 all numbers are floats
#[cfg(not(any(feature = "lua54", feature = "lua53", feature = "luajit")))]
{
let negative_zero = lua.load("-0").eval::<f64>()?;
assert_eq!(negative_zero, 0.0);
assert!(negative_zero.is_sign_negative());
}
Ok(())
}
@@ -882,6 +898,7 @@ fn test_registry_value_reuse() -> Result<()> {
}
#[test]
#[cfg(not(panic = "abort"))]
fn test_application_data() -> Result<()> {
let lua = Lua::new();
@@ -1227,6 +1244,17 @@ fn test_register_module() -> Result<()> {
res.unwrap_err().to_string(),
"runtime error: module name must begin with '@'"
);
// Luau registered modules (aliases) are case-insensitive
let res = lua.register_module("@My_Module", &t);
assert!(res.is_ok());
lua.load(
r#"
local my_module = require("@MY_MODule")
assert(my_module.name == "my_module")
"#,
)
.exec()?;
}
Ok(())
@@ -1362,6 +1390,80 @@ fn test_inspect_stack() -> Result<()> {
Ok(())
}
#[test]
fn test_traceback() -> Result<()> {
let lua = Lua::new();
// Test traceback at level 0 (not inside any function)
let traceback = lua.traceback(None, 0)?.to_string_lossy();
assert!(traceback.contains("stack traceback:"));
// Test traceback with a message prefix
let traceback = lua.traceback(Some("error occurred"), 0)?.to_string_lossy();
assert!(traceback.starts_with("error occurred"));
assert!(traceback.contains("stack traceback:"));
// Test traceback inside a function
let get_traceback = lua.create_function(|lua, (msg, level): (Option<StdString>, usize)| {
lua.traceback(msg.as_deref(), level)
})?;
lua.globals().set("get_traceback", get_traceback)?;
lua.load(
r#"
local function foo()
-- Level 1 is inside foo (the caller)
local traceback = get_traceback(nil, 1)
return traceback
end
local function bar()
local result = foo()
return result
end
local function baz()
local result = bar()
return result
end
local traceback = baz()
assert(traceback:match("in %a+ 'foo'"))
assert(traceback:match("in %a+ 'bar'"))
assert(traceback:match("in %a+ 'baz'"))
"#,
)
.exec()?;
// Test traceback at different levels
lua.load(
r#"
local function foo()
local tb0 = get_traceback(nil, 0)
local tb1 = get_traceback(nil, 1)
local tb2 = get_traceback(nil, 2)
return tb0, tb1, tb2
end
local function bar()
local tb0, tb1, tb2 = foo()
return tb0, tb1, tb2
end
local tb0, tb1, tb2 = bar()
assert(tb0:match("in %a+ 'get_traceback'"))
assert(tb0:match("in %a+ 'foo'"))
assert(not tb1:match("in %a+ 'get_traceback'"))
assert(tb1:match("in %a+ 'foo'"))
assert(not tb2:match("in %a+ 'foo'"))
assert(tb1:match("in %a+ 'bar'"))
"#,
)
.exec()?;
Ok(())
}
#[test]
fn test_multi_states() -> Result<()> {
let lua = Lua::new();
+23 -1
View File
@@ -1,6 +1,6 @@
use std::panic::catch_unwind;
use mlua::{Error, Function, Lua, Result, Thread, ThreadStatus};
use mlua::{Error, Function, IntoLua, Lua, Result, Thread, ThreadStatus, Value};
#[test]
fn test_thread() -> Result<()> {
@@ -199,6 +199,7 @@ fn test_coroutine_from_closure() -> Result<()> {
}
#[test]
#[cfg(not(panic = "abort"))]
fn test_coroutine_panic() {
match catch_unwind(|| -> Result<()> {
// check that coroutines propagate panics correctly
@@ -252,3 +253,24 @@ fn test_thread_resume_error() -> Result<()> {
Ok(())
}
#[test]
fn test_thread_resume_bad_arg() -> Result<()> {
let lua = Lua::new();
struct BadArg;
impl IntoLua for BadArg {
fn into_lua(self, _lua: &Lua) -> Result<Value> {
Err(Error::runtime("bad arg"))
}
}
let f = lua.create_thread(lua.create_function(|_, ()| Ok("okay"))?)?;
let res = f.resume::<()>((123, BadArg));
assert!(matches!(res, Err(Error::RuntimeError(msg)) if msg == "bad arg"));
let res = f.resume::<String>(()).unwrap();
assert_eq!(res, "okay");
Ok(())
}
+14 -6
View File
@@ -31,7 +31,9 @@ fn test_boolean_type_metatable() -> Result<()> {
let mt = lua.create_table()?;
mt.set("__add", Function::wrap(|a, b| Ok(a || b)))?;
lua.set_type_metatable::<bool>(Some(mt));
assert_eq!(lua.type_metatable::<bool>(), None);
lua.set_type_metatable::<bool>(Some(mt.clone()));
assert_eq!(lua.type_metatable::<bool>().unwrap(), mt);
lua.load(r#"assert(true + true == true)"#).exec().unwrap();
lua.load(r#"assert(true + false == true)"#).exec().unwrap();
@@ -52,7 +54,8 @@ fn test_lightuserdata_type_metatable() -> Result<()> {
Ok(LightUserData((a.0 as usize + b.0 as usize) as *mut c_void))
}),
)?;
lua.set_type_metatable::<LightUserData>(Some(mt));
lua.set_type_metatable::<LightUserData>(Some(mt.clone()));
assert_eq!(lua.type_metatable::<LightUserData>().unwrap(), mt);
let res = lua
.load(
@@ -77,7 +80,9 @@ fn test_number_type_metatable() -> Result<()> {
let mt = lua.create_table()?;
mt.set("__call", Function::wrap(|n1: f64, n2: f64| Ok(n1 * n2)))?;
lua.set_type_metatable::<Number>(Some(mt));
lua.set_type_metatable::<Number>(Some(mt.clone()));
assert_eq!(lua.type_metatable::<Number>().unwrap(), mt);
lua.load(r#"assert((1.5)(3.0) == 4.5)"#).exec().unwrap();
lua.load(r#"assert((5)(5) == 25)"#).exec().unwrap();
@@ -93,7 +98,8 @@ fn test_string_type_metatable() -> Result<()> {
"__add",
Function::wrap(|a: String, b: String| Ok(format!("{a}{b}"))),
)?;
lua.set_type_metatable::<LuaString>(Some(mt));
lua.set_type_metatable::<LuaString>(Some(mt.clone()));
assert_eq!(lua.type_metatable::<LuaString>().unwrap(), mt);
lua.load(r#"assert(("foo" + "bar") == "foobar")"#).exec().unwrap();
@@ -109,7 +115,8 @@ fn test_function_type_metatable() -> Result<()> {
"__index",
Function::wrap(|_: Function, key: String| Ok(format!("function.{key}"))),
)?;
lua.set_type_metatable::<Function>(Some(mt));
lua.set_type_metatable::<Function>(Some(mt.clone()));
assert_eq!(lua.type_metatable::<Function>(), Some(mt));
lua.load(r#"assert((function() end).foo == "function.foo")"#)
.exec()
@@ -127,7 +134,8 @@ fn test_thread_type_metatable() -> Result<()> {
"__index",
Function::wrap(|_: Thread, key: String| Ok(format!("thread.{key}"))),
)?;
lua.set_type_metatable::<Thread>(Some(mt));
lua.set_type_metatable::<Thread>(Some(mt.clone()));
assert_eq!(lua.type_metatable::<Thread>(), Some(mt));
lua.load(r#"assert((coroutine.create(function() end)).foo == "thread.foo")"#)
.exec()
+100 -1
View File
@@ -8,7 +8,7 @@ use std::sync::atomic::{AtomicI64, Ordering};
use mlua::{
AnyUserData, Error, ExternalError, Function, Lua, MetaMethod, Nil, ObjectLike, Result, String, UserData,
UserDataFields, UserDataMethods, UserDataRef, Value, Variadic,
UserDataFields, UserDataMethods, UserDataRef, UserDataRegistry, Value, Variadic,
};
#[test]
@@ -428,6 +428,39 @@ fn test_userdata_destroy() -> Result<()> {
Ok(())
}
#[test]
fn test_userdata_method_once() -> Result<()> {
struct MyUserdata(Arc<i64>);
impl UserData for MyUserdata {
fn add_methods<M: UserDataMethods<Self>>(methods: &mut M) {
methods.add_method_once("take_value", |_, this, ()| Ok(*this.0));
}
}
let lua = Lua::new();
let rc = Arc::new(42);
let userdata = lua.create_userdata(MyUserdata(rc.clone()))?;
lua.globals().set("userdata", &userdata)?;
// Control userdata
let userdata2 = lua.create_userdata(MyUserdata(rc.clone()))?;
lua.globals().set("userdata2", userdata2)?;
assert_eq!(lua.load("userdata:take_value()").eval::<i64>()?, 42);
match lua.load("userdata2.take_value(userdata)").eval::<i64>() {
Err(Error::CallbackError { cause, .. }) => {
let err = cause.to_string();
assert!(err.contains("bad argument `self` to `MyUserdata.take_value`"));
assert!(err.contains("userdata has been destructed"));
}
r => panic!("expected Err(CallbackError), got {r:?}"),
}
assert_eq!(Arc::strong_count(&rc), 2);
Ok(())
}
#[test]
fn test_user_values() -> Result<()> {
struct MyUserData;
@@ -525,6 +558,11 @@ fn test_fields() -> Result<()> {
Ok(())
});
// Field that emulates method
fields.add_field_function_get("val_fget", |lua, ud| {
lua.create_function(move |_, ()| Ok(ud.borrow::<MyUserData>()?.0))
});
// Use userdata "uservalue" storage
fields.add_field_function_get("uval", |_, ud| ud.user_value::<Option<String>>());
fields.add_field_function_set("uval", |_, ud, s: Option<String>| ud.set_user_value(s));
@@ -537,6 +575,10 @@ fn test_fields() -> Result<()> {
})
})
}
fn add_methods<M: UserDataMethods<Self>>(methods: &mut M) {
methods.add_method("dummy", |_, _, ()| Ok(()));
}
}
globals.set("ud", MyUserData(7))?;
@@ -546,6 +588,7 @@ fn test_fields() -> Result<()> {
assert(ud.val == 7)
ud.val = 10
assert(ud.val == 10)
assert(ud:val_fget() == 10)
assert(ud.uval == nil)
ud.uval = "hello"
@@ -914,6 +957,7 @@ fn test_nested_userdata_gc() -> Result<()> {
}
#[cfg(feature = "userdata-wrappers")]
#[rustversion::stable]
#[test]
fn test_userdata_wrappers() -> Result<()> {
#[derive(Debug)]
@@ -1297,3 +1341,58 @@ fn test_userdata_wrappers() -> Result<()> {
Ok(())
}
#[cfg(feature = "luau")]
#[test]
fn test_userdata_namecall() -> Result<()> {
let lua = Lua::new();
struct MyUserData;
impl UserData for MyUserData {
fn register(registry: &mut mlua::UserDataRegistry<Self>) {
registry.add_method("method", |_, _, ()| Ok("method called"));
registry.add_field_method_get("field", |_, _| Ok("field value"));
registry.add_meta_method(MetaMethod::Index, |_, _, key: StdString| Ok(key));
registry.enable_namecall();
}
}
let ud = lua.create_userdata(MyUserData)?;
lua.globals().set("ud", &ud)?;
lua.load(
r#"
assert(ud:method() == "method called")
assert(ud.field == "field value")
assert(ud.dynamic_field == "dynamic_field")
local ok, err = pcall(function() return ud:dynamic_field() end)
assert(tostring(err):find("attempt to call an unknown method 'dynamic_field'") ~= nil)
"#,
)
.exec()?;
ud.destroy()?;
let err = lua.load("ud:method()").exec().unwrap_err();
assert!(err.to_string().contains("userdata has been destructed"));
Ok(())
}
#[test]
fn test_userdata_get_path() -> Result<()> {
let lua = Lua::new();
struct MyUd;
impl UserData for MyUd {
fn register(registry: &mut UserDataRegistry<Self>) {
registry.add_field("value", "userdata_value");
}
}
let ud = lua.create_userdata(MyUd)?;
assert_eq!(ud.get_path::<String>(".value")?, "userdata_value");
Ok(())
}
+1 -1
View File
@@ -1,5 +1,5 @@
[default]
extend-ignore-identifiers-re = ["catched", "2nd", "ser"]
extend-ignore-identifiers-re = ["2nd", "ser"]
[default.extend-words]
thr = "thr"