use serde_json::Value;
use crate::protocol::{Operation, OperationKey, OperationKind};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RuleImpl {
Passthrough,
TransformTo,
Local,
Unsupported,
}
#[derive(Debug, Clone)]
pub struct CompiledRoutingRule {
pub operation: Operation,
pub kind: OperationKind,
pub implementation: RuleImpl,
pub dest_operation: Option<Operation>,
pub dest_kind: Option<OperationKind>,
}
pub struct RoutingRuleSpec<'a> {
pub id: i64,
pub provider_id: i64,
pub operation: &'a str,
pub kind: &'a str,
pub implementation: &'a str,
pub dest_operation: Option<&'a str>,
pub dest_kind: Option<&'a str>,
pub sort_order: i64,
pub enabled: bool,
}
pub fn compile(rows: &[RoutingRuleSpec<'_>]) -> Vec<CompiledRoutingRule> {
let mut rows: Vec<&RoutingRuleSpec<'_>> = rows.iter().filter(|r| r.enabled).collect();
rows.sort_by_key(|r| r.sort_order);
let mut out = Vec::new();
for row in rows {
match compile_row(row) {
Some(rule) => out.push(rule),
None => tracing::warn!(
rule_id = row.id,
provider_id = row.provider_id,
"skipping unparsable routing rule"
),
}
}
out
}
fn compile_row(row: &RoutingRuleSpec<'_>) -> Option<CompiledRoutingRule> {
Some(CompiledRoutingRule {
operation: parse_str(row.operation)?,
kind: parse_str(row.kind)?,
implementation: match row.implementation {
"passthrough" => RuleImpl::Passthrough,
"transform_to" => RuleImpl::TransformTo,
"local" => RuleImpl::Local,
"unsupported" => RuleImpl::Unsupported,
_ => return None,
},
dest_operation: match row.dest_operation {
Some(s) => Some(parse_str(s)?),
None => None,
},
dest_kind: match row.dest_kind {
Some(s) => Some(parse_str(s)?),
None => None,
},
})
}
fn parse_str<T: serde::de::DeserializeOwned>(s: &str) -> Option<T> {
serde_json::from_value(Value::String(s.to_owned())).ok()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RoutingDecision {
Passthrough,
TransformTo(OperationKey),
Local,
Unsupported,
}
pub fn decide(rules: &[CompiledRoutingRule], source: OperationKey) -> RoutingDecision {
if let Some(rule) = rules
.iter()
.find(|r| r.operation == source.operation && r.kind == source.kind)
{
return match rule.implementation {
RuleImpl::Passthrough => RoutingDecision::Passthrough,
RuleImpl::Local => RoutingDecision::Local,
RuleImpl::Unsupported => RoutingDecision::Unsupported,
RuleImpl::TransformTo => match rule.dest_kind {
Some(kind) => RoutingDecision::TransformTo(OperationKey {
operation: rule.dest_operation.unwrap_or(source.operation),
kind,
}),
None => RoutingDecision::Unsupported,
},
};
}
RoutingDecision::Unsupported
}
#[cfg(test)]
mod tests {
use super::*;
use crate::protocol::ContentGenerationKind;
fn cg(op: Operation, k: ContentGenerationKind) -> OperationKey {
OperationKey::content_generation(op, k)
}
#[test]
fn no_rule_is_unsupported() {
use ContentGenerationKind as K;
let src = cg(Operation::GenerateContent, K::ClaudeMessages);
assert_eq!(decide(&[], src), RoutingDecision::Unsupported);
}
#[test]
fn explicit_rule_wins() {
let rule = CompiledRoutingRule {
operation: Operation::GenerateContent,
kind: OperationKind::ContentGeneration(ContentGenerationKind::ClaudeMessages),
implementation: RuleImpl::Unsupported,
dest_operation: None,
dest_kind: None,
};
let src = cg(
Operation::GenerateContent,
ContentGenerationKind::ClaudeMessages,
);
assert_eq!(decide(&[rule], src), RoutingDecision::Unsupported);
}
}