1use serde::{Deserialize, Serialize};
2use serde_json::Value;
3
4use super::Citation;
5use super::cache::CacheInfo;
6use super::tool::ToolCall;
7
8#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
10pub struct TokenUsage {
11 pub prompt_tokens: u32,
12 pub completion_tokens: u32,
13 pub total_tokens: u32,
14 pub cached_tokens: Option<u32>,
16 #[serde(skip_serializing_if = "Option::is_none")]
18 pub reasoning_tokens: Option<u32>,
19 #[serde(skip_serializing_if = "Option::is_none")]
21 pub prompt_cache_hit_tokens: Option<u32>,
22 #[serde(skip_serializing_if = "Option::is_none")]
24 pub prompt_cache_miss_tokens: Option<u32>,
25 #[serde(skip_serializing_if = "Option::is_none")]
27 pub audio_tokens: Option<u32>,
28 #[serde(skip_serializing_if = "Option::is_none")]
30 pub cache_write_5m_input_tokens: Option<u32>,
31 #[serde(skip_serializing_if = "Option::is_none")]
33 pub cache_write_1h_input_tokens: Option<u32>,
34}
35
36impl TokenUsage {
37 pub fn new(prompt_tokens: u32, completion_tokens: u32) -> Self {
38 Self {
39 prompt_tokens,
40 completion_tokens,
41 total_tokens: prompt_tokens + completion_tokens,
42 cached_tokens: None,
43 reasoning_tokens: None,
44 prompt_cache_hit_tokens: None,
45 prompt_cache_miss_tokens: None,
46 audio_tokens: None,
47 cache_write_5m_input_tokens: None,
48 cache_write_1h_input_tokens: None,
49 }
50 }
51}
52
53#[derive(Debug, Clone, Default, Serialize, Deserialize)]
55pub struct CompletionResponse {
56 pub content: Option<String>,
58 pub thinking: Option<String>,
60 #[serde(default)]
62 pub tool_calls: Vec<ToolCall>,
63 pub usage: TokenUsage,
65 pub model: String,
67 pub finish_reason: FinishReason,
69 pub latency_ms: u64,
71 pub cache_info: Option<CacheInfo>,
73 #[serde(skip_serializing_if = "Option::is_none")]
75 pub id: Option<String>,
76 #[serde(skip_serializing_if = "Option::is_none")]
78 pub created: Option<u64>,
79 #[serde(skip_serializing_if = "Option::is_none")]
81 pub system_fingerprint: Option<String>,
82 #[serde(skip_serializing_if = "Option::is_none")]
84 pub refusal: Option<String>,
85 #[serde(skip_serializing_if = "Option::is_none")]
87 pub signature: Option<String>,
88 #[serde(skip_serializing_if = "Option::is_none")]
90 pub redacted_thinking: Option<String>,
91 #[serde(skip_serializing_if = "Option::is_none")]
93 pub citations: Option<Vec<Citation>>,
94}
95
96#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
98pub enum FinishReason {
99 #[default]
101 #[serde(rename = "stop")]
102 Stop,
103 #[serde(rename = "tool_call")]
105 ToolCall,
106 #[serde(rename = "max_tokens")]
108 MaxTokens,
109 #[serde(rename = "content_filter")]
111 ContentFilter,
112 #[serde(rename = "pause_turn")]
114 PauseTurn,
115 #[serde(rename = "refusal")]
117 Refusal,
118}
119
120#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
122#[serde(tag = "type")]
123pub enum StreamEvent {
124 #[serde(rename = "content_delta")]
126 ContentDelta { delta: String },
127
128 #[serde(rename = "tool_call_delta")]
130 ToolCallDelta {
131 index: usize,
132 #[serde(skip_serializing_if = "Option::is_none")]
133 id: Option<String>,
134 #[serde(skip_serializing_if = "Option::is_none")]
135 function_name: Option<String>,
136 arguments_delta: String,
137 },
138
139 #[serde(rename = "thinking_delta")]
141 ThinkingDelta { delta: String },
142
143 #[serde(rename = "image_delta")]
145 ImageDelta { media_type: String, delta: String },
146
147 #[serde(rename = "usage")]
149 Usage { usage: TokenUsage },
150
151 #[serde(rename = "done")]
153 Done {
154 finish_reason: FinishReason,
155 #[serde(skip_serializing_if = "Option::is_none")]
156 usage: Option<TokenUsage>,
157 },
158
159 #[serde(rename = "signature_delta")]
161 SignatureDelta { signature: String },
162
163 #[serde(rename = "citations_delta")]
165 CitationsDelta { citations: Value },
166
167 #[serde(rename = "redacted_thinking_delta")]
169 RedactedThinkingDelta { data: String },
170
171 #[serde(rename = "custom")]
173 Custom { event: String, data: Value },
174}
175
176#[derive(Debug, Clone, Serialize, Deserialize)]
180pub struct StreamChunk {
181 pub delta: String,
182 pub finish_reason: Option<String>,
183 pub usage: Option<TokenUsage>,
184}
185
186#[derive(Debug, Clone, Serialize, Deserialize)]
188pub struct StructuredResponse<T> {
189 pub parsed: T,
190 pub raw: String,
191 pub usage: TokenUsage,
192 pub model: String,
193 pub latency_ms: u64,
194}
195
196#[cfg(test)]
197mod tests {
198 use super::*;
199
200 #[test]
201 fn test_token_usage_new() {
202 let usage = TokenUsage::new(100, 50);
203 assert_eq!(usage.prompt_tokens, 100);
204 assert_eq!(usage.completion_tokens, 50);
205 assert_eq!(usage.total_tokens, 150);
206 assert!(usage.cached_tokens.is_none());
207 }
208
209 #[test]
210 fn test_token_usage_new_zero() {
211 let usage = TokenUsage::new(0, 0);
212 assert_eq!(usage.total_tokens, 0);
213 assert_eq!(usage.prompt_tokens, 0);
214 assert_eq!(usage.completion_tokens, 0);
215 }
216
217 #[test]
218 fn test_finish_reason_serde() {
219 let cases = vec![
220 (FinishReason::Stop, "\"stop\""),
221 (FinishReason::ToolCall, "\"tool_call\""),
222 (FinishReason::MaxTokens, "\"max_tokens\""),
223 (FinishReason::ContentFilter, "\"content_filter\""),
224 ];
225 for (reason, expected) in cases {
226 let json = serde_json::to_string(&reason).unwrap();
227 assert_eq!(json, expected);
228 let deserialized: FinishReason = serde_json::from_str(&json).unwrap();
229 assert!(
230 matches!(&deserialized, r if std::mem::discriminant(&reason) == std::mem::discriminant(r))
231 );
232 }
233 }
234
235 #[test]
236 fn test_stream_event_content_delta() {
237 let evt = StreamEvent::ContentDelta { delta: "Hello".into() };
238 match evt {
239 StreamEvent::ContentDelta { delta } => assert_eq!(delta, "Hello"),
240 _ => panic!("Wrong variant"),
241 }
242 }
243
244 #[test]
245 fn test_stream_event_tool_call_delta() {
246 let evt = StreamEvent::ToolCallDelta {
247 index: 0,
248 id: Some("call_1".into()),
249 function_name: Some("search".into()),
250 arguments_delta: "{}".into(),
251 };
252 match evt {
253 StreamEvent::ToolCallDelta { index, id, function_name, arguments_delta } => {
254 assert_eq!(index, 0);
255 assert_eq!(id.unwrap(), "call_1");
256 assert_eq!(function_name.unwrap(), "search");
257 assert_eq!(arguments_delta, "{}");
258 }
259 _ => panic!("Wrong variant"),
260 }
261 }
262
263 #[test]
264 fn test_stream_event_thinking_delta() {
265 let evt = StreamEvent::ThinkingDelta { delta: "thinking...".into() };
266 match evt {
267 StreamEvent::ThinkingDelta { delta } => assert_eq!(delta, "thinking..."),
268 _ => panic!("Wrong variant"),
269 }
270 }
271
272 #[test]
273 fn test_stream_event_usage() {
274 let usage = TokenUsage::new(10, 20);
275 let evt = StreamEvent::Usage { usage: usage.clone() };
276 match evt {
277 StreamEvent::Usage { usage: u } => {
278 assert_eq!(u.prompt_tokens, 10);
279 assert_eq!(u.completion_tokens, 20);
280 }
281 _ => panic!("Wrong variant"),
282 }
283 }
284
285 #[test]
286 fn test_stream_event_done() {
287 let evt = StreamEvent::Done { finish_reason: FinishReason::Stop, usage: None };
288 match evt {
289 StreamEvent::Done { finish_reason, usage } => {
290 assert!(matches!(finish_reason, FinishReason::Stop));
291 assert!(usage.is_none());
292 }
293 _ => panic!("Wrong variant"),
294 }
295 }
296
297 #[test]
298 fn test_stream_event_done_with_usage() {
299 let usage = TokenUsage::new(5, 10);
300 let evt = StreamEvent::Done { finish_reason: FinishReason::MaxTokens, usage: Some(usage) };
301 match evt {
302 StreamEvent::Done { finish_reason, usage } => {
303 assert!(matches!(finish_reason, FinishReason::MaxTokens));
304 assert_eq!(usage.unwrap().total_tokens, 15);
305 }
306 _ => panic!("Wrong variant"),
307 }
308 }
309
310 #[test]
311 fn test_stream_event_custom() {
312 let evt =
313 StreamEvent::Custom { event: "ping".into(), data: serde_json::json!({"key": "value"}) };
314 match evt {
315 StreamEvent::Custom { event, data } => {
316 assert_eq!(event, "ping");
317 assert_eq!(data["key"], "value");
318 }
319 _ => panic!("Wrong variant"),
320 }
321 }
322
323 #[test]
324 fn test_token_usage_serde() {
325 let usage = TokenUsage::new(100, 50);
326 let json = serde_json::to_string(&usage).unwrap();
327 let deserialized: TokenUsage = serde_json::from_str(&json).unwrap();
328 assert_eq!(deserialized.prompt_tokens, 100);
329 assert_eq!(deserialized.completion_tokens, 50);
330 assert_eq!(deserialized.total_tokens, 150);
331 }
332
333 #[test]
334 fn test_completion_response_fields() {
335 let usage = TokenUsage::new(10, 20);
336 let resp = CompletionResponse {
337 content: Some("Hello".into()),
338 thinking: None,
339 tool_calls: vec![],
340 usage,
341 model: "gpt-4".into(),
342 finish_reason: FinishReason::Stop,
343 latency_ms: 100,
344 cache_info: None,
345 id: Some("chatcmpl-123".into()),
346 created: Some(1700000000),
347 system_fingerprint: Some("fp_abc".into()),
348 refusal: None,
349 ..Default::default()
350 };
351 assert_eq!(resp.content.unwrap(), "Hello");
352 assert_eq!(resp.model, "gpt-4");
353 assert_eq!(resp.latency_ms, 100);
354 assert_eq!(resp.id.unwrap(), "chatcmpl-123");
355 assert_eq!(resp.created.unwrap(), 1700000000);
356 assert_eq!(resp.system_fingerprint.unwrap(), "fp_abc");
357 assert!(resp.refusal.is_none());
358 }
359
360 #[test]
361 fn test_completion_response_new_fields_serde() {
362 let usage = TokenUsage::new(10, 20);
363 let resp = CompletionResponse {
364 content: Some("Hi".into()),
365 thinking: None,
366 tool_calls: vec![],
367 usage,
368 model: "gpt-4o".into(),
369 finish_reason: FinishReason::Stop,
370 latency_ms: 200,
371 cache_info: None,
372 id: Some("chatcmpl-456".into()),
373 created: Some(1700000001),
374 system_fingerprint: Some("fp_xyz".into()),
375 refusal: Some("I cannot answer that.".into()),
376 ..Default::default()
377 };
378 let json = serde_json::to_string(&resp).unwrap();
379 let deserialized: CompletionResponse = serde_json::from_str(&json).unwrap();
380 assert_eq!(deserialized.id.unwrap(), "chatcmpl-456");
381 assert_eq!(deserialized.created.unwrap(), 1700000001);
382 assert_eq!(deserialized.system_fingerprint.unwrap(), "fp_xyz");
383 assert_eq!(deserialized.refusal.unwrap(), "I cannot answer that.");
384 }
385
386 #[test]
387 fn test_completion_response_new_fields_defaults() {
388 let usage = TokenUsage::new(10, 20);
389 let resp = CompletionResponse {
390 content: Some("Hi".into()),
391 thinking: None,
392 tool_calls: vec![],
393 usage,
394 model: "gpt-4o".into(),
395 finish_reason: FinishReason::Stop,
396 latency_ms: 200,
397 cache_info: None,
398 id: None,
399 created: None,
400 system_fingerprint: None,
401 refusal: None,
402 ..Default::default()
403 };
404 let json = serde_json::to_string(&resp).unwrap();
405 assert!(!json.contains("id"));
406 assert!(!json.contains("created"));
407 assert!(!json.contains("system_fingerprint"));
408 assert!(!json.contains("refusal"));
409 let deserialized: CompletionResponse = serde_json::from_str(&json).unwrap();
410 assert!(deserialized.id.is_none());
411 assert!(deserialized.created.is_none());
412 assert!(deserialized.system_fingerprint.is_none());
413 assert!(deserialized.refusal.is_none());
414 }
415
416 #[test]
417 fn test_token_usage_new_fields_serde() {
418 let mut usage = TokenUsage::new(100, 50);
419 usage.reasoning_tokens = Some(30);
420 usage.prompt_cache_hit_tokens = Some(20);
421 usage.prompt_cache_miss_tokens = Some(80);
422 usage.audio_tokens = Some(10);
423 let json = serde_json::to_string(&usage).unwrap();
424 let deserialized: TokenUsage = serde_json::from_str(&json).unwrap();
425 assert_eq!(deserialized.reasoning_tokens.unwrap(), 30);
426 assert_eq!(deserialized.prompt_cache_hit_tokens.unwrap(), 20);
427 assert_eq!(deserialized.prompt_cache_miss_tokens.unwrap(), 80);
428 assert_eq!(deserialized.audio_tokens.unwrap(), 10);
429 assert_eq!(deserialized.total_tokens, 150);
430 }
431
432 #[test]
433 fn test_token_usage_new_fields_defaults() {
434 let usage = TokenUsage::new(50, 25);
435 let json = serde_json::to_string(&usage).unwrap();
436 assert!(!json.contains("reasoning_tokens"));
437 assert!(!json.contains("prompt_cache_hit_tokens"));
438 assert!(!json.contains("prompt_cache_miss_tokens"));
439 assert!(!json.contains("audio_tokens"));
440 let deserialized: TokenUsage = serde_json::from_str(&json).unwrap();
441 assert!(deserialized.reasoning_tokens.is_none());
442 assert!(deserialized.prompt_cache_hit_tokens.is_none());
443 assert!(deserialized.prompt_cache_miss_tokens.is_none());
444 assert!(deserialized.audio_tokens.is_none());
445 }
446
447 #[test]
448 fn test_finish_reason_new_variants_serde() {
449 let cases = vec![
450 (FinishReason::PauseTurn, "\"pause_turn\""),
451 (FinishReason::Refusal, "\"refusal\""),
452 ];
453 for (reason, expected) in cases {
454 let json = serde_json::to_string(&reason).unwrap();
455 assert_eq!(json, expected);
456 let deserialized: FinishReason = serde_json::from_str(&json).unwrap();
457 assert!(
458 matches!(&deserialized, r if std::mem::discriminant(&reason) == std::mem::discriminant(r))
459 );
460 }
461 }
462
463 #[test]
464 fn test_stream_event_signature_delta() {
465 let evt = StreamEvent::SignatureDelta { signature: "sig_abc123".into() };
466 let json = serde_json::to_string(&evt).unwrap();
467 let deserialized: StreamEvent = serde_json::from_str(&json).unwrap();
468 match deserialized {
469 StreamEvent::SignatureDelta { signature } => {
470 assert_eq!(signature, "sig_abc123");
471 }
472 _ => panic!("Wrong variant"),
473 }
474 }
475
476 #[test]
477 fn test_stream_event_citations_delta() {
478 let citations = serde_json::json!([{"url": "https://example.com", "title": "Example"}]);
479 let evt = StreamEvent::CitationsDelta { citations: citations.clone() };
480 let json = serde_json::to_string(&evt).unwrap();
481 let deserialized: StreamEvent = serde_json::from_str(&json).unwrap();
482 match deserialized {
483 StreamEvent::CitationsDelta { citations: c } => {
484 assert_eq!(c[0]["url"], "https://example.com");
485 assert_eq!(c[0]["title"], "Example");
486 }
487 _ => panic!("Wrong variant"),
488 }
489 }
490
491 #[test]
492 fn test_stream_event_redacted_thinking_delta() {
493 let evt = StreamEvent::RedactedThinkingDelta { data: "redacted_thought".into() };
494 let json = serde_json::to_string(&evt).unwrap();
495 let deserialized: StreamEvent = serde_json::from_str(&json).unwrap();
496 match deserialized {
497 StreamEvent::RedactedThinkingDelta { data } => {
498 assert_eq!(data, "redacted_thought");
499 }
500 _ => panic!("Wrong variant"),
501 }
502 }
503}