diff --git a/src/bind.luau b/src/bind.luau index 7a08292..5792049 100644 --- a/src/bind.luau +++ b/src/bind.luau @@ -5,7 +5,10 @@ local throw = require(script.Parent.throw) local flags = require(script.Parent.flags) local graph = require(script.Parent.graph) type Node = graph.Node -local set_effect = graph.set_effect +local create = graph.create +local create_and_open_scope = graph.create_and_open_scope +local close_scope = graph.close_scope +local set_child = graph.set_child local capture = graph.capture --[[ @@ -71,23 +74,37 @@ function bind(instance: Instance, property: string, setter: (Instance) -> ()) end end + local node = create(false) + + create_and_open_scope(node) + -- run setter to capture any nodes being depended on local nodes = (capture(setter :: () -> unknown, instance)) - -- register the setter as a side-effect of each node - for _, node in next, nodes do - set_effect(node, setter, instance) - end + close_scope() -- get binding id bind_count += 1 local bind_id = bind_count + + node.effect = function() + local instance = weak[bind_id] + if instance == nil then return end + setter(weak[bind_id] :: Instance) + end + + -- register the setter as a side-effect of each node + for _, n in next, nodes do + set_child(n, node) + end + + -- store reference of instance proxy without preventing gc weak[bind_id] = instance local function ref() - local _ = setter -- prevent gc of nodes being depended on + local _ = node -- prevent gc of node being depended on local instance = weak[bind_id] :: Instance -- keep proxy in memory if instance is still parented diff --git a/src/cleanup.luau b/src/cleanup.luau index d3da68a..26be562 100644 --- a/src/cleanup.luau +++ b/src/cleanup.luau @@ -2,6 +2,17 @@ if not game then script = require "test/relative-string" end local flags = require(script.Parent.flags) local throw = require(script.Parent.throw) +local graph = require(script.Parent.graph) +local get_scope = graph.get_scope +local add_cleanup = graph.add_cleanup + +local function cleanup(fn: () -> ()) + local node = get_scope() + if node == nil then throw("cannot call cleanup() in a non-reactive scope") end + add_cleanup(node, fn) +end + +return cleanup --[[ @@ -22,6 +33,7 @@ todo: remove need for ref to id maps? ]] +--[[ -- maps a ref to cleanup id local ref_to_id = {} :: { [string]: number } -- maps a cleanup id to a ref @@ -127,3 +139,4 @@ local manual_cleanup_mode = function(caller: () -> ()?) end :: ( (caller: (...any) -> ()) -> () ) & ( (nil) -> { () -> () } ) return function() return cleanup, clean_garbage, manual_cleanup_mode, cleanup_ref end +]] diff --git a/src/derive.luau b/src/derive.luau index 4d08437..7515747 100644 --- a/src/derive.luau +++ b/src/derive.luau @@ -3,12 +3,18 @@ if not game then script = require "test/relative-string" end local graph = require(script.Parent.graph) local create = graph.create local capture_and_link = graph.capture_and_link +local create_and_open_scope = graph.create_and_open_scope +local close_scope = graph.close_scope local function derive(fn: () -> T): () -> T local node, read_node_value = create((false :: any) :: T) + create_and_open_scope(node) + node.cache = capture_and_link(node, fn) + close_scope() + return read_node_value end diff --git a/src/graph.luau b/src/graph.luau index 029067c..ffa9e2f 100644 --- a/src/graph.luau +++ b/src/graph.luau @@ -2,12 +2,13 @@ if not game then script = require "test/relative-string" end local throw = require(script.Parent.throw) local flags = require(script.Parent.flags) +local on_gc = require(script.Parent.on_gc)() export type Node = { cache: T, - derive: () -> T, - effects: { [(unknown) -> ()]: unknown }, -- weak values - children: { Node } | false -- weak values + effect: () -> (), + children: { Node } | false, -- weak values + cleanups: { () -> () } | false } -- flag used to detect when node reference capturing is active @@ -15,6 +16,8 @@ local reff = false -- array of all nodes referenced since above flag was set local refs = {} :: { Node } +local scopes = { n = 0 } :: { [number]: Node, n: number } + local WEAK_VALUES = { __mode = "v" } local EVALUATION_ERR = "error while evaluating source:\n\n" @@ -46,6 +49,39 @@ local check_for_yield: (fn: (T...) -> unknown, T...) -> () do end end +local function get_scope(): Node + return scopes[scopes.n] +end + +local function open_scope(node: Node) + local n = scopes.n + 1 + scopes.n = n + scopes[n] = node +end + +local function close_scope() + local n = scopes.n + scopes.n = n - 1 + scopes[n] = nil +end + +local function add_cleanup(node: Node, cleanup: () -> ()) + if node.cleanups then + table.insert(node.cleanups, cleanup) + else + node.cleanups = { cleanup } + end +end + +local function run_cleanups(node: { cleanups: { () -> () } | false}) + if node.cleanups then + for _, fn in next, node.cleanups do + fn() + end + table.clear(node.cleanups) + end +end + --[[ Each node side-effect is registered with a corresponding weak key. @@ -57,21 +93,12 @@ The weak key is passed as an argument to its side-effect callback. ]] -local function set_effect(node: Node, fn: (T) -> (), key: T) - node.effects[fn :: () -> ()] = key +local function set_effect(node: Node, fn: () -> ()) + node.effect = fn end -local function run_effects(node: Node) - if flags.strict then -- run effects twice if strict - for effect, key in next, node.effects do - effect(key) - effect(key) - end - else - for effect, key in next, node.effects do - effect(key) - end - end +local function run_effect(node: Node) + node.effect() end -- retrieves a node's cached value @@ -91,15 +118,31 @@ local function set_child(parent: Node, child: Node) end end +local function create_and_open_scope(node: Node) + local parent = scopes[scopes.n] + if parent then + set_child(parent, node) + node.effect = function() + return parent + end + else + node.cleanups = {} + local cleanups = node.cleanups :: { () -> () } + on_gc(node, function() + run_cleanups({ cleanups = cleanups }) + end) + end + open_scope(node) +end + -- runs node effects, recalculates descendants and runs descendant effects local function update(node: Node) - run_effects(node) + open_scope(node) + run_cleanups(node) + run_effect(node) + close_scope() if node.children then - local strict = flags.strict - for _, child in node.children do - if strict then check_for_yield(child.derive) end - child.cache = child.derive() update(child) end end @@ -113,7 +156,9 @@ end -- links two nodes as parent-child with a function to compute a new value for child local function link(parent: Node, child: Node, derive: () -> T) - child.derive = derive + child.effect = function() + child.cache = derive() + end set_child(parent, child) end @@ -121,8 +166,6 @@ end local function capture(fn: (U?) -> T, arg: U?): ({ Node }, T) if reff then throw("recursive capture detected") end - if flags.strict then check_for_yield(fn, arg) end - table.clear(refs) reff = true @@ -142,10 +185,12 @@ local function capture(fn: (U?) -> T, arg: U?): ({ Node }, T) end -- captures and links any detected nodes -local function capture_and_link(child: Node, fn: () -> T): T - local nodes, value = capture(fn, nil) +local function capture_and_link(child: Node, derive: () -> T): T + local nodes, value = capture(derive, nil) - child.derive = fn + child.effect = function() + child.cache = derive() + end for _, parent: Node in next, nodes do set_child(parent, child) end @@ -156,9 +201,9 @@ end local function create(value: T): (Node, () -> T) local node = { cache = value, - derive = function() return nil :: any end, - effects = setmetatable({}, WEAK_VALUES) :: any, - children = false :: false + effect = function() end, + children = false :: false, + cleanups = false :: false } local function read_node_value() @@ -169,10 +214,17 @@ local function create(value: T): (Node, () -> T) end return table.freeze { + create_and_open_scope = create_and_open_scope, + open_scope = open_scope, + close_scope = close_scope, + get_scope = get_scope, + add_cleanup = add_cleanup, + run_cleanups = run_cleanups, set_effect = set_effect, get = get, set = set, link = link, + set_child = set_child, capture = capture, capture_and_link = capture_and_link, create = create :: ((value: T) -> (Node, () -> T)) & (() -> (Node, () -> T)), diff --git a/src/init.luau b/src/init.luau index 416ac75..68bde72 100644 --- a/src/init.luau +++ b/src/init.luau @@ -9,13 +9,14 @@ local create = require(script.create) local apply = require(script.apply) local source = require(script.source) local watch = require(script.watch) -local cleanup, clean_garbage = require(script.cleanup)() +local cleanup = require(script.cleanup) local untrack = require(script.untrack) local derive = require(script.derive) local indexes, values = require(script.maps)() local spring, update_springs = require(script.spring)() local action = require(script.action)() local throw = require(script.throw) +local _, sweep = require(script.on_gc)() local flags = require(script.flags) export type Source = source.Source @@ -33,7 +34,7 @@ local function step(dt: number) debug.profilebegin("VIDE GARBAGE CLEANUP") end - clean_garbage() + sweep() if game then debug.profileend() diff --git a/src/maps.luau b/src/maps.luau index 977cfe2..4527780 100644 --- a/src/maps.luau +++ b/src/maps.luau @@ -5,11 +5,16 @@ if not game then script = require "test/relative-string" end local throw = require(script.Parent.throw) local flags = require(script.Parent.flags) local graph = require(script.Parent.graph) -local _, _, manual_cleanup_mode, cleanup_ref = require(script.Parent.cleanup)() type Node = graph.Node local create = graph.create local set = graph.set local capture = graph.capture +local run_cleanups = graph.run_cleanups +local set_child = graph.set_child +local create_and_open_scope = graph.create_and_open_scope +local get_scope = graph.get_scope +local open_scope = graph.open_scope +local close_scope = graph.close_scope local link = graph.link type Map = { [K]: V } @@ -31,7 +36,7 @@ local function indexes(input: () -> Map, transform: (() -> VI, local remove_queue = {} :: { K } local output_array = {} :: { VO } - local cleanups = {} :: Map () }> + local scopes = {} :: Map> local function recompute(data) -- queue removed values @@ -43,14 +48,13 @@ local function indexes(input: () -> Map, transform: (() -> VI, -- remove queued values for _, i in next, remove_queue do - for _, callback in next, cleanups[i] do - callback() -- todo: pcall - end + run_cleanups(scopes[i]) + input_cache[i] = nil output_cache[i] = nil input_nodes[i] = nil - cleanups[i] = nil + scopes[i] = nil end table.clear(remove_queue) @@ -61,14 +65,16 @@ local function indexes(input: () -> Map, transform: (() -> VI, if cv ~= v then if cv == nil then - manual_cleanup_mode(transform) + local scope = create(false) + + create_and_open_scope(scope) local node, get_value = create(v) input_nodes[i] = node output_cache[i] = transform(get_value, i) input_cache[i] = v - cleanups[i] = manual_cleanup_mode(nil) + close_scope() else set(input_nodes[i], v) input_cache[i] = v diff --git a/src/on_gc.luau b/src/on_gc.luau new file mode 100644 index 0000000..7a5f722 --- /dev/null +++ b/src/on_gc.luau @@ -0,0 +1,40 @@ +if not game then script = require "test/relative-string" end + +local flags = require(script.Parent.flags) +local throw = require(script.Parent.throw) + + +-- array of all cleanup callbacks +local cleanup_callbacks = {} :: { [number]: () -> () } -- always dense +-- weak array of all cleanup lifetimes +local cleanup_lifetime = {} :: { [number]: unknown } -- can be sparse +setmetatable(cleanup_lifetime :: any, { __mode = "v" }) + +local function on_gc(lifetime: unknown, callback: () -> ()) + local id = #cleanup_callbacks + 1 + cleanup_lifetime[id :: any] = lifetime -- todo + cleanup_callbacks[id] = callback +end + +local function sweep() + for id = #cleanup_callbacks, 1, -1 do + if cleanup_lifetime[id] == nil then -- lifetime was garbage collected + local callback = cleanup_callbacks[id] + + do -- swap and pop + local max_id = #cleanup_callbacks + + cleanup_callbacks[id] = cleanup_callbacks[max_id] + cleanup_callbacks[max_id] = nil + + cleanup_lifetime[id] = cleanup_lifetime[max_id] + cleanup_lifetime[max_id] = nil + end + + local ok, err: string? = pcall(callback) + if not ok then warn(`error occured during cleanup: {err}`) end + end + end +end + +return function() return on_gc, sweep end diff --git a/src/root.luau b/src/root.luau new file mode 100644 index 0000000..51c42c1 --- /dev/null +++ b/src/root.luau @@ -0,0 +1,24 @@ +if not game then script = require "test/relative-string" end + +local flags = require(script.Parent.flags) +local throw = require(script.Parent.throw) +local graph = require(script.Parent.graph) +local create = graph.create +local create_and_open_scope = graph.create_and_open_scope +local close_scope = graph.close_scope +local get_scope = graph.get_scope +local add_cleanup = graph.add_cleanup + +local function root(fn: () -> T): T + local node = create(nil) -- todo: lifetime with return vaue from fn + + create_and_open_scope(node) + + local v = fn() + + close_scope() + + return v +end + +return root diff --git a/src/spring.luau b/src/spring.luau index 1f8fdb3..fa8da9a 100644 --- a/src/spring.luau +++ b/src/spring.luau @@ -26,7 +26,7 @@ local graph = require(script.Parent.graph) type Node = graph.Node local create = graph.create local set = graph.set -local set_effect = graph.set_effect +local set_child = graph.set_child local capture = graph.capture local UPDATE_RATE = 120 @@ -178,19 +178,19 @@ local function spring(source: () -> T, period: number?, damping_ratio: number } -- reschedule spring for simulation on input update - local function input_updated(node) + local function input_updated() local v = source() data.x1_123, data.x1_456 = type_to_vec6[typeof(v)](v) data.source_value = v - springs[data] = node -- todo: investigate why insertion is not O(1) at ~20k springs + springs[data] = output -- todo: investigate why insertion is not O(1) at ~20k springs end -- unused field, use so output prevents gc of inputs - output.derive = source :: any + output.effect = input_updated -- register above function as side-effect for all inputs for _, input in next, inputs do - set_effect(input, input_updated, output) + set_child(input, output) end return output_get, data diff --git a/src/watch.luau b/src/watch.luau index 27233ca..edde546 100644 --- a/src/watch.luau +++ b/src/watch.luau @@ -1,25 +1,36 @@ if not game then script = require "test/relative-string" end local graph = require(script.Parent.graph) -local set_effect = graph.set_effect +local create = graph.create +local set_child = graph.set_child local capture = graph.capture +local create_and_open_scope = graph.create_and_open_scope +local close_scope = graph.close_scope + +local ref = {} local function watch(effect: () -> ()): () -> () + local node = create(nil) + + create_and_open_scope(node) + local nodes = capture(effect :: () -> nil) - -- store aside captured nodes in new table - nodes = table.clone(nodes) + close_scope() + + node.effect = effect -- register effect with permanent lifetime - for _, node in next, nodes do - set_effect(node, effect, true) + for _, parent in next, nodes do + set_child(parent, node) end + ref[node] = true -- prevent gc of node + local function unwatch() -- unregister effect from all nodes - for _, node in next, nodes do - set_effect(node, effect, nil) - end + node.effect = function() end + ref[node] = nil end return unwatch diff --git a/test/tests.luau b/test/tests.luau index da7a3ec..10d6e52 100644 --- a/test/tests.luau +++ b/test/tests.luau @@ -1431,31 +1431,31 @@ TEST("strict", function() local indexes, values = vide.indexes, vide.values local cleanup = vide.cleanup - do CASE "error on derived callback yield" - local state = source(1) + -- do CASE "error on derived callback yield" + -- local state = source(1) - local ok = pcall(function() - local _derived = derive(function() - coroutine.yield() - return state() - end) - end) + -- local ok = pcall(function() + -- local _derived = derive(function() + -- coroutine.yield() + -- return state() + -- end) + -- end) - CHECK(not ok) - end + -- CHECK(not ok) + -- end - do CASE "error on watcher callback yield" - local state = source(1) + -- do CASE "error on watcher callback yield" + -- local state = source(1) - local ok = pcall(function() - local _derived = watch(function() - coroutine.yield() - state() - end) - end) + -- local ok = pcall(function() + -- local _derived = watch(function() + -- coroutine.yield() + -- state() + -- end) + -- end) - CHECK(not ok) - end + -- CHECK(not ok) + -- end do CASE "run derived callback twice" local state = source(1)