1use 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#[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#[derive(Debug, Clone, Serialize, Deserialize)]
37pub struct RouteDecision {
38 pub target: RoutingTarget,
39 pub difficulty: Difficulty,
40 pub tier: RouteTier,
41 #[serde(default, skip_serializing_if = "Option::is_none")]
43 pub reasoning: Option<String>,
44 #[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 Explicit,
55 Heuristic,
57 Semantic,
59}
60
61pub 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 pub fn policy_target(&self, d: Difficulty) -> RoutingTarget {
74 self.policy.choose(d)
75 }
76
77 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 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 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 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}