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