From 9769525f014c84dcd14bd612b1336e7c45afc0cb Mon Sep 17 00:00:00 2001 From: Nathan Flurry Date: Fri, 4 Sep 2026 01:32:56 -0700 Subject: [PATCH] feat(rivetkit): add Node actor worker threads --- Cargo.lock | 1 + .../docs/general/node-worker-threads.mdx | 51 + .../docs/general/registry-configuration.mdx | 3 + docs/content/docs/general/runtime-modes.mdx | 2 + docs/sidebar.json | 4 + examples/docs/general-worker-threads/basic.ts | 16 + rivetkit-asyncapi/asyncapi.json | 2 +- .../actor.worker_acquire_timed_out.json | 5 + ...ctor.worker_pool_actor_not_registered.json | 5 + .../errors/actor.worker_pool_closed.json | 5 + ...ctor.worker_pool_duplicate_assignment.json | 5 + .../actor.worker_pool_invalid_config.json | 5 + .../actor.worker_registration_rejected.json | 5 + .../errors/actor.worker_spawn_failed.json | 5 + .../errors/actor.worker_thread_lost.json | 5 + .../rivetkit-core/src/actor/config.rs | 14 + .../packages/rivetkit-core/src/actor/task.rs | 47 +- .../packages/rivetkit-core/src/error.rs | 7 + .../src/registry/envoy_callbacks.rs | 6 +- .../rivetkit-core/src/registry/mod.rs | 305 ++- .../rivetkit-core/src/registry/worker_pool.rs | 1655 +++++++++++++++++ .../packages/rivetkit-core/src/serverless.rs | 11 +- .../artifacts/registry-config.json | 8 +- .../packages/rivetkit-napi/Cargo.toml | 1 + .../packages/rivetkit-napi/index.d.ts | 10 + .../rivetkit-napi/src/actor_factory.rs | 168 +- .../packages/rivetkit-napi/src/lib.rs | 1 + .../rivetkit-napi/src/napi_actor_events.rs | 33 + .../packages/rivetkit-napi/src/registry.rs | 259 ++- .../packages/rivetkit-napi/src/worker_pool.rs | 266 +++ .../rivetkit-napi/tests/napi_actor_events.rs | 29 + .../rivetkit/src/registry/config/index.ts | 18 + .../packages/rivetkit/src/registry/index.ts | 65 +- .../rivetkit/src/registry/napi-runtime.ts | 66 + .../packages/rivetkit/src/registry/native.ts | 125 +- .../rivetkit/src/registry/node-worker-pool.ts | 431 +++++ .../packages/rivetkit/src/registry/runtime.ts | 47 + .../tests/napi-worker-environments.test.ts | 45 + .../rivetkit/tests/runtime-selection.test.ts | 55 + .../rivetkit/tests/worker-threads.test.ts | 103 + 40 files changed, 3800 insertions(+), 94 deletions(-) create mode 100644 docs/content/docs/general/node-worker-threads.mdx create mode 100644 examples/docs/general-worker-threads/basic.ts create mode 100644 rivetkit-rust/engine/artifacts/errors/actor.worker_acquire_timed_out.json create mode 100644 rivetkit-rust/engine/artifacts/errors/actor.worker_pool_actor_not_registered.json create mode 100644 rivetkit-rust/engine/artifacts/errors/actor.worker_pool_closed.json create mode 100644 rivetkit-rust/engine/artifacts/errors/actor.worker_pool_duplicate_assignment.json create mode 100644 rivetkit-rust/engine/artifacts/errors/actor.worker_pool_invalid_config.json create mode 100644 rivetkit-rust/engine/artifacts/errors/actor.worker_registration_rejected.json create mode 100644 rivetkit-rust/engine/artifacts/errors/actor.worker_spawn_failed.json create mode 100644 rivetkit-rust/engine/artifacts/errors/actor.worker_thread_lost.json create mode 100644 rivetkit-rust/packages/rivetkit-core/src/registry/worker_pool.rs create mode 100644 rivetkit-typescript/packages/rivetkit-napi/src/worker_pool.rs create mode 100644 rivetkit-typescript/packages/rivetkit/src/registry/node-worker-pool.ts create mode 100644 rivetkit-typescript/packages/rivetkit/tests/napi-worker-environments.test.ts create mode 100644 rivetkit-typescript/packages/rivetkit/tests/worker-threads.test.ts diff --git a/Cargo.lock b/Cargo.lock index 4562be0684..28dc5ebc8d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -6356,6 +6356,7 @@ dependencies = [ "tracing-logfmt", "tracing-stackdriver", "tracing-subscriber", + "uuid", ] [[package]] diff --git a/docs/content/docs/general/node-worker-threads.mdx b/docs/content/docs/general/node-worker-threads.mdx new file mode 100644 index 0000000000..54db2900c1 --- /dev/null +++ b/docs/content/docs/general/node-worker-threads.mdx @@ -0,0 +1,51 @@ +--- +title: "Node.js Worker Threads" +description: "Run Rivet Actors across isolated Node.js event loops in a long-running Node.js process." +--- + +Node.js normally runs every actor in a process on one JavaScript event loop. Set `actorsPerThread` to distribute actor JavaScript across worker threads when one actor must not block every other actor in the process. + + + +`actorsPerThread` is a hard limit that includes actors which are starting, running, or stopping. Use `1` to give every actor generation its own event loop. Larger values reduce memory usage by allowing several actors to share one event loop. + +## Placement and scaling + +RivetKit creates baseline threads on demand, up to `os.availableParallelism()`, and spreads the first actors across them. After reaching that baseline, RivetKit fills existing threads before creating overflow threads. There is no configurable maximum thread count. + +An actor generation stays on its selected thread for its entire lifetime. RivetKit does not move running actors to compact the pool. An overflow thread becomes eligible to exit after its final actor completely stops, with a short idle delay to prevent thread churn during bursts. A sleeping actor is no longer resident and does not occupy a thread slot. + +Requests for an existing actor are delivered directly from the native runtime to that actor's worker thread. They do not pass through the main JavaScript event loop. If the main event loop is blocked, existing actors on other threads can continue running, but RivetKit cannot create or retire threads until the main event loop becomes available. + +## Requirements and constraints + +Worker threads require the native runtime and a persistent Node.js process. They are not supported with serverless handlers, the wasm runtime, Bun, Deno, `node -e`, stdin, or the Node.js REPL. + +RivetKit loads the file named by `process.argv[1]` in every new worker. That entrypoint must build the same actor registry and call `registry.start()` or `registry.startAndWait()` while the module is evaluated. Do not guard the registry startup call with `isMainThread`. + +Each worker validates its actor names and runtime configuration against the main registry before it can receive an actor. The waiting actor start fails if the entrypoint cannot load, does not start a registry, loads a different actor configuration, or cannot initialize the native runtime. Worker acquisition times out after 30 seconds; a worker that finishes booting later remains available for the next actor start. + +The complete entrypoint module executes independently in every worker. Guard unrelated main-only side effects individually, but do not guard the RivetKit registry startup call. Only one worker-enabled registry may be started from an entrypoint in this first version. Development runners that replace `process.argv[1]` with a launcher or wrapper are not supported unless that file reloads the same registry entrypoint. + +## Isolation and memory + +Each worker has its own V8 isolate, event loop, module cache, globals, and module-level singletons. JavaScript objects are not shared between workers. Actor definitions are evaluated separately in each worker, while durable actor state continues to use RivetKit state, KV, and SQLite normally. + +In one local Node.js 24 Linux benchmark, an idle worker with the RivetKit native addon added approximately 9.2 MiB PSS or 11.0 MiB RSS after the first worker. In a separate run, a minimal live actor added approximately 0.53 MiB PSS or 0.54 MiB RSS. At full occupancy, a lower-bound estimate per actor is `actor memory + thread memory / actorsPerThread`: + +| `actorsPerThread` | Approximate PSS per actor | Approximate RSS per actor | +|---:|---:|---:| +| 1 | 9.7 MiB | 11.6 MiB | +| 4 | 2.8 MiB | 3.3 MiB | +| 8 | 1.7 MiB | 1.9 MiB | +| 16 | 1.1 MiB | 1.2 MiB | + +These are local lower bounds, not product guarantees. Application modules, actor definitions, state, queues, SQLite, and caches add to them. The table also assumes full threads. For a pool-wide estimate, use `live actors × actor memory + live threads × thread memory`. + +Node.js and the system allocator may retain high-water RSS after a worker exits. If one actor blocks its worker, other actors assigned to that worker are also delayed. + +## Worker failures + +If a worker exits unexpectedly, its actors fail through the normal Rivet lifecycle. The control plane decides whether and where to start replacement generations. RivetKit never moves or locally resurrects the failed generation. + +Worker-pool failures use the `actor.worker_*` error codes, including `actor.worker_acquire_timed_out`, `actor.worker_spawn_failed`, and `actor.worker_thread_lost`. diff --git a/docs/content/docs/general/registry-configuration.mdx b/docs/content/docs/general/registry-configuration.mdx index 58762d2a76..3cb5562e32 100644 --- a/docs/content/docs/general/registry-configuration.mdx +++ b/docs/content/docs/general/registry-configuration.mdx @@ -43,6 +43,8 @@ After configuring your registry, start it: See [Runtime Modes](/actors/docs/general/runtime-modes) for details on when to use each mode. +Long-running Node.js registries can also set `actorsPerThread` to isolate actor JavaScript across multiple event loops. See [Node.js Worker Threads](/actors/docs/general/node-worker-threads) for scheduling, entrypoint, memory, and runtime constraints. + ## Environment Variables Many configuration options can be set via environment variables. See [Environment Variables](/actors/docs/general/environment-variables) for a complete reference. @@ -55,4 +57,5 @@ Many configuration options can be set via environment variables. See [Environmen ## Related - [Actor Configuration](/actors/docs/general/actor-configuration): Configure individual actors +- [Node.js Worker Threads](/actors/docs/general/node-worker-threads): Run actors across isolated Node.js event loops - [HTTP Server Setup](/actors/docs/general/http-server): Set up HTTP routing and middleware diff --git a/docs/content/docs/general/runtime-modes.mdx b/docs/content/docs/general/runtime-modes.mdx index 1dfee06e0a..a9cbb8b859 100644 --- a/docs/content/docs/general/runtime-modes.mdx +++ b/docs/content/docs/general/runtime-modes.mdx @@ -21,6 +21,8 @@ Runner is the default mode. Your app runs as a long-running process that opens a - **No public endpoint**: Your app connects out to Rivet, so it does not need to be publicly reachable or registered in the dashboard. - **Custom scaling**: You control how runner processes are pooled and scaled. +Runner deployments on Node.js can optionally distribute actor JavaScript across worker threads. See [Node.js Worker Threads](/actors/docs/general/node-worker-threads) for the entrypoint and isolation constraints. + ### Example diff --git a/docs/sidebar.json b/docs/sidebar.json index 13992a1f4c..f6efdb7322 100644 --- a/docs/sidebar.json +++ b/docs/sidebar.json @@ -296,6 +296,10 @@ "title": "WASM vs Native SDK", "href": "/actors/docs/general/wasm-vs-native-sdk" }, + { + "title": "Node.js Worker Threads", + "href": "/actors/docs/general/node-worker-threads" + }, { "title": "Registry Configuration", "href": "/actors/docs/general/registry-configuration" diff --git a/examples/docs/general-worker-threads/basic.ts b/examples/docs/general-worker-threads/basic.ts new file mode 100644 index 0000000000..eeef998a93 --- /dev/null +++ b/examples/docs/general-worker-threads/basic.ts @@ -0,0 +1,16 @@ +import { actor, setup } from "rivetkit"; + +const job = actor({ + state: {}, + actions: {}, +}); + +const registry = setup({ + use: { job }, + // Each worker thread can host at most four live actors. + actorsPerThread: 4, +}); + +// Keep this startup call unguarded so worker threads can attach their copy of +// the registry instead of starting another runner connection. +registry.start(); diff --git a/rivetkit-asyncapi/asyncapi.json b/rivetkit-asyncapi/asyncapi.json index 32b33213c3..aa8f74eec3 100644 --- a/rivetkit-asyncapi/asyncapi.json +++ b/rivetkit-asyncapi/asyncapi.json @@ -556,4 +556,4 @@ } } } -} \ No newline at end of file +} diff --git a/rivetkit-rust/engine/artifacts/errors/actor.worker_acquire_timed_out.json b/rivetkit-rust/engine/artifacts/errors/actor.worker_acquire_timed_out.json new file mode 100644 index 0000000000..e0b57ffcab --- /dev/null +++ b/rivetkit-rust/engine/artifacts/errors/actor.worker_acquire_timed_out.json @@ -0,0 +1,5 @@ +{ + "code": "worker_acquire_timed_out", + "group": "actor", + "message": "Timed out waiting for a worker thread" +} \ No newline at end of file diff --git a/rivetkit-rust/engine/artifacts/errors/actor.worker_pool_actor_not_registered.json b/rivetkit-rust/engine/artifacts/errors/actor.worker_pool_actor_not_registered.json new file mode 100644 index 0000000000..be9f9b7872 --- /dev/null +++ b/rivetkit-rust/engine/artifacts/errors/actor.worker_pool_actor_not_registered.json @@ -0,0 +1,5 @@ +{ + "code": "worker_pool_actor_not_registered", + "group": "actor", + "message": "Actor is not registered" +} \ No newline at end of file diff --git a/rivetkit-rust/engine/artifacts/errors/actor.worker_pool_closed.json b/rivetkit-rust/engine/artifacts/errors/actor.worker_pool_closed.json new file mode 100644 index 0000000000..004b015699 --- /dev/null +++ b/rivetkit-rust/engine/artifacts/errors/actor.worker_pool_closed.json @@ -0,0 +1,5 @@ +{ + "code": "worker_pool_closed", + "group": "actor", + "message": "Worker pool is closed" +} \ No newline at end of file diff --git a/rivetkit-rust/engine/artifacts/errors/actor.worker_pool_duplicate_assignment.json b/rivetkit-rust/engine/artifacts/errors/actor.worker_pool_duplicate_assignment.json new file mode 100644 index 0000000000..690e9197e5 --- /dev/null +++ b/rivetkit-rust/engine/artifacts/errors/actor.worker_pool_duplicate_assignment.json @@ -0,0 +1,5 @@ +{ + "code": "worker_pool_duplicate_assignment", + "group": "actor", + "message": "Actor generation is already assigned" +} \ No newline at end of file diff --git a/rivetkit-rust/engine/artifacts/errors/actor.worker_pool_invalid_config.json b/rivetkit-rust/engine/artifacts/errors/actor.worker_pool_invalid_config.json new file mode 100644 index 0000000000..4ba705d265 --- /dev/null +++ b/rivetkit-rust/engine/artifacts/errors/actor.worker_pool_invalid_config.json @@ -0,0 +1,5 @@ +{ + "code": "worker_pool_invalid_config", + "group": "actor", + "message": "Invalid worker pool configuration" +} \ No newline at end of file diff --git a/rivetkit-rust/engine/artifacts/errors/actor.worker_registration_rejected.json b/rivetkit-rust/engine/artifacts/errors/actor.worker_registration_rejected.json new file mode 100644 index 0000000000..08fb145ca6 --- /dev/null +++ b/rivetkit-rust/engine/artifacts/errors/actor.worker_registration_rejected.json @@ -0,0 +1,5 @@ +{ + "code": "worker_registration_rejected", + "group": "actor", + "message": "Worker thread registration was rejected" +} \ No newline at end of file diff --git a/rivetkit-rust/engine/artifacts/errors/actor.worker_spawn_failed.json b/rivetkit-rust/engine/artifacts/errors/actor.worker_spawn_failed.json new file mode 100644 index 0000000000..7ae979e884 --- /dev/null +++ b/rivetkit-rust/engine/artifacts/errors/actor.worker_spawn_failed.json @@ -0,0 +1,5 @@ +{ + "code": "worker_spawn_failed", + "group": "actor", + "message": "Worker thread failed to start" +} \ No newline at end of file diff --git a/rivetkit-rust/engine/artifacts/errors/actor.worker_thread_lost.json b/rivetkit-rust/engine/artifacts/errors/actor.worker_thread_lost.json new file mode 100644 index 0000000000..70dbd30ac9 --- /dev/null +++ b/rivetkit-rust/engine/artifacts/errors/actor.worker_thread_lost.json @@ -0,0 +1,5 @@ +{ + "code": "worker_thread_lost", + "group": "actor", + "message": "Actor worker thread exited." +} \ No newline at end of file diff --git a/rivetkit-rust/packages/rivetkit-core/src/actor/config.rs b/rivetkit-rust/packages/rivetkit-core/src/actor/config.rs index e9519f8824..c1f5c11757 100644 --- a/rivetkit-rust/packages/rivetkit-core/src/actor/config.rs +++ b/rivetkit-rust/packages/rivetkit-core/src/actor/config.rs @@ -1,8 +1,10 @@ use std::fmt; +use std::fmt::Write as _; use std::sync::Arc; use std::time::Duration; use rivet_envoy_client::config::HttpRequest; +use sha2::{Digest, Sha256}; use crate::inspector::InspectorTabEntry; @@ -216,6 +218,18 @@ pub struct ActorConfigInput { } impl ActorConfig { + /// Stable within one runtime build and process. Used to reject worker + /// environments that evaluated a different actor configuration before their + /// callback factories become schedulable. + pub fn worker_pool_fingerprint(&self) -> String { + let digest = Sha256::digest(format!("{self:#?}").as_bytes()); + let mut encoded = String::with_capacity(digest.len() * 2); + for byte in digest { + let _ = write!(encoded, "{byte:02x}"); + } + encoded + } + pub fn from_input(config: ActorConfigInput) -> Self { let mut actor_config = Self { name: config.name, diff --git a/rivetkit-rust/packages/rivetkit-core/src/actor/task.rs b/rivetkit-rust/packages/rivetkit-core/src/actor/task.rs index c3cdbe3d22..a7fb392d91 100644 --- a/rivetkit-rust/packages/rivetkit-core/src/actor/task.rs +++ b/rivetkit-rust/packages/rivetkit-core/src/actor/task.rs @@ -40,6 +40,7 @@ use futures::FutureExt; use parking_lot::Mutex; use tokio::sync::{broadcast, mpsc, oneshot}; use tokio::task::{JoinError, JoinHandle}; +use tokio_util::sync::CancellationToken; use tracing::{Instrument, instrument::WithSubscriber}; use crate::actor::action::ActionDispatchError; @@ -348,6 +349,8 @@ pub struct ActorTask { pub lifecycle: LifecycleState, pub factory: Arc, pub ctx: ActorContext, + /// Cancels the foreign-runtime adapter when its owning environment exits. + runtime_lost: Option, // === STARTUP === pub start_input: Option>, @@ -402,6 +405,31 @@ impl ActorTask { factory: Arc, ctx: ActorContext, start_input: Option>, + ) -> Self { + Self::new_with_runtime_loss( + actor_id, + generation, + lifecycle_inbox, + dispatch_inbox, + lifecycle_events, + factory, + ctx, + start_input, + None, + ) + } + + #[allow(clippy::too_many_arguments)] + pub fn new_with_runtime_loss( + actor_id: String, + generation: u32, + lifecycle_inbox: mpsc::UnboundedReceiver, + dispatch_inbox: mpsc::UnboundedReceiver, + lifecycle_events: mpsc::UnboundedReceiver, + factory: Arc, + ctx: ActorContext, + start_input: Option>, + runtime_lost: Option, ) -> Self { let (actor_event_tx, actor_event_rx) = mpsc::unbounded_channel(); let (inspector_overlay_tx, _) = broadcast::channel(INSPECTOR_OVERLAY_CHANNEL_CAPACITY); @@ -428,6 +456,7 @@ impl ActorTask { lifecycle: LifecycleState::default(), factory, ctx, + runtime_lost, start_input, actor_event_tx: Some(actor_event_tx), actor_event_rx: Some(actor_event_rx), @@ -1347,10 +1376,26 @@ impl ActorTask { startup_ready: startup_ready_tx, }; let factory = self.factory.clone(); + let runtime_lost = self.runtime_lost.clone(); let run_dispatch = tracing::dispatcher::get_default(Clone::clone); self.run_handle = Some(RuntimeSpawner::spawn( async move { - match AssertUnwindSafe(factory.start(start)).catch_unwind().await { + let outcome = if let Some(runtime_lost) = runtime_lost { + if runtime_lost.is_cancelled() { + return Err(crate::error::ActorRuntime::WorkerThreadLost.build()); + } + let run = AssertUnwindSafe(factory.start(start)).catch_unwind(); + tokio::select! { + biased; + _ = runtime_lost.cancelled() => { + return Err(crate::error::ActorRuntime::WorkerThreadLost.build()); + } + outcome = run => outcome, + } + } else { + AssertUnwindSafe(factory.start(start)).catch_unwind().await + }; + match outcome { Ok(result) => result, Err(_) => Err(ActorRuntime::Panicked { operation: "run handler".to_owned(), diff --git a/rivetkit-rust/packages/rivetkit-core/src/error.rs b/rivetkit-rust/packages/rivetkit-core/src/error.rs index 195d7ff837..820c42212c 100644 --- a/rivetkit-rust/packages/rivetkit-core/src/error.rs +++ b/rivetkit-rust/packages/rivetkit-core/src/error.rs @@ -234,6 +234,13 @@ pub enum ActorRuntime { "Actor task panicked while running {operation}." )] Panicked { operation: String }, + + #[error( + "worker_thread_lost", + "Actor worker thread exited.", + "The Node.js worker thread hosting this actor exited." + )] + WorkerThreadLost, } #[derive(RivetError, Debug, Clone, Deserialize, Serialize)] diff --git a/rivetkit-rust/packages/rivetkit-core/src/registry/envoy_callbacks.rs b/rivetkit-rust/packages/rivetkit-core/src/registry/envoy_callbacks.rs index 8ed87ce9f4..a8bc6095cb 100644 --- a/rivetkit-rust/packages/rivetkit-core/src/registry/envoy_callbacks.rs +++ b/rivetkit-rust/packages/rivetkit-core/src/registry/envoy_callbacks.rs @@ -17,10 +17,10 @@ impl EnvoyCallbacks for RegistryCallbacks { let actor_name = config.name.clone(); let key = actor_key_from_protocol(config.key.clone()); let input = config.input.clone(); - let factory = dispatcher.factories.get(&actor_name).cloned(); + let actor_config = dispatcher.actor_config(&actor_name).cloned(); Box::pin(async move { - let factory = factory.ok_or_else(|| { + let actor_config = actor_config.ok_or_else(|| { ActorRuntime::NotRegistered { actor_name: actor_name.clone(), } @@ -32,7 +32,7 @@ impl EnvoyCallbacks for RegistryCallbacks { generation, &actor_name, key, - factory.as_ref(), + &actor_config, )?; dispatcher diff --git a/rivetkit-rust/packages/rivetkit-core/src/registry/mod.rs b/rivetkit-rust/packages/rivetkit-core/src/registry/mod.rs index 24c897a7be..09d3378d59 100644 --- a/rivetkit-rust/packages/rivetkit-core/src/registry/mod.rs +++ b/rivetkit-rust/packages/rivetkit-core/src/registry/mod.rs @@ -34,7 +34,7 @@ use url::Url; use vbare::OwnedVersionedData; use crate::actor::action::ActionDispatchError; -use crate::actor::config::CanHibernateWebSocket; +use crate::actor::config::{ActorConfig, CanHibernateWebSocket}; use crate::actor::connection::{ConnHandle, HibernatableConnectionMetadata}; use crate::actor::context::{ActorContext, InspectorAttachGuard}; use crate::actor::factory::ActorFactory; @@ -67,13 +67,35 @@ mod inspector_ws; #[cfg(feature = "native-runtime")] mod runner_config; mod websocket; +#[cfg(feature = "native-runtime")] +pub mod worker_pool; use inspector::build_actor_inspector; use websocket::is_actor_connect_path; +#[cfg(feature = "native-runtime")] +use worker_pool::{ + ActorFactoryLease, ActorWorkerPool, ActorWorkerPoolCallbacks, ActorWorkerPoolConfig, +}; #[derive(Default)] pub struct CoreRegistry { factories: HashMap>, + actor_configs: HashMap, + #[cfg(feature = "native-runtime")] + worker_pool: Option>, +} + +pub(crate) enum ActorFactoryProvider { + Static(HashMap>), + #[cfg(feature = "native-runtime")] + WorkerPool(Arc), +} + +struct ActorFactorySelection { + factory: Arc, + runtime_lost: Option, + #[cfg(feature = "native-runtime")] + lease: Option, } #[derive(Clone)] @@ -131,6 +153,17 @@ struct ActorTaskHandle { lifecycle: mpsc::UnboundedSender, dispatch: mpsc::UnboundedSender, join: Arc>>>>, + #[cfg(feature = "native-runtime")] + worker_lease: Arc>>, +} + +impl ActorTaskHandle { + #[cfg(feature = "native-runtime")] + fn release_worker_lease(&self) { + if let Some(lease) = self.worker_lease.lock().take() { + lease.release(); + } + } } type ActiveActorInstance = Arc; @@ -165,14 +198,21 @@ struct PendingStop { } pub(crate) struct RegistryDispatcher { - pub(crate) factories: HashMap>, + factory_provider: ActorFactoryProvider, + actor_configs: HashMap, actor_instances: SccHashMap, - starting_instances: SccHashMap>, + starting_instances: SccHashMap, pending_stops: SccHashMap, region: String, handle_inspector_http_in_runtime: bool, } +#[derive(Clone)] +struct StartingActorInstance { + generation: u32, + notify: Arc, +} + pub(crate) struct RegistryCallbacks { pub(crate) dispatcher: Arc, } @@ -542,16 +582,58 @@ impl CoreRegistry { } pub fn register(&mut self, name: &str, factory: ActorFactory) { + self.actor_configs + .insert(name.to_owned(), factory.config().clone()); self.factories.insert(name.to_owned(), Arc::new(factory)); } pub fn register_shared(&mut self, name: &str, factory: Arc) { + self.actor_configs + .insert(name.to_owned(), factory.config().clone()); self.factories.insert(name.to_owned(), factory); } + pub fn register_config(&mut self, name: &str, config: ActorConfig) { + self.actor_configs.insert(name.to_owned(), config); + } + + #[cfg(feature = "native-runtime")] + pub fn enable_worker_pool( + &mut self, + config: ActorWorkerPoolConfig, + callbacks: ActorWorkerPoolCallbacks, + ) -> Result> { + if self.worker_pool.is_some() { + anyhow::bail!("actor worker pool is already configured"); + } + if !self.factories.is_empty() { + anyhow::bail!( + "main worker-pool registry must register actor configs without callback factories" + ); + } + let expected = self + .actor_configs + .iter() + .map(|(name, config)| (name.clone(), config.worker_pool_fingerprint())); + let pool = ActorWorkerPool::new(config, expected, callbacks); + self.worker_pool = Some(pool.clone()); + Ok(pool) + } + + #[cfg(feature = "native-runtime")] + pub fn into_worker_factories(self) -> Result>> { + if self.worker_pool.is_some() { + anyhow::bail!("a worker registration registry cannot host another worker pool"); + } + if self.factories.len() != self.actor_configs.len() { + anyhow::bail!("worker registry is missing actor callback factories"); + } + Ok(self.factories) + } + pub fn normal_metadata_payload(&self, config: &ServeConfig) -> ServerlessMetadataPayload { serverless_metadata_payload( - build_actor_metadata_map_from_factories(&self.factories), + build_actor_metadata_map_from_configs(&self.actor_configs), config, ServerlessMetadataEnvoyKind::Normal {}, ) @@ -559,7 +641,7 @@ impl CoreRegistry { pub fn serverless_metadata_payload(&self, config: &ServeConfig) -> ServerlessMetadataPayload { serverless_metadata_payload( - build_actor_metadata_map_from_factories(&self.factories), + build_actor_metadata_map_from_configs(&self.actor_configs), config, ServerlessMetadataEnvoyKind::Serverless {}, ) @@ -592,6 +674,8 @@ impl CoreRegistry { config.pool_name.clone(), ); + #[cfg(feature = "native-runtime")] + let worker_pool = self.worker_pool.clone(); let dispatcher = self.into_dispatcher(&config); let manage_engine = should_manage_engine(&config.endpoint, config.engine_spawn)?; #[cfg(feature = "native-runtime")] @@ -675,13 +759,25 @@ impl CoreRegistry { tokio::join!(shutdown_envoy, shutdown_development_processes); #[cfg(not(feature = "native-runtime"))] shutdown_envoy.await; + #[cfg(feature = "native-runtime")] + if let Some(worker_pool) = worker_pool { + worker_pool.shutdown(); + } Ok(()) } fn into_dispatcher(self, config: &ServeConfig) -> Arc { + #[cfg(feature = "native-runtime")] + let factory_provider = match self.worker_pool { + Some(pool) => ActorFactoryProvider::WorkerPool(pool), + None => ActorFactoryProvider::Static(self.factories), + }; + #[cfg(not(feature = "native-runtime"))] + let factory_provider = ActorFactoryProvider::Static(self.factories); Arc::new(RegistryDispatcher::new( - self.factories, + factory_provider, + self.actor_configs, config.handle_inspector_http_in_runtime, )) } @@ -690,17 +786,23 @@ impl CoreRegistry { self, config: ServeConfig, ) -> Result { + #[cfg(feature = "native-runtime")] + if self.worker_pool.is_some() { + anyhow::bail!("actor worker threads are not supported in serverless mode"); + } crate::serverless::CoreServerlessRuntime::new(self.factories, config).await } } impl RegistryDispatcher { pub(crate) fn new( - factories: HashMap>, + factory_provider: ActorFactoryProvider, + actor_configs: HashMap, handle_inspector_http_in_runtime: bool, ) -> Self { Self { - factories, + factory_provider, + actor_configs, actor_instances: SccHashMap::new(), starting_instances: SccHashMap::new(), pending_stops: SccHashMap::new(), @@ -710,7 +812,46 @@ impl RegistryDispatcher { } pub(crate) fn build_actor_metadata_map(&self) -> HashMap { - build_actor_metadata_map_from_factories(&self.factories) + build_actor_metadata_map_from_configs(&self.actor_configs) + } + + fn actor_config(&self, actor_name: &str) -> Option<&ActorConfig> { + self.actor_configs.get(actor_name) + } + + async fn acquire_factory( + &self, + actor_id: &str, + generation: u32, + actor_name: &str, + ) -> Result { + #[cfg(not(feature = "native-runtime"))] + let _ = (actor_id, generation); + match &self.factory_provider { + ActorFactoryProvider::Static(factories) => { + let factory = factories.get(actor_name).cloned().ok_or_else(|| { + ActorRuntime::NotRegistered { + actor_name: actor_name.to_owned(), + } + .build() + })?; + Ok(ActorFactorySelection { + factory, + runtime_lost: None, + #[cfg(feature = "native-runtime")] + lease: None, + }) + } + #[cfg(feature = "native-runtime")] + ActorFactoryProvider::WorkerPool(pool) => { + let lease = pool.acquire(actor_id, generation, actor_name).await?; + Ok(ActorFactorySelection { + factory: lease.factory(), + runtime_lost: Some(lease.worker_lost()), + lease: Some(lease), + }) + } + } } } @@ -747,13 +888,12 @@ pub(crate) fn serverless_metadata_payload( } } -fn build_actor_metadata_map_from_factories( - factories: &HashMap>, +fn build_actor_metadata_map_from_configs( + configs: &HashMap, ) -> HashMap { - factories + configs .iter() - .map(|(actor_name, factory)| { - let config = factory.config(); + .map(|(actor_name, config)| { let mut metadata = serde_json::Map::new(); if let Some(icon) = &config.icon { metadata.insert("icon".to_owned(), json!(icon)); @@ -766,23 +906,57 @@ fn build_actor_metadata_map_from_factories( .collect() } +#[cfg(test)] +fn build_actor_metadata_map_from_factories( + factories: &HashMap>, +) -> HashMap { + let configs = factories + .iter() + .map(|(name, factory)| (name.clone(), factory.config().clone())) + .collect(); + build_actor_metadata_map_from_configs(&configs) +} + impl RegistryDispatcher { async fn start_actor(self: &Arc, request: StartActorRequest) -> Result<()> { let startup_notify = Arc::new(Notify::new()); let _ = self .starting_instances - .insert_async(request.actor_id.clone(), startup_notify.clone()) + .insert_async( + request.actor_id.clone(), + StartingActorInstance { + generation: request.generation, + notify: startup_notify.clone(), + }, + ) .await; - let factory = self - .factories - .get(&request.actor_name) - .cloned() - .ok_or_else(|| { - ActorRuntime::NotRegistered { - actor_name: request.actor_name.clone(), + let selection = match self + .acquire_factory(&request.actor_id, request.generation, &request.actor_name) + .await + { + Ok(selection) => selection, + Err(error) => { + let pending_stop = self + .pending_stops + .remove_async(&request.actor_id.clone()) + .await + .map(|(_, pending_stop)| pending_stop); + if let Some(pending_stop) = pending_stop { + let _ = pending_stop + .stop_handle + .fail(anyhow::Error::new(RivetError::extract(&error))); } - .build() - })?; + self.finish_starting_actor(&request.actor_id, request.generation) + .await; + return Err(error); + } + }; + let ActorFactorySelection { + factory, + runtime_lost, + #[cfg(feature = "native-runtime")] + lease, + } = selection; let (lifecycle_tx, lifecycle_rx) = mpsc::unbounded_channel(); let (dispatch_tx, dispatch_rx) = mpsc::unbounded_channel(); let (lifecycle_events_tx, lifecycle_events_rx) = mpsc::unbounded_channel(); @@ -807,7 +981,7 @@ impl RegistryDispatcher { }) } }))); - let task = ActorTask::new( + let task = ActorTask::new_with_runtime_loss( request.actor_id.clone(), request.generation, lifecycle_rx, @@ -816,8 +990,11 @@ impl RegistryDispatcher { factory.clone(), request.ctx.clone(), request.input, + runtime_lost, ); - let join = RuntimeSpawner::spawn(task.run()); + let join = Arc::new(TokioMutex::new(Some(RuntimeSpawner::spawn(task.run())))); + #[cfg(feature = "native-runtime")] + let worker_lease = Arc::new(Mutex::new(lease)); let (start_tx, start_rx) = oneshot::channel(); let result: Result> = async { @@ -836,9 +1013,11 @@ impl RegistryDispatcher { ctx: request.ctx.clone(), factory, inspector, - lifecycle: lifecycle_tx, - dispatch: dispatch_tx, - join: Arc::new(TokioMutex::new(Some(join))), + lifecycle: lifecycle_tx.clone(), + dispatch: dispatch_tx.clone(), + join: join.clone(), + #[cfg(feature = "native-runtime")] + worker_lease: worker_lease.clone(), })) } .await @@ -865,9 +1044,7 @@ impl RegistryDispatcher { }, ) .await; - let _ = self - .starting_instances - .remove_async(&request.actor_id.clone()) + self.finish_starting_actor(&request.actor_id, request.generation) .await; let dispatcher = self.clone(); @@ -891,8 +1068,6 @@ impl RegistryDispatcher { .remove_stopping_actor_instance(&actor_id, &instance) .await; }); - startup_notify.notify_waiters(); - Ok(()) } else { self.set_actor_instance_state( @@ -900,25 +1075,61 @@ impl RegistryDispatcher { ActorInstanceState::Active(instance), ) .await; - let _ = self - .starting_instances - .remove_async(&request.actor_id.clone()) + self.finish_starting_actor(&request.actor_id, request.generation) .await; - startup_notify.notify_waiters(); Ok(()) } } Err(error) => { - let _ = self - .starting_instances - .remove_async(&request.actor_id.clone()) + request.ctx.set_local_alarm_callback(None); + request.ctx.configure_lifecycle_events(None); + drop(lifecycle_tx); + drop(dispatch_tx); + if let Some(join) = join.lock().await.take() + && let Err(join_error) = join.await + { + tracing::warn!( + actor_id = %request.actor_id, + ?join_error, + "failed to join actor task after startup failure", + ); + } + #[cfg(feature = "native-runtime")] + if let Some(lease) = worker_lease.lock().take() { + lease.release(); + } + self.finish_starting_actor(&request.actor_id, request.generation) .await; - startup_notify.notify_waiters(); Err(error) } } } + async fn finish_starting_actor(&self, actor_id: &str, generation: u32) { + let Some((_, starting)) = self + .starting_instances + .remove_if_async(&actor_id.to_owned(), |starting| { + starting.generation == generation + }) + .await + else { + if let Some(starting) = self + .starting_instances + .get_async(&actor_id.to_owned()) + .await + { + tracing::warn!( + actor_id, + expected_generation = generation, + actual_generation = starting.generation, + "refused to remove a different actor generation from the startup tracker", + ); + } + return; + }; + starting.notify.notify_waiters(); + } + async fn set_actor_instance_state(&self, actor_id: String, state: ActorInstanceState) { match self.actor_instances.entry_async(actor_id).await { SccEntry::Occupied(mut entry) => { @@ -1166,6 +1377,8 @@ impl RegistryDispatcher { Ok(()) }; instance.ctx.configure_lifecycle_events(None); + #[cfg(feature = "native-runtime")] + instance.release_worker_lease(); if let (Err(shutdown_error), Err(join_error)) = (&shutdown_result, &join_result) { tracing::warn!( @@ -1218,7 +1431,7 @@ impl RegistryDispatcher { generation: u32, actor_name: &str, key: ActorKey, - factory: &ActorFactory, + config: &ActorConfig, ) -> Result { let formatted_key = format_actor_key(&key); let ctx = ActorContext::build( @@ -1228,15 +1441,15 @@ impl RegistryDispatcher { self.region.clone(), Some(generation), handle.get_envoy_key().to_owned(), - factory.config().clone(), + config.clone(), LegacyActorKv::new(handle.clone(), actor_id.to_owned()), SqliteDb::new_with_remote_sqlite( handle.clone(), actor_id.to_owned(), Some(formatted_key), Some(generation as u64), - factory.config().has_database, - factory.config().remote_sqlite, + config.has_database, + config.remote_sqlite, )?, ); ctx.configure_envoy(handle, Some(generation)); diff --git a/rivetkit-rust/packages/rivetkit-core/src/registry/worker_pool.rs b/rivetkit-rust/packages/rivetkit-core/src/registry/worker_pool.rs new file mode 100644 index 0000000000..7ef581cca8 --- /dev/null +++ b/rivetkit-rust/packages/rivetkit-core/src/registry/worker_pool.rs @@ -0,0 +1,1655 @@ +use std::cmp::Reverse; +use std::collections::{BTreeMap, BTreeSet, HashMap, VecDeque}; +use std::sync::{Arc, LazyLock, Weak}; +use std::time::Duration; + +use anyhow::Result; +use parking_lot::Mutex; +use rivet_error::RivetError; +use rivet_metrics::prometheus::{ + Histogram, HistogramOpts, IntCounter, IntCounterVec, IntGauge, IntGaugeVec, Opts, Registry, +}; +use serde::Serialize; +use tokio::sync::Notify; +use tokio_util::sync::CancellationToken; +use uuid::Uuid; + +use crate::actor::factory::ActorFactory; +#[cfg(not(feature = "native-runtime"))] +use crate::runtime::RuntimeSpawner; +use crate::time::{Instant, sleep_until}; + +pub type WorkerId = u64; +pub type WorkerRegistrationEpoch = u64; + +const DEFAULT_ACQUIRE_TIMEOUT: Duration = Duration::from_secs(30); +const DEFAULT_IDLE_RETIRE_DELAY: Duration = Duration::from_secs(30); + +struct WorkerPoolMetrics { + workers: IntGaugeVec, + leases: IntGaugeVec, + available_slots: IntGaugeVec, + queued_acquires: IntGauge, + acquire_duration_seconds: Histogram, + events: IntCounterVec, + actors_failed_worker_loss: IntCounter, +} + +static METRICS: LazyLock = LazyLock::new(WorkerPoolMetrics::new); + +impl WorkerPoolMetrics { + fn new() -> Self { + let workers = IntGaugeVec::new( + Opts::new( + "rivetkit_actor_worker_threads", + "Node actor worker threads by class and lifecycle state", + ), + &["class", "state"], + ) + .expect("create worker thread gauge"); + let leases = IntGaugeVec::new( + Opts::new( + "rivetkit_actor_worker_leases", + "Actor generations assigned to Node workers by class", + ), + &["class"], + ) + .expect("create worker lease gauge"); + let available_slots = IntGaugeVec::new( + Opts::new( + "rivetkit_actor_worker_available_slots", + "Unreserved slots on ready Node workers by class", + ), + &["class"], + ) + .expect("create worker available-slot gauge"); + let queued_acquires = IntGauge::new( + "rivetkit_actor_worker_queued_acquires", + "Actor starts waiting to acquire a Node worker slot", + ) + .expect("create worker queued-acquire gauge"); + let acquire_duration_seconds = Histogram::with_opts( + HistogramOpts::new( + "rivetkit_actor_worker_acquire_duration_seconds", + "Time spent acquiring a Node worker slot", + ) + .buckets(rivet_metrics::MICRO_BUCKETS.to_vec()), + ) + .expect("create worker acquire duration histogram"); + let events = IntCounterVec::new( + Opts::new( + "rivetkit_actor_worker_events_total", + "Node actor worker lifecycle events", + ), + &["event", "class"], + ) + .expect("create worker lifecycle counter"); + let actors_failed_worker_loss = IntCounter::new( + "rivetkit_actor_worker_loss_actor_generations_total", + "Actor generations failed because their Node worker exited", + ) + .expect("create worker-loss actor counter"); + + register_metric(&rivet_metrics::REGISTRY, workers.clone()); + register_metric(&rivet_metrics::REGISTRY, leases.clone()); + register_metric(&rivet_metrics::REGISTRY, available_slots.clone()); + register_metric(&rivet_metrics::REGISTRY, queued_acquires.clone()); + register_metric(&rivet_metrics::REGISTRY, acquire_duration_seconds.clone()); + register_metric(&rivet_metrics::REGISTRY, events.clone()); + register_metric(&rivet_metrics::REGISTRY, actors_failed_worker_loss.clone()); + + Self { + workers, + leases, + available_slots, + queued_acquires, + acquire_duration_seconds, + events, + actors_failed_worker_loss, + } + } +} + +fn register_metric(registry: &Registry, metric: M) +where + M: rivet_metrics::prometheus::core::Collector + Clone + Send + Sync + 'static, +{ + if let Err(error) = registry.register(Box::new(metric)) { + tracing::warn!(?error, "worker-pool metric registration failed"); + } +} + +fn worker_class_label(class: WorkerClass) -> &'static str { + match class { + WorkerClass::Baseline => "baseline", + WorkerClass::Overflow => "overflow", + } +} + +struct AcquireMetricGuard { + started: Instant, +} + +impl Drop for AcquireMetricGuard { + fn drop(&mut self) { + METRICS + .acquire_duration_seconds + .observe(self.started.elapsed().as_secs_f64()); + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq, Ord, PartialOrd)] +pub enum WorkerClass { + Baseline, + Overflow, +} + +#[derive(Clone, Debug)] +pub struct WorkerSpawnRequest { + pub worker_id: WorkerId, + pub spawn_token: String, + pub class: WorkerClass, +} + +#[derive(Clone)] +pub struct ActorWorkerPoolCallbacks { + request_spawns: Arc) -> Result<()> + Send + Sync>, + retire_worker: Arc Result<()> + Send + Sync>, +} + +impl ActorWorkerPoolCallbacks { + pub fn new( + request_spawns: impl Fn(Vec) -> Result<()> + Send + Sync + 'static, + retire_worker: impl Fn(WorkerId, WorkerRegistrationEpoch) -> Result<()> + Send + Sync + 'static, + ) -> Self { + Self { + request_spawns: Arc::new(request_spawns), + retire_worker: Arc::new(retire_worker), + } + } +} + +#[derive(Clone, Debug)] +pub struct ActorWorkerPoolConfig { + pub actors_per_thread: usize, + pub baseline_worker_limit: usize, + pub acquire_timeout: Duration, + pub idle_retire_delay: Duration, +} + +impl ActorWorkerPoolConfig { + pub fn new(actors_per_thread: usize, baseline_worker_limit: usize) -> Result { + if actors_per_thread == 0 { + return Err(WorkerPoolInvalidConfig { + reason: "actors_per_thread must be greater than zero".to_owned(), + } + .build()); + } + if baseline_worker_limit == 0 { + return Err(WorkerPoolInvalidConfig { + reason: "baseline_worker_limit must be greater than zero".to_owned(), + } + .build()); + } + Ok(Self { + actors_per_thread, + baseline_worker_limit, + acquire_timeout: DEFAULT_ACQUIRE_TIMEOUT, + idle_retire_delay: DEFAULT_IDLE_RETIRE_DELAY, + }) + } + + #[cfg(test)] + pub(crate) fn with_timeouts( + mut self, + acquire_timeout: Duration, + idle_retire_delay: Duration, + ) -> Self { + self.acquire_timeout = acquire_timeout; + self.idle_retire_delay = idle_retire_delay; + self + } +} + +#[derive(Clone, Debug, Eq, PartialEq, Ord, PartialOrd)] +struct ActorGenerationKey { + actor_id: String, + generation: u32, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum WorkerState { + Ready, + Draining, +} + +struct WorkerRecord { + id: WorkerId, + epoch: WorkerRegistrationEpoch, + class: WorkerClass, + state: WorkerState, + factories: Arc>>, + assignments: BTreeSet, + lost: CancellationToken, + created_sequence: u64, + last_selected_sequence: u64, + retirement_epoch: u64, +} + +impl WorkerRecord { + fn baseline_key(&self) -> (usize, u64, WorkerId) { + (self.assignments.len(), self.last_selected_sequence, self.id) + } + + fn overflow_key(&self) -> (Reverse, u64, WorkerId) { + ( + Reverse(self.assignments.len()), + self.created_sequence, + self.id, + ) + } +} + +struct PendingSpawn { + request: WorkerSpawnRequest, +} + +#[derive(Default)] +struct SchedulerState { + workers: BTreeMap, + assignment_owners: BTreeMap, + baseline_available: BTreeSet<(usize, u64, WorkerId)>, + overflow_available: BTreeSet<(Reverse, u64, WorkerId)>, + pending_spawns: BTreeMap, + queued_acquires: usize, + baseline_target: usize, + baseline_workers: usize, + pending_baseline: usize, + pending_overflow: usize, + ready_free_slots: usize, + next_worker_id: WorkerId, + next_sequence: u64, + next_registration_epoch: WorkerRegistrationEpoch, + spawn_failures: VecDeque, + shutting_down: bool, +} + +pub struct ActorWorkerPool { + config: ActorWorkerPoolConfig, + expected_factories: BTreeMap, + callbacks: ActorWorkerPoolCallbacks, + state: Mutex, + changed: Notify, +} + +pub struct ActorFactoryLease { + pool: Weak, + actor: ActorGenerationKey, + worker_id: WorkerId, + worker_epoch: WorkerRegistrationEpoch, + factory: Arc, + worker_lost: CancellationToken, + released: bool, +} + +impl ActorFactoryLease { + pub fn factory(&self) -> Arc { + self.factory.clone() + } + + pub fn worker_lost(&self) -> CancellationToken { + self.worker_lost.clone() + } + + pub fn worker_id(&self) -> WorkerId { + self.worker_id + } + + pub fn release(mut self) { + self.release_inner(); + } + + fn release_inner(&mut self) { + if self.released { + return; + } + self.released = true; + if let Some(pool) = self.pool.upgrade() { + pool.release_assignment(&self.actor, self.worker_id, self.worker_epoch); + } + } +} + +impl Drop for ActorFactoryLease { + fn drop(&mut self) { + self.release_inner(); + } +} + +#[derive(Clone)] +pub struct WorkerRegistrationHandle { + pool: Weak, + worker_id: WorkerId, + worker_epoch: WorkerRegistrationEpoch, +} + +impl WorkerRegistrationHandle { + pub fn worker_id(&self) -> WorkerId { + self.worker_id + } + + pub fn worker_epoch(&self) -> WorkerRegistrationEpoch { + self.worker_epoch + } + + pub fn environment_dropped(&self) { + if let Some(pool) = self.pool.upgrade() { + pool.worker_lost(self.worker_id, self.worker_epoch); + } + } + + pub fn detach(&self) { + if let Some(pool) = self.pool.upgrade() { + pool.detach_worker(self.worker_id, self.worker_epoch); + } + } +} + +struct AcquireWaiter { + pool: Weak, + active: bool, +} + +impl AcquireWaiter { + fn finish_locked(&mut self, state: &mut SchedulerState) { + if !self.active { + return; + } + self.active = false; + state.queued_acquires = state.queued_acquires.saturating_sub(1); + METRICS.queued_acquires.dec(); + state.spawn_failures.truncate(state.queued_acquires); + } + + fn finish(&mut self) { + if !self.active { + return; + } + if let Some(pool) = self.pool.upgrade() { + let mut state = pool.state.lock(); + self.finish_locked(&mut state); + drop(state); + pool.changed.notify_waiters(); + } else { + self.active = false; + } + } +} + +impl Drop for AcquireWaiter { + fn drop(&mut self) { + self.finish(); + } +} + +impl ActorWorkerPool { + pub fn new( + config: ActorWorkerPoolConfig, + expected_factories: impl IntoIterator, + callbacks: ActorWorkerPoolCallbacks, + ) -> Arc { + Arc::new(Self { + config, + expected_factories: expected_factories.into_iter().collect(), + callbacks, + state: Mutex::new(SchedulerState { + next_worker_id: 1, + next_sequence: 1, + next_registration_epoch: 1, + ..SchedulerState::default() + }), + changed: Notify::new(), + }) + } + + pub async fn acquire( + self: &Arc, + actor_id: &str, + generation: u32, + actor_name: &str, + ) -> Result { + let _metric_guard = AcquireMetricGuard { + started: Instant::now(), + }; + if !self.expected_factories.contains_key(actor_name) { + return Err(WorkerPoolActorNotRegistered { + actor_name: actor_name.to_owned(), + } + .build()); + } + + let mut waiter = { + let mut state = self.state.lock(); + if state.shutting_down { + return Err(WorkerPoolClosed.build()); + } + state.queued_acquires += 1; + METRICS.queued_acquires.inc(); + state.baseline_target = state.baseline_target.max( + (state.assignment_owners.len() + state.queued_acquires) + .min(self.config.baseline_worker_limit), + ); + AcquireWaiter { + pool: Arc::downgrade(self), + active: true, + } + }; + let deadline = Instant::now() + self.config.acquire_timeout; + let actor = ActorGenerationKey { + actor_id: actor_id.to_owned(), + generation, + }; + + loop { + let notified = self.changed.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); + + let (lease, spawn_requests, failure) = { + let mut state = self.state.lock(); + if state.shutting_down { + return Err(WorkerPoolClosed.build()); + } + if state.assignment_owners.contains_key(&actor) { + return Err(WorkerPoolDuplicateAssignment { + actor_id: actor.actor_id.clone(), + generation: actor.generation, + } + .build()); + } + + let desired_baseline = self.desired_baseline_workers(&state); + let ready_baseline = state.baseline_workers; + let should_wait_for_baseline = ready_baseline < desired_baseline; + let selected = if should_wait_for_baseline { + None + } else { + self.reserve_available_worker(&mut state, &actor, actor_name) + }; + if let Some(lease) = selected { + // Decrement the waiter in the same critical section as the + // reservation. Spawn planning can never observe both the new + // assignment and a stale queued-acquire count. + waiter.finish_locked(&mut state); + (Some(lease), Vec::new(), None) + } else { + let failure = state.spawn_failures.pop_front(); + let requests = if failure.is_none() { + self.plan_spawns(&mut state) + } else { + Vec::new() + }; + (None, requests, failure) + } + }; + + if let Some(lease) = lease { + return Ok(lease); + } + if let Some(reason) = failure { + return Err(WorkerPoolSpawnFailed { reason }.build()); + } + if !spawn_requests.is_empty() + && let Err(error) = (self.callbacks.request_spawns)(spawn_requests.clone()) + { + self.fail_spawn_requests(&spawn_requests, format!("{error:#}")); + continue; + } + + tokio::select! { + _ = notified => {} + _ = sleep_until(deadline) => { + return Err(WorkerPoolAcquireTimedOut { + actor_id: actor.actor_id.clone(), + generation: actor.generation, + }.build()); + } + } + } + } + + pub fn register_worker( + self: &Arc, + worker_id: WorkerId, + spawn_token: &str, + class: WorkerClass, + factories: HashMap>, + ) -> Result { + if let Err(error) = self.validate_factories(&factories) { + METRICS + .events + .with_label_values(&["registration_rejected", worker_class_label(class)]) + .inc(); + return Err(error); + } + let mut state = self.state.lock(); + if state.shutting_down { + return Err(WorkerPoolClosed.build()); + } + let pending = state.pending_spawns.get(&worker_id).ok_or_else(|| { + WorkerPoolRegistrationRejected { + reason: format!("worker {worker_id} has no pending spawn"), + } + .build() + })?; + if pending.request.spawn_token != spawn_token || pending.request.class != class { + return Err(WorkerPoolRegistrationRejected { + reason: format!("worker {worker_id} spawn identity did not match"), + } + .build()); + } + Self::remove_pending_spawn(&mut state, worker_id); + if state.workers.contains_key(&worker_id) { + return Err(WorkerPoolRegistrationRejected { + reason: format!("worker {worker_id} is already registered"), + } + .build()); + } + + let epoch = state.next_registration_epoch; + state.next_registration_epoch += 1; + let created_sequence = state.next_sequence; + state.next_sequence += 1; + let record = WorkerRecord { + id: worker_id, + epoch, + class, + state: WorkerState::Ready, + factories: Arc::new(factories), + assignments: BTreeSet::new(), + lost: CancellationToken::new(), + created_sequence, + last_selected_sequence: 0, + retirement_epoch: 0, + }; + self.insert_availability(&mut state, &record); + state.workers.insert(worker_id, record); + state.ready_free_slots += self.config.actors_per_thread; + if class == WorkerClass::Baseline { + state.baseline_workers += 1; + } + let should_arm_idle_retirement = class == WorkerClass::Overflow; + let class_label = worker_class_label(class); + METRICS + .workers + .with_label_values(&[class_label, "ready"]) + .inc(); + METRICS + .available_slots + .with_label_values(&[class_label]) + .add(self.config.actors_per_thread as i64); + METRICS + .events + .with_label_values(&["ready", class_label]) + .inc(); + drop(state); + self.changed.notify_waiters(); + if should_arm_idle_retirement { + self.schedule_retirement(worker_id, epoch, 0); + } + Ok(WorkerRegistrationHandle { + pool: Arc::downgrade(self), + worker_id, + worker_epoch: epoch, + }) + } + + pub fn fail_worker_spawn(&self, worker_id: WorkerId, spawn_token: &str, reason: String) { + let mut state = self.state.lock(); + let class = state + .pending_spawns + .get(&worker_id) + .filter(|pending| pending.request.spawn_token == spawn_token) + .map(|pending| pending.request.class); + let Some(class) = class else { + return; + }; + Self::remove_pending_spawn(&mut state, worker_id); + METRICS + .events + .with_label_values(&["bootstrap_failure", worker_class_label(class)]) + .inc(); + if state.spawn_failures.len() < state.queued_acquires { + state.spawn_failures.push_back(reason); + } + drop(state); + self.changed.notify_waiters(); + } + + pub fn worker_lost(&self, worker_id: WorkerId, epoch: WorkerRegistrationEpoch) { + let (worker, spawn_requests) = { + let mut state = self.state.lock(); + let matches = state + .workers + .get(&worker_id) + .is_some_and(|worker| worker.epoch == epoch); + if !matches { + return; + } + let worker = state + .workers + .remove(&worker_id) + .expect("worker checked above"); + self.remove_availability(&mut state, &worker); + let class_label = worker_class_label(worker.class); + let state_label = match worker.state { + WorkerState::Ready => "ready", + WorkerState::Draining => "draining", + }; + METRICS + .workers + .with_label_values(&[class_label, state_label]) + .dec(); + METRICS + .events + .with_label_values(&[ + if worker.state == WorkerState::Ready { + "unexpected_exit" + } else { + "retired" + }, + class_label, + ]) + .inc(); + if !worker.assignments.is_empty() { + METRICS + .leases + .with_label_values(&[class_label]) + .sub(worker.assignments.len() as i64); + METRICS + .actors_failed_worker_loss + .inc_by(worker.assignments.len() as u64); + } + if worker.state == WorkerState::Ready { + METRICS + .available_slots + .with_label_values(&[class_label]) + .sub((self.config.actors_per_thread - worker.assignments.len()) as i64); + state.ready_free_slots = state + .ready_free_slots + .saturating_sub(self.config.actors_per_thread - worker.assignments.len()); + } + if worker.class == WorkerClass::Baseline { + state.baseline_workers = state.baseline_workers.saturating_sub(1); + } + for actor in &worker.assignments { + state.assignment_owners.remove(actor); + } + let requests = self.plan_spawns(&mut state); + (worker, requests) + }; + worker.lost.cancel(); + self.changed.notify_waiters(); + if !spawn_requests.is_empty() + && let Err(error) = (self.callbacks.request_spawns)(spawn_requests.clone()) + { + self.fail_spawn_requests(&spawn_requests, format!("{error:#}")); + } + } + + pub fn shutdown(&self) { + let workers = { + let mut state = self.state.lock(); + if state.shutting_down { + return; + } + state.shutting_down = true; + for (class, count) in [ + (WorkerClass::Baseline, state.pending_baseline), + (WorkerClass::Overflow, state.pending_overflow), + ] { + METRICS + .workers + .with_label_values(&[worker_class_label(class), "starting"]) + .sub(count as i64); + } + state.pending_spawns.clear(); + state.pending_baseline = 0; + state.pending_overflow = 0; + state.spawn_failures.clear(); + state.baseline_available.clear(); + state.overflow_available.clear(); + state.ready_free_slots = 0; + state + .workers + .values_mut() + .map(|worker| { + let class_label = worker_class_label(worker.class); + if worker.state == WorkerState::Ready { + METRICS + .workers + .with_label_values(&[class_label, "ready"]) + .dec(); + METRICS + .workers + .with_label_values(&[class_label, "draining"]) + .inc(); + METRICS + .available_slots + .with_label_values(&[class_label]) + .sub((self.config.actors_per_thread - worker.assignments.len()) as i64); + } + worker.state = WorkerState::Draining; + (worker.id, worker.epoch) + }) + .collect::>() + }; + self.changed.notify_waiters(); + for (worker_id, epoch) in workers { + if let Err(error) = (self.callbacks.retire_worker)(worker_id, epoch) { + tracing::warn!( + worker_id, + epoch, + ?error, + "failed to retire worker during shutdown" + ); + self.worker_lost(worker_id, epoch); + } + } + } + + fn desired_baseline_workers(&self, state: &SchedulerState) -> usize { + state.baseline_target.max( + (state.assignment_owners.len() + state.queued_acquires) + .min(self.config.baseline_worker_limit), + ) + } + + fn reserve_available_worker( + self: &Arc, + state: &mut SchedulerState, + actor: &ActorGenerationKey, + actor_name: &str, + ) -> Option { + let worker_id = state + .baseline_available + .first() + .map(|key| key.2) + .or_else(|| state.overflow_available.first().map(|key| key.2))?; + let worker = state.workers.get_mut(&worker_id)?; + if worker.assignments.contains(actor) { + return None; + } + let old_baseline_key = worker.baseline_key(); + let old_overflow_key = worker.overflow_key(); + match worker.class { + WorkerClass::Baseline => { + state.baseline_available.remove(&old_baseline_key); + } + WorkerClass::Overflow => { + state.overflow_available.remove(&old_overflow_key); + } + } + let factory = worker + .factories + .get(actor_name) + .expect("registered workers were validated against expected factories") + .clone(); + let inserted = worker.assignments.insert(actor.clone()); + debug_assert!(inserted, "actor assignment was checked before insertion"); + state + .assignment_owners + .insert(actor.clone(), (worker_id, worker.epoch)); + state.ready_free_slots = state.ready_free_slots.saturating_sub(1); + worker.retirement_epoch += 1; + worker.last_selected_sequence = state.next_sequence; + state.next_sequence += 1; + let lease = ActorFactoryLease { + pool: Arc::downgrade(self), + actor: actor.clone(), + worker_id, + worker_epoch: worker.epoch, + factory, + worker_lost: worker.lost.clone(), + released: false, + }; + let baseline_key = worker.baseline_key(); + let overflow_key = worker.overflow_key(); + let should_insert = worker.assignments.len() < self.config.actors_per_thread; + let class = worker.class; + let class_label = worker_class_label(class); + METRICS.leases.with_label_values(&[class_label]).inc(); + METRICS + .available_slots + .with_label_values(&[class_label]) + .dec(); + if should_insert { + match class { + WorkerClass::Baseline => { + state.baseline_available.insert(baseline_key); + } + WorkerClass::Overflow => { + state.overflow_available.insert(overflow_key); + } + } + } + Some(lease) + } + + fn plan_spawns(&self, state: &mut SchedulerState) -> Vec { + if state.shutting_down { + return Vec::new(); + } + let desired_baseline = self.desired_baseline_workers(state); + let current_baseline = state.baseline_workers + state.pending_baseline; + let baseline_to_spawn = desired_baseline.saturating_sub(current_baseline); + let mut requests = Vec::new(); + for _ in 0..baseline_to_spawn { + requests.push(Self::insert_pending_spawn(state, WorkerClass::Baseline)); + } + + let pending_slots = + (state.pending_baseline + state.pending_overflow) * self.config.actors_per_thread; + let uncovered = state + .queued_acquires + .saturating_sub(state.ready_free_slots + pending_slots); + let overflow_to_spawn = uncovered.div_ceil(self.config.actors_per_thread); + for _ in 0..overflow_to_spawn { + requests.push(Self::insert_pending_spawn(state, WorkerClass::Overflow)); + } + requests + } + + fn insert_pending_spawn(state: &mut SchedulerState, class: WorkerClass) -> WorkerSpawnRequest { + let worker_id = state.next_worker_id; + state.next_worker_id += 1; + let request = WorkerSpawnRequest { + worker_id, + spawn_token: Uuid::new_v4().to_string(), + class, + }; + state.pending_spawns.insert( + worker_id, + PendingSpawn { + request: request.clone(), + }, + ); + match class { + WorkerClass::Baseline => state.pending_baseline += 1, + WorkerClass::Overflow => state.pending_overflow += 1, + } + let class_label = worker_class_label(class); + METRICS + .workers + .with_label_values(&[class_label, "starting"]) + .inc(); + METRICS + .events + .with_label_values(&["spawn_requested", class_label]) + .inc(); + request + } + + fn remove_pending_spawn( + state: &mut SchedulerState, + worker_id: WorkerId, + ) -> Option { + let pending = state.pending_spawns.remove(&worker_id)?; + match pending.request.class { + WorkerClass::Baseline => { + state.pending_baseline = state.pending_baseline.saturating_sub(1); + } + WorkerClass::Overflow => { + state.pending_overflow = state.pending_overflow.saturating_sub(1); + } + } + METRICS + .workers + .with_label_values(&[worker_class_label(pending.request.class), "starting"]) + .dec(); + Some(pending) + } + + fn fail_spawn_requests(&self, requests: &[WorkerSpawnRequest], reason: String) { + let mut state = self.state.lock(); + for request in requests { + Self::remove_pending_spawn(&mut state, request.worker_id); + METRICS + .events + .with_label_values(&["bootstrap_failure", worker_class_label(request.class)]) + .inc(); + if state.spawn_failures.len() < state.queued_acquires { + state.spawn_failures.push_back(reason.clone()); + } + } + drop(state); + self.changed.notify_waiters(); + } + + fn validate_factories(&self, factories: &HashMap>) -> Result<()> { + if factories.len() != self.expected_factories.len() { + return Err(WorkerPoolRegistrationRejected { + reason: "actor factory names did not match the main registry".to_owned(), + } + .build()); + } + for (name, expected_fingerprint) in &self.expected_factories { + let factory = factories.get(name).ok_or_else(|| { + WorkerPoolRegistrationRejected { + reason: format!("actor factory {name:?} was missing"), + } + .build() + })?; + if factory.config().worker_pool_fingerprint() != *expected_fingerprint { + return Err(WorkerPoolRegistrationRejected { + reason: format!("actor factory {name:?} configuration did not match"), + } + .build()); + } + } + Ok(()) + } + + fn insert_availability(&self, state: &mut SchedulerState, worker: &WorkerRecord) { + if worker.state != WorkerState::Ready + || worker.assignments.len() >= self.config.actors_per_thread + { + return; + } + match worker.class { + WorkerClass::Baseline => { + state.baseline_available.insert(worker.baseline_key()); + } + WorkerClass::Overflow => { + state.overflow_available.insert(worker.overflow_key()); + } + } + } + + fn remove_availability(&self, state: &mut SchedulerState, worker: &WorkerRecord) { + match worker.class { + WorkerClass::Baseline => { + state.baseline_available.remove(&worker.baseline_key()); + } + WorkerClass::Overflow => { + state.overflow_available.remove(&worker.overflow_key()); + } + } + } + + fn release_assignment( + self: &Arc, + actor: &ActorGenerationKey, + worker_id: WorkerId, + epoch: WorkerRegistrationEpoch, + ) { + let retire = { + let mut state = self.state.lock(); + if state.assignment_owners.get(actor) != Some(&(worker_id, epoch)) { + return; + } + let Some(worker) = state.workers.get(&worker_id) else { + return; + }; + if worker.epoch != epoch || !worker.assignments.contains(actor) { + return; + } + let old_baseline_key = worker.baseline_key(); + let old_overflow_key = worker.overflow_key(); + let class = worker.class; + match class { + WorkerClass::Baseline => { + state.baseline_available.remove(&old_baseline_key); + } + WorkerClass::Overflow => { + state.overflow_available.remove(&old_overflow_key); + } + } + state.assignment_owners.remove(actor); + let worker = state + .workers + .get_mut(&worker_id) + .expect("worker checked above"); + worker.assignments.remove(actor); + let class_label = worker_class_label(class); + METRICS.leases.with_label_values(&[class_label]).dec(); + worker.retirement_epoch += 1; + let retirement_epoch = worker.retirement_epoch; + let should_retire = worker.class == WorkerClass::Overflow + && worker.assignments.is_empty() + && worker.state == WorkerState::Ready; + let baseline_key = worker.baseline_key(); + let overflow_key = worker.overflow_key(); + if worker.state == WorkerState::Ready { + state.ready_free_slots += 1; + METRICS + .available_slots + .with_label_values(&[class_label]) + .inc(); + match class { + WorkerClass::Baseline => { + state.baseline_available.insert(baseline_key); + } + WorkerClass::Overflow => { + state.overflow_available.insert(overflow_key); + } + } + } + if should_retire { + Some((worker_id, epoch, retirement_epoch)) + } else { + None + } + }; + self.changed.notify_waiters(); + if let Some((worker_id, epoch, retirement_epoch)) = retire { + self.schedule_retirement(worker_id, epoch, retirement_epoch); + } + } + + fn schedule_retirement( + self: &Arc, + worker_id: WorkerId, + epoch: WorkerRegistrationEpoch, + retirement_epoch: u64, + ) { + let pool = Arc::clone(self); + let deadline = Instant::now() + self.config.idle_retire_delay; + #[cfg(feature = "native-runtime")] + let future = async move { + sleep_until(deadline).await; + pool.retire_if_still_idle(worker_id, epoch, retirement_epoch); + }; + #[cfg(feature = "native-runtime")] + match tokio::runtime::Handle::try_current() { + Ok(runtime) => { + runtime.spawn(future); + } + Err(error) => { + // A lease can be released by a final synchronous owner during + // environment teardown. Shutdown retires all workers separately; + // skipping this idle timer is safer than panicking off-runtime. + tracing::warn!( + worker_id, + epoch, + ?error, + "could not arm worker idle retirement outside the async runtime", + ); + } + } + #[cfg(not(feature = "native-runtime"))] + RuntimeSpawner::spawn(async move { + sleep_until(deadline).await; + pool.retire_if_still_idle(worker_id, epoch, retirement_epoch); + }); + } + + fn retire_if_still_idle( + self: &Arc, + worker_id: WorkerId, + epoch: WorkerRegistrationEpoch, + retirement_epoch: u64, + ) { + let (should_retire, should_retry) = { + let mut state = self.state.lock(); + if state.shutting_down { + return; + } + if state.queued_acquires > 0 { + (true, true) + } else { + let Some(worker) = state.workers.get_mut(&worker_id) else { + return; + }; + if worker.epoch != epoch + || worker.retirement_epoch != retirement_epoch + || worker.class != WorkerClass::Overflow + || worker.state != WorkerState::Ready + || !worker.assignments.is_empty() + { + return; + } + let key = worker.overflow_key(); + worker.state = WorkerState::Draining; + METRICS + .workers + .with_label_values(&["overflow", "ready"]) + .dec(); + METRICS + .workers + .with_label_values(&["overflow", "draining"]) + .inc(); + METRICS + .available_slots + .with_label_values(&["overflow"]) + .sub(self.config.actors_per_thread as i64); + METRICS + .events + .with_label_values(&["retire_requested", "overflow"]) + .inc(); + state.overflow_available.remove(&key); + state.ready_free_slots = state + .ready_free_slots + .saturating_sub(self.config.actors_per_thread); + (true, false) + } + }; + if should_retry { + self.schedule_retirement(worker_id, epoch, retirement_epoch); + return; + } + if should_retire && let Err(error) = (self.callbacks.retire_worker)(worker_id, epoch) { + tracing::warn!( + worker_id, + epoch, + ?error, + "failed to request idle worker retirement" + ); + self.worker_lost(worker_id, epoch); + } + } + + fn detach_worker(&self, worker_id: WorkerId, epoch: WorkerRegistrationEpoch) { + self.worker_lost(worker_id, epoch); + } +} + +#[derive(RivetError, Serialize)] +#[error( + "actor", + "worker_pool_invalid_config", + "Invalid worker pool configuration", + "Invalid worker pool configuration: {reason}" +)] +struct WorkerPoolInvalidConfig { + reason: String, +} + +#[derive(RivetError)] +#[error("actor", "worker_pool_closed", "Worker pool is closed")] +struct WorkerPoolClosed; + +#[derive(RivetError, Serialize)] +#[error( + "actor", + "worker_pool_actor_not_registered", + "Actor is not registered", + "Actor {actor_name:?} is not registered in the worker pool" +)] +struct WorkerPoolActorNotRegistered { + actor_name: String, +} + +#[derive(RivetError, Serialize)] +#[error( + "actor", + "worker_pool_duplicate_assignment", + "Actor generation is already assigned", + "Actor {actor_id:?} generation {generation} is already assigned to a worker" +)] +struct WorkerPoolDuplicateAssignment { + actor_id: String, + generation: u32, +} + +#[derive(RivetError, Serialize)] +#[error( + "actor", + "worker_spawn_failed", + "Worker thread failed to start", + "Worker thread failed to start: {reason}" +)] +struct WorkerPoolSpawnFailed { + reason: String, +} + +#[derive(RivetError, Serialize)] +#[error( + "actor", + "worker_acquire_timed_out", + "Timed out waiting for a worker thread", + "Timed out waiting for a worker thread for actor {actor_id:?} generation {generation}" +)] +struct WorkerPoolAcquireTimedOut { + actor_id: String, + generation: u32, +} + +#[derive(RivetError, Serialize)] +#[error( + "actor", + "worker_registration_rejected", + "Worker thread registration was rejected", + "Worker thread registration was rejected: {reason}" +)] +struct WorkerPoolRegistrationRejected { + reason: String, +} + +#[cfg(test)] +mod tests { + use std::future; + + use tokio::sync::mpsc; + + use super::*; + use crate::ActorConfig; + + const ACTOR_NAME: &str = "counter"; + + fn actor_factories() -> HashMap> { + HashMap::from([( + ACTOR_NAME.to_owned(), + Arc::new(ActorFactory::new(ActorConfig::default(), |_start| { + Box::pin(future::pending()) + })), + )]) + } + + fn test_pool( + actors_per_thread: usize, + baseline_worker_limit: usize, + idle_retire_delay: Duration, + ) -> ( + Arc, + mpsc::UnboundedReceiver, + mpsc::UnboundedReceiver<(WorkerId, WorkerRegistrationEpoch)>, + ) { + let (spawn_tx, spawn_rx) = mpsc::unbounded_channel(); + let (retire_tx, retire_rx) = mpsc::unbounded_channel(); + let config = ActorWorkerPoolConfig::new(actors_per_thread, baseline_worker_limit) + .unwrap() + .with_timeouts(Duration::from_secs(1), idle_retire_delay); + let expected = [( + ACTOR_NAME.to_owned(), + ActorConfig::default().worker_pool_fingerprint(), + )]; + let pool = ActorWorkerPool::new( + config, + expected, + ActorWorkerPoolCallbacks::new( + move |requests| { + for request in requests { + spawn_tx.send(request)?; + } + Ok(()) + }, + move |worker_id, epoch| { + retire_tx.send((worker_id, epoch))?; + Ok(()) + }, + ), + ); + (pool, spawn_rx, retire_rx) + } + + async fn acquire_with_spawn( + pool: &Arc, + spawn_rx: &mut mpsc::UnboundedReceiver, + actor_id: &str, + generation: u32, + ) -> (ActorFactoryLease, WorkerRegistrationHandle) { + let acquire = tokio::spawn({ + let pool = pool.clone(); + let actor_id = actor_id.to_owned(); + async move { pool.acquire(&actor_id, generation, ACTOR_NAME).await } + }); + let request = spawn_rx.recv().await.expect("spawn request"); + let registration = pool + .register_worker( + request.worker_id, + &request.spawn_token, + request.class, + actor_factories(), + ) + .expect("register worker"); + let lease = acquire.await.expect("join acquire").expect("acquire"); + (lease, registration) + } + + #[tokio::test] + async fn spreads_baseline_then_bin_packs_overflow() { + let (pool, mut spawn_rx, _retire_rx) = test_pool(2, 2, Duration::from_secs(60)); + let (first, _first_registration) = + acquire_with_spawn(&pool, &mut spawn_rx, "actor-1", 1).await; + let (second, _second_registration) = + acquire_with_spawn(&pool, &mut spawn_rx, "actor-2", 1).await; + assert_ne!(first.worker_id(), second.worker_id()); + + let third = pool.acquire("actor-3", 1, ACTOR_NAME).await.unwrap(); + let fourth = pool.acquire("actor-4", 1, ACTOR_NAME).await.unwrap(); + assert_eq!(third.worker_id(), first.worker_id()); + assert_eq!(fourth.worker_id(), second.worker_id()); + + let (fifth, _overflow_registration) = + acquire_with_spawn(&pool, &mut spawn_rx, "actor-5", 1).await; + assert_ne!(fifth.worker_id(), first.worker_id()); + assert_ne!(fifth.worker_id(), second.worker_id()); + } + + #[tokio::test] + async fn actors_per_thread_is_a_hard_limit() { + let (pool, mut spawn_rx, _retire_rx) = test_pool(1, 1, Duration::from_secs(60)); + let (first, _first_registration) = + acquire_with_spawn(&pool, &mut spawn_rx, "actor-1", 1).await; + let (second, _second_registration) = + acquire_with_spawn(&pool, &mut spawn_rx, "actor-2", 1).await; + assert_ne!(first.worker_id(), second.worker_id()); + } + + #[tokio::test] + async fn concurrent_acquires_spawn_only_required_capacity() { + let (pool, mut spawn_rx, _retire_rx) = test_pool(2, 2, Duration::from_secs(60)); + let acquires = (0..5) + .map(|index| { + let pool = pool.clone(); + tokio::spawn( + async move { pool.acquire(&format!("actor-{index}"), 1, ACTOR_NAME).await }, + ) + }) + .collect::>(); + let mut registrations = Vec::new(); + let mut requests = Vec::new(); + for _ in 0..3 { + let request = spawn_rx.recv().await.expect("spawn request"); + registrations.push( + pool.register_worker( + request.worker_id, + &request.spawn_token, + request.class, + actor_factories(), + ) + .unwrap(), + ); + requests.push(request); + } + assert!(spawn_rx.try_recv().is_err()); + assert_eq!( + requests + .iter() + .filter(|request| request.class == WorkerClass::Baseline) + .count(), + 2, + ); + assert_eq!( + requests + .iter() + .filter(|request| request.class == WorkerClass::Overflow) + .count(), + 1, + ); + + let mut occupancy = BTreeMap::new(); + for acquire in acquires { + let lease = acquire.await.unwrap().unwrap(); + *occupancy.entry(lease.worker_id()).or_insert(0) += 1; + } + assert_eq!(occupancy.values().sum::(), 5); + assert!(occupancy.values().all(|count| *count <= 2)); + drop(registrations); + } + + #[tokio::test] + async fn invalid_registration_does_not_consume_spawn_token() { + let (pool, mut spawn_rx, _retire_rx) = test_pool(1, 1, Duration::from_secs(60)); + let acquire = tokio::spawn({ + let pool = pool.clone(); + async move { pool.acquire("actor-1", 1, ACTOR_NAME).await } + }); + let request = spawn_rx.recv().await.unwrap(); + assert!( + pool.register_worker( + request.worker_id, + "wrong-token", + request.class, + actor_factories(), + ) + .is_err() + ); + pool.register_worker( + request.worker_id, + &request.spawn_token, + request.class, + actor_factories(), + ) + .unwrap(); + assert!(acquire.await.unwrap().is_ok()); + } + + #[tokio::test] + async fn timed_out_acquire_keeps_late_worker_for_next_actor() { + let (spawn_tx, mut spawn_rx) = mpsc::unbounded_channel(); + let config = ActorWorkerPoolConfig::new(1, 1) + .unwrap() + .with_timeouts(Duration::from_millis(10), Duration::from_secs(60)); + let pool = ActorWorkerPool::new( + config, + [( + ACTOR_NAME.to_owned(), + ActorConfig::default().worker_pool_fingerprint(), + )], + ActorWorkerPoolCallbacks::new( + move |requests| { + for request in requests { + spawn_tx.send(request)?; + } + Ok(()) + }, + |_, _| Ok(()), + ), + ); + let acquire = tokio::spawn({ + let pool = pool.clone(); + async move { pool.acquire("actor-1", 1, ACTOR_NAME).await } + }); + let request = spawn_rx.recv().await.unwrap(); + assert!(acquire.await.unwrap().is_err()); + let _registration = pool + .register_worker( + request.worker_id, + &request.spawn_token, + request.class, + actor_factories(), + ) + .unwrap(); + let lease = pool.acquire("actor-2", 1, ACTOR_NAME).await.unwrap(); + assert_eq!(lease.worker_id(), request.worker_id); + } + + #[tokio::test] + async fn losing_worker_cancels_existing_leases() { + let (pool, mut spawn_rx, _retire_rx) = test_pool(1, 1, Duration::from_secs(60)); + let (lease, registration) = acquire_with_spawn(&pool, &mut spawn_rx, "actor-1", 1).await; + assert!(!lease.worker_lost().is_cancelled()); + registration.environment_dropped(); + assert!(lease.worker_lost().is_cancelled()); + } + + #[tokio::test] + async fn lost_baseline_worker_is_replaced_without_new_demand() { + let (pool, mut spawn_rx, _retire_rx) = test_pool(1, 1, Duration::from_secs(60)); + let (_lease, registration) = acquire_with_spawn(&pool, &mut spawn_rx, "actor-1", 1).await; + registration.environment_dropped(); + let replacement = spawn_rx.recv().await.expect("replacement spawn"); + assert_eq!(replacement.class, WorkerClass::Baseline); + assert_ne!(replacement.worker_id, registration.worker_id()); + } + + #[tokio::test] + async fn shutdown_fails_queued_acquire() { + let (pool, mut spawn_rx, _retire_rx) = test_pool(1, 1, Duration::from_secs(60)); + let acquire = tokio::spawn({ + let pool = pool.clone(); + async move { pool.acquire("actor-1", 1, ACTOR_NAME).await } + }); + let _pending = spawn_rx.recv().await.expect("spawn request"); + pool.shutdown(); + let error = match acquire.await.unwrap() { + Ok(_) => panic!("acquire unexpectedly succeeded"), + Err(error) => error, + }; + assert!(error.to_string().contains("Worker pool is closed")); + } + + #[tokio::test] + async fn one_spawn_failure_does_not_fail_every_waiter() { + let (pool, mut spawn_rx, _retire_rx) = test_pool(1, 1, Duration::from_secs(60)); + let acquires = (0..2) + .map(|index| { + let pool = pool.clone(); + tokio::spawn( + async move { pool.acquire(&format!("actor-{index}"), 1, ACTOR_NAME).await }, + ) + }) + .collect::>(); + let first = spawn_rx.recv().await.unwrap(); + let second = spawn_rx.recv().await.unwrap(); + pool.fail_worker_spawn(first.worker_id, &first.spawn_token, "boom".to_owned()); + let _second_registration = pool + .register_worker( + second.worker_id, + &second.spawn_token, + second.class, + actor_factories(), + ) + .unwrap(); + let replacement = spawn_rx.recv().await.unwrap(); + let _replacement_registration = pool + .register_worker( + replacement.worker_id, + &replacement.spawn_token, + replacement.class, + actor_factories(), + ) + .unwrap(); + + let mut results = Vec::new(); + for acquire in acquires { + results.push(acquire.await.unwrap()); + } + assert_eq!(results.iter().filter(|result| result.is_err()).count(), 1); + assert_eq!(results.iter().filter(|result| result.is_ok()).count(), 1); + } + + #[tokio::test] + async fn empty_overflow_worker_retires_after_idle_delay() { + let (pool, mut spawn_rx, mut retire_rx) = test_pool(1, 1, Duration::from_millis(10)); + let (_baseline, _baseline_registration) = + acquire_with_spawn(&pool, &mut spawn_rx, "actor-1", 1).await; + let (overflow, overflow_registration) = + acquire_with_spawn(&pool, &mut spawn_rx, "actor-2", 1).await; + let expected = (overflow.worker_id(), overflow_registration.worker_epoch()); + overflow.release(); + assert_eq!(retire_rx.recv().await, Some(expected)); + } + + #[tokio::test] + async fn baseline_worker_never_retires_from_idleness() { + let (pool, mut spawn_rx, mut retire_rx) = test_pool(1, 1, Duration::from_millis(10)); + let (baseline, _registration) = + acquire_with_spawn(&pool, &mut spawn_rx, "actor-1", 1).await; + baseline.release(); + assert!( + tokio::time::timeout(Duration::from_millis(30), retire_rx.recv()) + .await + .is_err(), + ); + } + + #[tokio::test] + async fn stale_generation_release_cannot_free_current_generation() { + let (pool, mut spawn_rx, _retire_rx) = test_pool(1, 1, Duration::from_secs(60)); + let (first, registration) = acquire_with_spawn(&pool, &mut spawn_rx, "actor-1", 1).await; + let worker_id = first.worker_id(); + first.release(); + let second = pool.acquire("actor-1", 2, ACTOR_NAME).await.unwrap(); + assert_eq!(second.worker_id(), worker_id); + pool.release_assignment( + &ActorGenerationKey { + actor_id: "actor-1".to_owned(), + generation: 1, + }, + worker_id, + registration.worker_epoch(), + ); + assert!( + pool.state + .lock() + .assignment_owners + .contains_key(&ActorGenerationKey { + actor_id: "actor-1".to_owned(), + generation: 2, + }), + ); + } + + #[tokio::test] + async fn stale_retirement_timer_cannot_drain_reused_worker() { + let (pool, mut spawn_rx, mut retire_rx) = test_pool(1, 1, Duration::from_millis(10)); + let (_baseline, _baseline_registration) = + acquire_with_spawn(&pool, &mut spawn_rx, "actor-1", 1).await; + let (overflow, overflow_registration) = + acquire_with_spawn(&pool, &mut spawn_rx, "actor-2", 1).await; + let overflow_id = overflow.worker_id(); + overflow.release(); + let reused = pool.acquire("actor-3", 1, ACTOR_NAME).await.unwrap(); + assert_eq!(reused.worker_id(), overflow_id); + assert!( + tokio::time::timeout(Duration::from_millis(30), retire_rx.recv()) + .await + .is_err(), + ); + drop(overflow_registration); + } + + #[tokio::test] + async fn late_empty_overflow_worker_still_retires() { + let (spawn_tx, mut spawn_rx) = mpsc::unbounded_channel(); + let (retire_tx, mut retire_rx) = mpsc::unbounded_channel(); + let config = ActorWorkerPoolConfig::new(1, 1) + .unwrap() + .with_timeouts(Duration::from_millis(10), Duration::from_millis(10)); + let pool = ActorWorkerPool::new( + config, + [( + ACTOR_NAME.to_owned(), + ActorConfig::default().worker_pool_fingerprint(), + )], + ActorWorkerPoolCallbacks::new( + move |requests| { + for request in requests { + spawn_tx.send(request)?; + } + Ok(()) + }, + move |worker_id, epoch| { + retire_tx.send((worker_id, epoch))?; + Ok(()) + }, + ), + ); + let (_baseline, _baseline_registration) = + acquire_with_spawn(&pool, &mut spawn_rx, "actor-1", 1).await; + let acquire = tokio::spawn({ + let pool = pool.clone(); + async move { pool.acquire("actor-2", 1, ACTOR_NAME).await } + }); + let request = spawn_rx.recv().await.unwrap(); + assert_eq!(request.class, WorkerClass::Overflow); + assert!(acquire.await.unwrap().is_err()); + let registration = pool + .register_worker( + request.worker_id, + &request.spawn_token, + request.class, + actor_factories(), + ) + .unwrap(); + assert_eq!( + tokio::time::timeout(Duration::from_secs(1), retire_rx.recv()) + .await + .expect("late overflow worker should retire"), + Some((request.worker_id, registration.worker_epoch())), + ); + } +} diff --git a/rivetkit-rust/packages/rivetkit-core/src/serverless.rs b/rivetkit-rust/packages/rivetkit-core/src/serverless.rs index b25bab5ef4..6d7747c11a 100644 --- a/rivetkit-rust/packages/rivetkit-core/src/serverless.rs +++ b/rivetkit-rust/packages/rivetkit-core/src/serverless.rs @@ -22,8 +22,8 @@ use crate::actor::factory::ActorFactory; #[cfg(feature = "native-runtime")] use crate::development_process::DevelopmentProcessManager; use crate::registry::{ - CoreEnvoyHandle, CoreEnvoyStatus, RegistryCallbacks, RegistryDispatcher, ServeConfig, - should_manage_engine, + ActorFactoryProvider, CoreEnvoyHandle, CoreEnvoyStatus, RegistryCallbacks, RegistryDispatcher, + ServeConfig, should_manage_engine, }; use crate::runtime::RuntimeSpawner; use crate::time::{sleep, timeout}; @@ -168,8 +168,13 @@ impl CoreServerlessRuntime { anyhow::bail!("engine process spawning requires the `native-runtime` feature"); } + let actor_configs = factories + .iter() + .map(|(name, factory)| (name.clone(), factory.config().clone())) + .collect(); let dispatcher = Arc::new(RegistryDispatcher::new( - factories, + ActorFactoryProvider::Static(factories), + actor_configs, config.handle_inspector_http_in_runtime, )); let base_path = normalize_base_path(config.serverless_base_path.as_deref()); diff --git a/rivetkit-typescript/artifacts/registry-config.json b/rivetkit-typescript/artifacts/registry-config.json index 4375317507..d3306c2e9a 100644 --- a/rivetkit-typescript/artifacts/registry-config.json +++ b/rivetkit-typescript/artifacts/registry-config.json @@ -18,6 +18,12 @@ "description": "Maximum size of outgoing WebSocket messages in bytes. Default: 1048576", "type": "number" }, + "actorsPerThread": { + "description": "Hard limit on resident actor generations per Node.js worker thread, including actors that are starting, running, or stopping. Omit to run actor JavaScript on the main thread.", + "type": "integer", + "exclusiveMinimum": 0, + "maximum": 9007199254740991 + }, "noWelcome": { "description": "Disable the welcome message on startup. Default: false", "type": "boolean" @@ -192,4 +198,4 @@ ], "additionalProperties": false, "title": "RivetKit Registry Configuration" -} \ No newline at end of file +} diff --git a/rivetkit-typescript/packages/rivetkit-napi/Cargo.toml b/rivetkit-typescript/packages/rivetkit-napi/Cargo.toml index 6e3cca239f..d2ea79a780 100644 --- a/rivetkit-typescript/packages/rivetkit-napi/Cargo.toml +++ b/rivetkit-typescript/packages/rivetkit-napi/Cargo.toml @@ -26,6 +26,7 @@ tracing-stackdriver.workspace = true tracing-subscriber.workspace = true parking_lot.workspace = true scc.workspace = true +uuid.workspace = true hex.workspace = true http.workspace = true rivet-error.workspace = true diff --git a/rivetkit-typescript/packages/rivetkit-napi/index.d.ts b/rivetkit-typescript/packages/rivetkit-napi/index.d.ts index 6e980d29c6..9746267cf6 100644 --- a/rivetkit-typescript/packages/rivetkit-napi/index.d.ts +++ b/rivetkit-typescript/packages/rivetkit-napi/index.d.ts @@ -299,6 +299,10 @@ export interface JsKvEntry { key: Buffer value: Buffer } +export interface JsWorkerRegistration { + workerId: number + workerEpoch: number +} /** N-API wrapper around `rivetkit-core::ActorContext`. */ export declare class ActorContext { state(): Buffer @@ -443,6 +447,12 @@ export declare class QueueMessage { export declare class CoreRegistry { constructor() register(name: string, factory: NapiActorFactory): void + registerActorConfig(name: string, config: JsActorConfig): void + configureWorkerPool(actorsPerThread: number, baselineWorkerLimit: number, requestSpawns: (...args: any[]) => any, retireWorker: (...args: any[]) => any): string + attachWorker(poolId: string, workerId: number, spawnToken: string, workerClass: string): JsWorkerRegistration + detachWorker(): void + workerSpawnFailed(workerId: number, spawnToken: string, reason: string): void + workerExited(workerId: number, workerEpoch: number): void serve(config: JsServeConfig): Promise /** * Wait until the serverful envoy has completed its Engine registration. diff --git a/rivetkit-typescript/packages/rivetkit-napi/src/actor_factory.rs b/rivetkit-typescript/packages/rivetkit-napi/src/actor_factory.rs index d01c1c0187..e342ed4bb4 100644 --- a/rivetkit-typescript/packages/rivetkit-napi/src/actor_factory.rs +++ b/rivetkit-typescript/packages/rivetkit-napi/src/actor_factory.rs @@ -1,11 +1,12 @@ use std::collections::HashMap; use std::sync::Arc; +use std::sync::atomic::{AtomicBool, Ordering}; use std::time::Duration; use anyhow::Result; use napi::bindgen_prelude::{Buffer, Promise}; use napi::threadsafe_function::{ErrorStrategy, ThreadSafeCallContext, ThreadsafeFunction}; -use napi::{Env, JsFunction, JsObject}; +use napi::{Env, JsFunction, JsObject, Status}; use napi_derive::napi; use rivet_error::{ActorSpecifier, RivetError, RivetErrorKind}; use rivetkit_core::inspector::InspectorTabEntry; @@ -23,7 +24,55 @@ use crate::napi_actor_events::run_adapter_loop; use crate::websocket::WebSocket; use crate::{BRIDGE_RIVET_ERROR_PREFIX, NapiInvalidArgument, napi_anyhow_error}; -pub(crate) type CallbackTsfn = ThreadsafeFunction; +/// Shared liveness for every callback created in one Node environment. +/// +/// Worker cleanup closes this before Node starts finalizing TSFNs. Callback +/// wrappers intentionally share the TSFN through an `Arc` so moving work to a +/// task never calls napi-rs' panic-on-clone `ThreadsafeFunction::clone` after +/// finalization has begun. +pub(crate) struct CallbackEnvironment { + live: AtomicBool, +} + +impl CallbackEnvironment { + pub(crate) fn new() -> Self { + Self { + live: AtomicBool::new(true), + } + } + + pub(crate) fn close(&self) { + self.live.store(false, Ordering::Release); + } +} + +pub(crate) struct CallbackTsfn { + inner: Arc>, + environment: Arc, +} + +impl Clone for CallbackTsfn { + fn clone(&self) -> Self { + Self { + inner: self.inner.clone(), + environment: self.environment.clone(), + } + } +} + +impl CallbackTsfn { + pub(crate) async fn call_async( + &self, + value: napi::Result, + ) -> napi::Result { + if !self.environment.live.load(Ordering::Acquire) { + return Err(napi::Error::from_status(Status::Closing)); + } + // `call_async` checks the TSFN's own aborted flag under napi-rs' lock. + // Closing between our liveness check and this call is therefore safe. + self.inner.call_async(value).await + } +} pub(crate) trait TsfnPayloadSummary { fn payload_summary(&self) -> String; @@ -245,6 +294,7 @@ pub(crate) struct AdapterConfig { #[allow(dead_code)] pub(crate) struct CallbackBindings { + pub(crate) environment: Arc, pub(crate) create_state: Option>, pub(crate) on_create: Option>, pub(crate) create_conn_state: Option>, @@ -316,6 +366,16 @@ impl NapiActorFactory { pub(crate) fn actor_factory(&self) -> Arc { Arc::clone(&self.inner) } + + pub(crate) fn callback_environment(&self) -> Arc { + self._bindings.environment.clone() + } +} + +pub(crate) fn actor_config_from_js(config: JsActorConfig) -> napi::Result { + let config = ActorConfig::from_input(ActorConfigInput::from(config)); + config.validate().map_err(napi_anyhow_error)?; + Ok(config) } #[napi] @@ -334,10 +394,9 @@ impl NapiActorFactory { let adapter_config = Arc::new(AdapterConfig::from_js_config(&js_config)); let adapter_bindings = Arc::clone(&bindings); let loop_config = Arc::clone(&adapter_config); - let actor_config = ActorConfig::from_input(ActorConfigInput::from(js_config)); // Reject malformed config (empty ids/labels, duplicate ids, custom // tabs colliding with built-in ids, etc.) before the actor starts. - actor_config.validate().map_err(napi_anyhow_error)?; + let actor_config = actor_config_from_js(js_config)?; let inner = Arc::new( CoreActorFactory::new_with_manual_startup_ready(actor_config, move |start| { let bindings = Arc::clone(&adapter_bindings); @@ -386,6 +445,7 @@ impl AdapterConfig { impl CallbackBindings { fn from_js(callbacks: JsObject) -> napi::Result { + let environment = Arc::new(CallbackEnvironment::new()); let actions = if let Some(actions) = callbacks.get::<_, JsObject>("actions")? { let mut mapped = HashMap::new(); for name in JsObject::keys(&actions)? { @@ -398,7 +458,10 @@ impl CallbackBindings { .build(), ) })?; - mapped.insert(name, create_tsfn(callback, build_action_payload)?); + mapped.insert( + name, + create_tsfn(callback, build_action_payload, environment.clone())?, + ); } mapped } else { @@ -406,68 +469,124 @@ impl CallbackBindings { }; Ok(Self { - create_state: optional_tsfn(&callbacks, "createState", build_create_state_payload)?, - on_create: optional_tsfn(&callbacks, "onCreate", build_create_state_payload)?, + environment: environment.clone(), + create_state: optional_tsfn( + &callbacks, + "createState", + build_create_state_payload, + &environment, + )?, + on_create: optional_tsfn( + &callbacks, + "onCreate", + build_create_state_payload, + &environment, + )?, create_conn_state: optional_tsfn( &callbacks, "createConnState", build_create_conn_state_payload, + &environment, )?, - create_vars: optional_tsfn(&callbacks, "createVars", build_lifecycle_payload)?, - on_migrate: optional_tsfn(&callbacks, "onMigrate", build_migrate_payload)?, - on_wake: optional_tsfn(&callbacks, "onWake", build_lifecycle_payload)?, + create_vars: optional_tsfn( + &callbacks, + "createVars", + build_lifecycle_payload, + &environment, + )?, + on_migrate: optional_tsfn( + &callbacks, + "onMigrate", + build_migrate_payload, + &environment, + )?, + on_wake: optional_tsfn(&callbacks, "onWake", build_lifecycle_payload, &environment)?, on_before_actor_start: optional_tsfn( &callbacks, "onBeforeActorStart", build_lifecycle_payload, + &environment, + )?, + on_sleep: optional_tsfn(&callbacks, "onSleep", build_lifecycle_payload, &environment)?, + on_destroy: optional_tsfn( + &callbacks, + "onDestroy", + build_lifecycle_payload, + &environment, )?, - on_sleep: optional_tsfn(&callbacks, "onSleep", build_lifecycle_payload)?, - on_destroy: optional_tsfn(&callbacks, "onDestroy", build_lifecycle_payload)?, on_before_connect: optional_tsfn( &callbacks, "onBeforeConnect", build_before_connect_payload, + &environment, + )?, + on_connect: optional_tsfn( + &callbacks, + "onConnect", + build_connection_payload, + &environment, )?, - on_connect: optional_tsfn(&callbacks, "onConnect", build_connection_payload)?, on_disconnect_final: optional_tsfn( &callbacks, "onDisconnectFinal", build_connection_payload, + &environment, )? .or(optional_tsfn( &callbacks, "onDisconnect", build_connection_payload, + &environment, )?), on_before_subscribe: optional_tsfn( &callbacks, "onBeforeSubscribe", build_before_subscribe_payload, + &environment, )?, actions, on_before_action_response: optional_tsfn( &callbacks, "onBeforeActionResponse", build_before_action_response_payload, + &environment, + )?, + on_request: optional_tsfn( + &callbacks, + "onRequest", + build_http_request_payload, + &environment, )?, - on_request: optional_tsfn(&callbacks, "onRequest", build_http_request_payload)?, - on_queue_send: optional_tsfn(&callbacks, "onQueueSend", build_queue_send_payload)?, - on_websocket: optional_tsfn(&callbacks, "onWebSocket", build_websocket_payload)?, - run: optional_tsfn(&callbacks, "run", build_lifecycle_payload)?, + on_queue_send: optional_tsfn( + &callbacks, + "onQueueSend", + build_queue_send_payload, + &environment, + )?, + on_websocket: optional_tsfn( + &callbacks, + "onWebSocket", + build_websocket_payload, + &environment, + )?, + run: optional_tsfn(&callbacks, "run", build_lifecycle_payload, &environment)?, get_workflow_history: optional_tsfn( &callbacks, "getWorkflowHistory", build_workflow_history_payload, + &environment, )?, replay_workflow: optional_tsfn( &callbacks, "replayWorkflow", build_workflow_replay_payload, + &environment, )?, serialize_state: optional_tsfn( &callbacks, "serializeState", build_serialize_state_payload, + &environment, )?, }) } @@ -477,6 +596,7 @@ fn optional_tsfn( callbacks: &JsObject, name: &str, build_args: F, + environment: &Arc, ) -> napi::Result>> where T: Send + 'static, @@ -485,17 +605,25 @@ where let Some(callback) = callbacks.get::<_, JsFunction>(name)? else { return Ok(None); }; - create_tsfn(callback, build_args).map(Some) + create_tsfn(callback, build_args, environment.clone()).map(Some) } -fn create_tsfn(callback: JsFunction, build_args: F) -> napi::Result> +fn create_tsfn( + callback: JsFunction, + build_args: F, + environment: Arc, +) -> napi::Result> where T: Send + 'static, F: Fn(&Env, T) -> napi::Result> + Send + Sync + 'static, { let build_args = Arc::new(build_args); - callback.create_threadsafe_function(0, move |ctx: ThreadSafeCallContext| { + let inner = callback.create_threadsafe_function(0, move |ctx: ThreadSafeCallContext| { build_args(&ctx.env, ctx.value) + })?; + Ok(CallbackTsfn { + inner: Arc::new(inner), + environment, }) } diff --git a/rivetkit-typescript/packages/rivetkit-napi/src/lib.rs b/rivetkit-typescript/packages/rivetkit-napi/src/lib.rs index 1c1c1b0a86..24f7f53bf1 100644 --- a/rivetkit-typescript/packages/rivetkit-napi/src/lib.rs +++ b/rivetkit-typescript/packages/rivetkit-napi/src/lib.rs @@ -11,6 +11,7 @@ pub mod registry; pub mod schedule; pub mod types; pub mod websocket; +mod worker_pool; use std::sync::Once; diff --git a/rivetkit-typescript/packages/rivetkit-napi/src/napi_actor_events.rs b/rivetkit-typescript/packages/rivetkit-napi/src/napi_actor_events.rs index 619f8d49b2..0c6b063204 100644 --- a/rivetkit-typescript/packages/rivetkit-napi/src/napi_actor_events.rs +++ b/rivetkit-typescript/packages/rivetkit-napi/src/napi_actor_events.rs @@ -20,6 +20,8 @@ use crate::NapiInvalidState; #[cfg(test)] use crate::actor_context::EndReason; use crate::actor_context::{ActorContext, RegisteredTask, state_deltas_from_payload}; +#[cfg(test)] +use crate::actor_factory::CallbackEnvironment; use crate::actor_factory::{ ActionPayload, AdapterConfig, BeforeActionResponsePayload, BeforeConnectPayload, BeforeSubscribePayload, CallbackBindings, ConnectionPayload, CreateConnStatePayload, @@ -41,6 +43,14 @@ struct DispatchCancelGuard { token: CancellationToken, } +struct AdapterAbortGuard { + token: CancellationToken, +} + +struct RunHandlerAbortGuard { + run_handler: RunHandlerSlot, +} + impl RunHandlerActiveGuard { fn new(ctx: CoreActorContext) -> Self { ctx.begin_run_handler(); @@ -91,6 +101,20 @@ impl Drop for DispatchCancelGuard { } } +impl Drop for AdapterAbortGuard { + fn drop(&mut self) { + self.token.cancel(); + } +} + +impl Drop for RunHandlerAbortGuard { + fn drop(&mut self) { + if let Some(handle) = self.run_handler.lock().take() { + handle.abort(); + } + } +} + static ACTION_TIMED_OUT_SCHEMA: RivetErrorSchema = RivetErrorSchema { group: "actor", code: "action_timed_out", @@ -125,6 +149,9 @@ pub(crate) async fn run_adapter_loop( let ctx = ActorContext::new(core_ctx.clone()); ctx.reset_runtime_shared_state(); let abort = CancellationToken::new(); + let _adapter_abort_guard = AdapterAbortGuard { + token: abort.clone(), + }; ctx.attach_napi_abort_token(abort.clone()); let (registered_task_tx, mut registered_task_rx) = unbounded_channel(); ctx.attach_task_sender(registered_task_tx); @@ -163,6 +190,12 @@ pub(crate) async fn run_adapter_loop( return Err(error); } }; + // Core intentionally drops this entire adapter future when its worker + // environment disappears. Keep an abort-on-drop owner so the separately + // spawned `run` task cannot detach while retaining TSFNs and actor context. + let _run_handler_abort_guard = RunHandlerAbortGuard { + run_handler: run_handler.clone(), + }; run_event_loop( &bindings, diff --git a/rivetkit-typescript/packages/rivetkit-napi/src/registry.rs b/rivetkit-typescript/packages/rivetkit-napi/src/registry.rs index c410453cc7..95c1677826 100644 --- a/rivetkit-typescript/packages/rivetkit-napi/src/registry.rs +++ b/rivetkit-typescript/packages/rivetkit-napi/src/registry.rs @@ -11,6 +11,7 @@ use rivetkit_core::{ CoreRegistry as NativeCoreRegistry, CoreServerlessRuntime, EngineSpawnMode, HTTP_BODY_STREAM_CHANNEL_CAPACITY, ServeConfig, ServerlessRequest, registry::CoreEnvoyHandle, + registry::worker_pool::{ActorWorkerPoolConfig, WorkerRegistrationHandle}, serverless::ServerlessStreamError, serverless_http::{ self, ApplicationFetch, ApplicationRequest, ApplicationResponse, ApplicationResponseBody, @@ -20,9 +21,15 @@ use rivetkit_core::{ use tokio::sync::{Mutex as TokioMutex, Notify, mpsc}; use tokio_util::sync::CancellationToken as CoreCancellationToken; -use crate::actor_factory::NapiActorFactory; +use crate::actor_factory::{ + CallbackEnvironment, JsActorConfig, NapiActorFactory, actor_config_from_js, +}; use crate::cancellation_token::CancellationToken; use crate::http::HttpResponseBodyStream; +use crate::worker_pool::{ + JsWorkerRegistration, WorkerPoolHost, create_callbacks, lookup_pool, parse_positive_usize, + parse_worker_class, parse_worker_epoch, parse_worker_id, registration_result, +}; use crate::{NapiInvalidState, napi_anyhow_error}; #[napi(object)] @@ -109,13 +116,40 @@ enum ServerlessStreamEvent { /// /// Mode A (`serve`) and Mode B (`handle_serverless_request` -> `Serverless(...)`) /// are mutually exclusive per instance: both transition out of `Registering`. +/// A worker registry instead transitions to `AttachedWorker` and can only detach +/// or shut down after its factory map is consumed by the process-global pool. /// `BuildingServerless` is a sentinel held across the `into_serverless_runtime` /// `.await` so a concurrent `shutdown()` can observe an in-flight build and /// either wait for it to settle into `Serverless(_)` (then tear it down) or /// transition directly to `ShutDown` while the build-side checks terminal /// state before installing. +#[derive(Clone)] +struct AttachedWorker { + handle: WorkerRegistrationHandle, + callback_environments: Arc>>, +} + +impl AttachedWorker { + fn close_callbacks(&self) { + for environment in self.callback_environments.iter() { + environment.close(); + } + } + + fn environment_dropped(&self) { + self.close_callbacks(); + self.handle.environment_dropped(); + } + + fn detach(&self) { + self.close_callbacks(); + self.handle.detach(); + } +} + enum RegistryState { Registering(NativeCoreRegistry), + AttachedWorker(AttachedWorker), BuildingServerless, Serving, Serverless(CoreServerlessRuntime), @@ -138,6 +172,8 @@ pub struct CoreRegistry { /// a build wait for the build to settle and then re-check the fast path /// instead of erroring with a misleading mode-conflict. build_complete: Arc, + worker_pool_host: Arc>>, + callback_environments: Arc>>>, } #[napi] @@ -156,11 +192,16 @@ impl CoreRegistry { route_package_version: Arc::new(ParkingMutex::new(None)), shutdown_token: CoreCancellationToken::new(), build_complete: Arc::new(Notify::new()), + worker_pool_host: Arc::new(ParkingMutex::new(None)), + callback_environments: Arc::new(ParkingMutex::new(Vec::new())), } } #[napi] pub fn register(&self, name: String, factory: &NapiActorFactory) -> napi::Result<()> { + if self.worker_pool_host.lock().is_some() { + return Err(registry_worker_pool_registrations_frozen_error()); + } // Registration runs on the sync N-API thread before any async work. // `try_lock` must always succeed here: no other path holds the lock at // this point. If somehow contended, surface the structured error rather @@ -172,9 +213,13 @@ impl CoreRegistry { match &mut *guard { RegistryState::Registering(registry) => { registry.register_shared(&name, factory.actor_factory()); + self.callback_environments + .lock() + .push(factory.callback_environment()); Ok(()) } RegistryState::BuildingServerless + | RegistryState::AttachedWorker(_) | RegistryState::Serving | RegistryState::Serverless(_) | RegistryState::ShuttingDown @@ -182,6 +227,157 @@ impl CoreRegistry { } } + #[napi(js_name = "registerActorConfig")] + pub fn register_actor_config(&self, name: String, config: JsActorConfig) -> napi::Result<()> { + if self.worker_pool_host.lock().is_some() { + return Err(registry_worker_pool_registrations_frozen_error()); + } + let config = actor_config_from_js(config)?; + let mut guard = self + .state + .try_lock() + .map_err(|_| registry_register_busy_error())?; + match &mut *guard { + RegistryState::Registering(registry) => { + registry.register_config(&name, config); + Ok(()) + } + _ => Err(registry_not_registering_error()), + } + } + + #[napi(js_name = "configureWorkerPool")] + pub fn configure_worker_pool( + &self, + env: Env, + actors_per_thread: f64, + baseline_worker_limit: f64, + request_spawns: JsFunction, + retire_worker: JsFunction, + ) -> napi::Result { + if self.worker_pool_host.lock().is_some() { + return Err(registry_worker_pool_already_configured_error()); + } + let actors_per_thread = parse_positive_usize(actors_per_thread, "actorsPerThread")?; + let baseline_worker_limit = + parse_positive_usize(baseline_worker_limit, "baselineWorkerLimit")?; + let config = ActorWorkerPoolConfig::new(actors_per_thread, baseline_worker_limit) + .map_err(napi_anyhow_error)?; + let callbacks = create_callbacks(&env, request_spawns, retire_worker)?; + let mut guard = self + .state + .try_lock() + .map_err(|_| registry_register_busy_error())?; + let RegistryState::Registering(registry) = &mut *guard else { + return Err(registry_not_registering_error()); + }; + let pool = registry + .enable_worker_pool(config, callbacks) + .map_err(napi_anyhow_error)?; + let pool_id = uuid::Uuid::new_v4().to_string(); + let host = WorkerPoolHost::new(pool_id.clone(), pool)?; + *self.worker_pool_host.lock() = Some(host); + Ok(pool_id) + } + + #[napi(js_name = "attachWorker")] + pub fn attach_worker( + &self, + mut env: Env, + pool_id: String, + worker_id: f64, + spawn_token: String, + worker_class: String, + ) -> napi::Result { + let pool = lookup_pool(&pool_id)?; + let worker_id = parse_worker_id(worker_id)?; + let class = parse_worker_class(&worker_class)?; + let mut guard = self + .state + .try_lock() + .map_err(|_| registry_register_busy_error())?; + let registry = match std::mem::replace(&mut *guard, RegistryState::ShutDown) { + RegistryState::Registering(registry) => registry, + other => { + *guard = other; + return Err(registry_not_registering_error()); + } + }; + let factories = registry + .into_worker_factories() + .map_err(napi_anyhow_error)?; + let handle = pool + .register_worker(worker_id, &spawn_token, class, factories) + .map_err(napi_anyhow_error)?; + // Factories created their TSFNs before this hook was registered. Node runs + // environment cleanup hooks in reverse registration order, so this removes + // the worker from scheduling and cancels its lease signals before the + // environment finalizes those callback resources. + let attached = AttachedWorker { + handle, + callback_environments: Arc::new(self.callback_environments.lock().drain(..).collect()), + }; + if let Err(error) = env.add_env_cleanup_hook(attached.clone(), |attached| { + attached.environment_dropped(); + }) { + attached.detach(); + return Err(error); + } + let result = registration_result(&attached.handle)?; + *guard = RegistryState::AttachedWorker(attached); + Ok(result) + } + + #[napi(js_name = "detachWorker")] + pub fn detach_worker(&self) -> napi::Result<()> { + let mut guard = self + .state + .try_lock() + .map_err(|_| registry_register_busy_error())?; + match std::mem::replace(&mut *guard, RegistryState::ShutDown) { + RegistryState::AttachedWorker(handle) => { + handle.detach(); + Ok(()) + } + RegistryState::ShutDown => Ok(()), + other => { + *guard = other; + Err(registry_wrong_mode_error()) + } + } + } + + #[napi(js_name = "workerSpawnFailed")] + pub fn worker_spawn_failed( + &self, + worker_id: f64, + spawn_token: String, + reason: String, + ) -> napi::Result<()> { + let worker_id = parse_worker_id(worker_id)?; + let Some(host) = self.worker_pool_host.lock().clone() else { + // A queued Worker constructor can fail after registry shutdown has + // already removed the pool host. Its pending reservation is gone. + return Ok(()); + }; + host.pool() + .fail_worker_spawn(worker_id, &spawn_token, reason); + Ok(()) + } + + #[napi(js_name = "workerExited")] + pub fn worker_exited(&self, worker_id: f64, worker_epoch: f64) -> napi::Result<()> { + let worker_id = parse_worker_id(worker_id)?; + let worker_epoch = parse_worker_epoch(worker_epoch)?; + let Some(host) = self.worker_pool_host.lock().clone() else { + // Worker exit is a fallback cleanup signal and can race the main + // registry dropping its already-shut-down pool host. + return Ok(()); + }; + host.pool().worker_lost(worker_id, worker_epoch); + Ok(()) + } + #[napi] pub async fn serve(&self, config: JsServeConfig) -> napi::Result<()> { tracing::debug!( @@ -231,6 +427,7 @@ impl CoreRegistry { *guard = RegistryState::ShutDown; } } + self.finish_worker_pool_host(false); self.serving_envoy_ready.notify_waiters(); result.map_err(napi_anyhow_error) } @@ -258,6 +455,9 @@ impl CoreRegistry { // N-API call reaches its state transition. The already-armed Notify // closes that race without polling. RegistryState::Registering(_) => {} + RegistryState::AttachedWorker(_) => { + return Err(registry_wrong_mode_error()); + } RegistryState::BuildingServerless => { return Err(registry_wrong_mode_error()); } @@ -285,16 +485,18 @@ impl CoreRegistry { // already past the state transition observes cancel promptly. self.shutdown_token.cancel(); - let (runtime, was_building) = { + let (runtime, attached_worker, was_building, shutdown_pool_now) = { let mut guard = self.state.lock().await; match std::mem::replace(&mut *guard, RegistryState::ShuttingDown) { - RegistryState::Registering(_) | RegistryState::Serving => (None, false), - RegistryState::Serverless(runtime) => (Some(runtime), false), + RegistryState::Registering(_) => (None, None, false, true), + RegistryState::AttachedWorker(handle) => (None, Some(handle), false, false), + RegistryState::Serving => (None, None, false, false), + RegistryState::Serverless(runtime) => (Some(runtime), None, false, false), RegistryState::BuildingServerless => { // An `ensure_serverless_runtime` call is mid-build. Its // post-build re-check will observe `shutdown_token` and // tear down the runtime itself before settling state. - (None, true) + (None, None, true, false) } RegistryState::ShuttingDown | RegistryState::ShutDown => { // Already in progress / done. @@ -304,6 +506,16 @@ impl CoreRegistry { } } }; + if let Some(handle) = attached_worker { + handle.detach(); + } + if shutdown_pool_now { + self.finish_worker_pool_host(true); + } else if let Some(host) = self.worker_pool_host.lock().clone() { + // Prevent any new worker environment from attaching once shutdown + // starts. Existing registration handles remain usable for cleanup. + host.unregister_directory(); + } if let Some(runtime) = runtime { runtime.shutdown().await; @@ -333,6 +545,7 @@ impl CoreRegistry { RegistryState::Serving => (self.serving_envoy.lock().clone(), None), RegistryState::Serverless(runtime) => (None, Some(runtime.clone())), RegistryState::Registering(_) + | RegistryState::AttachedWorker(_) | RegistryState::BuildingServerless | RegistryState::ShuttingDown | RegistryState::ShutDown => (None, None), @@ -358,6 +571,9 @@ impl CoreRegistry { RegistryState::Registering(_) => { return Ok(health_response(503, "not_started", &version)); } + RegistryState::AttachedWorker(_) => { + return Ok(health_response(503, "attached_worker", &version)); + } RegistryState::BuildingServerless => { return Ok(health_response(503, "starting", &version)); } @@ -602,6 +818,9 @@ impl CoreRegistry { if matches!(*guard, RegistryState::Serving) { return Err(registry_wrong_mode_error()); } + if matches!(*guard, RegistryState::AttachedWorker(_)) { + return Err(registry_wrong_mode_error()); + } if matches!(*guard, RegistryState::BuildingServerless) { // Another caller is building. Arm the notification before // dropping the lock so a completion we race against still @@ -680,6 +899,7 @@ impl CoreRegistry { // leaving `BuildingServerless` would produce. *guard = RegistryState::ShutDown; drop(guard); + self.finish_worker_pool_host(true); Err(napi_anyhow_error(error)) } }; @@ -696,6 +916,15 @@ impl CoreRegistry { .clone() .unwrap_or_else(|| env!("CARGO_PKG_VERSION").to_owned()) } + + fn finish_worker_pool_host(&self, shutdown: bool) { + if let Some(host) = self.worker_pool_host.lock().take() { + host.unregister_directory(); + if shutdown { + host.pool().shutdown(); + } + } + } } struct ApplicationFetchPayload { @@ -897,3 +1126,23 @@ fn registry_register_busy_error() -> napi::Error { .build(), ) } + +fn registry_worker_pool_already_configured_error() -> napi::Error { + napi_anyhow_error( + NapiInvalidState { + state: "core registry worker pool".to_owned(), + reason: "already configured".to_owned(), + } + .build(), + ) +} + +fn registry_worker_pool_registrations_frozen_error() -> napi::Error { + napi_anyhow_error( + NapiInvalidState { + state: "core registry worker pool".to_owned(), + reason: "actor registrations are frozen after the worker pool is configured".to_owned(), + } + .build(), + ) +} diff --git a/rivetkit-typescript/packages/rivetkit-napi/src/worker_pool.rs b/rivetkit-typescript/packages/rivetkit-napi/src/worker_pool.rs new file mode 100644 index 0000000000..cb898019d8 --- /dev/null +++ b/rivetkit-typescript/packages/rivetkit-napi/src/worker_pool.rs @@ -0,0 +1,266 @@ +use std::sync::{Arc, LazyLock, Weak}; + +use napi::threadsafe_function::{ + ErrorStrategy, ThreadSafeCallContext, ThreadsafeFunction, ThreadsafeFunctionCallMode, +}; +use napi::{Env, JsFunction, Status}; +use rivetkit_core::registry::worker_pool::{ + ActorWorkerPool, ActorWorkerPoolCallbacks, WorkerClass, WorkerId, WorkerRegistrationEpoch, + WorkerRegistrationHandle, WorkerSpawnRequest, +}; +use scc::HashMap as SccHashMap; + +use crate::{NapiInvalidArgument, napi_anyhow_error}; + +type SpawnTsfn = ThreadsafeFunction, ErrorStrategy::Fatal>; +type RetireTsfn = ThreadsafeFunction<(WorkerId, WorkerRegistrationEpoch), ErrorStrategy::Fatal>; + +static WORKER_POOLS: LazyLock>> = + LazyLock::new(SccHashMap::new); +const JAVASCRIPT_MAX_SAFE_INTEGER: u64 = 9_007_199_254_740_991; + +#[derive(Clone)] +pub(crate) struct WorkerPoolHost { + pool_id: String, + pool: Arc, +} + +impl WorkerPoolHost { + pub(crate) fn new(pool_id: String, pool: Arc) -> napi::Result { + WORKER_POOLS.retain_sync(|_, pool| pool.strong_count() > 0); + WORKER_POOLS + .insert_sync(pool_id.clone(), Arc::downgrade(&pool)) + .map_err(|_| invalid_argument("poolId", "worker pool id is already registered"))?; + Ok(Self { pool_id, pool }) + } + + pub(crate) fn pool(&self) -> &Arc { + &self.pool + } + + pub(crate) fn unregister_directory(&self) { + let should_remove = WORKER_POOLS.get_sync(&self.pool_id).is_some_and(|entry| { + entry + .get() + .upgrade() + .is_some_and(|pool| Arc::ptr_eq(&pool, &self.pool)) + }); + if should_remove { + WORKER_POOLS.remove_sync(&self.pool_id); + } + } +} + +pub(crate) fn lookup_pool(pool_id: &str) -> napi::Result> { + let pool = WORKER_POOLS + .get_sync(pool_id) + .and_then(|entry| entry.get().upgrade()); + match pool { + Some(pool) => Ok(pool), + None => { + WORKER_POOLS.remove_sync(pool_id); + Err(invalid_argument( + "poolId", + "worker pool is missing, shut down, or belongs to another process", + )) + } + } +} + +pub(crate) fn create_callbacks( + env: &Env, + request_spawns: JsFunction, + retire_worker: JsFunction, +) -> napi::Result { + let mut request_spawns = create_spawn_tsfn(request_spawns)?; + request_spawns.unref(env)?; + let mut retire_worker = create_retire_tsfn(retire_worker)?; + retire_worker.unref(env)?; + Ok(ActorWorkerPoolCallbacks::new( + move |requests| { + check_tsfn_status( + request_spawns.call(requests, ThreadsafeFunctionCallMode::NonBlocking), + ) + }, + move |worker_id, worker_epoch| { + check_tsfn_status(retire_worker.call( + (worker_id, worker_epoch), + ThreadsafeFunctionCallMode::NonBlocking, + )) + }, + )) +} + +fn create_spawn_tsfn(callback: JsFunction) -> napi::Result { + callback.create_threadsafe_function(0, |ctx: ThreadSafeCallContext>| { + let mut array = ctx.env.create_array_with_length(ctx.value.len())?; + for (index, request) in ctx.value.into_iter().enumerate() { + let mut object = ctx.env.create_object()?; + object.set( + "workerId", + u64_to_js_integer(request.worker_id, "workerId")?, + )?; + object.set("spawnToken", request.spawn_token)?; + object.set("class", worker_class_name(request.class))?; + array.set_element(index as u32, object)?; + } + Ok(vec![array.into_unknown()]) + }) +} + +fn create_retire_tsfn(callback: JsFunction) -> napi::Result { + callback.create_threadsafe_function( + 0, + |ctx: ThreadSafeCallContext<(WorkerId, WorkerRegistrationEpoch)>| { + let mut object = ctx.env.create_object()?; + object.set("workerId", u64_to_js_integer(ctx.value.0, "workerId")?)?; + object.set( + "workerEpoch", + u64_to_js_integer(ctx.value.1, "workerEpoch")?, + )?; + Ok(vec![object.into_unknown()]) + }, + ) +} + +pub(crate) fn parse_worker_class(value: &str) -> napi::Result { + match value { + "baseline" => Ok(WorkerClass::Baseline), + "overflow" => Ok(WorkerClass::Overflow), + _ => Err(invalid_argument( + "class", + "must be either \"baseline\" or \"overflow\"", + )), + } +} + +pub(crate) fn parse_worker_id(value: f64) -> napi::Result { + parse_js_safe_integer(value, "workerId") +} + +pub(crate) fn parse_worker_epoch(value: f64) -> napi::Result { + parse_js_safe_integer(value, "workerEpoch") +} + +pub(crate) fn parse_positive_usize(value: f64, argument: &str) -> napi::Result { + let value = parse_js_safe_integer(value, argument)?; + if value == 0 { + return Err(invalid_argument(argument, "must be greater than zero")); + } + value + .try_into() + .map_err(|_| invalid_argument(argument, "is too large for this platform")) +} + +pub(crate) fn registration_result( + handle: &WorkerRegistrationHandle, +) -> napi::Result { + Ok(JsWorkerRegistration { + worker_id: u64_to_js_integer(handle.worker_id(), "workerId")?, + worker_epoch: u64_to_js_integer(handle.worker_epoch(), "workerEpoch")?, + }) +} + +#[napi_derive::napi(object)] +pub struct JsWorkerRegistration { + pub worker_id: i64, + pub worker_epoch: i64, +} + +fn worker_class_name(class: WorkerClass) -> &'static str { + match class { + WorkerClass::Baseline => "baseline", + WorkerClass::Overflow => "overflow", + } +} + +fn check_tsfn_status(status: Status) -> anyhow::Result<()> { + if status == Status::Ok { + Ok(()) + } else { + anyhow::bail!("worker pool control callback is unavailable: {status:?}") + } +} + +fn parse_js_safe_integer(value: f64, argument: &str) -> napi::Result { + if !value.is_finite() + || value < 0.0 + || value.fract() != 0.0 + || value > JAVASCRIPT_MAX_SAFE_INTEGER as f64 + { + return Err(invalid_argument( + argument, + "must be a non-negative JavaScript safe integer", + )); + } + Ok(value as u64) +} + +fn u64_to_js_integer(value: u64, argument: &str) -> napi::Result { + if value > JAVASCRIPT_MAX_SAFE_INTEGER { + return Err(invalid_argument( + argument, + "exceeded JavaScript safe integer range", + )); + } + Ok(value as i64) +} + +fn invalid_argument(argument: &str, reason: &str) -> napi::Error { + napi_anyhow_error( + NapiInvalidArgument { + argument: argument.to_owned(), + reason: reason.to_owned(), + } + .build(), + ) +} + +#[cfg(test)] +mod tests { + use super::*; + use rivetkit_core::registry::worker_pool::ActorWorkerPoolConfig; + + fn test_pool() -> Arc { + ActorWorkerPool::new( + ActorWorkerPoolConfig::new(1, 1).expect("valid test config"), + [], + ActorWorkerPoolCallbacks::new(|_| Ok(()), |_, _| Ok(())), + ) + } + + #[test] + fn directory_shares_and_unregisters_pool() { + let pool = test_pool(); + let pool_id = uuid::Uuid::new_v4().to_string(); + let host = WorkerPoolHost::new(pool_id.clone(), pool.clone()).expect("register pool"); + let found = lookup_pool(&pool_id).expect("find pool"); + assert!(Arc::ptr_eq(&pool, &found)); + + host.unregister_directory(); + assert!(lookup_pool(&pool_id).is_err()); + } + + #[test] + fn worker_class_parser_rejects_unknown_values() { + assert_eq!( + parse_worker_class("baseline").unwrap(), + WorkerClass::Baseline + ); + assert_eq!( + parse_worker_class("overflow").unwrap(), + WorkerClass::Overflow + ); + assert!(parse_worker_class("other").is_err()); + } + + #[test] + fn numeric_boundaries_require_safe_integers() { + assert_eq!(parse_worker_id(42.0).unwrap(), 42); + assert!(parse_worker_id(-1.0).is_err()); + assert!(parse_worker_id(1.5).is_err()); + assert!(parse_worker_id(f64::NAN).is_err()); + assert!(parse_worker_id(JAVASCRIPT_MAX_SAFE_INTEGER as f64 + 1.0).is_err()); + assert!(parse_positive_usize(0.0, "actorsPerThread").is_err()); + } +} diff --git a/rivetkit-typescript/packages/rivetkit-napi/tests/napi_actor_events.rs b/rivetkit-typescript/packages/rivetkit-napi/tests/napi_actor_events.rs index 171d0dc2fb..c75f91909d 100644 --- a/rivetkit-typescript/packages/rivetkit-napi/tests/napi_actor_events.rs +++ b/rivetkit-typescript/packages/rivetkit-napi/tests/napi_actor_events.rs @@ -31,6 +31,7 @@ mod moved_tests { fn empty_bindings() -> CallbackBindings { CallbackBindings { + environment: StdArc::new(CallbackEnvironment::new()), create_state: None, on_create: None, create_conn_state: None, @@ -56,6 +57,34 @@ mod moved_tests { } } + #[tokio::test] + async fn dropping_run_handler_guard_aborts_owned_task() { + struct DropSignal(Option>); + + impl Drop for DropSignal { + fn drop(&mut self) { + if let Some(sender) = self.0.take() { + let _ = sender.send(()); + } + } + } + + let (dropped_tx, dropped_rx) = oneshot::channel(); + let handle = tokio::spawn(async move { + let _drop_signal = DropSignal(Some(dropped_tx)); + std::future::pending::<()>().await; + }); + tokio::task::yield_now().await; + let guard = RunHandlerAbortGuard { + run_handler: StdArc::new(Mutex::new(Some(handle))), + }; + drop(guard); + tokio::time::timeout(Duration::from_secs(1), dropped_rx) + .await + .expect("run handler should be aborted") + .expect("drop signal should be sent"); + } + fn assert_error_code(error: anyhow::Error, code: &str) { let error = RivetTransportError::extract(&error); assert_eq!(error.code(), code); diff --git a/rivetkit-typescript/packages/rivetkit/src/registry/config/index.ts b/rivetkit-typescript/packages/rivetkit/src/registry/config/index.ts index 4911cd079b..e8799ce31f 100644 --- a/rivetkit-typescript/packages/rivetkit/src/registry/config/index.ts +++ b/rivetkit-typescript/packages/rivetkit/src/registry/config/index.ts @@ -126,6 +126,15 @@ export const RegistryConfigSchema = z * */ wasm: WasmRuntimeConfigSchema.optional().default(() => ({})), + /** + * Runs actor JavaScript on Node.js worker threads, with this many live + * actors allowed on each thread. Actors remain pinned to their thread for + * the lifetime of that actor generation. + * + * @experimental + */ + actorsPerThread: z.number().int().positive().safe().optional(), + /** * @experimental * @@ -582,6 +591,15 @@ export const DocRegistryConfigSchema = z .describe( "Maximum size of outgoing WebSocket messages in bytes. Default: 1048576", ), + actorsPerThread: z + .number() + .int() + .positive() + .safe() + .optional() + .describe( + "Hard limit on resident actor generations per Node.js worker thread, including actors that are starting, running, or stopping. Omit to run actor JavaScript on the main thread.", + ), noWelcome: z .boolean() .optional() diff --git a/rivetkit-typescript/packages/rivetkit/src/registry/index.ts b/rivetkit-typescript/packages/rivetkit/src/registry/index.ts index 088f49af8e..0e271086a5 100644 --- a/rivetkit-typescript/packages/rivetkit/src/registry/index.ts +++ b/rivetkit-typescript/packages/rivetkit/src/registry/index.ts @@ -15,8 +15,12 @@ import { RegistryConfigSchema, } from "./config"; import { logger } from "./log"; -import { buildConfiguredRegistry } from "./native"; +import { attachActorWorkerRegistry, buildConfiguredRegistry } from "./native"; import { convertNativeHttpResponse } from "./native-http"; +import { + claimActorWorkerBootstrap, + setActorWorkerAttachPromise, +} from "./node-worker-pool"; import type { CoreRuntime, RuntimeApplicationFetch, @@ -154,6 +158,7 @@ export class Registry { #runtimeServePromise?: Promise; #runtimeServeLifecyclePromise?: Promise; #runtimeReadyPromise?: Promise; + #workerAttachPromise?: Promise; #runtimeServeConfiguredPromise?: ReturnType; #runtimeServerlessPromise?: ReturnType; #applicationListenerPromise?: Promise; @@ -214,6 +219,11 @@ export class Registry { */ public async handler(request: Request): Promise { const config = this.parseConfig(); + if (config.actorsPerThread !== undefined) { + throw new Error( + "actorsPerThread is only supported by persistent Node.js registry.start() deployments", + ); + } this.#printWelcome(config, "serverless"); if (!this.#runtimeServerlessPromise) { @@ -439,6 +449,14 @@ export class Registry { const port = opts.port ?? parsePortEnv(process.env.RIVET_PORT) ?? 3000; const publicDir = opts.publicDir ?? getRivetkitPublicDir(); const config = this.parseConfig(); + if ( + config.actorsPerThread !== undefined && + (!opts.application || getRivetkitRuntimeMode() === "serverless") + ) { + throw new Error( + "actorsPerThread is only supported by persistent Node.js registry.start() deployments", + ); + } if (opts.application && getRivetkitRuntimeMode() !== "serverless") { if (config.runtime === "wasm") { @@ -446,13 +464,16 @@ export class Registry { "registry.listen() requires the native runtime; use an application-owned HTTP server with WebAssembly", ); } - this.#installSignalHandlers(config); + const workerAttachPromise = this.#startEnvoy(config, false); + if (workerAttachPromise) { + await workerAttachPromise; + return; + } this.#printWelcome(config, "serverful", { port, host: opts.host, publicDir, }); - this.#startEnvoy(config, true); const readyPromise = this.startAndWait(); const configuredRegistryPromise = this.#runtimeServeConfiguredPromise; @@ -622,7 +643,18 @@ export class Registry { /** * Starts an actor envoy for standalone server deployments. */ - #startEnvoy(config: RegistryConfig, printWelcome: boolean) { + #startEnvoy( + config: RegistryConfig, + printWelcome: boolean, + ): Promise | undefined { + if (this.#workerAttachPromise) return this.#workerAttachPromise; + const bootstrap = claimActorWorkerBootstrap(); + if (bootstrap) { + const attachPromise = attachActorWorkerRegistry(config, bootstrap); + this.#workerAttachPromise = attachPromise; + setActorWorkerAttachPromise(bootstrap, attachPromise); + return attachPromise; + } if (!this.#runtimeServePromise) { const configuredRegistryPromise = this.#buildConfiguredRegistry(config); @@ -650,6 +682,7 @@ export class Registry { if (printWelcome) { this.#printWelcome(config, "serverful"); } + return undefined; } #installSignalHandlers(config: RegistryConfig): void { @@ -750,8 +783,13 @@ export class Registry { registries.push( (async () => { try { - const { runtime, registry } = await modeAPromise; - await runtime.shutdownRegistry(registry); + const { runtime, registry, workerPool } = + await modeAPromise; + try { + await runtime.shutdownRegistry(registry); + } finally { + await workerPool?.close(); + } } catch (err) { logger().warn( { err }, @@ -765,8 +803,13 @@ export class Registry { registries.push( (async () => { try { - const { runtime, registry } = await modeBPromise; - await runtime.shutdownRegistry(registry); + const { runtime, registry, workerPool } = + await modeBPromise; + try { + await runtime.shutdownRegistry(registry); + } finally { + await workerPool?.close(); + } } catch (err) { logger().warn( { error: err }, @@ -869,7 +912,11 @@ export class Registry { ), ); } - this.#startEnvoy(config, true); + const workerAttachPromise = this.#startEnvoy(config, true); + if (workerAttachPromise) { + this.#runtimeReadyPromise = workerAttachPromise; + return workerAttachPromise; + } const configuredRegistryPromise = this.#runtimeServeConfiguredPromise; const serveLifecyclePromise = this.#runtimeServeLifecyclePromise; if (!configuredRegistryPromise || !serveLifecyclePromise) { diff --git a/rivetkit-typescript/packages/rivetkit/src/registry/napi-runtime.ts b/rivetkit-typescript/packages/rivetkit/src/registry/napi-runtime.ts index 3d589d5242..4cab63bba8 100644 --- a/rivetkit-typescript/packages/rivetkit/src/registry/napi-runtime.ts +++ b/rivetkit-typescript/packages/rivetkit/src/registry/napi-runtime.ts @@ -48,6 +48,9 @@ import type { RuntimeStateDeltaPayload, RuntimeWebSocketEvent, RuntimeWorkflowKvWrite, + RuntimeWorkerRegistration, + RuntimeWorkerRetireRequest, + RuntimeWorkerSpawnRequest, SqliteTransactionHandle, WebSocketHandle, } from "./runtime"; @@ -282,6 +285,69 @@ export class NapiCoreRuntime implements CoreRuntime { asNativeRegistry(registry).register(name, asNativeFactory(factory)); } + registerActorConfig( + registry: RegistryHandle, + name: string, + config: RuntimeActorConfig, + ): void { + asNativeRegistry(registry).registerActorConfig(name, config); + } + + configureWorkerPool( + registry: RegistryHandle, + actorsPerThread: number, + baselineWorkerLimit: number, + requestSpawns: (requests: RuntimeWorkerSpawnRequest[]) => void, + retireWorker: (request: RuntimeWorkerRetireRequest) => void, + ): string { + return asNativeRegistry(registry).configureWorkerPool( + actorsPerThread, + baselineWorkerLimit, + requestSpawns, + retireWorker, + ); + } + + attachWorker( + registry: RegistryHandle, + poolId: string, + workerId: number, + spawnToken: string, + workerClass: "baseline" | "overflow", + ): RuntimeWorkerRegistration { + return asNativeRegistry(registry).attachWorker( + poolId, + workerId, + spawnToken, + workerClass, + ); + } + + detachWorker(registry: RegistryHandle): void { + asNativeRegistry(registry).detachWorker(); + } + + workerSpawnFailed( + registry: RegistryHandle, + workerId: number, + spawnToken: string, + reason: string, + ): void { + asNativeRegistry(registry).workerSpawnFailed( + workerId, + spawnToken, + reason, + ); + } + + workerExited( + registry: RegistryHandle, + workerId: number, + workerEpoch: number, + ): void { + asNativeRegistry(registry).workerExited(workerId, workerEpoch); + } + async serveRegistry( registry: RegistryHandle, config: RuntimeServeConfig, diff --git a/rivetkit-typescript/packages/rivetkit/src/registry/native.ts b/rivetkit-typescript/packages/rivetkit/src/registry/native.ts index cee301a9e2..aeaec791e5 100644 --- a/rivetkit-typescript/packages/rivetkit/src/registry/native.ts +++ b/rivetkit-typescript/packages/rivetkit/src/registry/native.ts @@ -236,6 +236,42 @@ export async function loadConfiguredRuntime( loaders: RuntimeLoaders = defaultRuntimeLoaders, ): Promise { const requested = resolveRuntimeKind(config.runtime); + if (config.actorsPerThread !== undefined) { + if (requested === "wasm") { + throw new RivetError( + "config", + "worker_threads_require_native", + "actorsPerThread requires the native Node.js runtime.", + { public: true, statusCode: 400 }, + ); + } + if (loaders.detectHost() !== "node-like") { + throw new RivetError( + "config", + "worker_threads_require_node", + "actorsPerThread is only supported in Node.js server deployments.", + { public: true, statusCode: 400 }, + ); + } + const workerGlobal = globalThis as typeof globalThis & { + Bun?: unknown; + Deno?: unknown; + process?: { versions?: { node?: string } }; + }; + if ( + workerGlobal.Bun !== undefined || + workerGlobal.Deno !== undefined || + typeof workerGlobal.process?.versions?.node !== "string" + ) { + throw new RivetError( + "config", + "worker_threads_require_node", + "actorsPerThread requires Node.js; Bun and Deno are not supported.", + { public: true, statusCode: 400 }, + ); + } + return (await loaders.loadNative()).runtime; + } if (requested === "native") { return (await loaders.loadNative()).runtime; @@ -3791,7 +3827,7 @@ function withConnContext( }); } -function buildActorConfig( +export function buildActorConfig( definition: AnyActorDefinition, registryConfig: RegistryConfig, runtimeKind: "napi" | "wasm", @@ -5505,6 +5541,8 @@ export async function buildRegistryWithRuntime( runtime: CoreRuntime; registry: RegistryHandle; serveConfig: RuntimeServeConfig; + workerPoolId?: string; + workerPool?: import("./node-worker-pool").NodeActorWorkerPool; }> { if ( config.test?.enabled && @@ -5521,26 +5559,101 @@ export async function buildRegistryWithRuntime( } const registry = runtime.createRegistry(); + let workerPoolId: string | undefined; + let workerPool: + | import("./node-worker-pool").NodeActorWorkerPool + | undefined; - for (const [name, definition] of Object.entries(config.use)) { - runtime.registerActor( + if (config.actorsPerThread !== undefined) { + if (runtime.kind !== "napi") { + throw new Error( + "actorsPerThread requires the native Node.js runtime", + ); + } + if (!runtime.registerActorConfig) { + throw new Error( + "actorsPerThread requires native actor-config registration support", + ); + } + for (const [name, definition] of Object.entries(config.use)) { + runtime.registerActorConfig( + registry, + name, + buildActorConfig(definition, config, runtime.kind), + ); + } + const { configureNodeActorWorkerPool } = await import( + "./node-worker-pool" + ); + workerPool = await configureNodeActorWorkerPool( + runtime, registry, - name, - buildNativeFactory(runtime, config, definition), + config.actorsPerThread, ); + workerPoolId = workerPool.poolId; + } else { + for (const [name, definition] of Object.entries(config.use)) { + runtime.registerActor( + registry, + name, + buildNativeFactory(runtime, config, definition), + ); + } } return { runtime, registry, serveConfig: await buildServeConfig(config, runtime.kind === "napi"), + workerPoolId, + workerPool, }; } +export async function attachActorWorkerRegistry( + config: RegistryConfig, + bootstrap: import("./node-worker-pool").ActorWorkerBootstrapState, +): Promise { + if (config.actorsPerThread === undefined) { + throw new Error( + "The worker entrypoint registry must configure actorsPerThread", + ); + } + const { runtime } = await loadNapiRuntime(); + const normalized = normalizeRuntimeConfigForKind(config, "native"); + const registry = runtime.createRegistry(); + for (const [name, definition] of Object.entries(normalized.use)) { + runtime.registerActor( + registry, + name, + buildNativeFactory(runtime, normalized, definition), + ); + } + if (!runtime.attachWorker) { + throw new Error( + "The native runtime does not support actor worker threads", + ); + } + const registration = runtime.attachWorker( + registry, + bootstrap.poolId, + bootstrap.workerId, + bootstrap.spawnToken, + bootstrap.class, + ); + bootstrap.runtime = runtime; + bootstrap.registry = registry; + bootstrap.registration = registration; + const { postActorWorkerReady } = await import("./node-worker-pool"); + await postActorWorkerReady(registration); +} + export async function buildNativeRegistry(config: RegistryConfig): Promise<{ runtime: CoreRuntime; registry: RegistryHandle; serveConfig: RuntimeServeConfig; + workerPoolId?: string; + workerPool?: import("./node-worker-pool").NodeActorWorkerPool; }> { const { runtime } = await loadNapiRuntime(); return buildRegistryWithRuntime( @@ -5553,6 +5666,8 @@ export async function buildConfiguredRegistry(config: RegistryConfig): Promise<{ runtime: CoreRuntime; registry: RegistryHandle; serveConfig: RuntimeServeConfig; + workerPoolId?: string; + workerPool?: import("./node-worker-pool").NodeActorWorkerPool; }> { const runtime = await loadConfiguredRuntime(config); return buildRegistryWithRuntime( diff --git a/rivetkit-typescript/packages/rivetkit/src/registry/node-worker-pool.ts b/rivetkit-typescript/packages/rivetkit/src/registry/node-worker-pool.ts new file mode 100644 index 0000000000..31029d7178 --- /dev/null +++ b/rivetkit-typescript/packages/rivetkit/src/registry/node-worker-pool.ts @@ -0,0 +1,431 @@ +import type { + CoreRuntime, + RegistryHandle, + RuntimeWorkerRegistration, + RuntimeWorkerRetireRequest, + RuntimeWorkerSpawnRequest, +} from "./runtime"; +import { logger } from "./log"; + +const WORKER_BOOTSTRAP_SYMBOL = Symbol.for( + "rivetkit.actorWorkerThread.bootstrap", +); + +export interface ActorWorkerBootstrapData { + poolId: string; + workerId: number; + spawnToken: string; + class: "baseline" | "overflow"; + entrypoint: string; +} + +export interface ActorWorkerBootstrapState extends ActorWorkerBootstrapData { + claimed: boolean; + attachPromise?: Promise; + runtime?: CoreRuntime; + registry?: RegistryHandle; + registration?: RuntimeWorkerRegistration; +} + +type WorkerStatusMessage = + | ({ kind: "ready" } & RuntimeWorkerRegistration) + | { kind: "bootstrapError"; reason: string } + | ({ kind: "retired" } & RuntimeWorkerRegistration); + +interface ManagedWorker { + worker: import("node:worker_threads").Worker; + request: RuntimeWorkerSpawnRequest; + registration?: RuntimeWorkerRegistration; + spawnFailureReported: boolean; + bootstrapTimeout: ReturnType; + retireFallback?: ReturnType; + pendingRetire?: RuntimeWorkerRetireRequest; + retirementRequested: boolean; + workerError?: string; + exited: Promise; + resolveExited: () => void; +} + +export interface NodeActorWorkerPool { + poolId: string; + close: () => Promise; +} + +const WORKER_BOOTSTRAP_TIMEOUT_MS = 60_000; +const WORKER_RETIRE_TIMEOUT_MS = 5_000; + +function requireWorkerRuntimeMethod(method: T | undefined, name: string): T { + if (method === undefined) { + throw new Error( + `actorsPerThread requires a native RivetKit runtime with ${name} support`, + ); + } + return method; +} + +function stringifyWorkerError(error: unknown): string { + if (error instanceof Error) return error.stack ?? error.message; + return String(error); +} + +export function getActorWorkerBootstrap(): + | ActorWorkerBootstrapState + | undefined { + return ( + globalThis as typeof globalThis & { + [WORKER_BOOTSTRAP_SYMBOL]?: ActorWorkerBootstrapState; + } + )[WORKER_BOOTSTRAP_SYMBOL]; +} + +export function claimActorWorkerBootstrap(): + | ActorWorkerBootstrapState + | undefined { + const bootstrap = getActorWorkerBootstrap(); + if (!bootstrap) return undefined; + if (bootstrap.claimed) { + throw new Error( + "A worker-thread entrypoint attempted to start more than one RivetKit registry", + ); + } + bootstrap.claimed = true; + return bootstrap; +} + +export function setActorWorkerAttachPromise( + bootstrap: ActorWorkerBootstrapState, + promise: Promise, +): void { + bootstrap.attachPromise = promise; +} + +export async function postActorWorkerReady( + registration: RuntimeWorkerRegistration, +): Promise { + const { parentPort } = await import("node:worker_threads"); + if (!parentPort) { + throw new Error("RivetKit actor worker has no parent MessagePort"); + } + parentPort.postMessage({ kind: "ready", ...registration }); +} + +function actorWorkerBootstrapSource(): string { + return ` +import { parentPort, workerData } from "node:worker_threads"; +const symbol = Symbol.for("rivetkit.actorWorkerThread.bootstrap"); +const state = { ...workerData, claimed: false }; +globalThis[symbol] = state; +try { + if (!parentPort) { + throw new Error("RivetKit actor worker has no parent MessagePort"); + } + await import(workerData.entrypoint); + if (!state.claimed || !state.attachPromise) { + throw new Error("The worker entrypoint did not start its RivetKit registry"); + } + await state.attachPromise; + parentPort.on("message", (message) => { + if (message?.kind !== "retire") return; + if (!state.runtime || !state.registry || !state.registration) { + throw new Error("RivetKit actor worker retired before registration"); + } + if ( + message.workerId !== state.registration.workerId || + message.workerEpoch !== state.registration.workerEpoch + ) { + throw new Error("RivetKit actor worker received a stale retire request"); + } + if (!state.runtime.detachWorker) { + throw new Error("The native runtime does not support worker detach"); + } + state.runtime.detachWorker(state.registry); + parentPort.postMessage({ + kind: "retired", + ...state.registration, + }); + }); +} catch (error) { + const reason = error instanceof Error ? (error.stack || error.message) : String(error); + parentPort?.postMessage({ kind: "bootstrapError", reason }); + throw error; +} +`; +} + +export async function configureNodeActorWorkerPool( + runtime: CoreRuntime, + registry: RegistryHandle, + actorsPerThread: number, +): Promise { + const configureWorkerPool = requireWorkerRuntimeMethod( + runtime.configureWorkerPool?.bind(runtime), + "worker-pool configuration", + ); + const workerSpawnFailed = requireWorkerRuntimeMethod( + runtime.workerSpawnFailed?.bind(runtime), + "worker spawn failure reporting", + ); + const workerExited = requireWorkerRuntimeMethod( + runtime.workerExited?.bind(runtime), + "worker exit reporting", + ); + const [{ Worker }, { availableParallelism }, { pathToFileURL }] = + await Promise.all([ + import("node:worker_threads"), + import("node:os"), + import("node:url"), + ]); + const entrypointPath = process.argv[1]; + if (!entrypointPath) { + throw new Error( + "actorsPerThread requires a file-based Node.js entrypoint in process.argv[1]", + ); + } + const entrypoint = pathToFileURL(entrypointPath).href; + const workers = new Map(); + const queuedWorkerIds = new Set(); + const spawnQueue: RuntimeWorkerSpawnRequest[] = []; + let spawnDrainScheduled = false; + let closing = false; + let poolId: string; + + const reportSpawnFailure = ( + managed: ManagedWorker, + reason: string, + ): void => { + if (managed.registration || managed.spawnFailureReported) return; + managed.spawnFailureReported = true; + logger().error( + { workerId: managed.request.workerId, error: reason }, + "actor worker thread failed to start", + ); + workerSpawnFailed( + registry, + managed.request.workerId, + managed.request.spawnToken, + reason, + ); + }; + + const spawnWorker = (request: RuntimeWorkerSpawnRequest): void => { + if (closing) return; + let worker: import("node:worker_threads").Worker; + try { + worker = new Worker( + new URL( + `data:text/javascript,${encodeURIComponent(actorWorkerBootstrapSource())}`, + ), + { + name: `rivetkit-actors-${request.workerId}`, + workerData: { ...request, poolId, entrypoint }, + }, + ); + } catch (error) { + workerSpawnFailed( + registry, + request.workerId, + request.spawnToken, + stringifyWorkerError(error), + ); + return; + } + let resolveExited!: () => void; + const exited = new Promise((resolve) => { + resolveExited = resolve; + }); + const managed: ManagedWorker = { + worker, + request, + spawnFailureReported: false, + retirementRequested: false, + bootstrapTimeout: setTimeout(() => { + reportSpawnFailure( + managed, + `worker did not register within ${WORKER_BOOTSTRAP_TIMEOUT_MS}ms`, + ); + void worker.terminate(); + }, WORKER_BOOTSTRAP_TIMEOUT_MS), + exited, + resolveExited, + }; + managed.bootstrapTimeout.unref?.(); + workers.set(request.workerId, managed); + worker.on("message", (message: WorkerStatusMessage) => { + if (message.kind === "ready") { + if ( + message.workerId !== request.workerId || + managed.registration + ) { + void worker.terminate(); + return; + } + managed.registration = message; + clearTimeout(managed.bootstrapTimeout); + if (managed.pendingRetire) { + const pendingRetire = managed.pendingRetire; + managed.pendingRetire = undefined; + if (managed.retireFallback) { + clearTimeout(managed.retireFallback); + managed.retireFallback = undefined; + } + requestRetirement(managed, pendingRetire); + } + } else if (message.kind === "bootstrapError") { + reportSpawnFailure(managed, message.reason); + } else if ( + managed.registration && + message.workerId === managed.registration.workerId && + message.workerEpoch === managed.registration.workerEpoch + ) { + if (managed.retireFallback) + clearTimeout(managed.retireFallback); + void worker.terminate(); + } + }); + worker.on("error", (error) => { + const reason = stringifyWorkerError(error); + if (managed.registration) { + managed.workerError = reason; + } else { + reportSpawnFailure(managed, reason); + } + }); + worker.on("exit", (code) => { + clearTimeout(managed.bootstrapTimeout); + if (managed.retireFallback) clearTimeout(managed.retireFallback); + workers.delete(request.workerId); + managed.resolveExited(); + if (managed.registration) { + if (!closing && !managed.retirementRequested) { + logger().error( + { + workerId: managed.registration.workerId, + workerEpoch: managed.registration.workerEpoch, + exitCode: code, + error: managed.workerError, + }, + "actor worker thread exited unexpectedly", + ); + } + workerExited( + registry, + managed.registration.workerId, + managed.registration.workerEpoch, + ); + } else { + reportSpawnFailure( + managed, + `worker exited before registration with code ${code}`, + ); + } + }); + }; + + const scheduleSpawnDrain = (): void => { + if (closing || spawnDrainScheduled || spawnQueue.length === 0) return; + spawnDrainScheduled = true; + setImmediate(() => { + spawnDrainScheduled = false; + const request = spawnQueue.shift(); + if (request) { + queuedWorkerIds.delete(request.workerId); + spawnWorker(request); + } + scheduleSpawnDrain(); + }); + }; + + const spawnWorkers = (requests: RuntimeWorkerSpawnRequest[]): void => { + for (const request of requests) { + if (closing) return; + if ( + workers.has(request.workerId) || + queuedWorkerIds.has(request.workerId) + ) { + workerSpawnFailed( + registry, + request.workerId, + request.spawnToken, + `RivetKit requested duplicate worker id ${request.workerId}`, + ); + continue; + } + queuedWorkerIds.add(request.workerId); + spawnQueue.push(request); + } + scheduleSpawnDrain(); + }; + + const requestRetirement = ( + managed: ManagedWorker, + request: RuntimeWorkerRetireRequest, + ): void => { + if (managed.retireFallback) return; + if (!managed.registration) { + managed.pendingRetire = request; + managed.retireFallback = setTimeout(() => { + void managed.worker.terminate(); + }, WORKER_RETIRE_TIMEOUT_MS); + managed.retireFallback.unref?.(); + return; + } else if (managed.registration.workerEpoch !== request.workerEpoch) { + return; + } + managed.retirementRequested = true; + try { + managed.worker.postMessage({ kind: "retire", ...request }); + managed.retireFallback = setTimeout(() => { + void managed.worker.terminate(); + }, WORKER_RETIRE_TIMEOUT_MS); + managed.retireFallback.unref?.(); + } catch { + void managed.worker.terminate(); + } + }; + + const retireWorker = (request: RuntimeWorkerRetireRequest): void => { + const managed = workers.get(request.workerId); + if (!managed) return; + requestRetirement(managed, request); + }; + + poolId = configureWorkerPool( + registry, + actorsPerThread, + availableParallelism(), + spawnWorkers, + retireWorker, + ); + return { + poolId, + close: async () => { + if (closing) { + await Promise.all( + [...workers.values()].map((worker) => worker.exited), + ); + return; + } + closing = true; + spawnQueue.length = 0; + queuedWorkerIds.clear(); + for (const managed of workers.values()) { + if (!managed.registration) void managed.worker.terminate(); + } + const deadline = new Promise((resolve) => { + const timeout = setTimeout(resolve, WORKER_RETIRE_TIMEOUT_MS); + timeout.unref?.(); + }); + await Promise.race([ + Promise.all( + [...workers.values()].map((worker) => worker.exited), + ), + deadline, + ]); + await Promise.all( + [...workers.values()].map((managed) => + managed.worker.terminate(), + ), + ); + }, + }; +} diff --git a/rivetkit-typescript/packages/rivetkit/src/registry/runtime.ts b/rivetkit-typescript/packages/rivetkit/src/registry/runtime.ts index 383cec1d1b..583cbc395b 100644 --- a/rivetkit-typescript/packages/rivetkit/src/registry/runtime.ts +++ b/rivetkit-typescript/packages/rivetkit/src/registry/runtime.ts @@ -321,6 +321,22 @@ export interface RuntimeActorConfig { inspectorTabs?: Array; } +export interface RuntimeWorkerSpawnRequest { + workerId: number; + spawnToken: string; + class: "baseline" | "overflow"; +} + +export interface RuntimeWorkerRetireRequest { + workerId: number; + workerEpoch: number; +} + +export interface RuntimeWorkerRegistration { + workerId: number; + workerEpoch: number; +} + export interface RuntimeInspectorTabEntry { id: string; /** Required for custom entries; omitted for built-in hides. */ @@ -446,6 +462,37 @@ export interface CoreRuntime { name: string, factory: ActorFactoryHandle, ): void; + registerActorConfig?( + registry: RegistryHandle, + name: string, + config: RuntimeActorConfig, + ): void; + configureWorkerPool?( + registry: RegistryHandle, + actorsPerThread: number, + baselineWorkerLimit: number, + requestSpawns: (requests: RuntimeWorkerSpawnRequest[]) => void, + retireWorker: (request: RuntimeWorkerRetireRequest) => void, + ): string; + attachWorker?( + registry: RegistryHandle, + poolId: string, + workerId: number, + spawnToken: string, + workerClass: "baseline" | "overflow", + ): RuntimeWorkerRegistration; + detachWorker?(registry: RegistryHandle): void; + workerSpawnFailed?( + registry: RegistryHandle, + workerId: number, + spawnToken: string, + reason: string, + ): void; + workerExited?( + registry: RegistryHandle, + workerId: number, + workerEpoch: number, + ): void; serveRegistry( registry: RegistryHandle, config: RuntimeServeConfig, diff --git a/rivetkit-typescript/packages/rivetkit/tests/napi-worker-environments.test.ts b/rivetkit-typescript/packages/rivetkit/tests/napi-worker-environments.test.ts new file mode 100644 index 0000000000..4c4a8739a9 --- /dev/null +++ b/rivetkit-typescript/packages/rivetkit/tests/napi-worker-environments.test.ts @@ -0,0 +1,45 @@ +import { Worker } from "node:worker_threads"; +import { CoreRegistry } from "../../rivetkit-napi/index.js"; +import { expect, test } from "vitest"; + +test("NAPI worker environments resolve the main environment's pool", async () => { + const registry = new CoreRegistry(); + const poolId = registry.configureWorkerPool( + 1, + 1, + () => {}, + () => {}, + ); + const addonUrl = new URL("../../rivetkit-napi/index.js", import.meta.url) + .href; + const source = ` +import { parentPort, workerData } from "node:worker_threads"; +const imported = await import(workerData.addonUrl); +const binding = imported.default ?? imported; +const registry = new binding.CoreRegistry(); +try { + registry.attachWorker(workerData.poolId, 1, "not-pending", "baseline"); + parentPort.postMessage({ unexpectedSuccess: true }); +} catch (error) { + parentPort.postMessage({ error: String(error) }); +} +`; + const worker = new Worker( + new URL(`data:text/javascript,${encodeURIComponent(source)}`), + { workerData: { addonUrl, poolId } }, + ); + + try { + const message = await new Promise<{ error?: string }>( + (resolve, reject) => { + worker.once("message", resolve); + worker.once("error", reject); + }, + ); + expect(message.error).toMatch(/has no pending spawn/); + expect(message.error).not.toMatch(/pool is missing/); + } finally { + await worker.terminate(); + await registry.shutdown(); + } +}, 10_000); diff --git a/rivetkit-typescript/packages/rivetkit/tests/runtime-selection.test.ts b/rivetkit-typescript/packages/rivetkit/tests/runtime-selection.test.ts index 56ca3d9476..a907f6f733 100644 --- a/rivetkit-typescript/packages/rivetkit/tests/runtime-selection.test.ts +++ b/rivetkit-typescript/packages/rivetkit/tests/runtime-selection.test.ts @@ -144,6 +144,61 @@ describe("runtime selection", () => { expect(nativeLoads).toBe(0); }); + test("actorsPerThread forces the native runtime without wasm fallback", async () => { + const nativeRuntime = fakeRuntime("napi"); + let wasmLoads = 0; + const runtime = await loadConfiguredRuntime( + parseConfig({ runtime: "auto", actorsPerThread: 4 }), + fakeLoaders({ + nativeRuntime, + onLoadWasm: () => { + wasmLoads += 1; + }, + }), + ); + + expect(runtime).toBe(nativeRuntime); + expect(wasmLoads).toBe(0); + }); + + test("actorsPerThread rejects the wasm runtime", async () => { + await expect( + loadConfiguredRuntime( + parseConfig({ runtime: "wasm", actorsPerThread: 1 }), + fakeLoaders({}), + ), + ).rejects.toThrow(/requires the native Node.js runtime/); + }); + + test("actorsPerThread rejects edge-like hosts", async () => { + await expect( + loadConfiguredRuntime( + parseConfig({ actorsPerThread: 1 }), + fakeLoaders({ host: "edge-like" }), + ), + ).rejects.toThrow(/only supported in Node.js/); + }); + + test.each([ + "Bun", + "Deno", + ] as const)("actorsPerThread rejects the %s Node compatibility layer", async (runtimeGlobal) => { + const config = parseConfig({ actorsPerThread: 1 }); + const globals = globalThis as typeof globalThis & + Record<"Bun" | "Deno", unknown>; + globals[runtimeGlobal] = {}; + try { + await expect( + loadConfiguredRuntime( + config, + fakeLoaders({ host: "node-like" }), + ), + ).rejects.toThrow(/Bun and Deno are not supported/); + } finally { + delete globals[runtimeGlobal]; + } + }); + test("passes explicit wasm init input to the wasm loader", async () => { const initInput = new Uint8Array([0, 1, 2]); let observedInitInput: unknown; diff --git a/rivetkit-typescript/packages/rivetkit/tests/worker-threads.test.ts b/rivetkit-typescript/packages/rivetkit/tests/worker-threads.test.ts new file mode 100644 index 0000000000..43cefe41af --- /dev/null +++ b/rivetkit-typescript/packages/rivetkit/tests/worker-threads.test.ts @@ -0,0 +1,103 @@ +import { availableParallelism } from "node:os"; +import { afterEach, describe, expect, test, vi } from "vitest"; +import { actor } from "@/actor/mod"; +import { RegistryConfigSchema } from "@/registry/config"; +import { buildRegistryWithRuntime } from "@/registry/native"; +import { + claimActorWorkerBootstrap, + setActorWorkerAttachPromise, +} from "@/registry/node-worker-pool"; +import type { CoreRuntime } from "@/registry/runtime"; + +const bootstrapSymbol = Symbol.for("rivetkit.actorWorkerThread.bootstrap"); + +afterEach(() => { + delete ( + globalThis as typeof globalThis & { + [bootstrapSymbol]?: unknown; + } + )[bootstrapSymbol]; +}); + +describe("Node actor worker threads", () => { + test("validates actorsPerThread as a positive safe integer", () => { + const input = { use: {}, startEngine: false }; + expect( + RegistryConfigSchema.parse({ ...input, actorsPerThread: 2 }) + .actorsPerThread, + ).toBe(2); + for (const actorsPerThread of [ + 0, + -1, + 1.5, + Number.MAX_SAFE_INTEGER + 1, + ]) { + expect( + RegistryConfigSchema.safeParse({ ...input, actorsPerThread }) + .success, + ).toBe(false); + } + }); + + test("registers metadata on the main registry and configures the hybrid pool", async () => { + const registerActor = vi.fn(); + const registerActorConfig = vi.fn(); + const configureWorkerPool = vi.fn(() => "pool-1"); + const runtime = { + kind: "napi", + createRegistry: () => ({ registry: true }), + registerActor, + registerActorConfig, + configureWorkerPool, + workerSpawnFailed: vi.fn(), + workerExited: vi.fn(), + } as unknown as CoreRuntime; + const definition = actor({ state: {}, actions: {} }); + const config = RegistryConfigSchema.parse({ + use: { counter: definition }, + startEngine: false, + actorsPerThread: 3, + }); + + const result = await buildRegistryWithRuntime(config, runtime); + + expect(registerActor).not.toHaveBeenCalled(); + expect(registerActorConfig).toHaveBeenCalledOnce(); + expect(registerActorConfig).toHaveBeenCalledWith( + result.registry, + "counter", + expect.objectContaining({ actions: [] }), + ); + expect(configureWorkerPool).toHaveBeenCalledWith( + result.registry, + 3, + availableParallelism(), + expect.any(Function), + expect.any(Function), + ); + expect(result.workerPoolId).toBe("pool-1"); + }); + + test("worker bootstrap can only be claimed once", () => { + const state = { + poolId: "pool-1", + workerId: 2, + spawnToken: "token", + class: "baseline" as const, + entrypoint: "file:///app.js", + claimed: false, + }; + ( + globalThis as typeof globalThis & { + [bootstrapSymbol]?: typeof state; + } + )[bootstrapSymbol] = state; + + const claimed = claimActorWorkerBootstrap(); + expect(claimed).toBe(state); + const attached = Promise.resolve(); + setActorWorkerAttachPromise(claimed!, attached); + expect(state).toMatchObject({ claimed: true, attachPromise: attached }); + expect(() => claimActorWorkerBootstrap()).toThrow(/more than one/); + }); +});