校验图片请求模型与 provider 绑定一致

executor 三个请求入口改用 ensure_provider_matches_model,模型不属于当前 provider 时直接返回 InvalidRequest

新增单测覆盖 provider 不一致的快速失败,并把 GPT 模型相关集成测试的 settings 调整为 tiantoken
This commit is contained in:
2026-09-19 13:46:47 +08:00
parent 0c54cca81c
commit 42fdb57efe
2 changed files with 75 additions and 24 deletions
@@ -69,12 +69,7 @@ pub async fn create_image_generation_with_model(
failure_context: &str,
) -> Result<GeneratedImages, PlatformImageError> {
let requested_model = normalize_image_model(model);
resolve_image_provider(requested_model).map_err(|message| {
PlatformImageError::InvalidRequest {
provider: settings.provider.as_str(),
message: format!("{failure_context}{message}{requested_model}"),
}
})?;
ensure_provider_matches_model(settings, requested_model, failure_context)?;
if !reference_images.is_empty() {
let resolved_references = resolve_reference_images(
http_client,
@@ -268,10 +263,7 @@ pub async fn create_nanobanana_generate_content(
failure_context: &str,
) -> Result<GeneratedImages, PlatformImageError> {
let model = normalize_image_model(model);
resolve_image_provider(model).map_err(|message| PlatformImageError::InvalidRequest {
provider: settings.provider.as_str(),
message: format!("{failure_context}{message}{model}"),
})?;
ensure_provider_matches_model(settings, model, failure_context)?;
if reference_images.len() > VECTOR_ENGINE_NANOBANANA_MAX_REFERENCE_IMAGES {
return Err(PlatformImageError::InvalidRequest {
provider: settings.provider.as_str(),
@@ -492,12 +484,7 @@ pub async fn create_image_edit_with_references_and_model(
failure_context: &str,
) -> Result<GeneratedImages, PlatformImageError> {
let requested_model = normalize_image_model(model);
resolve_image_provider(requested_model).map_err(|message| {
PlatformImageError::InvalidRequest {
provider: settings.provider.as_str(),
message: format!("{failure_context}{message}{requested_model}"),
}
})?;
ensure_provider_matches_model(settings, requested_model, failure_context)?;
if reference_images.is_empty() {
return Err(missing_reference_images_error(
settings.provider,
@@ -704,6 +691,34 @@ pub async fn create_image_edit_with_references_and_model(
}
}
/// 校验请求模型与 settings 绑定的 provider 是否一致。
///
/// executor 用 `settings.base_url`/`settings.api_key` 发请求,如果调用方传入了属于
/// 另一个 provider 的模型,就会把错误的密钥/网关当成目标端点。这里显式 fail fast,
/// 避免静默地把请求发到不匹配的 provider。
fn ensure_provider_matches_model(
settings: &ImageProviderSettings,
model: &str,
failure_context: &str,
) -> Result<(), PlatformImageError> {
let resolved =
resolve_image_provider(model).map_err(|message| PlatformImageError::InvalidRequest {
provider: settings.provider.as_str(),
message: format!("{failure_context}{message}{model}"),
})?;
if resolved != settings.provider {
return Err(PlatformImageError::InvalidRequest {
provider: settings.provider.as_str(),
message: format!(
"{failure_context}:模型 {model} 属于 {},与当前 {provider} 设置不一致",
resolved.as_str(),
provider = settings.provider.as_str()
),
});
}
Ok(())
}
fn preferred_image_upstream_model(requested_model: &str) -> &str {
// Provider routing is selected by api-server at the task boundary. Keep the
// concrete value intact here so retries stay on the same model and legacy
@@ -922,7 +937,7 @@ mod tests {
}
}
fn test_settings() -> ImageProviderSettings {
fn vector_engine_settings() -> ImageProviderSettings {
ImageProviderSettings {
provider: ImageProvider::VectorEngine,
base_url: "http://127.0.0.1:9".to_string(),
@@ -932,12 +947,22 @@ mod tests {
}
}
fn tiantoken_settings() -> ImageProviderSettings {
ImageProviderSettings {
provider: ImageProvider::Tiantoken,
base_url: "http://127.0.0.1:9".to_string(),
api_key: "test-key".to_string(),
request_timeout_ms: 1_000,
request_deadline: None,
}
}
#[tokio::test]
async fn gpt_image_edit_rejects_six_references_before_network_send() {
let references = (0..6).map(reference_image).collect::<Vec<_>>();
let error = create_image_edit_with_references_and_model(
&reqwest::Client::new(),
&test_settings(),
&tiantoken_settings(),
GPT_IMAGE_2_MODEL,
"测试提示词",
None,
@@ -957,7 +982,7 @@ mod tests {
let references = (0..15).map(reference_image).collect::<Vec<_>>();
let error = create_nanobanana_generate_content(
&reqwest::Client::new(),
&test_settings(),
&vector_engine_settings(),
crate::image_provider::constants::NANOBANANA_2_MODEL,
"测试提示词",
None,
@@ -975,7 +1000,7 @@ mod tests {
#[tokio::test]
async fn expired_deadline_stops_generation_before_network_send() {
let settings = ImageProviderSettings {
provider: ImageProvider::VectorEngine,
provider: ImageProvider::Tiantoken,
base_url: "http://127.0.0.1:9".to_string(),
api_key: "test-key".to_string(),
request_timeout_ms: 1_000,
@@ -1007,6 +1032,32 @@ mod tests {
);
}
#[tokio::test]
async fn image_generation_rejects_model_from_another_provider() {
let error = create_image_generation(
&reqwest::Client::new(),
&vector_engine_settings(),
"测试提示词",
None,
"1024x1024",
1,
&[],
"测试图片生成失败",
)
.await
.expect_err("provider mismatch must fail fast");
match error {
PlatformImageError::InvalidRequest { message, .. } => {
assert!(
message.contains("设置不一致"),
"unexpected message: {message}"
);
}
other => panic!("expected InvalidRequest, got {other:?}"),
}
}
#[test]
fn image_provider_send_retry_policy_allows_four_retries_before_final_attempt() {
assert_eq!(VECTOR_ENGINE_SEND_MAX_ATTEMPTS, 5);
@@ -245,7 +245,7 @@ async fn image_edit_retries_send_timeout_once_and_succeeds() {
});
let settings = ImageProviderSettings {
provider: ImageProvider::VectorEngine,
provider: ImageProvider::Tiantoken,
base_url: format!("http://{server_addr}/v1"),
api_key: "test-key".to_string(),
request_timeout_ms: 40,
@@ -346,7 +346,7 @@ async fn image_provider_deadline_clips_stalled_attempt_and_prevents_retry() {
});
let mut settings = ImageProviderSettings {
provider: ImageProvider::VectorEngine,
provider: ImageProvider::Tiantoken,
base_url: format!("http://{server_addr}/v1"),
api_key: "test-key".to_string(),
request_timeout_ms: 5_000,
@@ -506,7 +506,7 @@ async fn image_generation_stays_on_model_after_upstream_502() {
});
let settings = ImageProviderSettings {
provider: ImageProvider::VectorEngine,
provider: ImageProvider::Tiantoken,
base_url: format!("http://{server_addr}/v1"),
api_key: "test-key".to_string(),
request_timeout_ms: 1_000,
@@ -793,7 +793,7 @@ async fn start_http_response_sequence(
fn test_image_provider_settings(base_url: String) -> ImageProviderSettings {
ImageProviderSettings {
provider: ImageProvider::VectorEngine,
provider: ImageProvider::Tiantoken,
base_url,
api_key: "test-key".to_string(),
request_timeout_ms: 1_000,