use std::marker::PhantomData;
use luau_vm::Thread as VmThread;
use luau_vm::internal::userdata::{UserdataTypeRegistration, UserdataTypeRegistryAccess};
use luau_vm::native::{NativeCallContext, NativeCallResult};
use luau_vm::thread::{LUA_REFNIL, StackGuard, upvalue_index};
use luau_vm::types::{LUA_TFUNCTION, LUA_TNIL, LUA_TTABLE};
use crate::callback::Callback;
use crate::error::Error;
use crate::lua::runtime::RuntimeData;
use crate::lua::{LuaRef, RegistryKey};
use crate::thread::Thread;
use crate::value::{IntoLua, Value};
use super::MetaMethod;
use super::callback::push_userdata_callback;
pub struct UserdataRegistry<'lua, T> {
pub(super) state: Option<UserdataRegistryState<'lua>>,
pub(super) _marker: PhantomData<fn() -> T>,
}
pub(super) struct UserdataRegistryState<'lua> {
thread: &'lua VmThread,
runtime: &'lua RuntimeData,
definition: UserdataDefinition,
next_value: i32,
values: Option<RegistryKey>,
result: Result<(), Error>,
}
pub(super) struct UserdataDefinition {
members: Vec<RegistrationEntry>,
getters: Vec<RegistrationEntry>,
setters: Vec<RegistrationEntry>,
metatable: Vec<RegistrationEntry>,
index_overridden: bool,
new_index_overridden: bool,
values: Option<RegistryKey>,
}
struct RegistrationEntry {
name: Vec<u8>,
value: RegistrationValue,
}
enum RegistrationValue {
Stored(i32),
Callback(Box<dyn Callback>),
}
impl<'lua, T> UserdataRegistry<'lua, T> {
pub(super) fn new(thread: &'lua VmThread, runtime: &'lua RuntimeData) -> Self {
Self {
state: Some(UserdataRegistryState {
thread,
runtime,
definition: UserdataDefinition {
members: Vec::new(),
getters: Vec::new(),
setters: Vec::new(),
metatable: Vec::new(),
index_overridden: false,
new_index_overridden: false,
values: None,
},
next_value: 1,
values: None,
result: Ok(()),
}),
_marker: PhantomData,
}
}
pub(super) fn finish(mut self) -> Result<UserdataDefinition, Error> {
let mut state = self
.state
.take()
.expect("userdata registry state must be present");
match core::mem::replace(&mut state.result, Ok(())) {
Ok(()) => {
state.definition.values = state.values.take();
Ok(state.definition)
}
Err(error) => {
state.release_values();
Err(error)
}
}
}
pub(super) fn register_proxy_target<R>(&mut self)
where
R: super::Userdata + 'static,
{
let state = self
.state
.take()
.expect("userdata registry state must be present");
let mut target = UserdataRegistry::<R> {
state: Some(state),
_marker: PhantomData,
};
R::register(&mut target);
self.state = target.state.take();
}
pub(super) fn set_member<V>(&mut self, name: impl AsRef<[u8]>, value: V)
where
V: for<'value> IntoLua<'value>,
{
self.store_value(RegistrationTarget::Members, name, value);
}
pub(super) fn set_metatable_value<V>(&mut self, name: impl AsRef<[u8]>, value: V)
where
V: for<'value> IntoLua<'value>,
{
self.note_meta_field(name.as_ref());
self.store_value(RegistrationTarget::Metatable, name, value);
}
fn store_value<V>(&mut self, target: RegistrationTarget, name: impl AsRef<[u8]>, value: V)
where
V: for<'value> IntoLua<'value>,
{
let state = self.state_mut();
if state.result.is_err() {
return;
}
match state.store_pending_value(|thread| unsafe { value.push_into_stack(thread) }) {
Ok(slot) => target
.entries(&mut state.definition)
.push(RegistrationEntry {
name: name.as_ref().to_vec(),
value: RegistrationValue::Stored(slot),
}),
Err(error) => state.record(Err(error)),
}
}
pub(super) fn set_member_callback(
&mut self,
name: impl AsRef<[u8]>,
callback: Box<dyn Callback>,
) {
self.store_callback(RegistrationTarget::Members, name, callback);
}
pub(super) fn set_member_value_with<F>(&mut self, name: impl AsRef<[u8]>, field: F)
where
F: for<'value> FnOnce(LuaRef<'value>) -> Result<Value<'value>, Error> + 'static,
{
self.store_value_with(RegistrationTarget::Members, name, field);
}
pub(super) fn set_metatable_callback(
&mut self,
name: impl AsRef<[u8]>,
callback: Box<dyn Callback>,
) {
self.note_meta_field(name.as_ref());
self.store_callback(RegistrationTarget::Metatable, name, callback);
}
fn store_callback(
&mut self,
target: RegistrationTarget,
name: impl AsRef<[u8]>,
callback: Box<dyn Callback>,
) {
let state = self.state_mut();
if state.result.is_err() {
return;
}
target
.entries(&mut state.definition)
.push(RegistrationEntry {
name: name.as_ref().to_vec(),
value: RegistrationValue::Callback(callback),
});
}
pub(super) fn set_getter_callback(
&mut self,
name: impl AsRef<[u8]>,
callback: Box<dyn Callback>,
) {
self.store_callback(RegistrationTarget::Getters, name, callback);
}
pub(super) fn set_setter_callback(
&mut self,
name: impl AsRef<[u8]>,
callback: Box<dyn Callback>,
) {
self.store_callback(RegistrationTarget::Setters, name, callback);
}
pub(super) fn set_metatable_value_with<F>(&mut self, name: impl AsRef<[u8]>, field: F)
where
F: for<'value> FnOnce(LuaRef<'value>) -> Result<Value<'value>, Error> + 'static,
{
self.note_meta_field(name.as_ref());
self.store_value_with(RegistrationTarget::Metatable, name, field);
}
fn store_value_with<F>(&mut self, target: RegistrationTarget, name: impl AsRef<[u8]>, field: F)
where
F: for<'value> FnOnce(LuaRef<'value>) -> Result<Value<'value>, Error> + 'static,
{
let state = self.state_mut();
if state.result.is_err() {
return;
}
match field(LuaRef::new(state.thread, state.runtime)) {
Ok(value) => {
match state.store_pending_value(|thread| unsafe { value.push_into_stack(thread) }) {
Ok(slot) => target
.entries(&mut state.definition)
.push(RegistrationEntry {
name: name.as_ref().to_vec(),
value: RegistrationValue::Stored(slot),
}),
Err(error) => state.record(Err(error)),
}
}
Err(error) => state.record(Err(error)),
}
}
fn note_meta_field(&mut self, name: &[u8]) {
let state = self.state_mut();
if let Err(error) = MetaMethod::validate(name) {
state.record(Err(error));
return;
}
if name == MetaMethod::Index.name().as_bytes() {
state.definition.index_overridden = true;
} else if name == MetaMethod::NewIndex.name().as_bytes() {
state.definition.new_index_overridden = true;
}
}
fn state_mut(&mut self) -> &mut UserdataRegistryState<'lua> {
self.state
.as_mut()
.expect("userdata registry state must be present")
}
}
impl UserdataRegistryState<'_> {
fn store_pending_value(
&mut self,
push: impl FnOnce(&Thread<'_>) -> Result<(), Error>,
) -> Result<i32, Error> {
let slot = self.next_value;
let next_value = slot
.checked_add(1)
.ok_or_else(|| Error::runtime("too many registered userdata values"))?;
let safe_thread = Thread::new(self.thread, self.runtime);
unsafe {
let _stack = StackGuard::new(self.thread);
match &self.values {
Some(values) => safe_thread.push_registry_value(values)?,
None => self
.thread
.new_table()
.map_err(|error| Error::from_thread_exit(self.thread, error))?,
}
let values = self.thread.abs_index(-1);
push(&safe_thread)?;
self.thread
.raw_seti(values, slot)
.map_err(|error| Error::from_thread_exit(self.thread, error))?;
if self.values.is_none() {
self.values = Some(safe_thread.create_registry_key_from_top()?);
}
self.next_value = next_value;
Ok(slot)
}
}
fn release_values(&mut self) {
if let Some(values) = self.values.take() {
let reference = values.take();
if reference > LUA_REFNIL {
unsafe { self.thread.unref_value(reference) };
}
}
}
fn record(&mut self, result: Result<(), Error>) {
if self.result.is_ok() {
self.result = result;
}
}
}
enum RegistrationTarget {
Members,
Getters,
Setters,
Metatable,
}
impl RegistrationTarget {
fn entries<'a>(
&self,
definition: &'a mut UserdataDefinition,
) -> &'a mut Vec<RegistrationEntry> {
match self {
Self::Members => &mut definition.members,
Self::Getters => &mut definition.getters,
Self::Setters => &mut definition.setters,
Self::Metatable => &mut definition.metatable,
}
}
}
impl UserdataDefinition {
pub(super) fn release(mut self, thread: &VmThread) {
if let Some(values) = self.values.take() {
let reference = values.take();
if reference > LUA_REFNIL {
unsafe { thread.unref_value(reference) };
}
}
}
pub(super) fn materialize<T: 'static>(
mut self,
thread: &Thread<'_>,
name: String,
) -> Result<UserdataTypeRegistration, Error> {
unsafe {
let vm_thread = thread.as_vm();
let _stack = StackGuard::new(vm_thread);
let values = match self.values.take() {
Some(values) => {
thread.push_registry_value(&values)?;
let index = Some(vm_thread.abs_index(-1));
let reference = values.take();
if reference > LUA_REFNIL {
vm_thread.unref_value(reference);
}
index
}
None => None,
};
vm_thread
.new_table()
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
let metatable = vm_thread.abs_index(-1);
vm_thread
.push_string(&name)
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
vm_thread
.raw_set_field(metatable, MetaMethod::Type)
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
vm_thread
.new_table()
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
let members = vm_thread.abs_index(-1);
vm_thread
.new_table()
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
let getters = vm_thread.abs_index(-1);
vm_thread
.new_table()
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
let setters = vm_thread.abs_index(-1);
let has_getters = !self.getters.is_empty();
let has_setters = !self.setters.is_empty();
populate_registration_table(vm_thread, members, self.members, values)?;
populate_registration_table(vm_thread, getters, self.getters, values)?;
populate_registration_table(vm_thread, setters, self.setters, values)?;
populate_registration_table(vm_thread, metatable, self.metatable, values)?;
if self.index_overridden || has_getters {
vm_thread
.push_value(members)
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
vm_thread
.push_value(getters)
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
if self.index_overridden {
vm_thread
.raw_get_field(metatable, MetaMethod::Index)
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
} else {
vm_thread
.push_nil()
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
}
vm_thread
.push_native_closure(userdata_index, Some("__index"), 3)
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
} else {
vm_thread
.push_value(members)
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
}
vm_thread
.raw_set_field(metatable, MetaMethod::Index)
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
if has_setters {
vm_thread
.push_value(setters)
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
if self.new_index_overridden {
vm_thread
.raw_get_field(metatable, MetaMethod::NewIndex)
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
} else {
vm_thread
.push_nil()
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
}
vm_thread
.push_native_closure(userdata_new_index, Some("__newindex"), 2)
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
vm_thread
.raw_set_field(metatable, MetaMethod::NewIndex)
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
}
vm_thread
.push_boolean(0)
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
vm_thread
.raw_set_field(metatable, "__metatable")
.map_err(|error| Error::from_thread_exit(vm_thread, error))?;
vm_thread
.register_userdata_type::<T>(metatable)
.map_err(|error| Error::from_thread_exit(vm_thread, error))
}
}
}
fn populate_registration_table(
thread: &VmThread,
table: i32,
entries: Vec<RegistrationEntry>,
values: Option<i32>,
) -> Result<(), Error> {
for entry in entries {
match entry.value {
RegistrationValue::Stored(slot) => {
let values = values.expect("stored userdata registration value table must exist");
unsafe {
thread
.raw_geti(values, slot)
.map_err(|error| Error::from_thread_exit(thread, error))?;
}
}
RegistrationValue::Callback(callback) => push_userdata_callback(thread, callback)?,
}
unsafe {
thread
.raw_set_field(table, entry.name)
.map_err(|error| Error::from_thread_exit(thread, error))?;
}
}
Ok(())
}
fn userdata_index(context: NativeCallContext<'_>) -> NativeCallResult {
unsafe {
let thread = context.raw_thread();
thread.push_value(2)?;
if thread.raw_get(upvalue_index(1)) != LUA_TNIL {
return Ok(1);
}
thread.pop(1);
thread.push_value(2)?;
if thread.raw_get(upvalue_index(2)) == LUA_TFUNCTION {
thread.push_value(1)?;
thread.call(1, 1)?;
return Ok(1);
}
thread.pop(1);
match thread.type_of(upvalue_index(3)) {
LUA_TNIL => {
thread.push_nil()?;
Ok(1)
}
LUA_TFUNCTION => {
thread.push_value(upvalue_index(3))?;
thread.push_value(1)?;
thread.push_value(2)?;
thread.call(2, 1)?;
Ok(1)
}
LUA_TTABLE => {
thread.push_value(2)?;
thread.get_table(upvalue_index(3))?;
Ok(1)
}
_ => {
thread.push_nil()?;
Ok(1)
}
}
}
}
fn userdata_new_index(context: NativeCallContext<'_>) -> NativeCallResult {
unsafe {
let thread = context.raw_thread();
thread.push_value(2)?;
if thread.raw_get(upvalue_index(1)) == LUA_TFUNCTION {
thread.push_value(1)?;
thread.push_value(3)?;
thread.call(2, 0)?;
return Ok(0);
}
thread.pop(1);
match thread.type_of(upvalue_index(2)) {
LUA_TNIL => {
luau_vm::error!(thread, "attempt to set unknown userdata field").map_err(Into::into)
}
LUA_TFUNCTION => {
thread.push_value(upvalue_index(2))?;
thread.push_value(1)?;
thread.push_value(2)?;
thread.push_value(3)?;
thread.call(3, 0)?;
Ok(0)
}
LUA_TTABLE => {
thread.push_value(2)?;
thread.push_value(3)?;
thread.set_table(upvalue_index(2))?;
Ok(0)
}
_ => {
luau_vm::error!(thread, "attempt to set unknown userdata field").map_err(Into::into)
}
}
}
}