Skip to main content

rivetkit_core/registry/
envoy_callbacks.rs

1use tracing::Instrument;
2
3use super::*;
4use crate::error::ActorRuntime;
5use crate::runtime::RuntimeSpawner;
6
7impl EnvoyCallbacks for RegistryCallbacks {
8	fn on_actor_start(
9		&self,
10		handle: EnvoyHandle,
11		actor_id: String,
12		generation: u32,
13		config: protocol::ActorConfig,
14		preloaded_kv: Option<protocol::PreloadedKv>,
15	) -> EnvoyBoxFuture<anyhow::Result<()>> {
16		let dispatcher = self.dispatcher.clone();
17		let actor_name = config.name.clone();
18		let key = actor_key_from_protocol(config.key.clone());
19		let preload_persisted_actor = decode_preloaded_persisted_actor(preloaded_kv.as_ref());
20		let preloaded_kv = preloaded_kv.map(preloaded_kv_from_protocol);
21		let input = config.input.clone();
22		let factory = dispatcher.factories.get(&actor_name).cloned();
23
24		Box::pin(async move {
25			let factory = factory.ok_or_else(|| {
26				ActorRuntime::NotRegistered {
27					actor_name: actor_name.clone(),
28				}
29				.build()
30			})?;
31			let ctx = dispatcher.build_actor_context(
32				handle,
33				&actor_id,
34				generation,
35				&actor_name,
36				key,
37				factory.as_ref(),
38			);
39
40			dispatcher
41				.start_actor(StartActorRequest {
42					actor_id: actor_id.clone(),
43					generation,
44					actor_name,
45					input,
46					preload_persisted_actor: preload_persisted_actor?,
47					preloaded_kv,
48					ctx,
49				})
50				.await?;
51
52			Ok(())
53		})
54	}
55
56	fn on_actor_stop_with_completion(
57		&self,
58		_handle: EnvoyHandle,
59		actor_id: String,
60		_generation: u32,
61		reason: protocol::StopActorReason,
62		stop_handle: ActorStopHandle,
63	) -> EnvoyBoxFuture<anyhow::Result<()>> {
64		let dispatcher = self.dispatcher.clone();
65		Box::pin(async move {
66			RuntimeSpawner::spawn(
67				async move {
68					if let Err(error) = dispatcher.stop_actor(&actor_id, reason, stop_handle).await
69					{
70						tracing::error!(
71							?error,
72							"actor stop failed after asynchronous completion handoff",
73						);
74					}
75				}
76				.in_current_span(),
77			);
78			Ok(())
79		})
80	}
81
82	fn on_shutdown(&self) {}
83
84	fn fetch(
85		&self,
86		_handle: EnvoyHandle,
87		actor_id: String,
88		_gateway_id: protocol::GatewayId,
89		_request_id: protocol::RequestId,
90		request: HttpRequest,
91	) -> EnvoyBoxFuture<anyhow::Result<HttpResponse>> {
92		tracing::info!(
93			method = %request.method,
94			path = %request.path,
95			"envoy callback: fetch request"
96		);
97		let dispatcher = self.dispatcher.clone();
98		Box::pin(async move { dispatcher.handle_fetch(&actor_id, request).await })
99	}
100
101	fn websocket(
102		&self,
103		_handle: EnvoyHandle,
104		actor_id: String,
105		_gateway_id: protocol::GatewayId,
106		_request_id: protocol::RequestId,
107		_request: HttpRequest,
108		_path: String,
109		_headers: HashMap<String, String>,
110		_is_hibernatable: bool,
111		is_restoring_hibernatable: bool,
112		sender: WebSocketSender,
113	) -> EnvoyBoxFuture<anyhow::Result<WebSocketHandler>> {
114		tracing::info!(
115			path = %_path,
116			is_hibernatable = _is_hibernatable,
117			is_restoring_hibernatable,
118			"envoy callback: websocket request"
119		);
120		let dispatcher = self.dispatcher.clone();
121		Box::pin(async move {
122			dispatcher
123				.handle_websocket(
124					&actor_id,
125					&_request,
126					&_path,
127					&_headers,
128					&_gateway_id,
129					&_request_id,
130					_is_hibernatable,
131					is_restoring_hibernatable,
132					sender,
133				)
134				.await
135		})
136	}
137
138	fn can_hibernate(
139		&self,
140		actor_id: &str,
141		_gateway_id: &protocol::GatewayId,
142		_request_id: &protocol::RequestId,
143		request: &HttpRequest,
144	) -> EnvoyBoxFuture<anyhow::Result<bool>> {
145		let can_hibernate = self.dispatcher.can_hibernate(actor_id, request);
146		Box::pin(async move { Ok(can_hibernate) })
147	}
148}
149
150impl ServeSettings {
151	fn from_env() -> Self {
152		let engine_host = env::var("RIVET_RUN_ENGINE_HOST").ok();
153		let engine_port = env::var("RIVET_RUN_ENGINE_PORT")
154			.ok()
155			.and_then(|value| value.parse().ok());
156		let endpoint = env::var("RIVET_ENDPOINT").unwrap_or_else(|_| {
157			default_engine_endpoint(
158				engine_host.as_deref().unwrap_or("127.0.0.1"),
159				engine_port.unwrap_or(6420),
160			)
161		});
162
163		Self {
164			version: env::var("RIVET_ENVOY_VERSION")
165				.ok()
166				.and_then(|value| value.parse().ok())
167				.unwrap_or(1),
168			endpoint,
169			token: Some(env::var("RIVET_TOKEN").unwrap_or_else(|_| "dev".to_owned())),
170			namespace: env::var("RIVET_NAMESPACE").unwrap_or_else(|_| "default".to_owned()),
171			pool_name: env::var("RIVET_POOL_NAME").unwrap_or_else(|_| "rivetkit-rust".to_owned()),
172			engine_binary_path: env::var_os("RIVET_ENGINE_BINARY_PATH").map(PathBuf::from),
173			engine_host,
174			engine_port,
175			engine_spawn: super::EngineSpawnMode::from_env(),
176			engine_auto_download: matches!(
177				env::var("RIVETKIT_ENGINE_AUTO_DOWNLOAD").as_deref(),
178				Ok("1") | Ok("true") | Ok("TRUE") | Ok("yes") | Ok("YES")
179			),
180			handle_inspector_http_in_runtime: false,
181			serverless_base_path: None,
182			serverless_package_version: env!("CARGO_PKG_VERSION").to_owned(),
183			serverless_client_endpoint: None,
184			serverless_client_namespace: None,
185			serverless_client_token: None,
186			serverless_validate_endpoint: true,
187			serverless_max_start_payload_bytes: 1_048_576,
188		}
189	}
190}
191
192fn default_engine_endpoint(host: &str, port: u16) -> String {
193	let url_host = if host.contains(':') && !host.starts_with('[') {
194		format!("[{host}]")
195	} else {
196		host.to_owned()
197	};
198	format!("http://{url_host}:{port}")
199}
200
201impl ServeConfig {
202	pub fn from_env() -> Self {
203		let settings = ServeSettings::from_env();
204		Self {
205			version: settings.version,
206			endpoint: settings.endpoint,
207			token: settings.token,
208			namespace: settings.namespace,
209			pool_name: settings.pool_name,
210			engine_binary_path: settings.engine_binary_path,
211			engine_host: settings.engine_host,
212			engine_port: settings.engine_port,
213			engine_spawn: settings.engine_spawn,
214			engine_auto_download: settings.engine_auto_download,
215			handle_inspector_http_in_runtime: settings.handle_inspector_http_in_runtime,
216			serverless_base_path: settings.serverless_base_path,
217			serverless_package_version: settings.serverless_package_version,
218			serverless_client_endpoint: settings.serverless_client_endpoint,
219			serverless_client_namespace: settings.serverless_client_namespace,
220			serverless_client_token: settings.serverless_client_token,
221			serverless_validate_endpoint: settings.serverless_validate_endpoint,
222			serverless_max_start_payload_bytes: settings.serverless_max_start_payload_bytes,
223			serverless_cache_envoy: true,
224			..Default::default()
225		}
226	}
227}
228
229fn actor_key_from_protocol(key: Option<String>) -> ActorKey {
230	key.as_deref()
231		.map(deserialize_actor_key_from_protocol)
232		.unwrap_or_default()
233}
234
235fn deserialize_actor_key_from_protocol(key: &str) -> ActorKey {
236	const EMPTY_KEY: &str = "/";
237	const KEY_SEPARATOR: char = '/';
238
239	if key.is_empty() || key == EMPTY_KEY {
240		return Vec::new();
241	}
242
243	let mut parts = Vec::new();
244	let mut current_part = String::new();
245	let mut escaping = false;
246	let mut empty_string_marker = false;
247
248	for ch in key.chars() {
249		if escaping {
250			if ch == '0' {
251				empty_string_marker = true;
252			} else {
253				current_part.push(ch);
254			}
255			escaping = false;
256		} else if ch == '\\' {
257			escaping = true;
258		} else if ch == KEY_SEPARATOR {
259			if empty_string_marker {
260				parts.push(String::new());
261				empty_string_marker = false;
262			} else {
263				parts.push(std::mem::take(&mut current_part));
264			}
265		} else {
266			current_part.push(ch);
267		}
268	}
269
270	if escaping {
271		current_part.push('\\');
272		parts.push(current_part);
273	} else if empty_string_marker {
274		parts.push(String::new());
275	} else if !current_part.is_empty() || !parts.is_empty() {
276		parts.push(current_part);
277	}
278
279	parts.into_iter().map(ActorKeySegment::String).collect()
280}
281
282fn decode_preloaded_persisted_actor(
283	preloaded_kv: Option<&protocol::PreloadedKv>,
284) -> Result<PreloadedPersistedActor> {
285	let Some(preloaded_kv) = preloaded_kv else {
286		return Ok(PreloadedPersistedActor::NoBundle);
287	};
288	let Some(entry) = preloaded_kv
289		.entries
290		.iter()
291		.find(|entry| entry.key == PERSIST_DATA_KEY)
292	else {
293		return Ok(
294			if preloaded_kv
295				.requested_get_keys
296				.iter()
297				.any(|key| key == PERSIST_DATA_KEY)
298			{
299				PreloadedPersistedActor::BundleExistsButEmpty
300			} else {
301				PreloadedPersistedActor::NoBundle
302			},
303		);
304	};
305
306	decode_persisted_actor(&entry.value)
307		.map(PreloadedPersistedActor::Some)
308		.context("decode preloaded persisted actor")
309}
310
311fn preloaded_kv_from_protocol(preloaded_kv: protocol::PreloadedKv) -> PreloadedKv {
312	PreloadedKv::new_with_requested_get_keys(
313		preloaded_kv
314			.entries
315			.into_iter()
316			.map(|entry| (entry.key, entry.value)),
317		preloaded_kv.requested_get_keys,
318		preloaded_kv.requested_prefixes,
319	)
320}
321
322// Test shim keeps moved tests in crate-root tests/ with private-module access.
323#[cfg(test)]
324#[path = "../../tests/envoy_callbacks.rs"]
325mod tests;