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                && let Some(err) = parse_connect_error(&frame.payload)
69            {
70                return Err(CursorDecodeError::ConnectEnd(err));
71            }
72            events.push(CursorStreamEvent::End);
73            continue;
74        }
75
76        let decompressed;
77        let payload = if frame.flags & crate::providers::cursor::connect::FLAG_GZIP != 0 {
78            decompressed = crate::providers::cursor::connect::decode_gzip_frame(&frame.payload)
79                .map_err(|error| CursorDecodeError::Decode(format!("gzip decompress: {error}")))?;
80            &decompressed[..]
81        } else {
82            &frame.payload[..]
83        };
84        if events_from_current_payload(payload, &mut events) {
85            continue;
86        }
87
88        let msg = match decode_frame_payload(frame) {
89            Ok(message) => message,
90            Err(_) => continue,
91        };
92
93        events_from_message(&msg, &mut events);
94    }
95
96    Ok(events)
97}
98
99/// Build an accumulated Anthropic response JSON from upstream bytes for
100/// non-streaming mode.
101pub fn decode_cursor_upstream(
102    upstream: &CursorUpstreamResponse,
103    message_id: &str,
104    model: &str,
105) -> Result<serde_json::Value, CursorDecodeError> {
106    let events = decode_upstream_response(&upstream.body)?;
107
108    let mut text_content = String::new();
109    let mut final_input_tokens: u64 = 0;
110    let mut final_output_tokens: u64 = 0;
111
112    for event in &events {
113        match event {
114            CursorStreamEvent::TextDelta { text } => text_content.push_str(text),
115            CursorStreamEvent::Usage {
116                input_tokens,
117                output_tokens,
118                ..
119            } => {
120                final_input_tokens = *input_tokens;
121                final_output_tokens = *output_tokens;
122            }
123            CursorStreamEvent::End => break,
124            _ => {}
125        }
126    }
127
128    let input_tokens = final_input_tokens.max(estimate_input_tokens(&text_content));
129
130    Ok(serde_json::json!({
131        "id": message_id,
132        "type": "message",
133        "role": "assistant",
134        "content": [
135            {"type": "text", "text": text_content}
136        ],
137        "model": model,
138        "stop_reason": "end_turn",
139        "stop_sequence": null,
140        "usage": {
141            "input_tokens": input_tokens,
142            "output_tokens": final_output_tokens,
143            "cache_creation_input_tokens": 0,
144            "cache_read_input_tokens": 0
145        }
146    }))
147}
148
149fn estimate_input_tokens(_content: &str) -> u64 {
150    // Rough upper bound: 4 chars per token for input estimation
151    (_content.len() / 4) as u64
152}
153
154struct ProtoField<'a> {
155    number: u64,
156    wire_type: u8,
157    data: &'a [u8],
158    value: u64,
159}
160
161fn read_varint(bytes: &[u8]) -> Option<(u64, &[u8])> {
162    let mut value = 0u64;
163    let mut shift = 0;
164    for (index, byte) in bytes.iter().copied().enumerate() {
165        value |= u64::from(byte & 0x7f) << shift;
166        if byte & 0x80 == 0 {
167            return Some((value, &bytes[index + 1..]));
168        }
169        shift += 7;
170        if shift >= 64 {
171            return None;
172        }
173    }
174    None
175}
176
177fn proto_fields(mut bytes: &[u8]) -> impl Iterator<Item = ProtoField<'_>> {
178    std::iter::from_fn(move || {
179        let (tag, rest) = read_varint(bytes)?;
180        bytes = rest;
181        let number = tag >> 3;
182        let wire_type = (tag & 7) as u8;
183        match wire_type {
184            0 => {
185                let (value, rest) = read_varint(bytes)?;
186                bytes = rest;
187                Some(ProtoField {
188                    number,
189                    wire_type,
190                    data: &[],
191                    value,
192                })
193            }
194            1 if bytes.len() >= 8 => {
195                bytes = &bytes[8..];
196                Some(ProtoField {
197                    number,
198                    wire_type,
199                    data: &[],
200                    value: 0,
201                })
202            }
203            2 => {
204                let (length, rest) = read_varint(bytes)?;
205                let length = usize::try_from(length).ok()?;
206                if rest.len() < length {
207                    return None;
208                }
209                let data = &rest[..length];
210                bytes = &rest[length..];
211                Some(ProtoField {
212                    number,
213                    wire_type,
214                    data,
215                    value: 0,
216                })
217            }
218            5 if bytes.len() >= 4 => {
219                bytes = &bytes[4..];
220                Some(ProtoField {
221                    number,
222                    wire_type,
223                    data: &[],
224                    value: 0,
225                })
226            }
227            _ => None,
228        }
229    })
230}
231
232fn nested_text(payload: &[u8], update_type: u64) -> Option<String> {
233    let interaction =
234        proto_fields(payload).find(|field| field.number == 1 && field.wire_type == 2)?;
235    let update = proto_fields(interaction.data)
236        .find(|field| field.number == update_type && field.wire_type == 2)?;
237    let text = proto_fields(update.data).find(|field| field.number == 1 && field.wire_type == 2)?;
238    let text = std::str::from_utf8(text.data).ok()?;
239    (!text.is_empty()).then(|| text.to_string())
240}
241
242fn current_usage(payload: &[u8]) -> Option<(u64, u64, u64, u64)> {
243    let interaction =
244        proto_fields(payload).find(|field| field.number == 1 && field.wire_type == 2)?;
245    let ended =
246        proto_fields(interaction.data).find(|field| field.number == 14 && field.wire_type == 2)?;
247    let mut usage = [0; 4];
248    for field in proto_fields(ended.data) {
249        if field.wire_type == 0 && (1..=4).contains(&field.number) {
250            usage[field.number as usize - 1] = field.value;
251        }
252    }
253    Some((usage[0], usage[1], usage[2], usage[3]))
254}
255
256fn events_from_current_payload(payload: &[u8], events: &mut Vec<CursorStreamEvent>) -> bool {
257    let mut decoded = false;
258    if let Some(text) = nested_text(payload, 4) {
259        events.push(CursorStreamEvent::ThinkingDelta { text });
260        decoded = true;
261    }
262    if let Some(text) = nested_text(payload, 1) {
263        events.push(CursorStreamEvent::TextDelta { text });
264        decoded = true;
265    }
266    if let Some((input_tokens, output_tokens, cache_read_tokens, cache_write_tokens)) =
267        current_usage(payload)
268    {
269        events.push(CursorStreamEvent::Usage {
270            input_tokens,
271            output_tokens,
272            cache_read_tokens,
273            cache_write_tokens,
274        });
275        events.push(CursorStreamEvent::End);
276        decoded = true;
277    }
278    decoded
279}
280
281fn events_from_message(msg: &AgentServerMessage, events: &mut Vec<CursorStreamEvent>) {
282    // Check for exec_server_message with session info
283    if let Some(ref exec) = msg.exec_server_message
284        && let Some(ref session_id) = exec.notes_session_id
285        && !session_id.is_empty()
286    {
287        events.push(CursorStreamEvent::Session {
288            session_id: session_id.clone(),
289        });
290    }
291
292    if let Some(ref update) = msg.interaction_update {
293        // Thinking delta
294        if let Some(ref td) = update.thinking_delta
295            && !td.text.is_empty()
296        {
297            events.push(CursorStreamEvent::ThinkingDelta {
298                text: td.text.clone(),
299            });
300        }
301
302        // Text delta
303        if let Some(ref td) = update.text_delta
304            && !td.text.is_empty()
305        {
306            events.push(CursorStreamEvent::TextDelta {
307                text: td.text.clone(),
308            });
309        }
310
311        // Turn ended (usage + end)
312        if let Some(ref te) = update.turn_ended {
313            events.push(CursorStreamEvent::Usage {
314                input_tokens: te.input_tokens,
315                output_tokens: te.output_tokens,
316                cache_read_tokens: te.cache_read_tokens,
317                cache_write_tokens: te.cache_write_tokens,
318            });
319            events.push(CursorStreamEvent::End);
320        }
321    }
322}
323
324/// Extract an estimate of input tokens from a MessagesRequest for usage
325/// reporting. This is a rough heuristic based on JSON string length.
326pub fn estimate_request_input_tokens(req: &MessagesRequest) -> u64 {
327    let prompt = super::request::render_cursor_prompt(req);
328    (prompt.len() / 4).max(1) as u64
329}
330
331#[cfg(test)]
332mod tests {
333    use super::*;
334    use crate::providers::cursor::connect::encode_connect_frame;
335    use crate::providers::cursor::proto::*;
336    use crate::providers::cursor::test_frames;
337    use prost::Message;
338
339    #[test]
340    fn decodes_text_and_usage_events() {
341        let mut body = Vec::new();
342        body.extend_from_slice(&test_frames::text_frame("Hello"));
343        body.extend_from_slice(&test_frames::text_frame(" world"));
344        body.extend_from_slice(&test_frames::usage_frame(10, 5));
345        body.extend_from_slice(&test_frames::end_frame());
346
347        let events = decode_upstream_response(&body).unwrap();
348        assert_eq!(events.len(), 5);
349        assert!(matches!(events[0], CursorStreamEvent::TextDelta { .. }));
350        assert!(matches!(events[1], CursorStreamEvent::TextDelta { .. }));
351        assert!(matches!(events[2], CursorStreamEvent::Usage { .. }));
352        assert!(matches!(events[3], CursorStreamEvent::End));
353        assert!(matches!(events[4], CursorStreamEvent::End));
354    }
355
356    #[test]
357    fn decodes_thinking_delta() {
358        let body = test_frames::thinking_frame("thinking...");
359
360        let events = decode_upstream_response(&body).unwrap();
361        assert_eq!(events.len(), 1);
362        if let CursorStreamEvent::ThinkingDelta { text } = &events[0] {
363            assert_eq!(text, "thinking...");
364        } else {
365            panic!("expected ThinkingDelta");
366        }
367    }
368
369    #[test]
370    fn decodes_session_event() {
371        let msg = AgentServerMessage {
372            interaction_update: None,
373            exec_server_message: Some(ExecServerMessage {
374                notes_session_id: Some("session-123".to_string()),
375            }),
376        };
377        let mut payload = Vec::new();
378        msg.encode(&mut payload).unwrap();
379        let body = encode_connect_frame(&payload, 0).to_vec();
380
381        let events = decode_upstream_response(&body).unwrap();
382        assert_eq!(events.len(), 1);
383        if let CursorStreamEvent::Session { session_id } = &events[0] {
384            assert_eq!(session_id, "session-123");
385        } else {
386            panic!("expected Session");
387        }
388    }
389
390    #[test]
391    fn accumulate_response_produces_anthropic_json() {
392        let mut body = Vec::new();
393        body.extend_from_slice(&test_frames::text_frame("Hello world"));
394        body.extend_from_slice(&test_frames::usage_frame(15, 3));
395        body.extend_from_slice(&test_frames::end_frame());
396
397        let upstream = CursorUpstreamResponse {
398            status: 200,
399            body,
400            error_detail: None,
401        };
402
403        let json = decode_cursor_upstream(&upstream, "msg_test", "cursor-test").unwrap();
404        assert_eq!(json["id"], "msg_test");
405        assert_eq!(json["content"][0]["text"], "Hello world");
406        assert_eq!(json["usage"]["input_tokens"].as_u64(), Some(15));
407        assert_eq!(json["usage"]["output_tokens"].as_u64(), Some(3));
408        assert_eq!(
409            json["usage"]["cache_creation_input_tokens"].as_u64(),
410            Some(0)
411        );
412        assert_eq!(json["usage"]["cache_read_input_tokens"].as_u64(), Some(0));
413        assert_eq!(json["stop_reason"], "end_turn");
414    }
415
416    #[test]
417    fn empty_upstream_produces_empty_response() {
418        let upstream = CursorUpstreamResponse {
419            status: 200,
420            body: Vec::new(),
421            error_detail: None,
422        };
423        let json = decode_cursor_upstream(&upstream, "msg_empty", "cursor-test").unwrap();
424        assert_eq!(json["content"][0]["text"], "");
425    }
426
427    #[test]
428    fn connect_end_frame_with_error_is_rejected() {
429        let json_err = serde_json::json!({
430            "error": {"code": "resource_exhausted", "message": "quota exceeded"}
431        });
432        let payload = serde_json::to_vec(&json_err).unwrap();
433        let frame = encode_connect_frame(&payload, FLAG_END);
434        let result = decode_upstream_response(&frame);
435        assert!(result.is_err());
436        let err = result.unwrap_err();
437        assert_eq!(err.status(), Some(429));
438        assert!(err.to_string().contains("quota exceeded"));
439    }
440
441    #[test]
442    fn multiple_text_deltas_accumulate() {
443        let mut body = Vec::new();
444        body.extend_from_slice(&test_frames::text_frame("Hello "));
445        body.extend_from_slice(&test_frames::text_frame("world"));
446        body.extend_from_slice(&test_frames::usage_frame(10, 2));
447        body.extend_from_slice(&test_frames::end_frame());
448
449        let events = decode_upstream_response(&body).unwrap();
450        let text: String = events
451            .iter()
452            .filter_map(|e| {
453                if let CursorStreamEvent::TextDelta { text } = e {
454                    Some(text.as_str())
455                } else {
456                    None
457                }
458            })
459            .collect();
460        assert_eq!(text, "Hello world");
461    }
462}