agentic_server/handler/websocket/
responses.rs1use 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 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
97async 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
146async 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 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}