Skip to main content

claude_codex/providers/cursor/
response.rs

1use crate::anthropic::schema::MessagesRequest;
2use crate::providers::cursor::client::{
3    CursorUpstreamResponse, decode_frame_payload, decode_upstream_frames,
4};
5use crate::providers::cursor::connect::{ConnectEndError, FLAG_END, parse_connect_error};
6use crate::providers::cursor::proto::AgentServerMessage;
7
8/// A decoded event from the Cursor upstream response stream.
9#[derive(Debug, Clone)]
10pub enum CursorStreamEvent {
11    Session {
12        session_id: String,
13    },
14    ThinkingDelta {
15        text: String,
16    },
17    TextDelta {
18        text: String,
19    },
20    Usage {
21        input_tokens: u64,
22        output_tokens: u64,
23        cache_read_tokens: u64,
24        cache_write_tokens: u64,
25    },
26    End,
27}
28
29#[derive(Debug, Clone)]
30pub enum CursorDecodeError {
31    ConnectEnd(ConnectEndError),
32    Decode(String),
33}
34
35impl CursorDecodeError {
36    pub fn status(&self) -> Option<u16> {
37        match self {
38            CursorDecodeError::ConnectEnd(err) => Some(err.status),
39            CursorDecodeError::Decode(_) => None,
40        }
41    }
42}
43
44impl std::fmt::Display for CursorDecodeError {
45    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
46        match self {
47            CursorDecodeError::ConnectEnd(err) => write!(f, "{err}"),
48            CursorDecodeError::Decode(message) => write!(f, "{message}"),
49        }
50    }
51}
52
53impl std::error::Error for CursorDecodeError {}
54
55/// Decode upstream response bytes into a sequence of CursorStreamEvents.
56///
57/// Returns both the events and the final usage for the response, since the
58/// upstream may send multiple update frames.
59pub fn decode_upstream_response(body: &[u8]) -> Result<Vec<CursorStreamEvent>, CursorDecodeError> {
60    let frames =
61        decode_upstream_frames(body).map_err(|e| CursorDecodeError::Decode(e.to_string()))?;
62    let mut events = Vec::new();
63
64    for frame in &frames {
65        if frame.flags & FLAG_END != 0 {
66            // Check for Connect error in end frame
67            if !frame.payload.is_empty() {
68                if let Some(err) = parse_connect_error(&frame.payload) {
69                    return Err(CursorDecodeError::ConnectEnd(err));
70                }
71            }
72            events.push(CursorStreamEvent::End);
73            continue;
74        }
75
76        let msg = match decode_frame_payload(frame) {
77            Ok(m) => m,
78            Err(_) => continue,
79        };
80
81        events_from_message(&msg, &mut events);
82    }
83
84    Ok(events)
85}
86
87/// Build an accumulated Anthropic response JSON from upstream bytes for
88/// non-streaming mode.
89pub fn decode_cursor_upstream(
90    upstream: &CursorUpstreamResponse,
91    message_id: &str,
92    model: &str,
93) -> Result<serde_json::Value, CursorDecodeError> {
94    let events = decode_upstream_response(&upstream.body)?;
95
96    let mut text_content = String::new();
97    let mut final_input_tokens: u64 = 0;
98    let mut final_output_tokens: u64 = 0;
99
100    for event in &events {
101        match event {
102            CursorStreamEvent::TextDelta { text } => text_content.push_str(text),
103            CursorStreamEvent::Usage {
104                input_tokens,
105                output_tokens,
106                ..
107            } => {
108                final_input_tokens = *input_tokens;
109                final_output_tokens = *output_tokens;
110            }
111            CursorStreamEvent::End => break,
112            _ => {}
113        }
114    }
115
116    let input_tokens = final_input_tokens.max(estimate_input_tokens(&text_content));
117
118    Ok(serde_json::json!({
119        "id": message_id,
120        "type": "message",
121        "role": "assistant",
122        "content": [
123            {"type": "text", "text": text_content}
124        ],
125        "model": model,
126        "stop_reason": "end_turn",
127        "stop_sequence": null,
128        "usage": {
129            "input_tokens": input_tokens,
130            "output_tokens": final_output_tokens,
131            "cache_creation_input_tokens": 0,
132            "cache_read_input_tokens": 0
133        }
134    }))
135}
136
137fn estimate_input_tokens(_content: &str) -> u64 {
138    // Rough upper bound: 4 chars per token for input estimation
139    (_content.len() / 4) as u64
140}
141
142fn events_from_message(msg: &AgentServerMessage, events: &mut Vec<CursorStreamEvent>) {
143    // Check for exec_server_message with session info
144    if let Some(ref exec) = msg.exec_server_message {
145        if let Some(ref session_id) = exec.notes_session_id {
146            if !session_id.is_empty() {
147                events.push(CursorStreamEvent::Session {
148                    session_id: session_id.clone(),
149                });
150            }
151        }
152    }
153
154    if let Some(ref update) = msg.interaction_update {
155        // Thinking delta
156        if let Some(ref td) = update.thinking_delta {
157            if !td.text.is_empty() {
158                events.push(CursorStreamEvent::ThinkingDelta {
159                    text: td.text.clone(),
160                });
161            }
162        }
163
164        // Text delta
165        if let Some(ref td) = update.text_delta {
166            if !td.text.is_empty() {
167                events.push(CursorStreamEvent::TextDelta {
168                    text: td.text.clone(),
169                });
170            }
171        }
172
173        // Turn ended (usage + end)
174        if let Some(ref te) = update.turn_ended {
175            events.push(CursorStreamEvent::Usage {
176                input_tokens: te.input_tokens,
177                output_tokens: te.output_tokens,
178                cache_read_tokens: te.cache_read_tokens,
179                cache_write_tokens: te.cache_write_tokens,
180            });
181            events.push(CursorStreamEvent::End);
182        }
183    }
184}
185
186/// Extract an estimate of input tokens from a MessagesRequest for usage
187/// reporting. This is a rough heuristic based on JSON string length.
188pub fn estimate_request_input_tokens(req: &MessagesRequest) -> u64 {
189    let prompt = super::request::render_cursor_prompt(req);
190    (prompt.len() / 4).max(1) as u64
191}
192
193#[cfg(test)]
194mod tests {
195    use super::*;
196    use crate::providers::cursor::connect::encode_connect_frame;
197    use crate::providers::cursor::proto::*;
198    use crate::providers::cursor::test_frames;
199    use prost::Message;
200
201    #[test]
202    fn decodes_text_and_usage_events() {
203        let mut body = Vec::new();
204        body.extend_from_slice(&test_frames::text_frame("Hello"));
205        body.extend_from_slice(&test_frames::text_frame(" world"));
206        body.extend_from_slice(&test_frames::usage_frame(10, 5));
207        body.extend_from_slice(&test_frames::end_frame());
208
209        let events = decode_upstream_response(&body).unwrap();
210        assert_eq!(events.len(), 5);
211        assert!(matches!(events[0], CursorStreamEvent::TextDelta { .. }));
212        assert!(matches!(events[1], CursorStreamEvent::TextDelta { .. }));
213        assert!(matches!(events[2], CursorStreamEvent::Usage { .. }));
214        assert!(matches!(events[3], CursorStreamEvent::End));
215        assert!(matches!(events[4], CursorStreamEvent::End));
216    }
217
218    #[test]
219    fn decodes_thinking_delta() {
220        let body = test_frames::thinking_frame("thinking...");
221
222        let events = decode_upstream_response(&body).unwrap();
223        assert_eq!(events.len(), 1);
224        if let CursorStreamEvent::ThinkingDelta { text } = &events[0] {
225            assert_eq!(text, "thinking...");
226        } else {
227            panic!("expected ThinkingDelta");
228        }
229    }
230
231    #[test]
232    fn decodes_session_event() {
233        let msg = AgentServerMessage {
234            interaction_update: None,
235            exec_server_message: Some(ExecServerMessage {
236                notes_session_id: Some("session-123".to_string()),
237            }),
238        };
239        let mut payload = Vec::new();
240        msg.encode(&mut payload).unwrap();
241        let body = encode_connect_frame(&payload, 0).to_vec();
242
243        let events = decode_upstream_response(&body).unwrap();
244        assert_eq!(events.len(), 1);
245        if let CursorStreamEvent::Session { session_id } = &events[0] {
246            assert_eq!(session_id, "session-123");
247        } else {
248            panic!("expected Session");
249        }
250    }
251
252    #[test]
253    fn accumulate_response_produces_anthropic_json() {
254        let mut body = Vec::new();
255        body.extend_from_slice(&test_frames::text_frame("Hello world"));
256        body.extend_from_slice(&test_frames::usage_frame(15, 3));
257        body.extend_from_slice(&test_frames::end_frame());
258
259        let upstream = CursorUpstreamResponse {
260            status: 200,
261            body,
262            error_detail: None,
263        };
264
265        let json = decode_cursor_upstream(&upstream, "msg_test", "cursor-test").unwrap();
266        assert_eq!(json["id"], "msg_test");
267        assert_eq!(json["content"][0]["text"], "Hello world");
268        assert_eq!(json["usage"]["input_tokens"].as_u64(), Some(15));
269        assert_eq!(json["usage"]["output_tokens"].as_u64(), Some(3));
270        assert_eq!(
271            json["usage"]["cache_creation_input_tokens"].as_u64(),
272            Some(0)
273        );
274        assert_eq!(json["usage"]["cache_read_input_tokens"].as_u64(), Some(0));
275        assert_eq!(json["stop_reason"], "end_turn");
276    }
277
278    #[test]
279    fn empty_upstream_produces_empty_response() {
280        let upstream = CursorUpstreamResponse {
281            status: 200,
282            body: Vec::new(),
283            error_detail: None,
284        };
285        let json = decode_cursor_upstream(&upstream, "msg_empty", "cursor-test").unwrap();
286        assert_eq!(json["content"][0]["text"], "");
287    }
288
289    #[test]
290    fn connect_end_frame_with_error_is_rejected() {
291        let json_err = serde_json::json!({
292            "error": {"code": "resource_exhausted", "message": "quota exceeded"}
293        });
294        let payload = serde_json::to_vec(&json_err).unwrap();
295        let frame = encode_connect_frame(&payload, FLAG_END);
296        let result = decode_upstream_response(&frame);
297        assert!(result.is_err());
298        let err = result.unwrap_err();
299        assert_eq!(err.status(), Some(429));
300        assert!(err.to_string().contains("quota exceeded"));
301    }
302
303    #[test]
304    fn multiple_text_deltas_accumulate() {
305        let mut body = Vec::new();
306        body.extend_from_slice(&test_frames::text_frame("Hello "));
307        body.extend_from_slice(&test_frames::text_frame("world"));
308        body.extend_from_slice(&test_frames::usage_frame(10, 2));
309        body.extend_from_slice(&test_frames::end_frame());
310
311        let events = decode_upstream_response(&body).unwrap();
312        let text: String = events
313            .iter()
314            .filter_map(|e| {
315                if let CursorStreamEvent::TextDelta { text } = e {
316                    Some(text.as_str())
317                } else {
318                    None
319                }
320            })
321            .collect();
322        assert_eq!(text, "Hello world");
323    }
324}