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