Skip to main content

claude_codex/providers/grok/
client.rs

1use 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}