Skip to main content

agentic_server/handler/websocket/
responses.rs

1use std::collections::VecDeque;
2use std::sync::Arc;
3
4use axum::extract::State;
5use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
6use axum::http::HeaderMap;
7use axum::response::Response;
8use either::Either;
9use futures::stream::{SplitSink, SplitStream};
10use futures::{SinkExt, StreamExt};
11use serde_json::Value;
12use tokio_util::sync::CancellationToken;
13use tracing::{debug, warn};
14
15use agentic_core::executor::{BoxStream, ExecuteRequest, ExecutorError};
16use agentic_core::types::request_response::RequestPayload;
17
18use super::super::common::{MAX_BODY_SIZE, extract_bearer};
19use super::error::WsError;
20use crate::app::AppState;
21
22type WsSender = SplitSink<WebSocket, Message>;
23type WsReceiver = SplitStream<WebSocket>;
24
25pub async fn responses_ws(State(state): State<AppState>, headers: HeaderMap, ws: WebSocketUpgrade) -> Response {
26    ws.max_message_size(MAX_BODY_SIZE)
27        .max_frame_size(MAX_BODY_SIZE)
28        .on_upgrade(move |socket| responses_ws_loop(socket, state, headers))
29}
30
31async fn responses_ws_loop(socket: WebSocket, state: AppState, headers: HeaderMap) {
32    debug!("responses websocket session opened");
33    let shutdown_token = state.shutdown_token.clone();
34    let (mut sender, mut receiver) = socket.split();
35
36    // Requests received while a stream is active, processed in order after it completes.
37    let mut queue: VecDeque<String> = VecDeque::new();
38
39    loop {
40        let text = if let Some(buffered) = queue.pop_front() {
41            buffered
42        } else {
43            let message = tokio::select! {
44                () = shutdown_token.cancelled() => break,
45                message = receiver.next() => message,
46            };
47
48            let Some(message) = message else {
49                break;
50            };
51
52            match message {
53                Ok(Message::Text(text)) => text.to_string(),
54                Ok(Message::Binary(_)) => {
55                    if !handle_ws_error(&mut sender, WsError::BinaryFrame).await {
56                        break;
57                    }
58                    continue;
59                }
60                Ok(Message::Close(_)) => break,
61                Ok(Message::Ping(payload)) => {
62                    if sender.send(Message::Pong(payload)).await.is_err() {
63                        break;
64                    }
65                    continue;
66                }
67                Ok(Message::Pong(_)) => continue,
68                Err(e) => {
69                    warn!("responses websocket receive error: {e}");
70                    break;
71                }
72            }
73        };
74
75        match handle_ws_text(
76            &mut sender,
77            &mut receiver,
78            &state,
79            &headers,
80            &text,
81            &shutdown_token,
82            &mut queue,
83        )
84        .await
85        {
86            Ok(()) => {}
87            Err(err) => {
88                if !handle_ws_error(&mut sender, err).await {
89                    break;
90                }
91            }
92        }
93    }
94    debug!("responses websocket session closed");
95}
96
97/// Process one `response.create` message.
98///
99/// Any requests received from the client while the stream is active are
100/// pushed onto `queue` and processed by the caller in order after this returns.
101async fn handle_ws_text(
102    sender: &mut WsSender,
103    receiver: &mut WsReceiver,
104    state: &AppState,
105    headers: &HeaderMap,
106    text: &str,
107    shutdown_token: &CancellationToken,
108    queue: &mut VecDeque<String>,
109) -> Result<(), WsError> {
110    let value = serde_json::from_str::<Value>(text).map_err(WsError::InvalidJson)?;
111
112    if value.get("type").and_then(Value::as_str) != Some("response.create") {
113        return Err(WsError::UnexpectedType);
114    }
115
116    let mut payload = serde_json::from_value::<RequestPayload>(value).map_err(ExecutorError::from)?;
117    let requested_stream = payload.stream;
118    let requested_store = payload.store;
119    payload.stream = true;
120    payload.store = true;
121    debug!(
122        requested_stream,
123        requested_store,
124        forced_stream = payload.stream,
125        forced_store = payload.store,
126        has_previous_response_id = payload.previous_response_id.is_some(),
127        has_conversation_id = payload.conversation_id.is_some(),
128        tools = payload.tools.as_ref().map_or(0, Vec::len),
129        "accepted websocket response.create"
130    );
131
132    let auth = extract_bearer(headers, state.openai_api_key.as_deref());
133    let result = ExecuteRequest::new(payload, Arc::clone(&state.exec_ctx))
134        .with_auth(auth)
135        .run()
136        .await?;
137    let Either::Right(stream) = result else {
138        return Err(WsError::Executor(ExecutorError::InvalidRequest(
139            "websocket response.create must produce a stream".to_owned(),
140        )));
141    };
142
143    stream_ws_response(sender, receiver, stream, shutdown_token, queue).await
144}
145
146/// Stream a response from the executor to the client.
147///
148/// Requests arriving from the client while the stream is active are pushed
149/// onto `queue` so the caller can process them in order after this returns.
150async fn stream_ws_response(
151    sender: &mut WsSender,
152    receiver: &mut WsReceiver,
153    mut stream: BoxStream,
154    shutdown_token: &CancellationToken,
155    queue: &mut VecDeque<String>,
156) -> Result<(), WsError> {
157    'stream: loop {
158        let next_line = tokio::select! {
159            () = shutdown_token.cancelled() => return Err(WsError::Shutdown),
160            message = receiver.next() => {
161                match message {
162                    None | Some(Ok(Message::Close(_))) => return Err(WsError::ClientDisconnected),
163                    Some(Ok(Message::Ping(payload))) => {
164                        sender.send(Message::Pong(payload)).await.map_err(|_| WsError::SendFailed)?;
165                        continue 'stream;
166                    }
167                    Some(Ok(Message::Pong(_))) => continue 'stream,
168                    Some(Ok(Message::Binary(_))) => return Err(WsError::BinaryFrame),
169                    Some(Ok(Message::Text(text))) => {
170                        // Client pipelined the next request while we are still streaming.
171                        // Enqueue it and keep draining the current stream.
172                        queue.push_back(text.to_string());
173                        debug!(
174                            queued_requests = queue.len(),
175                            "queued pipelined websocket response.create while stream is active"
176                        );
177                        continue 'stream;
178                    }
179                    Some(Err(e)) => return Err(WsError::Receive(e.to_string())),
180                }
181            }
182            line = stream.next() => line,
183        };
184        let Some(line) = next_line else {
185            break;
186        };
187        let Some(data) = line.strip_prefix("data: ") else {
188            continue;
189        };
190        let data = data.trim();
191        if data == "[DONE]" {
192            continue;
193        }
194        let value = match serde_json::from_str::<Value>(data) {
195            Ok(value) => value,
196            Err(e) => return Err(WsError::Executor(ExecutorError::from(e))),
197        };
198        send_ws_json(sender, value).await?;
199    }
200
201    Ok(())
202}
203
204async fn handle_ws_error(sender: &mut WsSender, err: WsError) -> bool {
205    match err {
206        WsError::Shutdown | WsError::ClientDisconnected | WsError::SendFailed => false,
207        WsError::Receive(message) => {
208            warn!("responses websocket receive error: {message}");
209            false
210        }
211        err => send_ws_error(sender, &err).await.is_ok(),
212    }
213}
214
215async fn send_ws_error(sender: &mut WsSender, err: &WsError) -> Result<(), WsError> {
216    let Some(frame) = err.to_ws_frame() else {
217        return Err(WsError::SendFailed);
218    };
219    send_ws_json(sender, frame).await
220}
221
222async fn send_ws_json(sender: &mut WsSender, value: Value) -> Result<(), WsError> {
223    let text = serde_json::to_string(&value).map_err(WsError::SerializeJson)?;
224    sender
225        .send(Message::Text(text.into()))
226        .await
227        .map_err(|_| WsError::SendFailed)
228}