use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::core::config::{RoutingRules, parse_route_target};
use crate::core::intent_engine::{self, TaskClassification};
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct RoutingDecision {
pub requested_model: String,
pub routed_model: String,
pub routed_provider: Option<String>,
pub tier: String,
pub confidence: f64,
pub reasoning: String,
pub model_changed: bool,
pub estimated_cost_ratio: Option<f64>,
}
pub fn route(body: &Value, rules: &RoutingRules) -> Option<RoutingDecision> {
if !rules.is_active() || rules.tiers.is_empty() {
return None;
}
let requested_model = extract_model(body)?;
if rules.aliases.contains_key(&requested_model) {
return None;
}
let messages = body.get("messages")?;
let last_user_content = extract_last_user_content(messages)?;
let classification = intent_engine::classify(&last_user_content);
let route = intent_engine::route_intent(&last_user_content, &classification);
let tier_key = route.model_tier.as_str();
let target = rules.tiers.get(tier_key)?;
if target.is_empty() {
return Some(RoutingDecision {
requested_model: requested_model.clone(),
routed_model: requested_model,
routed_provider: None,
tier: tier_key.to_string(),
confidence: route.confidence,
reasoning: route.reasoning,
model_changed: false,
estimated_cost_ratio: None,
});
}
let (provider, model) = parse_route_target(target)?;
let cost_ratio = estimate_cost_ratio(&requested_model, model);
Some(RoutingDecision {
requested_model: requested_model.clone(),
routed_model: model.to_string(),
routed_provider: provider.map(str::to_string),
tier: tier_key.to_string(),
confidence: route.confidence,
reasoning: route.reasoning,
model_changed: requested_model != model,
estimated_cost_ratio: cost_ratio,
})
}
pub fn apply_decision(body: &mut Value, decision: &RoutingDecision) {
if !decision.model_changed {
return;
}
if let Some(obj) = body.as_object_mut() {
obj.insert(
"model".to_string(),
Value::String(decision.routed_model.clone()),
);
}
}
pub fn classify_only(body: &Value) -> Option<(TaskClassification, intent_engine::IntentRoute)> {
let messages = body.get("messages")?;
let content = extract_last_user_content(messages)?;
let classification = intent_engine::classify(&content);
let route = intent_engine::route_intent(&content, &classification);
Some((classification, route))
}
fn extract_model(body: &Value) -> Option<String> {
body.get("model")
.and_then(Value::as_str)
.map(str::to_string)
}
fn extract_last_user_content(messages: &Value) -> Option<String> {
let arr = messages.as_array()?;
for msg in arr.iter().rev() {
let role = msg.get("role").and_then(Value::as_str)?;
if role != "user" {
continue;
}
match msg.get("content") {
Some(Value::String(s)) => return Some(s.clone()),
Some(Value::Array(blocks)) => {
let text: String = blocks
.iter()
.filter_map(|b| {
if b.get("type").and_then(Value::as_str) == Some("text") {
b.get("text").and_then(Value::as_str)
} else {
None
}
})
.collect::<Vec<_>>()
.join("\n");
if !text.is_empty() {
return Some(text);
}
}
_ => {}
}
}
None
}
fn estimate_cost_ratio(original: &str, routed: &str) -> Option<f64> {
let orig_cost = model_cost_tier(original)?;
let routed_cost = model_cost_tier(routed)?;
if orig_cost == 0.0 {
return None;
}
Some(routed_cost / orig_cost)
}
fn model_cost_tier(model: &str) -> Option<f64> {
let m = model.to_lowercase();
if m.contains("nano") || m.contains("gpt-4.1-nano") {
Some(0.04)
} else if m.contains("haiku")
|| m.contains("flash")
|| m.contains("4o-mini")
|| m.contains("4.1-mini")
|| m.contains("deepseek")
{
Some(0.2)
} else if m.contains("opus") || m.contains("o3-pro") || m.contains("o1-pro") {
Some(5.0)
} else if m.contains("o3") || m.contains("o1") || m.contains("gpt-5") {
Some(2.5)
} else if m.contains("sonnet")
|| m.contains("gpt-4o")
|| m.contains("gemini-2.5-pro")
|| m.contains("gemini-2.0-pro")
{
Some(1.0)
} else {
None
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
use std::collections::BTreeMap;
fn active_rules(tiers: &[(&str, &str)]) -> RoutingRules {
RoutingRules {
enabled: Some(true),
aliases: BTreeMap::new(),
tiers: tiers
.iter()
.map(|(k, v)| (k.to_string(), v.to_string()))
.collect(),
}
}
fn request_body(model: &str, user_message: &str) -> Value {
json!({
"model": model,
"messages": [
{"role": "user", "content": user_message}
]
})
}
#[test]
fn inactive_routing_returns_none() {
let rules = RoutingRules::default();
let body = request_body("claude-sonnet-4", "fix the bug");
assert!(route(&body, &rules).is_none());
}
#[test]
fn empty_tiers_returns_none() {
let rules = RoutingRules {
enabled: Some(true),
aliases: BTreeMap::new(),
tiers: BTreeMap::new(),
};
let body = request_body("claude-sonnet-4", "fix the bug");
assert!(route(&body, &rules).is_none());
}
#[test]
fn no_model_field_returns_none() {
let body = json!({"messages": [{"role": "user", "content": "hi"}]});
let rules = active_rules(&[("fast", "x")]);
assert!(route(&body, &rules).is_none());
}
#[test]
fn no_user_messages_returns_none() {
let body = json!({
"model": "claude-sonnet-4",
"messages": [{"role": "assistant", "content": "hello"}]
});
let rules = active_rules(&[("fast", "x")]);
assert!(route(&body, &rules).is_none());
}
#[test]
fn alias_takes_priority_over_tier() {
let mut rules = active_rules(&[("fast", "anthropic:claude-haiku-4-5")]);
rules
.aliases
.insert("my-model".to_string(), "openai:gpt-4o".to_string());
let body = request_body("my-model", "explain the code");
assert!(route(&body, &rules).is_none(), "alias exempts from tiers");
}
#[test]
fn fast_tier_downgrades_explore_queries() {
let rules = active_rules(&[("fast", "anthropic:claude-haiku-4-5")]);
let body = request_body("claude-sonnet-4", "explain how the cache works");
let decision = route(&body, &rules).expect("should route");
assert_eq!(decision.requested_model, "claude-sonnet-4");
assert_eq!(decision.routed_model, "claude-haiku-4-5");
assert_eq!(decision.routed_provider.as_deref(), Some("anthropic"));
assert_eq!(decision.tier, "fast");
assert!(decision.model_changed);
assert!(decision.confidence > 0.5);
}
#[test]
fn premium_tier_upgrades_generation_tasks() {
let rules = active_rules(&[("premium", "anthropic:claude-opus-4")]);
let body = request_body("claude-sonnet-4", "implement a new auth module with JWT");
let decision = route(&body, &rules).expect("should route");
assert_eq!(decision.routed_model, "claude-opus-4");
assert_eq!(decision.tier, "premium");
assert!(decision.model_changed);
}
#[test]
fn standard_tier_for_fixbug() {
let rules = active_rules(&[
("fast", "claude-haiku-4-5"),
("standard", "claude-sonnet-4"),
("premium", "claude-opus-4"),
]);
let body = request_body("claude-opus-4", "fix the bug in auth.rs");
let decision = route(&body, &rules).expect("should route");
assert_eq!(decision.tier, "standard");
assert_eq!(decision.routed_model, "claude-sonnet-4");
}
#[test]
fn missing_tier_key_is_passthrough() {
let rules = active_rules(&[("fast", "anthropic:claude-haiku-4-5")]);
let body = request_body("claude-sonnet-4", "fix the null pointer bug in auth.rs");
assert!(route(&body, &rules).is_none());
}
#[test]
fn empty_tier_target_keeps_model() {
let rules = active_rules(&[("fast", "")]);
let body = request_body("claude-sonnet-4", "explain and describe this function");
let decision = route(&body, &rules).expect("should route");
assert_eq!(decision.routed_model, "claude-sonnet-4");
assert!(!decision.model_changed);
assert_eq!(decision.tier, "fast");
}
#[test]
fn model_only_target_keeps_provider() {
let rules = active_rules(&[("fast", "claude-haiku-4-5")]);
let body = request_body("claude-sonnet-4", "explain what this function does");
let decision = route(&body, &rules).expect("should route");
assert_eq!(decision.routed_model, "claude-haiku-4-5");
assert_eq!(decision.routed_provider, None, "no provider override");
assert!(decision.model_changed);
}
#[test]
fn apply_decision_rewrites_body() {
let mut body = request_body("claude-sonnet-4", "explain");
let decision = RoutingDecision {
requested_model: "claude-sonnet-4".into(),
routed_model: "claude-haiku-4-5".into(),
routed_provider: Some("anthropic".into()),
tier: "fast".into(),
confidence: 0.85,
reasoning: "explore(what) + low complexity -> fast".into(),
model_changed: true,
estimated_cost_ratio: Some(0.2),
};
apply_decision(&mut body, &decision);
assert_eq!(body["model"], "claude-haiku-4-5");
}
#[test]
fn apply_decision_noop_when_unchanged() {
let mut body = request_body("claude-sonnet-4", "explain");
let decision = RoutingDecision {
requested_model: "claude-sonnet-4".into(),
routed_model: "claude-sonnet-4".into(),
routed_provider: None,
tier: "standard".into(),
confidence: 0.8,
reasoning: "fix_bug(how) -> standard".into(),
model_changed: false,
estimated_cost_ratio: None,
};
apply_decision(&mut body, &decision);
assert_eq!(body["model"], "claude-sonnet-4");
}
#[test]
fn cost_ratio_estimates_downgrade_savings() {
let rules = active_rules(&[("fast", "claude-haiku-4-5")]);
let body = request_body("claude-sonnet-4", "explain what this module does");
let decision = route(&body, &rules).expect("should route");
let ratio = decision.estimated_cost_ratio.expect("known models");
assert!(ratio < 1.0, "haiku cheaper than sonnet: {ratio}");
assert!(ratio > 0.0);
}
#[test]
fn cost_ratio_none_for_unknown_models() {
let rules = active_rules(&[("fast", "custom-local-model")]);
let body = request_body("claude-sonnet-4", "explain what this code does");
let decision = route(&body, &rules).expect("should route");
assert_eq!(decision.estimated_cost_ratio, None);
}
#[test]
fn cost_tiers_are_ordered() {
assert!(
model_cost_tier("claude-opus-4").unwrap() > model_cost_tier("claude-sonnet-4").unwrap()
);
assert!(
model_cost_tier("claude-sonnet-4").unwrap()
> model_cost_tier("claude-haiku-4-5").unwrap()
);
assert!(model_cost_tier("gpt-4o").unwrap() > model_cost_tier("gpt-4o-mini").unwrap());
}
#[test]
fn multipart_content_blocks_extracted() {
let body = json!({
"model": "claude-sonnet-4",
"messages": [{
"role": "user",
"content": [
{"type": "text", "text": "explain how the session cache works"},
{"type": "image", "source": {"type": "base64"}}
]
}]
});
let rules = active_rules(&[("fast", "claude-haiku-4-5")]);
let decision = route(&body, &rules).expect("should extract text blocks");
assert_eq!(decision.tier, "fast");
}
#[test]
fn decision_is_deterministic() {
let rules = active_rules(&[
("fast", "claude-haiku-4-5"),
("standard", "claude-sonnet-4"),
("premium", "claude-opus-4"),
]);
let body = request_body("claude-sonnet-4", "explain how the proxy routing works");
let d1 = route(&body, &rules);
let d2 = route(&body, &rules);
assert_eq!(d1, d2, "routing must be deterministic (#498)");
}
}