use std::collections::HashMap;
use std::time::Duration;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use super::message::Message;
use super::tool::ToolDefinition;
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ServiceTier {
Auto,
Default,
Flex,
Scale,
Priority,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct CompletionRequest {
pub model: Option<String>,
pub messages: Vec<Message>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<Vec<ToolDefinition>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_choice: Option<ToolChoice>,
#[serde(skip_serializing_if = "Option::is_none")]
pub response_format: Option<ResponseFormat>,
pub temperature: Option<f32>,
pub max_tokens: Option<usize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_completion_tokens: Option<usize>,
pub top_p: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_k: Option<u32>,
pub stop: Option<Vec<String>>,
pub frequency_penalty: Option<f32>,
pub presence_penalty: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub seed: Option<u64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub reasoning_effort: Option<ReasoningEffort>,
#[serde(skip_serializing_if = "Option::is_none")]
pub logprobs: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub logit_bias: Option<HashMap<String, f32>>,
pub stream_include_usage: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub parallel_tool_calls: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub metadata: Option<HashMap<String, Value>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub store: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub service_tier: Option<ServiceTier>,
#[serde(skip_serializing_if = "Option::is_none")]
pub thinking: Option<ThinkingConfig>,
#[serde(skip)]
pub request_id: String,
}
impl CompletionRequest {
pub fn new(model: impl Into<String>, messages: Vec<Message>) -> Self {
Self {
model: Some(model.into()),
messages,
tools: None,
tool_choice: None,
response_format: None,
temperature: None,
max_tokens: None,
max_completion_tokens: None,
top_p: None,
top_k: None,
stop: None,
frequency_penalty: None,
presence_penalty: None,
seed: None,
reasoning_effort: None,
logprobs: None,
logit_bias: None,
stream_include_usage: None,
parallel_tool_calls: None,
user: None,
metadata: None,
store: None,
service_tier: None,
thinking: None,
request_id: uuid::Uuid::new_v4().to_string(),
}
}
}
impl Default for CompletionRequest {
fn default() -> Self {
Self {
model: None,
messages: Vec::new(),
tools: None,
tool_choice: None,
response_format: None,
temperature: None,
max_tokens: None,
max_completion_tokens: None,
top_p: None,
top_k: None,
stop: None,
frequency_penalty: None,
presence_penalty: None,
seed: None,
reasoning_effort: None,
logprobs: None,
logit_bias: None,
stream_include_usage: None,
parallel_tool_calls: None,
user: None,
metadata: None,
store: None,
service_tier: None,
thinking: None,
request_id: uuid::Uuid::new_v4().to_string(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ToolChoice {
Auto,
Required,
Disabled,
Specific { name: String },
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type")]
pub enum ResponseFormat {
#[serde(rename = "json_object")]
Json,
JsonSchema { schema: Value, name: String },
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum ReasoningEffort {
#[serde(rename = "low")]
Low,
#[serde(rename = "medium")]
Medium,
#[serde(rename = "high")]
High,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum ThinkingType {
Enabled {
#[serde(skip_serializing_if = "Option::is_none")]
budget_tokens: Option<u32>,
},
Disabled,
Adaptive,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ThinkingDisplay {
Summarized,
Omitted,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ThinkingConfig {
#[serde(flatten)]
pub thinking_type: ThinkingType,
#[serde(skip_serializing_if = "Option::is_none")]
pub display: Option<ThinkingDisplay>,
}
#[derive(Debug, Clone, Default)]
pub struct RequestOptions {
pub timeout: Option<Duration>,
pub cancel: Option<crate::cancel::CancellationToken>,
pub metadata: Option<HashMap<String, Value>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StructuredRequest {
pub model: String,
pub messages: Vec<Message>,
pub response_schema: Value,
pub temperature: Option<f32>,
pub max_tokens: Option<usize>,
pub request_id: String,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_completion_request_new() {
let req = CompletionRequest::new("gpt-4", vec![Message::user("Hello")]);
assert_eq!(req.model.as_deref(), Some("gpt-4"));
assert_eq!(req.messages.len(), 1);
assert!(!req.request_id.is_empty());
assert!(req.tools.is_none());
assert!(req.temperature.is_none());
assert!(req.max_tokens.is_none());
assert!(req.top_p.is_none());
assert!(req.stop.is_none());
}
#[test]
fn test_completion_request_new_empty_messages() {
let req = CompletionRequest::new("gpt-4", vec![]);
assert_eq!(req.model.as_deref(), Some("gpt-4"));
assert!(req.messages.is_empty());
}
#[test]
fn test_completion_request_default_model_none() {
let req = CompletionRequest::default();
assert!(req.model.is_none());
}
#[test]
fn test_completion_request_unique_request_id() {
let req1 = CompletionRequest::new("gpt-4", vec![]);
let req2 = CompletionRequest::new("gpt-4", vec![]);
assert_ne!(req1.request_id, req2.request_id);
}
#[test]
fn test_request_options_default() {
let opts = RequestOptions::default();
assert!(opts.timeout.is_none());
assert!(opts.cancel.is_none());
assert!(opts.metadata.is_none());
}
#[test]
fn test_tool_choice_serde() {
let choices = vec![
(ToolChoice::Auto, r#""auto""#),
(ToolChoice::Required, r#""required""#),
(ToolChoice::Disabled, r#""disabled""#),
];
for (choice, expected) in choices {
let json = serde_json::to_string(&choice).unwrap();
assert_eq!(json, expected);
}
}
#[test]
fn test_tool_choice_specific_serde() {
let choice = ToolChoice::Specific { name: "search".into() };
let json = serde_json::to_string(&choice).unwrap();
assert!(json.contains("search"));
}
#[test]
fn test_response_format_json() {
let fmt = ResponseFormat::Json;
let json = serde_json::to_string(&fmt).unwrap();
assert_eq!(json, r#"{"type":"json_object"}"#);
}
#[test]
fn test_response_format_json_schema() {
let schema = serde_json::json!({"type": "object"});
let fmt = ResponseFormat::JsonSchema { schema: schema.clone(), name: "MySchema".into() };
let json = serde_json::to_string(&fmt).unwrap();
assert!(json.contains("MySchema"));
assert!(json.contains("type"));
}
#[test]
fn test_completion_request_serialize() {
let req = CompletionRequest::new("gpt-4", vec![Message::user("Hi")]);
let json = serde_json::to_string(&req).unwrap();
assert!(json.contains("gpt-4"));
assert!(json.contains("Hi"));
assert!(!json.contains("request_id"));
}
#[test]
fn test_reasoning_effort_serde() {
assert_eq!(serde_json::to_string(&ReasoningEffort::Low).unwrap(), r#""low""#);
assert_eq!(serde_json::to_string(&ReasoningEffort::Medium).unwrap(), r#""medium""#);
assert_eq!(serde_json::to_string(&ReasoningEffort::High).unwrap(), r#""high""#);
}
#[test]
fn test_completion_request_new_fields() {
let mut req = CompletionRequest::new("gpt-4", vec![]);
req.seed = Some(42);
req.reasoning_effort = Some(ReasoningEffort::Medium);
req.max_completion_tokens = Some(4000);
req.top_k = Some(50);
req.logprobs = Some(true);
req.logit_bias = Some(HashMap::from([("hello".into(), 0.5)]));
let json = serde_json::to_string(&req).unwrap();
assert!(json.contains("42"));
assert!(json.contains("medium"));
assert!(json.contains("4000"));
assert!(json.contains("50"));
}
#[test]
fn test_service_tier_serde() {
assert_eq!(serde_json::to_string(&ServiceTier::Auto).unwrap(), r#""auto""#);
assert_eq!(serde_json::to_string(&ServiceTier::Default).unwrap(), r#""default""#);
assert_eq!(serde_json::to_string(&ServiceTier::Flex).unwrap(), r#""flex""#);
assert_eq!(serde_json::to_string(&ServiceTier::Scale).unwrap(), r#""scale""#);
assert_eq!(serde_json::to_string(&ServiceTier::Priority).unwrap(), r#""priority""#);
}
#[test]
fn test_thinking_type_enabled_serde() {
let enabled = ThinkingType::Enabled { budget_tokens: Some(4096) };
let json = serde_json::to_string(&enabled).unwrap();
assert!(json.contains(r#""type":"enabled""#));
assert!(json.contains("4096"));
}
#[test]
fn test_thinking_type_disabled_serde() {
let disabled = ThinkingType::Disabled;
let json = serde_json::to_string(&disabled).unwrap();
assert_eq!(json, r#"{"type":"disabled"}"#);
}
#[test]
fn test_thinking_type_adaptive_serde() {
let adaptive = ThinkingType::Adaptive;
let json = serde_json::to_string(&adaptive).unwrap();
assert_eq!(json, r#"{"type":"adaptive"}"#);
}
#[test]
fn test_thinking_display_serde() {
assert_eq!(serde_json::to_string(&ThinkingDisplay::Summarized).unwrap(), r#""summarized""#);
assert_eq!(serde_json::to_string(&ThinkingDisplay::Omitted).unwrap(), r#""omitted""#);
}
#[test]
fn test_thinking_config_serde() {
let config = ThinkingConfig {
thinking_type: ThinkingType::Enabled { budget_tokens: Some(2048) },
display: Some(ThinkingDisplay::Summarized),
};
let json = serde_json::to_string(&config).unwrap();
assert!(json.contains(r#""type":"enabled""#));
assert!(json.contains("2048"));
assert!(json.contains(r#""display":"summarized""#));
}
#[test]
fn test_completion_request_new_fields_serialize() {
let mut req = CompletionRequest::new("gpt-4", vec![Message::user("Hello")]);
req.parallel_tool_calls = Some(true);
req.user = Some("user-123".into());
req.metadata = Some(HashMap::from([("session_id".into(), Value::String("abc".into()))]));
req.store = Some(true);
req.service_tier = Some(ServiceTier::Auto);
req.thinking = Some(ThinkingConfig {
thinking_type: ThinkingType::Enabled { budget_tokens: Some(4096) },
display: None,
});
let json = serde_json::to_string(&req).unwrap();
assert!(json.contains("true")); assert!(json.contains("user-123"));
assert!(json.contains("session_id"));
assert!(json.contains("auto")); assert!(json.contains("enabled")); assert!(json.contains("4096")); }
#[test]
fn test_completion_request_new_fields_absent_when_none() {
let req = CompletionRequest::new("gpt-4", vec![Message::user("Hi")]);
let json = serde_json::to_string(&req).unwrap();
assert!(!json.contains("parallel_tool_calls"));
assert!(!json.contains("\"user\":"));
assert!(!json.contains("metadata"));
assert!(!json.contains("store"));
assert!(!json.contains("service_tier"));
assert!(!json.contains("thinking"));
}
}