Skip to main content

tokenmiser_router/
lib.rs

1//! Classify prompt difficulty and pick a model.
2//!
3//! Tier 0 (`tier0.rs`) is a sub-microsecond heuristic over length, keyword,
4//! role and JSON-mode signals. Tier 1 (`tier1.rs`) is an exemplar-based
5//! semantic classifier sharing the L2 cache's bge-small embedder, used when
6//! Tier 0 is not decisive.
7
8use serde::{Deserialize, Serialize};
9use tokenmiser_providers::ChatRequest;
10
11pub mod dsl;
12pub mod policy;
13pub mod replay;
14pub mod tier0;
15pub mod tier1;
16pub mod tier2;
17
18pub use dsl::{PolicyEngine, RequestView};
19pub use policy::{RoutingPolicy, RoutingTarget};
20pub use replay::{replay, ReplayResult};
21pub use tier0::tier0_difficulty;
22pub use tier1::Tier1Classifier;
23pub use tier2::{should_escalate, CascadeConfig, EscalateDecision};
24
25/// Coarse difficulty class used by the policy to pick a model.
26#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
27#[serde(rename_all = "lowercase")]
28pub enum Difficulty {
29    Easy,
30    Medium,
31    Hard,
32}
33
34/// What model to call, why, and the counterfactual model used for cost
35/// accounting and shadow A/B.
36#[derive(Debug, Clone, Serialize, Deserialize)]
37pub struct RouteDecision {
38    pub target: RoutingTarget,
39    pub difficulty: Difficulty,
40    pub tier: RouteTier,
41    /// Reasoning trace; not yet populated.
42    #[serde(default, skip_serializing_if = "Option::is_none")]
43    pub reasoning: Option<String>,
44    /// The model a frontier-only route would have called; drives the
45    /// counterfactual savings ledger.
46    #[serde(default, skip_serializing_if = "Option::is_none")]
47    pub counterfactual_model: Option<String>,
48}
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
51#[serde(rename_all = "lowercase")]
52pub enum RouteTier {
53    /// Caller named a model; no classification ran.
54    Explicit,
55    /// Tier 0 heuristic was sufficient.
56    Heuristic,
57    /// Tier 1 semantic classifier was used.
58    Semantic,
59}
60
61/// The orchestrator that runs Tier 0 → Tier 1 → policy in order.
62pub struct Router {
63    policy: RoutingPolicy,
64    tier1: Option<Tier1Classifier>,
65}
66
67impl Router {
68    pub fn new(policy: RoutingPolicy, tier1: Option<Tier1Classifier>) -> Self {
69        Self { policy, tier1 }
70    }
71
72    /// The policy target for a difficulty band.
73    pub fn policy_target(&self, d: Difficulty) -> RoutingTarget {
74        self.policy.choose(d)
75    }
76
77    /// Decide where to send `req`, honoring an explicitly named model.
78    pub fn decide(&self, req: &ChatRequest) -> RouteDecision {
79        let requested = req.model.as_str();
80        let auto = requested == "auto" || requested == "tokenmiser:auto";
81
82        if !auto {
83            // Respect the caller's model but still classify, so /stats keeps
84            // reporting difficulty.
85            let difficulty = tier0_difficulty(req);
86            return RouteDecision {
87                target: RoutingTarget::passthrough(requested),
88                difficulty,
89                tier: RouteTier::Explicit,
90                reasoning: None,
91                counterfactual_model: self.policy.frontier_for(Difficulty::Hard),
92            };
93        }
94
95        // Tier 0 first, Tier 1 only when Tier 0 is not decisive.
96        let t0 = tier0_difficulty(req);
97        let (difficulty, tier) = match (t0, &self.tier1) {
98            (Difficulty::Medium, Some(t1)) => (t1.classify(req), RouteTier::Semantic),
99            _ => (t0, RouteTier::Heuristic),
100        };
101
102        let target = self.policy.choose(difficulty);
103        let counterfactual = self.policy.frontier_for(Difficulty::Hard);
104
105        RouteDecision {
106            target,
107            difficulty,
108            tier,
109            reasoning: None,
110            counterfactual_model: counterfactual,
111        }
112    }
113}
114
115#[cfg(test)]
116mod tests {
117    use super::*;
118    use tokenmiser_providers::ChatMessage;
119
120    fn user(model: &str, s: &str) -> ChatRequest {
121        ChatRequest {
122            model: model.into(),
123            messages: vec![ChatMessage {
124                role: "user".into(),
125                content: serde_json::Value::String(s.into()),
126                extra: Default::default(),
127            }],
128            temperature: None,
129            max_tokens: None,
130            top_p: None,
131            stream: None,
132            extra: Default::default(),
133        }
134    }
135
136    #[test]
137    fn explicit_model_request_is_passthrough() {
138        let router = Router::new(RoutingPolicy::default(), None);
139        let d = router.decide(&user("gpt-5", "anything"));
140        assert!(matches!(d.tier, RouteTier::Explicit));
141        assert_eq!(d.target.model, "gpt-5");
142    }
143
144    #[test]
145    fn auto_easy_routes_to_local() {
146        let router = Router::new(RoutingPolicy::default(), None);
147        let d = router.decide(&user("auto", "what is 2+2?"));
148        assert_eq!(d.difficulty, Difficulty::Easy);
149        // default policy maps Easy → ollama:llama3.2
150        assert!(d.target.model.contains("llama") || d.target.provider == "ollama");
151    }
152
153    #[test]
154    fn auto_hard_routes_to_frontier() {
155        let router = Router::new(RoutingPolicy::default(), None);
156        let d = router.decide(&user("auto", "refactor this auth middleware to JWT"));
157        assert_eq!(d.difficulty, Difficulty::Hard);
158        assert!(d.target.model.contains("opus") || d.target.model.contains("gpt"));
159    }
160}