Skip to main content

tokio_prompt_orchestrator/
prompt_router.rs

1//! Route prompts to different handlers based on content and intent.
2//!
3//! Rules are evaluated in descending priority order; the first matching rule
4//! determines the [`RouteTarget`].  If no rule matches, a built-in default
5//! rule that always routes to `DirectModel("default")` is used.
6
7/// Where a matched prompt should be sent.
8#[derive(Debug, Clone, PartialEq)]
9pub enum RouteTarget {
10    /// Expand via a named template.
11    Template(String),
12    /// Send directly to the named model.
13    DirectModel(String),
14    /// Run through a named sequence of pipeline stages.
15    Pipeline(Vec<String>),
16    /// Reject the prompt with an explanation.
17    Reject { reason: String },
18}
19
20/// A condition that can be tested against a prompt string and optional intent.
21#[derive(Debug, Clone)]
22pub enum RoutingCondition {
23    /// Prompt contains the given keyword (case-insensitive).
24    ContainsKeyword(String),
25    /// Prompt is longer than N characters.
26    LongerThan(usize),
27    /// Prompt is shorter than N characters.
28    ShorterThan(usize),
29    /// Prompt matches a glob pattern (`*` = any substring, `?` = any char).
30    MatchesRegex(String),
31    /// The supplied intent string equals this value (case-insensitive).
32    HasIntent(String),
33    /// Always evaluates to `true`.
34    Always,
35}
36
37/// A single routing rule with priority and conditions.
38#[derive(Debug, Clone)]
39pub struct RoutingRule {
40    /// Human-readable name for logging and debugging.
41    pub name: String,
42    /// Higher values are evaluated first.
43    pub priority: u32,
44    /// Conditions that must be satisfied.
45    pub conditions: Vec<RoutingCondition>,
46    /// Where to send the prompt if this rule fires.
47    pub target: RouteTarget,
48    /// `true` → all conditions must match (AND); `false` → any condition (OR).
49    pub match_all: bool,
50}
51
52/// The result of routing a prompt.
53#[derive(Debug, Clone)]
54pub struct RoutingDecision {
55    /// Name of the rule that fired.
56    pub rule_name: String,
57    /// Resolved target.
58    pub target: RouteTarget,
59    /// Confidence in [0.0, 1.0].
60    pub confidence: f64,
61    /// Which condition descriptions contributed to the match.
62    pub matched_conditions: Vec<String>,
63}
64
65// ── Router ────────────────────────────────────────────────────────────────────
66
67/// Routes prompts to targets by evaluating a sorted list of [`RoutingRule`]s.
68pub struct PromptRouter {
69    rules: Vec<RoutingRule>,
70}
71
72impl PromptRouter {
73    /// Create an empty router (no rules).
74    pub fn new() -> Self {
75        Self { rules: Vec::new() }
76    }
77
78    /// Add a rule; the internal list is kept sorted by priority descending.
79    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    /// Route `prompt` (with optional `intent`) and return the first matching
85    /// [`RoutingDecision`].  Falls back to `DirectModel("default")` if no rule
86    /// matches.
87    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        // Default fallback
107        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    /// For every rule, return `(rule_name, matched, reason_strings)`.
116    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
136// ── Condition evaluation ──────────────────────────────────────────────────────
137
138/// Evaluate a single condition; returns `(matched, description)`.
139pub 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
169/// Evaluate an entire rule against `prompt` + `intent`.
170/// Returns `(overall_match, list_of_descriptions_for_matched_conditions)`.
171fn 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        // No conditions → always fire.
180        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
198// ── Glob matching ─────────────────────────────────────────────────────────────
199
200/// Minimal glob-style matching: `*` matches any substring, `?` matches any
201/// single character.  Case-insensitive.
202fn 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            // Star: match zero characters or consume one text character.
214            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
224// ── Builder ───────────────────────────────────────────────────────────────────
225
226/// Fluent builder for [`PromptRouter`].
227pub 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    /// Start a new builder.
237    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    /// Name the next rule.
248    pub fn named(mut self, name: &str) -> Self {
249        self.pending_name = name.to_string();
250        self
251    }
252
253    /// Set priority for the next rule.
254    pub fn with_priority(mut self, priority: u32) -> Self {
255        self.pending_priority = priority;
256        self
257    }
258
259    /// Require all conditions (AND mode).
260    pub fn all_of(mut self) -> Self {
261        self.pending_match_all = true;
262        self
263    }
264
265    /// Add a condition to the pending rule.
266    pub fn when(mut self, condition: RoutingCondition) -> Self {
267        self.pending_conditions.push(condition);
268        self
269    }
270
271    /// Commit the pending rule with the given target and reset pending state.
272    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        // Reset pending state.
281        self.pending_name = "rule".to_string();
282        self.pending_priority = 0;
283        self.pending_match_all = false;
284        self
285    }
286
287    /// Consume the builder and return a [`PromptRouter`].
288    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// ── Tests ─────────────────────────────────────────────────────────────────────
304
305#[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        // Only "hello" → AND rule should fail, OR rule should fire.
382        let decision = router.route("hello there", None);
383        assert_eq!(decision.rule_name, "or-rule");
384
385        // Both keywords → AND rule fires (higher priority).
386        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}