Skip to main content

claude_codex/providers/grok/translate/
stream.rs

1use super::reducer::{Reducer, ReducerEvent};
2use crate::anthropic::sse::{SseEvent, encode_sse_event};
3
4pub const MAX_SSE_FRAME_BYTES: usize = 1024 * 1024;
5
6#[derive(Default)]
7pub struct SseDecoder {
8    frame: Vec<u8>,
9    line_start: usize,
10    skip_lf: bool,
11}
12
13impl SseDecoder {
14    pub fn push(&mut self, input: &[u8]) -> anyhow::Result<Vec<SseEvent>> {
15        let mut events = Vec::new();
16        for &byte in input {
17            if self.skip_lf {
18                self.skip_lf = false;
19                if byte == b'\n' {
20                    continue;
21                }
22            }
23            match byte {
24                b'\n' => self.end_line(&mut events)?,
25                b'\r' => {
26                    self.end_line(&mut events)?;
27                    self.skip_lf = true;
28                }
29                _ => self.push_byte(byte)?,
30            }
31        }
32        Ok(events)
33    }
34
35    pub fn finish(&mut self) -> anyhow::Result<()> {
36        if self.frame.is_empty() {
37            Ok(())
38        } else {
39            anyhow::bail!("Grok SSE stream ended with an incomplete frame")
40        }
41    }
42
43    fn push_byte(&mut self, byte: u8) -> anyhow::Result<()> {
44        if self.frame.len() >= MAX_SSE_FRAME_BYTES {
45            anyhow::bail!("Grok SSE frame exceeds the size limit");
46        }
47        self.frame.push(byte);
48        Ok(())
49    }
50
51    fn end_line(&mut self, events: &mut Vec<SseEvent>) -> anyhow::Result<()> {
52        if self.frame.len() == self.line_start {
53            if !self.frame.is_empty()
54                && let Some(event) = parse_frame(&self.frame)?
55            {
56                events.push(event);
57            }
58            self.frame.clear();
59            self.line_start = 0;
60            return Ok(());
61        }
62        self.push_byte(b'\n')?;
63        self.line_start = self.frame.len();
64        Ok(())
65    }
66}
67
68fn parse_frame(frame: &[u8]) -> anyhow::Result<Option<SseEvent>> {
69    let frame = std::str::from_utf8(frame)
70        .map_err(|_| anyhow::anyhow!("Grok SSE frame contains invalid UTF-8"))?;
71    let mut event = None;
72    let mut data = Vec::new();
73    for line in frame.lines() {
74        if line.starts_with(':') {
75            continue;
76        }
77        let (field, value) = line.split_once(':').unwrap_or((line, ""));
78        let value = value.strip_prefix(' ').unwrap_or(value);
79        match field {
80            "event" => event = Some(value.to_owned()),
81            "data" => data.push(value),
82            _ => {}
83        }
84    }
85    if data.is_empty() {
86        return Ok(None);
87    }
88    Ok(Some(SseEvent {
89        event,
90        data: data.join("\n"),
91    }))
92}
93
94pub struct StreamTranslator {
95    message_id: String,
96    model: String,
97    started: bool,
98    finished: bool,
99}
100
101pub struct LiveStreamTranslator {
102    decoder: SseDecoder,
103    reducer: Reducer,
104    renderer: StreamTranslator,
105}
106
107impl LiveStreamTranslator {
108    pub fn new(message_id: String, model: String) -> Self {
109        Self {
110            decoder: SseDecoder::default(),
111            reducer: Reducer::default(),
112            renderer: StreamTranslator::new(message_id, model),
113        }
114    }
115
116    pub fn push(&mut self, chunk: &[u8]) -> anyhow::Result<Vec<u8>> {
117        let mut out = Vec::new();
118        for event in self.decoder.push(chunk)? {
119            let value = serde_json::from_str(&event.data)
120                .map_err(|_| anyhow::anyhow!("malformed Grok SSE event"))?;
121            out.extend(self.renderer.render(self.reducer.push(value)?)?);
122        }
123        Ok(out)
124    }
125
126    pub fn finish(mut self) -> anyhow::Result<()> {
127        self.decoder.finish()?;
128        if !self.reducer.finished() {
129            anyhow::bail!("Grok stream ended without completion");
130        }
131        Ok(())
132    }
133}
134
135impl StreamTranslator {
136    pub fn new(message_id: String, model: String) -> Self {
137        Self {
138            message_id,
139            model,
140            started: false,
141            finished: false,
142        }
143    }
144
145    pub fn render(&mut self, events: Vec<ReducerEvent>) -> anyhow::Result<Vec<u8>> {
146        if self.finished && !events.is_empty() {
147            anyhow::bail!("event after terminal completion");
148        }
149        let mut out = Vec::new();
150        for event in events {
151            if !self.started
152                && matches!(
153                    event,
154                    ReducerEvent::ThinkingStart(_)
155                        | ReducerEvent::TextStart(_)
156                        | ReducerEvent::ToolStart(_, _, _)
157                        | ReducerEvent::HostedSearch { .. }
158                        | ReducerEvent::Finish { .. }
159                )
160            {
161                self.started = true;
162                emit(
163                    &mut out,
164                    "message_start",
165                    serde_json::json!({"type":"message_start","message":{"id":self.message_id,"type":"message","role":"assistant","model":self.model,"content":[],"stop_reason":null,"stop_sequence":null,"usage":{"input_tokens":0,"output_tokens":0}}}),
166                );
167            }
168            if matches!(event, ReducerEvent::Finish { .. }) {
169                self.finished = true;
170            }
171            render(&mut out, event);
172        }
173        Ok(out)
174    }
175}
176
177pub fn translate_stream_bytes(
178    upstream: &[u8],
179    message_id: &str,
180    model: &str,
181) -> anyhow::Result<Vec<u8>> {
182    let mut decoder = SseDecoder::default();
183    let mut reducer = Reducer::default();
184    let mut translator = StreamTranslator::new(message_id.into(), model.into());
185    let mut out = Vec::new();
186    for event in decoder.push(upstream)? {
187        let value = serde_json::from_str(&event.data)
188            .map_err(|_| anyhow::anyhow!("malformed Grok SSE event"))?;
189        out.extend(translator.render(reducer.push(value)?)?);
190    }
191    decoder.finish()?;
192    if !reducer.finished() {
193        anyhow::bail!("Grok stream ended without completion");
194    }
195    Ok(out)
196}
197
198pub fn stream_error() -> Vec<u8> {
199    let data = serde_json::json!({"type":"error","error":{"type":"api_error","message":"Grok stream is invalid"}});
200    encode_sse_event(Some("error"), &data.to_string())
201}
202
203fn emit(out: &mut Vec<u8>, event: &str, data: serde_json::Value) {
204    out.extend(encode_sse_event(Some(event), &data.to_string()));
205}
206
207fn render(out: &mut Vec<u8>, event: ReducerEvent) {
208    match event {
209        ReducerEvent::ThinkingStart(i) => emit(
210            out,
211            "content_block_start",
212            serde_json::json!({"type":"content_block_start","index":i,"content_block":{"type":"thinking","thinking":"","signature":""}}),
213        ),
214        ReducerEvent::ThinkingDelta(i, t) => emit(
215            out,
216            "content_block_delta",
217            serde_json::json!({"type":"content_block_delta","index":i,"delta":{"type":"thinking_delta","thinking":t}}),
218        ),
219        ReducerEvent::ThinkingStop(i) | ReducerEvent::TextStop(i) | ReducerEvent::ToolStop(i) => {
220            emit(
221                out,
222                "content_block_stop",
223                serde_json::json!({"type":"content_block_stop","index":i}),
224            )
225        }
226        ReducerEvent::TextStart(i) => emit(
227            out,
228            "content_block_start",
229            serde_json::json!({"type":"content_block_start","index":i,"content_block":{"type":"text","text":""}}),
230        ),
231        ReducerEvent::TextDelta(i, t) => emit(
232            out,
233            "content_block_delta",
234            serde_json::json!({"type":"content_block_delta","index":i,"delta":{"type":"text_delta","text":t}}),
235        ),
236        ReducerEvent::ToolStart(i, id, name) => emit(
237            out,
238            "content_block_start",
239            serde_json::json!({"type":"content_block_start","index":i,"content_block":{"type":"tool_use","id":id,"name":name,"input":{}}}),
240        ),
241        ReducerEvent::ToolDelta(i, t) => emit(
242            out,
243            "content_block_delta",
244            serde_json::json!({"type":"content_block_delta","index":i,"delta":{"type":"input_json_delta","partial_json":t}}),
245        ),
246        ReducerEvent::HostedSearch {
247            index,
248            result_index,
249            id,
250            name,
251            query,
252        } => {
253            let result_type = format!("{name}_tool_result");
254            emit(
255                out,
256                "content_block_start",
257                serde_json::json!({"type":"content_block_start","index":index,"content_block":{"type":"server_tool_use","id":id,"name":name,"input":{}}}),
258            );
259            emit(
260                out,
261                "content_block_delta",
262                serde_json::json!({"type":"content_block_delta","index":index,"delta":{"type":"input_json_delta","partial_json":serde_json::json!({"query":query}).to_string()}}),
263            );
264            emit(
265                out,
266                "content_block_stop",
267                serde_json::json!({"type":"content_block_stop","index":index}),
268            );
269            emit(
270                out,
271                "content_block_start",
272                serde_json::json!({"type":"content_block_start","index":result_index,"content_block":{"type":result_type,"tool_use_id":id,"content":[]}}),
273            );
274            emit(
275                out,
276                "content_block_stop",
277                serde_json::json!({"type":"content_block_stop","index":result_index}),
278            );
279        }
280        ReducerEvent::Citation(i, annotation) => {
281            let citation = serde_json::json!({
282                "type":"web_search_result_location",
283                "url":annotation.get("url").and_then(serde_json::Value::as_str).unwrap_or_default(),
284                "title":annotation.get("title").and_then(serde_json::Value::as_str).unwrap_or_default(),
285                "cited_text":annotation.get("text").and_then(serde_json::Value::as_str).unwrap_or_default()
286            });
287            emit(
288                out,
289                "content_block_delta",
290                serde_json::json!({"type":"content_block_delta","index":i,"delta":{"type":"citations_delta","citation":citation}}),
291            );
292        }
293        ReducerEvent::Finish {
294            stop_reason,
295            output_tokens,
296            web_search_requests,
297            x_search_requests,
298            ..
299        } => {
300            let hosted_search_requests = web_search_requests + x_search_requests;
301            emit(
302                out,
303                "message_delta",
304                serde_json::json!({"type":"message_delta","delta":{"stop_reason":stop_reason,"stop_sequence":null},"usage":{"output_tokens":output_tokens,"server_tool_use":{"web_search_requests":hosted_search_requests,"x_search_requests":x_search_requests}}}),
305            );
306            emit(
307                out,
308                "message_stop",
309                serde_json::json!({"type":"message_stop"}),
310            );
311        }
312    }
313}
314
315#[cfg(test)]
316mod tests {
317    use super::*;
318
319    #[test]
320    fn decoder_accepts_every_boundary_and_line_ending() {
321        let input = b": note\r\nevent: ignored\r\ndata: first\r\ndata: second\r\n\r\n";
322        let expected = vec![SseEvent {
323            event: Some("ignored".into()),
324            data: "first\nsecond".into(),
325        }];
326        for split in 0..=input.len() {
327            let mut decoder = SseDecoder::default();
328            let mut events = decoder.push(&input[..split]).unwrap();
329            events.extend(decoder.push(&input[split..]).unwrap());
330            decoder.finish().unwrap();
331            assert_eq!(events, expected);
332        }
333    }
334
335    #[test]
336    fn decoder_ignores_data_less_frames_at_every_boundary() {
337        let input = b": keepalive\n\nid: 42\n\nevent: ignored\n\nretry: 5000\n\ndata: complete\n\n";
338        let expected = vec![SseEvent {
339            event: None,
340            data: "complete".into(),
341        }];
342        for split in 0..=input.len() {
343            let mut decoder = SseDecoder::default();
344            let mut events = decoder.push(&input[..split]).unwrap();
345            events.extend(decoder.push(&input[split..]).unwrap());
346            decoder.finish().unwrap();
347            assert_eq!(events, expected);
348        }
349    }
350
351    #[test]
352    fn decoder_requires_terminated_valid_frames_and_bounds_them() {
353        assert!(SseDecoder::default().push(b"data: \xff\n\n").is_err());
354        let mut decoder = SseDecoder::default();
355        decoder.push(b"data: incomplete").unwrap();
356        assert!(decoder.finish().is_err());
357        let mut decoder = SseDecoder::default();
358        let exact = vec![b'x'; MAX_SSE_FRAME_BYTES - b"data: \n".len()];
359        assert!(decoder.push(b"data: ").is_ok());
360        assert!(decoder.push(&exact).is_ok());
361        let events = decoder.push(b"\n\n").unwrap();
362        assert_eq!(events[0].data.len(), exact.len());
363        decoder.finish().unwrap();
364        let mut decoder = SseDecoder::default();
365        assert!(decoder.push(b"data: ").is_ok());
366        assert!(decoder.push(&vec![b'x'; exact.len() + 1]).is_ok());
367        assert!(decoder.push(b"\n").is_err());
368    }
369
370    #[test]
371    fn stream_translates_hosted_web_search_and_citations() {
372        let input = b"data: {\"type\":\"response.output_item.added\",\"item\":{\"type\":\"web_search_call\",\"id\":\"ws_1\"}}\n\ndata: {\"type\":\"response.web_search_call.in_progress\",\"item_id\":\"ws_1\"}\n\ndata: {\"type\":\"response.web_search_call.searching\",\"item_id\":\"ws_1\"}\n\ndata: {\"type\":\"response.web_search_call.completed\",\"item_id\":\"ws_1\"}\n\ndata: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"web_search_call\",\"id\":\"ws_1\",\"action\":{\"query\":\"rust news\"}}}\n\ndata: {\"type\":\"response.output_text.delta\",\"delta\":\"Result\"}\n\ndata: {\"type\":\"response.output_text.annotation.added\",\"annotation\":{\"type\":\"url_citation\",\"url\":\"https://example.com\",\"title\":\"Example\"}}\n\ndata: {\"type\":\"response.output_text.done\"}\n\ndata: {\"type\":\"response.completed\",\"response\":{\"usage\":{\"input_tokens\":3,\"output_tokens\":2}}}\n\n";
373        let output =
374            String::from_utf8(translate_stream_bytes(input, "msg_1", "grok-4.5").unwrap()).unwrap();
375        assert!(output.contains("server_tool_use"));
376        assert!(output.contains("web_search_tool_result"));
377        assert!(output.contains("citations_delta"));
378        assert!(output.contains("https://example.com"));
379        assert!(output.contains("\"web_search_requests\":1"));
380    }
381
382    #[test]
383    fn stream_translates_hosted_x_search_usage_and_citations() {
384        let input = b"data: {\"type\":\"response.output_item.added\",\"item\":{\"type\":\"custom_tool_call\",\"name\":\"x_search\",\"id\":\"xs_1\"}}\n\ndata: {\"type\":\"response.custom_tool_call_input.delta\",\"item_id\":\"xs_1\",\"delta\":\"{\\\"query\\\":\\\"claude-code-proxy\\\"}\"}\n\ndata: {\"type\":\"response.custom_tool_call_input.done\",\"item_id\":\"xs_1\"}\n\ndata: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"custom_tool_call\",\"name\":\"x_search\",\"id\":\"xs_1\"}}\n\ndata: {\"type\":\"response.output_text.delta\",\"delta\":\"Recent post\"}\n\ndata: {\"type\":\"response.output_text.annotation.added\",\"annotation\":{\"type\":\"url_citation\",\"url\":\"https://x.com/example/status/1\",\"title\":\"Example post\"}}\n\ndata: {\"type\":\"response.output_text.done\"}\n\ndata: {\"type\":\"response.completed\",\"response\":{\"usage\":{\"input_tokens\":4,\"output_tokens\":3}}}\n\n";
385        let output =
386            String::from_utf8(translate_stream_bytes(input, "msg_1", "grok-4.5").unwrap()).unwrap();
387        assert!(output.contains("\"name\":\"x_search\""));
388        assert!(output.contains("x_search_tool_result"));
389        assert!(output.contains("https://x.com/example/status/1"));
390        assert!(output.contains("\"web_search_requests\":1"));
391        assert!(output.contains("\"x_search_requests\":1"));
392        assert!(!output.contains("\"name\":\"Bash\""));
393    }
394
395    #[test]
396    fn live_translator_emits_first_event_before_upstream_completion() {
397        let mut translator = LiveStreamTranslator::new("msg_1".into(), "grok-4.5".into());
398        let output = translator
399            .push(b"data: {\"type\":\"response.output_text.delta\",\"delta\":\"first\"}\n\n")
400            .unwrap();
401        assert!(String::from_utf8(output).unwrap().contains("first"));
402    }
403}