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}