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                .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}