Skip to main content

synapse/
ai_task_type.rs

1//! Static route-alias → AI task type table.
2//!
3//! Mirrors `PricingTable`: a flat TOML root map, loaded once at startup. The
4//! file is keyed by task type so a taxonomy reads as one line per type:
5//!
6//! ```toml
7//! conversation = ["conversation", "planning", "nl-plan"]
8//! extraction   = ["extract", "doc-extract"]
9//! ```
10//!
11//! It is inverted to an alias → task lookup at load, so resolution is O(1) on
12//! the hot path.
13
14use std::collections::HashMap;
15
16/// Recorded when neither a caller header nor the table supplies a task type.
17pub 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    /// Parse `task = [aliases]` and invert to an alias → task map. An alias
26    /// claimed by two task types is ambiguous and fails at startup rather than
27    /// silently resolving to whichever the map iterated last.
28    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                        // Sorted so the message is stable regardless of map order.
41                        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    /// The task type configured for `alias`, else [`DEFAULT_AI_TASK_TYPE`].
55    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}