1use std::collections::HashMap;
15
16pub const DEFAULT_AI_TASK_TYPE: &str = "simple";
18
19#[derive(Debug, Clone, Default)]
20pub struct AiTaskTypeTable {
21 by_alias: HashMap<String, String>,
22}
23
24impl AiTaskTypeTable {
25 pub fn from_toml_str(s: &str) -> anyhow::Result<Self> {
29 toml::from_str::<HashMap<String, Vec<String>>>(s)?
30 .into_iter()
31 .flat_map(|(task, aliases)| {
32 aliases
33 .into_iter()
34 .map(move |alias| (alias, task.clone()))
35 .collect::<Vec<_>>()
36 })
37 .try_fold(HashMap::new(), |mut acc, (alias, task)| {
38 match acc.insert(alias.clone(), task.clone()) {
39 Some(prev) if prev != task => {
40 let mut pair = [prev, task];
42 pair.sort();
43 let [a, b] = pair;
44 anyhow::bail!(
45 "route alias '{alias}' is mapped to more than one ai task type ('{a}' and '{b}')"
46 )
47 }
48 _ => Ok(acc),
49 }
50 })
51 .map(|by_alias| Self { by_alias })
52 }
53
54 pub fn resolve(&self, alias: &str) -> &str {
56 self.by_alias
57 .get(alias)
58 .map(String::as_str)
59 .unwrap_or(DEFAULT_AI_TASK_TYPE)
60 }
61}
62
63#[cfg(test)]
64mod tests {
65 use super::*;
66
67 const SAMPLE: &str = r#"
68 conversation = ["conversation", "planning", "nl-plan"]
69 extraction = ["extract", "doc-extract"]
70 "#;
71
72 #[test]
73 fn resolves_every_alias_of_a_task_type() {
74 let t = AiTaskTypeTable::from_toml_str(SAMPLE).unwrap();
75 assert_eq!(t.resolve("conversation"), "conversation");
76 assert_eq!(t.resolve("planning"), "conversation");
77 assert_eq!(t.resolve("nl-plan"), "conversation");
78 assert_eq!(t.resolve("doc-extract"), "extraction");
79 }
80
81 #[test]
82 fn unmapped_alias_falls_back_to_simple() {
83 let t = AiTaskTypeTable::from_toml_str(SAMPLE).unwrap();
84 assert_eq!(t.resolve("fast"), "simple");
85 }
86
87 #[test]
88 fn empty_table_resolves_everything_to_simple() {
89 let t = AiTaskTypeTable::default();
90 assert_eq!(t.resolve("conversation"), "simple");
91 }
92
93 #[test]
94 fn an_alias_under_two_task_types_is_a_startup_error() {
95 let err = AiTaskTypeTable::from_toml_str(
96 r#"
97 conversation = ["planning"]
98 extraction = ["planning"]
99 "#,
100 )
101 .unwrap_err()
102 .to_string();
103 assert!(err.contains("planning"), "got {err}");
104 assert!(
105 err.contains("conversation") && err.contains("extraction"),
106 "got {err}"
107 );
108 }
109
110 #[test]
111 fn an_alias_repeated_under_one_task_type_is_allowed() {
112 let t =
113 AiTaskTypeTable::from_toml_str(r#"conversation = ["planning", "planning"]"#).unwrap();
114 assert_eq!(t.resolve("planning"), "conversation");
115 }
116
117 #[test]
118 fn malformed_toml_is_an_error() {
119 assert!(AiTaskTypeTable::from_toml_str("conversation = 5").is_err());
120 }
121}