Skip to main content

claude_codex/providers/grok/
mod.rs

1pub mod auth;
2pub mod client;
3pub mod count_tokens;
4pub mod translate;
5
6use std::convert::Infallible;
7use std::sync::Arc;
8use std::time::{SystemTime, UNIX_EPOCH};
9
10use async_trait::async_trait;
11use axum::{
12    Json,
13    body::Body,
14    http::StatusCode,
15    response::{IntoResponse, Response},
16};
17use bytes::Bytes;
18use futures_util::{Stream, StreamExt};
19
20use crate::anthropic::{
21    error::json_error,
22    schema::{CountTokensResponse, MessagesRequest},
23};
24use crate::monitor::MonitorHandle;
25use crate::provider::{
26    CliHandlers, Generation, GenerationBody, Provider, ProviderError, ProviderErrorKind,
27    RequestContext,
28};
29use crate::{registry::GROK_MODELS, traffic::StreamTrafficCapture};
30
31use self::auth::token_store::file_store;
32use self::translate::{
33    accumulate::accumulate_response_with_traffic,
34    model_allowlist::{assert_allowed_model, resolve_model},
35    request::translate_request,
36    stream::{SseDecoder, StreamTranslator, stream_error},
37};
38
39pub struct GrokProvider {
40    client: Arc<client::GrokClient>,
41}
42impl GrokProvider {
43    pub fn new() -> Self {
44        crate::config::warn_grok_tool_image_mode_once(&crate::logging::create_logger("grok"));
45        Self {
46            client: Arc::new(
47                client::GrokClient::new(
48                    crate::config::grok_base_url(),
49                    crate::config::grok_client_version(),
50                )
51                .expect("Grok transport is unavailable"),
52            ),
53        }
54    }
55
56    pub fn with_client(client: client::GrokClient) -> Self {
57        Self {
58            client: Arc::new(client),
59        }
60    }
61}
62impl Default for GrokProvider {
63    fn default() -> Self {
64        Self::new()
65    }
66}
67
68#[async_trait]
69impl Provider for GrokProvider {
70    fn name(&self) -> &'static str {
71        "grok"
72    }
73    fn supported_models(&self) -> Vec<String> {
74        GROK_MODELS
75            .iter()
76            .map(|model| (*model).to_string())
77            .collect()
78    }
79    fn cli(&self) -> &'static dyn CliHandlers {
80        &GROK_CLI
81    }
82    async fn handle_messages(&self, body: MessagesRequest, ctx: RequestContext) -> Response {
83        let requested = body.model.clone().unwrap_or_else(|| "grok-4.5".into());
84        let resolved = resolve_model(&requested);
85        if let Err(error) = assert_allowed_model(&resolved) {
86            return json_error(
87                StatusCode::BAD_REQUEST,
88                "invalid_request_error",
89                error.to_string(),
90            );
91        }
92        let translated = match translate_request(&body, resolved.clone()) {
93            Ok(value) => value,
94            Err(error) => {
95                return json_error(
96                    StatusCode::BAD_REQUEST,
97                    "invalid_request_error",
98                    error.to_string(),
99                );
100            }
101        };
102        if let Some(monitor) = &ctx.monitor {
103            monitor.model_resolved(&ctx.req_id, &resolved);
104            monitor.upstream_started(&ctx.req_id);
105        }
106        let upstream = match self.client.post(&translated, ctx.traffic.clone()).await {
107            Ok(response) => response,
108            Err(error) => return map_error(error),
109        };
110        if body.stream {
111            stream_response(
112                upstream,
113                format!("msg_{}", uuid::Uuid::new_v4().simple()),
114                requested,
115                ctx.monitor.clone(),
116                ctx.req_id.clone(),
117                ctx.traffic.clone(),
118            )
119        } else {
120            let upstream_bytes = match upstream.into_bytes().await {
121                Ok(bytes) => bytes,
122                Err(error) => {
123                    write_error(ctx.traffic.as_deref(), "body_read", "transport");
124                    return map_error(error);
125                }
126            };
127            match accumulate_response_with_traffic(
128                &upstream_bytes,
129                &format!("msg_{}", uuid::Uuid::new_v4().simple()),
130                &requested,
131                ctx.traffic.as_deref(),
132            ) {
133                Ok(value) => {
134                    if let Some(traffic) = ctx.traffic.as_ref() {
135                        traffic.write_json("051-downstream-response", &value);
136                    }
137                    if let Some(monitor) = ctx.monitor.as_ref() {
138                        monitor.usage_updated(
139                            &ctx.req_id,
140                            value
141                                .pointer("/usage/input_tokens")
142                                .and_then(|v| v.as_u64()),
143                            value
144                                .pointer("/usage/output_tokens")
145                                .and_then(|v| v.as_u64()),
146                        );
147                    }
148                    (StatusCode::OK, Json(value)).into_response()
149                }
150                Err(_) => {
151                    write_error(ctx.traffic.as_deref(), "accumulate", "invalid_response");
152                    json_error(
153                        StatusCode::BAD_GATEWAY,
154                        "api_error",
155                        "Grok response is invalid",
156                    )
157                }
158            }
159        }
160    }
161    async fn handle_count_tokens(&self, body: MessagesRequest, ctx: RequestContext) -> Response {
162        let requested = body.model.clone().unwrap_or_else(|| "grok-4.5".into());
163        let resolved = resolve_model(&requested);
164        if let Err(error) = assert_allowed_model(&resolved) {
165            return json_error(
166                StatusCode::BAD_REQUEST,
167                "invalid_request_error",
168                error.to_string(),
169            );
170        }
171        let translated = match translate_request(&body, resolved) {
172            Ok(value) => value,
173            Err(error) => {
174                return json_error(
175                    StatusCode::BAD_REQUEST,
176                    "invalid_request_error",
177                    error.to_string(),
178                );
179            }
180        };
181        let tokens = count_tokens::count_tokens(&translated);
182        if let Some(monitor) = ctx.monitor.as_ref() {
183            monitor.usage_updated(&ctx.req_id, Some(tokens), None);
184        }
185        (
186            StatusCode::OK,
187            Json(CountTokensResponse {
188                input_tokens: tokens,
189            }),
190        )
191            .into_response()
192    }
193
194    async fn generate_anthropic_stream(
195        &self,
196        mut body: MessagesRequest,
197        ctx: RequestContext,
198    ) -> Result<Generation, ProviderError> {
199        body.stream = true;
200        let requested = body.model.clone().unwrap_or_else(|| "grok-4.5".into());
201        let resolved = resolve_model(&requested);
202        assert_allowed_model(&resolved).map_err(|error| {
203            ProviderError::new(
204                StatusCode::BAD_REQUEST,
205                ProviderErrorKind::InvalidRequest,
206                error.to_string(),
207            )
208        })?;
209        let translated = translate_request(&body, resolved.clone()).map_err(|error| {
210            ProviderError::new(
211                StatusCode::BAD_REQUEST,
212                ProviderErrorKind::InvalidRequest,
213                error.to_string(),
214            )
215        })?;
216        if let Some(monitor) = ctx.monitor.as_ref() {
217            monitor.model_resolved(&ctx.req_id, &resolved);
218            monitor.upstream_started(&ctx.req_id);
219        }
220        let upstream = self
221            .client
222            .post(&translated, ctx.traffic.clone())
223            .await
224            .map_err(grok_provider_error)?;
225        let response = stream_response(
226            upstream,
227            format!("msg_{}", uuid::Uuid::new_v4().simple()),
228            requested,
229            ctx.monitor.clone(),
230            ctx.req_id.clone(),
231            ctx.traffic.clone(),
232        );
233        Ok(Generation {
234            body: GenerationBody::LiveSse(response.into_body()),
235            resolved_model: resolved,
236        })
237    }
238}
239
240fn stream_response(
241    response: client::GrokResponse,
242    message_id: String,
243    model: String,
244    monitor: Option<MonitorHandle>,
245    req_id: String,
246    traffic: Option<Arc<crate::traffic::TrafficCapture>>,
247) -> Response {
248    stream_body(
249        response.into_stream(),
250        message_id,
251        model,
252        monitor,
253        req_id,
254        traffic,
255    )
256}
257
258fn stream_body<S>(
259    upstream: S,
260    message_id: String,
261    model: String,
262    monitor: Option<MonitorHandle>,
263    req_id: String,
264    traffic: Option<Arc<crate::traffic::TrafficCapture>>,
265) -> Response
266where
267    S: Stream<Item = Result<Bytes, client::GrokError>> + Unpin + Send + 'static,
268{
269    let state = GrokStreamState {
270        upstream,
271        decoder: SseDecoder::default(),
272        reducer: translate::reducer::Reducer::default(),
273        translator: StreamTranslator::new(message_id, model),
274        terminal: false,
275        error_sent: false,
276        monitor,
277        req_id,
278        bytes: 0,
279        chunks: 0,
280        stream_capture: traffic.as_ref().map(|traffic| traffic.stream_capture()),
281        traffic,
282    };
283    let stream = futures_util::stream::unfold(state, |mut state| async move {
284        state
285            .next_output()
286            .await
287            .map(|bytes| (Ok::<Bytes, Infallible>(Bytes::from(bytes)), state))
288    });
289    (
290        [
291            (http::header::CONTENT_TYPE, "text/event-stream"),
292            (http::header::CACHE_CONTROL, "no-cache"),
293        ],
294        Body::from_stream(stream),
295    )
296        .into_response()
297}
298
299struct GrokStreamState<S> {
300    upstream: S,
301    decoder: SseDecoder,
302    reducer: translate::reducer::Reducer,
303    translator: StreamTranslator,
304    terminal: bool,
305    error_sent: bool,
306    monitor: Option<MonitorHandle>,
307    req_id: String,
308    bytes: u64,
309    chunks: u64,
310    stream_capture: Option<StreamTrafficCapture>,
311    traffic: Option<Arc<crate::traffic::TrafficCapture>>,
312}
313
314impl<S> GrokStreamState<S>
315where
316    S: Stream<Item = Result<Bytes, client::GrokError>> + Unpin,
317{
318    async fn next_output(&mut self) -> Option<Vec<u8>> {
319        if self.terminal {
320            return None;
321        }
322        if self.error_sent {
323            self.terminal = true;
324            return None;
325        }
326        loop {
327            let chunk = match self.upstream.next().await {
328                Some(Ok(chunk)) => chunk,
329                Some(Err(_)) => return Some(self.fail_at("transport", "upstream_stream")),
330                None => {
331                    if self.decoder.finish().is_err() || !self.reducer.finished() {
332                        return Some(self.fail_at("decoder", "incomplete_stream"));
333                    }
334                    self.terminal = true;
335                    self.finish_capture(true);
336                    return None;
337                }
338            };
339            if self.bytes == 0
340                && let Some(monitor) = self.monitor.as_ref()
341            {
342                monitor.generation_started(&self.req_id);
343            }
344            self.bytes = self.bytes.saturating_add(chunk.len() as u64);
345            self.chunks = self.chunks.saturating_add(1);
346            if let Some(monitor) = self.monitor.as_ref() {
347                monitor.stream_progress(&self.req_id, chunk.len() as u64, 1, None, None);
348            }
349            let events = match self.decoder.push(&chunk) {
350                Ok(events) => events,
351                Err(_) => return Some(self.fail_at("decoder", "malformed_sse")),
352            };
353            let mut out = Vec::new();
354            for event in events {
355                let value: serde_json::Value = match serde_json::from_str(&event.data) {
356                    Ok(value) => value,
357                    Err(_) => {
358                        if let Some(capture) = self.stream_capture.as_mut() {
359                            capture.malformed("json", "malformed_event");
360                        }
361                        return Some(self.fail_at("json", "malformed_event"));
362                    }
363                };
364                if let Some(capture) = self.stream_capture.as_mut() {
365                    capture.upstream_event(event.event.as_deref(), &value);
366                }
367                let reduced = match self.reducer.push(value) {
368                    Ok(events) => events,
369                    Err(_) => return Some(self.fail_at("reducer", "invalid_event")),
370                };
371                let usage = reduced.iter().find_map(|event| match event {
372                    translate::reducer::ReducerEvent::Finish {
373                        input_tokens,
374                        output_tokens,
375                        ..
376                    } => Some((*input_tokens, *output_tokens)),
377                    _ => None,
378                });
379                match self.translator.render(reduced) {
380                    Ok(bytes) => out.extend(bytes),
381                    Err(_) => return Some(self.fail_at("render", "invalid_event")),
382                }
383                if let Some((input_tokens, output_tokens)) = usage
384                    && let Some(monitor) = self.monitor.as_ref()
385                {
386                    monitor.usage_updated(&self.req_id, Some(input_tokens), Some(output_tokens));
387                }
388                if self.reducer.finished() {
389                    self.terminal = true;
390                    self.capture_downstream(&out);
391                    self.finish_capture(true);
392                    return if out.is_empty() { None } else { Some(out) };
393                }
394            }
395            if !out.is_empty() {
396                self.capture_downstream(&out);
397                return Some(out);
398            }
399        }
400    }
401
402    fn fail_at(&mut self, stage: &str, kind: &str) -> Vec<u8> {
403        self.error_sent = true;
404        let mut fields = serde_json::Map::new();
405        fields.insert("reqId".into(), serde_json::json!(self.req_id));
406        fields.insert("stage".into(), serde_json::json!(stage));
407        fields.insert("kind".into(), serde_json::json!(kind));
408        fields.insert("bytes".into(), serde_json::json!(self.bytes));
409        fields.insert("chunks".into(), serde_json::json!(self.chunks));
410        crate::logging::create_logger("grok").warn("grok_stream_failed", Some(fields));
411        if let Some(capture) = self.stream_capture.as_mut() {
412            capture.malformed(stage, kind);
413            capture.downstream_event("error", serde_json::json!({"type":"error","error":{"type":"api_error","message":"Grok stream is invalid"}}));
414        }
415        if let Some(traffic) = self.traffic.as_ref() {
416            traffic.write_json("060-grok-stream-error", &serde_json::json!({"stage":stage,"kind":kind,"bytes":self.bytes,"chunks":self.chunks}));
417        }
418        self.finish_capture(false);
419        stream_error()
420    }
421
422    fn capture_downstream(&mut self, bytes: &[u8]) {
423        let Some(capture) = self.stream_capture.as_mut() else {
424            return;
425        };
426        let mut decoder = SseDecoder::default();
427        if let Ok(events) = decoder.push(bytes) {
428            for event in events {
429                if let Ok(data) = serde_json::from_str(&event.data) {
430                    capture.downstream_event(event.event.as_deref().unwrap_or("message"), data);
431                }
432            }
433        }
434    }
435
436    fn finish_capture(&mut self, completed: bool) {
437        if let (Some(capture), Some(traffic)) = (self.stream_capture.take(), self.traffic.as_ref())
438        {
439            capture.finish(
440                traffic,
441                serde_json::json!({
442                    "kind": if completed { "stream_completion" } else { "stream_error" },
443                    "bytes": self.bytes,
444                    "chunks": self.chunks,
445                }),
446            );
447        }
448    }
449}
450
451impl<S> Drop for GrokStreamState<S> {
452    fn drop(&mut self) {
453        if self.terminal || self.stream_capture.is_none() {
454            return;
455        }
456        if let Some(traffic) = self.traffic.as_ref() {
457            traffic.write_json(
458                "060-grok-stream-abandoned",
459                &serde_json::json!({
460                    "stage": "downstream",
461                    "kind": "client_disconnect",
462                    "reason": "downstream_body_dropped",
463                    "bytes": self.bytes,
464                    "chunks": self.chunks,
465                }),
466            );
467        }
468        if let (Some(capture), Some(traffic)) = (self.stream_capture.take(), self.traffic.as_ref())
469        {
470            capture.finish(
471                traffic,
472                serde_json::json!({
473                    "kind": "stream_abandoned",
474                    "reason": "downstream_body_dropped",
475                    "bytes": self.bytes,
476                    "chunks": self.chunks,
477                }),
478            );
479        }
480    }
481}
482
483fn write_error(traffic: Option<&crate::traffic::TrafficCapture>, stage: &str, kind: &str) {
484    if let Some(traffic) = traffic {
485        traffic.write_json(
486            "060-grok-stream-error",
487            &serde_json::json!({"stage":stage,"kind":kind}),
488        );
489    }
490}
491
492fn grok_provider_error(error: client::GrokError) -> ProviderError {
493    let kind = match error.status {
494        StatusCode::UNAUTHORIZED => ProviderErrorKind::Authentication,
495        StatusCode::TOO_MANY_REQUESTS => ProviderErrorKind::RateLimit,
496        StatusCode::PAYMENT_REQUIRED | StatusCode::FORBIDDEN => ProviderErrorKind::Permission,
497        _ => ProviderErrorKind::Api,
498    };
499    let status = match kind {
500        ProviderErrorKind::Api => StatusCode::BAD_GATEWAY,
501        _ => error.status,
502    };
503    let mut mapped = ProviderError::new(status, kind, error.message);
504    mapped.retry_after = error.retry_after;
505    mapped
506}
507
508fn map_error(error: client::GrokError) -> Response {
509    match error.status {
510        StatusCode::UNAUTHORIZED => json_error(
511            StatusCode::UNAUTHORIZED,
512            "authentication_error",
513            error.message,
514        ),
515        StatusCode::TOO_MANY_REQUESTS => {
516            let response = json_error(
517                StatusCode::TOO_MANY_REQUESTS,
518                "rate_limit_error",
519                error.message,
520            );
521            if let Some(retry_after) = error.retry_after {
522                ([(http::header::RETRY_AFTER, retry_after)], response).into_response()
523            } else {
524                response
525            }
526        }
527        StatusCode::PAYMENT_REQUIRED | StatusCode::FORBIDDEN => {
528            json_error(error.status, "permission_error", error.message)
529        }
530        _ => json_error(StatusCode::BAD_GATEWAY, "api_error", error.message),
531    }
532}
533
534pub struct GrokCli;
535pub static GROK_CLI: GrokCli = GrokCli;
536impl CliHandlers for GrokCli {
537    fn login(&self) -> anyhow::Result<()> {
538        let store = file_store();
539        auth::login::login(&store)?;
540        println!("Grok authentication saved in {}", store.auth_path());
541        Ok(())
542    }
543    fn device(&self) -> anyhow::Result<()> {
544        let store = file_store();
545        auth::device::device_login(&store)?;
546        println!("Grok authentication saved in {}", store.auth_path());
547        Ok(())
548    }
549    fn status(&self) -> anyhow::Result<()> {
550        let store = file_store();
551        match store.load_auth()? {
552            Some(auth) => {
553                println!("Auth path: {}", store.auth_path());
554                println!("Authenticated: true");
555                println!(
556                    "Expires in {}s",
557                    auth.expires_at_ms.saturating_sub(now_ms()) / 1000
558                );
559                Ok(())
560            }
561            None => anyhow::bail!("Not authenticated"),
562        }
563    }
564    fn logout(&self) -> anyhow::Result<()> {
565        let store = file_store();
566        store.clear_auth()?;
567        println!("Grok proxy credentials removed");
568        Ok(())
569    }
570}
571fn now_ms() -> u64 {
572    SystemTime::now()
573        .duration_since(UNIX_EPOCH)
574        .unwrap_or_default()
575        .as_millis() as u64
576}
577
578#[cfg(test)]
579mod tests {
580    use std::pin::Pin;
581    use std::task::{Context, Poll};
582    use std::time::Duration;
583
584    use crate::monitor::{EndpointKind, MonitorHandle};
585    use crate::traffic::test_capture;
586    use http_body_util::BodyExt;
587    use tempfile::TempDir;
588    use tokio::sync::mpsc;
589
590    use super::*;
591
592    struct ChannelStream(mpsc::Receiver<Result<Bytes, client::GrokError>>);
593
594    impl Stream for ChannelStream {
595        type Item = Result<Bytes, client::GrokError>;
596
597        fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
598            self.0.poll_recv(cx)
599        }
600    }
601
602    #[tokio::test]
603    async fn streaming_usage_updates_completed_monitor_request() {
604        let monitor = MonitorHandle::new(10);
605        monitor.request_started(
606            "req_1",
607            Some("session_1".into()),
608            Some(1),
609            EndpointKind::Messages,
610        );
611        monitor.provider_selected("req_1", "grok", "grok-4.5", None);
612        monitor.request_completed("req_1", 200, None, None);
613
614        let upstream = futures_util::stream::iter(vec![Ok(Bytes::from_static(
615            b"data: {\"type\":\"response.output_text.delta\",\"delta\":\"ok\"}\n\ndata: {\"type\":\"response.output_text.done\"}\n\ndata: {\"type\":\"response.completed\",\"response\":{\"usage\":{\"input_tokens\":12,\"output_tokens\":3}}}\n\n",
616        ))]);
617        let response = stream_body(
618            upstream,
619            "msg_1".into(),
620            "grok-4.5".into(),
621            Some(monitor.clone()),
622            "req_1".into(),
623            None,
624        );
625        let _ = response.into_body().collect().await.unwrap();
626
627        let snapshot = monitor.snapshot();
628        let request = snapshot
629            .recent
630            .iter()
631            .find(|request| request.request_id == "req_1")
632            .unwrap();
633        assert_eq!(request.input_tokens, Some(12));
634        assert_eq!(request.output_tokens, Some(3));
635        assert!(request.streamed_bytes > 0);
636        assert!(request.stream_chunks > 0);
637        let session = snapshot
638            .sessions
639            .iter()
640            .find(|session| session.session_id.as_deref() == Some("session_1"))
641            .unwrap();
642        assert_eq!(session.input_tokens, 12);
643        assert_eq!(session.output_tokens, 3);
644    }
645
646    #[tokio::test]
647    async fn downstream_event_arrives_before_upstream_completion() {
648        let (tx, rx) = mpsc::channel(2);
649        let response = stream_body(
650            ChannelStream(rx),
651            "msg_1".into(),
652            "grok-4.5".into(),
653            None,
654            "req_1".into(),
655            None,
656        );
657        let mut body = response.into_body();
658
659        tx.send(Ok(Bytes::from_static(
660            b"data: {\"type\":\"response.output_text.delta\",\"delta\":\"first\"}\n\n",
661        )))
662        .await
663        .unwrap();
664
665        let first = tokio::time::timeout(Duration::from_millis(250), body.frame())
666            .await
667            .expect("downstream body waited for upstream completion")
668            .expect("downstream body ended before its first event")
669            .expect("downstream body frame failed")
670            .into_data()
671            .expect("first downstream frame was not data");
672        let first = String::from_utf8(first.to_vec()).unwrap();
673        assert!(first.contains("event: message_start"));
674        assert!(first.contains("first"));
675
676        tx.send(Ok(Bytes::from_static(
677            b"data: {\"type\":\"response.completed\",\"response\":{\"usage\":{}}}\n\n",
678        )))
679        .await
680        .unwrap();
681        let terminal = tokio::time::timeout(Duration::from_millis(250), body.frame())
682            .await
683            .expect("downstream completion timed out")
684            .expect("downstream completion was missing")
685            .expect("downstream completion frame failed")
686            .into_data()
687            .expect("downstream completion frame was not data");
688        assert!(
689            String::from_utf8(terminal.to_vec())
690                .unwrap()
691                .contains("event: message_stop")
692        );
693        assert!(
694            tokio::time::timeout(Duration::from_millis(250), body.frame())
695                .await
696                .expect("downstream EOF waited for upstream EOF")
697                .is_none()
698        );
699    }
700
701    #[tokio::test]
702    async fn data_less_upstream_frame_does_not_interrupt_stream() {
703        let upstream = futures_util::stream::iter(vec![
704            Ok(Bytes::from_static(
705                b"data: {\"type\":\"response.output_text.delta\",\"delta\":\"first\"}\n\n",
706            )),
707            Ok(Bytes::from_static(b": keepalive\n\n")),
708            Ok(Bytes::from_static(
709                b"data: {\"type\":\"response.completed\",\"response\":{\"usage\":{}}}\n\n",
710            )),
711        ]);
712        let response = stream_body(
713            upstream,
714            "msg_1".into(),
715            "grok-4.5".into(),
716            None,
717            "req_1".into(),
718            None,
719        );
720        let body = response.into_body().collect().await.unwrap().to_bytes();
721        let body = String::from_utf8(body.to_vec()).unwrap();
722        assert!(body.contains("first"));
723        assert!(body.contains("event: message_stop"));
724        assert!(!body.contains("event: error"));
725    }
726
727    #[tokio::test]
728    async fn dropped_downstream_body_finalizes_partial_capture() {
729        let temp = TempDir::new().unwrap();
730        let traffic = Arc::new(test_capture(temp.path().join("traffic")));
731        let (tx, rx) = mpsc::channel(2);
732        let response = stream_body(
733            ChannelStream(rx),
734            "msg_1".into(),
735            "grok-4.5".into(),
736            None,
737            "req_1".into(),
738            Some(traffic),
739        );
740        let mut body = response.into_body();
741
742        tx.send(Ok(Bytes::from_static(
743            b"data: {\"type\":\"response.output_text.delta\",\"delta\":\"partial\"}\n\n",
744        )))
745        .await
746        .unwrap();
747        let _ = body.frame().await.unwrap().unwrap();
748        drop(body);
749
750        let entries: Vec<_> = std::fs::read_dir(temp.path().join("traffic"))
751            .unwrap()
752            .map(|entry| entry.unwrap().path())
753            .collect();
754        let abandoned = entries
755            .iter()
756            .find(|path| path.to_string_lossy().contains("060-grok-stream-abandoned"))
757            .unwrap();
758        let summary = entries
759            .iter()
760            .find(|path| path.to_string_lossy().contains("061-grok-stream-summary"))
761            .unwrap();
762        let abandoned: serde_json::Value =
763            serde_json::from_slice(&std::fs::read(abandoned).unwrap()).unwrap();
764        let summary: serde_json::Value =
765            serde_json::from_slice(&std::fs::read(summary).unwrap()).unwrap();
766        assert_eq!(abandoned["kind"], "client_disconnect");
767        assert_eq!(summary["completion"]["kind"], "stream_abandoned");
768        assert_eq!(summary["completion"]["reason"], "downstream_body_dropped");
769        assert_eq!(summary["completion"]["chunks"], 1);
770        assert!(summary["upstream_events"]["captured"].as_u64().unwrap() > 0);
771        assert!(summary["downstream_events"]["captured"].as_u64().unwrap() > 0);
772    }
773
774    #[tokio::test]
775    async fn streaming_capture_writes_redacted_complete_artifacts() {
776        let temp = TempDir::new().unwrap();
777        let traffic = Arc::new(test_capture(temp.path().join("traffic")));
778        let upstream = futures_util::stream::iter(vec![Ok(Bytes::from_static(
779            b"data: {\"type\":\"response.output_text.delta\",\"delta\":\"first\"}\n\ndata: {\"type\":\"response.output_item.added\",\"item\":{\"type\":\"function_call\",\"call_id\":\"call_1\",\"name\":\"lookup\"}}\n\ndata: {\"type\":\"response.function_call_arguments.delta\",\"call_id\":\"call_1\",\"delta\":\"{}\"}\n\ndata: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"function_call\",\"call_id\":\"call_1\"}}\n\ndata: {\"type\":\"response.completed\",\"response\":{\"usage\":{}}}\n\n",
780        ))]);
781        let response = stream_body(
782            upstream,
783            "msg_1".into(),
784            "grok-4.5".into(),
785            None,
786            "req_1".into(),
787            Some(traffic),
788        );
789        let body = response.into_body().collect().await.unwrap().to_bytes();
790        assert!(String::from_utf8_lossy(&body).contains("tool_use"));
791        let names: Vec<_> = std::fs::read_dir(temp.path().join("traffic"))
792            .unwrap()
793            .map(|entry| entry.unwrap().file_name().to_string_lossy().into_owned())
794            .collect();
795        assert!(
796            names
797                .iter()
798                .any(|name| name.contains("032-upstream-response-body.sse"))
799        );
800        assert!(
801            names
802                .iter()
803                .any(|name| name.contains("061-grok-stream-summary"))
804        );
805    }
806
807    #[tokio::test]
808    async fn streaming_capture_records_fragmented_search_and_tool_events() {
809        let temp = TempDir::new().unwrap();
810        let traffic = Arc::new(test_capture(temp.path().join("traffic")));
811        let upstream = futures_util::stream::iter(vec![
812            Ok(Bytes::from_static(
813                b"data: {\"type\":\"response.output_text.delta\",\"delta\":\"fir",
814            )),
815            Ok(Bytes::from_static(
816                b"st\"}\n\ndata: {\"type\":\"response.output_item.added\",\"item\":{\"type\":\"custom_tool_call\",\"name\":\"x_search\",\"id\":\"search_1\"}}\n\ndata: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"custom_tool_call\",\"name\":\"x_search\",\"id\":\"search_1\"}}\n\ndata: {\"type\":\"response.output_item.added\",\"item\":{\"type\":\"function_call\",\"call_id\":\"call_1\",\"name\":\"lookup\"}}\n\ndata: {\"type\":\"response.function_call_arguments.delta\",\"call_id\":\"call_1\",\"delta\":\"{}\"}\n\ndata: {\"type\":\"response.output_item.done\",\"item\":{\"type\":\"function_call\",\"call_id\":\"call_1\"}}\n\ndata: {\"type\":\"response.completed\",\"response\":{\"usage\":{}}}\n\n",
817            )),
818        ]);
819        let response = stream_body(
820            upstream,
821            "msg_1".into(),
822            "grok-4.5".into(),
823            None,
824            "req_1".into(),
825            Some(traffic),
826        );
827        let body = response.into_body().collect().await.unwrap().to_bytes();
828        assert!(String::from_utf8_lossy(&body).contains("tool_use"));
829        let captured = capture_contents(temp.path().join("traffic"));
830        assert!(captured.contains("x_search"));
831        assert!(captured.contains("function_call"));
832        assert!(captured.contains("stream_completion"));
833    }
834
835    #[tokio::test]
836    async fn streaming_capture_records_malformed_and_failed_streams() {
837        for (payload, stage) in [
838            (b"data: {bad json}\n\n".as_slice(), "json"),
839            (
840                b"data: {\"type\":\"response.failed\",\"response\":{}}\n\n".as_slice(),
841                "reducer",
842            ),
843        ] {
844            let temp = TempDir::new().unwrap();
845            let traffic = Arc::new(test_capture(temp.path().join("traffic")));
846            let response = stream_body(
847                futures_util::stream::iter(vec![Ok(Bytes::copy_from_slice(payload))]),
848                "msg_1".into(),
849                "grok-4.5".into(),
850                None,
851                "req_1".into(),
852                Some(traffic),
853            );
854            let body = response.into_body().collect().await.unwrap().to_bytes();
855            assert!(String::from_utf8_lossy(&body).contains("event: error"));
856            let captured = capture_contents(temp.path().join("traffic"));
857            assert!(captured.contains(&format!("\"stage\": \"{stage}\"")));
858            assert!(captured.contains("stream_error"));
859        }
860    }
861
862    #[test]
863    fn non_streaming_malformed_capture_keeps_diagnostics() {
864        let temp = TempDir::new().unwrap();
865        let traffic = test_capture(temp.path().join("traffic"));
866        assert!(
867            accumulate_response_with_traffic(
868                b"data: {bad json}\n\n",
869                "msg_1",
870                "grok-4.5",
871                Some(&traffic),
872            )
873            .is_err()
874        );
875        let captured = capture_contents(temp.path().join("traffic"));
876        assert!(captured.contains("malformed_event"));
877        assert!(captured.contains("\"outcome\": \"error\""));
878    }
879
880    #[test]
881    fn transport_failure_capture_contains_no_credentials() {
882        let temp = TempDir::new().unwrap();
883        let traffic = test_capture(temp.path().join("traffic"));
884        client::capture_failure(Some(&traffic), "transport", "transport", 1);
885        let captured = capture_contents(temp.path().join("traffic"));
886        assert!(captured.contains("transport"));
887        for secret in [
888            "Bearer token",
889            "refresh-secret",
890            "oauth-code",
891            "person@example.com",
892        ] {
893            assert!(!captured.contains(secret));
894        }
895    }
896
897    fn capture_contents(root: std::path::PathBuf) -> String {
898        let mut captured = String::new();
899        let mut pending = vec![root];
900        while let Some(path) = pending.pop() {
901            for entry in std::fs::read_dir(path).unwrap() {
902                let path = entry.unwrap().path();
903                if path.is_dir() {
904                    pending.push(path);
905                } else {
906                    captured.push_str(&std::fs::read_to_string(path).unwrap());
907                }
908            }
909        }
910        captured
911    }
912
913    #[test]
914    fn non_streaming_capture_writes_response_and_redacts_secrets() {
915        let temp = TempDir::new().unwrap();
916        let traffic = test_capture(temp.path().join("traffic"));
917        let upstream = b"data: {\"type\":\"response.output_text.delta\",\"delta\":\"ok\",\"access_token\":\"secret\"}\n\ndata: {\"type\":\"response.completed\",\"response\":{}}\n\n";
918        let value = accumulate_response_with_traffic(upstream, "msg_1", "grok-4.5", Some(&traffic))
919            .unwrap();
920        traffic.write_json("051-downstream-response", &value);
921        let mut captured = String::new();
922        for entry in std::fs::read_dir(temp.path().join("traffic")).unwrap() {
923            let path = entry.unwrap().path();
924            if path.is_file() {
925                captured.push_str(&std::fs::read_to_string(path).unwrap());
926            }
927        }
928        assert!(captured.contains("[redacted len=6]"));
929        assert!(!captured.contains("secret"));
930    }
931}