use serde::{Deserialize, Serialize};
use serde_json::Value;
use super::Citation;
use super::cache::CacheInfo;
use super::tool::ToolCall;
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct TokenUsage {
pub prompt_tokens: u32,
pub completion_tokens: u32,
pub total_tokens: u32,
pub cached_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub reasoning_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub prompt_cache_hit_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub prompt_cache_miss_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub audio_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_write_5m_input_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_write_1h_input_tokens: Option<u32>,
}
impl TokenUsage {
pub fn new(prompt_tokens: u32, completion_tokens: u32) -> Self {
Self {
prompt_tokens,
completion_tokens,
total_tokens: prompt_tokens + completion_tokens,
cached_tokens: None,
reasoning_tokens: None,
prompt_cache_hit_tokens: None,
prompt_cache_miss_tokens: None,
audio_tokens: None,
cache_write_5m_input_tokens: None,
cache_write_1h_input_tokens: None,
}
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
pub struct CompletionResponse {
pub content: Option<String>,
pub thinking: Option<String>,
#[serde(default)]
pub tool_calls: Vec<ToolCall>,
pub usage: TokenUsage,
pub model: String,
pub finish_reason: FinishReason,
pub latency_ms: u64,
pub cache_info: Option<CacheInfo>,
#[serde(skip_serializing_if = "Option::is_none")]
pub id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub created: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub system_fingerprint: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub refusal: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub signature: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub redacted_thinking: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub citations: Option<Vec<Citation>>,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub enum FinishReason {
#[default]
#[serde(rename = "stop")]
Stop,
#[serde(rename = "tool_call")]
ToolCall,
#[serde(rename = "max_tokens")]
MaxTokens,
#[serde(rename = "content_filter")]
ContentFilter,
#[serde(rename = "pause_turn")]
PauseTurn,
#[serde(rename = "refusal")]
Refusal,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum StreamEvent {
#[serde(rename = "content_delta")]
ContentDelta { delta: String },
#[serde(rename = "tool_call_delta")]
ToolCallDelta {
index: usize,
#[serde(skip_serializing_if = "Option::is_none")]
id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
function_name: Option<String>,
arguments_delta: String,
},
#[serde(rename = "thinking_delta")]
ThinkingDelta { delta: String },
#[serde(rename = "image_delta")]
ImageDelta { media_type: String, delta: String },
#[serde(rename = "usage")]
Usage { usage: TokenUsage },
#[serde(rename = "done")]
Done {
finish_reason: FinishReason,
#[serde(skip_serializing_if = "Option::is_none")]
usage: Option<TokenUsage>,
},
#[serde(rename = "signature_delta")]
SignatureDelta { signature: String },
#[serde(rename = "citations_delta")]
CitationsDelta { citations: Value },
#[serde(rename = "redacted_thinking_delta")]
RedactedThinkingDelta { data: String },
#[serde(rename = "custom")]
Custom { event: String, data: Value },
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StreamChunk {
pub delta: String,
pub finish_reason: Option<String>,
pub usage: Option<TokenUsage>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StructuredResponse<T> {
pub parsed: T,
pub raw: String,
pub usage: TokenUsage,
pub model: String,
pub latency_ms: u64,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_token_usage_new() {
let usage = TokenUsage::new(100, 50);
assert_eq!(usage.prompt_tokens, 100);
assert_eq!(usage.completion_tokens, 50);
assert_eq!(usage.total_tokens, 150);
assert!(usage.cached_tokens.is_none());
}
#[test]
fn test_token_usage_new_zero() {
let usage = TokenUsage::new(0, 0);
assert_eq!(usage.total_tokens, 0);
assert_eq!(usage.prompt_tokens, 0);
assert_eq!(usage.completion_tokens, 0);
}
#[test]
fn test_finish_reason_serde() {
let cases = vec![
(FinishReason::Stop, "\"stop\""),
(FinishReason::ToolCall, "\"tool_call\""),
(FinishReason::MaxTokens, "\"max_tokens\""),
(FinishReason::ContentFilter, "\"content_filter\""),
];
for (reason, expected) in cases {
let json = serde_json::to_string(&reason).unwrap();
assert_eq!(json, expected);
let deserialized: FinishReason = serde_json::from_str(&json).unwrap();
assert!(
matches!(&deserialized, r if std::mem::discriminant(&reason) == std::mem::discriminant(r))
);
}
}
#[test]
fn test_stream_event_content_delta() {
let evt = StreamEvent::ContentDelta { delta: "Hello".into() };
match evt {
StreamEvent::ContentDelta { delta } => assert_eq!(delta, "Hello"),
_ => panic!("Wrong variant"),
}
}
#[test]
fn test_stream_event_tool_call_delta() {
let evt = StreamEvent::ToolCallDelta {
index: 0,
id: Some("call_1".into()),
function_name: Some("search".into()),
arguments_delta: "{}".into(),
};
match evt {
StreamEvent::ToolCallDelta { index, id, function_name, arguments_delta } => {
assert_eq!(index, 0);
assert_eq!(id.unwrap(), "call_1");
assert_eq!(function_name.unwrap(), "search");
assert_eq!(arguments_delta, "{}");
}
_ => panic!("Wrong variant"),
}
}
#[test]
fn test_stream_event_thinking_delta() {
let evt = StreamEvent::ThinkingDelta { delta: "thinking...".into() };
match evt {
StreamEvent::ThinkingDelta { delta } => assert_eq!(delta, "thinking..."),
_ => panic!("Wrong variant"),
}
}
#[test]
fn test_stream_event_usage() {
let usage = TokenUsage::new(10, 20);
let evt = StreamEvent::Usage { usage: usage.clone() };
match evt {
StreamEvent::Usage { usage: u } => {
assert_eq!(u.prompt_tokens, 10);
assert_eq!(u.completion_tokens, 20);
}
_ => panic!("Wrong variant"),
}
}
#[test]
fn test_stream_event_done() {
let evt = StreamEvent::Done { finish_reason: FinishReason::Stop, usage: None };
match evt {
StreamEvent::Done { finish_reason, usage } => {
assert!(matches!(finish_reason, FinishReason::Stop));
assert!(usage.is_none());
}
_ => panic!("Wrong variant"),
}
}
#[test]
fn test_stream_event_done_with_usage() {
let usage = TokenUsage::new(5, 10);
let evt = StreamEvent::Done { finish_reason: FinishReason::MaxTokens, usage: Some(usage) };
match evt {
StreamEvent::Done { finish_reason, usage } => {
assert!(matches!(finish_reason, FinishReason::MaxTokens));
assert_eq!(usage.unwrap().total_tokens, 15);
}
_ => panic!("Wrong variant"),
}
}
#[test]
fn test_stream_event_custom() {
let evt =
StreamEvent::Custom { event: "ping".into(), data: serde_json::json!({"key": "value"}) };
match evt {
StreamEvent::Custom { event, data } => {
assert_eq!(event, "ping");
assert_eq!(data["key"], "value");
}
_ => panic!("Wrong variant"),
}
}
#[test]
fn test_token_usage_serde() {
let usage = TokenUsage::new(100, 50);
let json = serde_json::to_string(&usage).unwrap();
let deserialized: TokenUsage = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized.prompt_tokens, 100);
assert_eq!(deserialized.completion_tokens, 50);
assert_eq!(deserialized.total_tokens, 150);
}
#[test]
fn test_completion_response_fields() {
let usage = TokenUsage::new(10, 20);
let resp = CompletionResponse {
content: Some("Hello".into()),
thinking: None,
tool_calls: vec![],
usage,
model: "gpt-4".into(),
finish_reason: FinishReason::Stop,
latency_ms: 100,
cache_info: None,
id: Some("chatcmpl-123".into()),
created: Some(1700000000),
system_fingerprint: Some("fp_abc".into()),
refusal: None,
..Default::default()
};
assert_eq!(resp.content.unwrap(), "Hello");
assert_eq!(resp.model, "gpt-4");
assert_eq!(resp.latency_ms, 100);
assert_eq!(resp.id.unwrap(), "chatcmpl-123");
assert_eq!(resp.created.unwrap(), 1700000000);
assert_eq!(resp.system_fingerprint.unwrap(), "fp_abc");
assert!(resp.refusal.is_none());
}
#[test]
fn test_completion_response_new_fields_serde() {
let usage = TokenUsage::new(10, 20);
let resp = CompletionResponse {
content: Some("Hi".into()),
thinking: None,
tool_calls: vec![],
usage,
model: "gpt-4o".into(),
finish_reason: FinishReason::Stop,
latency_ms: 200,
cache_info: None,
id: Some("chatcmpl-456".into()),
created: Some(1700000001),
system_fingerprint: Some("fp_xyz".into()),
refusal: Some("I cannot answer that.".into()),
..Default::default()
};
let json = serde_json::to_string(&resp).unwrap();
let deserialized: CompletionResponse = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized.id.unwrap(), "chatcmpl-456");
assert_eq!(deserialized.created.unwrap(), 1700000001);
assert_eq!(deserialized.system_fingerprint.unwrap(), "fp_xyz");
assert_eq!(deserialized.refusal.unwrap(), "I cannot answer that.");
}
#[test]
fn test_completion_response_new_fields_defaults() {
let usage = TokenUsage::new(10, 20);
let resp = CompletionResponse {
content: Some("Hi".into()),
thinking: None,
tool_calls: vec![],
usage,
model: "gpt-4o".into(),
finish_reason: FinishReason::Stop,
latency_ms: 200,
cache_info: None,
id: None,
created: None,
system_fingerprint: None,
refusal: None,
..Default::default()
};
let json = serde_json::to_string(&resp).unwrap();
assert!(!json.contains("id"));
assert!(!json.contains("created"));
assert!(!json.contains("system_fingerprint"));
assert!(!json.contains("refusal"));
let deserialized: CompletionResponse = serde_json::from_str(&json).unwrap();
assert!(deserialized.id.is_none());
assert!(deserialized.created.is_none());
assert!(deserialized.system_fingerprint.is_none());
assert!(deserialized.refusal.is_none());
}
#[test]
fn test_token_usage_new_fields_serde() {
let mut usage = TokenUsage::new(100, 50);
usage.reasoning_tokens = Some(30);
usage.prompt_cache_hit_tokens = Some(20);
usage.prompt_cache_miss_tokens = Some(80);
usage.audio_tokens = Some(10);
let json = serde_json::to_string(&usage).unwrap();
let deserialized: TokenUsage = serde_json::from_str(&json).unwrap();
assert_eq!(deserialized.reasoning_tokens.unwrap(), 30);
assert_eq!(deserialized.prompt_cache_hit_tokens.unwrap(), 20);
assert_eq!(deserialized.prompt_cache_miss_tokens.unwrap(), 80);
assert_eq!(deserialized.audio_tokens.unwrap(), 10);
assert_eq!(deserialized.total_tokens, 150);
}
#[test]
fn test_token_usage_new_fields_defaults() {
let usage = TokenUsage::new(50, 25);
let json = serde_json::to_string(&usage).unwrap();
assert!(!json.contains("reasoning_tokens"));
assert!(!json.contains("prompt_cache_hit_tokens"));
assert!(!json.contains("prompt_cache_miss_tokens"));
assert!(!json.contains("audio_tokens"));
let deserialized: TokenUsage = serde_json::from_str(&json).unwrap();
assert!(deserialized.reasoning_tokens.is_none());
assert!(deserialized.prompt_cache_hit_tokens.is_none());
assert!(deserialized.prompt_cache_miss_tokens.is_none());
assert!(deserialized.audio_tokens.is_none());
}
#[test]
fn test_finish_reason_new_variants_serde() {
let cases = vec![
(FinishReason::PauseTurn, "\"pause_turn\""),
(FinishReason::Refusal, "\"refusal\""),
];
for (reason, expected) in cases {
let json = serde_json::to_string(&reason).unwrap();
assert_eq!(json, expected);
let deserialized: FinishReason = serde_json::from_str(&json).unwrap();
assert!(
matches!(&deserialized, r if std::mem::discriminant(&reason) == std::mem::discriminant(r))
);
}
}
#[test]
fn test_stream_event_signature_delta() {
let evt = StreamEvent::SignatureDelta { signature: "sig_abc123".into() };
let json = serde_json::to_string(&evt).unwrap();
let deserialized: StreamEvent = serde_json::from_str(&json).unwrap();
match deserialized {
StreamEvent::SignatureDelta { signature } => {
assert_eq!(signature, "sig_abc123");
}
_ => panic!("Wrong variant"),
}
}
#[test]
fn test_stream_event_citations_delta() {
let citations = serde_json::json!([{"url": "https://example.com", "title": "Example"}]);
let evt = StreamEvent::CitationsDelta { citations: citations.clone() };
let json = serde_json::to_string(&evt).unwrap();
let deserialized: StreamEvent = serde_json::from_str(&json).unwrap();
match deserialized {
StreamEvent::CitationsDelta { citations: c } => {
assert_eq!(c[0]["url"], "https://example.com");
assert_eq!(c[0]["title"], "Example");
}
_ => panic!("Wrong variant"),
}
}
#[test]
fn test_stream_event_redacted_thinking_delta() {
let evt = StreamEvent::RedactedThinkingDelta { data: "redacted_thought".into() };
let json = serde_json::to_string(&evt).unwrap();
let deserialized: StreamEvent = serde_json::from_str(&json).unwrap();
match deserialized {
StreamEvent::RedactedThinkingDelta { data } => {
assert_eq!(data, "redacted_thought");
}
_ => panic!("Wrong variant"),
}
}
}