claude_codex/providers/cursor/
response.rs1use 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#[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
55pub 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 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
99pub 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 (_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 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 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 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 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
324pub 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}