Skip to main content

claude_codex/providers/cursor/
client.rs

1use base64::Engine;
2use bytes::Bytes;
3use futures_util::StreamExt;
4use prost::Message;
5use tokio::sync::mpsc;
6
7use crate::config;
8use crate::providers::cursor::connect::{
9    ConnectFrame, ConnectFrameDecoder, FLAG_END, FLAG_GZIP, encode_connect_frame,
10    parse_connect_error,
11};
12use crate::providers::cursor::model::CursorModelResolution;
13use crate::providers::cursor::proto;
14use crate::providers::cursor::request::CursorSelectedImage;
15
16/// Upstream response from the Cursor API.
17///
18/// Contains the raw response bytes (or body bytes for streaming) and the
19/// HTTP status.
20pub struct CursorUpstreamResponse {
21    pub status: u16,
22    pub body: Vec<u8>,
23    pub error_detail: Option<String>,
24}
25
26impl CursorUpstreamResponse {
27    pub fn is_success(&self) -> bool {
28        self.status >= 200 && self.status < 300
29    }
30}
31
32/// HTTP/2 client for the Cursor AgentService/Run endpoint.
33pub struct CursorHttpClient {
34    client: reqwest::Client,
35    base_url: String,
36}
37
38impl Default for CursorHttpClient {
39    fn default() -> Self {
40        Self::new()
41    }
42}
43
44impl CursorHttpClient {
45    pub fn new() -> Self {
46        // Use HTTP/2 prior knowledge for cleartext URLs (mock testing) and
47        // standard TLS for https URLs.
48        let base_url = config::cursor_base_url();
49        let is_cleartext = base_url.starts_with("http://");
50
51        let mut builder = reqwest::Client::builder()
52            .http2_keep_alive_timeout(std::time::Duration::from_secs(30))
53            .http2_keep_alive_while_idle(true);
54
55        if is_cleartext {
56            builder = builder.http2_prior_knowledge();
57        }
58
59        let client = builder.build().expect("CursorHttpClient: reqwest client");
60
61        Self { client, base_url }
62    }
63
64    /// Run the Cursor agent with the given prompt and token.
65    ///
66    /// Opens a bidirectional Connect stream, sends the agent request frames,
67    /// and keeps the request body alive while collecting the response frames.
68    pub async fn run_agent(
69        &self,
70        token: &str,
71        prompt: &str,
72        model: &str,
73        images: &[CursorSelectedImage],
74    ) -> Result<CursorUpstreamResponse, CursorError> {
75        let resolved = super::model::resolve_cursor_model(model)
76            .map_err(|e| CursorError::internal(format!("model resolution: {e}")))?;
77
78        let request_id = uuid::Uuid::new_v4().to_string();
79        let frames = build_run_frames(prompt, &resolved, images, &request_id);
80        let (tx, rx) = mpsc::channel::<Result<Bytes, std::io::Error>>(8);
81        let sender = tokio::spawn(async move {
82            for (index, frame) in frames.into_iter().enumerate() {
83                if tx.send(Ok(frame)).await.is_err() {
84                    return;
85                }
86                let delay = match index {
87                    0 => std::time::Duration::from_millis(1500),
88                    1 => std::time::Duration::from_millis(800),
89                    _ => std::time::Duration::from_millis(400),
90                };
91                tokio::time::sleep(delay).await;
92            }
93
94            let mut heartbeat = tokio::time::interval(std::time::Duration::from_secs(5));
95            heartbeat.tick().await;
96            loop {
97                heartbeat.tick().await;
98                if tx.send(Ok(heartbeat_frame())).await.is_err() {
99                    return;
100                }
101            }
102        });
103        let body =
104            reqwest::Body::wrap_stream(futures_util::stream::unfold(rx, |mut rx| async move {
105                rx.recv().await.map(|item| (item, rx))
106            }));
107
108        let url = format!(
109            "{}/agent.v1.AgentService/Run",
110            self.base_url.trim_end_matches('/')
111        );
112        let client_version = config::cursor_client_version();
113
114        let response = self
115            .client
116            .post(&url)
117            .bearer_auth(token)
118            .header("content-type", "application/connect+proto")
119            .header("connect-protocol-version", "1")
120            .header("connect-accept-encoding", "gzip,br")
121            .header("user-agent", "connect-es/1.6.1")
122            .header("x-cursor-client-type", "cli")
123            .header("x-cursor-client-version", &client_version)
124            .header("x-ghost-mode", "true")
125            .header("x-request-id", &request_id)
126            .header("x-original-request-id", &request_id)
127            .header("x-cursor-streaming", "true")
128            .header("te", "trailers")
129            .body(body)
130            .send()
131            .await
132            .map_err(CursorError::from_reqwest)?;
133
134        let status = response.status().as_u16();
135        let headers = response.headers().clone();
136        let error_detail = response
137            .headers()
138            .get("grpc-message")
139            .and_then(|value| value.to_str().ok())
140            .map(str::to_string);
141        let mut stream = response.bytes_stream();
142        let mut body_bytes = Vec::new();
143        let mut received_data = false;
144
145        loop {
146            let timeout = if received_data {
147                std::time::Duration::from_secs(5)
148            } else {
149                std::time::Duration::from_secs(60)
150            };
151            match tokio::time::timeout(timeout, stream.next()).await {
152                Ok(Some(Ok(chunk))) => {
153                    received_data = true;
154                    body_bytes.extend_from_slice(&chunk);
155                    if contains_end_frame(&body_bytes) {
156                        break;
157                    }
158                }
159                Ok(Some(Err(error))) => {
160                    sender.abort();
161                    return Err(CursorError::internal(format!("read body: {error}")));
162                }
163                Ok(None) => break,
164                Err(_) if received_data => break,
165                Err(_) => {
166                    sender.abort();
167                    return Err(CursorError::internal(
168                        "Cursor upstream timed out before sending a response",
169                    ));
170                }
171            }
172        }
173        sender.abort();
174
175        if status >= 400 {
176            let detail = parse_error_body(&body_bytes, &headers);
177            return Err(CursorError::new(status, "Cursor upstream error", detail));
178        }
179
180        Ok(CursorUpstreamResponse {
181            status,
182            body: body_bytes,
183            error_detail,
184        })
185    }
186}
187
188fn encode_varint(mut value: u64, out: &mut Vec<u8>) {
189    while value >= 0x80 {
190        out.push(((value as u8) & 0x7f) | 0x80);
191        value >>= 7;
192    }
193    out.push(value as u8);
194}
195
196fn field_bytes(field: u64, value: &[u8]) -> Vec<u8> {
197    let mut out = Vec::with_capacity(value.len() + 4);
198    encode_varint((field << 3) | 2, &mut out);
199    encode_varint(value.len() as u64, &mut out);
200    out.extend_from_slice(value);
201    out
202}
203
204fn field_string(field: u64, value: &str) -> Vec<u8> {
205    field_bytes(field, value.as_bytes())
206}
207
208fn field_varint(field: u64, value: u64) -> Vec<u8> {
209    let mut out = Vec::new();
210    encode_varint(field << 3, &mut out);
211    encode_varint(value, &mut out);
212    out
213}
214
215fn model_message(model: &str, fast: bool) -> Vec<u8> {
216    let mut out = field_string(1, model);
217    let mut parameter = field_string(1, "fast");
218    parameter.extend(field_string(2, if fast { "true" } else { "false" }));
219    out.extend(field_bytes(3, &parameter));
220    out
221}
222
223fn mode_value(resolved: &CursorModelResolution) -> u64 {
224    match resolved.mode {
225        super::model::CursorAgentMode::Agent => 1,
226        super::model::CursorAgentMode::Ask => 2,
227        super::model::CursorAgentMode::Plan => 3,
228    }
229}
230
231fn selected_context(images: &[CursorSelectedImage]) -> Option<Vec<u8>> {
232    if images.is_empty() {
233        return None;
234    }
235
236    let mut context = Vec::new();
237    for image in images {
238        let data = base64::engine::general_purpose::STANDARD
239            .decode(&image.data)
240            .unwrap_or_default();
241        let mut selected = field_string(2, &image.uuid);
242        selected.extend(field_string(3, &image.path));
243        selected.extend(field_string(7, &image.mime_type));
244        selected.extend(field_bytes(8, &data));
245        context.extend(field_bytes(1, &selected));
246    }
247    Some(context)
248}
249
250fn build_run_frames(
251    prompt: &str,
252    resolved: &CursorModelResolution,
253    images: &[CursorSelectedImage],
254    request_id: &str,
255) -> Vec<Bytes> {
256    let conversation_id = uuid::Uuid::new_v4().to_string();
257    let mut user_message = field_string(1, prompt);
258    user_message.extend(field_string(2, request_id));
259    if let Some(context) = selected_context(images) {
260        user_message.extend(field_bytes(3, &context));
261    } else {
262        user_message.extend(field_bytes(3, &[]));
263    }
264    user_message.extend(field_varint(4, mode_value(resolved)));
265
266    let action = field_bytes(1, &field_bytes(1, &user_message));
267    let mut request = field_bytes(1, &[]);
268    request.extend(field_bytes(2, &action));
269    request.extend(field_bytes(4, &[]));
270    request.extend(field_string(5, &conversation_id));
271    request.extend(field_bytes(
272        9,
273        &model_message(&resolved.model_id, resolved.fast),
274    ));
275    request.extend(field_varint(12, 0));
276    request.extend(field_bytes(14, &field_string(1, "default")));
277    request.extend(field_bytes(
278        14,
279        &model_message(&resolved.model_id, resolved.fast),
280    ));
281    request.extend(field_string(16, &conversation_id));
282    let first = encode_connect_frame(field_bytes(1, &request), 0);
283
284    let cwd = std::env::current_dir()
285        .ok()
286        .and_then(|path| path.to_str().map(str::to_string))
287        .unwrap_or_default();
288    let mut environment = field_string(1, std::env::consts::OS);
289    environment.extend(field_string(2, &cwd));
290    environment.extend(field_string(
291        3,
292        if cfg!(windows) { "powershell" } else { "bash" },
293    ));
294    environment.extend(field_string(10, "UTC"));
295    environment.extend(field_string(11, &cwd));
296    environment.extend(field_varint(14, 1));
297    environment.extend(field_varint(16, 1));
298    environment.extend(field_varint(19, 0));
299    environment.extend(field_varint(20, 0));
300    environment.extend(field_string(21, &cwd));
301    environment.extend(field_varint(22, 0));
302    let context = field_bytes(
303        2,
304        &field_bytes(
305            10,
306            &field_bytes(1, &field_bytes(1, &field_bytes(4, &environment))),
307        ),
308    );
309
310    let mut frames = vec![first, encode_connect_frame(context, 0)];
311    frames.push(encode_connect_frame(
312        field_bytes(5, &field_string(1, "")),
313        0,
314    ));
315    frames.push(encode_connect_frame(
316        field_bytes(3, &field_string(3, "")),
317        0,
318    ));
319    for sequence in 1..=8 {
320        let mut marker = field_varint(1, sequence);
321        marker.extend(field_string(3, ""));
322        frames.push(encode_connect_frame(field_bytes(3, &marker), 0));
323    }
324    frames
325}
326
327fn heartbeat_frame() -> Bytes {
328    encode_connect_frame(field_bytes(7, &[]), 0)
329}
330
331fn contains_end_frame(body: &[u8]) -> bool {
332    let mut offset = 0;
333    while body.len().saturating_sub(offset) >= 5 {
334        let length = u32::from_be_bytes([
335            body[offset + 1],
336            body[offset + 2],
337            body[offset + 3],
338            body[offset + 4],
339        ]) as usize;
340        if body.len().saturating_sub(offset) < 5 + length {
341            return false;
342        }
343        if body[offset] & FLAG_END != 0 {
344            return true;
345        }
346        offset += 5 + length;
347    }
348    false
349}
350
351fn parse_error_body(body_bytes: &[u8], _headers: &reqwest::header::HeaderMap) -> Option<String> {
352    if body_bytes.len() < 5 {
353        return None;
354    }
355    // Try to parse as Connect end frame with JSON error
356    if body_bytes.len() >= 5 {
357        let flags = body_bytes[0];
358        let len = u32::from_be_bytes([body_bytes[1], body_bytes[2], body_bytes[3], body_bytes[4]])
359            as usize;
360        if flags & FLAG_END != 0 && body_bytes.len() >= 5 + len {
361            let payload = &body_bytes[5..5 + len];
362            let err = parse_connect_error(payload);
363            if err.is_some() {
364                return err.map(|e| e.detail);
365            }
366        }
367    }
368
369    // Try plain text error
370    if let Ok(text) = String::from_utf8(body_bytes.to_vec())
371        && !text.is_empty()
372    {
373        return Some(text);
374    }
375    None
376}
377
378/// Decode upstream response bytes into Connect frames containing
379/// AgentServerMessage values.
380pub fn decode_upstream_frames(body: &[u8]) -> Result<Vec<ConnectFrame>, CursorError> {
381    let mut decoder = ConnectFrameDecoder::new();
382    let frames = decoder
383        .push(body)
384        .map_err(|e| CursorError::internal(format!("frame decode: {e}")))?;
385    Ok(frames)
386}
387
388/// Decode a single Connect frame payload into an AgentServerMessage.
389/// Handles gzip decompression if the FLAG_GZIP bit is set.
390pub fn decode_frame_payload(
391    frame: &ConnectFrame,
392) -> Result<proto::AgentServerMessage, CursorError> {
393    let payload = if frame.flags & FLAG_GZIP != 0 {
394        super::connect::decode_gzip_frame(&frame.payload)
395            .map_err(|e| CursorError::internal(format!("gzip decompress: {e}")))?
396    } else {
397        frame.payload.to_vec()
398    };
399
400    proto::AgentServerMessage::decode(&payload[..])
401        .map_err(|e| CursorError::internal(format!("prost decode: {e}")))
402}
403
404// ---------------------------------------------------------------------------
405// Error type
406// ---------------------------------------------------------------------------
407
408#[derive(Debug, Clone)]
409pub struct CursorError {
410    pub status: u16,
411    pub message: String,
412    pub detail: Option<String>,
413    pub retry_after: Option<String>,
414}
415
416impl CursorError {
417    pub fn new(status: u16, message: impl Into<String>, detail: Option<String>) -> Self {
418        Self {
419            status,
420            message: message.into(),
421            detail,
422            retry_after: None,
423        }
424    }
425
426    pub fn internal(message: impl Into<String>) -> Self {
427        Self {
428            status: 502,
429            message: message.into(),
430            detail: None,
431            retry_after: None,
432        }
433    }
434
435    pub fn from_reqwest(e: reqwest::Error) -> Self {
436        let status = e.status().map(|s| s.as_u16()).unwrap_or(502);
437        Self {
438            status,
439            message: e.to_string(),
440            detail: None,
441            retry_after: None,
442        }
443    }
444}
445
446impl std::fmt::Display for CursorError {
447    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
448        write!(f, "Cursor error {}: {}", self.status, self.message)
449    }
450}
451
452impl std::error::Error for CursorError {}