1use std::{collections::VecDeque, convert::Infallible, sync::Arc};
2
3use axum::{
4 Json,
5 body::Body,
6 http::StatusCode,
7 response::{IntoResponse, Response},
8};
9use bytes::{Bytes, BytesMut};
10use http_body_util::BodyExt;
11use serde_json::{Value, json};
12
13use crate::{
14 provider::{Generation, GenerationBody},
15 providers::codex::native::NativeResponseOutcome,
16 traffic::{MAX_SSE_CAPTURE_BYTES, TrafficCapture},
17};
18
19use super::{
20 MAX_PROVIDER_STREAM_BYTES, MAX_SSE_EVENT_BYTES, OpenAiError, OpenAiResponseMetadata,
21 OpenAiSurface,
22 response::{
23 AnthropicAccumulator, BlockKind, SseEvent, buffered_response, chat_citation,
24 chat_finish_reason, hosted_search_action, normalized_arguments, responses_citation,
25 responses_response,
26 },
27};
28
29fn upstream_invalid(message: impl Into<String>, _param: Option<impl Into<String>>) -> OpenAiError {
30 OpenAiError::upstream_protocol(message)
31}
32
33#[derive(Default)]
34pub struct SseDecoder {
35 pending: BytesMut,
36}
37
38impl SseDecoder {
39 pub fn push(&mut self, bytes: &[u8]) -> Result<Vec<SseEvent>, OpenAiError> {
40 self.pending.extend_from_slice(bytes);
41 let mut events = Vec::new();
42 while let Some((position, delimiter_len)) = find_event_delimiter(&self.pending) {
43 if position > MAX_SSE_EVENT_BYTES {
44 return Err(OpenAiError::upstream_protocol(
45 "Provider SSE event exceeded the size limit",
46 ));
47 }
48 let frame = self.pending.split_to(position + delimiter_len);
49 let payload = &frame[..position];
50 if let Some(event) = parse_frame(payload)? {
51 events.push(event);
52 }
53 }
54 if self.pending.len() > MAX_SSE_EVENT_BYTES {
55 return Err(OpenAiError::upstream_protocol(
56 "Provider SSE event exceeded the size limit",
57 ));
58 }
59 Ok(events)
60 }
61
62 pub fn finish(&self) -> Result<(), OpenAiError> {
63 if self.pending.iter().all(u8::is_ascii_whitespace) {
64 Ok(())
65 } else {
66 Err(OpenAiError::upstream_protocol(
67 "Provider SSE stream ended with an incomplete event",
68 ))
69 }
70 }
71}
72
73pub async fn openai_response(
74 surface: OpenAiSurface,
75 generation: Generation,
76 stream: bool,
77 include_usage: bool,
78 response_metadata: OpenAiResponseMetadata,
79 traffic: Option<Arc<TrafficCapture>>,
80) -> Result<Response, OpenAiError> {
81 let model = generation.resolved_model.clone();
82 let response_id = match surface {
83 OpenAiSurface::ChatCompletions => format!("chatcmpl-{}", uuid::Uuid::new_v4().simple()),
84 OpenAiSurface::Responses => format!("resp_{}", uuid::Uuid::new_v4().simple()),
85 };
86 let created = current_seconds();
87 if stream {
88 return Ok(streaming_response(
89 generation.body,
90 Renderer::new(
91 surface,
92 include_usage,
93 response_id,
94 model,
95 created,
96 response_metadata,
97 ),
98 traffic,
99 ));
100 }
101 let bytes = collect_generation(generation.body).await?;
102 let events = decode_all(&bytes)?;
103 let value = buffered_response(
104 surface,
105 &events,
106 &response_id,
107 &model,
108 created,
109 &response_metadata,
110 )?;
111 if let Some(capture) = traffic.as_ref() {
112 capture.write_json("071-openai-downstream-response", &value);
113 }
114 Ok((StatusCode::OK, Json(value)).into_response())
115}
116
117async fn collect_generation(body: GenerationBody) -> Result<Bytes, OpenAiError> {
118 match body {
119 GenerationBody::BufferedSse(bytes) => {
120 if bytes.len() > MAX_PROVIDER_STREAM_BYTES {
121 Err(OpenAiError::upstream_protocol(
122 "Provider response exceeded the size limit",
123 ))
124 } else {
125 Ok(bytes)
126 }
127 }
128 GenerationBody::LiveSse(mut body) => {
129 let mut output = BytesMut::new();
130 while let Some(frame) = body.frame().await {
131 let frame = frame.map_err(|error| OpenAiError {
132 status: StatusCode::BAD_GATEWAY,
133 kind: "api_error".into(),
134 message: format!("Provider stream read failed: {error}").into(),
135 param: None,
136 code: None,
137 retry_after: None,
138 })?;
139 if let Ok(data) = frame.into_data() {
140 if output.len().saturating_add(data.len()) > MAX_PROVIDER_STREAM_BYTES {
141 return Err(OpenAiError::upstream_protocol(
142 "Provider response exceeded the size limit",
143 ));
144 }
145 output.extend_from_slice(&data);
146 }
147 }
148 Ok(output.freeze())
149 }
150 }
151}
152
153fn decode_all(bytes: &[u8]) -> Result<Vec<SseEvent>, OpenAiError> {
154 let mut decoder = SseDecoder::default();
155 let events = decoder.push(bytes)?;
156 decoder.finish()?;
157 Ok(events)
158}
159
160fn streaming_response(
161 body: GenerationBody,
162 renderer: Renderer,
163 traffic: Option<Arc<TrafficCapture>>,
164) -> Response {
165 let outcome = NativeResponseOutcome::default();
166 let state = StreamState {
167 body: match body {
168 GenerationBody::BufferedSse(bytes) => Body::from(bytes),
169 GenerationBody::LiveSse(body) => body,
170 },
171 decoder: SseDecoder::default(),
172 renderer,
173 pending: VecDeque::new(),
174 finished: false,
175 bytes: 0,
176 outcome: outcome.clone(),
177 traffic,
178 downstream: Vec::new(),
179 downstream_truncated: false,
180 capture_finished: false,
181 };
182 let stream = futures_util::stream::unfold(state, |mut state| async move {
183 state
184 .next()
185 .await
186 .map(|bytes| (Ok::<Bytes, Infallible>(bytes), state))
187 });
188 let mut response = (
189 [
190 (http::header::CONTENT_TYPE, "text/event-stream"),
191 (http::header::CACHE_CONTROL, "no-cache"),
192 (http::header::CONNECTION, "keep-alive"),
193 ],
194 Body::from_stream(stream),
195 )
196 .into_response();
197 response.extensions_mut().insert(outcome);
198 response
199}
200
201struct StreamState {
202 body: Body,
203 decoder: SseDecoder,
204 renderer: Renderer,
205 pending: VecDeque<Bytes>,
206 finished: bool,
207 bytes: usize,
208 outcome: NativeResponseOutcome,
209 traffic: Option<Arc<TrafficCapture>>,
210 downstream: Vec<u8>,
211 downstream_truncated: bool,
212 capture_finished: bool,
213}
214
215impl StreamState {
216 async fn next(&mut self) -> Option<Bytes> {
217 loop {
218 if let Some(bytes) = self.pending.pop_front() {
219 self.capture_bytes(&bytes);
220 return Some(bytes);
221 }
222 if self.finished {
223 let outcome = if self.outcome.failure().is_some() {
224 "failed"
225 } else {
226 "completed"
227 };
228 self.finish_capture(outcome);
229 return None;
230 }
231 match self.body.frame().await {
232 Some(Ok(frame)) => {
233 let Ok(data) = frame.into_data() else {
234 continue;
235 };
236 self.bytes = self.bytes.saturating_add(data.len());
237 if self.bytes > MAX_PROVIDER_STREAM_BYTES {
238 self.fail(OpenAiError::upstream_protocol(
239 "Provider response exceeded the size limit",
240 ));
241 continue;
242 }
243 match self.decoder.push(&data) {
244 Ok(events) => {
245 for event in events {
246 match self.renderer.render(&event) {
247 Ok(frames) => self.pending.extend(frames),
248 Err(error) => {
249 self.fail(error);
250 break;
251 }
252 }
253 }
254 }
255 Err(error) => self.fail(error),
256 }
257 }
258 Some(Err(error)) => self.fail(OpenAiError {
259 status: StatusCode::BAD_GATEWAY,
260 kind: "api_error".into(),
261 message: format!("Provider stream read failed: {error}").into(),
262 param: None,
263 code: None,
264 retry_after: None,
265 }),
266 None => {
267 if let Err(error) = self.decoder.finish() {
268 self.fail(error);
269 } else if !self.renderer.state.stopped {
270 self.fail(upstream_invalid(
271 "Provider stream ended before message_stop",
272 None::<String>,
273 ));
274 } else {
275 self.finished = true;
276 }
277 }
278 }
279 }
280 }
281
282 fn capture_bytes(&mut self, bytes: &[u8]) {
283 let remaining = MAX_SSE_CAPTURE_BYTES.saturating_sub(self.downstream.len());
284 self.downstream
285 .extend_from_slice(&bytes[..bytes.len().min(remaining)]);
286 if bytes.len() > remaining {
287 self.downstream_truncated = true;
288 }
289 }
290
291 fn finish_capture(&mut self, outcome: &str) {
292 if self.capture_finished {
293 return;
294 }
295 self.capture_finished = true;
296 if let Some(traffic) = self.traffic.as_ref() {
297 traffic.write_bytes("071-openai-downstream.sse", &self.downstream);
298 traffic.write_json(
299 "072-openai-stream-summary",
300 &json!({
301 "outcome":outcome,
302 "capturedBytes":self.downstream.len(),
303 "truncated":self.downstream_truncated,
304 }),
305 );
306 }
307 }
308
309 fn fail(&mut self, error: OpenAiError) {
310 self.outcome.fail(error.message.to_string());
311 self.pending.extend(self.renderer.failure(&error));
312 self.finished = true;
313 }
314}
315
316impl Drop for StreamState {
317 fn drop(&mut self) {
318 if !self.capture_finished {
319 self.finish_capture("abandoned");
320 }
321 }
322}
323
324struct Renderer {
325 surface: OpenAiSurface,
326 include_usage: bool,
327 response_id: String,
328 model: String,
329 created: u64,
330 sequence: u64,
331 state: AnthropicAccumulator,
332 response_metadata: OpenAiResponseMetadata,
333}
334
335impl Renderer {
336 fn new(
337 surface: OpenAiSurface,
338 include_usage: bool,
339 response_id: String,
340 model: String,
341 created: u64,
342 response_metadata: OpenAiResponseMetadata,
343 ) -> Self {
344 Self {
345 surface,
346 include_usage,
347 response_id,
348 model,
349 created,
350 sequence: 0,
351 state: AnthropicAccumulator::default(),
352 response_metadata,
353 }
354 }
355
356 fn render(&mut self, event: &SseEvent) -> Result<Vec<Bytes>, OpenAiError> {
357 self.state.apply(event)?;
358 match self.surface {
359 OpenAiSurface::ChatCompletions => self.render_chat(event),
360 OpenAiSurface::Responses => self.render_responses(event),
361 }
362 }
363
364 fn render_chat(&self, event: &SseEvent) -> Result<Vec<Bytes>, OpenAiError> {
365 let kind = event
366 .data
367 .get("type")
368 .and_then(Value::as_str)
369 .unwrap_or_default();
370 let mut out = Vec::new();
371 match kind {
372 "message_start" => out.push(chat_data(json!({
373 "id":self.response_id,
374 "object":"chat.completion.chunk",
375 "created":self.created,
376 "model":self.model,
377 "choices":[{"index":0,"delta":{"role":"assistant","content":""},"finish_reason":null,"logprobs":null}],
378 }))),
379 "content_block_start" => {
380 let index = event.data.get("index").and_then(Value::as_u64).unwrap_or_default();
381 if let Some(block) = self.state.blocks.iter().rev().find(|block| block.index == index as usize)
382 && let BlockKind::Tool { id, name, .. } = &block.kind
383 {
384 let tool_index = self.tool_index(index as usize);
385 out.push(chat_data(json!({
386 "id":self.response_id,
387 "object":"chat.completion.chunk",
388 "created":self.created,
389 "model":self.model,
390 "choices":[{"index":0,"delta":{"tool_calls":[{"index":tool_index,"id":id,"type":"function","function":{"name":name,"arguments":""}}]},"finish_reason":null,"logprobs":null}],
391 })));
392 }
393 }
394 "content_block_delta" => {
395 let index = event.data.get("index").and_then(Value::as_u64).unwrap_or_default();
396 let delta = event.data.get("delta").unwrap_or(&Value::Null);
397 let payload = match delta.get("type").and_then(Value::as_str) {
398 Some("text_delta") => json!({"content":delta.get("text").and_then(Value::as_str).unwrap_or_default()}),
399 Some("thinking_delta") => json!({"reasoning_content":delta.get("thinking").and_then(Value::as_str).unwrap_or_default()}),
400 Some("input_json_delta") => {
401 if self.state.blocks.iter().any(|block| {
402 block.index == index as usize
403 && matches!(block.kind, BlockKind::Tool { .. })
404 }) {
405 let tool_index = self.tool_index(index as usize);
406 json!({"tool_calls":[{"index":tool_index,"function":{"arguments":delta.get("partial_json").and_then(Value::as_str).unwrap_or_default()}}]})
407 } else {
408 return Ok(out);
409 }
410 }
411 Some("citations_delta") => json!({"annotations":self.state.citations.iter().map(responses_citation).collect::<Vec<_>>().last().map(chat_citation).into_iter().collect::<Vec<_>>()}),
412 Some("signature_delta") => return Ok(out),
413 _ => return Err(upstream_invalid("Unsupported provider content delta", None::<String>)),
414 };
415 out.push(chat_data(json!({
416 "id":self.response_id,
417 "object":"chat.completion.chunk",
418 "created":self.created,
419 "model":self.model,
420 "choices":[{"index":0,"delta":payload,"finish_reason":null,"logprobs":null}],
421 })));
422 }
423 "message_delta" => {
424 let mut chunk = json!({
425 "id":self.response_id,
426 "object":"chat.completion.chunk",
427 "created":self.created,
428 "model":self.model,
429 "choices":[{"index":0,"delta":{},"finish_reason":chat_finish_reason(self.state.stop_reason.as_deref()),"logprobs":null}],
430 });
431 if self.include_usage {
432 chunk["usage"] = self.state.usage.chat_value();
433 }
434 out.push(chat_data(chunk));
435 }
436 "message_stop" => out.push(Bytes::from_static(b"data: [DONE]\n\n")),
437 _ => {}
438 }
439 Ok(out)
440 }
441
442 fn render_responses(&mut self, event: &SseEvent) -> Result<Vec<Bytes>, OpenAiError> {
443 let kind = event
444 .data
445 .get("type")
446 .and_then(Value::as_str)
447 .unwrap_or_default();
448 let mut out = Vec::new();
449 match kind {
450 "message_start" => {
451 let response = response_shell(
452 &self.response_id,
453 &self.model,
454 self.created,
455 "in_progress",
456 &self.response_metadata,
457 );
458 out.push(self.responses_event("response.created", json!({"response":response})));
459 out.push(
460 self.responses_event("response.in_progress", json!({"response":response})),
461 );
462 }
463 "content_block_start" => {
464 let index = event
465 .data
466 .get("index")
467 .and_then(Value::as_u64)
468 .unwrap_or_default() as usize;
469 let output_index = self.output_index(index);
470 let block = self
471 .state
472 .blocks
473 .iter()
474 .rev()
475 .find(|block| block.index == index)
476 .cloned()
477 .ok_or_else(|| {
478 upstream_invalid("Provider content block is missing", None::<String>)
479 })?;
480 match &block.kind {
481 BlockKind::Text { .. } => {
482 let item_id = self.message_item_id(index);
483 out.push(self.responses_event("response.output_item.added", json!({
484 "output_index":output_index,
485 "item":{"id":item_id,"type":"message","role":"assistant","status":"in_progress","content":[]},
486 })));
487 out.push(self.responses_event("response.content_part.added", json!({
488 "item_id":item_id,"output_index":output_index,"content_index":0,
489 "part":{"type":"output_text","text":"","annotations":[]},
490 })));
491 }
492 BlockKind::Thinking { .. } => {
493 out.push(self.responses_event("response.output_item.added", json!({
494 "output_index":output_index,
495 "item":{"id":self.block_item_id("rs", index),"type":"reasoning","status":"in_progress","summary":[]},
496 })));
497 out.push(self.responses_event("response.reasoning_summary_part.added", json!({
498 "item_id":self.block_item_id("rs", index),"output_index":output_index,"summary_index":0,
499 "part":{"type":"summary_text","text":""},
500 })));
501 }
502 BlockKind::Tool { id, name, .. } => out.push(self.responses_event("response.output_item.added", json!({
503 "output_index":output_index,
504 "item":{"id":self.block_item_id("fc", index),"type":"function_call","call_id":id,"name":name,"arguments":"","status":"in_progress"},
505 }))),
506 BlockKind::HostedSearch { id, .. } => out.push(self.responses_event("response.output_item.added", json!({
507 "output_index":output_index,
508 "item":{"id":id,"type":"web_search_call","status":"in_progress"},
509 }))),
510 BlockKind::HostedResult => {}
511 }
512 }
513 "content_block_delta" => {
514 let index = event
515 .data
516 .get("index")
517 .and_then(Value::as_u64)
518 .unwrap_or_default() as usize;
519 let output_index = self.output_index(index);
520 let delta = event.data.get("delta").unwrap_or(&Value::Null);
521 match delta.get("type").and_then(Value::as_str) {
522 Some("text_delta") => out.push(self.responses_event("response.output_text.delta", json!({
523 "item_id":self.message_item_id(index),"output_index":output_index,"content_index":0,
524 "delta":delta.get("text").and_then(Value::as_str).unwrap_or_default(),"logprobs":[],
525 }))),
526 Some("thinking_delta") => out.push(self.responses_event("response.reasoning_summary_text.delta", json!({
527 "item_id":self.block_item_id("rs", index),"output_index":output_index,"summary_index":0,
528 "delta":delta.get("thinking").and_then(Value::as_str).unwrap_or_default(),
529 }))),
530 Some("input_json_delta") => {
531 if self.state.blocks.iter().any(|block| {
532 block.index == index && matches!(block.kind, BlockKind::Tool { .. })
533 }) {
534 out.push(self.responses_event("response.function_call_arguments.delta", json!({
535 "item_id":self.block_item_id("fc", index),"output_index":output_index,
536 "delta":delta.get("partial_json").and_then(Value::as_str).unwrap_or_default(),
537 })));
538 }
539 }
540 Some("citations_delta") => out.push(self.responses_event("response.output_text.annotation.added", json!({
541 "item_id":self.message_item_id(index),"output_index":output_index,"content_index":0,
542 "annotation_index":self.state.citations.len().saturating_sub(1),
543 "annotation":self.state.citations.last().map(responses_citation).unwrap_or(Value::Null),
544 }))),
545 Some("signature_delta") => {}
546 _ => return Err(upstream_invalid("Unsupported provider content delta", None::<String>)),
547 }
548 }
549 "content_block_stop" => {
550 let index = event
551 .data
552 .get("index")
553 .and_then(Value::as_u64)
554 .unwrap_or_default() as usize;
555 let output_index = self.output_index(index);
556 let block = self
557 .state
558 .blocks
559 .iter()
560 .rev()
561 .find(|block| block.index == index)
562 .cloned()
563 .ok_or_else(|| {
564 upstream_invalid("Provider content block is missing", None::<String>)
565 })?;
566 match block.kind {
567 BlockKind::Text { text } => {
568 out.push(self.responses_event("response.output_text.done", json!({
569 "item_id":self.message_item_id(index),"output_index":output_index,"content_index":0,"text":text,"logprobs":[],
570 })));
571 out.push(self.responses_event("response.content_part.done", json!({
572 "item_id":self.message_item_id(index),"output_index":output_index,"content_index":0,
573 "part":{"type":"output_text","text":text,"annotations":self.state.citations.iter().map(responses_citation).collect::<Vec<_>>()},
574 })));
575 out.push(self.responses_event("response.output_item.done", json!({
576 "output_index":output_index,
577 "item":{"id":self.message_item_id(index),"type":"message","role":"assistant","status":"completed","content":[{"type":"output_text","text":text,"annotations":self.state.citations.iter().map(responses_citation).collect::<Vec<_>>()}]},
578 })));
579 }
580 BlockKind::Thinking { text } => {
581 out.push(self.responses_event("response.reasoning_summary_text.done", json!({
582 "item_id":self.block_item_id("rs", index),"output_index":output_index,"summary_index":0,"text":text,
583 })));
584 out.push(self.responses_event("response.reasoning_summary_part.done", json!({
585 "item_id":self.block_item_id("rs", index),"output_index":output_index,"summary_index":0,
586 "part":{"type":"summary_text","text":text},
587 })));
588 out.push(self.responses_event("response.output_item.done", json!({
589 "output_index":output_index,
590 "item":{"id":self.block_item_id("rs", index),"type":"reasoning","status":"completed","summary":[{"type":"summary_text","text":text}]},
591 })));
592 }
593 BlockKind::Tool {
594 id,
595 name,
596 arguments,
597 } => {
598 let arguments = normalized_arguments(&arguments);
599 out.push(self.responses_event("response.function_call_arguments.done", json!({
600 "item_id":self.block_item_id("fc", index),"output_index":output_index,"arguments":arguments,
601 })));
602 out.push(self.responses_event("response.output_item.done", json!({
603 "output_index":output_index,
604 "item":{"id":self.block_item_id("fc", index),"type":"function_call","call_id":id,"name":name,"arguments":arguments,"status":"completed"},
605 })));
606 }
607 BlockKind::HostedSearch {
608 id,
609 name,
610 arguments,
611 } => out.push(self.responses_event("response.output_item.done", json!({
612 "output_index":output_index,
613 "item":{"id":id,"type":"web_search_call","status":"completed","action":hosted_search_action(&name, &arguments)},
614 }))),
615 BlockKind::HostedResult => {}
616 }
617 }
618 "message_stop" => {
619 let response = responses_response(
620 &self.state,
621 &self.response_id,
622 &self.model,
623 self.created,
624 &self.response_metadata,
625 );
626 let kind = if self.state.stop_reason.as_deref() == Some("max_tokens") {
627 "response.incomplete"
628 } else {
629 "response.completed"
630 };
631 out.push(self.responses_event(kind, json!({"response":response})));
632 }
633 _ => {}
634 }
635 Ok(out)
636 }
637
638 fn failure(&mut self, error: &OpenAiError) -> Vec<Bytes> {
639 match self.surface {
640 OpenAiSurface::ChatCompletions => vec![
641 chat_data(json!({
642 "error":{"message":error.message,"type":error.kind,"param":error.param,"code":error.code}
643 })),
644 Bytes::from_static(b"data: [DONE]\n\n"),
645 ],
646 OpenAiSurface::Responses => {
647 let mut response = response_shell(
648 &self.response_id,
649 &self.model,
650 self.created,
651 "failed",
652 &self.response_metadata,
653 );
654 response["error"] = json!({"code":error.code,"message":error.message});
655 vec![self.responses_event("response.failed", json!({"response":response}))]
656 }
657 }
658 }
659
660 fn responses_event(&mut self, kind: &str, fields: Value) -> Bytes {
661 let sequence = self.sequence;
662 self.sequence = self.sequence.saturating_add(1);
663 let mut value = fields.as_object().cloned().unwrap_or_default();
664 value.insert("type".to_string(), Value::String(kind.to_string()));
665 value.insert("sequence_number".to_string(), json!(sequence));
666 named_sse(kind, Value::Object(value))
667 }
668
669 fn message_item_id(&self, block_index: usize) -> String {
670 self.block_item_id("msg", block_index)
671 }
672
673 fn output_index(&self, block_index: usize) -> usize {
674 self.state
675 .blocks
676 .iter()
677 .filter(|block| {
678 block.index < block_index && !matches!(block.kind, BlockKind::HostedResult)
679 })
680 .count()
681 }
682
683 fn tool_index(&self, block_index: usize) -> usize {
684 self.state
685 .blocks
686 .iter()
687 .filter(|block| {
688 block.index < block_index && matches!(block.kind, BlockKind::Tool { .. })
689 })
690 .count()
691 }
692
693 fn block_item_id(&self, prefix: &str, index: usize) -> String {
694 format!(
695 "{prefix}_{}_{index}",
696 self.response_id.trim_start_matches("resp_")
697 )
698 }
699}
700
701fn chat_data(value: Value) -> Bytes {
702 Bytes::from(format!("data: {}\n\n", value))
703}
704
705fn named_sse(kind: &str, value: Value) -> Bytes {
706 Bytes::from(format!("event: {kind}\ndata: {value}\n\n"))
707}
708
709fn response_shell(
710 id: &str,
711 model: &str,
712 created: u64,
713 status: &str,
714 response_metadata: &OpenAiResponseMetadata,
715) -> Value {
716 json!({
717 "id":id,
718 "object":"response",
719 "created_at":created,
720 "status":status,
721 "model":model,
722 "output":[],
723 "parallel_tool_calls":false,
724 "tool_choice":response_metadata.tool_choice,
725 "tools":response_metadata.tools,
726 "error":null,
727 "incomplete_details":null,
728 "usage":null,
729 })
730}
731
732fn parse_frame(frame: &[u8]) -> Result<Option<SseEvent>, OpenAiError> {
733 let text = std::str::from_utf8(frame)
734 .map_err(|_| OpenAiError::upstream_protocol("Provider SSE event is not UTF-8"))?;
735 let mut event = None;
736 let mut data = Vec::new();
737 for line in text.lines() {
738 let line = line.trim_end_matches('\r');
739 if line.starts_with(':') || line.is_empty() {
740 continue;
741 }
742 if let Some(value) = line.strip_prefix("event:") {
743 event = Some(value.trim_start().to_string());
744 } else if let Some(value) = line.strip_prefix("data:") {
745 data.push(value.trim_start());
746 }
747 }
748 if data.is_empty() {
749 return Ok(None);
750 }
751 let data = data.join("\n");
752 if data == "[DONE]" {
753 return Ok(None);
754 }
755 let data = serde_json::from_str(&data).map_err(|error| {
756 OpenAiError::upstream_protocol(format!("Provider SSE event contains invalid JSON: {error}"))
757 })?;
758 Ok(Some(SseEvent { event, data }))
759}
760
761fn find_event_delimiter(bytes: &[u8]) -> Option<(usize, usize)> {
762 bytes
763 .windows(2)
764 .position(|window| window == b"\n\n")
765 .map(|position| (position, 2))
766 .or_else(|| {
767 bytes
768 .windows(4)
769 .position(|window| window == b"\r\n\r\n")
770 .map(|position| (position, 4))
771 })
772}
773
774fn current_seconds() -> u64 {
775 use std::time::{SystemTime, UNIX_EPOCH};
776 SystemTime::now()
777 .duration_since(UNIX_EPOCH)
778 .unwrap_or_default()
779 .as_secs()
780}
781
782#[cfg(test)]
783mod tests {
784 use super::*;
785 use http_body_util::BodyExt;
786
787 #[test]
788 fn decoder_handles_fragmented_and_batched_events() {
789 let mut decoder = SseDecoder::default();
790 assert!(
791 decoder
792 .push(b"event: message_start\nda")
793 .unwrap()
794 .is_empty()
795 );
796 let events = decoder
797 .push(b"ta: {\"type\":\"message_start\",\"message\":{}}\n\nevent: ping\ndata: {\"type\":\"ping\"}\n\n")
798 .unwrap();
799 assert_eq!(events.len(), 2);
800 assert_eq!(events[0].event.as_deref(), Some("message_start"));
801 decoder.finish().unwrap();
802 }
803
804 #[test]
805 fn decoder_accepts_large_batches_of_small_events() {
806 let event = "event: ping\ndata: {\"type\":\"ping\"}\n\n";
807 let count = MAX_SSE_EVENT_BYTES / event.len() + 1;
808 let mut decoder = SseDecoder::default();
809 let events = decoder.push(event.repeat(count).as_bytes()).unwrap();
810 assert_eq!(events.len(), count);
811 decoder.finish().unwrap();
812 }
813
814 #[tokio::test]
815 async fn chat_stream_emits_tool_deltas_usage_and_done() {
816 let input = concat!(
817 "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"id\":\"msg_1\",\"usage\":{\"input_tokens\":2}}}\n\n",
818 "event: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"tool_use\",\"id\":\"call_1\",\"name\":\"lookup\",\"input\":{}}}\n\n",
819 "event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"input_json_delta\",\"partial_json\":\"{}\"}}\n\n",
820 "event: content_block_stop\ndata: {\"type\":\"content_block_stop\",\"index\":0}\n\n",
821 "event: message_delta\ndata: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"tool_use\"},\"usage\":{\"output_tokens\":1}}\n\n",
822 "event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n",
823 );
824 let response = streaming_response(
825 GenerationBody::BufferedSse(Bytes::from_static(input.as_bytes())),
826 Renderer::new(
827 OpenAiSurface::ChatCompletions,
828 true,
829 "chatcmpl_test".into(),
830 "kimi-k2.6".into(),
831 1,
832 OpenAiResponseMetadata::default(),
833 ),
834 None,
835 );
836 let bytes = response.into_body().collect().await.unwrap().to_bytes();
837 let text = String::from_utf8(bytes.to_vec()).unwrap();
838 assert!(text.contains("tool_calls"));
839 assert!(text.contains("\"finish_reason\":\"tool_calls\""));
840 assert!(text.contains("\"total_tokens\":3"));
841 assert!(text.ends_with("data: [DONE]\n\n"));
842 }
843
844 #[tokio::test]
845 async fn responses_stream_numbers_events_and_completes() {
846 let input = concat!(
847 "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"id\":\"msg_1\",\"usage\":{}}}\n\n",
848 "event: content_block_start\ndata: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n",
849 "event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"hi\"}}\n\n",
850 "event: content_block_stop\ndata: {\"type\":\"content_block_stop\",\"index\":0}\n\n",
851 "event: message_delta\ndata: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"end_turn\"},\"usage\":{\"output_tokens\":1}}\n\n",
852 "event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n",
853 );
854 let response = streaming_response(
855 GenerationBody::BufferedSse(Bytes::from_static(input.as_bytes())),
856 Renderer::new(
857 OpenAiSurface::Responses,
858 false,
859 "resp_test".into(),
860 "grok-4.5".into(),
861 1,
862 OpenAiResponseMetadata::default(),
863 ),
864 None,
865 );
866 let bytes = response.into_body().collect().await.unwrap().to_bytes();
867 let text = String::from_utf8(bytes.to_vec()).unwrap();
868 assert!(text.contains("event: response.created"));
869 assert!(text.contains("event: response.output_text.delta"));
870 assert!(text.contains("event: response.completed"));
871 assert!(text.contains("\"sequence_number\":0"));
872 }
873}