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