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 logit_bias: Option<HashMap<u32, i32>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub moderation: Option<ChatModerationParam>,
#[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 prompt_cache_options: Option<PromptCacheOptions>,
#[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: MessageContent,
#[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(untagged)]
pub enum MessageContent {
Text(String),
Parts(Vec<ContentPart>),
}
impl From<&str> for MessageContent {
fn from(value: &str) -> Self {
Self::Text(value.to_string())
}
}
impl From<String> for MessageContent {
fn from(value: String) -> Self {
Self::Text(value)
}
}
impl From<Vec<ContentPart>> for MessageContent {
fn from(value: Vec<ContentPart>) -> Self {
Self::Parts(value)
}
}
impl Default for MessageContent {
fn default() -> Self {
Self::Text(String::new())
}
}
#[derive(Debug, Serialize, Clone)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum ContentPart {
Text {
text: String,
#[serde(skip_serializing_if = "Option::is_none")]
prompt_cache_breakpoint: Option<PromptCacheBreakpoint>,
},
ImageUrl {
image_url: ContentPartImageUrl,
#[serde(skip_serializing_if = "Option::is_none")]
prompt_cache_breakpoint: Option<PromptCacheBreakpoint>,
},
InputAudio {
input_audio: ContentPartInputAudio,
#[serde(skip_serializing_if = "Option::is_none")]
prompt_cache_breakpoint: Option<PromptCacheBreakpoint>,
},
File {
file: ContentPartFile,
#[serde(skip_serializing_if = "Option::is_none")]
prompt_cache_breakpoint: Option<PromptCacheBreakpoint>,
},
}
#[derive(Debug, Serialize, Clone)]
pub struct PromptCacheBreakpoint {
pub mode: PromptCacheBreakpointMode,
}
#[derive(Debug, Serialize, Clone)]
#[serde(rename_all = "lowercase")]
pub enum PromptCacheBreakpointMode {
Explicit,
}
#[derive(Debug, Serialize, Clone)]
pub struct ContentPartImageUrl {
pub url: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub detail: Option<ImageDetail>,
}
#[derive(Debug, Serialize, Clone, Copy)]
#[serde(rename_all = "lowercase")]
pub enum ImageDetail {
Auto,
Low,
High,
}
#[derive(Debug, Serialize, Clone)]
pub struct ContentPartInputAudio {
pub data: String,
pub format: InputAudioFormat,
}
#[derive(Debug, Serialize, Clone, Copy)]
#[serde(rename_all = "lowercase")]
pub enum InputAudioFormat {
Wav,
Mp3,
}
#[derive(Debug, Serialize, Clone, Default)]
pub struct ContentPartFile {
#[serde(skip_serializing_if = "Option::is_none")]
pub file_data: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub file_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub filename: Option<String>,
}
#[derive(Debug, Serialize, Clone)]
pub struct ChatModerationParam {
pub model: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub policy: Option<ModerationPolicyParam>,
}
#[derive(Debug, Serialize, Clone, Default)]
pub struct ModerationPolicyParam {
#[serde(skip_serializing_if = "Option::is_none")]
pub input: Option<ModerationPolicySideParam>,
#[serde(skip_serializing_if = "Option::is_none")]
pub output: Option<ModerationPolicySideParam>,
}
#[derive(Debug, Serialize, Clone)]
pub struct ModerationPolicySideParam {
pub mode: ModerationPolicyMode,
}
#[derive(Debug, Serialize, Clone, Copy)]
#[serde(rename_all = "lowercase")]
pub enum ModerationPolicyMode {
Score,
Block,
}
#[derive(Debug, Serialize, Clone, Default)]
pub struct PromptCacheOptions {
#[serde(skip_serializing_if = "Option::is_none")]
pub mode: Option<PromptCacheMode>,
#[serde(skip_serializing_if = "Option::is_none")]
pub ttl: Option<PromptCacheTtl>,
}
#[derive(Debug, Serialize, Clone, Copy)]
#[serde(rename_all = "lowercase")]
pub enum PromptCacheMode {
Implicit,
Explicit,
}
#[derive(Debug, Serialize, Clone, Copy)]
pub enum PromptCacheTtl {
#[serde(rename = "30m")]
ThirtyMinutes,
}
#[derive(Debug, Serialize, Clone)]
#[serde(tag = "type", 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,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub schema: Option<serde_json::Map<String, serde_json::Value>>,
#[serde(skip_serializing_if = "Option::is_none")]
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,
#[serde(rename = "type")]
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, Default)]
pub struct WebSearchOptions {
#[serde(skip_serializing_if = "Option::is_none")]
pub search_context_size: Option<LowMediumHighEnum>,
#[serde(skip_serializing_if = "Option::is_none")]
pub user_location: Option<WebSearchOptionsUserLocation>,
}
#[derive(Serialize, Debug, Clone)]
#[serde(tag = "type", rename_all = "snake_case")]
pub enum WebSearchOptionsUserLocation {
Approximate {
approximate: WebSearchOptionsUserLocationApproximate,
},
}
#[derive(Serialize, Debug, Clone, Default)]
pub struct WebSearchOptionsUserLocationApproximate {
#[serde(skip_serializing_if = "Option::is_none")]
pub city: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub country: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub region: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub timezone: Option<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,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub parameters: Option<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,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub format: Option<ToolCustomFormat>,
}
#[derive(Serialize, Debug, Clone)]
#[serde(rename_all = "snake_case", tag = "type")]
pub enum ToolCustomFormat {
Text,
Grammar {
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: Vec<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?".into(),
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?".into(),
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 assistant_tool_call_serialization() {
let function_call = AssistantToolCall::Function {
id: "call_abc".to_string(),
function: ToolCallFunction {
arguments: "{\"city\":\"paris\"}".to_string(),
name: "get_weather".to_string(),
},
};
let json = serde_json::to_string(&function_call).unwrap();
assert!(json.contains(r#""type":"function""#), "json: {json}");
assert!(!json.contains(r#""role""#), "json: {json}");
let custom_call = AssistantToolCall::Custom {
id: "call_def".to_string(),
custom: ToolCallCustom {
input: "2+2".to_string(),
name: "calculator".to_string(),
},
};
let json = serde_json::to_string(&custom_call).unwrap();
assert!(json.contains(r#""type":"custom""#), "json: {json}");
assert!(!json.contains(r#""role""#), "json: {json}");
}
#[test]
fn prediction_type_serialization() {
let prediction = ChatCompletionPredictionContentParam {
content: ChatCompletionPredictionContentParamContent::Text(
"The capital of France is Paris.".to_string(),
),
type_: ChatCompletionPredictionContentParamType::Content,
};
let json = serde_json::to_string(&prediction).unwrap();
assert!(json.contains(r#""type":"content""#), "json: {json}");
assert!(!json.contains("type_"), "json: {json}");
}
#[test]
fn allowed_tools_choice_serialization() {
let mut weather = serde_json::Map::new();
weather.insert("type".to_string(), serde_json::json!("function"));
weather.insert(
"function".to_string(),
serde_json::json!({ "name": "get_weather" }),
);
let choice = ToolChoiceSpecific::AllowedTools {
allowed_tools: ToolChoiceAllowedTools {
mode: ToolChoiceAllowedToolsMode::Required,
tools: vec![weather],
},
};
let json = serde_json::to_string(&choice).unwrap();
assert!(json.contains(r#""type":"allowed_tools""#), "json: {json}");
assert!(json.contains(r#""mode":"required""#), "json: {json}");
assert!(json.contains(r#""tools":[{"#), "json: {json}");
}
#[test]
fn web_search_options_serialization() {
let options = WebSearchOptions {
search_context_size: None,
user_location: Some(WebSearchOptionsUserLocation::Approximate {
approximate: WebSearchOptionsUserLocationApproximate {
city: Some("San Francisco".to_string()),
country: None,
region: None,
timezone: None,
},
}),
};
let json = serde_json::to_string(&options).unwrap();
assert!(!json.contains("search_context_size"), "json: {json}");
assert!(json.contains(r#""type":"approximate""#), "json: {json}");
assert!(
json.contains(r#""approximate":{"city":"San Francisco"}"#),
"json: {json}"
);
}
#[test]
fn json_schema_optional_fields_serialization() {
let schema = JSONSchema {
name: "Answer".to_string(),
description: None,
schema: None,
strict: None,
};
let json = serde_json::to_string(&schema).unwrap();
assert_eq!(json, r#"{"name":"Answer"}"#);
let function = ToolFunction {
name: "get_weather".to_string(),
description: None,
parameters: None,
strict: None,
};
let json = serde_json::to_string(&function).unwrap();
assert_eq!(json, r#"{"name":"get_weather"}"#);
}
#[test]
fn user_text_content_serialization() {
let request = RequestBody {
messages: vec![Message::User {
content: "Hi".into(),
name: None,
}],
model: "gpt-4o".to_string(),
..Default::default()
};
let json = serde_json::to_string(&request).unwrap();
assert!(json.contains(r#""content":"Hi""#), "json: {json}");
}
#[test]
fn multimodal_content_serialization() {
let request = RequestBody {
messages: vec![Message::User {
content: MessageContent::Parts(vec![
ContentPart::ImageUrl {
image_url: ContentPartImageUrl {
url: "https://example.com/cat.png".to_string(),
detail: Some(ImageDetail::High),
},
prompt_cache_breakpoint: None,
},
ContentPart::Text {
text: "What's in this image?".to_string(),
prompt_cache_breakpoint: Some(PromptCacheBreakpoint {
mode: PromptCacheBreakpointMode::Explicit,
}),
},
]),
name: None,
}],
model: "gpt-4o".to_string(),
..Default::default()
};
let json = serde_json::to_string(&request).unwrap();
assert!(json.contains(r#""type":"image_url""#), "json: {json}");
assert!(
json.contains(r#""url":"https://example.com/cat.png""#),
"json: {json}"
);
assert!(json.contains(r#""detail":"high""#), "json: {json}");
assert!(json.contains(r#""type":"text""#), "json: {json}");
assert!(
json.contains(r#""prompt_cache_breakpoint":{"mode":"explicit"}"#),
"json: {json}"
);
}
#[test]
fn audio_and_file_content_serialization() {
let content = MessageContent::Parts(vec![
ContentPart::InputAudio {
input_audio: ContentPartInputAudio {
data: "aGVsbG8=".to_string(),
format: InputAudioFormat::Wav,
},
prompt_cache_breakpoint: None,
},
ContentPart::File {
file: ContentPartFile {
file_id: Some("file-abc".to_string()),
..Default::default()
},
prompt_cache_breakpoint: None,
},
]);
let json = serde_json::to_string(&content).unwrap();
assert!(json.contains(r#""type":"input_audio""#), "json: {json}");
assert!(json.contains(r#""data":"aGVsbG8=""#), "json: {json}");
assert!(json.contains(r#""format":"wav""#), "json: {json}");
assert!(json.contains(r#""type":"file""#), "json: {json}");
assert!(
json.contains(r#""file":{"file_id":"file-abc"}"#),
"json: {json}"
);
assert!(!json.contains("file_data"), "json: {json}");
}
#[test]
fn new_params_serialization() {
let mut logit_bias = HashMap::new();
logit_bias.insert(40u32, -100i32);
let request = RequestBody {
messages: vec![Message::User {
content: "Hi".into(),
name: None,
}],
model: "gpt-5".to_string(),
logit_bias: Some(logit_bias),
moderation: Some(ChatModerationParam {
model: "omni-moderation-latest".to_string(),
policy: Some(ModerationPolicyParam {
input: Some(ModerationPolicySideParam {
mode: ModerationPolicyMode::Block,
}),
output: None,
}),
}),
prompt_cache_options: Some(PromptCacheOptions {
mode: Some(PromptCacheMode::Explicit),
ttl: Some(PromptCacheTtl::ThirtyMinutes),
}),
..Default::default()
};
let json = serde_json::to_string(&request).unwrap();
assert!(json.contains(r#""logit_bias":{"40":-100}"#), "json: {json}");
assert!(
json.contains(
r#""moderation":{"model":"omni-moderation-latest","policy":{"input":{"mode":"block"}}}"#
),
"json: {json}"
);
assert!(
json.contains(r#""prompt_cache_options":{"mode":"explicit","ttl":"30m"}"#),
"json: {json}"
);
}
#[test]
fn reasoning_effort_serialization() {
let request = RequestBody {
messages: vec![Message::User {
content: "What's your name?".into(),
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".into(),
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?".into(),
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?".into(),
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}");
}
const QWEN_CHAT_URL: &str = "https://dashscope.aliyuncs.com/compatible-mode/v1";
const QWEN_MULTIMODAL_MODEL: &str = "qwen3.8-flash";
fn qwen_api_key() -> Option<String> {
std::env::var("QWEN_API_KEY")
.ok()
.map(|key| key.trim().to_string())
.filter(|key| !key.is_empty())
}
#[tokio::test]
async fn test_qwen_image_input() -> Result<(), anyhow::Error> {
let Some(api_key) = qwen_api_key() else {
println!("Skipping: set QWEN_API_KEY to run this test");
return Ok(());
};
let request = RequestBody {
messages: vec![
Message::System {
content: "This is a request of test purpose. Reply briefly".to_string(),
name: None,
},
Message::User {
content: MessageContent::Parts(vec![
ContentPart::ImageUrl {
image_url: ContentPartImageUrl {
url: "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/xzsgiz/football1.jpg"
.to_string(),
detail: None,
},
prompt_cache_breakpoint: None,
},
ContentPart::Text {
text: "What is shown in this image? Answer with one short sentence."
.to_string(),
prompt_cache_breakpoint: None,
},
]),
name: None,
},
],
model: QWEN_MULTIMODAL_MODEL.to_string(),
..Default::default()
};
let response = request
.get_response(&crate::rest::default_client(), QWEN_CHAT_URL, &api_key)
.await?;
let content = response.choices[0]
.message
.content
.clone()
.unwrap_or_default();
println!("image response: {content}");
assert!(
!content.trim().is_empty(),
"empty content for a valid image request"
);
Ok(())
}
#[tokio::test]
async fn test_qwen_audio_input() -> Result<(), anyhow::Error> {
let Some(api_key) = qwen_api_key() else {
println!("Skipping: set QWEN_API_KEY to run this test");
return Ok(());
};
let request = RequestBody {
messages: vec![Message::User {
content: MessageContent::Parts(vec![
ContentPart::InputAudio {
input_audio: ContentPartInputAudio {
data: "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20250211/tixcef/cherry.wav"
.to_string(),
format: InputAudioFormat::Wav,
},
prompt_cache_breakpoint: None,
},
ContentPart::Text {
text: "What does the speaker say in this audio? Reply briefly."
.to_string(),
prompt_cache_breakpoint: None,
},
]),
name: None,
}],
model: "qwen-omni-turbo".to_string(),
stream: Some(true),
modalities: Some(vec![Modality::Text]),
..Default::default()
};
let mut stream = request
.get_stream_response(&crate::rest::default_client(), QWEN_CHAT_URL, &api_key)
.await?;
let mut message = String::new();
while let Some(chunk) = stream.next().await {
let chunk = chunk?;
if let Some(choice) = chunk.choices.first()
&& let Some(content) = choice.delta.content.as_deref()
{
message.push_str(content);
}
}
println!("audio response: {message}");
assert!(
!message.trim().is_empty(),
"empty content for a valid audio request"
);
Ok(())
}
#[tokio::test]
async fn test_qwen_text_input() -> Result<(), anyhow::Error> {
let Some(api_key) = qwen_api_key() else {
println!("Skipping: set QWEN_API_KEY to run this test");
return Ok(());
};
let request = RequestBody {
messages: vec![Message::User {
content: "Reply with exactly one word.".into(),
name: None,
}],
model: QWEN_MULTIMODAL_MODEL.to_string(),
..Default::default()
};
let response = request
.get_response(&crate::rest::default_client(), QWEN_CHAT_URL, &api_key)
.await?;
let content = response.choices[0]
.message
.content
.clone()
.unwrap_or_default();
println!("text response: {content}");
assert!(!content.trim().is_empty(), "empty content for text input");
Ok(())
}
#[tokio::test]
async fn test_qwen_multimodal_stream() -> Result<(), anyhow::Error> {
let Some(api_key) = qwen_api_key() else {
println!("Skipping: set QWEN_API_KEY to run this test");
return Ok(());
};
let request = RequestBody {
messages: vec![Message::User {
content: MessageContent::Parts(vec![
ContentPart::ImageUrl {
image_url: ContentPartImageUrl {
url: "https://help-static-aliyun-doc.aliyuncs.com/file-manage-files/zh-CN/20241108/xzsgiz/football1.jpg"
.to_string(),
detail: None,
},
prompt_cache_breakpoint: None,
},
ContentPart::Text {
text: "What is shown in this image? Answer with one short sentence."
.to_string(),
prompt_cache_breakpoint: None,
},
]),
name: None,
}],
model: QWEN_MULTIMODAL_MODEL.to_string(),
stream: Some(true),
..Default::default()
};
let mut stream = request
.get_stream_response(&crate::rest::default_client(), QWEN_CHAT_URL, &api_key)
.await?;
let mut message = String::new();
while let Some(chunk) = stream.next().await {
let chunk = chunk?;
if let Some(choice) = chunk.choices.first()
&& let Some(content) = choice.delta.content.as_deref()
{
message.push_str(content);
}
}
println!("streamed message: {message}");
assert!(
!message.trim().is_empty(),
"empty streamed content for a valid image request"
);
Ok(())
}
}