1use 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#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)]
50#[serde(rename_all = "snake_case")]
51pub enum CliHitlStyle {
52 #[default]
54 Prompt,
55 AutoApprove,
57 AutoReject,
59}
60
61#[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#[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 #[serde(default, skip_serializing_if = "Option::is_none")]
90 pub base_url: Option<String>,
91
92 #[serde(default, skip_serializing_if = "Option::is_none")]
95 pub api_key_env: Option<String>,
96
97 #[serde(default, skip_serializing_if = "Option::is_none")]
99 pub timeout_seconds: Option<u64>,
100
101 #[serde(default, skip_serializing_if = "Option::is_none")]
103 pub reasoning: Option<bool>,
104
105 #[serde(default, skip_serializing_if = "Option::is_none")]
107 pub reasoning_effort: Option<String>,
108
109 #[serde(default, skip_serializing_if = "Option::is_none")]
111 pub reasoning_budget_tokens: Option<u32>,
112
113 #[serde(default, skip_serializing_if = "Option::is_none")]
115 pub function_calling: Option<bool>,
116
117 #[serde(default, skip_serializing_if = "Option::is_none")]
119 pub tool_choice: Option<ToolChoice>,
120
121 #[serde(default, skip_serializing_if = "Option::is_none")]
123 pub vision: Option<bool>,
124
125 #[serde(default, skip_serializing_if = "Option::is_none")]
127 pub json_mode: Option<bool>,
128
129 #[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 pub fn with_router_roles(mut self, config: RouterRolesConfig) -> Self {
198 self.router = Some(RouterSelector::Hierarchical(Box::new(config)));
199 self
200 }
201
202 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); assert_eq!(config.max_tokens, 2000); }
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 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}