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}