Skip to main content

ai_agents_runtime/spec/
llm.rs

1//! LLM configuration types
2
3use ai_agents_core::ToolChoice;
4use serde::{Deserialize, Serialize};
5use std::collections::HashMap;
6
7#[derive(Debug, Clone, Serialize, Deserialize, Default)]
8pub struct CliMetadata {
9    #[serde(default, skip_serializing_if = "Option::is_none")]
10    pub welcome: Option<String>,
11
12    #[serde(default, skip_serializing_if = "Vec::is_empty")]
13    pub hints: Vec<String>,
14
15    #[serde(default, skip_serializing_if = "Option::is_none")]
16    pub show_tools: Option<bool>,
17
18    #[serde(default, skip_serializing_if = "Option::is_none")]
19    pub show_state: Option<bool>,
20
21    #[serde(default, skip_serializing_if = "Option::is_none")]
22    pub show_timing: Option<bool>,
23
24    #[serde(default, skip_serializing_if = "Option::is_none")]
25    pub streaming: Option<bool>,
26
27    #[serde(default, skip_serializing_if = "Option::is_none")]
28    pub prompt_style: Option<CliPromptStyle>,
29
30    #[serde(default, skip_serializing_if = "Option::is_none")]
31    pub disable_builtin_commands: Option<bool>,
32
33    #[serde(default, skip_serializing_if = "Option::is_none")]
34    pub hitl: Option<CliHitlMetadata>,
35
36    #[serde(default, skip_serializing_if = "Option::is_none")]
37    pub theme: Option<String>,
38}
39
40#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
41#[serde(rename_all = "snake_case")]
42pub enum CliPromptStyle {
43    Simple,
44    WithState,
45}
46
47/// Controls how the CLI handles HITL approval requests at runtime.
48#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)]
49#[serde(rename_all = "snake_case")]
50pub enum CliHitlStyle {
51    /// Interactive y/N prompt in the terminal (default).
52    #[default]
53    Prompt,
54    /// Silently approve all requests.
55    AutoApprove,
56    /// Silently reject all requests.
57    AutoReject,
58}
59
60/// CLI-specific HITL display settings from `metadata.cli.hitl`.
61#[derive(Debug, Clone, Serialize, Deserialize, Default)]
62pub struct CliHitlMetadata {
63    #[serde(default, skip_serializing_if = "Option::is_none")]
64    pub style: Option<CliHitlStyle>,
65
66    #[serde(default, skip_serializing_if = "Option::is_none")]
67    pub show_context: Option<bool>,
68}
69
70/// Configuration for LLM provider
71#[derive(Debug, Clone, Serialize, Deserialize)]
72pub struct LLMConfig {
73    pub provider: String,
74
75    pub model: String,
76
77    #[serde(default = "default_temperature")]
78    pub temperature: f32,
79
80    #[serde(default = "default_max_tokens")]
81    pub max_tokens: u32,
82
83    #[serde(skip_serializing_if = "Option::is_none")]
84    pub top_p: Option<f32>,
85
86    /// Base URL for the LLM provider API.
87    /// Required for `openai-compatible`; optional override for other providers.
88    #[serde(default, skip_serializing_if = "Option::is_none")]
89    pub base_url: Option<String>,
90
91    /// Environment variable name containing the API key.
92    /// Overrides the provider's default env var (e.g. OPENAI_API_KEY).
93    #[serde(default, skip_serializing_if = "Option::is_none")]
94    pub api_key_env: Option<String>,
95
96    /// Request timeout in seconds.
97    #[serde(default, skip_serializing_if = "Option::is_none")]
98    pub timeout_seconds: Option<u64>,
99
100    /// Enable extended thinking / reasoning mode.
101    #[serde(default, skip_serializing_if = "Option::is_none")]
102    pub reasoning: Option<bool>,
103
104    /// Reasoning effort level: "low", "medium", or "high".
105    #[serde(default, skip_serializing_if = "Option::is_none")]
106    pub reasoning_effort: Option<String>,
107
108    /// Maximum token budget for reasoning.
109    #[serde(default, skip_serializing_if = "Option::is_none")]
110    pub reasoning_budget_tokens: Option<u32>,
111
112    /// Override whether the provider supports function/tool calling.
113    #[serde(default, skip_serializing_if = "Option::is_none")]
114    pub function_calling: Option<bool>,
115
116    /// Opt in to provider-native or runtime-enforced tool selection.
117    #[serde(default, skip_serializing_if = "Option::is_none")]
118    pub tool_choice: Option<ToolChoice>,
119
120    /// Override whether the provider supports vision inputs.
121    #[serde(default, skip_serializing_if = "Option::is_none")]
122    pub vision: Option<bool>,
123
124    /// Override whether the provider supports JSON mode.
125    #[serde(default, skip_serializing_if = "Option::is_none")]
126    pub json_mode: Option<bool>,
127
128    /// Additional provider-specific configuration
129    #[serde(flatten)]
130    pub extra: HashMap<String, serde_json::Value>,
131}
132
133fn default_temperature() -> f32 {
134    0.7
135}
136
137fn default_max_tokens() -> u32 {
138    2000
139}
140
141impl Default for LLMConfig {
142    fn default() -> Self {
143        Self {
144            provider: "openai".to_string(),
145            model: "gpt-4".to_string(),
146            temperature: default_temperature(),
147            max_tokens: default_max_tokens(),
148            top_p: None,
149            base_url: None,
150            api_key_env: None,
151            timeout_seconds: None,
152            reasoning: None,
153            reasoning_effort: None,
154            reasoning_budget_tokens: None,
155            function_calling: None,
156            tool_choice: None,
157            vision: None,
158            json_mode: None,
159            extra: HashMap::new(),
160        }
161    }
162}
163
164#[derive(Debug, Clone, Serialize, Deserialize)]
165#[serde(deny_unknown_fields)]
166pub struct LLMSelector {
167    #[serde(default = "default_alias")]
168    pub default: String,
169    #[serde(default)]
170    pub router: Option<String>,
171}
172
173fn default_alias() -> String {
174    "default".to_string()
175}
176
177impl Default for LLMSelector {
178    fn default() -> Self {
179        Self {
180            default: default_alias(),
181            router: None,
182        }
183    }
184}
185
186impl LLMSelector {
187    pub fn new(default: impl Into<String>) -> Self {
188        Self {
189            default: default.into(),
190            router: None,
191        }
192    }
193
194    pub fn with_router(mut self, router: impl Into<String>) -> Self {
195        self.router = Some(router.into());
196        self
197    }
198}
199
200#[cfg(test)]
201mod tests {
202    use super::*;
203
204    #[test]
205    fn test_cli_metadata_deserialize() {
206        let yaml = r#"
207welcome: "=== Demo ==="
208hints:
209  - "Try: hello"
210  - "Try: help"
211show_tools: true
212show_state: false
213show_timing: true
214streaming: true
215prompt_style: with_state
216disable_builtin_commands: false
217"#;
218        let metadata: CliMetadata = serde_yaml::from_str(yaml).unwrap();
219        assert_eq!(metadata.welcome.as_deref(), Some("=== Demo ==="));
220        assert_eq!(metadata.hints.len(), 2);
221        assert_eq!(metadata.show_tools, Some(true));
222        assert_eq!(metadata.show_state, Some(false));
223        assert_eq!(metadata.show_timing, Some(true));
224        assert_eq!(metadata.streaming, Some(true));
225        assert_eq!(metadata.prompt_style, Some(CliPromptStyle::WithState));
226        assert_eq!(metadata.disable_builtin_commands, Some(false));
227        assert!(metadata.hitl.is_none());
228    }
229
230    #[test]
231    fn test_llm_config_default() {
232        let config = LLMConfig::default();
233        assert_eq!(config.provider, "openai");
234        assert_eq!(config.model, "gpt-4");
235        assert_eq!(config.temperature, 0.7);
236        assert_eq!(config.max_tokens, 2000);
237        assert_eq!(config.base_url, None);
238        assert_eq!(config.api_key_env, None);
239    }
240
241    #[test]
242    fn test_llm_config_with_base_url() {
243        let yaml = r#"
244provider: openai-compatible
245model: llama3.2
246base_url: http://localhost:1234/v1
247"#;
248        let config: LLMConfig = serde_yaml::from_str(yaml).unwrap();
249        assert_eq!(config.provider, "openai-compatible");
250        assert_eq!(
251            config.base_url,
252            Some("http://localhost:1234/v1".to_string())
253        );
254    }
255
256    #[test]
257    fn test_llm_config_with_api_key_env() {
258        let yaml = r#"
259provider: openai-compatible
260model: my-model
261base_url: http://my-server:8080/v1
262api_key_env: MY_SERVER_KEY
263"#;
264        let config: LLMConfig = serde_yaml::from_str(yaml).unwrap();
265        assert_eq!(config.api_key_env, Some("MY_SERVER_KEY".to_string()));
266    }
267
268    #[test]
269    fn test_llm_config_base_url_does_not_leak_to_extra() {
270        let yaml = r#"
271provider: openai-compatible
272model: my-model
273base_url: http://localhost:1234/v1
274"#;
275        let config: LLMConfig = serde_yaml::from_str(yaml).unwrap();
276        assert!(!config.extra.contains_key("base_url"));
277    }
278
279    #[test]
280    fn test_llm_config_deserialize() {
281        let yaml = r#"
282provider: openai
283model: gpt-3.5-turbo
284temperature: 0.5
285max_tokens: 1000
286"#;
287        let config: LLMConfig = serde_yaml::from_str(yaml).unwrap();
288        assert_eq!(config.provider, "openai");
289        assert_eq!(config.model, "gpt-3.5-turbo");
290        assert_eq!(config.temperature, 0.5);
291        assert_eq!(config.max_tokens, 1000);
292    }
293
294    #[test]
295    fn test_llm_config_tool_choice_deserialize() {
296        let required: LLMConfig =
297            serde_yaml::from_str("provider: openai\nmodel: gpt-5.1-mini\ntool_choice: required\n")
298                .unwrap();
299        assert_eq!(required.tool_choice, Some(ToolChoice::Required));
300        assert!(!required.extra.contains_key("tool_choice"));
301
302        let specific: LLMConfig = serde_yaml::from_str(
303            "provider: openai\nmodel: gpt-5.1-mini\ntool_choice:\n  specific: random\n",
304        )
305        .unwrap();
306        assert_eq!(
307            specific.tool_choice,
308            Some(ToolChoice::Specific("random".to_string()))
309        );
310    }
311
312    #[test]
313    fn test_llm_config_with_defaults() {
314        let yaml = r#"
315provider: openai
316model: gpt-4
317"#;
318        let config: LLMConfig = serde_yaml::from_str(yaml).unwrap();
319        assert_eq!(config.temperature, 0.7); // default
320        assert_eq!(config.max_tokens, 2000); // default
321    }
322
323    #[test]
324    fn test_llm_config_extra_fields() {
325        let yaml = r#"
326provider: openai
327model: gpt-4
328custom_field: "value"
329another_field: 123
330"#;
331        let config: LLMConfig = serde_yaml::from_str(yaml).unwrap();
332        assert!(config.extra.contains_key("custom_field"));
333        assert!(config.extra.contains_key("another_field"));
334    }
335
336    #[test]
337    fn test_llm_selector_default() {
338        let selector = LLMSelector::default();
339        assert_eq!(selector.default, "default");
340        assert!(selector.router.is_none());
341    }
342
343    #[test]
344    fn test_llm_selector_with_router() {
345        let selector = LLMSelector::new("main").with_router("cheap");
346        assert_eq!(selector.default, "main");
347        assert_eq!(selector.router, Some("cheap".to_string()));
348    }
349
350    #[test]
351    fn test_cli_hitl_metadata_deserialize() {
352        let yaml = r#"
353style: auto_approve
354show_context: false
355"#;
356        let meta: CliHitlMetadata = serde_yaml::from_str(yaml).unwrap();
357        assert_eq!(meta.style, Some(CliHitlStyle::AutoApprove));
358        assert_eq!(meta.show_context, Some(false));
359    }
360
361    #[test]
362    fn test_cli_hitl_style_default() {
363        assert_eq!(CliHitlStyle::default(), CliHitlStyle::Prompt);
364    }
365
366    #[test]
367    fn test_cli_metadata_with_hitl() {
368        let yaml = r#"
369welcome: "Hello"
370hints: []
371hitl:
372  style: prompt
373  show_context: true
374"#;
375        let meta: CliMetadata = serde_yaml::from_str(yaml).unwrap();
376        let hitl = meta.hitl.unwrap();
377        assert_eq!(hitl.style, Some(CliHitlStyle::Prompt));
378        assert_eq!(hitl.show_context, Some(true));
379    }
380
381    #[test]
382    fn test_llm_selector_deserialize() {
383        let yaml = r#"
384default: main
385router: router_llm
386"#;
387        let selector: LLMSelector = serde_yaml::from_str(yaml).unwrap();
388        assert_eq!(selector.default, "main");
389        assert_eq!(selector.router, Some("router_llm".to_string()));
390    }
391
392    #[test]
393    fn test_llm_config_reasoning_fields_deser() {
394        let yaml = r#"
395provider: openai
396model: o3
397timeout_seconds: 120
398reasoning: true
399reasoning_effort: high
400reasoning_budget_tokens: 16384
401"#;
402        let config: LLMConfig = serde_yaml::from_str(yaml).unwrap();
403        assert_eq!(config.timeout_seconds, Some(120));
404        assert_eq!(config.reasoning, Some(true));
405        assert_eq!(config.reasoning_effort.as_deref(), Some("high"));
406        assert_eq!(config.reasoning_budget_tokens, Some(16384));
407        // Must NOT leak into extra
408        assert!(!config.extra.contains_key("timeout_seconds"));
409        assert!(!config.extra.contains_key("reasoning"));
410        assert!(!config.extra.contains_key("reasoning_effort"));
411        assert!(!config.extra.contains_key("reasoning_budget_tokens"));
412    }
413
414    #[test]
415    fn test_ollama_named_fields_land_in_extra() {
416        let yaml = r#"
417provider: ollama
418model: llama3.1
419num_ctx: 8192
420num_gpu: -1
421keep_alive: 5m
422"#;
423        let config: LLMConfig = serde_yaml::from_str(yaml).unwrap();
424        assert_eq!(config.extra.get("num_ctx"), Some(&serde_json::json!(8192)));
425        assert_eq!(config.extra.get("num_gpu"), Some(&serde_json::json!(-1)));
426        assert_eq!(
427            config.extra.get("keep_alive"),
428            Some(&serde_json::json!("5m"))
429        );
430    }
431
432    #[test]
433    fn test_llm_config_feature_override_fields_deser() {
434        let yaml = r#"
435provider: openai-compatible
436model: qwen3:8b
437base_url: http://localhost:11434/v1
438function_calling: true
439vision: false
440json_mode: true
441"#;
442        let config: LLMConfig = serde_yaml::from_str(yaml).unwrap();
443        assert_eq!(config.function_calling, Some(true));
444        assert_eq!(config.vision, Some(false));
445        assert_eq!(config.json_mode, Some(true));
446        assert!(!config.extra.contains_key("function_calling"));
447        assert!(!config.extra.contains_key("vision"));
448        assert!(!config.extra.contains_key("json_mode"));
449    }
450}