diff --git a/crates/aether-data/adapters/postgres/src/candidate_selection.rs b/crates/aether-data/adapters/postgres/src/candidate_selection.rs index dddec6d1b..ee7690826 100644 --- a/crates/aether-data/adapters/postgres/src/candidate_selection.rs +++ b/crates/aether-data/adapters/postgres/src/candidate_selection.rs @@ -1050,17 +1050,6 @@ fn requested_model_selection_sql() -> String { ) ) ) - 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 ( @@ -1068,17 +1057,6 @@ fn requested_model_selection_sql() -> String { 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( @@ -1108,6 +1086,47 @@ fn requested_model_selection_sql() -> String { ) ) ) + 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 + ) + AND 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 __AETHER_PROVIDER_MODEL_MAPPING_API_FORMAT_MATCH__ + ) + ) + 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 ( @@ -1706,7 +1725,7 @@ mod tests { let sql = requested_model_selection_sql(); let compatibility = PROVIDER_MODEL_MAPPING_API_FORMAT_MATCH_SQL; - assert_eq!(sql.matches(compatibility).count(), 3); + assert_eq!(sql.matches(compatibility).count(), 4); assert!(!sql.contains(PROVIDER_MODEL_MAPPING_API_FORMAT_MATCH_MARKER)); assert!(compatibility.contains("LOWER(BTRIM(p.provider_type)) = 'codex'")); assert!(compatibility.contains("LOWER($4) = 'codex:live'")); @@ -1730,6 +1749,21 @@ mod tests { } } + #[test] + fn requested_model_sql_confines_default_name_to_mapping_scope() { + let sql = requested_model_selection_sql(); + + // 默认名(provider_model_name)的可用性由映射的端点/格式范围决定: + // 分支 A(全局名匹配)不再单独放行“无默认名映射”的行; + // 分支 B(provider 名匹配)要求默认名映射命中,或映射范围内至少有一条映射命中。 + assert_eq!(sql.matches("->> 'name' = m.provider_model_name").count(), 2); + assert_eq!( + sql.matches(PROVIDER_MODEL_MAPPING_API_FORMAT_MATCH_SQL) + .count(), + 4 + ); + } + #[test] fn candidate_selection_sql_allows_grok_oauth_chat_auth() { let requested_model_sql = requested_model_selection_sql(); diff --git a/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs b/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs index cf70cc1bb..211ad8e3c 100644 --- a/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs +++ b/crates/aether-data/runtime/src/repository/candidate_selection/memory.rs @@ -284,7 +284,12 @@ fn row_default_provider_model_name_available( return true; } } - !has_explicit_default_mapping + if has_explicit_default_mapping { + return false; + } + // 与调度核心保持一致:映射一旦限定了端点或 API 格式范围, + // 默认模型名就只在范围内可用;范围之外整行不再匹配。 + row_mapping_matches_scope(row, api_format) } fn row_mapping_matches_scope(row: &StoredMinimalCandidateSelectionRow, api_format: &str) -> bool { @@ -691,6 +696,45 @@ mod tests { assert_eq!(rows[0].endpoint_id, "endpoint-openai"); } + #[tokio::test] + async fn requested_model_filter_confines_rows_to_mapping_endpoint_scope() { + // 场景:映射只把上游名绑定到 openai:chat 端点; + // 其他端点(claude:messages)不应回退到默认名,整行不可匹配。 + let mut selected = sample_row("provider-1", "openai:chat", "deepseek-flash", 10); + selected.endpoint_id = "endpoint-chat".to_string(); + selected.model_provider_model_name = "deepseek-flash".to_string(); + selected.model_provider_model_mappings = Some(vec![StoredProviderModelMapping { + name: "gpt-6-luna".to_string(), + priority: 1, + api_formats: None, + endpoint_ids: Some(vec!["endpoint-chat".to_string()]), + operations: None, + }]); + + 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-flash") + .await + .expect("list should succeed"); + assert!(rows.is_empty()); + + let rows = repository + .list_for_exact_api_format_and_requested_model("openai:chat", "deepseek-flash") + .await + .expect("list should succeed"); + assert_eq!(rows.len(), 1); + assert_eq!(rows[0].endpoint_id, "endpoint-chat"); + } + #[tokio::test] async fn requested_model_page_returns_requested_slice_only() { let mut rows = Vec::new(); diff --git a/crates/aether-scheduler-core/src/model.rs b/crates/aether-scheduler-core/src/model.rs index 1d1c099f5..c2b2d6299 100644 --- a/crates/aether-scheduler-core/src/model.rs +++ b/crates/aether-scheduler-core/src/model.rs @@ -393,14 +393,20 @@ fn row_default_provider_model_name_available( return true; } } - !has_explicit_default_mapping + if has_explicit_default_mapping { + return false; + } + // 映射一旦限定了端点或 API 格式范围,默认模型名就只在范围内可用; + // 范围之外整行不再参与候选,避免把默认名发到不认识它的上游端点。 + mappings + .iter() + .any(|mapping| mapping_endpoint_and_api_format_scope_matches(mapping, row, api_format)) } -fn mapping_scope_matches( +fn mapping_endpoint_and_api_format_scope_matches( mapping: &StoredProviderModelMapping, row: &StoredMinimalCandidateSelectionRow, api_format: &str, - request_operation: Option<&str>, ) -> bool { let api_format_matches_scope = mapping.api_formats.as_ref().is_none_or(|api_formats| { api_formats.iter().any(|value| { @@ -411,24 +417,29 @@ fn mapping_scope_matches( return false; } - let endpoint_matches_scope = mapping.endpoint_ids.as_ref().is_none_or(|endpoint_ids| { + mapping.endpoint_ids.as_ref().is_none_or(|endpoint_ids| { endpoint_ids .iter() .any(|endpoint_id| endpoint_id == &row.endpoint_id) - }); - if !endpoint_matches_scope { - return false; - } - - mapping.operations.as_ref().is_none_or(|operations| { - request_operation.is_some_and(|request_operation| { - operations - .iter() - .any(|operation| operation.eq_ignore_ascii_case(request_operation)) - }) }) } +fn mapping_scope_matches( + mapping: &StoredProviderModelMapping, + row: &StoredMinimalCandidateSelectionRow, + api_format: &str, + request_operation: Option<&str>, +) -> bool { + mapping_endpoint_and_api_format_scope_matches(mapping, row, api_format) + && mapping.operations.as_ref().is_none_or(|operations| { + request_operation.is_some_and(|request_operation| { + operations + .iter() + .any(|operation| operation.eq_ignore_ascii_case(request_operation)) + }) + }) +} + fn mapping_operation_scope_rank(mapping: &StoredProviderModelMapping) -> u8 { u8::from(mapping.operations.is_some()) } @@ -1010,6 +1021,55 @@ mod tests { ); } + #[test] + fn endpoint_scoped_mapping_confines_row_availability_to_scoped_endpoints() { + // 场景:某行只把上游名 gpt-6-luna 映射到 openai:chat 端点; + // 其他端点(如 claude:messages)不应回退到默认名,整行不可用。 + let mut row = sample_row("deepseek-flash", "deepseek-flash"); + row.endpoint_id = "endpoint-chat".to_string(); + row.endpoint_api_format = "openai:chat".to_string(); + row.model_provider_model_mappings = Some(vec![StoredProviderModelMapping { + name: "gpt-6-luna".to_string(), + priority: 1, + api_formats: None, + endpoint_ids: Some(vec!["endpoint-chat".to_string()]), + operations: None, + }]); + + assert!(row_supports_requested_model( + &row, + "deepseek-flash", + "openai:chat" + )); + assert_eq!( + resolve_provider_model_name(&row, "deepseek-flash", "openai:chat") + .map(|resolved| resolved.0), + Some("gpt-6-luna".to_string()) + ); + + let mut claude_row = row.clone(); + claude_row.endpoint_id = "endpoint-claude".to_string(); + claude_row.endpoint_api_format = "claude:messages".to_string(); + + assert!(!row_supports_requested_model( + &claude_row, + "deepseek-flash", + "claude:messages" + )); + assert!( + resolve_provider_model_name(&claude_row, "deepseek-flash", "claude:messages").is_none() + ); + assert_eq!( + resolve_requested_global_model_name_with_model_directives( + &[claude_row], + "deepseek-flash", + "claude:messages", + false, + ), + None + ); + } + #[test] fn reserved_global_model_name_keeps_its_own_rows_addressable() { let row = cursor_alias_row(); @@ -1069,6 +1129,48 @@ mod tests { ); } + #[test] + fn operation_scoped_mapping_keeps_default_name_available() { + // 只限定“适用请求”(如 compact)的映射不应把整行从其他请求里排除掉。 + let mut row = sample_row("gpt-5.6-sol", "gpt-5.6-sol"); + row.endpoint_api_format = "openai:responses".to_string(); + row.model_provider_model_mappings = Some(vec![StoredProviderModelMapping { + name: "gpt-5.6-terra".to_string(), + priority: 1, + api_formats: Some(vec!["openai:responses".to_string()]), + endpoint_ids: None, + operations: Some(vec!["compact".to_string()]), + }]); + + assert!(row_supports_requested_model( + &row, + "gpt-5.6-sol", + "openai:responses" + )); + assert_eq!( + resolve_provider_model_name_with_model_directives_and_request_operation( + &row, + "gpt-5.6-sol", + "openai:responses", + false, + None, + ) + .map(|resolved| resolved.0), + Some("gpt-5.6-sol".to_string()) + ); + assert_eq!( + resolve_provider_model_name_with_model_directives_and_request_operation( + &row, + "gpt-5.6-sol", + "openai:responses", + false, + Some("compact"), + ) + .map(|resolved| resolved.0), + Some("gpt-5.6-terra".to_string()) + ); + } + fn sample_row( global_model_name: &str, model_provider_model_name: &str,