1#[derive(Debug, Clone, PartialEq)]
9pub enum RouteTarget {
10 Template(String),
12 DirectModel(String),
14 Pipeline(Vec<String>),
16 Reject { reason: String },
18}
19
20#[derive(Debug, Clone)]
22pub enum RoutingCondition {
23 ContainsKeyword(String),
25 LongerThan(usize),
27 ShorterThan(usize),
29 MatchesRegex(String),
31 HasIntent(String),
33 Always,
35}
36
37#[derive(Debug, Clone)]
39pub struct RoutingRule {
40 pub name: String,
42 pub priority: u32,
44 pub conditions: Vec<RoutingCondition>,
46 pub target: RouteTarget,
48 pub match_all: bool,
50}
51
52#[derive(Debug, Clone)]
54pub struct RoutingDecision {
55 pub rule_name: String,
57 pub target: RouteTarget,
59 pub confidence: f64,
61 pub matched_conditions: Vec<String>,
63}
64
65pub struct PromptRouter {
69 rules: Vec<RoutingRule>,
70}
71
72impl PromptRouter {
73 pub fn new() -> Self {
75 Self { rules: Vec::new() }
76 }
77
78 pub fn add_rule(&mut self, rule: RoutingRule) {
80 self.rules.push(rule);
81 self.rules.sort_by_key(|r| std::cmp::Reverse(r.priority));
82 }
83
84 pub fn route(&self, prompt: &str, intent: Option<&str>) -> RoutingDecision {
88 for rule in &self.rules {
89 let (matched, matched_conds) =
90 evaluate_rule(rule, prompt, intent);
91 if matched {
92 let confidence = if rule.conditions.is_empty() {
93 0.5
94 } else {
95 matched_conds.len() as f64 / rule.conditions.len() as f64
96 };
97 return RoutingDecision {
98 rule_name: rule.name.clone(),
99 target: rule.target.clone(),
100 confidence,
101 matched_conditions: matched_conds,
102 };
103 }
104 }
105
106 RoutingDecision {
108 rule_name: "default".to_string(),
109 target: RouteTarget::DirectModel("default".to_string()),
110 confidence: 1.0,
111 matched_conditions: vec!["Always (fallback)".to_string()],
112 }
113 }
114
115 pub fn explain_routing(
117 &self,
118 prompt: &str,
119 ) -> Vec<(String, bool, Vec<String>)> {
120 self.rules
121 .iter()
122 .map(|rule| {
123 let (matched, conds) = evaluate_rule(rule, prompt, None);
124 (rule.name.clone(), matched, conds)
125 })
126 .collect()
127 }
128}
129
130impl Default for PromptRouter {
131 fn default() -> Self {
132 Self::new()
133 }
134}
135
136pub fn evaluate_condition(
140 condition: &RoutingCondition,
141 prompt: &str,
142 intent: Option<&str>,
143) -> bool {
144 match condition {
145 RoutingCondition::ContainsKeyword(kw) => {
146 prompt.to_lowercase().contains(&kw.to_lowercase())
147 }
148 RoutingCondition::LongerThan(n) => prompt.len() > *n,
149 RoutingCondition::ShorterThan(n) => prompt.len() < *n,
150 RoutingCondition::MatchesRegex(pattern) => glob_match(pattern, prompt),
151 RoutingCondition::HasIntent(expected) => intent
152 .map(|i| i.to_lowercase() == expected.to_lowercase())
153 .unwrap_or(false),
154 RoutingCondition::Always => true,
155 }
156}
157
158fn condition_description(condition: &RoutingCondition) -> String {
159 match condition {
160 RoutingCondition::ContainsKeyword(kw) => format!("contains keyword '{kw}'"),
161 RoutingCondition::LongerThan(n) => format!("longer than {n} chars"),
162 RoutingCondition::ShorterThan(n) => format!("shorter than {n} chars"),
163 RoutingCondition::MatchesRegex(p) => format!("matches glob '{p}'"),
164 RoutingCondition::HasIntent(i) => format!("has intent '{i}'"),
165 RoutingCondition::Always => "always".to_string(),
166 }
167}
168
169fn evaluate_rule(
172 rule: &RoutingRule,
173 prompt: &str,
174 intent: Option<&str>,
175) -> (bool, Vec<String>) {
176 let mut matched_conds: Vec<String> = Vec::new();
177
178 if rule.conditions.is_empty() {
179 return (true, vec!["(no conditions)".to_string()]);
181 }
182
183 for cond in &rule.conditions {
184 if evaluate_condition(cond, prompt, intent) {
185 matched_conds.push(condition_description(cond));
186 }
187 }
188
189 let overall = if rule.match_all {
190 matched_conds.len() == rule.conditions.len()
191 } else {
192 !matched_conds.is_empty()
193 };
194
195 (overall, matched_conds)
196}
197
198fn glob_match(pattern: &str, text: &str) -> bool {
203 let p: Vec<char> = pattern.to_lowercase().chars().collect();
204 let t: Vec<char> = text.to_lowercase().chars().collect();
205 glob_match_inner(&p, &t)
206}
207
208fn glob_match_inner(pattern: &[char], text: &[char]) -> bool {
209 match (pattern.first(), text.first()) {
210 (None, None) => true,
211 (None, Some(_)) => false,
212 (Some('*'), _) => {
213 glob_match_inner(&pattern[1..], text)
215 || (!text.is_empty() && glob_match_inner(pattern, &text[1..]))
216 }
217 (Some('?'), Some(_)) => glob_match_inner(&pattern[1..], &text[1..]),
218 (Some('?'), None) => false,
219 (Some(p), Some(t)) => p == t && glob_match_inner(&pattern[1..], &text[1..]),
220 (Some(_), None) => false,
221 }
222}
223
224pub struct PromptRouterBuilder {
228 pending_conditions: Vec<RoutingCondition>,
229 pending_name: String,
230 pending_priority: u32,
231 pending_match_all: bool,
232 rules: Vec<RoutingRule>,
233}
234
235impl PromptRouterBuilder {
236 pub fn new() -> Self {
238 Self {
239 pending_conditions: Vec::new(),
240 pending_name: "rule".to_string(),
241 pending_priority: 0,
242 pending_match_all: false,
243 rules: Vec::new(),
244 }
245 }
246
247 pub fn named(mut self, name: &str) -> Self {
249 self.pending_name = name.to_string();
250 self
251 }
252
253 pub fn with_priority(mut self, priority: u32) -> Self {
255 self.pending_priority = priority;
256 self
257 }
258
259 pub fn all_of(mut self) -> Self {
261 self.pending_match_all = true;
262 self
263 }
264
265 pub fn when(mut self, condition: RoutingCondition) -> Self {
267 self.pending_conditions.push(condition);
268 self
269 }
270
271 pub fn route_to(mut self, target: RouteTarget) -> Self {
273 self.rules.push(RoutingRule {
274 name: self.pending_name.clone(),
275 priority: self.pending_priority,
276 conditions: std::mem::take(&mut self.pending_conditions),
277 target,
278 match_all: self.pending_match_all,
279 });
280 self.pending_name = "rule".to_string();
282 self.pending_priority = 0;
283 self.pending_match_all = false;
284 self
285 }
286
287 pub fn build(self) -> PromptRouter {
289 let mut router = PromptRouter::new();
290 for rule in self.rules {
291 router.add_rule(rule);
292 }
293 router
294 }
295}
296
297impl Default for PromptRouterBuilder {
298 fn default() -> Self {
299 Self::new()
300 }
301}
302
303#[cfg(test)]
306mod tests {
307 use super::*;
308
309 fn keyword_rule(kw: &str, priority: u32, target: RouteTarget) -> RoutingRule {
310 RoutingRule {
311 name: format!("kw:{kw}"),
312 priority,
313 conditions: vec![RoutingCondition::ContainsKeyword(kw.to_string())],
314 target,
315 match_all: false,
316 }
317 }
318
319 #[test]
320 fn keyword_match_routing() {
321 let mut router = PromptRouter::new();
322 router.add_rule(keyword_rule("summarize", 10, RouteTarget::Template("summary".to_string())));
323
324 let decision = router.route("Please summarize this document.", None);
325 assert_eq!(decision.rule_name, "kw:summarize");
326 match decision.target {
327 RouteTarget::Template(t) => assert_eq!(t, "summary"),
328 other => panic!("unexpected target: {other:?}"),
329 }
330 assert!(!decision.matched_conditions.is_empty());
331 }
332
333 #[test]
334 fn priority_ordering() {
335 let mut router = PromptRouter::new();
336 router.add_rule(keyword_rule(
337 "code",
338 5,
339 RouteTarget::DirectModel("low-priority-model".to_string()),
340 ));
341 router.add_rule(keyword_rule(
342 "code",
343 20,
344 RouteTarget::DirectModel("high-priority-model".to_string()),
345 ));
346
347 let decision = router.route("write some code please", None);
348 match &decision.target {
349 RouteTarget::DirectModel(m) => assert_eq!(m, "high-priority-model"),
350 other => panic!("unexpected target: {other:?}"),
351 }
352 }
353
354 #[test]
355 fn and_vs_or_conditions() {
356 let and_rule = RoutingRule {
357 name: "and-rule".to_string(),
358 priority: 10,
359 conditions: vec![
360 RoutingCondition::ContainsKeyword("hello".to_string()),
361 RoutingCondition::ContainsKeyword("world".to_string()),
362 ],
363 target: RouteTarget::DirectModel("and-model".to_string()),
364 match_all: true,
365 };
366 let or_rule = RoutingRule {
367 name: "or-rule".to_string(),
368 priority: 5,
369 conditions: vec![
370 RoutingCondition::ContainsKeyword("hello".to_string()),
371 RoutingCondition::ContainsKeyword("world".to_string()),
372 ],
373 target: RouteTarget::DirectModel("or-model".to_string()),
374 match_all: false,
375 };
376
377 let mut router = PromptRouter::new();
378 router.add_rule(and_rule);
379 router.add_rule(or_rule);
380
381 let decision = router.route("hello there", None);
383 assert_eq!(decision.rule_name, "or-rule");
384
385 let decision2 = router.route("hello world", None);
387 assert_eq!(decision2.rule_name, "and-rule");
388 }
389
390 #[test]
391 fn no_match_falls_to_default() {
392 let router = PromptRouter::new();
393 let decision = router.route("some random prompt", None);
394 assert_eq!(decision.rule_name, "default");
395 match decision.target {
396 RouteTarget::DirectModel(m) => assert_eq!(m, "default"),
397 other => panic!("unexpected: {other:?}"),
398 }
399 }
400
401 #[test]
402 fn reject_rule() {
403 let mut router = PromptRouter::new();
404 router.add_rule(RoutingRule {
405 name: "reject-profanity".to_string(),
406 priority: 100,
407 conditions: vec![RoutingCondition::ContainsKeyword("badword".to_string())],
408 target: RouteTarget::Reject { reason: "content policy".to_string() },
409 match_all: false,
410 });
411
412 let decision = router.route("this contains badword unfortunately", None);
413 match decision.target {
414 RouteTarget::Reject { reason } => assert_eq!(reason, "content policy"),
415 other => panic!("expected Reject, got {other:?}"),
416 }
417 }
418
419 #[test]
420 fn builder_api() {
421 let router = PromptRouterBuilder::new()
422 .named("summary-rule")
423 .with_priority(10)
424 .when(RoutingCondition::ContainsKeyword("summarize".to_string()))
425 .route_to(RouteTarget::Template("summary".to_string()))
426 .named("reject-rule")
427 .with_priority(50)
428 .when(RoutingCondition::ContainsKeyword("spam".to_string()))
429 .route_to(RouteTarget::Reject { reason: "spam detected".to_string() })
430 .build();
431
432 let d = router.route("please summarize", None);
433 assert_eq!(d.rule_name, "summary-rule");
434
435 let d2 = router.route("buy my spam product", None);
436 assert_eq!(d2.rule_name, "reject-rule");
437 }
438
439 #[test]
440 fn glob_matching() {
441 assert!(glob_match("hello*", "hello world"));
442 assert!(glob_match("*world", "hello world"));
443 assert!(glob_match("hell? world", "hello world"));
444 assert!(!glob_match("hell? world", "hell world"));
445 assert!(glob_match("*", "anything at all"));
446 }
447}