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			start_services: matches!(env::var("RIVET_RUN_SERVICES").as_deref(), Ok("1")),
170			services_binary_path: env::var_os("RIVET_SERVICES_BINARY").map(PathBuf::from),
171			engine_host,
172			engine_port,
173			engine_spawn: super::EngineSpawnMode::from_env(),
174			engine_auto_download: matches!(
175				env::var("RIVETKIT_ENGINE_AUTO_DOWNLOAD").as_deref(),
176				Ok("1") | Ok("true") | Ok("TRUE") | Ok("yes") | Ok("YES")
177			),
178			handle_inspector_http_in_runtime: false,
179			serverless_base_path: None,
180			serverless_package_version: env!("CARGO_PKG_VERSION").to_owned(),
181			serverless_client_endpoint: None,
182			serverless_client_namespace: None,
183			serverless_client_token: None,
184			serverless_validate_endpoint: true,
185			serverless_max_start_payload_bytes: 1_048_576,
186		}
187	}
188}
189
190fn default_engine_endpoint(host: &str, port: u16) -> String {
191	let url_host = if host.contains(':') && !host.starts_with('[') {
192		format!("[{host}]")
193	} else {
194		host.to_owned()
195	};
196	format!("http://{url_host}:{port}")
197}
198
199impl ServeConfig {
200	pub fn from_env() -> Self {
201		let settings = ServeSettings::from_env();
202		Self {
203			version: settings.version,
204			endpoint: settings.endpoint,
205			token: settings.token,
206			namespace: settings.namespace,
207			pool_name: settings.pool_name,
208			engine_binary_path: settings.engine_binary_path,
209			start_services: settings.start_services,
210			services_binary_path: settings.services_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}