mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
feat: add model directive management
This commit is contained in:
@@ -63,23 +63,44 @@ pub fn auth_constraints_allow_model(
|
||||
constraints: Option<&SchedulerAuthConstraints>,
|
||||
requested_model_name: &str,
|
||||
resolved_global_model_name: &str,
|
||||
) -> bool {
|
||||
auth_constraints_allow_model_with_model_directives(
|
||||
constraints,
|
||||
requested_model_name,
|
||||
resolved_global_model_name,
|
||||
false,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn auth_constraints_allow_model_with_model_directives(
|
||||
constraints: Option<&SchedulerAuthConstraints>,
|
||||
requested_model_name: &str,
|
||||
resolved_global_model_name: &str,
|
||||
enable_model_directives: bool,
|
||||
) -> bool {
|
||||
let Some(allowed) = constraints.and_then(|constraints| constraints.allowed_models.as_deref())
|
||||
else {
|
||||
return true;
|
||||
};
|
||||
|
||||
allowed
|
||||
.iter()
|
||||
.any(|value| value == requested_model_name || value == resolved_global_model_name)
|
||||
let base_model = enable_model_directives
|
||||
.then(|| aether_ai_formats::model_directive_base_model(requested_model_name))
|
||||
.flatten();
|
||||
allowed.iter().any(|value| {
|
||||
value == requested_model_name
|
||||
|| value == resolved_global_model_name
|
||||
|| base_model
|
||||
.as_ref()
|
||||
.is_some_and(|base_model| value == base_model)
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
api_format_matches_allowed_value, auth_constraints_allow_api_format,
|
||||
auth_constraints_allow_model, auth_constraints_allow_provider,
|
||||
provider_matches_allowed_value, SchedulerAuthConstraints,
|
||||
auth_constraints_allow_model, auth_constraints_allow_model_with_model_directives,
|
||||
auth_constraints_allow_provider, provider_matches_allowed_value, SchedulerAuthConstraints,
|
||||
};
|
||||
|
||||
fn sample_constraints() -> SchedulerAuthConstraints {
|
||||
@@ -200,6 +221,23 @@ mod tests {
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn model_directive_base_model_requires_explicit_enablement() {
|
||||
let constraints = sample_constraints();
|
||||
|
||||
assert!(!auth_constraints_allow_model(
|
||||
Some(&constraints),
|
||||
"gpt-5-high",
|
||||
"gpt-5-high"
|
||||
));
|
||||
assert!(auth_constraints_allow_model_with_model_directives(
|
||||
Some(&constraints),
|
||||
"gpt-5-high",
|
||||
"gpt-5-high",
|
||||
true
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn api_format_allowed_value_matches_current_signatures_only() {
|
||||
assert!(api_format_matches_allowed_value(
|
||||
|
||||
@@ -9,6 +9,20 @@ use super::types::{
|
||||
|
||||
pub fn enumerate_minimal_candidate_selection(
|
||||
input: EnumerateMinimalCandidateSelectionInput<'_>,
|
||||
) -> Result<Vec<SchedulerMinimalCandidateSelectionCandidate>, DataLayerError> {
|
||||
enumerate_minimal_candidate_selection_inner(input, false)
|
||||
}
|
||||
|
||||
pub fn enumerate_minimal_candidate_selection_with_model_directives(
|
||||
input: EnumerateMinimalCandidateSelectionInput<'_>,
|
||||
enable_model_directives: bool,
|
||||
) -> Result<Vec<SchedulerMinimalCandidateSelectionCandidate>, DataLayerError> {
|
||||
enumerate_minimal_candidate_selection_inner(input, enable_model_directives)
|
||||
}
|
||||
|
||||
fn enumerate_minimal_candidate_selection_inner(
|
||||
input: EnumerateMinimalCandidateSelectionInput<'_>,
|
||||
enable_model_directives: bool,
|
||||
) -> Result<Vec<SchedulerMinimalCandidateSelectionCandidate>, DataLayerError> {
|
||||
let EnumerateMinimalCandidateSelectionInput {
|
||||
rows,
|
||||
@@ -26,10 +40,11 @@ pub fn enumerate_minimal_candidate_selection(
|
||||
if !crate::auth_constraints_allow_api_format(auth_constraints, normalized_api_format) {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
if !crate::auth_constraints_allow_model(
|
||||
if !crate::auth_constraints_allow_model_with_model_directives(
|
||||
auth_constraints,
|
||||
requested_model_name,
|
||||
resolved_global_model_name,
|
||||
enable_model_directives,
|
||||
) {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
@@ -48,7 +63,12 @@ pub fn enumerate_minimal_candidate_selection(
|
||||
continue;
|
||||
}
|
||||
let Some((selected_provider_model_name, mapping_matched_model)) =
|
||||
crate::resolve_provider_model_name(&row, requested_model_name, normalized_api_format)
|
||||
crate::resolve_provider_model_name_with_model_directives(
|
||||
&row,
|
||||
requested_model_name,
|
||||
normalized_api_format,
|
||||
enable_model_directives,
|
||||
)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
@@ -8,6 +8,7 @@ pub use capability::{
|
||||
};
|
||||
pub use enumeration::{
|
||||
collect_global_model_names_for_required_capability, enumerate_minimal_candidate_selection,
|
||||
enumerate_minimal_candidate_selection_with_model_directives,
|
||||
};
|
||||
pub use selectability::{
|
||||
auth_api_key_concurrency_limit_reached, candidate_is_selectable_with_runtime_state,
|
||||
|
||||
@@ -13,13 +13,14 @@ pub use affinity::{
|
||||
};
|
||||
pub use auth::{
|
||||
api_format_matches_allowed_value, auth_constraints_allow_api_format,
|
||||
auth_constraints_allow_model, auth_constraints_allow_provider, provider_matches_allowed_value,
|
||||
SchedulerAuthConstraints,
|
||||
auth_constraints_allow_model, auth_constraints_allow_model_with_model_directives,
|
||||
auth_constraints_allow_provider, provider_matches_allowed_value, SchedulerAuthConstraints,
|
||||
};
|
||||
pub use candidate::{
|
||||
auth_api_key_concurrency_limit_reached, candidate_is_selectable_with_runtime_state,
|
||||
candidate_runtime_skip_reason_with_state, candidate_supports_required_capability,
|
||||
collect_global_model_names_for_required_capability, enumerate_minimal_candidate_selection,
|
||||
enumerate_minimal_candidate_selection_with_model_directives,
|
||||
requested_capability_priority_for_candidate, CandidateRuntimeSelectabilityInput,
|
||||
EnumerateMinimalCandidateSelectionInput, SchedulerMinimalCandidateSelectionCandidate,
|
||||
SchedulerPriorityMode,
|
||||
@@ -35,8 +36,11 @@ pub use health::{
|
||||
};
|
||||
pub use model::{
|
||||
candidate_model_names, extract_global_priority_for_format, matches_model_mapping,
|
||||
normalize_api_format, resolve_provider_model_name, resolve_requested_global_model_name,
|
||||
row_supports_requested_model, row_supports_required_capability, select_provider_model_name,
|
||||
normalize_api_format, resolve_provider_model_name,
|
||||
resolve_provider_model_name_with_model_directives, resolve_requested_global_model_name,
|
||||
resolve_requested_global_model_name_with_model_directives, row_supports_requested_model,
|
||||
row_supports_requested_model_with_model_directives, row_supports_required_capability,
|
||||
select_provider_model_name,
|
||||
};
|
||||
pub use provider::{build_provider_concurrent_limit_map, should_skip_provider_quota};
|
||||
pub use ranking::{
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
use std::borrow::Cow;
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
@@ -11,39 +12,79 @@ pub fn resolve_requested_global_model_name(
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
) -> Option<String> {
|
||||
resolve_global_model_name_by(rows, |row| row.global_model_name == requested_model_name)
|
||||
.or_else(|| {
|
||||
resolve_global_model_name_by(rows, |row| {
|
||||
row.model_provider_model_name == requested_model_name
|
||||
})
|
||||
})
|
||||
.or_else(|| {
|
||||
resolve_global_model_name_by(rows, |row| {
|
||||
row.model_provider_model_mappings
|
||||
.as_ref()
|
||||
.is_some_and(|mappings| {
|
||||
mappings.iter().any(|mapping| {
|
||||
mapping_scope_matches(mapping, api_format)
|
||||
&& mapping.name == requested_model_name
|
||||
resolve_requested_global_model_name_with_model_directives(
|
||||
rows,
|
||||
requested_model_name,
|
||||
api_format,
|
||||
false,
|
||||
)
|
||||
}
|
||||
|
||||
pub fn resolve_requested_global_model_name_with_model_directives(
|
||||
rows: &[StoredMinimalCandidateSelectionRow],
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
enable_model_directives: bool,
|
||||
) -> Option<String> {
|
||||
requested_model_name_candidates(requested_model_name, enable_model_directives).find_map(
|
||||
|requested_model_name| {
|
||||
let requested_model_name = requested_model_name.as_ref();
|
||||
resolve_global_model_name_by(rows, |row| row.global_model_name == requested_model_name)
|
||||
.or_else(|| {
|
||||
resolve_global_model_name_by(rows, |row| {
|
||||
row.model_provider_model_name == requested_model_name
|
||||
})
|
||||
})
|
||||
.or_else(|| {
|
||||
resolve_global_model_name_by(rows, |row| {
|
||||
row.model_provider_model_mappings
|
||||
.as_ref()
|
||||
.is_some_and(|mappings| {
|
||||
mappings.iter().any(|mapping| {
|
||||
mapping_scope_matches(mapping, api_format)
|
||||
&& mapping.name == requested_model_name
|
||||
})
|
||||
})
|
||||
})
|
||||
})
|
||||
.or_else(|| {
|
||||
resolve_global_model_name_by(rows, |row| {
|
||||
row.global_model_mappings.as_ref().is_some_and(|patterns| {
|
||||
patterns
|
||||
.iter()
|
||||
.any(|pattern| matches_model_mapping(pattern, requested_model_name))
|
||||
})
|
||||
})
|
||||
})
|
||||
})
|
||||
.or_else(|| {
|
||||
resolve_global_model_name_by(rows, |row| {
|
||||
row.global_model_mappings.as_ref().is_some_and(|patterns| {
|
||||
patterns
|
||||
.iter()
|
||||
.any(|pattern| matches_model_mapping(pattern, requested_model_name))
|
||||
})
|
||||
})
|
||||
})
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
pub fn row_supports_requested_model(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
row_supports_requested_model_with_model_directives(row, requested_model_name, api_format, false)
|
||||
}
|
||||
|
||||
pub fn row_supports_requested_model_with_model_directives(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
enable_model_directives: bool,
|
||||
) -> bool {
|
||||
requested_model_name_candidates(requested_model_name, enable_model_directives).any(
|
||||
|requested_model_name| {
|
||||
row_supports_requested_model_exact(row, requested_model_name.as_ref(), api_format)
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
fn row_supports_requested_model_exact(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
row.global_model_name == requested_model_name
|
||||
|| row.model_provider_model_name == requested_model_name
|
||||
@@ -87,6 +128,15 @@ pub fn resolve_provider_model_name(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
) -> Option<(String, Option<String>)> {
|
||||
resolve_provider_model_name_with_model_directives(row, requested_model_name, api_format, false)
|
||||
}
|
||||
|
||||
pub fn resolve_provider_model_name_with_model_directives(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
enable_model_directives: bool,
|
||||
) -> Option<(String, Option<String>)> {
|
||||
let selected_provider_model_name = select_provider_model_name(row, api_format);
|
||||
let Some(key_allowed_models) = row.key_allowed_models.as_ref() else {
|
||||
@@ -103,6 +153,16 @@ pub fn resolve_provider_model_name(
|
||||
return Some((selected_provider_model_name, None));
|
||||
}
|
||||
|
||||
if enable_model_directives {
|
||||
if let Some(base_model) =
|
||||
aether_ai_formats::model_directive_base_model(requested_model_name)
|
||||
{
|
||||
if key_allowed_models.iter().any(|value| value == &base_model) {
|
||||
return Some((selected_provider_model_name, Some(base_model)));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut sorted_allowed_models = key_allowed_models
|
||||
.iter()
|
||||
.map(String::as_str)
|
||||
@@ -302,9 +362,26 @@ fn api_format_matches(left: &str, right: &str) -> bool {
|
||||
normalize_api_format(left) == normalize_api_format(right)
|
||||
}
|
||||
|
||||
fn requested_model_name_candidates(
|
||||
requested_model_name: &str,
|
||||
enable_model_directives: bool,
|
||||
) -> impl Iterator<Item = Cow<'_, str>> {
|
||||
let requested_model_name = requested_model_name.trim();
|
||||
let base_model = enable_model_directives
|
||||
.then(|| aether_ai_formats::model_directive_base_model(requested_model_name))
|
||||
.flatten();
|
||||
std::iter::once(Cow::Borrowed(requested_model_name)).chain(base_model.map(Cow::Owned))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::matches_model_mapping;
|
||||
use super::{
|
||||
matches_model_mapping, resolve_provider_model_name,
|
||||
resolve_provider_model_name_with_model_directives,
|
||||
resolve_requested_global_model_name_with_model_directives, row_supports_requested_model,
|
||||
row_supports_requested_model_with_model_directives,
|
||||
};
|
||||
use aether_data_contracts::repository::candidate_selection::StoredMinimalCandidateSelectionRow;
|
||||
|
||||
#[test]
|
||||
fn model_mapping_match_is_case_insensitive() {
|
||||
@@ -322,4 +399,103 @@ mod tests {
|
||||
fn invalid_model_mapping_pattern_returns_false() {
|
||||
assert!(!matches_model_mapping("([a-z", "gpt-4o"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn model_directive_suffix_matches_base_model_as_fallback() {
|
||||
let row = sample_row("gpt-5.4", "gpt-5.4-upstream");
|
||||
|
||||
assert!(!row_supports_requested_model(
|
||||
&row,
|
||||
"gpt-5.4-xhigh",
|
||||
"openai:chat"
|
||||
));
|
||||
assert!(row_supports_requested_model_with_model_directives(
|
||||
&row,
|
||||
"gpt-5.4-xhigh",
|
||||
"openai:chat",
|
||||
true
|
||||
));
|
||||
assert_eq!(
|
||||
resolve_requested_global_model_name_with_model_directives(
|
||||
&[row],
|
||||
"gpt-5.4-xhigh",
|
||||
"openai:chat",
|
||||
true
|
||||
)
|
||||
.as_deref(),
|
||||
Some("gpt-5.4")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn model_directive_suffix_prefers_exact_model_before_base_fallback() {
|
||||
let exact = sample_row("gpt-5.4-high", "gpt-5.4-high-upstream");
|
||||
let base = sample_row("gpt-5.4", "gpt-5.4-upstream");
|
||||
|
||||
assert_eq!(
|
||||
resolve_requested_global_model_name_with_model_directives(
|
||||
&[base, exact],
|
||||
"gpt-5.4-high",
|
||||
"openai:chat",
|
||||
true
|
||||
)
|
||||
.as_deref(),
|
||||
Some("gpt-5.4-high")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn model_directive_base_model_satisfies_key_allowed_models() {
|
||||
let mut row = sample_row("gpt-5.4", "gpt-5.4-upstream");
|
||||
row.key_allowed_models = Some(vec!["gpt-5.4".to_string()]);
|
||||
|
||||
assert!(resolve_provider_model_name(&row, "gpt-5.4-max", "openai:chat").is_none());
|
||||
let resolved = resolve_provider_model_name_with_model_directives(
|
||||
&row,
|
||||
"gpt-5.4-max",
|
||||
"openai:chat",
|
||||
true,
|
||||
)
|
||||
.expect("base model should satisfy key allowed models");
|
||||
|
||||
assert_eq!(resolved.0, "gpt-5.4-upstream");
|
||||
assert_eq!(resolved.1.as_deref(), Some("gpt-5.4"));
|
||||
}
|
||||
|
||||
fn sample_row(
|
||||
global_model_name: &str,
|
||||
model_provider_model_name: &str,
|
||||
) -> StoredMinimalCandidateSelectionRow {
|
||||
StoredMinimalCandidateSelectionRow {
|
||||
provider_id: "provider-1".to_string(),
|
||||
provider_name: "Provider".to_string(),
|
||||
provider_type: "openai".to_string(),
|
||||
provider_priority: 0,
|
||||
provider_is_active: true,
|
||||
endpoint_id: "endpoint-1".to_string(),
|
||||
endpoint_api_format: "openai:chat".to_string(),
|
||||
endpoint_api_family: None,
|
||||
endpoint_kind: None,
|
||||
endpoint_is_active: true,
|
||||
key_id: "key-1".to_string(),
|
||||
key_name: "Key".to_string(),
|
||||
key_auth_type: "api_key".to_string(),
|
||||
key_is_active: true,
|
||||
key_api_formats: None,
|
||||
key_allowed_models: None,
|
||||
key_capabilities: None,
|
||||
key_internal_priority: 0,
|
||||
key_global_priority_by_format: None,
|
||||
model_id: format!("model-{global_model_name}"),
|
||||
global_model_id: format!("global-{global_model_name}"),
|
||||
global_model_name: global_model_name.to_string(),
|
||||
global_model_mappings: None,
|
||||
global_model_supports_streaming: Some(true),
|
||||
model_provider_model_name: model_provider_model_name.to_string(),
|
||||
model_provider_model_mappings: None,
|
||||
model_supports_streaming: Some(true),
|
||||
model_is_active: true,
|
||||
model_is_available: true,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user