1use 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#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)]
49#[serde(rename_all = "snake_case")]
50pub enum CliHitlStyle {
51 #[default]
53 Prompt,
54 AutoApprove,
56 AutoReject,
58}
59
60#[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#[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 #[serde(default, skip_serializing_if = "Option::is_none")]
89 pub base_url: Option<String>,
90
91 #[serde(default, skip_serializing_if = "Option::is_none")]
94 pub api_key_env: Option<String>,
95
96 #[serde(default, skip_serializing_if = "Option::is_none")]
98 pub timeout_seconds: Option<u64>,
99
100 #[serde(default, skip_serializing_if = "Option::is_none")]
102 pub reasoning: Option<bool>,
103
104 #[serde(default, skip_serializing_if = "Option::is_none")]
106 pub reasoning_effort: Option<String>,
107
108 #[serde(default, skip_serializing_if = "Option::is_none")]
110 pub reasoning_budget_tokens: Option<u32>,
111
112 #[serde(default, skip_serializing_if = "Option::is_none")]
114 pub function_calling: Option<bool>,
115
116 #[serde(default, skip_serializing_if = "Option::is_none")]
118 pub tool_choice: Option<ToolChoice>,
119
120 #[serde(default, skip_serializing_if = "Option::is_none")]
122 pub vision: Option<bool>,
123
124 #[serde(default, skip_serializing_if = "Option::is_none")]
126 pub json_mode: Option<bool>,
127
128 #[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); assert_eq!(config.max_tokens, 2000); }
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 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}