Skip to main content

claude_codex/openai_compat/
stream.rs

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}