mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 18:59:50 +08:00
feat: add selectable routing groups and composite billing
Support per-model provider enablement and compact model editing. Capture request-time billing factors, charge customer costs separately, and preserve historical statistics without backfills.
This commit is contained in:
@@ -132,6 +132,86 @@ fn is_false(value: &bool) -> bool {
|
||||
mod execution_policy_tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn group_visibility_is_opt_in_and_round_trips_without_losing_policy() {
|
||||
let legacy: RoutingGroupConfig = serde_json::from_str("{}").unwrap();
|
||||
assert!(!legacy.user_visible);
|
||||
assert!(!RoutingGroupConfig::default().user_visible);
|
||||
for user_visible in [false, true] {
|
||||
let config: RoutingGroupConfig = serde_json::from_value(serde_json::json!({
|
||||
"user_visible": user_visible,
|
||||
"billing_multiplier": 0.5,
|
||||
"disabled_providers": ["private-provider"],
|
||||
"default_policy": { "scheduling_mode": "fixed_order" }
|
||||
}))
|
||||
.unwrap();
|
||||
let encoded = serde_json::to_value(&config).unwrap();
|
||||
assert_eq!(encoded["user_visible"], user_visible);
|
||||
assert_eq!(config.billing_multiplier, 0.5);
|
||||
assert_eq!(config.disabled_providers, ["private-provider"]);
|
||||
assert_eq!(
|
||||
serde_json::from_value::<RoutingGroupConfig>(encoded).unwrap(),
|
||||
config
|
||||
);
|
||||
}
|
||||
for invalid in [
|
||||
serde_json::json!(null),
|
||||
serde_json::json!("true"),
|
||||
serde_json::json!(1),
|
||||
] {
|
||||
assert!(
|
||||
serde_json::from_value::<RoutingGroupConfig>(serde_json::json!({
|
||||
"user_visible": invalid
|
||||
}))
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn group_billing_multiplier_defaults_to_one_and_rejects_invalid_values() {
|
||||
let legacy: RoutingGroupConfig = serde_json::from_str("{}").unwrap();
|
||||
assert_eq!(legacy.billing_multiplier, 1.0);
|
||||
assert_eq!(RoutingGroupConfig::default().billing_multiplier, 1.0);
|
||||
for multiplier in [0.0, 0.25, 1.0, 2.5] {
|
||||
let config: RoutingGroupConfig = serde_json::from_value(serde_json::json!({
|
||||
"billing_multiplier": multiplier
|
||||
}))
|
||||
.unwrap();
|
||||
crate::validate_routing_group_config(&config).unwrap();
|
||||
assert_eq!(config.billing_multiplier, multiplier);
|
||||
assert_eq!(
|
||||
serde_json::from_value::<RoutingGroupConfig>(
|
||||
serde_json::to_value(&config).unwrap()
|
||||
)
|
||||
.unwrap(),
|
||||
config
|
||||
);
|
||||
}
|
||||
for multiplier in [-1.0, f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
|
||||
let config = RoutingGroupConfig {
|
||||
billing_multiplier: multiplier,
|
||||
..RoutingGroupConfig::default()
|
||||
};
|
||||
assert!(matches!(
|
||||
crate::validate_routing_group_config(&config),
|
||||
Err(crate::RoutingValidationError::InvalidBillingMultiplier)
|
||||
));
|
||||
}
|
||||
for value in [
|
||||
serde_json::json!(null),
|
||||
serde_json::json!("2"),
|
||||
serde_json::json!(false),
|
||||
] {
|
||||
assert!(
|
||||
serde_json::from_value::<RoutingGroupConfig>(serde_json::json!({
|
||||
"billing_multiplier": value
|
||||
}))
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn routing_failover_configuration_round_trips_and_validates() {
|
||||
let config: RoutingGroupConfig = serde_json::from_value(serde_json::json!({
|
||||
@@ -188,6 +268,11 @@ pub struct RoutingModelPolicy {
|
||||
pub allowed_providers: Vec<String>,
|
||||
#[serde(default)]
|
||||
pub allowed_keys: Vec<String>,
|
||||
/// Per-model provider enablement. A `false` value adds a provider to this
|
||||
/// model's exclusions and `true` removes an inherited exclusion, including
|
||||
/// one from the legacy group-wide `disabled_providers` baseline.
|
||||
#[serde(default)]
|
||||
pub provider_enabled_overrides: BTreeMap<String, bool>,
|
||||
#[serde(default)]
|
||||
pub provider_priority_overrides: BTreeMap<String, i32>,
|
||||
#[serde(default)]
|
||||
@@ -222,10 +307,17 @@ pub struct RoutingRule {
|
||||
pub stop_processing: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct RoutingGroupConfig {
|
||||
/// Providers excluded from every model in this group, including providers
|
||||
/// otherwise selected by model policies or routing rules.
|
||||
/// Whether authenticated users can discover and explicitly select this
|
||||
/// group. Private bindings and automatic defaults remain independent.
|
||||
#[serde(default)]
|
||||
pub user_visible: bool,
|
||||
/// Group-wide billing multiplier, snapshotted when a request is routed.
|
||||
#[serde(default = "default_billing_multiplier")]
|
||||
pub billing_multiplier: f64,
|
||||
/// Legacy provider exclusion baseline for the group. Explicit per-model
|
||||
/// enablement overrides may change it; allowlists and rules cannot.
|
||||
#[serde(default)]
|
||||
pub disabled_providers: Vec<String>,
|
||||
/// The default policy is global for the selected strategy group. Model
|
||||
@@ -238,6 +330,23 @@ pub struct RoutingGroupConfig {
|
||||
pub rules: Vec<RoutingRule>,
|
||||
}
|
||||
|
||||
pub(crate) fn default_billing_multiplier() -> f64 {
|
||||
1.0
|
||||
}
|
||||
|
||||
impl Default for RoutingGroupConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
user_visible: false,
|
||||
billing_multiplier: default_billing_multiplier(),
|
||||
disabled_providers: Vec::new(),
|
||||
default_policy: RoutingDefaultPolicy::default(),
|
||||
model_policies: Vec::new(),
|
||||
rules: Vec::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RoutingGroupRecord {
|
||||
pub id: String,
|
||||
|
||||
@@ -47,10 +47,15 @@ pub struct MatchedRoutingRule {
|
||||
pub phase: RoutingRulePhase,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct ResolvedRoutingPolicy {
|
||||
#[serde(default = "crate::model::default_billing_multiplier")]
|
||||
pub billing_multiplier: f64,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub group_id: Option<String>,
|
||||
/// Display name captured alongside the selected group by the gateway.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub group_name: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub group_version: Option<i64>,
|
||||
pub selection_source: String,
|
||||
@@ -80,7 +85,9 @@ pub fn resolve_routing_policy(
|
||||
.map_err(|error| RoutingPolicyError::InvalidConfig(error.to_string()))?;
|
||||
|
||||
let mut policy = ResolvedRoutingPolicy {
|
||||
billing_multiplier: config.billing_multiplier,
|
||||
group_id: input.group_id.map(str::to_string),
|
||||
group_name: None,
|
||||
group_version: input.group_version,
|
||||
selection_source: input.selection_source.to_string(),
|
||||
requested_model: input.requested_model.to_string(),
|
||||
@@ -105,6 +112,11 @@ pub fn resolve_routing_policy(
|
||||
{
|
||||
apply_model_policy(&mut policy, model_policy);
|
||||
}
|
||||
for model_policy in
|
||||
matching_provider_enablement_policies(config, input.requested_model, input.resolved_model)
|
||||
{
|
||||
apply_provider_enable_overrides(&mut policy, model_policy);
|
||||
}
|
||||
|
||||
let condition_context = RoutingConditionContext {
|
||||
model: input.requested_model,
|
||||
@@ -284,6 +296,57 @@ fn matching_model_policies<'a>(
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn matching_provider_enablement_policies<'a>(
|
||||
config: &'a RoutingGroupConfig,
|
||||
requested_model: &str,
|
||||
resolved_model: &str,
|
||||
) -> Vec<&'a RoutingModelPolicy> {
|
||||
let mut matches = matching_model_policies(config, requested_model, resolved_model)
|
||||
.into_iter()
|
||||
.filter(|policy| !policy.provider_enabled_overrides.is_empty())
|
||||
.collect::<Vec<_>>();
|
||||
// Enablement is a layered exception map: broad defaults first, then
|
||||
// prefixes, then exact model entries. Other model-policy fields retain
|
||||
// their historical configured-order merge semantics.
|
||||
matches.sort_by_key(|policy| model_pattern_specificity(&policy.model));
|
||||
matches
|
||||
}
|
||||
|
||||
fn apply_provider_enable_overrides(
|
||||
policy: &mut ResolvedRoutingPolicy,
|
||||
model_policy: &RoutingModelPolicy,
|
||||
) {
|
||||
for (provider_id, enabled) in &model_policy.provider_enabled_overrides {
|
||||
if *enabled {
|
||||
policy
|
||||
.ranking_overlay
|
||||
.disabled_providers
|
||||
.retain(|disabled| disabled != provider_id);
|
||||
} else if !policy
|
||||
.ranking_overlay
|
||||
.disabled_providers
|
||||
.iter()
|
||||
.any(|disabled| disabled == provider_id)
|
||||
{
|
||||
policy
|
||||
.ranking_overlay
|
||||
.disabled_providers
|
||||
.push(provider_id.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn model_pattern_specificity(pattern: &str) -> (u8, usize) {
|
||||
let pattern = pattern.trim();
|
||||
if pattern == "*" {
|
||||
(0, 0)
|
||||
} else if let Some(prefix) = pattern.strip_suffix('*') {
|
||||
(1, prefix.len())
|
||||
} else {
|
||||
(2, 0)
|
||||
}
|
||||
}
|
||||
|
||||
fn model_allowed(patterns: &[String], requested_model: &str) -> bool {
|
||||
patterns.is_empty()
|
||||
|| patterns
|
||||
@@ -405,7 +468,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn group_disabled_providers_apply_to_every_model_and_cannot_be_reenabled() {
|
||||
fn legacy_group_exclusions_cannot_be_bypassed_by_allowlists_or_rule_actions() {
|
||||
let config: RoutingGroupConfig = serde_json::from_value(json!({
|
||||
"disabled_providers": ["provider-disabled"],
|
||||
"model_policies": [{
|
||||
@@ -478,6 +541,106 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn model_provider_enablement_is_scoped_and_specific_overrides_win() {
|
||||
let mut config: RoutingGroupConfig = serde_json::from_value(json!({
|
||||
"disabled_providers": ["provider-root"],
|
||||
"model_policies": [
|
||||
{
|
||||
"model": "*",
|
||||
"provider_enabled_overrides": {
|
||||
"provider-model": false,
|
||||
"provider-specific": false
|
||||
}
|
||||
},
|
||||
{
|
||||
"model": "model-*",
|
||||
"provider_enabled_overrides": {
|
||||
"provider-model": true,
|
||||
"provider-specific": true
|
||||
}
|
||||
},
|
||||
{
|
||||
"model": "model-exact",
|
||||
"provider_enabled_overrides": {
|
||||
"provider-specific": false,
|
||||
"provider-exact": true,
|
||||
"provider-root": true
|
||||
}
|
||||
}
|
||||
]
|
||||
}))
|
||||
.unwrap();
|
||||
|
||||
// Persisted order need not put broad defaults first. Only the new
|
||||
// enablement map follows specificity; priorities retain their old order.
|
||||
config.model_policies.reverse();
|
||||
config.model_policies[0]
|
||||
.provider_priority_overrides
|
||||
.insert("provider-model".into(), 1);
|
||||
config.model_policies[2]
|
||||
.provider_priority_overrides
|
||||
.insert("provider-model".into(), 9);
|
||||
|
||||
let for_model = |model: &str| {
|
||||
resolve_routing_policy(
|
||||
&config,
|
||||
RoutingPolicyInput {
|
||||
group_id: Some("group-1"),
|
||||
group_version: Some(1),
|
||||
selection_source: "test",
|
||||
requested_model: model,
|
||||
resolved_model: model,
|
||||
api_format: "openai:chat",
|
||||
user_id: None,
|
||||
api_key_id: None,
|
||||
headers: &json!({}),
|
||||
body: &json!({}),
|
||||
phase: RoutingRulePhase::ClientRequest,
|
||||
},
|
||||
)
|
||||
.unwrap()
|
||||
};
|
||||
|
||||
let exact = for_model("model-exact");
|
||||
assert!(exact.ranking_overlay.provider_allowed("provider-model"));
|
||||
assert_eq!(
|
||||
exact.ranking_overlay.provider_priority_overrides["provider-model"],
|
||||
9
|
||||
);
|
||||
assert!(!exact.ranking_overlay.provider_allowed("provider-specific"));
|
||||
assert!(exact.ranking_overlay.provider_allowed("provider-exact"));
|
||||
assert!(exact.ranking_overlay.provider_allowed("provider-root"));
|
||||
|
||||
let wildcard_prefix = for_model("model-other");
|
||||
assert!(wildcard_prefix
|
||||
.ranking_overlay
|
||||
.provider_allowed("provider-model"));
|
||||
assert!(wildcard_prefix
|
||||
.ranking_overlay
|
||||
.provider_allowed("provider-specific"));
|
||||
assert!(!wildcard_prefix
|
||||
.ranking_overlay
|
||||
.provider_allowed("provider-root"));
|
||||
|
||||
let unrelated = for_model("other-model");
|
||||
assert!(!unrelated.ranking_overlay.provider_allowed("provider-model"));
|
||||
assert!(!unrelated
|
||||
.ranking_overlay
|
||||
.provider_allowed("provider-specific"));
|
||||
assert!(!unrelated.ranking_overlay.provider_allowed("provider-root"));
|
||||
|
||||
let encoded = serde_json::to_value(&config).unwrap();
|
||||
assert_eq!(
|
||||
encoded["model_policies"][2]["provider_enabled_overrides"]["provider-model"],
|
||||
false
|
||||
);
|
||||
assert_eq!(
|
||||
serde_json::from_value::<RoutingGroupConfig>(encoded).unwrap(),
|
||||
config
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn all_model_scheduling_and_rankings_apply_to_future_models() {
|
||||
let config: RoutingGroupConfig = serde_json::from_value(json!({
|
||||
@@ -599,6 +762,8 @@ mod tests {
|
||||
#[test]
|
||||
fn resolves_model_policy_and_matching_rule() {
|
||||
let config = RoutingGroupConfig {
|
||||
user_visible: false,
|
||||
billing_multiplier: 1.0,
|
||||
disabled_providers: vec![],
|
||||
default_policy: RoutingDefaultPolicy::default(),
|
||||
model_policies: vec![RoutingModelPolicy {
|
||||
@@ -668,6 +833,8 @@ mod tests {
|
||||
#[test]
|
||||
fn default_policy_applies_to_models_without_an_override() {
|
||||
let config = RoutingGroupConfig {
|
||||
user_visible: false,
|
||||
billing_multiplier: 2.5,
|
||||
disabled_providers: vec![],
|
||||
default_policy: RoutingDefaultPolicy {
|
||||
priority_mode: RoutingSetPriorityMode::GlobalKey,
|
||||
@@ -707,6 +874,7 @@ mod tests {
|
||||
assert_eq!(special.scheduling_mode, RoutingSchedulingMode::LoadBalance);
|
||||
assert!(special.keep_priority_on_conversion);
|
||||
assert_eq!(special.sticky_key_attempts, 3);
|
||||
assert_eq!(special.billing_multiplier, 2.5);
|
||||
assert_eq!(
|
||||
special.ranking_overlay.allowed_providers,
|
||||
vec!["provider-special"]
|
||||
@@ -741,6 +909,7 @@ mod tests {
|
||||
assert_eq!(ordinary.scheduling_mode, RoutingSchedulingMode::LoadBalance);
|
||||
assert!(ordinary.keep_priority_on_conversion);
|
||||
assert_eq!(ordinary.sticky_key_attempts, 3);
|
||||
assert_eq!(ordinary.billing_multiplier, 2.5);
|
||||
assert!(ordinary.ranking_overlay.allowed_providers.is_empty());
|
||||
assert!(ordinary.ranking_overlay.allowed_keys.is_empty());
|
||||
assert!(ordinary
|
||||
|
||||
@@ -13,7 +13,8 @@ pub enum CandidateKind {
|
||||
|
||||
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub struct RankingOverlay {
|
||||
/// Group-wide exclusions take precedence over every provider allowlist.
|
||||
/// Effective provider exclusions after model overrides. These take
|
||||
/// precedence over every provider allowlist.
|
||||
#[serde(default)]
|
||||
pub disabled_providers: Vec<String>,
|
||||
#[serde(default)]
|
||||
|
||||
@@ -55,11 +55,15 @@ pub struct RoutingRuntimeFacts {
|
||||
pub priority_mode: Option<RoutingSetPriorityMode>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub struct RoutingDecisionTrace {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub billing_multiplier: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub group_id: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub group_name: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub group_version: Option<i64>,
|
||||
pub selection_source: String,
|
||||
#[serde(default)]
|
||||
|
||||
@@ -28,6 +28,8 @@ const ROUTING_POOL_PRESETS: &[&str] = &[
|
||||
|
||||
#[derive(Debug, Error, Clone, PartialEq, Eq)]
|
||||
pub enum RoutingValidationError {
|
||||
#[error("routing group billing multiplier must be a non-negative finite number")]
|
||||
InvalidBillingMultiplier,
|
||||
#[error("routing failover rules are invalid: {0}")]
|
||||
InvalidFailoverRules(String),
|
||||
#[error("routing rule id is empty")]
|
||||
@@ -71,6 +73,9 @@ pub enum RoutingValidationError {
|
||||
pub fn validate_routing_group_config(
|
||||
config: &RoutingGroupConfig,
|
||||
) -> Result<(), RoutingValidationError> {
|
||||
if !config.billing_multiplier.is_finite() || config.billing_multiplier < 0.0 {
|
||||
return Err(RoutingValidationError::InvalidBillingMultiplier);
|
||||
}
|
||||
crate::validate_routing_failover_rules(&config.default_policy.execution_policy.failover_rules)
|
||||
.map_err(RoutingValidationError::InvalidFailoverRules)?;
|
||||
let mut rule_ids = BTreeSet::new();
|
||||
|
||||
Reference in New Issue
Block a user