Skip to main content

appcore_gateway/
socket.rs

1// =============================================================================
2//        #######
3//     ###       ###     F: socket.rs
4//    ##   ## ##   ##    P: AppCore-Runtime
5//         ## ##
6//                       C: 2026/08/02 12:48:56 by dnettoRaw
7//    ##   ## ##   ##    U: 2026/08/02 12:48:56 by dnettoRaw
8//      ###########      S: 1.0.1-rc.8
9// =============================================================================
10
11//! Bounded Gateway worker and client WebSocket loops.
12
13use 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;
34// appcore-norm: allow(global-state) reason: process-wide semaphore enforces the configured request limit
35static 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    // appcore-norm: allow(global-state) reason: atomic sequence prevents process-local temporary path collisions
396    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}