mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
fix: tighten model candidate matching scopes
This commit is contained in:
@@ -185,8 +185,10 @@ fn row_matches_requested_model(
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
row.global_model_name == requested_model_name
|
||||
|| row.model_provider_model_name == requested_model_name
|
||||
(row_has_available_provider_model(row, api_format)
|
||||
&& row.global_model_name == requested_model_name)
|
||||
|| (row_default_provider_model_name_available(row, api_format)
|
||||
&& row.model_provider_model_name == requested_model_name)
|
||||
|| row
|
||||
.model_provider_model_mappings
|
||||
.as_ref()
|
||||
@@ -205,6 +207,60 @@ fn row_matches_requested_model(
|
||||
})
|
||||
}
|
||||
|
||||
fn row_has_available_provider_model(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
row_mapping_matches_scope(row, api_format)
|
||||
|| row_default_provider_model_name_available(row, api_format)
|
||||
}
|
||||
|
||||
fn row_default_provider_model_name_available(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
let Some(mappings) = row.model_provider_model_mappings.as_ref() else {
|
||||
return true;
|
||||
};
|
||||
let mut has_explicit_default_mapping = false;
|
||||
for mapping in mappings {
|
||||
if mapping.name != row.model_provider_model_name {
|
||||
continue;
|
||||
}
|
||||
has_explicit_default_mapping = true;
|
||||
if mapping_scope_matches(mapping, row, api_format) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
!has_explicit_default_mapping
|
||||
}
|
||||
|
||||
fn row_mapping_matches_scope(row: &StoredMinimalCandidateSelectionRow, api_format: &str) -> bool {
|
||||
row.model_provider_model_mappings
|
||||
.as_ref()
|
||||
.is_some_and(|mappings| {
|
||||
mappings
|
||||
.iter()
|
||||
.any(|mapping| mapping_scope_matches(mapping, row, api_format))
|
||||
})
|
||||
}
|
||||
|
||||
fn mapping_scope_matches(
|
||||
mapping: &super::StoredProviderModelMapping,
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
mapping.api_formats.as_ref().is_none_or(|formats| {
|
||||
formats
|
||||
.iter()
|
||||
.any(|value| api_format_matches(value, api_format))
|
||||
}) && mapping.endpoint_ids.as_ref().is_none_or(|endpoint_ids| {
|
||||
endpoint_ids
|
||||
.iter()
|
||||
.any(|endpoint_id| endpoint_id == &row.endpoint_id)
|
||||
})
|
||||
}
|
||||
|
||||
fn key_auth_channel_matches(row: &StoredMinimalCandidateSelectionRow, api_format: &str) -> bool {
|
||||
let provider_type = row.provider_type.trim().to_ascii_lowercase();
|
||||
let auth_type = row.key_auth_type.trim().to_ascii_lowercase();
|
||||
@@ -244,7 +300,7 @@ mod tests {
|
||||
use super::InMemoryMinimalCandidateSelectionReadRepository;
|
||||
use crate::repository::candidate_selection::{
|
||||
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
|
||||
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsQuery,
|
||||
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
|
||||
StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
|
||||
@@ -310,14 +366,12 @@ mod tests {
|
||||
async fn filters_by_exact_api_format_and_requested_model_aliases() {
|
||||
let mut mapped = sample_row("provider-1", "openai:chat", "gpt-4.1", 10);
|
||||
mapped.model_provider_model_name = "provider-gpt-4.1".to_string();
|
||||
mapped.model_provider_model_mappings = Some(vec![
|
||||
crate::repository::candidate_selection::StoredProviderModelMapping {
|
||||
name: "alias-gpt-4.1".to_string(),
|
||||
priority: 0,
|
||||
api_formats: Some(vec!["openai:chat".to_string()]),
|
||||
endpoint_ids: None,
|
||||
},
|
||||
]);
|
||||
mapped.model_provider_model_mappings = Some(vec![StoredProviderModelMapping {
|
||||
name: "alias-gpt-4.1".to_string(),
|
||||
priority: 0,
|
||||
api_formats: Some(vec!["openai:chat".to_string()]),
|
||||
endpoint_ids: None,
|
||||
}]);
|
||||
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||
mapped,
|
||||
sample_row("provider-2", "openai:chat", "gpt-4.1-mini", 20),
|
||||
@@ -332,6 +386,42 @@ mod tests {
|
||||
assert_eq!(rows[0].provider_id, "provider-1");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn requested_model_filter_respects_endpoint_scoped_default_mapping() {
|
||||
let mut selected = sample_row("provider-1", "openai:chat", "deepseek-v4-pro", 10);
|
||||
selected.endpoint_id = "endpoint-openai".to_string();
|
||||
selected.model_provider_model_name = "deepseek-v4-pro".to_string();
|
||||
selected.model_provider_model_mappings = Some(vec![StoredProviderModelMapping {
|
||||
name: "deepseek-v4-pro".to_string(),
|
||||
priority: 1,
|
||||
api_formats: None,
|
||||
endpoint_ids: Some(vec!["endpoint-openai".to_string()]),
|
||||
}]);
|
||||
|
||||
let mut scoped_out = selected.clone();
|
||||
scoped_out.provider_id = "provider-2".to_string();
|
||||
scoped_out.endpoint_id = "endpoint-claude".to_string();
|
||||
scoped_out.endpoint_api_format = "claude:messages".to_string();
|
||||
scoped_out.key_id = "key-provider-2".to_string();
|
||||
scoped_out.key_api_formats = Some(vec!["claude:messages".to_string()]);
|
||||
|
||||
let repository =
|
||||
InMemoryMinimalCandidateSelectionReadRepository::seed(vec![scoped_out, selected]);
|
||||
|
||||
let rows = repository
|
||||
.list_for_exact_api_format_and_requested_model("claude:messages", "deepseek-v4-pro")
|
||||
.await
|
||||
.expect("list should succeed");
|
||||
assert!(rows.is_empty());
|
||||
|
||||
let rows = repository
|
||||
.list_for_exact_api_format_and_requested_model("openai:chat", "deepseek-v4-pro")
|
||||
.await
|
||||
.expect("list should succeed");
|
||||
assert_eq!(rows.len(), 1);
|
||||
assert_eq!(rows[0].endpoint_id, "endpoint-openai");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn requested_model_page_returns_requested_slice_only() {
|
||||
let mut rows = Vec::new();
|
||||
|
||||
@@ -313,26 +313,75 @@ fn row_matches_requested_model(
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
row.global_model_name == requested_model_name
|
||||
|| row.model_provider_model_name == requested_model_name
|
||||
(row_has_available_provider_model(row, api_format)
|
||||
&& row.global_model_name == requested_model_name)
|
||||
|| (row_default_provider_model_name_available(row, api_format)
|
||||
&& row.model_provider_model_name == requested_model_name)
|
||||
|| row
|
||||
.model_provider_model_mappings
|
||||
.as_ref()
|
||||
.is_some_and(|mappings| {
|
||||
mappings.iter().any(|mapping| {
|
||||
mapping.api_formats.as_ref().is_none_or(|formats| {
|
||||
formats
|
||||
.iter()
|
||||
.any(|value| api_format_matches(value, api_format))
|
||||
}) && mapping.endpoint_ids.as_ref().is_none_or(|endpoint_ids| {
|
||||
endpoint_ids
|
||||
.iter()
|
||||
.any(|endpoint_id| endpoint_id == &row.endpoint_id)
|
||||
}) && mapping.name == requested_model_name
|
||||
mapping_scope_matches(mapping, row, api_format)
|
||||
&& mapping.name == requested_model_name
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn row_has_available_provider_model(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
row_mapping_matches_scope(row, api_format)
|
||||
|| row_default_provider_model_name_available(row, api_format)
|
||||
}
|
||||
|
||||
fn row_default_provider_model_name_available(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
let Some(mappings) = row.model_provider_model_mappings.as_ref() else {
|
||||
return true;
|
||||
};
|
||||
let mut has_explicit_default_mapping = false;
|
||||
for mapping in mappings {
|
||||
if mapping.name != row.model_provider_model_name {
|
||||
continue;
|
||||
}
|
||||
has_explicit_default_mapping = true;
|
||||
if mapping_scope_matches(mapping, row, api_format) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
!has_explicit_default_mapping
|
||||
}
|
||||
|
||||
fn row_mapping_matches_scope(row: &StoredMinimalCandidateSelectionRow, api_format: &str) -> bool {
|
||||
row.model_provider_model_mappings
|
||||
.as_ref()
|
||||
.is_some_and(|mappings| {
|
||||
mappings
|
||||
.iter()
|
||||
.any(|mapping| mapping_scope_matches(mapping, row, api_format))
|
||||
})
|
||||
}
|
||||
|
||||
fn mapping_scope_matches(
|
||||
mapping: &super::StoredProviderModelMapping,
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
mapping.api_formats.as_ref().is_none_or(|formats| {
|
||||
formats
|
||||
.iter()
|
||||
.any(|value| api_format_matches(value, api_format))
|
||||
}) && mapping.endpoint_ids.as_ref().is_none_or(|endpoint_ids| {
|
||||
endpoint_ids
|
||||
.iter()
|
||||
.any(|endpoint_id| endpoint_id == &row.endpoint_id)
|
||||
})
|
||||
}
|
||||
|
||||
fn key_auth_channel_matches(row: &CandidateSelectionRow, api_format: &str) -> bool {
|
||||
let provider_type = row.row.provider_type.trim().to_ascii_lowercase();
|
||||
let auth_type = row.row.key_auth_type.trim().to_ascii_lowercase();
|
||||
|
||||
@@ -711,13 +711,110 @@ fn requested_model_selection_sql() -> String {
|
||||
.replace(
|
||||
"AND gm.name = $2",
|
||||
r#"AND (
|
||||
gm.name = $2
|
||||
OR m.provider_model_name = $2
|
||||
(
|
||||
gm.name = $2
|
||||
AND (
|
||||
m.provider_model_mappings IS NULL
|
||||
OR jsonb_typeof(m.provider_model_mappings) <> 'array'
|
||||
OR EXISTS (
|
||||
SELECT 1
|
||||
FROM jsonb_array_elements(
|
||||
CASE
|
||||
WHEN jsonb_typeof(m.provider_model_mappings) = 'array'
|
||||
THEN m.provider_model_mappings
|
||||
ELSE '[]'::jsonb
|
||||
END
|
||||
) AS mapping(value)
|
||||
WHERE (
|
||||
mapping.value -> 'api_formats' IS NULL
|
||||
OR jsonb_typeof(mapping.value -> 'api_formats') <> 'array'
|
||||
OR EXISTS (
|
||||
SELECT 1
|
||||
FROM jsonb_array_elements_text(mapping.value -> 'api_formats') AS fmt(value)
|
||||
WHERE LOWER(BTRIM(fmt.value)) = ANY($3::text[])
|
||||
)
|
||||
)
|
||||
AND (
|
||||
mapping.value -> 'endpoint_ids' IS NULL
|
||||
OR jsonb_typeof(mapping.value -> 'endpoint_ids') <> 'array'
|
||||
OR EXISTS (
|
||||
SELECT 1
|
||||
FROM jsonb_array_elements_text(mapping.value -> 'endpoint_ids') AS endpoint(value)
|
||||
WHERE endpoint.value = pe.id
|
||||
)
|
||||
)
|
||||
)
|
||||
OR NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM jsonb_array_elements(
|
||||
CASE
|
||||
WHEN jsonb_typeof(m.provider_model_mappings) = 'array'
|
||||
THEN m.provider_model_mappings
|
||||
ELSE '[]'::jsonb
|
||||
END
|
||||
) AS mapping(value)
|
||||
WHERE mapping.value ->> 'name' = m.provider_model_name
|
||||
)
|
||||
)
|
||||
)
|
||||
OR (
|
||||
m.provider_model_name = $2
|
||||
AND (
|
||||
m.provider_model_mappings IS NULL
|
||||
OR jsonb_typeof(m.provider_model_mappings) <> 'array'
|
||||
OR NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM jsonb_array_elements(
|
||||
CASE
|
||||
WHEN jsonb_typeof(m.provider_model_mappings) = 'array'
|
||||
THEN m.provider_model_mappings
|
||||
ELSE '[]'::jsonb
|
||||
END
|
||||
) AS mapping(value)
|
||||
WHERE mapping.value ->> 'name' = m.provider_model_name
|
||||
)
|
||||
OR EXISTS (
|
||||
SELECT 1
|
||||
FROM jsonb_array_elements(
|
||||
CASE
|
||||
WHEN jsonb_typeof(m.provider_model_mappings) = 'array'
|
||||
THEN m.provider_model_mappings
|
||||
ELSE '[]'::jsonb
|
||||
END
|
||||
) AS mapping(value)
|
||||
WHERE mapping.value ->> 'name' = m.provider_model_name
|
||||
AND (
|
||||
mapping.value -> 'api_formats' IS NULL
|
||||
OR jsonb_typeof(mapping.value -> 'api_formats') <> 'array'
|
||||
OR EXISTS (
|
||||
SELECT 1
|
||||
FROM jsonb_array_elements_text(mapping.value -> 'api_formats') AS fmt(value)
|
||||
WHERE LOWER(BTRIM(fmt.value)) = ANY($3::text[])
|
||||
)
|
||||
)
|
||||
AND (
|
||||
mapping.value -> 'endpoint_ids' IS NULL
|
||||
OR jsonb_typeof(mapping.value -> 'endpoint_ids') <> 'array'
|
||||
OR EXISTS (
|
||||
SELECT 1
|
||||
FROM jsonb_array_elements_text(mapping.value -> 'endpoint_ids') AS endpoint(value)
|
||||
WHERE endpoint.value = pe.id
|
||||
)
|
||||
)
|
||||
)
|
||||
)
|
||||
)
|
||||
OR (
|
||||
jsonb_typeof(m.provider_model_mappings) = 'array'
|
||||
AND EXISTS (
|
||||
SELECT 1
|
||||
FROM jsonb_array_elements(m.provider_model_mappings) AS mapping(value)
|
||||
FROM jsonb_array_elements(
|
||||
CASE
|
||||
WHEN jsonb_typeof(m.provider_model_mappings) = 'array'
|
||||
THEN m.provider_model_mappings
|
||||
ELSE '[]'::jsonb
|
||||
END
|
||||
) AS mapping(value)
|
||||
WHERE mapping.value ->> 'name' = $2
|
||||
AND (
|
||||
mapping.value -> 'api_formats' IS NULL
|
||||
@@ -1108,7 +1205,7 @@ mod tests {
|
||||
let sql = requested_model_selection_sql();
|
||||
|
||||
assert!(sql.contains("m.provider_model_name = $2"));
|
||||
assert!(sql.contains("jsonb_array_elements(m.provider_model_mappings)"));
|
||||
assert!(sql.contains("jsonb_array_elements("));
|
||||
assert!(!sql.contains("json_typeof(m.provider_model_mappings)"));
|
||||
assert!(!sql.contains("json_array_elements_text(gm.config -> 'model_mappings')"));
|
||||
assert!(sql.contains("ORDER BY\n global_model_name ASC,"));
|
||||
|
||||
@@ -313,26 +313,75 @@ fn row_matches_requested_model(
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
row.global_model_name == requested_model_name
|
||||
|| row.model_provider_model_name == requested_model_name
|
||||
(row_has_available_provider_model(row, api_format)
|
||||
&& row.global_model_name == requested_model_name)
|
||||
|| (row_default_provider_model_name_available(row, api_format)
|
||||
&& row.model_provider_model_name == requested_model_name)
|
||||
|| row
|
||||
.model_provider_model_mappings
|
||||
.as_ref()
|
||||
.is_some_and(|mappings| {
|
||||
mappings.iter().any(|mapping| {
|
||||
mapping.api_formats.as_ref().is_none_or(|formats| {
|
||||
formats
|
||||
.iter()
|
||||
.any(|value| api_format_matches(value, api_format))
|
||||
}) && mapping.endpoint_ids.as_ref().is_none_or(|endpoint_ids| {
|
||||
endpoint_ids
|
||||
.iter()
|
||||
.any(|endpoint_id| endpoint_id == &row.endpoint_id)
|
||||
}) && mapping.name == requested_model_name
|
||||
mapping_scope_matches(mapping, row, api_format)
|
||||
&& mapping.name == requested_model_name
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn row_has_available_provider_model(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
row_mapping_matches_scope(row, api_format)
|
||||
|| row_default_provider_model_name_available(row, api_format)
|
||||
}
|
||||
|
||||
fn row_default_provider_model_name_available(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
let Some(mappings) = row.model_provider_model_mappings.as_ref() else {
|
||||
return true;
|
||||
};
|
||||
let mut has_explicit_default_mapping = false;
|
||||
for mapping in mappings {
|
||||
if mapping.name != row.model_provider_model_name {
|
||||
continue;
|
||||
}
|
||||
has_explicit_default_mapping = true;
|
||||
if mapping_scope_matches(mapping, row, api_format) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
!has_explicit_default_mapping
|
||||
}
|
||||
|
||||
fn row_mapping_matches_scope(row: &StoredMinimalCandidateSelectionRow, api_format: &str) -> bool {
|
||||
row.model_provider_model_mappings
|
||||
.as_ref()
|
||||
.is_some_and(|mappings| {
|
||||
mappings
|
||||
.iter()
|
||||
.any(|mapping| mapping_scope_matches(mapping, row, api_format))
|
||||
})
|
||||
}
|
||||
|
||||
fn mapping_scope_matches(
|
||||
mapping: &super::StoredProviderModelMapping,
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
mapping.api_formats.as_ref().is_none_or(|formats| {
|
||||
formats
|
||||
.iter()
|
||||
.any(|value| api_format_matches(value, api_format))
|
||||
}) && mapping.endpoint_ids.as_ref().is_none_or(|endpoint_ids| {
|
||||
endpoint_ids
|
||||
.iter()
|
||||
.any(|endpoint_id| endpoint_id == &row.endpoint_id)
|
||||
})
|
||||
}
|
||||
|
||||
fn key_auth_channel_matches(row: &CandidateSelectionRow, api_format: &str) -> bool {
|
||||
let provider_type = row.row.provider_type.trim().to_ascii_lowercase();
|
||||
let auth_type = row.row.key_auth_type.trim().to_ascii_lowercase();
|
||||
|
||||
Reference in New Issue
Block a user