use serde::{Deserialize, Serialize};
use tokenmiser_providers::ChatRequest;
pub mod dsl;
pub mod policy;
pub mod replay;
pub mod tier0;
pub mod tier1;
pub mod tier2;
pub use dsl::{PolicyEngine, RequestView};
pub use policy::{RoutingPolicy, RoutingTarget};
pub use replay::{replay, ReplayResult};
pub use tier0::tier0_difficulty;
pub use tier1::Tier1Classifier;
pub use tier2::{should_escalate, CascadeConfig, EscalateDecision};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum Difficulty {
Easy,
Medium,
Hard,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct RouteDecision {
pub target: RoutingTarget,
pub difficulty: Difficulty,
pub tier: RouteTier,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub reasoning: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub counterfactual_model: Option<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum RouteTier {
Explicit,
Heuristic,
Semantic,
}
pub struct Router {
policy: RoutingPolicy,
tier1: Option<Tier1Classifier>,
}
impl Router {
pub fn new(policy: RoutingPolicy, tier1: Option<Tier1Classifier>) -> Self {
Self { policy, tier1 }
}
pub fn policy_target(&self, d: Difficulty) -> RoutingTarget {
self.policy.choose(d)
}
pub fn decide(&self, req: &ChatRequest) -> RouteDecision {
let requested = req.model.as_str();
let auto = requested == "auto" || requested == "tokenmiser:auto";
if !auto {
let difficulty = tier0_difficulty(req);
return RouteDecision {
target: RoutingTarget::passthrough(requested),
difficulty,
tier: RouteTier::Explicit,
reasoning: None,
counterfactual_model: self.policy.frontier_for(Difficulty::Hard),
};
}
let t0 = tier0_difficulty(req);
let (difficulty, tier) = match (t0, &self.tier1) {
(Difficulty::Medium, Some(t1)) => (t1.classify(req), RouteTier::Semantic),
_ => (t0, RouteTier::Heuristic),
};
let target = self.policy.choose(difficulty);
let counterfactual = self.policy.frontier_for(Difficulty::Hard);
RouteDecision {
target,
difficulty,
tier,
reasoning: None,
counterfactual_model: counterfactual,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokenmiser_providers::ChatMessage;
fn user(model: &str, s: &str) -> ChatRequest {
ChatRequest {
model: model.into(),
messages: vec![ChatMessage {
role: "user".into(),
content: serde_json::Value::String(s.into()),
extra: Default::default(),
}],
temperature: None,
max_tokens: None,
top_p: None,
stream: None,
extra: Default::default(),
}
}
#[test]
fn explicit_model_request_is_passthrough() {
let router = Router::new(RoutingPolicy::default(), None);
let d = router.decide(&user("gpt-5", "anything"));
assert!(matches!(d.tier, RouteTier::Explicit));
assert_eq!(d.target.model, "gpt-5");
}
#[test]
fn auto_easy_routes_to_local() {
let router = Router::new(RoutingPolicy::default(), None);
let d = router.decide(&user("auto", "what is 2+2?"));
assert_eq!(d.difficulty, Difficulty::Easy);
assert!(d.target.model.contains("llama") || d.target.provider == "ollama");
}
#[test]
fn auto_hard_routes_to_frontier() {
let router = Router::new(RoutingPolicy::default(), None);
let d = router.decide(&user("auto", "refactor this auth middleware to JWT"));
assert_eq!(d.difficulty, Difficulty::Hard);
assert!(d.target.model.contains("opus") || d.target.model.contains("gpt"));
}
}