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#[cfg(test)]
324#[path = "../../tests/envoy_callbacks.rs"]
325mod tests;