Skip to main content

xz_provider/types/
request.rs

1use std::collections::HashMap;
2use std::time::Duration;
3
4use serde::{Deserialize, Serialize};
5use serde_json::Value;
6
7use super::message::Message;
8use super::tool::ToolDefinition;
9
10/// 服务处理层级 — 控制请求的处理优先级/延迟/吞吐量
11///
12/// OpenAI: `auto`, `default`;Anthropic: `default`;部分企业端点支持 `flex`/`scale`
13#[derive(Debug, Clone, Serialize, Deserialize)]
14#[serde(rename_all = "snake_case")]
15pub enum ServiceTier {
16    /// 由平台自动选择最合适的层级
17    Auto,
18    /// 标准处理层级(默认)
19    Default,
20    /// 灵活层级(较低优先级,适合背景任务)
21    Flex,
22    /// 扩展层级(更高吞吐量,适合批量场景)
23    Scale,
24    /// 优先级层级(最快响应,适合交互场景)
25    Priority,
26}
27
28/// 统一请求类型 —— 数据平面(发给 LLM 的内容)
29/// 所有能力通过 option 字段开启,不需要多个方法
30#[derive(Debug, Clone, Serialize, Deserialize)]
31pub struct CompletionRequest {
32    /// 目标模型。为 None 时由路由层根据 RouteContext 决定。
33    pub model: Option<String>,
34    pub messages: Vec<Message>,
35
36    // ── 工具调用 ──
37    /// 提供可用工具列表
38    #[serde(skip_serializing_if = "Option::is_none")]
39    pub tools: Option<Vec<ToolDefinition>>,
40    /// 控制工具调用行为
41    #[serde(skip_serializing_if = "Option::is_none")]
42    pub tool_choice: Option<ToolChoice>,
43
44    // ── 结构化输出 ──
45    /// 要求 LLM 按 JSON Schema 输出
46    #[serde(skip_serializing_if = "Option::is_none")]
47    pub response_format: Option<ResponseFormat>,
48
49    // ── 生成参数 ──
50    pub temperature: Option<f32>,
51    pub max_tokens: Option<usize>,
52    /// OpenAI o-series 使用 max_completion_tokens 而非 max_tokens
53    #[serde(skip_serializing_if = "Option::is_none")]
54    pub max_completion_tokens: Option<usize>,
55    pub top_p: Option<f32>,
56    /// top_k 采样参数(Claude、Gemini 支持)
57    #[serde(skip_serializing_if = "Option::is_none")]
58    pub top_k: Option<u32>,
59    pub stop: Option<Vec<String>>,
60    pub frequency_penalty: Option<f32>,
61    pub presence_penalty: Option<f32>,
62    /// 随机种子,用于可复现输出(评测、调试场景必需)
63    #[serde(skip_serializing_if = "Option::is_none")]
64    pub seed: Option<u64>,
65    /// 推理努力程度(OpenAI o-series reasoning_effort / Claude thinking budget)
66    #[serde(skip_serializing_if = "Option::is_none")]
67    pub reasoning_effort: Option<ReasoningEffort>,
68    /// 返回 log probabilities(调试、结构化分析场景)
69    #[serde(skip_serializing_if = "Option::is_none")]
70    pub logprobs: Option<bool>,
71    /// token 偏置(调整特定 token 的出现概率)
72    #[serde(skip_serializing_if = "Option::is_none")]
73    pub logit_bias: Option<HashMap<String, f32>>,
74
75    // ── 流式选项 ──
76    /// 流式模式下是否在末尾返回 usage
77    pub stream_include_usage: Option<bool>,
78
79    // ── 其他协议字段 ──
80    /// 并行工具调用控制(OpenAI parallel_tool_calls)
81    #[serde(skip_serializing_if = "Option::is_none")]
82    pub parallel_tool_calls: Option<bool>,
83    /// 终端用户标识(用于监控/滥用检测)
84    #[serde(skip_serializing_if = "Option::is_none")]
85    pub user: Option<String>,
86    /// 请求级元数据(OpenAI/Anthropic 支持,用于追踪/蒸馏)
87    #[serde(skip_serializing_if = "Option::is_none")]
88    pub metadata: Option<HashMap<String, Value>>,
89    /// 是否存储请求以供蒸馏/评测(OpenAI store)
90    #[serde(skip_serializing_if = "Option::is_none")]
91    pub store: Option<bool>,
92    /// 服务处理层级(OpenAI/Anthropic)
93    #[serde(skip_serializing_if = "Option::is_none")]
94    pub service_tier: Option<ServiceTier>,
95    /// 思考模式控制(DeepSeek R1 thinking)
96    #[serde(skip_serializing_if = "Option::is_none")]
97    pub thinking: Option<ThinkingConfig>,
98
99    /// 请求唯一标识(用于 tracing,自动生成)
100    #[serde(skip)]
101    pub request_id: String,
102}
103
104impl CompletionRequest {
105    pub fn new(model: impl Into<String>, messages: Vec<Message>) -> Self {
106        Self {
107            model: Some(model.into()),
108            messages,
109            tools: None,
110            tool_choice: None,
111            response_format: None,
112            temperature: None,
113            max_tokens: None,
114            max_completion_tokens: None,
115            top_p: None,
116            top_k: None,
117            stop: None,
118            frequency_penalty: None,
119            presence_penalty: None,
120            seed: None,
121            reasoning_effort: None,
122            logprobs: None,
123            logit_bias: None,
124            stream_include_usage: None,
125            parallel_tool_calls: None,
126            user: None,
127            metadata: None,
128            store: None,
129            service_tier: None,
130            thinking: None,
131            request_id: uuid::Uuid::new_v4().to_string(),
132        }
133    }
134}
135
136impl Default for CompletionRequest {
137    fn default() -> Self {
138        Self {
139            model: None,
140            messages: Vec::new(),
141            tools: None,
142            tool_choice: None,
143            response_format: None,
144            temperature: None,
145            max_tokens: None,
146            max_completion_tokens: None,
147            top_p: None,
148            top_k: None,
149            stop: None,
150            frequency_penalty: None,
151            presence_penalty: None,
152            seed: None,
153            reasoning_effort: None,
154            logprobs: None,
155            logit_bias: None,
156            stream_include_usage: None,
157            parallel_tool_calls: None,
158            user: None,
159            metadata: None,
160            store: None,
161            service_tier: None,
162            thinking: None,
163            request_id: uuid::Uuid::new_v4().to_string(),
164        }
165    }
166}
167
168/// 控制工具调用行为
169#[derive(Debug, Clone, Serialize, Deserialize)]
170#[serde(rename_all = "snake_case")]
171pub enum ToolChoice {
172    /// LLM 自主决定是否调用工具
173    Auto,
174    /// 必须调用工具
175    Required,
176    /// 禁止调用工具
177    Disabled,
178    /// 指定调用某个工具
179    Specific { name: String },
180}
181
182/// 结构化输出格式要求
183#[derive(Debug, Clone, Serialize, Deserialize)]
184#[serde(tag = "type")]
185pub enum ResponseFormat {
186    /// 要求输出合法 JSON(不限定 schema)
187    #[serde(rename = "json_object")]
188    Json,
189    /// 要求输出符合指定 JSON Schema(语法级保证)
190    JsonSchema { schema: Value, name: String },
191}
192
193/// 推理努力程度 — 控制 o-series / thinking 模型的推理深度
194#[derive(Debug, Clone, Serialize, Deserialize)]
195pub enum ReasoningEffort {
196    #[serde(rename = "low")]
197    Low,
198    #[serde(rename = "medium")]
199    Medium,
200    #[serde(rename = "high")]
201    High,
202}
203
204/// 思考模式类型 — 控制模型是否/如何展示推理过程
205#[derive(Debug, Clone, Serialize, Deserialize)]
206#[serde(tag = "type", rename_all = "snake_case")]
207pub enum ThinkingType {
208    /// 启用思考模式,可选择设置 token 预算
209    Enabled {
210        /// 思考过程 token 预算(DeepSeek R1)
211        #[serde(skip_serializing_if = "Option::is_none")]
212        budget_tokens: Option<u32>,
213    },
214    /// 禁用思考模式
215    Disabled,
216    /// 自适应思考模式(模型自行决定)
217    Adaptive,
218}
219
220/// 思考内容展示方式
221#[derive(Debug, Clone, Serialize, Deserialize)]
222#[serde(rename_all = "snake_case")]
223pub enum ThinkingDisplay {
224    /// 返回摘要形式的思考过程
225    Summarized,
226    /// 省略思考过程
227    Omitted,
228}
229
230/// 思考模式配置 — DeepSeek R1 等模型的思考过程控制
231#[derive(Debug, Clone, Serialize, Deserialize)]
232pub struct ThinkingConfig {
233    /// 思考模式类型(启用/禁用/自适应)
234    #[serde(flatten)]
235    pub thinking_type: ThinkingType,
236    /// 思考内容展示方式
237    #[serde(skip_serializing_if = "Option::is_none")]
238    pub display: Option<ThinkingDisplay>,
239}
240
241/// 请求选项 —— 控制平面(告诉 Provider 怎么执行)
242/// 与 CompletionRequest 分离,避免超时/取消等控制参数混入序列化数据
243#[derive(Debug, Clone, Default)]
244pub struct RequestOptions {
245    /// 请求级超时。None 表示使用 Provider 默认超时。
246    pub timeout: Option<Duration>,
247    /// 取消令牌
248    pub cancel: Option<crate::cancel::CancellationToken>,
249    /// 请求级元数据(用于透传 trace_id 等)
250    pub metadata: Option<HashMap<String, Value>>,
251}
252
253// ── Old StructuredRequest compat ──
254
255/// (Deprecated) 结构化输出请求 —— 使用 CompletionRequest.response_format 代替
256#[derive(Debug, Clone, Serialize, Deserialize)]
257pub struct StructuredRequest {
258    pub model: String,
259    pub messages: Vec<Message>,
260    pub response_schema: Value,
261    pub temperature: Option<f32>,
262    pub max_tokens: Option<usize>,
263    pub request_id: String,
264}
265
266#[cfg(test)]
267mod tests {
268    use super::*;
269
270    #[test]
271    fn test_completion_request_new() {
272        let req = CompletionRequest::new("gpt-4", vec![Message::user("Hello")]);
273        assert_eq!(req.model.as_deref(), Some("gpt-4"));
274        assert_eq!(req.messages.len(), 1);
275        assert!(!req.request_id.is_empty());
276        assert!(req.tools.is_none());
277        assert!(req.temperature.is_none());
278        assert!(req.max_tokens.is_none());
279        assert!(req.top_p.is_none());
280        assert!(req.stop.is_none());
281    }
282
283    #[test]
284    fn test_completion_request_new_empty_messages() {
285        let req = CompletionRequest::new("gpt-4", vec![]);
286        assert_eq!(req.model.as_deref(), Some("gpt-4"));
287        assert!(req.messages.is_empty());
288    }
289
290    #[test]
291    fn test_completion_request_default_model_none() {
292        let req = CompletionRequest::default();
293        assert!(req.model.is_none());
294    }
295
296    #[test]
297    fn test_completion_request_unique_request_id() {
298        let req1 = CompletionRequest::new("gpt-4", vec![]);
299        let req2 = CompletionRequest::new("gpt-4", vec![]);
300        assert_ne!(req1.request_id, req2.request_id);
301    }
302
303    #[test]
304    fn test_request_options_default() {
305        let opts = RequestOptions::default();
306        assert!(opts.timeout.is_none());
307        assert!(opts.cancel.is_none());
308        assert!(opts.metadata.is_none());
309    }
310
311    #[test]
312    fn test_tool_choice_serde() {
313        let choices = vec![
314            (ToolChoice::Auto, r#""auto""#),
315            (ToolChoice::Required, r#""required""#),
316            (ToolChoice::Disabled, r#""disabled""#),
317        ];
318        for (choice, expected) in choices {
319            let json = serde_json::to_string(&choice).unwrap();
320            assert_eq!(json, expected);
321        }
322    }
323
324    #[test]
325    fn test_tool_choice_specific_serde() {
326        let choice = ToolChoice::Specific { name: "search".into() };
327        let json = serde_json::to_string(&choice).unwrap();
328        assert!(json.contains("search"));
329    }
330
331    #[test]
332    fn test_response_format_json() {
333        let fmt = ResponseFormat::Json;
334        let json = serde_json::to_string(&fmt).unwrap();
335        assert_eq!(json, r#"{"type":"json_object"}"#);
336    }
337
338    #[test]
339    fn test_response_format_json_schema() {
340        let schema = serde_json::json!({"type": "object"});
341        let fmt = ResponseFormat::JsonSchema { schema: schema.clone(), name: "MySchema".into() };
342        let json = serde_json::to_string(&fmt).unwrap();
343        assert!(json.contains("MySchema"));
344        assert!(json.contains("type"));
345    }
346
347    #[test]
348    fn test_completion_request_serialize() {
349        let req = CompletionRequest::new("gpt-4", vec![Message::user("Hi")]);
350        let json = serde_json::to_string(&req).unwrap();
351        assert!(json.contains("gpt-4"));
352        assert!(json.contains("Hi"));
353        // request_id is skipped in serialization
354        assert!(!json.contains("request_id"));
355    }
356
357    #[test]
358    fn test_reasoning_effort_serde() {
359        assert_eq!(serde_json::to_string(&ReasoningEffort::Low).unwrap(), r#""low""#);
360        assert_eq!(serde_json::to_string(&ReasoningEffort::Medium).unwrap(), r#""medium""#);
361        assert_eq!(serde_json::to_string(&ReasoningEffort::High).unwrap(), r#""high""#);
362    }
363
364    #[test]
365    fn test_completion_request_new_fields() {
366        let mut req = CompletionRequest::new("gpt-4", vec![]);
367        req.seed = Some(42);
368        req.reasoning_effort = Some(ReasoningEffort::Medium);
369        req.max_completion_tokens = Some(4000);
370        req.top_k = Some(50);
371        req.logprobs = Some(true);
372        req.logit_bias = Some(HashMap::from([("hello".into(), 0.5)]));
373        let json = serde_json::to_string(&req).unwrap();
374        assert!(json.contains("42"));
375        assert!(json.contains("medium"));
376        assert!(json.contains("4000"));
377        assert!(json.contains("50"));
378    }
379
380    #[test]
381    fn test_service_tier_serde() {
382        assert_eq!(serde_json::to_string(&ServiceTier::Auto).unwrap(), r#""auto""#);
383        assert_eq!(serde_json::to_string(&ServiceTier::Default).unwrap(), r#""default""#);
384        assert_eq!(serde_json::to_string(&ServiceTier::Flex).unwrap(), r#""flex""#);
385        assert_eq!(serde_json::to_string(&ServiceTier::Scale).unwrap(), r#""scale""#);
386        assert_eq!(serde_json::to_string(&ServiceTier::Priority).unwrap(), r#""priority""#);
387    }
388
389    #[test]
390    fn test_thinking_type_enabled_serde() {
391        let enabled = ThinkingType::Enabled { budget_tokens: Some(4096) };
392        let json = serde_json::to_string(&enabled).unwrap();
393        assert!(json.contains(r#""type":"enabled""#));
394        assert!(json.contains("4096"));
395    }
396
397    #[test]
398    fn test_thinking_type_disabled_serde() {
399        let disabled = ThinkingType::Disabled;
400        let json = serde_json::to_string(&disabled).unwrap();
401        assert_eq!(json, r#"{"type":"disabled"}"#);
402    }
403
404    #[test]
405    fn test_thinking_type_adaptive_serde() {
406        let adaptive = ThinkingType::Adaptive;
407        let json = serde_json::to_string(&adaptive).unwrap();
408        assert_eq!(json, r#"{"type":"adaptive"}"#);
409    }
410
411    #[test]
412    fn test_thinking_display_serde() {
413        assert_eq!(serde_json::to_string(&ThinkingDisplay::Summarized).unwrap(), r#""summarized""#);
414        assert_eq!(serde_json::to_string(&ThinkingDisplay::Omitted).unwrap(), r#""omitted""#);
415    }
416
417    #[test]
418    fn test_thinking_config_serde() {
419        let config = ThinkingConfig {
420            thinking_type: ThinkingType::Enabled { budget_tokens: Some(2048) },
421            display: Some(ThinkingDisplay::Summarized),
422        };
423        let json = serde_json::to_string(&config).unwrap();
424        assert!(json.contains(r#""type":"enabled""#));
425        assert!(json.contains("2048"));
426        assert!(json.contains(r#""display":"summarized""#));
427    }
428
429    #[test]
430    fn test_completion_request_new_fields_serialize() {
431        let mut req = CompletionRequest::new("gpt-4", vec![Message::user("Hello")]);
432        req.parallel_tool_calls = Some(true);
433        req.user = Some("user-123".into());
434        req.metadata = Some(HashMap::from([("session_id".into(), Value::String("abc".into()))]));
435        req.store = Some(true);
436        req.service_tier = Some(ServiceTier::Auto);
437        req.thinking = Some(ThinkingConfig {
438            thinking_type: ThinkingType::Enabled { budget_tokens: Some(4096) },
439            display: None,
440        });
441        let json = serde_json::to_string(&req).unwrap();
442        assert!(json.contains("true")); // parallel_tool_calls
443        assert!(json.contains("user-123"));
444        assert!(json.contains("session_id"));
445        assert!(json.contains("auto")); // service_tier
446        assert!(json.contains("enabled")); // thinking type
447        assert!(json.contains("4096")); // budget_tokens
448    }
449
450    #[test]
451    fn test_completion_request_new_fields_absent_when_none() {
452        let req = CompletionRequest::new("gpt-4", vec![Message::user("Hi")]);
453        let json = serde_json::to_string(&req).unwrap();
454        assert!(!json.contains("parallel_tool_calls"));
455        assert!(!json.contains("\"user\":"));
456        assert!(!json.contains("metadata"));
457        assert!(!json.contains("store"));
458        assert!(!json.contains("service_tier"));
459        assert!(!json.contains("thinking"));
460    }
461}