Skip to main content

agentic_server/handler/websocket/
responses.rs

1use std::collections::VecDeque;
2use std::sync::Arc;
3
4use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
5use axum::extract::{Extension, State};
6use axum::http::HeaderMap;
7use axum::response::Response;
8use either::Either;
9use futures::stream::{SplitSink, SplitStream};
10use futures::{Sink, SinkExt, Stream, StreamExt};
11use serde_json::Value;
12use tokio_util::sync::CancellationToken;
13use tracing::{debug, warn};
14
15use agentic_core::ResponseUsage;
16use agentic_core::executor::{
17    BoxStream, ExecuteRequest, ExecutorError, RequestContext, persist_turn, rehydrate_conversation,
18};
19use agentic_core::types::request_response::RequestPayload;
20use agentic_core::utils::common::utcnow_str;
21
22use super::super::common::{MAX_BODY_SIZE, extract_bearer};
23use super::error::WsError;
24use crate::app::AppState;
25use crate::auth::AuthenticatedPrincipal;
26
27type WsSender = SplitSink<WebSocket, Message>;
28type WsReceiver = SplitStream<WebSocket>;
29
30pub async fn responses_ws(State(state): State<AppState>, headers: HeaderMap, ws: WebSocketUpgrade) -> Response {
31    upgrade_responses_ws(state, headers, ws, None)
32}
33
34pub(crate) async fn responses_ws_with_auth(
35    State(state): State<AppState>,
36    principal: Option<Extension<AuthenticatedPrincipal>>,
37    headers: HeaderMap,
38    ws: WebSocketUpgrade,
39) -> Response {
40    upgrade_responses_ws(state, headers, ws, principal.map(|Extension(principal)| principal))
41}
42
43fn upgrade_responses_ws(
44    state: AppState,
45    headers: HeaderMap,
46    ws: WebSocketUpgrade,
47    principal: Option<AuthenticatedPrincipal>,
48) -> Response {
49    let websocket_guard = state.websocket_tracker.track();
50    ws.max_message_size(MAX_BODY_SIZE)
51        .max_frame_size(MAX_BODY_SIZE)
52        .on_upgrade(move |socket| async move {
53            let _websocket_guard = websocket_guard;
54            responses_ws_loop(socket, state, headers, principal).await;
55        })
56}
57
58async fn responses_ws_loop(
59    socket: WebSocket,
60    state: AppState,
61    headers: HeaderMap,
62    principal: Option<AuthenticatedPrincipal>,
63) {
64    debug!("responses websocket session opened");
65    let shutdown_token = state.shutdown_token.clone();
66    let (mut sender, mut receiver) = socket.split();
67
68    // Requests received while a stream is active, processed in order after it completes.
69    let mut queue: VecDeque<String> = VecDeque::new();
70
71    loop {
72        if shutdown_token.is_cancelled() {
73            break;
74        }
75        let text = if let Some(buffered) = queue.pop_front() {
76            buffered
77        } else {
78            let message = next_ws_message(&shutdown_token, &mut receiver).await;
79
80            let Some(message) = message else {
81                break;
82            };
83
84            match message {
85                Ok(Message::Text(text)) => text.to_string(),
86                Ok(Message::Binary(_)) => {
87                    if !handle_ws_error(&mut sender, WsError::BinaryFrame).await {
88                        break;
89                    }
90                    continue;
91                }
92                Ok(Message::Close(_)) => break,
93                Ok(Message::Ping(payload)) => {
94                    if sender.send(Message::Pong(payload)).await.is_err() {
95                        break;
96                    }
97                    continue;
98                }
99                Ok(Message::Pong(_)) => continue,
100                Err(e) => {
101                    warn!("responses websocket receive error: {e}");
102                    break;
103                }
104            }
105        };
106
107        if let Some(event) = websocket_identity_error_event(principal.as_ref()) {
108            let _ = send_ws_json(&mut sender, event).await;
109            break;
110        }
111
112        match handle_ws_text(
113            &mut sender,
114            &mut receiver,
115            &state,
116            &headers,
117            &text,
118            &shutdown_token,
119            &mut queue,
120        )
121        .await
122        {
123            Ok(()) => {}
124            Err(err) => {
125                if !handle_ws_error(&mut sender, err).await {
126                    break;
127                }
128            }
129        }
130    }
131    close_ws(&mut sender, &mut receiver).await;
132    debug!("responses websocket session closed");
133}
134
135fn websocket_identity_error_event(principal: Option<&AuthenticatedPrincipal>) -> Option<Value> {
136    principal.is_some_and(AuthenticatedPrincipal::is_expired).then(|| {
137        serde_json::json!({
138            "type": "error",
139            "code": "invalid_token",
140            "message": "OIDC bearer token expired",
141            "param": null,
142            "sequence_number": 0,
143        })
144    })
145}
146
147async fn next_ws_message<Receiver>(
148    shutdown_token: &CancellationToken,
149    receiver: &mut Receiver,
150) -> Option<Receiver::Item>
151where
152    Receiver: Stream + Unpin,
153{
154    tokio::select! {
155        biased;
156        () = shutdown_token.cancelled() => None,
157        message = receiver.next() => {
158            if shutdown_token.is_cancelled() {
159                None
160            } else {
161                message
162            }
163        },
164    }
165}
166
167fn keep_if_running<T>(shutdown_token: &CancellationToken, value: T) -> Option<T> {
168    (!shutdown_token.is_cancelled()).then_some(value)
169}
170
171async fn close_ws<Sender, Receiver, SendError, ReceiveError>(sender: &mut Sender, receiver: &mut Receiver)
172where
173    Sender: Sink<Message, Error = SendError> + Unpin,
174    Receiver: Stream<Item = Result<Message, ReceiveError>> + Unpin,
175    SendError: std::fmt::Display,
176    ReceiveError: std::fmt::Display,
177{
178    if let Err(error) = sender.close().await {
179        debug!(%error, "failed to send responses websocket close frame");
180        return;
181    }
182
183    while let Some(message) = receiver.next().await {
184        match message {
185            Ok(Message::Close(_)) => break,
186            Ok(Message::Text(_) | Message::Binary(_) | Message::Ping(_) | Message::Pong(_)) => {}
187            Err(error) => {
188                debug!(%error, "responses websocket close handshake receive failed");
189                break;
190            }
191        }
192    }
193}
194
195/// Process one `response.create` message.
196///
197/// Any requests received from the client while the stream is active are
198/// pushed onto `queue` and processed by the caller in order after this returns.
199async fn handle_ws_text(
200    sender: &mut WsSender,
201    receiver: &mut WsReceiver,
202    state: &AppState,
203    headers: &HeaderMap,
204    text: &str,
205    shutdown_token: &CancellationToken,
206    queue: &mut VecDeque<String>,
207) -> Result<(), WsError> {
208    let value = serde_json::from_str::<Value>(text).map_err(WsError::InvalidJson)?;
209
210    if value.get("type").and_then(Value::as_str) != Some("response.create") {
211        return Err(WsError::UnexpectedType);
212    }
213
214    let generate = value.get("generate").and_then(Value::as_bool);
215    let mut payload = serde_json::from_value::<RequestPayload>(value).map_err(ExecutorError::from)?;
216    let requested_stream = payload.stream;
217    let requested_store = payload.store;
218    payload.stream = true;
219    payload.store = true;
220    debug!(
221        requested_stream,
222        requested_store,
223        forced_stream = payload.stream,
224        forced_store = payload.store,
225        has_previous_response_id = payload.previous_response_id.is_some(),
226        has_conversation_id = payload.conversation_id.is_some(),
227        ?generate,
228        tools = payload.tools.as_ref().map_or(0, Vec::len),
229        "accepted websocket response.create"
230    );
231
232    if generate == Some(false) {
233        debug!("handling non-generating websocket request locally");
234        return complete_without_inference(sender, state, payload).await;
235    }
236
237    let auth = extract_bearer(headers, state.openai_api_key.as_deref());
238    let result = ExecuteRequest::new(payload, Arc::clone(&state.exec_ctx))
239        .with_auth(auth)
240        .run()
241        .await?;
242    let Some(result) = keep_if_running(shutdown_token, result) else {
243        debug!("discarded websocket response initialized during shutdown");
244        return Ok(());
245    };
246    let Either::Right(stream) = result else {
247        return Err(WsError::Executor(Box::new(ExecutorError::InvalidRequest(
248            "websocket response.create must produce a stream".to_owned(),
249        ))));
250    };
251
252    stream_ws_response(sender, receiver, stream, shutdown_token, queue).await
253}
254
255async fn complete_without_inference(
256    sender: &mut WsSender,
257    state: &AppState,
258    payload: RequestPayload,
259) -> Result<(), WsError> {
260    let ctx = rehydrate_conversation(payload, &state.exec_ctx).await?;
261    let created_at = utcnow_str();
262    let created_event = empty_response_event(&ctx, created_at, "response.created", "in_progress", 0, None);
263    let completed_event = empty_response_event(
264        &ctx,
265        created_at,
266        "response.completed",
267        "completed",
268        1,
269        Some(ResponseUsage::default()),
270    );
271
272    #[cfg(debug_assertions)]
273    state.websocket_tracker.pause_local_completion_after_rehydration().await;
274    persist_turn(
275        ctx,
276        Vec::new(),
277        &state.exec_ctx.conv_handler,
278        &state.exec_ctx.resp_handler,
279    )
280    .await?;
281
282    send_ws_json(sender, created_event).await?;
283    send_ws_json(sender, completed_event).await
284}
285
286fn empty_response_event(
287    ctx: &RequestContext,
288    created_at: i64,
289    event_type: &str,
290    status: &str,
291    sequence_number: u32,
292    usage: Option<ResponseUsage>,
293) -> Value {
294    serde_json::json!({
295        "type": event_type,
296        "sequence_number": sequence_number,
297        "response": {
298            "id": &ctx.response_id,
299            "object": "response",
300            "created_at": created_at,
301            "model": &ctx.enriched_request.model,
302            "status": status,
303            "output": [],
304            "usage": usage,
305            "incomplete_details": null,
306            "error": null,
307            "previous_response_id": &ctx.original_request.previous_response_id,
308            "conversation_id": &ctx.conversation_id,
309            "instructions": &ctx.enriched_request.instructions,
310        },
311    })
312}
313
314enum ShutdownInput<ReceiverItem, UpstreamItem> {
315    Receiver(Option<ReceiverItem>),
316    Upstream(Option<UpstreamItem>),
317}
318
319async fn next_shutdown_input<Receiver, Upstream>(
320    receiver: &mut Receiver,
321    upstream: &mut Upstream,
322    prefer_receiver: bool,
323) -> ShutdownInput<Receiver::Item, Upstream::Item>
324where
325    Receiver: Stream + Unpin,
326    Upstream: Stream + Unpin,
327{
328    if prefer_receiver {
329        tokio::select! {
330            biased;
331            message = receiver.next() => ShutdownInput::Receiver(message),
332            line = upstream.next() => ShutdownInput::Upstream(line),
333        }
334    } else {
335        tokio::select! {
336            biased;
337            line = upstream.next() => ShutdownInput::Upstream(line),
338            message = receiver.next() => ShutdownInput::Receiver(message),
339        }
340    }
341}
342
343/// Stream a response from the executor to the client.
344///
345/// Requests arriving from the client while the stream is active are pushed
346/// onto `queue` so the caller can process them in order after this returns.
347async fn stream_ws_response(
348    sender: &mut WsSender,
349    receiver: &mut WsReceiver,
350    mut stream: BoxStream,
351    shutdown_token: &CancellationToken,
352    queue: &mut VecDeque<String>,
353) -> Result<(), WsError> {
354    let mut prefer_shutdown_receiver = true;
355    'stream: loop {
356        if shutdown_token.is_cancelled() {
357            match next_shutdown_input(receiver, &mut stream, prefer_shutdown_receiver).await {
358                ShutdownInput::Receiver(message) => {
359                    prefer_shutdown_receiver = false;
360                    match message {
361                        None | Some(Ok(Message::Close(_))) => return Err(WsError::ClientDisconnected),
362                        Some(Ok(Message::Ping(payload))) => {
363                            sender
364                                .send(Message::Pong(payload))
365                                .await
366                                .map_err(|_| WsError::SendFailed)?;
367                        }
368                        Some(Ok(Message::Text(_) | Message::Binary(_) | Message::Pong(_))) => {}
369                        Some(Err(error)) => return Err(WsError::Receive(error.to_string())),
370                    }
371                    continue 'stream;
372                }
373                ShutdownInput::Upstream(line) => {
374                    prefer_shutdown_receiver = true;
375                    let Some(line) = line else {
376                        break;
377                    };
378                    forward_ws_stream_chunk(sender, &line).await?;
379                }
380            }
381            continue;
382        }
383
384        let next_line = tokio::select! {
385            () = shutdown_token.cancelled() => continue 'stream,
386            message = receiver.next() => {
387                match message {
388                    None | Some(Ok(Message::Close(_))) => return Err(WsError::ClientDisconnected),
389                    Some(Ok(Message::Ping(payload))) => {
390                        sender.send(Message::Pong(payload)).await.map_err(|_| WsError::SendFailed)?;
391                        continue 'stream;
392                    }
393                    Some(Ok(Message::Pong(_))) => continue 'stream,
394                    Some(Ok(Message::Binary(_))) => return Err(WsError::BinaryFrame),
395                    Some(Ok(Message::Text(text))) => {
396                        // Client pipelined the next request while we are still streaming.
397                        // Enqueue it and keep draining the current stream.
398                        queue.push_back(text.to_string());
399                        debug!(
400                            queued_requests = queue.len(),
401                            "queued pipelined websocket response.create while stream is active"
402                        );
403                        continue 'stream;
404                    }
405                    Some(Err(e)) => return Err(WsError::Receive(e.to_string())),
406                }
407            }
408            line = stream.next() => line,
409        };
410        let Some(line) = next_line else {
411            break;
412        };
413        forward_ws_stream_chunk(sender, &line).await?;
414    }
415
416    Ok(())
417}
418
419fn sse_json_data_lines(chunk: &str) -> impl Iterator<Item = &str> {
420    chunk
421        .lines()
422        .filter_map(|line| line.strip_prefix("data: "))
423        .map(str::trim)
424        .filter(|data| *data != "[DONE]")
425}
426
427async fn forward_ws_stream_chunk(sender: &mut WsSender, chunk: &str) -> Result<(), WsError> {
428    for data in sse_json_data_lines(chunk) {
429        let value = serde_json::from_str::<Value>(data)
430            .map_err(ExecutorError::from)
431            .map_err(WsError::from)?;
432        send_ws_json(sender, value).await?;
433    }
434    Ok(())
435}
436
437async fn handle_ws_error(sender: &mut WsSender, err: WsError) -> bool {
438    match err {
439        WsError::ClientDisconnected | WsError::SendFailed => false,
440        WsError::Receive(message) => {
441            warn!("responses websocket receive error: {message}");
442            false
443        }
444        err => send_ws_error(sender, &err).await.is_ok(),
445    }
446}
447
448async fn send_ws_error(sender: &mut WsSender, err: &WsError) -> Result<(), WsError> {
449    let Some(frame) = err.to_ws_frame() else {
450        return Err(WsError::SendFailed);
451    };
452    send_ws_json(sender, frame).await
453}
454
455async fn send_ws_json(sender: &mut WsSender, value: Value) -> Result<(), WsError> {
456    let text = serde_json::to_string(&value).map_err(WsError::SerializeJson)?;
457    sender
458        .send(Message::Text(text.into()))
459        .await
460        .map_err(|_| WsError::SendFailed)
461}
462
463#[cfg(test)]
464mod tests {
465    use std::pin::Pin;
466    use std::task::{Context, Poll};
467
468    use axum::extract::ws::Message;
469    use futures::{Sink, Stream, StreamExt, sink, stream};
470    use serde_json::json;
471    use tokio_util::sync::CancellationToken;
472
473    use super::{
474        ShutdownInput, WsError, close_ws, keep_if_running, next_shutdown_input, next_ws_message, sse_json_data_lines,
475        websocket_identity_error_event,
476    };
477    use crate::auth::AuthenticatedPrincipal;
478
479    struct CloseErrorSink;
480
481    struct CancellingStream {
482        shutdown_token: CancellationToken,
483        item: Option<&'static str>,
484    }
485
486    #[test]
487    fn sse_json_data_lines_accept_named_and_data_only_frames() {
488        let chunk = concat!(
489            "event: response.completed\n",
490            "data: {\"type\":\"response.completed\"}\n\n",
491            "data: [DONE]\n\n",
492        );
493
494        assert_eq!(
495            sse_json_data_lines(chunk).collect::<Vec<_>>(),
496            [r#"{"type":"response.completed"}"#]
497        );
498    }
499
500    impl Stream for CancellingStream {
501        type Item = &'static str;
502
503        fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
504            self.shutdown_token.cancel();
505            Poll::Ready(self.item.take())
506        }
507    }
508
509    impl Sink<Message> for CloseErrorSink {
510        type Error = &'static str;
511
512        fn poll_ready(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
513            Poll::Ready(Ok(()))
514        }
515
516        fn start_send(self: Pin<&mut Self>, _item: Message) -> Result<(), Self::Error> {
517            Ok(())
518        }
519
520        fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
521            Poll::Ready(Ok(()))
522        }
523
524        fn poll_close(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
525            Poll::Ready(Err("close failed"))
526        }
527    }
528
529    #[tokio::test]
530    async fn cancelled_shutdown_wins_over_ready_websocket_message() {
531        let shutdown_token = CancellationToken::new();
532        shutdown_token.cancel();
533        let mut receiver = stream::iter(["must remain unread"]);
534
535        assert!(next_ws_message(&shutdown_token, &mut receiver).await.is_none());
536        assert_eq!(receiver.next().await, Some("must remain unread"));
537    }
538
539    #[tokio::test]
540    async fn cancellation_during_receive_discards_websocket_message() {
541        let shutdown_token = CancellationToken::new();
542        let mut receiver = CancellingStream {
543            shutdown_token: shutdown_token.clone(),
544            item: Some("must be discarded"),
545        };
546
547        assert!(next_ws_message(&shutdown_token, &mut receiver).await.is_none());
548        assert!(shutdown_token.is_cancelled());
549        assert_eq!(receiver.next().await, None);
550    }
551
552    #[test]
553    fn cancellation_after_request_setup_discards_unpolled_stream() {
554        let shutdown_token = CancellationToken::new();
555        shutdown_token.cancel();
556
557        assert_eq!(keep_if_running(&shutdown_token, "unpolled stream"), None);
558    }
559
560    #[test]
561    fn websocket_identity_expiry_uses_responses_error_event() {
562        assert!(websocket_identity_error_event(None).is_none());
563        let frame = websocket_identity_error_event(Some(&AuthenticatedPrincipal::expired_for_test()))
564            .expect("expired-token error event");
565
566        assert_eq!(
567            frame,
568            json!({
569                "type": "error",
570                "code": "invalid_token",
571                "message": "OIDC bearer token expired",
572                "param": null,
573                "sequence_number": 0,
574            })
575        );
576
577        let generic_frame = WsError::UnexpectedType
578            .to_ws_frame()
579            .expect("generic client-visible error frame");
580        assert_eq!(generic_frame["status"], 400);
581        assert_eq!(generic_frame["error"]["code"], "invalid_request_error");
582    }
583
584    #[tokio::test]
585    async fn close_ws_ignores_late_frames_until_peer_close() {
586        let mut sender = sink::drain();
587        let mut receiver = stream::iter([
588            Ok::<_, &'static str>(Message::Text("late request".into())),
589            Ok(Message::Binary(vec![1].into())),
590            Ok(Message::Close(None)),
591            Err("must remain unread"),
592        ]);
593
594        close_ws(&mut sender, &mut receiver).await;
595
596        assert!(matches!(receiver.next().await, Some(Err("must remain unread"))));
597    }
598
599    #[tokio::test]
600    async fn close_ws_returns_without_reading_when_close_send_fails() {
601        let mut sender = CloseErrorSink;
602        let mut receiver = stream::iter([Ok::<_, &'static str>(Message::Close(None))]);
603
604        close_ws(&mut sender, &mut receiver).await;
605
606        assert!(matches!(receiver.next().await, Some(Ok(Message::Close(None)))));
607    }
608
609    #[tokio::test]
610    async fn close_ws_stops_reading_after_receive_error() {
611        let mut sender = sink::drain();
612        let mut receiver = stream::iter([Err::<Message, _>("receive failed"), Ok(Message::Close(None))]);
613
614        close_ws(&mut sender, &mut receiver).await;
615
616        assert!(matches!(receiver.next().await, Some(Ok(Message::Close(None)))));
617    }
618
619    #[tokio::test]
620    async fn shutdown_input_priority_alternates_when_both_streams_are_ready() {
621        let mut receiver = stream::repeat(());
622        let mut upstream = stream::repeat(());
623
624        assert!(matches!(
625            next_shutdown_input(&mut receiver, &mut upstream, true).await,
626            ShutdownInput::Receiver(Some(()))
627        ));
628        assert!(matches!(
629            next_shutdown_input(&mut receiver, &mut upstream, false).await,
630            ShutdownInput::Upstream(Some(()))
631        ));
632    }
633}