1use crate::config::{MAX_GATEWAY_CONNECTIONS, MAX_GATEWAY_MESSAGE_BYTES, MAX_GATEWAY_TENANTS};
14use crate::connection::{
15 ClientConnection, WorkerConnection, WorkerConnectionKey, CONNECTION_BUFFER_CAPACITY,
16};
17use crate::{EnvelopeRouter, GatewaySession, GatewayState, MeshPeerResponse, TenantState};
18use appcore_contracts::InstallationId;
19use appcore_distributed_contracts::{PeerRpcEnvelope, PeerRpcResponse};
20use appcore_security::RuntimeTokenClaims;
21use appcore_types::{CapabilityName, ClusterId, CoreId, TenantId};
22use axum::extract::ws::{Message, WebSocket};
23use futures_util::stream::SplitSink;
24use futures_util::{SinkExt, StreamExt};
25use serde::Deserialize;
26use std::sync::atomic::{AtomicU64, Ordering};
27use std::sync::{Arc, LazyLock};
28use std::time::{Duration, SystemTime, UNIX_EPOCH};
29use tokio::sync::{mpsc, Semaphore};
30use tokio::task::{JoinHandle, JoinSet};
31use tracing::{info, warn};
32
33const MAX_IN_FLIGHT_CLIENT_REQUESTS: usize = 4_096;
34static REQUEST_SLOTS: LazyLock<Arc<Semaphore>> =
36 LazyLock::new(|| Arc::new(Semaphore::new(MAX_IN_FLIGHT_CLIENT_REQUESTS)));
37
38pub(crate) struct WorkerSocketContext {
39 pub(crate) tenant_id: TenantId,
40 pub(crate) cluster_id: ClusterId,
41 pub(crate) installation_id: InstallationId,
42 pub(crate) core_id: CoreId,
43 pub(crate) capabilities: Vec<CapabilityName>,
44 pub(crate) expires_at_ms: u64,
45}
46
47pub(crate) async fn handle_worker_socket(
48 state: Arc<GatewayState>,
49 context: WorkerSocketContext,
50 socket: WebSocket,
51) {
52 let WorkerSocketContext {
53 tenant_id,
54 cluster_id,
55 installation_id,
56 core_id,
57 capabilities,
58 expires_at_ms,
59 } = context;
60 let (sink, mut stream) = socket.split();
61 let (tx, rx) = mpsc::channel::<Message>(CONNECTION_BUFFER_CAPACITY);
62 let key = WorkerConnectionKey {
63 tenant_id: tenant_id.clone(),
64 installation_id: installation_id.clone(),
65 core_id: core_id.clone(),
66 };
67 let conn = WorkerConnection::new_in_cluster(key, cluster_id, tx, now_ms());
68 let replaced = {
69 let mut tenants = state.tenants.write();
70 if !tenants.contains_key(&tenant_id) && tenants.len() >= MAX_GATEWAY_TENANTS {
71 warn!("Gateway tenant limit rejected worker connection");
72 return;
73 }
74 if connection_count(&tenants) >= MAX_GATEWAY_CONNECTIONS {
75 warn!("Gateway global connection limit rejected worker connection");
76 return;
77 }
78 let tenant = tenants
79 .entry(tenant_id.clone())
80 .or_insert_with(|| TenantState::new(tenant_id.clone()));
81 let replaced = tenant.get_worker(&installation_id, &core_id).is_some();
82 if tenant.add_worker(conn.clone(), capabilities).is_err() {
83 warn!("Gateway worker limit rejected connection");
84 return;
85 }
86 replaced
87 };
88 if !replaced {
89 state.metrics.worker_connected();
90 }
91 info!(
92 "Worker connected: tenant={}, installation={}, core={}",
93 tenant_id.as_str(),
94 installation_id.as_str(),
95 core_id.as_str()
96 );
97
98 let writer_task = spawn_socket_writer(&state, sink, rx);
99 let mut socket_shutdown = state.subscribe_shutdown();
100 loop {
101 let message = tokio::select! {
102 biased;
103 result = socket_shutdown.changed() => {
104 if result.is_err() || *socket_shutdown.borrow() {
105 break;
106 }
107 continue;
108 }
109 result = tokio::time::timeout(
110 session_wait(state.config().heartbeat_timeout, expires_at_ms),
111 stream.next(),
112 ) => match result {
113 Ok(Some(Ok(message))) => message,
114 _ => break,
115 }
116 };
117 if !handle_worker_message(&state, &tenant_id, &conn, message) {
118 break;
119 }
120 }
121 let removed = {
122 let mut tenants = state.tenants.write();
123 tenants.get_mut(&tenant_id).is_some_and(|tenant| {
124 tenant.remove_worker_if_current(&installation_id, &core_id, conn.generation())
125 })
126 };
127 writer_task.abort();
128 let _ = writer_task.await;
129 if removed {
130 state.metrics.worker_disconnected();
131 }
132 info!(
133 "Worker disconnected: tenant={}, installation={}, core={}",
134 tenant_id.as_str(),
135 installation_id.as_str(),
136 core_id.as_str()
137 );
138}
139
140#[allow(clippy::too_many_arguments)]
141pub(crate) async fn handle_client_socket(
142 state: Arc<GatewayState>,
143 tenant_id: TenantId,
144 cluster_id: ClusterId,
145 claims: RuntimeTokenClaims,
146 socket: WebSocket,
147) {
148 let session_id = unique_id("sess");
149 let connection_id = unique_id("conn");
150 let (sink, mut stream) = socket.split();
151 let (tx, rx) = mpsc::channel::<Message>(CONNECTION_BUFFER_CAPACITY);
152 let connection = ClientConnection::new(
153 connection_id.clone(),
154 tenant_id.clone(),
155 session_id.clone(),
156 tx,
157 );
158 let boundary = ClientBoundary {
159 cluster_id,
160 expires_at_ms: claims.expires_at_ms,
161 };
162 let session = GatewaySession::new(
163 session_id.clone(),
164 tenant_id.clone(),
165 now_ms(),
166 claims.expires_at_ms,
167 claims.subject,
168 );
169 {
170 let mut tenants = state.tenants.write();
171 if !tenants.contains_key(&tenant_id) && tenants.len() >= MAX_GATEWAY_TENANTS {
172 warn!("Gateway tenant limit rejected client connection");
173 return;
174 }
175 if connection_count(&tenants) >= MAX_GATEWAY_CONNECTIONS {
176 warn!("Gateway global connection limit rejected client connection");
177 return;
178 }
179 let tenant = tenants
180 .entry(tenant_id.clone())
181 .or_insert_with(|| TenantState::new(tenant_id.clone()));
182 if tenant.try_add_client(connection.clone()).is_err() {
183 warn!("Gateway client limit rejected connection");
184 return;
185 }
186 tenant.sessions.insert(session_id.clone(), session);
187 }
188 state.metrics.client_connected();
189 info!(
190 "Client connected: tenant={}, connection_id={}",
191 tenant_id.as_str(),
192 connection_id
193 );
194
195 let writer_task = spawn_socket_writer(&state, sink, rx);
196 let mut request_tasks = JoinSet::new();
197 let mut socket_shutdown = state.subscribe_shutdown();
198 loop {
199 let message = tokio::select! {
200 biased;
201 result = socket_shutdown.changed() => {
202 if result.is_err() || *socket_shutdown.borrow() {
203 break;
204 }
205 continue;
206 }
207 result = tokio::time::timeout(
208 session_wait(state.config().heartbeat_timeout, boundary.expires_at_ms),
209 stream.next(),
210 ) => match result {
211 Ok(Some(Ok(message))) => message,
212 _ => break,
213 }
214 };
215 while request_tasks.try_join_next().is_some() {}
216 if !handle_client_message(&state, &connection, &boundary, message, &mut request_tasks) {
217 break;
218 }
219 }
220 {
221 let mut tenants = state.tenants.write();
222 if let Some(tenant) = tenants.get_mut(&tenant_id) {
223 tenant.remove_client(&connection_id);
224 tenant.sessions.remove(&session_id);
225 }
226 }
227 request_tasks.abort_all();
228 while request_tasks.join_next().await.is_some() {}
229 writer_task.abort();
230 let _ = writer_task.await;
231 state.metrics.client_disconnected();
232 info!(
233 "Client disconnected: tenant={}, connection_id={}",
234 tenant_id.as_str(),
235 connection_id
236 );
237}
238
239fn spawn_socket_writer(
240 state: &GatewayState,
241 mut sink: SplitSink<WebSocket, Message>,
242 mut receiver: mpsc::Receiver<Message>,
243) -> JoinHandle<()> {
244 let mut shutdown = state.subscribe_shutdown();
245 tokio::spawn(async move {
246 loop {
247 tokio::select! {
248 biased;
249 result = shutdown.changed() => {
250 if result.is_err() || *shutdown.borrow() {
251 break;
252 }
253 }
254 message = receiver.recv() => {
255 let Some(message) = message else {
256 break;
257 };
258 if sink.send(message).await.is_err() {
259 break;
260 }
261 }
262 }
263 }
264 })
265}
266
267fn handle_worker_message(
268 state: &Arc<GatewayState>,
269 tenant_id: &TenantId,
270 connection: &WorkerConnection,
271 message: Message,
272) -> bool {
273 match message {
274 Message::Text(text) if text.len() <= MAX_GATEWAY_MESSAGE_BYTES => {
275 if is_heartbeat(&text) {
276 connection.update_heartbeat(now_ms());
277 return true;
278 }
279 if let Ok(response) = serde_json::from_str::<MeshPeerResponse>(&text) {
280 connection.update_heartbeat(now_ms());
281 return EnvelopeRouter::handle_worker_mesh_response_from(
282 Arc::clone(state),
283 tenant_id,
284 connection,
285 response,
286 )
287 .is_ok();
288 }
289 if let Ok(response) = serde_json::from_str::<PeerRpcResponse>(&text) {
290 connection.update_heartbeat(now_ms());
291 return EnvelopeRouter::handle_worker_response_from(
292 Arc::clone(state),
293 tenant_id,
294 connection,
295 response,
296 )
297 .is_ok();
298 }
299 false
300 }
301 Message::Ping(payload) => {
302 connection.update_heartbeat(now_ms());
303 connection.send(Message::Pong(payload)).is_ok()
304 }
305 Message::Pong(_) => {
306 connection.update_heartbeat(now_ms());
307 true
308 }
309 Message::Close(_) => false,
310 _ => false,
311 }
312}
313
314fn handle_client_message(
315 state: &Arc<GatewayState>,
316 connection: &ClientConnection,
317 boundary: &ClientBoundary,
318 message: Message,
319 request_tasks: &mut JoinSet<()>,
320) -> bool {
321 if now_ms() >= boundary.expires_at_ms {
322 return false;
323 }
324 let text = match message {
325 Message::Text(text) => text,
326 Message::Ping(payload) => return connection.send(Message::Pong(payload)).is_ok(),
327 Message::Pong(_) => return true,
328 Message::Close(_) => return false,
329 _ => return false,
330 };
331 if text.len() > MAX_GATEWAY_MESSAGE_BYTES {
332 return false;
333 }
334 let Ok(envelope) = serde_json::from_str::<PeerRpcEnvelope>(&text) else {
335 return false;
336 };
337 if envelope.tenant_id != connection.tenant_id || envelope.cluster_id != boundary.cluster_id {
338 send_rejection(
339 connection,
340 envelope.request_id,
341 "session_boundary_violation",
342 );
343 return true;
344 }
345 let Ok(permit) = Arc::clone(&REQUEST_SLOTS).try_acquire_owned() else {
346 send_rejection(connection, envelope.request_id, "gateway_overloaded");
347 return true;
348 };
349 let state = Arc::clone(state);
350 let connection = connection.clone();
351 request_tasks.spawn(async move {
352 let _permit = permit;
353 let response =
354 EnvelopeRouter::route_request(state, envelope, Duration::from_secs(10)).await;
355 if let Ok(json) = serde_json::to_string(&response) {
356 let _ = connection.send(Message::Text(json.into()));
357 }
358 });
359 true
360}
361
362struct ClientBoundary {
363 cluster_id: ClusterId,
364 expires_at_ms: u64,
365}
366
367fn send_rejection(connection: &ClientConnection, request_id: String, reason: &'static str) {
368 let response = PeerRpcResponse::rejected(request_id, reason);
369 if let Ok(json) = serde_json::to_string(&response) {
370 let _ = connection.send(Message::Text(json.into()));
371 }
372}
373
374#[derive(Deserialize)]
375#[serde(deny_unknown_fields)]
376struct HeartbeatFrame {
377 #[serde(rename = "type")]
378 kind: String,
379}
380
381fn is_heartbeat(text: &str) -> bool {
382 serde_json::from_str::<HeartbeatFrame>(text)
383 .map(|heartbeat| heartbeat.kind == "heartbeat")
384 .unwrap_or(false)
385}
386
387fn now_ms() -> u64 {
388 SystemTime::now()
389 .duration_since(UNIX_EPOCH)
390 .map(|duration| duration.as_millis() as u64)
391 .unwrap_or(0)
392}
393
394fn unique_id(prefix: &str) -> String {
395 static COUNTER: AtomicU64 = AtomicU64::new(0);
397 format!(
398 "{}-{}-{}-{}",
399 prefix,
400 std::process::id(),
401 now_ms(),
402 COUNTER.fetch_add(1, Ordering::Relaxed)
403 )
404}
405
406fn session_wait(idle_timeout: Duration, expires_at_ms: u64) -> Duration {
407 let remaining = Duration::from_millis(expires_at_ms.saturating_sub(now_ms()).max(1));
408 idle_timeout.min(remaining)
409}
410
411fn connection_count(tenants: &std::collections::HashMap<TenantId, TenantState>) -> usize {
412 tenants.values().fold(0usize, |total, tenant| {
413 total
414 .saturating_add(tenant.workers.len())
415 .saturating_add(tenant.clients.len())
416 })
417}
418
419#[cfg(test)]
420mod tests {
421 use super::{handle_client_message, is_heartbeat, ClientBoundary};
422 use crate::{ClientConnection, GatewayConfig, GatewayState};
423 use appcore_security::HashTokenProvider;
424 use appcore_types::{ClusterId, TenantId};
425 use axum::extract::ws::Message;
426 use std::sync::Arc;
427 use tokio::sync::mpsc;
428 use tokio::task::JoinSet;
429
430 #[test]
431 fn heartbeat_requires_the_exact_schema() {
432 assert!(is_heartbeat(r#"{"type":"heartbeat"}"#));
433 assert!(!is_heartbeat("heartbeat"));
434 assert!(!is_heartbeat(r#"{"type":"not-heartbeat"}"#));
435 assert!(!is_heartbeat(
436 r#"{"type":"heartbeat","credential":"secret"}"#
437 ));
438 }
439
440 #[tokio::test]
441 async fn client_ping_receives_pong() {
442 let provider = HashTokenProvider::from_secret(vec![9; 32]).unwrap();
443 let state = Arc::new(
444 GatewayState::new(
445 GatewayConfig::new(([127, 0, 0, 1], 8080).into(), "gateway.test"),
446 provider,
447 )
448 .unwrap(),
449 );
450 let tenant = TenantId::new("tenant-a").unwrap();
451 let cluster = ClusterId::new("cluster-a").unwrap();
452 let (sender, mut receiver) = mpsc::channel(1);
453 let connection = ClientConnection::new(
454 "connection-a".to_string(),
455 tenant,
456 "session-a".to_string(),
457 sender,
458 );
459 let boundary = ClientBoundary {
460 cluster_id: cluster,
461 expires_at_ms: u64::MAX,
462 };
463 let mut request_tasks = JoinSet::new();
464
465 assert!(handle_client_message(
466 &state,
467 &connection,
468 &boundary,
469 Message::Ping(vec![1, 2, 3].into()),
470 &mut request_tasks,
471 ));
472 assert_eq!(
473 receiver.recv().await,
474 Some(Message::Pong(vec![1, 2, 3].into()))
475 );
476 }
477}