use core::ptr::NonNull;
use luau_common::flags;
use crate::Table;
use crate::call::{ProtectedCall, ThreadStack};
use crate::function::{Closure, Proto};
use crate::gc::{BLACK_BIT, GCS_ATOMIC, GCS_PROPAGATE, GcObject, WHITE_BITS, bit_mask};
use crate::handle::RawHandle;
use crate::metamethod::TmEvent;
use crate::state::{GlobalState, ThreadLifecycle, ThreadState};
use crate::string::TString;
use crate::thread::Thread;
use crate::types;
use crate::types::{
LUA_T_COUNT, LUA_TBUFFER, LUA_TCLASS, LUA_TFUNCTION, LUA_TOBJECT, LUA_TPROTO, LUA_TSTRING,
LUA_TTABLE, LUA_TTHREAD, LUA_TUPVALUE, LUA_TUSERDATA,
};
use crate::userdata::TypedUserdataAccess;
use crate::value::{RAW_TVALUE_NIL, TValue};
use crate::{Class, Object, VmError, VmErrorResult, VmExit, VmResult};
impl GlobalState {
pub(super) unsafe fn mark_value(&self, value: TValue) {
if value.is_collectable() {
let object = value.gc_value();
if unsafe { object.is_white() } {
unsafe { self.really_mark_object(object) };
}
}
}
pub(super) unsafe fn mark_object(&self, object: GcObject) {
if unsafe { object.is_white() } {
unsafe { self.really_mark_object(object) };
}
}
pub(super) unsafe fn mark_mt(&self) {
for tag in 0..LUA_T_COUNT {
if let Some(metatable) = self.metatable(tag) {
unsafe { self.mark_object(metatable.into()) };
}
}
}
pub(super) unsafe fn mark_tagged_userdata_metatables(&self) {
for tag in 0..crate::userdata::USERDATA_TAG_LIMIT {
if let Some(metatable) = self.userdata_metatable(tag) {
unsafe { self.mark_object(metatable.into()) };
}
}
}
pub(super) unsafe fn mark_userdata_type_metatables(&self) {
for metatable in unsafe { &*self.userdata_type_registry_ptr() }.rooted_metatables() {
unsafe { self.mark_object(metatable.into()) };
}
}
pub(super) unsafe fn mark_userdata_direct_access(&self) {
unsafe {
for direct_access in &self.as_ptr().as_ref().unwrap_unchecked().userdata_direct {
self.mark_value(TValue::from_ref(&direct_access.index_tm));
self.mark_value(TValue::from_ref(&direct_access.new_index_tm));
self.mark_value(TValue::from_ref(&direct_access.name_call_tm));
}
}
}
pub(super) unsafe fn mark_userdata_direct_fields(&self) {
unsafe {
for direct_fields in &self
.as_ptr()
.as_ref()
.unwrap_unchecked()
.userdata_direct_fields
{
if let Some(direct_fields) = NonNull::new(*direct_fields) {
self.mark_object(Table::from_raw(direct_fields).into());
}
}
}
}
pub(super) unsafe fn really_mark_object(&self, mut object_ref: GcObject) {
unsafe {
debug_assert!(object_ref.is_white());
debug_assert!(!self.is_dead(object_ref));
object_ref.white_to_gray();
match object_ref.as_ptr().as_ref().unwrap_unchecked().tt as i32 {
x if x == LUA_TSTRING => object_ref.gray_to_black(),
x if x == LUA_TUSERDATA => {
object_ref.gray_to_black();
let userdata = object_ref.to_userdata();
let tag = userdata.as_ptr().as_ref().unwrap_unchecked().tag as usize;
let thread = self.main_thread();
if flags::LuauGcTraceUdata.get()
&& tag < crate::userdata::USERDATA_TAG_LIMIT
&& let Some(mark) = self.userdata_mark(tag)
{
mark(&thread, userdata.data_mut_ptr().cast());
}
if let Some(metatable) = userdata.metatable() {
self.mark_object(metatable.into());
}
if let Some(userdata) = thread.typed_userdata(userdata) {
self.mark_value(thread.typed_userdata_value(userdata));
}
}
x if x == LUA_TUPVALUE => {
let upvalue = object_ref.to_upvalue();
self.mark_value(upvalue.value());
if !upvalue.is_open() {
object_ref.gray_to_black();
}
}
x if x == LUA_TFUNCTION => {
object_ref.to_closure().set_gc_list(self.gray());
self.set_gray(Some(object_ref));
}
x if x == LUA_TTABLE => {
object_ref.to_table().set_gc_list(self.gray());
self.set_gray(Some(object_ref));
}
x if x == LUA_TTHREAD => {
object_ref.to_state().set_gc_list(self.gray());
self.set_gray(Some(object_ref));
}
x if x == LUA_TPROTO => {
object_ref.to_proto().set_gc_list(self.gray());
self.set_gray(Some(object_ref));
}
x if x == LUA_TCLASS => {
object_ref.to_class().set_gc_list(self.gray());
self.set_gray(Some(object_ref));
}
x if x == LUA_TOBJECT => {
object_ref.to_object().set_gc_list(self.gray());
self.set_gray(Some(object_ref));
}
x if x == LUA_TBUFFER => {
object_ref.gray_to_black();
}
other => unreachable!("unsupported collectable type {}", other),
}
}
}
}
impl Thread {
pub(super) unsafe fn mark_root(&self) {
unsafe {
let global = self.global();
let main_thread = global.main_thread();
global.set_gray(None);
global.set_gray_again(None);
global.set_weak(None);
global.mark_object((&main_thread).into());
if !main_thread
.as_ptr()
.as_ref()
.unwrap_unchecked()
.gt
.is_null()
{
global.mark_object(main_thread.globals().into());
}
global.mark_value(TValue::from_ref(
&global.as_ptr().as_ref().unwrap_unchecked().registry,
));
if flags::LuauGcTraceUdata.get() {
global.mark_value(global.weak_registry());
if let Some(embedder_gc) = global.embedder_gc() {
embedder_gc(&main_thread, None);
}
}
if flags::LuauUdataDirectAccess6.get() {
global.mark_userdata_direct_access();
}
if flags::LuauDirectFieldGet.get() {
global.mark_userdata_direct_fields();
}
global.mark_mt();
global.mark_userdata_type_metatables();
if flags::LuauUdataMetatablePinned.get() {
global.mark_tagged_userdata_metatables();
}
global.set_gc_state(GCS_PROPAGATE);
}
}
}
impl TString {
pub(super) unsafe fn mark_gc(&self) {
unsafe {
let marked = &mut self.as_ptr().as_mut().unwrap_unchecked().marked;
*marked = (*marked & !WHITE_BITS) | bit_mask(BLACK_BIT);
}
}
}
impl Thread {
pub(super) unsafe fn clear_stack(&self) {
unsafe {
let stack_end = self
.stack()
.add(self.as_ptr().as_ref().unwrap_unchecked().stack_size as usize);
let mut slot = self.stack_top();
while slot < stack_end {
slot.value_unchecked().set_nil();
slot = slot.add(1);
}
}
}
pub(super) unsafe fn shrink_stack(&self) -> VmErrorResult {
unsafe {
let thread_state = self.as_ptr().as_ref().unwrap_unchecked();
let stack_start = self.restore_stack(0);
let stack_last = self.stack_last();
let mut limit = self.stack_top();
let mut call_info_cursor = self.base_call_info_cursor();
let current_call_info_cursor = self.current_call_info_cursor();
while call_info_cursor <= current_call_info_cursor {
let call_info_top = call_info_cursor.call_info_unchecked().top();
debug_assert!(call_info_top <= stack_last);
if limit < call_info_top {
limit = call_info_top;
}
call_info_cursor = call_info_cursor.add(1);
}
let call_info_used =
current_call_info_cursor.offset_from(self.base_call_info_cursor()) as i32;
let stack_used = limit.offset_from(stack_start) as i32;
let size_ci = thread_state.size_ci as usize;
let stack_size = thread_state.stack_size as usize;
if size_ci > crate::thread::LUAI_MAX_CALLS {
return Ok(());
}
if 3 * (call_info_used as usize) < size_ci && 2 * crate::state::BASIC_CI_SIZE < size_ci
{
self.realloc_ci(thread_state.size_ci / 2)?;
}
if 3 * (stack_used as usize) < stack_size
&& 2 * (crate::state::BASIC_STACK_SIZE + crate::state::EXTRA_STACK) < stack_size
{
self.realloc_stack(thread_state.stack_size / 2, false)?;
}
}
Ok(())
}
pub(super) unsafe fn shrink_stack_protected(&self) {
unsafe fn shrink_stack_trampoline(thread: &Thread, _: &mut ()) -> VmResult {
unsafe { thread.shrink_stack() }?;
Ok(())
}
let mut unit = ();
let result = unsafe { self.raw_run_protected(shrink_stack_trampoline, &mut unit) };
debug_assert!(matches!(
result,
Ok(()) | Err(VmExit::Error(VmError::Memory))
));
}
}
impl GlobalState {
unsafe fn traverse_table(&self, table: Table) -> bool {
unsafe {
let mut weak_key = false;
let mut weak_value = false;
let table_ref = table.as_ptr().as_ref().unwrap_unchecked();
if let Some(metatable) = NonNull::new(table_ref.metatable) {
self.mark_object(Table::from_raw(metatable).into());
let event_name = self.tm_name(TmEvent::Mode as usize);
if let Some(mode) = Table::from_raw(metatable).get_tm(TmEvent::Mode, event_name)
&& mode.is_string()
{
let mode_string = mode.string_value();
let mode_bytes = mode_string.as_bytes();
weak_key = mode_bytes.contains(&b'k');
weak_value = mode_bytes.contains(&b'v');
if weak_key || weak_value {
table.set_gc_list(self.weak());
self.set_weak(Some(table.into()));
}
}
}
if weak_key && weak_value {
return true;
}
if !weak_value {
for index in (0..table_ref.size_array as usize).rev() {
self.mark_value(table.array_slot(index));
}
}
for index in (0..table.node_count()).rev() {
let node = table.node(index as i32);
debug_assert!(
node.key().tt() != types::LUA_TDEADKEY || node.value_unchecked().is_nil()
);
if node.value_unchecked().is_nil() {
if node.key().is_collectable() {
node.key().set_tt(types::LUA_TDEADKEY);
}
} else {
debug_assert!(!node.key().is_nil());
if !weak_key {
let mut key = RAW_TVALUE_NIL;
node.write_key_to_value(TValue::from_mut(&mut key));
self.mark_value(TValue::from_ref(&key));
}
if !weak_value {
self.mark_value(node.value_unchecked());
}
}
}
weak_key || weak_value
}
}
unsafe fn traverse_proto(&self, proto: Proto) {
unsafe {
let proto_ref = proto.as_ptr().as_ref().unwrap_unchecked();
if let Some(source) = proto.source() {
source.mark_gc();
}
if let Some(debug_name) = proto.debug_name() {
debug_name.mark_gc();
}
for index in 0..proto_ref.size_k as usize {
self.mark_value(proto.constant(index));
}
for index in 0..proto_ref.size_upvalues as usize {
if let Some(upvalue) = proto.upvalue_name(index) {
upvalue.mark_gc();
}
}
for index in 0..proto_ref.size_p as usize {
self.mark_object(proto.child_proto(index).unwrap_unchecked().into());
}
for index in 0..proto_ref.size_loc_vars as usize {
if let Some(local) = proto.loc_var(index)
&& let Some(var_name) = local.name()
{
var_name.mark_gc();
}
}
}
}
unsafe fn traverse_closure(&self, closure: Closure) {
unsafe {
let closure_ref = closure.as_ptr().as_ref().unwrap_unchecked();
if let Some(environment) = NonNull::new(closure_ref.env) {
self.mark_object(Table::from_raw(environment).into());
}
if closure.is_native() {
if flags::LuauManagedDebugNames.get()
&& let Some(debug_name) = closure.native_debug_name()
&& let crate::string::LuaStringRepr::Interned(debug_name) = debug_name.0
{
debug_name.mark_gc();
}
for index in 0..closure_ref.n_upvalues as usize {
self.mark_value(closure.native_upvalue(index));
}
} else {
if let Some(proto) = closure.proto() {
debug_assert_eq!(
closure_ref.n_upvalues as i32,
proto.as_ptr().as_ref().unwrap_unchecked().n_ups as i32
);
self.mark_object(proto.into());
}
for index in 0..closure_ref.n_upvalues as usize {
self.mark_value(closure.lua_upvalue_ref(index));
}
}
}
}
unsafe fn traverse_class(&self, class_object: Class) {
unsafe {
let class_ref = class_object.as_ptr().as_ref().unwrap_unchecked();
self.mark_object(class_object.name().into());
self.mark_object(class_object.members_to_offset().into());
for index in 0..class_ref.number_of_all_members as usize {
self.mark_object(class_object.offset_to_member(index).into());
}
let static_member_count =
(class_ref.number_of_all_members - class_ref.number_of_instance_members) as usize;
for index in 0..static_member_count {
self.mark_value(class_object.static_member(index));
}
if let Some(metatable) = class_object.metatable() {
self.mark_object(metatable.into());
}
if let Some(instance_metatable) = class_object.instance_metatable() {
self.mark_object(instance_metatable.into());
}
}
}
unsafe fn traverse_object(&self, object_instance: Object) {
unsafe {
let object_ref = object_instance.as_ptr().as_ref().unwrap_unchecked();
self.mark_object(object_instance.class().into());
for index in 0..object_ref.number_of_members as usize {
self.mark_value(object_instance.member(index));
}
}
}
unsafe fn traverse_stack(&self, thread: &Thread) {
unsafe {
let thread_ref = thread.as_ptr().as_ref().unwrap_unchecked();
let stack_top = thread.stack_top();
let mut slot = thread.restore_stack(0);
if !thread_ref.gt.is_null() {
self.mark_object(thread.globals().into());
}
if let Some(name_call) = thread.name_call() {
name_call.mark_gc();
}
while slot < stack_top {
self.mark_value(slot.value_unchecked());
slot = slot.add(1);
}
let mut upvalue = thread.open_upvalue();
while let Some(current_upvalue) = upvalue {
debug_assert!(current_upvalue.is_open());
current_upvalue
.as_ptr()
.as_mut()
.unwrap_unchecked()
.marked_open = 1;
self.mark_object(current_upvalue.into());
upvalue = current_upvalue.open_data().thread_next();
}
}
}
pub(super) unsafe fn propagate_mark(&self) -> usize {
unsafe {
let mut object = self.gray().unwrap_unchecked();
debug_assert!(object.is_gray());
object.gray_to_black();
match object.as_ptr().as_ref().unwrap_unchecked().tt as i32 {
x if x == types::LUA_TTABLE => {
let table = object.to_table();
self.set_gray(table.gc_list());
if self.traverse_table(table) {
object.black_to_gray();
}
table.gc_work_size(!flags::LuauGcTableStepFix.get())
}
x if x == types::LUA_TFUNCTION => {
let closure = object.to_closure();
self.set_gray(closure.gc_list());
self.traverse_closure(closure);
closure.size()
}
x if x == types::LUA_TTHREAD => {
let thread = object.to_state();
self.set_gray(thread.gc_list());
let active = thread.as_ptr().as_ref().unwrap_unchecked().is_active
|| thread == self.main_thread();
self.traverse_stack(&thread);
if active {
thread.set_gc_list(self.gray_again());
self.set_gray_again(Some(object));
object.black_to_gray();
}
if !active || self.gc_state() == GCS_ATOMIC {
thread.clear_stack();
}
if self.gc_state() == GCS_PROPAGATE {
thread.shrink_stack_protected();
}
thread.allocation_size()
}
x if x == types::LUA_TPROTO => {
let proto = object.to_proto();
self.set_gray(proto.gc_list());
self.traverse_proto(proto);
proto.size()
}
x if x == types::LUA_TCLASS => {
let class_object = object.to_class();
self.set_gray(class_object.gc_list());
self.traverse_class(class_object);
class_object.allocation_size()
}
x if x == types::LUA_TOBJECT => {
let object_instance = object.to_object();
self.set_gray(object_instance.gc_list());
self.traverse_object(object_instance);
object_instance.allocation_size()
}
other => unreachable!("unexpected object type in gc propagation: {}", other),
}
}
}
pub(super) unsafe fn propagate_all(&self) -> usize {
let mut work = 0;
while self.gray().is_some() {
work += unsafe { self.propagate_mark() };
}
work
}
}