Skip to main content

claude_codex/providers/cursor/
mod.rs

1pub mod auth;
2pub mod client;
3pub mod connect;
4pub mod model;
5pub mod proto;
6pub mod request;
7pub mod response;
8pub mod sse;
9#[cfg(test)]
10pub(crate) mod test_frames;
11pub mod tool_bridge;
12pub mod tool_use_xml;
13
14use async_trait::async_trait;
15use axum::Json;
16use axum::response::{IntoResponse, Response};
17use http::StatusCode;
18
19use crate::anthropic::error::json_error;
20use crate::anthropic::schema::{CountTokensResponse, MessagesRequest};
21use crate::monitor::usage_from_anthropic_sse;
22use crate::provider::{
23    CliHandlers, Generation, GenerationBody, Provider, ProviderError, ProviderErrorKind,
24    RequestContext,
25};
26use crate::providers::cursor::auth::{
27    clear_cursor_auth, expired_auth_message, load_cursor_auth, missing_auth_message,
28    run_cursor_login,
29};
30use crate::providers::cursor::client::CursorHttpClient;
31use crate::providers::cursor::model::resolve_cursor_model;
32use crate::providers::cursor::request::render_cursor_prompt;
33use crate::providers::cursor::response::{
34    CursorDecodeError, decode_cursor_upstream, decode_upstream_response,
35};
36use crate::providers::cursor::tool_bridge::{
37    BridgeRegistry, advertised_tool_names, can_bridge_cursor_native_tools, find_tool_result,
38    resume_cursor_tool_bridge, start_cursor_tool_bridge,
39};
40
41// ---------------------------------------------------------------------------
42// Provider
43// ---------------------------------------------------------------------------
44
45pub struct CursorProvider;
46
47impl Default for CursorProvider {
48    fn default() -> Self {
49        Self::new()
50    }
51}
52
53impl CursorProvider {
54    pub fn new() -> Self {
55        Self
56    }
57}
58
59#[async_trait]
60impl Provider for CursorProvider {
61    fn name(&self) -> &'static str {
62        "cursor"
63    }
64
65    fn supported_models(&self) -> Vec<String> {
66        model::cursor_supported_models()
67    }
68
69    fn cli(&self) -> &'static dyn CliHandlers {
70        &CURSOR_CLI
71    }
72
73    async fn handle_messages(&self, body: MessagesRequest, ctx: RequestContext) -> Response {
74        let message_id = format!("msg_{}", uuid::Uuid::new_v4().to_string().replace('-', ""));
75        let want_stream = body.stream;
76        let model = body.model.as_deref().unwrap_or("cursor");
77
78        let resolved = resolve_cursor_model(model);
79        if let Err(e) = resolved {
80            return json_error(
81                StatusCode::BAD_REQUEST,
82                "invalid_request_error",
83                format!("Model \"{model}\" is not supported: {e}"),
84            );
85        }
86
87        if let Some(ref session_id) = ctx.session_id
88            && let Some(pending) = BridgeRegistry::pending_tool(session_id)
89            && let Some(result) = find_tool_result(&body, pending.tool_use_id())
90        {
91            let (_result_messages, sse_bytes) =
92                resume_cursor_tool_bridge(session_id, &message_id, model, result, &pending);
93            if let Some(monitor) = ctx.monitor.as_ref() {
94                let (input_tokens, output_tokens) = usage_from_anthropic_sse(&sse_bytes);
95                monitor.stream_progress(
96                    &ctx.req_id,
97                    sse_bytes.len() as u64,
98                    count_sse_events(&sse_bytes),
99                    input_tokens,
100                    output_tokens,
101                );
102            }
103            let headers = [
104                (http::header::CONTENT_TYPE, "text/event-stream"),
105                (http::header::CACHE_CONTROL, "no-cache"),
106                (http::header::CONNECTION, "keep-alive"),
107            ];
108            return (headers, sse_bytes).into_response();
109        }
110
111        let auth = match load_cursor_auth() {
112            Ok(Some(auth)) => auth,
113            Ok(None) => {
114                return json_error(
115                    StatusCode::UNAUTHORIZED,
116                    "authentication_error",
117                    missing_auth_message(),
118                );
119            }
120            Err(err) => {
121                return json_error(
122                    StatusCode::UNAUTHORIZED,
123                    "authentication_error",
124                    format!("Cursor auth failed: {err}"),
125                );
126            }
127        };
128
129        if matches!(auth.expires, Some(expires) if expires <= now_ms() + 60_000) {
130            return json_error(
131                StatusCode::UNAUTHORIZED,
132                "authentication_error",
133                expired_auth_message(&auth),
134            );
135        }
136
137        let token = auth.access_token;
138
139        let prompt = render_cursor_prompt(&body);
140        let images = request::cursor_selected_images(&body);
141
142        let client = CursorHttpClient::new();
143        if let Some(monitor) = ctx.monitor.as_ref() {
144            monitor.upstream_started(&ctx.req_id);
145        }
146        let upstream = match client.run_agent(&token, &prompt, model, &images).await {
147            Ok(r) => r,
148            Err(e) => {
149                return map_cursor_error_to_response(&e);
150            }
151        };
152
153        if want_stream {
154            let session_id = ctx.session_id.as_deref();
155            let bridge_eligible = can_bridge_cursor_native_tools(&body, session_id);
156
157            if bridge_eligible {
158                let events = match decode_upstream_response(&upstream.body) {
159                    Ok(e) => e,
160                    Err(e) => return map_cursor_decode_error_to_response(&e),
161                };
162
163                let allowed = advertised_tool_names(&body);
164                let (sse_bytes, _paused) = start_cursor_tool_bridge(
165                    &message_id,
166                    model,
167                    session_id.unwrap(),
168                    &events,
169                    allowed,
170                    Box::new(|| uuid::Uuid::new_v4().to_string().replace('-', "")),
171                );
172                if let Some(monitor) = ctx.monitor.as_ref() {
173                    let (input_tokens, output_tokens) = usage_from_anthropic_sse(&sse_bytes);
174                    monitor.stream_progress(
175                        &ctx.req_id,
176                        sse_bytes.len() as u64,
177                        count_sse_events(&sse_bytes),
178                        input_tokens,
179                        output_tokens,
180                    );
181                }
182
183                let headers = [
184                    (http::header::CONTENT_TYPE, "text/event-stream"),
185                    (http::header::CACHE_CONTROL, "no-cache"),
186                    (http::header::CONNECTION, "keep-alive"),
187                ];
188                (headers, sse_bytes).into_response()
189            } else {
190                let sse_bytes = sse::frame_cursor_stream(&upstream, &message_id, model);
191                if let Some(monitor) = ctx.monitor.as_ref() {
192                    let (input_tokens, output_tokens) = usage_from_anthropic_sse(&sse_bytes);
193                    monitor.stream_progress(
194                        &ctx.req_id,
195                        sse_bytes.len() as u64,
196                        count_sse_events(&sse_bytes),
197                        input_tokens,
198                        output_tokens,
199                    );
200                }
201                let headers = [
202                    (http::header::CONTENT_TYPE, "text/event-stream"),
203                    (http::header::CACHE_CONTROL, "no-cache"),
204                    (http::header::CONNECTION, "keep-alive"),
205                ];
206                (headers, sse_bytes).into_response()
207            }
208        } else {
209            match decode_cursor_upstream(&upstream, &message_id, model) {
210                Ok(json) => {
211                    if let Some(monitor) = ctx.monitor.as_ref() {
212                        monitor.usage_updated(
213                            &ctx.req_id,
214                            json.pointer("/usage/input_tokens").and_then(|v| v.as_u64()),
215                            json.pointer("/usage/output_tokens")
216                                .and_then(|v| v.as_u64()),
217                        );
218                    }
219                    (StatusCode::OK, Json(json)).into_response()
220                }
221                Err(e) => map_cursor_decode_error_to_response(&e),
222            }
223        }
224    }
225
226    async fn handle_count_tokens(&self, body: MessagesRequest, ctx: RequestContext) -> Response {
227        let prompt = render_cursor_prompt(&body);
228        let tokens = (prompt.len() / 4) as u64; // rough estimate
229        if let Some(monitor) = ctx.monitor.as_ref() {
230            monitor.usage_updated(&ctx.req_id, Some(tokens), None);
231        }
232        (
233            StatusCode::OK,
234            Json(CountTokensResponse {
235                input_tokens: tokens,
236            }),
237        )
238            .into_response()
239    }
240
241    async fn generate_anthropic_stream(
242        &self,
243        mut body: MessagesRequest,
244        ctx: RequestContext,
245    ) -> Result<Generation, ProviderError> {
246        body.stream = true;
247        let requested = body.model.clone().unwrap_or_else(|| "cursor".to_string());
248        let resolved = resolve_cursor_model(&requested).map_err(|error| {
249            ProviderError::new(
250                StatusCode::BAD_REQUEST,
251                ProviderErrorKind::InvalidRequest,
252                format!("Model \"{requested}\" is not supported: {error}"),
253            )
254        })?;
255        if let Some(monitor) = ctx.monitor.as_ref() {
256            monitor.model_resolved(&ctx.req_id, &resolved.model_id);
257        }
258        let message_id = format!("msg_{}", uuid::Uuid::new_v4().simple());
259        if let Some(session_id) = ctx.session_id.as_deref()
260            && let Some(pending) = BridgeRegistry::pending_tool(session_id)
261            && let Some(result) = find_tool_result(&body, pending.tool_use_id())
262        {
263            let (_, bytes) =
264                resume_cursor_tool_bridge(session_id, &message_id, &requested, result, &pending);
265            return Ok(Generation {
266                body: GenerationBody::BufferedSse(bytes.into()),
267                resolved_model: resolved.model_id,
268            });
269        }
270        let auth = load_cursor_auth()
271            .map_err(|error| {
272                ProviderError::new(
273                    StatusCode::UNAUTHORIZED,
274                    ProviderErrorKind::Authentication,
275                    format!("Cursor auth failed: {error}"),
276                )
277            })?
278            .ok_or_else(|| {
279                ProviderError::new(
280                    StatusCode::UNAUTHORIZED,
281                    ProviderErrorKind::Authentication,
282                    missing_auth_message(),
283                )
284            })?;
285        if matches!(auth.expires, Some(expires) if expires <= now_ms() + 60_000) {
286            return Err(ProviderError::new(
287                StatusCode::UNAUTHORIZED,
288                ProviderErrorKind::Authentication,
289                expired_auth_message(&auth),
290            ));
291        }
292        let prompt = render_cursor_prompt(&body);
293        let images = request::cursor_selected_images(&body);
294        if let Some(traffic) = ctx.traffic.as_ref() {
295            traffic.write_json(
296                "020-upstream-request",
297                &serde_json::json!({
298                    "model": requested,
299                    "prompt": prompt,
300                    "image_count": images.len(),
301                }),
302            );
303        }
304        if let Some(monitor) = ctx.monitor.as_ref() {
305            monitor.upstream_started(&ctx.req_id);
306        }
307        let upstream = CursorHttpClient::new()
308            .run_agent(&auth.access_token, &prompt, &requested, &images)
309            .await
310            .map_err(cursor_provider_error)?;
311        if let Some(traffic) = ctx.traffic.as_ref() {
312            traffic.write_bytes("032-upstream-response-body.bin", &upstream.body);
313        }
314        let bytes = if can_bridge_cursor_native_tools(&body, ctx.session_id.as_deref()) {
315            let events =
316                decode_upstream_response(&upstream.body).map_err(cursor_decode_provider_error)?;
317            let allowed = advertised_tool_names(&body);
318            start_cursor_tool_bridge(
319                &message_id,
320                &requested,
321                ctx.session_id.as_deref().expect("bridge session validated"),
322                &events,
323                allowed,
324                Box::new(|| uuid::Uuid::new_v4().simple().to_string()),
325            )
326            .0
327        } else {
328            sse::frame_cursor_stream(&upstream, &message_id, &requested)
329        };
330        if let Some(traffic) = ctx.traffic.as_ref() {
331            traffic.write_bytes("050-anthropic-intermediate.sse", &bytes);
332        }
333        if let Some(monitor) = ctx.monitor.as_ref() {
334            let (input_tokens, output_tokens) = usage_from_anthropic_sse(&bytes);
335            monitor.stream_progress(
336                &ctx.req_id,
337                bytes.len() as u64,
338                count_sse_events(&bytes),
339                input_tokens,
340                output_tokens,
341            );
342        }
343        Ok(Generation {
344            body: GenerationBody::BufferedSse(bytes.into()),
345            resolved_model: resolved.model_id,
346        })
347    }
348}
349
350fn count_sse_events(bytes: &[u8]) -> u64 {
351    String::from_utf8_lossy(bytes).matches("event:").count() as u64
352}
353
354fn now_ms() -> u64 {
355    std::time::SystemTime::now()
356        .duration_since(std::time::UNIX_EPOCH)
357        .unwrap_or_default()
358        .as_millis() as u64
359}
360
361// ---------------------------------------------------------------------------
362// Error mapping
363// ---------------------------------------------------------------------------
364
365fn cursor_provider_error(err: client::CursorError) -> ProviderError {
366    let (status, kind) = match err.status {
367        401 | 403 => (StatusCode::UNAUTHORIZED, ProviderErrorKind::Authentication),
368        429 => (StatusCode::TOO_MANY_REQUESTS, ProviderErrorKind::RateLimit),
369        _ => (StatusCode::BAD_GATEWAY, ProviderErrorKind::Api),
370    };
371    let mut error = ProviderError::new(status, kind, err.detail.unwrap_or(err.message));
372    if err.status == 429 {
373        error.retry_after = Some(err.retry_after.unwrap_or_else(|| "5".to_string()));
374    }
375    error
376}
377
378fn cursor_decode_provider_error(err: CursorDecodeError) -> ProviderError {
379    let (status, kind) = match err.status() {
380        Some(401 | 403) => (StatusCode::UNAUTHORIZED, ProviderErrorKind::Authentication),
381        Some(429) => (StatusCode::TOO_MANY_REQUESTS, ProviderErrorKind::RateLimit),
382        _ => (StatusCode::BAD_GATEWAY, ProviderErrorKind::Api),
383    };
384    ProviderError::new(status, kind, format!("Response decoding error: {err}"))
385}
386
387fn map_cursor_error_to_response(err: &client::CursorError) -> Response {
388    match err.status {
389        401 | 403 => json_error(
390            StatusCode::UNAUTHORIZED,
391            "authentication_error",
392            err.detail.as_deref().unwrap_or("Authentication failed"),
393        ),
394        429 => {
395            let retry_after = err.retry_after.as_deref().unwrap_or("5");
396            let resp = json_error(
397                StatusCode::TOO_MANY_REQUESTS,
398                "rate_limit_error",
399                &err.message,
400            );
401            let headers = [(http::header::RETRY_AFTER, retry_after)];
402            (headers, resp).into_response()
403        }
404        _ => json_error(
405            StatusCode::BAD_GATEWAY,
406            "api_error",
407            err.detail.as_deref().unwrap_or("Upstream error"),
408        ),
409    }
410}
411
412fn map_cursor_decode_error_to_response(err: &CursorDecodeError) -> Response {
413    match err.status() {
414        Some(401 | 403) => json_error(
415            StatusCode::UNAUTHORIZED,
416            "authentication_error",
417            err.to_string(),
418        ),
419        Some(429) => json_error(
420            StatusCode::TOO_MANY_REQUESTS,
421            "rate_limit_error",
422            err.to_string(),
423        ),
424        _ => json_error(
425            StatusCode::BAD_GATEWAY,
426            "api_error",
427            format!("Response decoding error: {err}"),
428        ),
429    }
430}
431
432// ---------------------------------------------------------------------------
433// CLI
434// ---------------------------------------------------------------------------
435
436pub(crate) struct CursorCli;
437
438impl CliHandlers for CursorCli {
439    fn login(&self) -> Result<(), anyhow::Error> {
440        let auth = run_cursor_login()?.ok_or_else(|| anyhow::anyhow!("Cursor login timed out"))?;
441        println!("Cursor auth saved in {}", auth.source);
442        if let Some(ref user_id) = auth.user_id {
443            println!("User: {user_id}");
444        }
445        if let Some(ref email) = auth.email {
446            println!("Email: {email}");
447        }
448        Ok(())
449    }
450
451    fn device(&self) -> Result<(), anyhow::Error> {
452        anyhow::bail!("cursor: device login not yet implemented");
453    }
454
455    fn status(&self) -> Result<(), anyhow::Error> {
456        match load_cursor_auth()? {
457            Some(auth) => {
458                println!("Auth source: {}", auth.source);
459                if let Some(ref user_id) = auth.user_id {
460                    println!("User: {user_id}");
461                }
462                if let Some(ref email) = auth.email {
463                    println!("Email: {email}");
464                }
465                if let Some(expires) = auth.expires {
466                    let remaining = expires.saturating_sub(now_ms()) / 1000;
467                    println!("Access token expires in: {remaining}s");
468                } else {
469                    println!("Access token expiry: unknown");
470                }
471                Ok(())
472            }
473            None => {
474                anyhow::bail!("Not authenticated");
475            }
476        }
477    }
478
479    fn logout(&self) -> Result<(), anyhow::Error> {
480        clear_cursor_auth()?;
481        println!(
482            "Cursor persistent auth cleared. Unset CCP_CURSOR_AUTH_TOKEN or CURSOR_AUTH_TOKEN if using env auth."
483        );
484        Ok(())
485    }
486}
487
488pub(crate) static CURSOR_CLI: CursorCli = CursorCli;
489
490#[cfg(test)]
491mod tests {
492    use super::*;
493
494    #[test]
495    fn supported_models_includes_legacy_and_agent() {
496        let provider = CursorProvider::new();
497        let models = provider.supported_models();
498        assert!(models.contains(&"cursor".to_string()));
499        assert!(models.contains(&"cursor-agent".to_string()));
500        assert!(models.contains(&"cursor-plan".to_string()));
501        assert!(models.contains(&"cursor-ask".to_string()));
502    }
503
504    #[test]
505    fn cursor_cli_logout_does_not_error() {
506        let result = CURSOR_CLI.logout();
507        assert!(result.is_ok());
508    }
509}