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
16pub 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
32pub 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 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 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, ¶meter));
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 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 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
378pub 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
388pub 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#[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 {}