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 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
87pub 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 (_content.len() / 4) as u64
140}
141
142fn events_from_message(msg: &AgentServerMessage, events: &mut Vec<CursorStreamEvent>) {
143 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 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 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 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
186pub 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}