diff --git a/rivetkit-rust/packages/rivetkit-core/src/actor/connection.rs b/rivetkit-rust/packages/rivetkit-core/src/actor/connection.rs index dc9626a394..028f6c9700 100644 --- a/rivetkit-rust/packages/rivetkit-core/src/actor/connection.rs +++ b/rivetkit-rust/packages/rivetkit-core/src/actor/connection.rs @@ -25,6 +25,7 @@ use crate::actor::persist::{ }; use crate::actor::state::RequestSaveOpts; use crate::error::ActorRuntime; +use crate::runtime::RuntimeSpawner; use crate::time::timeout; use crate::types::ConnId; @@ -143,6 +144,46 @@ pub(crate) fn decode_persisted_connection(payload: &[u8]) -> Result, +} + +impl DisconnectOnDrop { + pub(crate) fn new(conn: ConnHandle) -> Self { + Self { conn: Some(conn) } + } + + pub(crate) async fn disconnect(mut self) -> Result<()> { + match self.conn.take() { + Some(conn) => conn.disconnect(None).await, + None => Ok(()), + } + } + + pub(crate) fn disarm(mut self) { + self.conn = None; + } +} + +impl Drop for DisconnectOnDrop { + fn drop(&mut self) { + if let Some(conn) = self.conn.take() { + RuntimeSpawner::spawn(async move { + if let Err(error) = conn.disconnect(None).await { + tracing::warn!( + conn_id = conn.id(), + ?error, + "failed to disconnect a connection whose caller went away" + ); + } + }); + } + } +} + #[derive(Clone)] pub struct ConnHandle(Arc); @@ -691,10 +732,14 @@ impl ActorContext { .await?; self.insert_existing(conn.clone()); + // The caller can go away while onConnect runs. + let opening = DisconnectOnDrop::new(conn.clone()); if let Err(error) = self.emit_connection_open(&conn, request).await { + opening.disarm(); self.remove_existing(conn.id()); return Err(error); } + opening.disarm(); self.0.metrics.inc_connections_total(); self.record_connections_updated(); self.reset_sleep_timer(); diff --git a/rivetkit-rust/packages/rivetkit-core/src/registry/http.rs b/rivetkit-rust/packages/rivetkit-core/src/registry/http.rs index 8bc3643d9b..42dd5ad2b9 100644 --- a/rivetkit-rust/packages/rivetkit-core/src/registry/http.rs +++ b/rivetkit-rust/packages/rivetkit-core/src/registry/http.rs @@ -241,6 +241,7 @@ impl RegistryDispatcher { ); } }; + let request_connection = DisconnectOnDrop::new(conn.clone()); let dispatch_result = with_action_dispatch_timeout( config.action_timeout, @@ -253,7 +254,7 @@ impl RegistryDispatcher { ), ) .await; - let disconnect_result = conn.disconnect(None).await; + let disconnect_result = request_connection.disconnect().await; match dispatch_result { Ok(output) => { @@ -362,6 +363,7 @@ impl RegistryDispatcher { ); } }; + let request_connection = DisconnectOnDrop::new(conn.clone()); let incoming = crate::telemetry::IncomingInvocationContext::from_http_headers(request.headers()); @@ -391,7 +393,7 @@ impl RegistryDispatcher { } Err(error) => Err(error), }; - let disconnect_result = conn.disconnect(None).await; + let disconnect_result = request_connection.disconnect().await; match queue_result { Ok(result) => { diff --git a/rivetkit-rust/packages/rivetkit-core/src/registry/mod.rs b/rivetkit-rust/packages/rivetkit-core/src/registry/mod.rs index f63cc207e4..47b9b967d7 100644 --- a/rivetkit-rust/packages/rivetkit-core/src/registry/mod.rs +++ b/rivetkit-rust/packages/rivetkit-core/src/registry/mod.rs @@ -35,7 +35,7 @@ use vbare::OwnedVersionedData; use crate::actor::action::ActionDispatchError; use crate::actor::config::CanHibernateWebSocket; -use crate::actor::connection::{ConnHandle, HibernatableConnectionMetadata}; +use crate::actor::connection::{ConnHandle, DisconnectOnDrop, HibernatableConnectionMetadata}; use crate::actor::context::{ActorContext, InspectorAttachGuard}; use crate::actor::factory::ActorFactory; use crate::actor::kv::LegacyActorKv; @@ -69,7 +69,6 @@ mod runner_config; mod websocket; use inspector::build_actor_inspector; -use websocket::is_actor_connect_path; #[derive(Default)] pub struct CoreRegistry { @@ -1276,10 +1275,6 @@ impl RegistryDispatcher { impl RegistryDispatcher { fn can_hibernate(&self, actor_id: &str, request: &HttpRequest) -> bool { - if matches!(is_actor_connect_path(&request.path), Ok(true)) { - return true; - } - let Some(instance) = self .actor_instances .read_sync(actor_id, |_, state| state.active_instance()) diff --git a/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/actors/sleepWithSlowConnect.ts b/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/actors/sleepWithSlowConnect.ts new file mode 100644 index 0000000000..0956327bbe --- /dev/null +++ b/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/actors/sleepWithSlowConnect.ts @@ -0,0 +1,3 @@ +import { sleepWithSlowConnect } from "../sleep"; + +export default sleepWithSlowConnect; diff --git a/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/file-system-hibernation-cleanup.ts b/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/file-system-hibernation-cleanup.ts index bdf0438778..b46fc5217f 100644 --- a/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/file-system-hibernation-cleanup.ts +++ b/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/file-system-hibernation-cleanup.ts @@ -33,5 +33,6 @@ export const fileSystemHibernationCleanupActor = actor({ }, options: { sleepTimeout: 500, + canHibernateWebSocket: true, }, }); diff --git a/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/hibernation.ts b/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/hibernation.ts index f5000bd53f..b11ffdfcfc 100644 --- a/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/hibernation.ts +++ b/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/hibernation.ts @@ -75,6 +75,7 @@ export const hibernationActor = actor({ }, options: { sleepTimeout: HIBERNATION_SLEEP_TIMEOUT, + canHibernateWebSocket: true, }, }); @@ -108,5 +109,6 @@ export const hibernationSleepWindowActor = actor({ }, options: { sleepTimeout: HIBERNATION_SLEEP_TIMEOUT, + canHibernateWebSocket: true, }, }); diff --git a/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/registry-static.ts b/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/registry-static.ts index 622d752cc4..ac165dd402 100644 --- a/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/registry-static.ts +++ b/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/registry-static.ts @@ -126,6 +126,7 @@ import { sleepWithNoSleepOption, sleepWithRawHttp, sleepWithRawWebSocket, + sleepWithSlowConnect, sleepWithWaitUntilInOnWake, sleepWithWaitUntilMessage, } from "./sleep"; @@ -225,6 +226,7 @@ export const registry = setup({ // From sleep.ts sleep, sleepWithLongRpc, + sleepWithSlowConnect, sleepWithRawHttp, sleepWithRawWebSocket, sleepWithNoSleepOption, diff --git a/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/sleep-db.ts b/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/sleep-db.ts index 9dea3c743e..fbfb56f2de 100644 --- a/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/sleep-db.ts +++ b/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/sleep-db.ts @@ -403,6 +403,7 @@ export const sleepWithDbAction = actor({ }, options: { sleepTimeout: SLEEP_DB_TIMEOUT, + canHibernateWebSocket: true, }, }); diff --git a/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/sleep.ts b/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/sleep.ts index f96ec5015a..3c64538c1d 100644 --- a/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/sleep.ts +++ b/rivetkit-typescript/packages/rivetkit/fixtures/driver-test-suite/sleep.ts @@ -176,6 +176,42 @@ export const sleepWithLongRpc = actor({ }, }); +export const sleepWithSlowConnect = actor({ + state: { startCount: 0, sleepCount: 0 }, + createVars: () => ({}) as { releaseConnect?: () => void }, + onWake: (c) => { + c.state.startCount += 1; + }, + onSleep: (c) => { + c.state.sleepCount += 1; + }, + onConnect: async (c, conn) => { + if ( + !(conn.params as { holdConnect?: boolean } | undefined)?.holdConnect + ) + return; + const { promise, resolve } = promiseWithResolvers((reason) => + c.log.warn({ msg: "unhandled held connect rejection", reason }), + ); + c.vars.releaseConnect = resolve; + c.broadcast("connecting"); + await promise; + }, + actions: { + getCounts: (c) => { + return { + startCount: c.state.startCount, + sleepCount: c.state.sleepCount, + }; + }, + ping: () => "pong", + releaseConnect: (c) => c.vars.releaseConnect?.(), + }, + options: { + sleepTimeout: SLEEP_TIMEOUT, + }, +}); + export const sleepWithWaitUntilMessage = actor({ state: { startCount: 0, diff --git a/rivetkit-typescript/packages/rivetkit/src/actor/config.ts b/rivetkit-typescript/packages/rivetkit/src/actor/config.ts index d33ac26ed1..96d2d7af70 100644 --- a/rivetkit-typescript/packages/rivetkit/src/actor/config.ts +++ b/rivetkit-typescript/packages/rivetkit/src/actor/config.ts @@ -1271,9 +1271,9 @@ const GlobalActorOptionsBaseSchema = z /** Enables the experimental Actor Runtime Socket for this actor. */ enableActorRuntimeSocket: z.boolean().default(false), /** - * Can hibernate WebSockets for onWebSocket. - * - * WebSockets using actions/events are hibernatable by default. + * Can hibernate WebSockets, both client connections for actions and + * events and WebSockets for onWebSocket. A WebSocket that does not + * hibernate closes when the actor sleeps, and the client reconnects. * * @experimental **/ @@ -2325,7 +2325,7 @@ export const DocActorOptionsSchema = z .boolean() .optional() .describe( - "Whether WebSockets using onWebSocket can be hibernated. WebSockets using actions/events are hibernatable by default. Default: false", + "Whether WebSockets can be hibernated, both client connections for actions and events and WebSockets using onWebSocket. A WebSocket that does not hibernate closes when the actor sleeps. Default: false", ), }) .describe("Actor options for timeouts and behavior configuration."); diff --git a/rivetkit-typescript/packages/rivetkit/tests/driver/actor-conn-hibernation.test.ts b/rivetkit-typescript/packages/rivetkit/tests/driver/actor-conn-hibernation.test.ts index 7114dd9d8d..294f2628bd 100644 --- a/rivetkit-typescript/packages/rivetkit/tests/driver/actor-conn-hibernation.test.ts +++ b/rivetkit-typescript/packages/rivetkit/tests/driver/actor-conn-hibernation.test.ts @@ -140,6 +140,48 @@ describeDriverMatrix("Actor Conn Hibernation", (driverTestConfig) => { await hibernatingActor.dispose(); }); + test("without canHibernateWebSocket, a connection closes when its actor sleeps and the client reconnects", async (c) => { + const { client } = await setupDriverTest(c, driverTestConfig); + const connection = client.sleep + .getOrCreate(["no-hibernation"]) + .connect(); + + let openCount = 0; + connection.onOpen(() => { + openCount += 1; + }); + + // Poll until the connection handshake finishes and the async onOpen callback has fired. + await vi.waitFor( + () => { + expect(connection.isConnected).toBe(true); + expect(openCount).toBe(1); + }, + { + timeout: CONNECTION_READY_TIMEOUT_MS, + interval: 100, + }, + ); + + await connection.triggerSleep(); + // The client reconnects on its own after the sleep closes its WebSocket. + await vi.waitFor( + () => { + expect(openCount).toBe(2); + }, + { + timeout: CONNECTION_READY_TIMEOUT_MS, + interval: 100, + }, + ); + expect(await connection.getCounts()).toEqual({ + startCount: 2, + sleepCount: 1, + }); + + await connection.dispose(); + }); + test("closing connection during hibernation", async (c) => { const { client } = await setupDriverTest(c, driverTestConfig); diff --git a/rivetkit-typescript/packages/rivetkit/tests/driver/actor-sleep.test.ts b/rivetkit-typescript/packages/rivetkit/tests/driver/actor-sleep.test.ts index 5ff0ab7631..00d7828895 100644 --- a/rivetkit-typescript/packages/rivetkit/tests/driver/actor-sleep.test.ts +++ b/rivetkit-typescript/packages/rivetkit/tests/driver/actor-sleep.test.ts @@ -1070,3 +1070,82 @@ describeDriverMatrix("Actor Sleep", (driverTestConfig) => { ); }); }); + +// Gateway3 tells the actor when a caller disconnects before the response starts. +describeDriverMatrix( + "Actor Sleep with Gateway3", + (driverTestConfig) => { + test("an aborted rpc lets the actor sleep once it finishes", async (c) => { + const { client } = await setupDriverTest(c, driverTestConfig); + const key = [crypto.randomUUID()]; + + // A connection only observes that the rpc started, then leaves. + const observer = client.sleepWithLongRpc.getOrCreate(key).connect(); + const started = new Promise((resolve) => + observer.once("waiting", resolve), + ); + // The subscription and this call share one ordered connection, so the subscription is active before the rpc starts. + await observer.getCounts(); + const caller = new AbortController(); + const aborted = client.sleepWithLongRpc.getOrCreate(key).action({ + name: "longRunningRpc", + args: [], + signal: caller.signal, + }); + await started; + await observer.dispose(); + caller.abort(); + await expect(aborted).rejects.toThrow(); + + // The rpc keeps running after the abort. Once it finishes, nothing is left to keep the actor awake. + await client.sleepWithLongRpc + .getOrCreate(key) + .finishLongRunningRpc(); + await waitFor(driverTestConfig, SLEEP_TIMEOUT + 250); + + const { startCount, sleepCount } = await client.sleepWithLongRpc + .getOrCreate(key) + .getCounts(); + expect(sleepCount).toBe(1); + expect(startCount).toBe(2); + }); + + test("a caller that leaves during onConnect lets the actor sleep", async (c) => { + const { client } = await setupDriverTest(c, driverTestConfig); + const key = [crypto.randomUUID()]; + + // A connection only observes that onConnect started, then leaves. + const observer = client.sleepWithSlowConnect + .getOrCreate(key) + .connect(); + const connecting = new Promise((resolve) => + observer.once("connecting", resolve), + ); + // The subscription and this call share one ordered connection, so the subscription is active before the call below connects. + await observer.ping(); + const caller = new AbortController(); + const aborted = client.sleepWithSlowConnect + .getOrCreate(key, { params: { holdConnect: true } }) + .action({ name: "ping", args: [], signal: caller.signal }); + await connecting; + await observer.dispose(); + caller.abort(); + await expect(aborted).rejects.toThrow(); + + // onConnect finishes after the caller left. Nothing is left to keep the actor awake. + await client.sleepWithSlowConnect.getOrCreate(key).releaseConnect(); + await waitFor(driverTestConfig, SLEEP_TIMEOUT + 250); + + const { startCount, sleepCount } = await client.sleepWithSlowConnect + .getOrCreate(key) + .getCounts(); + expect(sleepCount).toBe(1); + expect(startCount).toBe(2); + }); + }, + { + runtimes: ["native"], + encodings: ["bare"], + config: { engine: { gateway3: true } }, + }, +);