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}