Skip to main content

ai_agents_runtime/spec/
llm.rs

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