use std::collections::HashMap;
use serde::Serialize;
use url::Url;
use crate::{
chat::ServiceTier,
errors::OapiError,
rest::post::{Post, PostNoStream, PostStream},
};
#[derive(Serialize, Debug, Default, Clone)]
pub struct RequestBody {
#[serde(skip_serializing_if = "Option::is_none")]
pub audio: Option<ChatCompletionAudioParam>,
#[serde(skip_serializing_if = "Option::is_none")]
pub frequency_penalty: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub logprobs: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_completion_tokens: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_tokens: Option<u32>,
pub messages: Vec<Message>,
#[serde(skip_serializing_if = "Option::is_none")]
pub metadata: Option<HashMap<String, String>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub modalities: Option<Vec<Modality>>,
pub model: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub n: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub parallel_tool_calls: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub prediction: Option<ChatCompletionPredictionContentParam>,
#[serde(skip_serializing_if = "Option::is_none")]
pub presence_penalty: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub prompt_cache_key: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub reasoning_effort: Option<ReasoningEffort>,
#[serde(skip_serializing_if = "Option::is_none")]
pub response_format: Option<ResponseFormat>,
#[serde(skip_serializing_if = "Option::is_none")]
pub safety_identifier: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub seed: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub service_tier: Option<ServiceTier>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stop: Option<StopKeywords>,
#[serde(skip_serializing_if = "Option::is_none")]
pub store: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stream: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stream_options: Option<StreamOptions>,
#[serde(skip_serializing_if = "Option::is_none")]
pub temperature: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_p: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_choice: Option<ToolChoice>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<Vec<RequestTool>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_logprobs: Option<u32>,
#[cfg(feature = "deepseek")]
#[serde(skip_serializing_if = "Option::is_none")]
pub thinking: Option<DeepSeekThinking>,
#[cfg(feature = "deepseek")]
#[serde(skip_serializing_if = "Option::is_none")]
pub user_id: Option<String>,
#[cfg(feature = "qwen")]
#[serde(skip_serializing_if = "Option::is_none")]
pub enable_thinking: Option<bool>,
#[cfg(feature = "qwen")]
#[serde(skip_serializing_if = "Option::is_none")]
pub thinking_budget: Option<u32>,
#[cfg(feature = "qwen")]
#[serde(skip_serializing_if = "Option::is_none")]
pub top_k: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub verbosity: Option<LowMediumHighEnum>,
#[serde(rename = "web_search_options", skip_serializing_if = "Option::is_none")]
pub web_search_options: Option<WebSearchOptions>,
#[serde(flatten, skip_serializing_if = "Option::is_none")]
pub extra_body_map: Option<serde_json::Map<String, serde_json::Value>>,
}
#[derive(Serialize, Debug, Clone)]
#[serde(tag = "role", rename_all = "lowercase")]
pub enum Message {
System {
content: String,
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
},
User {
content: String,
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
},
Assistant {
content: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
audio: Option<AssistantAudio>,
#[serde(skip_serializing_if = "Option::is_none")]
refusal: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
#[cfg(feature = "deepseek")]
#[serde(skip_serializing_if = "is_false")]
prefix: bool,
#[cfg(feature = "deepseek")]
#[serde(skip_serializing_if = "Option::is_none")]
reasoning_content: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
tool_calls: Option<Vec<AssistantToolCall>>,
},
Tool {
content: String,
tool_call_id: String,
},
Function {
content: String,
name: String,
},
Developer {
content: String,
name: Option<String>,
},
}
#[derive(Debug, Serialize, Clone)]
#[serde(tag = "role", rename_all = "lowercase")]
pub enum AssistantToolCall {
Function {
id: String,
function: ToolCallFunction,
},
Custom {
id: String,
custom: ToolCallCustom,
},
}
#[derive(Debug, Serialize, Clone)]
pub struct ToolCallFunction {
arguments: String,
name: String,
}
#[derive(Debug, Serialize, Clone)]
pub struct ToolCallCustom {
input: String,
name: String,
}
#[derive(Debug, Serialize, Clone)]
pub struct AssistantAudio {
pub id: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub data: Option<String>,
}
#[derive(Debug, Serialize, Clone)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum ResponseFormat {
JsonSchema {
json_schema: JSONSchema,
},
JsonObject,
Text,
}
#[derive(Debug, Serialize, Clone)]
pub struct JSONSchema {
pub name: String,
pub description: String,
pub schema: serde_json::Map<String, serde_json::Value>,
pub strict: Option<bool>,
}
#[derive(Serialize, Debug, Clone)]
#[serde(rename_all = "snake_case")]
pub enum Modality {
Text,
Audio,
}
#[derive(Serialize, Debug, Clone)]
pub struct ChatCompletionAudioParam {
pub format: AudioFormat,
pub voice: Voice,
}
#[derive(Serialize, Debug, Clone)]
#[serde(rename_all = "snake_case")]
pub enum AudioFormat {
Wav,
Aac,
Mp3,
Flac,
Opus,
Pcm16,
}
#[derive(Serialize, Debug, Clone)]
#[serde(untagged)]
pub enum Voice {
BuiltIn(String),
Custom {
id: String,
},
}
#[derive(Serialize, Debug, Clone)]
pub struct ChatCompletionPredictionContentParam {
pub content: ChatCompletionPredictionContentParamContent,
pub type_: ChatCompletionPredictionContentParamType,
}
#[derive(Serialize, Debug, Clone)]
#[serde(untagged)]
pub enum ChatCompletionPredictionContentParamContent {
Text(String),
ChatCompletionContentPartTextParam {
text: String,
#[serde(rename = "type")]
type_: ChatCompletionContentPartTextParamType,
},
}
#[derive(Serialize, Debug, Clone)]
#[serde(rename_all = "snake_case")]
pub enum ChatCompletionContentPartTextParamType {
Text,
}
#[derive(Serialize, Debug, Clone)]
#[serde(rename_all = "snake_case")]
pub enum ChatCompletionPredictionContentParamType {
Content,
}
#[cfg(feature = "deepseek")]
#[inline]
fn is_false(value: &bool) -> bool {
!value
}
#[derive(Serialize, Debug, Clone)]
#[serde(untagged)]
pub enum StopKeywords {
Word(String),
Words(Vec<String>),
}
#[derive(Serialize, Debug, Clone)]
#[serde(rename_all = "snake_case")]
pub enum LowMediumHighEnum {
Low,
Medium,
High,
}
#[derive(Serialize, Debug, Clone)]
pub struct WebSearchOptions {
pub search_context_size: LowMediumHighEnum,
pub user_location: Option<WebSearchOptionsUserLocation>,
}
#[derive(Serialize, Debug, Clone)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum WebSearchOptionsUserLocation {
Approximate(WebSearchOptionsUserLocationApproximate),
}
#[derive(Serialize, Debug, Clone)]
pub struct WebSearchOptionsUserLocationApproximate {
pub city: String,
pub country: String,
pub region: String,
pub timezone: String,
}
#[derive(Serialize, Debug, Clone)]
pub struct StreamOptions {
pub include_usage: bool,
}
#[derive(Serialize, Debug, Clone)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum RequestTool {
Function { function: ToolFunction },
Custom {
custom: ToolCustom,
},
}
#[derive(Serialize, Debug, Clone)]
pub struct ToolFunction {
pub name: String,
pub description: String,
pub parameters: serde_json::Map<String, serde_json::Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub strict: Option<bool>,
}
#[derive(Serialize, Debug, Clone)]
pub struct ToolCustom {
pub name: String,
pub description: String,
pub format: String,
}
#[derive(Serialize, Debug, Clone)]
#[serde(rename_all = "snake_case", tag = "type")]
pub enum ToolCustomFormat {
CustomFormatText,
CustomFormatGrammar {
grammar: ToolCustomFormatGrammarGrammar,
},
}
#[derive(Debug, Serialize, Clone)]
pub struct ToolCustomFormatGrammarGrammar {
pub definition: String,
pub syntax: ToolCustomFormatGrammarGrammarSyntax,
}
#[derive(Debug, Serialize, Clone)]
#[serde(rename_all = "snake_case")]
pub enum ToolCustomFormatGrammarGrammarSyntax {
Lark,
Regex,
}
#[derive(Debug, Serialize, Clone)]
#[serde(rename_all = "snake_case")]
pub enum ToolChoice {
None,
Auto,
Required,
#[serde(untagged)]
Specific(ToolChoiceSpecific),
}
#[derive(Debug, Serialize, Clone)]
#[serde(rename_all = "snake_case", tag = "type")]
pub enum ToolChoiceSpecific {
AllowedTools {
allowed_tools: ToolChoiceAllowedTools,
},
Function { function: ToolChoiceFunction },
Custom { custom: ToolChoiceCustom },
}
#[derive(Debug, Serialize, Clone)]
pub struct ToolChoiceAllowedTools {
pub mode: ToolChoiceAllowedToolsMode,
pub tools: serde_json::Map<String, serde_json::Value>,
}
#[derive(Debug, Serialize, Clone)]
#[serde(rename_all = "lowercase")]
pub enum ToolChoiceAllowedToolsMode {
Auto,
Required,
}
#[derive(Debug, Serialize, Clone)]
pub struct ToolChoiceFunction {
pub name: String,
}
#[derive(Debug, Serialize, Clone)]
pub struct ToolChoiceCustom {
pub name: String,
}
#[cfg(feature = "deepseek")]
#[derive(Debug, Serialize, Clone, PartialEq, Eq)]
pub struct DeepSeekThinking {
#[serde(rename = "type")]
pub type_: DeepSeekThinkingType,
}
#[cfg(feature = "deepseek")]
#[derive(Debug, Serialize, Clone, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum DeepSeekThinkingType {
Enabled,
Disabled,
}
#[derive(Debug, Serialize, Clone, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum ReasoningEffort {
None,
Minimal,
Low,
Medium,
High,
Xhigh,
Max,
}
impl RequestBody {
pub fn is_streaming(&self) -> bool {
self.stream.unwrap_or(false)
}
}
impl Post for RequestBody {
fn is_streaming(&self) -> bool {
RequestBody::is_streaming(self)
}
fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
let mut url = Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
url.path_segments_mut()
.map_err(|_| OapiError::UrlCannotBeBase(base_url.to_string()))?
.push("chat")
.push("completions");
Ok(url.to_string())
}
}
impl PostNoStream for RequestBody {
type Response = super::response::no_streaming::ChatCompletion;
}
impl PostStream for RequestBody {
type Response = super::response::streaming::ChatCompletionChunk;
}
#[cfg(test)]
mod request_test {
use futures_util::StreamExt;
use super::*;
const DEEPSEEK_CHAT_URL: &str = "https://api.deepseek.com";
const DEEPSEEK_MODEL: &str = "deepseek-v4-flash";
fn deepseek_api_key() -> Option<String> {
std::env::var("DEEPSEEK_API_KEY")
.ok()
.map(|key| key.trim().to_string())
.filter(|key| !key.is_empty())
}
#[tokio::test]
async fn test_deepseek_no_stream() {
let Some(api_key) = deepseek_api_key() else {
println!("Skipping: set DEEPSEEK_API_KEY to run this test");
return;
};
let request = RequestBody {
messages: vec![
Message::System {
content: "This is a request of test purpose. Reply briefly".to_string(),
name: None,
},
Message::User {
content: "What's your name?".to_string(),
name: None,
},
],
model: DEEPSEEK_MODEL.to_string(),
stream: Some(false),
..Default::default()
};
let response = request
.get_response_string(&crate::rest::default_client(), DEEPSEEK_CHAT_URL, &api_key)
.await
.unwrap();
println!("{}", response);
assert!(response.to_ascii_lowercase().contains("deepseek"));
}
#[tokio::test]
async fn test_deepseek_stream() {
let Some(api_key) = deepseek_api_key() else {
println!("Skipping: set DEEPSEEK_API_KEY to run this test");
return;
};
let request = RequestBody {
messages: vec![
Message::System {
content: "This is a request of test purpose. Reply briefly".to_string(),
name: None,
},
Message::User {
content: "Who are you?".to_string(),
name: None,
},
],
model: DEEPSEEK_MODEL.to_string(),
stream: Some(true),
..Default::default()
};
let mut response = request
.get_stream_response_string(&crate::rest::default_client(), DEEPSEEK_CHAT_URL, &api_key)
.await
.unwrap();
while let Some(chunk) = response.next().await {
println!("{}", chunk.unwrap());
}
}
#[test]
fn reasoning_effort_serialization() {
let request = RequestBody {
messages: vec![Message::User {
content: "What's your name?".to_string(),
name: None,
}],
model: "gpt-5".to_string(),
reasoning_effort: Some(ReasoningEffort::Xhigh),
..Default::default()
};
let json = serde_json::to_string(&request).unwrap();
assert!(
json.contains(r#""reasoning_effort":"xhigh""#),
"json: {json}"
);
}
#[cfg(feature = "deepseek")]
#[test]
fn deepseek_assistant_prefix_serialization() {
let request = RequestBody {
messages: vec![
Message::User {
content: "Please write quick sort code".to_string(),
name: None,
},
Message::Assistant {
content: Some("```python\n".to_string()),
audio: None,
refusal: None,
name: None,
prefix: true,
reasoning_content: None,
tool_calls: None,
},
],
model: DEEPSEEK_MODEL.to_string(),
..Default::default()
};
let json = serde_json::to_string(&request).unwrap();
assert!(json.contains(r#""prefix":true"#), "json: {json}");
}
#[cfg(feature = "deepseek")]
#[test]
fn deepseek_thinking_params_serialization() {
let request = RequestBody {
messages: vec![Message::User {
content: "What's your name?".to_string(),
name: None,
}],
model: DEEPSEEK_MODEL.to_string(),
thinking: Some(DeepSeekThinking {
type_: DeepSeekThinkingType::Disabled,
}),
user_id: Some("user-123".to_string()),
..Default::default()
};
let json = serde_json::to_string(&request).unwrap();
assert!(
json.contains(r#""thinking":{"type":"disabled"}"#),
"json: {json}"
);
assert!(json.contains(r#""user_id":"user-123""#), "json: {json}");
}
#[cfg(feature = "qwen")]
#[test]
fn qwen_params_serialization() {
let request = RequestBody {
messages: vec![Message::User {
content: "What's your name?".to_string(),
name: None,
}],
model: "qwen-plus".to_string(),
enable_thinking: Some(false),
thinking_budget: Some(1024),
top_k: Some(20),
..Default::default()
};
let json = serde_json::to_string(&request).unwrap();
assert!(json.contains(r#""enable_thinking":false"#), "json: {json}");
assert!(json.contains(r#""thinking_budget":1024"#), "json: {json}");
assert!(json.contains(r#""top_k":20"#), "json: {json}");
}
}