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#[derive(Debug, Clone, Serialize, Deserialize)]
14#[serde(rename_all = "snake_case")]
15pub enum ServiceTier {
16 Auto,
18 Default,
20 Flex,
22 Scale,
24 Priority,
26}
27
28#[derive(Debug, Clone, Serialize, Deserialize)]
31pub struct CompletionRequest {
32 pub model: Option<String>,
34 pub messages: Vec<Message>,
35
36 #[serde(skip_serializing_if = "Option::is_none")]
39 pub tools: Option<Vec<ToolDefinition>>,
40 #[serde(skip_serializing_if = "Option::is_none")]
42 pub tool_choice: Option<ToolChoice>,
43
44 #[serde(skip_serializing_if = "Option::is_none")]
47 pub response_format: Option<ResponseFormat>,
48
49 pub temperature: Option<f32>,
51 pub max_tokens: Option<usize>,
52 #[serde(skip_serializing_if = "Option::is_none")]
54 pub max_completion_tokens: Option<usize>,
55 pub top_p: Option<f32>,
56 #[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 #[serde(skip_serializing_if = "Option::is_none")]
64 pub seed: Option<u64>,
65 #[serde(skip_serializing_if = "Option::is_none")]
67 pub reasoning_effort: Option<ReasoningEffort>,
68 #[serde(skip_serializing_if = "Option::is_none")]
70 pub logprobs: Option<bool>,
71 #[serde(skip_serializing_if = "Option::is_none")]
73 pub logit_bias: Option<HashMap<String, f32>>,
74
75 pub stream_include_usage: Option<bool>,
78
79 #[serde(skip_serializing_if = "Option::is_none")]
82 pub parallel_tool_calls: Option<bool>,
83 #[serde(skip_serializing_if = "Option::is_none")]
85 pub user: Option<String>,
86 #[serde(skip_serializing_if = "Option::is_none")]
88 pub metadata: Option<HashMap<String, Value>>,
89 #[serde(skip_serializing_if = "Option::is_none")]
91 pub store: Option<bool>,
92 #[serde(skip_serializing_if = "Option::is_none")]
94 pub service_tier: Option<ServiceTier>,
95 #[serde(skip_serializing_if = "Option::is_none")]
97 pub thinking: Option<ThinkingConfig>,
98
99 #[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#[derive(Debug, Clone, Serialize, Deserialize)]
170#[serde(rename_all = "snake_case")]
171pub enum ToolChoice {
172 Auto,
174 Required,
176 Disabled,
178 Specific { name: String },
180}
181
182#[derive(Debug, Clone, Serialize, Deserialize)]
184#[serde(tag = "type")]
185pub enum ResponseFormat {
186 #[serde(rename = "json_object")]
188 Json,
189 JsonSchema { schema: Value, name: String },
191}
192
193#[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#[derive(Debug, Clone, Serialize, Deserialize)]
206#[serde(tag = "type", rename_all = "snake_case")]
207pub enum ThinkingType {
208 Enabled {
210 #[serde(skip_serializing_if = "Option::is_none")]
212 budget_tokens: Option<u32>,
213 },
214 Disabled,
216 Adaptive,
218}
219
220#[derive(Debug, Clone, Serialize, Deserialize)]
222#[serde(rename_all = "snake_case")]
223pub enum ThinkingDisplay {
224 Summarized,
226 Omitted,
228}
229
230#[derive(Debug, Clone, Serialize, Deserialize)]
232pub struct ThinkingConfig {
233 #[serde(flatten)]
235 pub thinking_type: ThinkingType,
236 #[serde(skip_serializing_if = "Option::is_none")]
238 pub display: Option<ThinkingDisplay>,
239}
240
241#[derive(Debug, Clone, Default)]
244pub struct RequestOptions {
245 pub timeout: Option<Duration>,
247 pub cancel: Option<crate::cancel::CancellationToken>,
249 pub metadata: Option<HashMap<String, Value>>,
251}
252
253#[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 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")); assert!(json.contains("user-123"));
444 assert!(json.contains("session_id"));
445 assert!(json.contains("auto")); assert!(json.contains("enabled")); assert!(json.contains("4096")); }
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}