1use crate::error::{GatewayError, GatewayResult};
14use crate::mesh::{MeshPeerRequest, MeshPeerResponse};
15use crate::state::GatewayState;
16use appcore_distributed_contracts::{PeerRpcEnvelope, PeerRpcResponse};
17use appcore_peer_rpc::payload_hash;
18use appcore_types::{ProtocolVersion, TenantId};
19use axum::extract::ws::Message;
20use std::collections::HashMap;
21use std::sync::{Arc, LazyLock, Mutex};
22use std::time::{Duration, SystemTime, UNIX_EPOCH};
23use tokio::sync::oneshot;
24
25#[derive(Debug, Clone, PartialEq, Eq, Hash)]
26struct PendingKey {
27 state_id: usize,
28 tenant_id: String,
29 request_id: String,
30 mesh: bool,
31}
32
33#[derive(Debug, Clone, Copy)]
34struct PendingMetadata {
35 worker_generation: u64,
36 response_limit: Option<usize>,
37}
38
39static PENDING_METADATA: LazyLock<Mutex<HashMap<PendingKey, PendingMetadata>>> =
41 LazyLock::new(|| Mutex::new(HashMap::new()));
42
43pub struct EnvelopeRouter;
45
46impl EnvelopeRouter {
47 pub async fn route_request(
52 state: Arc<GatewayState>,
53 envelope: PeerRpcEnvelope,
54 timeout: Duration,
55 ) -> PeerRpcResponse {
56 let request_id = envelope.request_id.clone();
57 let tenant_id = envelope.tenant_id.clone();
58 let capability = envelope.capability.clone();
59 if let Some(error) = validate_routing_envelope(&envelope) {
60 state.metrics.routing_failure();
61 return PeerRpcResponse::rejected(request_id, error);
62 }
63 let timeout = timeout.min(crate::config::MAX_GATEWAY_REQUEST_TIMEOUT);
64
65 let (rx, worker_conn) = {
67 let mut tenants = state.tenants.write();
68 let tenant_state = tenants
69 .entry(tenant_id.clone())
70 .or_insert_with(|| crate::tenant::TenantState::new(tenant_id.clone()));
71
72 let Some(worker_conn) = tenant_state
73 .get_worker_by_core(&envelope.target_core_id)
74 .filter(|worker| {
75 worker.cluster_id() == Some(&envelope.cluster_id)
76 && tenant_state
77 .registry
78 .resolve(&capability)
79 .is_some_and(|workers| workers.contains(&worker.key))
80 })
81 .cloned()
82 else {
83 state.metrics.routing_failure();
84 return PeerRpcResponse::rejected(
85 request_id,
86 format!("compatible_worker_unavailable: {}", capability.as_str()),
87 );
88 };
89 if !tenant_state.can_register_pending(&request_id, false) {
90 state.metrics.routing_failure();
91 return PeerRpcResponse::rejected(request_id, "pending_request_rejected");
92 }
93
94 let (tx, rx) = oneshot::channel();
95 tenant_state.pending_requests.insert(request_id.clone(), tx);
96 register_pending_metadata(
97 &state,
98 &tenant_id,
99 &request_id,
100 false,
101 worker_conn.generation(),
102 None,
103 );
104 (rx, worker_conn)
105 };
106 let _cleanup = PendingCleanup::new(
107 Arc::clone(&state),
108 tenant_id.clone(),
109 request_id.clone(),
110 false,
111 );
112
113 let payload = match serde_json::to_string(&envelope) {
115 Ok(json) => json,
116 Err(err) => {
117 state.metrics.routing_failure();
118 cleanup_pending_request(&state, &tenant_id, &request_id);
119 return PeerRpcResponse::rejected(
120 request_id,
121 format!("serialization_failed: {err}"),
122 );
123 }
124 };
125
126 if let Err(err) = worker_conn.send(Message::Text(payload.into())) {
127 state.metrics.routing_failure();
128 cleanup_pending_request(&state, &tenant_id, &request_id);
129 return PeerRpcResponse::rejected(request_id, format!("forward_failed: {err:?}"));
130 }
131
132 tokio::select! {
134 biased;
135 _ = state.wait_for_shutdown() => {
136 state.metrics.routing_failure();
137 PeerRpcResponse::rejected(request_id, "gateway_shutting_down")
138 }
139 result = tokio::time::timeout(timeout, rx) => match result {
140 Ok(Ok(response)) => {
141 state.metrics.message_routed();
142 response
143 }
144 Ok(Err(_)) => {
145 state.metrics.routing_failure();
146 PeerRpcResponse::rejected(request_id, "worker_connection_lost")
147 }
148 Err(_) => {
149 state.metrics.routing_failure();
150 PeerRpcResponse::rejected(request_id, "worker_response_timeout")
151 }
152 }
153 }
154 }
155
156 pub fn handle_worker_response(
158 state: Arc<GatewayState>,
159 tenant_id: &TenantId,
160 response: PeerRpcResponse,
161 ) -> GatewayResult<()> {
162 dispatch_worker_response(state, tenant_id, None, response)
163 }
164
165 pub fn handle_worker_response_from(
167 state: Arc<GatewayState>,
168 tenant_id: &TenantId,
169 worker: &crate::WorkerConnection,
170 response: PeerRpcResponse,
171 ) -> GatewayResult<()> {
172 dispatch_worker_response(state, tenant_id, Some(worker.generation()), response)
173 }
174
175 pub async fn route_mesh_request(
177 state: Arc<GatewayState>,
178 request: MeshPeerRequest,
179 timeout: Duration,
180 ) -> MeshPeerResponse {
181 if let Err(error) = request.validate_schema() {
182 state.metrics.routing_failure();
183 return MeshPeerResponse::rejected(request.request_id, error.to_string());
184 }
185 let request_id = request.request_id.clone();
186 let tenant_id = request.target_tenant_id.clone();
187 let timeout = timeout.min(crate::config::MAX_GATEWAY_REQUEST_TIMEOUT);
188 let peer_envelope = request.peer_envelope().ok();
189
190 let (rx, worker_conn) = {
191 let mut tenants = state.tenants.write();
192 let Some(tenant_state) = tenants.get_mut(&tenant_id) else {
193 state.metrics.routing_failure();
194 return MeshPeerResponse::rejected(request_id, "tenant_unavailable");
195 };
196 let Some(worker_conn) = tenant_state
197 .get_worker_by_core(&request.target_core_id)
198 .filter(|worker| {
199 peer_envelope.as_ref().is_none_or(|envelope| {
200 worker.cluster_id() == Some(&envelope.cluster_id)
201 && tenant_state
202 .registry
203 .resolve(&envelope.capability)
204 .is_some_and(|workers| workers.contains(&worker.key))
205 })
206 })
207 .cloned()
208 else {
209 state.metrics.routing_failure();
210 return MeshPeerResponse::rejected(request_id, "worker_offline");
211 };
212 if !tenant_state.can_register_pending(&request_id, true) {
213 state.metrics.routing_failure();
214 return MeshPeerResponse::rejected(request_id, "pending_request_rejected");
215 }
216 let (tx, rx) = oneshot::channel();
217 tenant_state
218 .pending_mesh_requests
219 .insert(request_id.clone(), tx);
220 register_pending_metadata(
221 &state,
222 &tenant_id,
223 &request_id,
224 true,
225 worker_conn.generation(),
226 Some(request.max_response_bytes),
227 );
228 (rx, worker_conn)
229 };
230 let _cleanup = PendingCleanup::new(
231 Arc::clone(&state),
232 tenant_id.clone(),
233 request_id.clone(),
234 true,
235 );
236
237 let payload = match serde_json::to_string(&request) {
238 Ok(json) => json,
239 Err(error) => {
240 state.metrics.routing_failure();
241 cleanup_pending_mesh_request(&state, &tenant_id, &request_id);
242 return MeshPeerResponse::rejected(
243 request_id,
244 format!("serialization_failed: {error}"),
245 );
246 }
247 };
248 if let Err(error) = worker_conn.send(Message::Text(payload.into())) {
249 state.metrics.routing_failure();
250 cleanup_pending_mesh_request(&state, &tenant_id, &request_id);
251 return MeshPeerResponse::rejected(request_id, format!("forward_failed: {error:?}"));
252 }
253
254 tokio::select! {
255 biased;
256 _ = state.wait_for_shutdown() => {
257 state.metrics.routing_failure();
258 MeshPeerResponse::rejected(request_id, "gateway_shutting_down")
259 }
260 result = tokio::time::timeout(timeout, rx) => match result {
261 Ok(Ok(response)) => {
262 state.metrics.message_routed();
263 response
264 }
265 Ok(Err(_)) => {
266 state.metrics.routing_failure();
267 MeshPeerResponse::rejected(request_id, "worker_connection_lost")
268 }
269 Err(_) => {
270 state.metrics.routing_failure();
271 MeshPeerResponse::rejected(request_id, "worker_response_timeout")
272 }
273 }
274 }
275 }
276
277 pub fn handle_worker_mesh_response(
279 state: Arc<GatewayState>,
280 tenant_id: &TenantId,
281 response: MeshPeerResponse,
282 ) -> GatewayResult<()> {
283 dispatch_worker_mesh_response(state, tenant_id, None, response)
284 }
285
286 pub fn handle_worker_mesh_response_from(
288 state: Arc<GatewayState>,
289 tenant_id: &TenantId,
290 worker: &crate::WorkerConnection,
291 response: MeshPeerResponse,
292 ) -> GatewayResult<()> {
293 dispatch_worker_mesh_response(state, tenant_id, Some(worker.generation()), response)
294 }
295}
296
297struct PendingCleanup {
298 state: Arc<GatewayState>,
299 tenant_id: TenantId,
300 request_id: String,
301 mesh: bool,
302}
303
304impl PendingCleanup {
305 fn new(state: Arc<GatewayState>, tenant_id: TenantId, request_id: String, mesh: bool) -> Self {
306 Self {
307 state,
308 tenant_id,
309 request_id,
310 mesh,
311 }
312 }
313}
314
315impl Drop for PendingCleanup {
316 fn drop(&mut self) {
317 if self.mesh {
318 cleanup_pending_mesh_request(&self.state, &self.tenant_id, &self.request_id);
319 } else {
320 cleanup_pending_request(&self.state, &self.tenant_id, &self.request_id);
321 }
322 }
323}
324
325fn cleanup_pending_request(state: &GatewayState, tenant_id: &TenantId, request_id: &str) {
326 let mut tenants = state.tenants.write();
327 if let Some(tenant_state) = tenants.get_mut(tenant_id) {
328 tenant_state.pending_requests.remove(request_id);
329 }
330 remove_pending_metadata(state, tenant_id, request_id, false);
331}
332
333fn cleanup_pending_mesh_request(state: &GatewayState, tenant_id: &TenantId, request_id: &str) {
334 let mut tenants = state.tenants.write();
335 if let Some(tenant_state) = tenants.get_mut(tenant_id) {
336 tenant_state.pending_mesh_requests.remove(request_id);
337 }
338 remove_pending_metadata(state, tenant_id, request_id, true);
339}
340
341fn dispatch_worker_response(
342 state: Arc<GatewayState>,
343 tenant_id: &TenantId,
344 generation: Option<u64>,
345 response: PeerRpcResponse,
346) -> GatewayResult<()> {
347 let request_id = response.request_id.clone();
348 let mut tenants = state.tenants.write();
349 if let Some(tenant_state) = tenants.get_mut(tenant_id) {
350 let expected = pending_metadata(&state, tenant_id, &request_id, false);
351 if expected.is_some_and(|metadata| {
352 generation.is_none_or(|value| value == metadata.worker_generation)
353 }) {
354 remove_pending_metadata(&state, tenant_id, &request_id, false);
355 if let Some(tx) = tenant_state.pending_requests.remove(&request_id) {
356 let _ = tx.send(response);
357 return Ok(());
358 }
359 }
360 }
361 Err(orphaned_response(tenant_id, &request_id, false))
362}
363
364fn dispatch_worker_mesh_response(
365 state: Arc<GatewayState>,
366 tenant_id: &TenantId,
367 generation: Option<u64>,
368 response: MeshPeerResponse,
369) -> GatewayResult<()> {
370 let request_id = response.request_id.clone();
371 let mut tenants = state.tenants.write();
372 if let Some(tenant_state) = tenants.get_mut(tenant_id) {
373 let expected = pending_metadata(&state, tenant_id, &request_id, true);
374 if expected.is_some_and(|metadata| {
375 generation.is_none_or(|value| value == metadata.worker_generation)
376 && metadata
377 .response_limit
378 .is_some_and(|limit| response.validate_for_request(&request_id, limit).is_ok())
379 }) {
380 remove_pending_metadata(&state, tenant_id, &request_id, true);
381 if let Some(tx) = tenant_state.pending_mesh_requests.remove(&request_id) {
382 let _ = tx.send(response);
383 return Ok(());
384 }
385 }
386 }
387 Err(orphaned_response(tenant_id, &request_id, true))
388}
389
390fn register_pending_metadata(
391 state: &GatewayState,
392 tenant_id: &TenantId,
393 request_id: &str,
394 mesh: bool,
395 worker_generation: u64,
396 response_limit: Option<usize>,
397) {
398 pending_metadata_map().insert(
399 pending_key(state, tenant_id, request_id, mesh),
400 PendingMetadata {
401 worker_generation,
402 response_limit,
403 },
404 );
405}
406
407fn pending_metadata(
408 state: &GatewayState,
409 tenant_id: &TenantId,
410 request_id: &str,
411 mesh: bool,
412) -> Option<PendingMetadata> {
413 pending_metadata_map()
414 .get(&pending_key(state, tenant_id, request_id, mesh))
415 .copied()
416}
417
418fn remove_pending_metadata(
419 state: &GatewayState,
420 tenant_id: &TenantId,
421 request_id: &str,
422 mesh: bool,
423) {
424 pending_metadata_map().remove(&pending_key(state, tenant_id, request_id, mesh));
425}
426
427fn pending_metadata_map() -> std::sync::MutexGuard<'static, HashMap<PendingKey, PendingMetadata>> {
428 PENDING_METADATA
429 .lock()
430 .unwrap_or_else(std::sync::PoisonError::into_inner)
431}
432
433fn pending_key(
434 state: &GatewayState,
435 tenant_id: &TenantId,
436 request_id: &str,
437 mesh: bool,
438) -> PendingKey {
439 PendingKey {
440 state_id: std::ptr::from_ref(state) as usize,
441 tenant_id: tenant_id.as_str().to_string(),
442 request_id: request_id.to_string(),
443 mesh,
444 }
445}
446
447fn orphaned_response(tenant_id: &TenantId, request_id: &str, mesh: bool) -> GatewayError {
448 GatewayError::Protocol(format!(
449 "orphaned {}response or timeout for tenant {} request: {}",
450 if mesh { "worker mesh " } else { "worker " },
451 tenant_id.as_str(),
452 request_id
453 ))
454}
455
456fn validate_routing_envelope(envelope: &PeerRpcEnvelope) -> Option<&'static str> {
457 if envelope.protocol_version != ProtocolVersion::default() {
458 return Some("protocol_version_unsupported");
459 }
460 if envelope.expires_at_ms <= envelope.timestamp_ms {
461 return Some("envelope_expiry_invalid");
462 }
463 if envelope.expires_at_ms <= now_ms() {
464 return Some("envelope_expired");
465 }
466 if envelope.body_hash != payload_hash(&envelope.payload) {
467 return Some("body_hash_invalid");
468 }
469 if envelope.payload.len() > crate::config::MAX_GATEWAY_MESSAGE_BYTES {
470 return Some("payload_too_large");
471 }
472 None
473}
474
475fn now_ms() -> u64 {
476 SystemTime::now()
477 .duration_since(UNIX_EPOCH)
478 .map(|duration| duration.as_millis() as u64)
479 .unwrap_or(0)
480}