1use std::sync::Arc;
4
5use ai_agents_core::{LLMError, LLMProvider};
6use serde::{Deserialize, Deserializer, Serialize};
7
8macro_rules! routing_schema {
9 ($( $group:ident : $config:ident { $( $field:ident => $role:ident ),+ } ),+ $(,)?) => {
10 $(
11 #[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
12 #[serde(deny_unknown_fields)]
13 pub struct $config {
14 #[serde(default, skip_serializing_if = "Option::is_none")]
15 pub default: Option<String>,
16 $(#[serde(default, skip_serializing_if = "Option::is_none")]
17 pub $field: Option<String>,)+
18 }
19 )+
20
21 #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
23 pub enum LLMRole { $($( $role, )+)+ }
24
25 #[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
27 #[serde(deny_unknown_fields)]
28 pub struct RouterRolesConfig {
29 #[serde(default, skip_serializing_if = "Option::is_none")]
30 pub default: Option<String>,
31 $(#[serde(default, skip_serializing_if = "Option::is_none")]
32 pub $group: Option<$config>,)+
33 }
34
35 impl LLMRole {
36 pub const ALL: &'static [Self] = &[$($(Self::$role,)+)+];
37
38 pub fn as_path(self) -> &'static str {
39 match self { $($(Self::$role => concat!(stringify!($group), ".", stringify!($field)),)+)+ }
40 }
41
42 pub fn group(self) -> &'static str {
43 match self { $($(Self::$role => stringify!($group),)+)+ }
44 }
45 }
46
47 impl RouterRolesConfig {
48 pub fn select<'a>(&'a self, role: LLMRole, local: Option<&'a str>, main: &'a str) -> (&'a str, LLMSelectionSource) {
50 if let Some(alias) = local { return (alias, LLMSelectionSource::Local); }
51 let (leaf, group) = match role {
52 $($(LLMRole::$role => self.$group.as_ref().map(|config| (config.$field.as_deref(), config.default.as_deref())).unwrap_or((None, None)),)+)+
53 };
54 if let Some(alias) = leaf { return (alias, LLMSelectionSource::Role); }
55 if let Some(alias) = group { return (alias, LLMSelectionSource::Group); }
56 if let Some(alias) = self.default.as_deref() { return (alias, LLMSelectionSource::Router); }
57 (main, LLMSelectionSource::Default)
58 }
59
60 pub fn configured_aliases(&self) -> Vec<(&'static str, &str)> {
62 let mut aliases = Vec::new();
63 if let Some(alias) = self.default.as_deref() { aliases.push(("llm.router.default", alias)); }
64 $(if let Some(config) = &self.$group {
65 if let Some(alias) = config.default.as_deref() { aliases.push((concat!("llm.router.", stringify!($group), ".default"), alias)); }
66 $(if let Some(alias) = config.$field.as_deref() { aliases.push((concat!("llm.router.", stringify!($group), ".", stringify!($field)), alias)); })+
67 })+
68 aliases
69 }
70
71 pub fn validate(&self) -> Result<(), LLMError> {
73 for (path, alias) in self.configured_aliases() {
74 if alias.trim().is_empty() { return Err(LLMError::Config(format!("Invalid {path}: alias must not be empty"))); }
75 }
76 Ok(())
77 }
78 }
79 }
80}
81
82routing_schema! {
83 state: StateRouterConfig { transition => StateTransition, extract => StateExtract },
84 skills: SkillsRouterConfig { selection => SkillsSelection },
85 tools: ToolsRouterConfig { condition => ToolsCondition },
86 process: ProcessRouterConfig { detect => ProcessDetect, extract => ProcessExtract, sanitize => ProcessSanitize, transform => ProcessTransform, validate => ProcessValidate },
87 disambiguation: DisambiguationRouterConfig { detection => DisambiguationDetection, skip => DisambiguationSkip, clarification => DisambiguationClarification, parse => DisambiguationParse, confirmation => DisambiguationConfirmation, confirmation_parse => DisambiguationConfirmationParse, response => DisambiguationResponse },
88 reasoning: ReasoningRouterConfig { selection => ReasoningSelection, planning => ReasoningPlanning, reflection_decision => ReasoningReflectionDecision, reflection_evaluation => ReasoningReflectionEvaluation },
89 memory: MemoryRouterConfig { summarize => MemorySummarize, merge => MemoryMerge, facts => MemoryFacts, relationships => MemoryRelationships },
90 context: ContextRouterConfig { summarize => ContextSummarize },
91 orchestration: OrchestrationRouterConfig { routing => OrchestrationRouting, handoff => OrchestrationHandoff, speaker => OrchestrationSpeaker, consensus => OrchestrationConsensus, synthesis => OrchestrationSynthesis, vote => OrchestrationVote, tiebreak => OrchestrationTiebreak, summary => OrchestrationSummary },
92 hitl: HitlRouterConfig { message => HitlMessage },
93 spawner: SpawnerRouterConfig { generation => SpawnerGeneration, repair => SpawnerRepair },
94 web: WebRouterConfig { extract => WebExtract },
95 evaluation: EvaluationRouterConfig { response => EvaluationResponse, facts => EvaluationFacts },
96}
97
98#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
100#[serde(untagged)]
101pub enum RouterSelector {
102 Alias(String),
103 Hierarchical(Box<RouterRolesConfig>),
104}
105
106impl<'de> Deserialize<'de> for RouterSelector {
107 fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
108 let value = serde_json::Value::deserialize(deserializer)?;
109 if let Some(alias) = value.as_str() {
110 return Ok(Self::Alias(alias.to_string()));
111 }
112 validate_tree_shape(&value).map_err(serde::de::Error::custom)?;
113 let config: RouterRolesConfig =
114 serde_json::from_value(value).map_err(serde::de::Error::custom)?;
115 config.validate().map_err(serde::de::Error::custom)?;
116 Ok(Self::Hierarchical(Box::new(config)))
117 }
118}
119
120fn validate_tree_shape(value: &serde_json::Value) -> Result<(), String> {
122 let root = value
123 .as_object()
124 .ok_or("Invalid llm.router: expected alias or mapping")?;
125 for (group, value) in root {
126 if group == "default" {
127 validate_alias_value("llm.router.default", value)?;
128 continue;
129 }
130 if !LLMRole::ALL.iter().any(|role| role.group() == group) {
131 return Err(format!("Invalid llm.router.{group}: unknown group"));
132 }
133 if value.is_null() {
134 continue;
135 }
136 let fields = value
137 .as_object()
138 .ok_or_else(|| format!("Invalid llm.router.{group}: expected mapping"))?;
139 for (field, value) in fields {
140 let path = format!("{group}.{field}");
141 if field != "default" && !LLMRole::ALL.iter().any(|role| role.as_path() == path) {
142 return Err(format!("Invalid llm.router.{path}: unknown field"));
143 }
144 validate_alias_value(&format!("llm.router.{path}"), value)?;
145 }
146 }
147 Ok(())
148}
149
150fn validate_alias_value(path: &str, value: &serde_json::Value) -> Result<(), String> {
151 if value.is_null() || value.as_str().is_some_and(|alias| !alias.trim().is_empty()) {
152 Ok(())
153 } else {
154 Err(format!("Invalid {path}: expected nonempty alias or null"))
155 }
156}
157
158#[derive(Debug, Clone, Copy, PartialEq, Eq)]
160pub enum LLMSelectionSource {
161 Local,
162 Role,
163 Group,
164 Router,
165 Default,
166}
167
168#[derive(Clone)]
170pub struct ResolvedRoleLLM {
171 pub role: LLMRole,
172 pub alias: String,
173 pub source: LLMSelectionSource,
174 pub provider: Arc<dyn LLMProvider>,
175}
176
177pub fn deserialize_present_alias<'de, D: Deserializer<'de>>(
179 deserializer: D,
180) -> Result<Option<String>, D::Error> {
181 match serde_json::Value::deserialize(deserializer)? {
182 serde_json::Value::String(alias) => Ok(Some(alias)),
183 _ => Err(serde::de::Error::custom(
184 "expected a present alias string; null is not supported",
185 )),
186 }
187}
188
189#[cfg(test)]
190mod tests {
191 use super::*;
192
193 #[test]
194 fn all_roles_inherit_and_preserve_sources() {
195 let mut tree = RouterRolesConfig::default();
196 assert_eq!(LLMRole::ALL.len(), 39);
197 for role in LLMRole::ALL {
198 assert_eq!(
199 tree.select(*role, None, "main"),
200 ("main", LLMSelectionSource::Default)
201 );
202 assert_eq!(
203 tree.select(*role, Some("router"), "main"),
204 ("router", LLMSelectionSource::Local)
205 );
206 }
207 tree.default = Some("fast".into());
208 for role in LLMRole::ALL {
209 assert_eq!(
210 tree.select(*role, None, "main"),
211 ("fast", LLMSelectionSource::Router)
212 );
213 }
214 tree.state = Some(StateRouterConfig {
215 default: Some("group".into()),
216 transition: Some("leaf".into()),
217 extract: None,
218 });
219 assert_eq!(
220 tree.select(LLMRole::StateTransition, None, "main"),
221 ("leaf", LLMSelectionSource::Role)
222 );
223 assert_eq!(
224 tree.select(LLMRole::StateExtract, None, "main"),
225 ("group", LLMSelectionSource::Group)
226 );
227 }
228
229 #[test]
230 fn strict_tree_and_empty_mapping_round_trip() {
231 let empty: RouterSelector = serde_json::from_str("{}").unwrap();
232 assert!(matches!(empty, RouterSelector::Hierarchical(_)));
233 assert_eq!(serde_json::to_string(&empty).unwrap(), "{}");
234 for input in [
235 r#"{"state":{"transiton":"fast"}}"#,
236 r#"{"state":"fast"}"#,
237 r#"{"state":{"transition":" "}}"#,
238 "false",
239 ] {
240 assert!(
241 serde_json::from_str::<RouterSelector>(input).is_err(),
242 "{input}"
243 );
244 }
245 let scalar: RouterSelector = serde_json::from_str(r#""router""#).unwrap();
246 assert_eq!(scalar, RouterSelector::Alias("router".into()));
247 }
248}