fix: tighten model candidate matching scopes

This commit is contained in:
fawney19
2026-05-12 00:52:19 +08:00
parent 0fa97595bf
commit 09146f8cdd
6 changed files with 563 additions and 71 deletions

View File

@@ -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();

View File

@@ -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();

View File

@@ -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,"));

View File

@@ -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();

View File

@@ -29,33 +29,38 @@ pub fn resolve_requested_global_model_name_with_model_directives(
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
})
resolve_global_model_name_by(rows, |row| {
row_has_available_provider_model(row, api_format)
&& row.global_model_name == requested_model_name
})
.or_else(|| {
resolve_global_model_name_by(rows, |row| {
row_default_provider_model_name_available(row, api_format)
&& 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, row, api_format)
&& mapping.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, row, 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| {
})
.or_else(|| {
resolve_global_model_name_by(rows, |row| {
row_has_available_provider_model(row, api_format)
&& row.global_model_mappings.as_ref().is_some_and(|patterns| {
patterns
.iter()
.any(|pattern| matches_model_mapping(pattern, requested_model_name))
})
})
})
})
},
)
}
@@ -86,8 +91,15 @@ fn row_supports_requested_model_exact(
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.global_model_mappings.as_ref().is_some_and(|patterns| {
patterns
.iter()
.any(|pattern| matches_model_mapping(pattern, requested_model_name))
}))
|| row
.model_provider_model_mappings
.as_ref()
@@ -97,11 +109,6 @@ fn row_supports_requested_model_exact(
&& mapping.name == requested_model_name
})
})
|| row.global_model_mappings.as_ref().is_some_and(|patterns| {
patterns
.iter()
.any(|pattern| matches_model_mapping(pattern, requested_model_name))
})
}
fn resolve_global_model_name_by<F>(
@@ -138,7 +145,7 @@ pub fn resolve_provider_model_name_with_model_directives(
api_format: &str,
enable_model_directives: bool,
) -> Option<(String, Option<String>)> {
let selected_provider_model_name = select_provider_model_name(row, api_format);
let selected_provider_model_name = resolve_selected_provider_model_name(row, api_format)?;
let Some(key_allowed_models) = row.key_allowed_models.as_ref() else {
return Some((selected_provider_model_name, None));
};
@@ -195,11 +202,19 @@ pub fn select_provider_model_name(
row: &StoredMinimalCandidateSelectionRow,
api_format: &str,
) -> String {
resolve_selected_provider_model_name(row, api_format)
.unwrap_or_else(|| row.model_provider_model_name.clone())
}
fn resolve_selected_provider_model_name(
row: &StoredMinimalCandidateSelectionRow,
api_format: &str,
) -> Option<String> {
let Some(mappings) = row.model_provider_model_mappings.as_ref() else {
return row.model_provider_model_name.clone();
return Some(row.model_provider_model_name.clone());
};
mappings
if let Some(mapping) = mappings
.iter()
.filter(|mapping| mapping_scope_matches(mapping, row, api_format))
.min_by(|left, right| {
@@ -207,15 +222,22 @@ pub fn select_provider_model_name(
.cmp(&right.priority)
.then(left.name.cmp(&right.name))
})
.map(|mapping| mapping.name.clone())
.unwrap_or_else(|| row.model_provider_model_name.clone())
{
return Some(mapping.name.clone());
}
row_default_provider_model_name_available(row, api_format)
.then(|| row.model_provider_model_name.clone())
}
pub fn candidate_model_names(
row: &StoredMinimalCandidateSelectionRow,
api_format: &str,
) -> BTreeSet<String> {
let mut names = BTreeSet::from([row.model_provider_model_name.clone()]);
let mut names = BTreeSet::new();
if row_default_provider_model_name_available(row, api_format) {
names.insert(row.model_provider_model_name.clone());
}
if let Some(mappings) = row.model_provider_model_mappings.as_ref() {
for mapping in mappings {
if mapping_scope_matches(mapping, row, api_format) {
@@ -226,6 +248,33 @@ pub fn candidate_model_names(
names
}
fn row_has_available_provider_model(
row: &StoredMinimalCandidateSelectionRow,
api_format: &str,
) -> bool {
resolve_selected_provider_model_name(row, api_format).is_some()
}
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 mapping_scope_matches(
mapping: &StoredProviderModelMapping,
row: &StoredMinimalCandidateSelectionRow,
@@ -358,7 +407,8 @@ fn row_has_candidate_model_name(
api_format: &str,
model_name: &str,
) -> bool {
row.model_provider_model_name == model_name
(row_default_provider_model_name_available(row, api_format)
&& row.model_provider_model_name == model_name)
|| row
.model_provider_model_mappings
.as_ref()
@@ -494,6 +544,54 @@ mod tests {
assert_eq!(resolved.1.as_deref(), Some("gpt-5.4"));
}
#[test]
fn endpoint_scoped_default_mapping_limits_exact_global_model_match() {
let mut row = sample_row("deepseek-v4-pro", "deepseek-v4-pro");
row.endpoint_id = "endpoint-claude".to_string();
row.endpoint_api_format = "claude:messages".to_string();
row.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()]),
}]);
assert!(!row_supports_requested_model(
&row,
"deepseek-v4-pro",
"claude:messages"
));
assert!(resolve_provider_model_name(&row, "deepseek-v4-pro", "claude:messages").is_none());
assert_eq!(
resolve_requested_global_model_name_with_model_directives(
&[row.clone()],
"deepseek-v4-pro",
"claude:messages",
false,
),
None
);
row.endpoint_id = "endpoint-openai".to_string();
row.endpoint_api_format = "openai:chat".to_string();
assert!(row_supports_requested_model(
&row,
"deepseek-v4-pro",
"openai:chat"
));
assert_eq!(
resolve_requested_global_model_name_with_model_directives(
&[row],
"deepseek-v4-pro",
"openai:chat",
false,
)
.as_deref(),
Some("deepseek-v4-pro")
);
}
fn sample_row(
global_model_name: &str,
model_provider_model_name: &str,