Skip to main content

vtcode_llm/providers/
zai.rs

1use reqwest::RequestBuilder;
2use serde_json::{Map, Value};
3use vtcode_config::constants::{env_vars, models, urls};
4
5use super::openai_compat::{OpenAiCompatCore, OpenAiCompatSpec, SystemPromptPlacement, impl_openai_compat_provider};
6use crate::provider::{LLMError, LLMRequest};
7
8pub struct ZaiSpec;
9
10fn resolve_zai_base_url(base_url: Option<String>) -> String {
11    if let Some(url) = base_url {
12        let trimmed = url.trim();
13        if !trimmed.is_empty() {
14            return trimmed.to_string();
15        }
16    }
17
18    if let Ok(value) = std::env::var(env_vars::ZAI_BASE_URL) {
19        let trimmed = value.trim();
20        if !trimmed.is_empty() {
21            return trimmed.to_string();
22        }
23    }
24
25    if let Ok(legacy) = std::env::var(env_vars::Z_AI_BASE_URL) {
26        let trimmed = legacy.trim();
27        if !trimmed.is_empty() {
28            return trimmed.to_string();
29        }
30    }
31
32    urls::ZAI_API_BASE.to_string()
33}
34
35impl OpenAiCompatSpec for ZaiSpec {
36    const NAME: &'static str = "Z.AI";
37    const KEY: &'static str = "zai";
38    const API_KEY_ENV: &'static str = "ZAI_API_KEY";
39    const DEFAULT_MODEL: &'static str = models::zai::DEFAULT_MODEL;
40    const DEFAULT_BASE_URL: &'static str = urls::ZAI_API_BASE;
41    const BASE_URL_ENV: Option<&'static str> = Some(env_vars::ZAI_BASE_URL);
42    const LISTED_MODELS: &'static [&'static str] = models::zai::SUPPORTED_MODELS;
43    const VALIDATION_ALLOWLIST: Option<&'static [&'static str]> = Some(models::zai::SUPPORTED_MODELS);
44
45    const SYSTEM_PROMPT: SystemPromptPlacement = SystemPromptPlacement::Omitted;
46    const SUPPRESS_SAMPLING_WHEN_REASONING: bool = false;
47    const DELTA_ORDER: super::shared::OpenAiDeltaOrder = super::shared::OpenAiDeltaOrder::ContentFirst;
48
49    fn resolve_base_url(_api_key: &str, base_url: Option<String>) -> String {
50        resolve_zai_base_url(base_url)
51    }
52
53    fn insert_tool_choice(_core: &OpenAiCompatCore<Self>, request: &LLMRequest, payload: &mut Map<String, Value>) {
54        if let Some(choice) = &request.tool_choice {
55            // Z.AI only supports "auto"; any other requested mode is coerced.
56            let tool_choice_value = match choice {
57                crate::provider::ToolChoice::Auto => choice.to_provider_format(Self::KEY),
58                _ => Value::String("auto".to_string()),
59            };
60            payload.insert("tool_choice".to_string(), tool_choice_value);
61        } else if request.tools.as_ref().is_some_and(|tools| !tools.is_empty()) {
62            payload.insert("tool_choice".to_string(), Value::String("auto".to_string()));
63        }
64    }
65
66    fn insert_reasoning(
67        core: &OpenAiCompatCore<Self>,
68        request: &LLMRequest,
69        payload: &mut Map<String, Value>,
70    ) -> Result<(), LLMError> {
71        let has_preserved_reasoning = request.messages.iter().any(|message| {
72            message.role == crate::provider::MessageRole::Assistant
73                && message.reasoning.as_ref().is_some_and(|reasoning| !reasoning.is_empty())
74        });
75
76        if let Some(effort) = request.reasoning_effort {
77            if effort == vtcode_config::types::ReasoningEffortLevel::None {
78                payload.insert("thinking".to_owned(), serde_json::json!({"type": "disabled"}));
79                return Ok(());
80            }
81
82            use crate::rig_adapter::RigProviderCapabilities;
83            use vtcode_config::models::Provider;
84            let supported = crate::provider::catalog_or_explicit_reasoning_efforts(
85                Self::KEY,
86                &request.model,
87                core.model_behavior
88                    .as_ref()
89                    .and_then(|behavior| behavior.model_supports_reasoning_effort)
90                    .unwrap_or(false),
91            );
92            let reasoning_params = RigProviderCapabilities::new(Provider::ZAI, &request.model)
93                .reasoning_parameters_for_supported_efforts(effort, supported)
94                ?
95                .ok_or_else(|| LLMError::InvalidRequest {
96                    message: format!(
97                        "Reasoning effort `{effort}` is unsupported for Z.AI model `{}`; choose one of low, high, or max",
98                        request.model
99                    ),
100                    metadata: None,
101                })?;
102            if let Some(params_obj) = reasoning_params.as_object() {
103                for (k, v) in params_obj {
104                    payload.insert(k.clone(), v.clone());
105                }
106            }
107        }
108
109        if has_preserved_reasoning {
110            if let Some(thinking) = payload.get_mut("thinking").and_then(Value::as_object_mut) {
111                thinking.insert("clear_thinking".to_owned(), Value::Bool(false));
112            } else {
113                payload.insert(
114                    "thinking".to_owned(),
115                    serde_json::json!({
116                        "type": "enabled",
117                        "clear_thinking": false
118                    }),
119                );
120            }
121        }
122
123        Ok(())
124    }
125
126    fn finish_payload(
127        _core: &OpenAiCompatCore<Self>,
128        request: &LLMRequest,
129        payload: &mut Map<String, Value>,
130    ) -> Result<(), LLMError> {
131        if let Some(do_sample) = request.do_sample {
132            payload.insert("do_sample".to_owned(), Value::Bool(do_sample));
133        }
134
135        if request.stream && request.tools.as_ref().is_some_and(|tools| !tools.is_empty()) {
136            payload.insert("tool_stream".to_string(), Value::Bool(true));
137        }
138
139        if request.output_format.is_some() {
140            payload.insert("response_format".to_owned(), serde_json::json!({ "type": "json_object" }));
141        }
142
143        Ok(())
144    }
145
146    fn apply_auth(core: &OpenAiCompatCore<Self>, builder: RequestBuilder) -> RequestBuilder {
147        builder.bearer_auth(&core.api_key).header("Accept-Language", "en-US,en")
148    }
149}
150
151impl_openai_compat_provider!(ZAIProvider, ZaiSpec, {
152    fn supports_reasoning(&self, model: &str) -> bool {
153        // Codex-inspired robustness: Setting model_supports_reasoning to false
154        // does NOT disable it for known reasoning models.
155        model.contains("glm")
156            || self
157                .core
158                .model_behavior
159                .as_ref()
160                .and_then(|b| b.model_supports_reasoning)
161                .unwrap_or(false)
162    }
163
164    fn supports_reasoning_effort(&self, model: &str) -> bool {
165        // Same robustness logic for reasoning effort
166        model.contains("glm")
167            || self
168                .core
169                .model_behavior
170                .as_ref()
171                .and_then(|b| b.model_supports_reasoning_effort)
172                .unwrap_or(false)
173    }
174});
175
176#[cfg(test)]
177mod tests {
178    use super::{ZAIProvider, resolve_zai_base_url};
179    use crate::provider::{LLMRequest, Message, ToolChoice, ToolDefinition};
180    use std::sync::Arc;
181    use vtcode_config::constants::models;
182    use vtcode_config::types::ReasoningEffortLevel;
183
184    #[test]
185    fn payload_includes_top_p() {
186        let provider = ZAIProvider::new("test-key".to_string());
187        let request = LLMRequest {
188            model: models::zai::GLM_5_3.to_string(),
189            messages: vec![Message::user("hello".to_string())].into(),
190            top_p: Some(0.95),
191            ..Default::default()
192        };
193
194        let payload = provider.core.convert_request(&request).expect("payload should be valid");
195        let top_p = payload.get("top_p").and_then(|v| v.as_f64()).expect("top_p should be present");
196        assert!((top_p - 0.95).abs() < 1e-6);
197    }
198
199    #[test]
200    fn payload_enables_tool_stream_when_streaming_with_tools() {
201        let provider = ZAIProvider::new("test-key".to_string());
202        let request = LLMRequest {
203            model: models::zai::GLM_5_3.to_string(),
204            messages: vec![Message::user("hello".to_string())].into(),
205            stream: true,
206            tools: Some(Arc::new(vec![ToolDefinition::function(
207                "get_weather".to_string(),
208                "Get weather".to_string(),
209                serde_json::json!({
210                    "type": "object",
211                    "properties": {
212                        "location": {"type": "string"}
213                    },
214                    "required": ["location"]
215                }),
216            )])),
217            ..Default::default()
218        };
219
220        let payload = provider.core.convert_request(&request).expect("payload should be valid");
221        assert_eq!(payload.get("stream").and_then(|v| v.as_bool()), Some(true));
222        assert_eq!(payload.get("tool_stream").and_then(|v| v.as_bool()), Some(true));
223    }
224
225    #[test]
226    fn payload_streaming_without_tools_does_not_set_tool_stream() {
227        let provider = ZAIProvider::new("test-key".to_string());
228        let request = LLMRequest {
229            model: models::zai::GLM_5_3.to_string(),
230            messages: vec![Message::user("hello".to_string())].into(),
231            stream: true,
232            ..Default::default()
233        };
234
235        let payload = provider.core.convert_request(&request).expect("payload should be valid");
236        assert_eq!(payload.get("stream").and_then(|v| v.as_bool()), Some(true));
237        assert!(payload.get("tool_stream").is_none());
238    }
239
240    #[test]
241    fn zai_base_url_uses_explicit_override() {
242        let resolved = resolve_zai_base_url(Some("https://api.z.ai/api/coding/paas/v4".to_string()));
243        assert_eq!(resolved, "https://api.z.ai/api/coding/paas/v4");
244    }
245
246    #[test]
247    fn payload_includes_do_sample() {
248        let provider = ZAIProvider::new("test-key".to_string());
249        let request = LLMRequest {
250            model: models::zai::GLM_5_3.to_string(),
251            messages: vec![Message::user("hello".to_string())].into(),
252            do_sample: Some(false),
253            ..Default::default()
254        };
255
256        let payload = provider.core.convert_request(&request).expect("payload should be valid");
257        assert_eq!(payload.get("do_sample").and_then(|v| v.as_bool()), Some(false));
258    }
259
260    #[test]
261    fn payload_disables_thinking_for_none_effort() {
262        let provider = ZAIProvider::new("test-key".to_string());
263        let request = LLMRequest {
264            model: models::zai::GLM_5_3.to_string(),
265            messages: vec![Message::user("hello".to_string())].into(),
266            reasoning_effort: Some(ReasoningEffortLevel::None),
267            ..Default::default()
268        };
269
270        let payload = provider.core.convert_request(&request).expect("payload should be valid");
271        assert_eq!(payload.get("thinking").and_then(|v| v.get("type")).and_then(|v| v.as_str()), Some("disabled"));
272    }
273
274    #[test]
275    fn payload_enables_thinking_for_low_effort() {
276        let provider = ZAIProvider::new("test-key".to_string());
277        let request = LLMRequest {
278            model: models::zai::GLM_5_3.to_string(),
279            messages: vec![Message::user("hello".to_string())].into(),
280            reasoning_effort: Some(ReasoningEffortLevel::Low),
281            ..Default::default()
282        };
283
284        let payload = provider.core.convert_request(&request).expect("payload should be valid");
285        assert_eq!(payload.get("thinking").and_then(|v| v.get("type")).and_then(|v| v.as_str()), Some("enabled"));
286        assert_eq!(payload.get("reasoning_effort").and_then(|v| v.as_str()), Some("low"));
287    }
288
289    #[test]
290    fn payload_rejects_unsupported_xhigh_and_preserves_max() {
291        let provider = ZAIProvider::new("test-key".to_string());
292        let unsupported = LLMRequest {
293            model: models::zai::GLM_5_3.to_string(),
294            messages: vec![Message::user("hello".to_string())].into(),
295            reasoning_effort: Some(ReasoningEffortLevel::XHigh),
296            ..Default::default()
297        };
298        let error = provider
299            .core
300            .convert_request(&unsupported)
301            .expect_err("unsupported effort must be blocked before transport");
302        assert!(error.to_string().contains("unsupported"));
303
304        let supported = LLMRequest {
305            reasoning_effort: Some(ReasoningEffortLevel::Max),
306            ..unsupported
307        };
308        let payload = provider.core.convert_request(&supported).expect("payload should be valid");
309        assert_eq!(payload.get("thinking").and_then(|v| v.get("type")).and_then(|v| v.as_str()), Some("enabled"));
310        assert_eq!(payload.get("reasoning_effort").and_then(|v| v.as_str()), Some("max"));
311    }
312
313    #[test]
314    fn payload_enables_preserved_thinking_when_reasoning_history_present() {
315        let provider = ZAIProvider::new("test-key".to_string());
316        let mut assistant = Message::assistant("tool planning".to_string());
317        assistant.reasoning = Some("reason step 1".to_string());
318
319        let request = LLMRequest {
320            model: models::zai::GLM_5_3.to_string(),
321            messages: vec![assistant].into(),
322            ..Default::default()
323        };
324
325        let payload = provider.core.convert_request(&request).expect("payload should be valid");
326        assert_eq!(payload.get("thinking").and_then(|v| v.get("type")).and_then(|v| v.as_str()), Some("enabled"));
327        assert_eq!(
328            payload
329                .get("thinking")
330                .and_then(|v| v.get("clear_thinking"))
331                .and_then(|v| v.as_bool()),
332            Some(false)
333        );
334    }
335
336    #[test]
337    fn payload_serializes_assistant_reasoning_content() {
338        let provider = ZAIProvider::new("test-key".to_string());
339        let mut assistant = Message::assistant("answer".to_string());
340        assistant.reasoning = Some("chain".to_string());
341
342        let request = LLMRequest {
343            model: models::zai::GLM_5_3.to_string(),
344            messages: vec![assistant].into(),
345            ..Default::default()
346        };
347
348        let payload = provider.core.convert_request(&request).expect("payload should be valid");
349        let messages = payload
350            .get("messages")
351            .and_then(|v| v.as_array())
352            .expect("messages should be serialized");
353        let first = messages.first().expect("at least one message");
354        assert_eq!(first.get("reasoning_content").and_then(|v| v.as_str()), Some("chain"));
355    }
356
357    #[test]
358    fn payload_serializes_web_search_tool() {
359        let provider = ZAIProvider::new("test-key".to_string());
360        let request = LLMRequest {
361            model: models::zai::GLM_5_3.to_string(),
362            messages: vec![Message::user("latest economic events".to_string())].into(),
363            tools: Some(Arc::new(vec![ToolDefinition::web_search(serde_json::json!({
364                "enable": true,
365                "search_engine": "search-prime",
366                "count": 5
367            }))])),
368            ..Default::default()
369        };
370
371        let payload = provider.core.convert_request(&request).expect("payload should be valid");
372        let tools = payload
373            .get("tools")
374            .and_then(|v| v.as_array())
375            .expect("tools should be serialized");
376        let first = tools.first().expect("at least one tool");
377        assert_eq!(first.get("type").and_then(|v| v.as_str()), Some("web_search"));
378        assert_eq!(
379            first
380                .get("web_search")
381                .and_then(|v| v.get("search_engine"))
382                .and_then(|v| v.as_str()),
383            Some("search-prime")
384        );
385    }
386
387    #[test]
388    fn payload_tool_choice_auto_when_requested() {
389        let provider = ZAIProvider::new("test-key".to_string());
390        let request = LLMRequest {
391            model: models::zai::GLM_5_3.to_string(),
392            messages: vec![Message::user("hello".to_string())].into(),
393            tool_choice: Some(ToolChoice::auto()),
394            ..Default::default()
395        };
396
397        let payload = provider.core.convert_request(&request).expect("payload should be valid");
398        assert_eq!(payload.get("tool_choice").and_then(|v| v.as_str()), Some("auto"));
399    }
400
401    #[test]
402    fn payload_forces_tool_choice_to_auto_for_non_auto_permissions() {
403        let provider = ZAIProvider::new("test-key".to_string());
404        let request = LLMRequest {
405            model: models::zai::GLM_5_3.to_string(),
406            messages: vec![Message::user("hello".to_string())].into(),
407            tool_choice: Some(ToolChoice::none()),
408            ..Default::default()
409        };
410
411        let payload = provider.core.convert_request(&request).expect("payload should be valid");
412        assert_eq!(payload.get("tool_choice").and_then(|v| v.as_str()), Some("auto"));
413    }
414
415    #[test]
416    fn payload_defaults_tool_choice_to_auto_when_tools_provided() {
417        let provider = ZAIProvider::new("test-key".to_string());
418        let request = LLMRequest {
419            model: models::zai::GLM_5_3.to_string(),
420            messages: vec![Message::user("hello".to_string())].into(),
421            tools: Some(Arc::new(vec![ToolDefinition::function(
422                "get_weather".to_string(),
423                "Get weather".to_string(),
424                serde_json::json!({
425                    "type": "object",
426                    "properties": {
427                        "location": {"type": "string"}
428                    },
429                    "required": ["location"]
430                }),
431            )])),
432            ..Default::default()
433        };
434
435        let payload = provider.core.convert_request(&request).expect("payload should be valid");
436        assert_eq!(payload.get("tool_choice").and_then(|v| v.as_str()), Some("auto"));
437    }
438
439    #[test]
440    fn payload_enables_json_mode_when_output_format_requested() {
441        let provider = ZAIProvider::new("test-key".to_string());
442        let request = LLMRequest {
443            model: models::zai::GLM_5_3.to_string(),
444            messages: vec![Message::user("return json".to_string())].into(),
445            output_format: Some(serde_json::json!({
446                "type": "object",
447                "properties": {
448                    "sentiment": {"type": "string"}
449                }
450            })),
451            ..Default::default()
452        };
453
454        let payload = provider.core.convert_request(&request).expect("payload should be valid");
455        assert_eq!(
456            payload
457                .get("response_format")
458                .and_then(|v| v.get("type"))
459                .and_then(|v| v.as_str()),
460            Some("json_object")
461        );
462    }
463
464    #[test]
465    fn payload_keeps_json_mode_when_thinking_disabled() {
466        let provider = ZAIProvider::new("test-key".to_string());
467        let request = LLMRequest {
468            model: models::zai::GLM_5_3.to_string(),
469            messages: vec![Message::user("return json".to_string())].into(),
470            output_format: Some(serde_json::json!({
471                "type": "object",
472                "properties": {
473                    "sentiment": {"type": "string"}
474                }
475            })),
476            reasoning_effort: Some(ReasoningEffortLevel::None),
477            ..Default::default()
478        };
479
480        let payload = provider.core.convert_request(&request).expect("payload should be valid");
481        assert_eq!(
482            payload
483                .get("response_format")
484                .and_then(|v| v.get("type"))
485                .and_then(|v| v.as_str()),
486            Some("json_object")
487        );
488        assert_eq!(payload.get("thinking").and_then(|v| v.get("type")).and_then(|v| v.as_str()), Some("disabled"));
489    }
490}