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 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
195async 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
343async 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 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}