claude_codex/providers/grok/
client.rs1use std::sync::Arc;
2use std::time::{Duration, Instant};
3
4use futures_util::StreamExt;
5use http::StatusCode;
6
7use super::auth::manager::GrokAuthManager;
8use super::auth::token_store::{StoredAuth, file_store};
9use super::translate::request::GrokResponsesRequest;
10use crate::traffic::TrafficCapture;
11
12const DEFAULT_BASE_URL: &str = "https://cli-chat-proxy.grok.com/v1";
13const MAX_BUFFERED_RESPONSE_BYTES: usize = 8 * 1024 * 1024;
14
15pub struct GrokClient {
16 client: Arc<reqwest::Client>,
17 auth: Arc<GrokAuthManager<crate::auth::FileAuthStore<StoredAuth>>>,
18 url: String,
19 client_version: String,
20}
21
22pub struct GrokResponse {
23 response: reqwest::Response,
24}
25pub struct GrokError {
26 pub status: StatusCode,
27 pub retry_after: Option<String>,
28 pub message: String,
29}
30
31impl GrokResponse {
32 pub fn into_response(self) -> reqwest::Response {
33 self.response
34 }
35
36 pub fn into_stream(
37 self,
38 ) -> impl futures_util::Stream<Item = Result<bytes::Bytes, GrokError>> + Send {
39 self.response.bytes_stream().map(|chunk| {
40 chunk.map_err(|_| GrokError {
41 status: StatusCode::BAD_GATEWAY,
42 retry_after: None,
43 message: "Grok upstream stream failed".into(),
44 })
45 })
46 }
47
48 pub async fn into_bytes(self) -> Result<Vec<u8>, GrokError> {
49 let mut stream = self.into_stream();
50 let mut bytes = Vec::new();
51 while let Some(chunk) = stream.next().await {
52 let chunk = chunk?;
53 if bytes.len().saturating_add(chunk.len()) > MAX_BUFFERED_RESPONSE_BYTES {
54 return Err(GrokError {
55 status: StatusCode::BAD_GATEWAY,
56 retry_after: None,
57 message: "Grok upstream response exceeds the size limit".into(),
58 });
59 }
60 bytes.extend_from_slice(&chunk);
61 }
62 Ok(bytes)
63 }
64}
65
66impl GrokClient {
67 pub fn new(base_url: String, client_version: String) -> anyhow::Result<Self> {
68 let client = Arc::new(
69 reqwest::Client::builder()
70 .redirect(reqwest::redirect::Policy::none())
71 .connect_timeout(Duration::from_secs(10))
72 .tcp_nodelay(true)
73 .http2_keep_alive_interval(Duration::from_secs(15))
74 .http2_keep_alive_timeout(Duration::from_secs(5))
75 .http2_keep_alive_while_idle(true)
76 .build()?,
77 );
78 let auth = Arc::new(GrokAuthManager::new(file_store())?);
79 Ok(Self::with_shared(
80 url_for(base_url)?,
81 client_version,
82 client,
83 auth,
84 ))
85 }
86
87 fn with_shared(
88 url: String,
89 client_version: String,
90 client: Arc<reqwest::Client>,
91 auth: Arc<GrokAuthManager<crate::auth::FileAuthStore<StoredAuth>>>,
92 ) -> Self {
93 Self {
94 client,
95 auth,
96 url,
97 client_version,
98 }
99 }
100
101 pub async fn post(
102 &self,
103 body: &GrokResponsesRequest,
104 traffic: Option<Arc<TrafficCapture>>,
105 ) -> Result<GrokResponse, GrokError> {
106 if let Some(capture) = traffic.as_ref() {
107 let body_value = serde_json::to_value(body).unwrap_or(serde_json::Value::Null);
108 capture.write_json("020-upstream-request", &body_value);
109 capture.write_json("021-upstream-request-metadata", &serde_json::json!({
110 "method": "POST", "url": safe_url(&self.url), "provider": "grok", "transport": "http",
111 "headers": {"accept":"text/event-stream", "content-type":"application/json", "authorization":"[redacted]", "x-xai-token-auth":"[redacted]"},
112 "body_bytes": serde_json::to_vec(body).map(|v| v.len()).unwrap_or(0),
113 }));
114 }
115 let auth = match self.auth.get_auth().await {
116 Ok(auth) => auth,
117 Err(error) => {
118 capture_failure(traffic.as_deref(), "auth", "authentication", 0);
119 return Err(auth_error(error));
120 }
121 };
122 let response = self
123 .attempt(&auth.access, body, 1, traffic.as_deref())
124 .await?;
125 if response.status() == StatusCode::UNAUTHORIZED {
126 let refreshed = self
127 .auth
128 .force_refresh(&auth.access)
129 .await
130 .map_err(|error| {
131 capture_failure(traffic.as_deref(), "auth", "refresh", 1);
132 auth_error(error)
133 })?;
134 let replay = self
135 .attempt(&refreshed.access, body, 2, traffic.as_deref())
136 .await?;
137 if replay.status() == StatusCode::UNAUTHORIZED {
138 capture_failure(traffic.as_deref(), "auth", "unauthorized", 2);
139 return Err(auth_error(anyhow::anyhow!("unauthorized")));
140 }
141 return Ok(self.captured_response(replay, traffic.as_deref()));
142 }
143 Ok(self.captured_response(response, traffic.as_deref()))
144 }
145
146 fn captured_response(
147 &self,
148 response: reqwest::Response,
149 traffic: Option<&TrafficCapture>,
150 ) -> GrokResponse {
151 if let Some(capture) = traffic.as_ref() {
152 capture.write_json("030-upstream-response-headers", &serde_json::json!({
153 "status": response.status().as_u16(), "headers": safe_headers(response.headers()),
154 }));
155 }
156 GrokResponse { response }
157 }
158
159 async fn attempt(
160 &self,
161 access: &str,
162 body: &GrokResponsesRequest,
163 attempt: u8,
164 traffic: Option<&TrafficCapture>,
165 ) -> Result<reqwest::Response, GrokError> {
166 let started = Instant::now();
167 let response = self
168 .client
169 .post(&self.url)
170 .header("accept", "text/event-stream")
171 .header("content-type", "application/json")
172 .header("authorization", format!("Bearer {access}"))
173 .header("x-xai-token-auth", "xai-grok-cli")
174 .header("x-grok-client-identifier", "grok-shell")
175 .header("x-grok-client-version", &self.client_version)
176 .json(body)
177 .send()
178 .await
179 .map_err(|_| {
180 capture_failure(traffic, "transport", "transport", attempt);
181 GrokError {
182 status: StatusCode::BAD_GATEWAY,
183 retry_after: None,
184 message: "Grok upstream request failed".into(),
185 }
186 })?;
187 let status = response.status();
188 if let Some(capture) = traffic {
189 capture.write_json("022-upstream-attempt", &serde_json::json!({"attempt":attempt,"status":status.as_u16(),"elapsed_ms":started.elapsed().as_millis(),"headers":safe_headers(response.headers())}));
190 }
191 if !status.is_success() && status != StatusCode::UNAUTHORIZED {
192 let retry_after = response
193 .headers()
194 .get("retry-after")
195 .and_then(|v| v.to_str().ok())
196 .map(str::to_string);
197 if let Some(capture) = traffic {
198 let (body, truncated) = read_rejected_body(response, 64 * 1024).await;
199 let detail = serde_json::from_slice::<serde_json::Value>(&body)
200 .unwrap_or_else(|_| serde_json::json!({"body_bytes": body.len()}));
201 capture.write_json(
202 "031-upstream-error-body",
203 &serde_json::json!({"attempt":attempt,"status":status.as_u16(),"truncated":truncated,"body":detail}),
204 );
205 }
206 return Err(GrokError {
207 status,
208 retry_after,
209 message: "Grok upstream rejected the request".into(),
210 });
211 }
212 Ok(response)
213 }
214}
215
216async fn read_rejected_body(response: reqwest::Response, limit: usize) -> (Vec<u8>, bool) {
217 let mut body = Vec::new();
218 let mut stream = response.bytes_stream();
219 while let Some(chunk) = stream.next().await {
220 let Ok(chunk) = chunk else { break };
221 let remaining = limit.saturating_sub(body.len());
222 if chunk.len() > remaining {
223 body.extend_from_slice(&chunk[..remaining]);
224 return (body, true);
225 }
226 body.extend_from_slice(&chunk);
227 }
228 (body, false)
229}
230
231pub(super) fn capture_failure(
232 traffic: Option<&TrafficCapture>,
233 stage: &str,
234 kind: &str,
235 attempt: u8,
236) {
237 if let Some(capture) = traffic {
238 capture.write_json(
239 "060-grok-stream-error",
240 &serde_json::json!({"stage":stage,"kind":kind,"attempt":attempt}),
241 );
242 }
243}
244
245fn safe_headers(headers: &reqwest::header::HeaderMap) -> serde_json::Value {
246 let mut result = serde_json::Map::new();
247 for name in [
248 "content-type",
249 "content-length",
250 "retry-after",
251 "x-request-id",
252 ] {
253 if let Some(value) = headers.get(name).and_then(|value| value.to_str().ok()) {
254 result.insert(
255 name.to_string(),
256 serde_json::Value::String(value.to_string()),
257 );
258 }
259 }
260 serde_json::Value::Object(result)
261}
262
263fn safe_url(raw: &str) -> String {
264 let Ok(mut url) = reqwest::Url::parse(raw) else {
265 return "[invalid-url]".into();
266 };
267 let _ = url.set_username("");
268 let _ = url.set_password(None);
269 url.set_query(None);
270 url.to_string()
271}
272
273fn url_for(base_url: String) -> anyhow::Result<String> {
274 responses_url(&base_url)
275}
276fn responses_url(base_url: &str) -> anyhow::Result<String> {
277 let base_url = if base_url.trim().is_empty() {
278 DEFAULT_BASE_URL
279 } else {
280 base_url.trim()
281 };
282 let mut url = reqwest::Url::parse(base_url)?;
283 let path = url.path().trim_end_matches('/');
284 if !path.ends_with("/responses") {
285 url.set_path(&format!("{path}/responses"));
286 }
287 Ok(url.to_string().trim_end_matches('/').to_string())
288}
289
290fn auth_error(_: anyhow::Error) -> GrokError {
291 GrokError {
292 status: StatusCode::UNAUTHORIZED,
293 retry_after: None,
294 message: "Grok authentication requires official CLI login and proxy import".into(),
295 }
296}
297
298#[cfg(test)]
299mod tests {
300 use super::responses_url;
301 #[test]
302 fn responses_url_appends_responses_to_base_path() {
303 assert_eq!(
304 responses_url("http://127.0.0.1:8080/v1").unwrap(),
305 "http://127.0.0.1:8080/v1/responses"
306 );
307 }
308 #[test]
309 fn responses_url_preserves_responses_endpoint() {
310 assert_eq!(
311 responses_url("https://example.com/custom/responses/").unwrap(),
312 "https://example.com/custom/responses"
313 );
314 }
315 #[test]
316 fn responses_url_rejects_invalid_url() {
317 assert!(responses_url(":invalid").is_err());
318 }
319}