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}