Skip to main content

xz_provider/types/
response.rs

1use serde::{Deserialize, Serialize};
2use serde_json::Value;
3
4use super::Citation;
5use super::cache::CacheInfo;
6use super::tool::ToolCall;
7
8/// Token 用量
9#[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    /// 缓存命中的 token 数(Prompt Caching)
15    pub cached_tokens: Option<u32>,
16    /// 推理/思考 token 数(DeepSeek reasoning、OpenAI o-series)
17    #[serde(skip_serializing_if = "Option::is_none")]
18    pub reasoning_tokens: Option<u32>,
19    /// 缓存命中(DeepSeek prompt cache hit)
20    #[serde(skip_serializing_if = "Option::is_none")]
21    pub prompt_cache_hit_tokens: Option<u32>,
22    /// 缓存未命中(DeepSeek prompt cache miss)
23    #[serde(skip_serializing_if = "Option::is_none")]
24    pub prompt_cache_miss_tokens: Option<u32>,
25    /// 音频输入 token 数
26    #[serde(skip_serializing_if = "Option::is_none")]
27    pub audio_tokens: Option<u32>,
28    /// 写入 5 分钟生命周期的缓存的 token 数
29    #[serde(skip_serializing_if = "Option::is_none")]
30    pub cache_write_5m_input_tokens: Option<u32>,
31    /// 写入 1 小时生命周期的缓存的 token 数
32    #[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/// 补全响应
54#[derive(Debug, Clone, Default, Serialize, Deserialize)]
55pub struct CompletionResponse {
56    /// 文本内容(可能为空,如 LLM 仅发起 tool call)
57    pub content: Option<String>,
58    /// 思考过程(Claude extended thinking / DeepSeek reasoning)
59    pub thinking: Option<String>,
60    /// LLM 发起的工具调用
61    #[serde(default)]
62    pub tool_calls: Vec<ToolCall>,
63    /// Token 用量
64    pub usage: TokenUsage,
65    /// 实际使用的模型名(可能与请求不同,如自动回退后)
66    pub model: String,
67    /// 结束原因
68    pub finish_reason: FinishReason,
69    /// 请求延迟(毫秒)
70    pub latency_ms: u64,
71    /// 缓存命中信息
72    pub cache_info: Option<CacheInfo>,
73    /// 响应标识符(如 "chatcmpl-123")
74    #[serde(skip_serializing_if = "Option::is_none")]
75    pub id: Option<String>,
76    /// Unix 时间戳
77    #[serde(skip_serializing_if = "Option::is_none")]
78    pub created: Option<u64>,
79    /// 后端指纹
80    #[serde(skip_serializing_if = "Option::is_none")]
81    pub system_fingerprint: Option<String>,
82    /// 内容拒绝消息
83    #[serde(skip_serializing_if = "Option::is_none")]
84    pub refusal: Option<String>,
85    /// 思考链签名(Anthropic extended thinking 验证用)
86    #[serde(skip_serializing_if = "Option::is_none")]
87    pub signature: Option<String>,
88    /// 加密思考链数据(API 返回已删节的加密内容)
89    #[serde(skip_serializing_if = "Option::is_none")]
90    pub redacted_thinking: Option<String>,
91    /// 引用来源(Anthropic Citations 功能)
92    #[serde(skip_serializing_if = "Option::is_none")]
93    pub citations: Option<Vec<Citation>>,
94}
95
96/// 结束原因
97#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
98pub enum FinishReason {
99    /// 正常结束
100    #[default]
101    #[serde(rename = "stop")]
102    Stop,
103    /// LLM 请求调用工具
104    #[serde(rename = "tool_call")]
105    ToolCall,
106    /// 达到 token 上限
107    #[serde(rename = "max_tokens")]
108    MaxTokens,
109    /// 内容过滤
110    #[serde(rename = "content_filter")]
111    ContentFilter,
112    /// 暂停对话回合(如 Claude 的 pause_turn)
113    #[serde(rename = "pause_turn")]
114    PauseTurn,
115    /// 拒绝回答(如内容策略拒绝)
116    #[serde(rename = "refusal")]
117    Refusal,
118}
119
120/// 流式响应事件类型
121#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
122#[serde(tag = "type")]
123pub enum StreamEvent {
124    /// 文本内容增量
125    #[serde(rename = "content_delta")]
126    ContentDelta { delta: String },
127
128    /// 工具调用增量 —— 函数名和参数片段逐步拼接
129    #[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    /// 思考过程增量(Claude extended thinking、DeepSeek reasoning)
140    #[serde(rename = "thinking_delta")]
141    ThinkingDelta { delta: String },
142
143    /// 图像增量(Gemini 图像生成等多模态输出场景)
144    #[serde(rename = "image_delta")]
145    ImageDelta { media_type: String, delta: String },
146
147    /// Token 用量更新
148    #[serde(rename = "usage")]
149    Usage { usage: TokenUsage },
150
151    /// 流结束信号
152    #[serde(rename = "done")]
153    Done {
154        finish_reason: FinishReason,
155        #[serde(skip_serializing_if = "Option::is_none")]
156        usage: Option<TokenUsage>,
157    },
158
159    /// 签名增量(模型输出签名,用于验证)
160    #[serde(rename = "signature_delta")]
161    SignatureDelta { signature: String },
162
163    /// 引用来源增量
164    #[serde(rename = "citations_delta")]
165    CitationsDelta { citations: Value },
166
167    /// 隐式思考过程增量(redacted thinking,API 返回已删节版本)
168    #[serde(rename = "redacted_thinking_delta")]
169    RedactedThinkingDelta { data: String },
170
171    /// Provider 特有的事件(可扩展)
172    #[serde(rename = "custom")]
173    Custom { event: String, data: Value },
174}
175
176// ── Old types compat ──
177
178/// (Deprecated) 流式补全块 —— 使用 StreamEvent 代替
179#[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/// (Deprecated) 结构化输出响应 —— 直接使用 CompletionResponse
187#[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}