#![cfg(feature = "zai")]
use serde::{Deserialize, Serialize};
#[derive(Serialize, Deserialize, Debug, Clone, Default)]
pub struct PlatformParams {
#[serde(skip_serializing_if = "Option::is_none")]
pub watermark_enabled: Option<bool>,
}
#[derive(Serialize, Deserialize, Debug, Clone)]
pub struct RetrievalTool {
pub knowledge_id: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub prompt_template: Option<String>,
}
#[derive(Serialize, Deserialize, Debug, Clone, Default)]
pub struct WebSearchTool {
#[serde(skip_serializing_if = "Option::is_none")]
pub search_engine: Option<SearchEngine>,
#[serde(skip_serializing_if = "Option::is_none")]
pub enable: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub search_query: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub count: Option<u32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub search_domain_filter: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub search_recency_filter: Option<SearchRecencyFilter>,
#[serde(skip_serializing_if = "Option::is_none")]
pub content_size: Option<ContentSize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub result_sequence: Option<ResultSequence>,
#[serde(skip_serializing_if = "Option::is_none")]
pub search_result: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub require_search: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub search_prompt: Option<String>,
}
crate::wire_string_enum! {
pub enum SearchEngine {
SearchProJina => "search_pro_jina",
}
}
crate::wire_string_enum! {
pub enum SearchRecencyFilter {
OneDay => "oneDay",
OneWeek => "oneWeek",
OneMonth => "oneMonth",
OneYear => "oneYear",
NoLimit => "noLimit",
}
}
crate::wire_string_enum! {
pub enum ContentSize {
Medium => "medium",
High => "high",
}
}
crate::wire_string_enum! {
pub enum ResultSequence {
Before => "before",
After => "after",
}
}
#[derive(Serialize, Deserialize, Debug, Clone, Default, PartialEq)]
pub struct WebSearchResult {
#[serde(skip_serializing_if = "Option::is_none")]
pub title: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub link: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub media: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub icon: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub refer: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub publish_date: Option<String>,
}
#[cfg(test)]
mod tests {
use super::*;
use crate::chat::create::request::{Message, RequestBody, RequestTool};
#[test]
fn platform_params_flatten_to_top_level() {
let request = RequestBody {
messages: vec![Message::user("Hello")],
model: "glm-4.6".to_string(),
do_sample: Some(false),
tool_stream: Some(true),
zai_platform: Some(PlatformParams {
watermark_enabled: Some(false),
}),
..Default::default()
};
let json = serde_json::to_value(&request).unwrap();
assert_eq!(json["do_sample"], serde_json::json!(false));
assert_eq!(json["tool_stream"], serde_json::json!(true));
assert_eq!(json["watermark_enabled"], serde_json::json!(false));
let parsed: RequestBody = serde_json::from_str(
r#"{
"model": "glm-4.6",
"messages": [{"role": "user", "content": "Hello"}],
"do_sample": false,
"watermark_enabled": false,
"some_future_zai_field": 42
}"#,
)
.unwrap();
assert_eq!(parsed.do_sample, Some(false));
assert_eq!(
parsed
.zai_platform
.as_ref()
.expect("zai_platform")
.watermark_enabled,
Some(false)
);
let extra = parsed.extra_body_map.as_ref().expect("extra_body_map");
assert_eq!(extra.len(), 1, "extra_body_map: {extra:?}");
assert_eq!(extra["some_future_zai_field"], 42);
}
#[test]
fn tools_serialize_with_type_tag() {
let request = RequestBody {
messages: vec![Message::user("Hello")],
model: "glm-4.6".to_string(),
tools: Some(vec![
RequestTool::Retrieval {
retrieval: RetrievalTool {
knowledge_id: "kb-123".to_string(),
prompt_template: Some("使用以下资料回答".to_string()),
},
},
RequestTool::WebSearch {
web_search: WebSearchTool {
search_engine: Some(SearchEngine::SearchProJina),
enable: Some(true),
count: Some(5),
search_recency_filter: Some(SearchRecencyFilter::OneWeek),
content_size: Some(ContentSize::High),
result_sequence: Some(ResultSequence::Before),
..Default::default()
},
},
]),
..Default::default()
};
let json = serde_json::to_value(&request).unwrap();
assert_eq!(json["tools"][0]["type"], serde_json::json!("retrieval"));
assert_eq!(
json["tools"][0]["retrieval"]["knowledge_id"],
serde_json::json!("kb-123")
);
assert_eq!(json["tools"][1]["type"], serde_json::json!("web_search"));
assert_eq!(
json["tools"][1]["web_search"]["search_engine"],
serde_json::json!("search_pro_jina")
);
assert_eq!(
json["tools"][1]["web_search"]["search_recency_filter"],
serde_json::json!("oneWeek")
);
assert!(json["tools"][1]["web_search"].get("search_query").is_none());
}
#[test]
fn web_search_result_parses_and_engine_tolerates_unknown() {
let result: WebSearchResult = serde_json::from_value(serde_json::json!({
"title": "Rust 1.88 released",
"link": "https://example.com/rust",
"refer": "1",
"publish_date": "2026-09-01"
}))
.unwrap();
assert_eq!(result.title.as_deref(), Some("Rust 1.88 released"));
assert_eq!(result.refer.as_deref(), Some("1"));
assert!(result.content.is_none());
let engine: SearchEngine = serde_json::from_str(r#""search_pro_jina""#).unwrap();
assert_eq!(engine, SearchEngine::SearchProJina);
let unknown: SearchEngine = serde_json::from_str(r#""some_future_engine""#).unwrap();
assert_eq!(
unknown,
SearchEngine::Unknown("some_future_engine".to_string())
);
assert_eq!(unknown.as_str(), "some_future_engine");
}
}