diff --git a/crates/api/src/routes/completions.rs b/crates/api/src/routes/completions.rs index f1aaa775..941781d3 100644 --- a/crates/api/src/routes/completions.rs +++ b/crates/api/src/routes/completions.rs @@ -3838,8 +3838,14 @@ mod tests { #[test] fn test_classify_provider_error_404_surfaces_message() { - let (status, error_type, message) = - classify_provider_error(404, "model 'foo' not found".to_string()); + let (status, error_type, message) = classify_provider_error( + 404, + "model 'foo' not found".to_string(), + ( + StatusCode::INTERNAL_SERVER_ERROR, + "Embeddings request failed. Please try again later.", + ), + ); assert_eq!(status, StatusCode::NOT_FOUND); assert_eq!(error_type, "not_found_error"); assert_eq!(message, "model 'foo' not found"); @@ -3847,8 +3853,14 @@ mod tests { #[test] fn test_classify_provider_error_429_surfaces_message() { - let (status, error_type, message) = - classify_provider_error(429, "too many concurrent requests".to_string()); + let (status, error_type, message) = classify_provider_error( + 429, + "too many concurrent requests".to_string(), + ( + StatusCode::INTERNAL_SERVER_ERROR, + "Embeddings request failed. Please try again later.", + ), + ); assert_eq!(status, StatusCode::TOO_MANY_REQUESTS); assert_eq!(error_type, "rate_limit_error"); assert_eq!(message, "too many concurrent requests"); @@ -3859,6 +3871,10 @@ mod tests { let (status, error_type, message) = classify_provider_error( 400, "dimensions is not supported for this model".to_string(), + ( + StatusCode::INTERNAL_SERVER_ERROR, + "Embeddings request failed. Please try again later.", + ), ); assert_eq!(status, StatusCode::BAD_REQUEST); assert_eq!(error_type, "invalid_request_error"); @@ -3867,8 +3883,14 @@ mod tests { #[test] fn test_classify_provider_error_422_preserves_upstream_status() { - let (status, error_type, message) = - classify_provider_error(422, "validation failed".to_string()); + let (status, error_type, message) = classify_provider_error( + 422, + "validation failed".to_string(), + ( + StatusCode::INTERNAL_SERVER_ERROR, + "Embeddings request failed. Please try again later.", + ), + ); assert_eq!(status, StatusCode::UNPROCESSABLE_ENTITY); assert_eq!(error_type, "invalid_request_error"); assert_eq!(message, "validation failed"); @@ -3878,8 +3900,14 @@ mod tests { fn test_classify_provider_error_401_masked_as_5xx() { // 401 from upstream means *our* credentials are wrong — the client // did nothing to cause it, so we must not echo the auth error. - let (status, error_type, message) = - classify_provider_error(401, "Invalid API key 'sk-***'".to_string()); + let (status, error_type, message) = classify_provider_error( + 401, + "Invalid API key 'sk-***'".to_string(), + ( + StatusCode::BAD_GATEWAY, + "Privacy classify request failed. Please try again later.", + ), + ); assert_eq!(status, StatusCode::INTERNAL_SERVER_ERROR); assert_eq!(error_type, "server_error"); assert!( @@ -3891,8 +3919,14 @@ mod tests { #[test] fn test_classify_provider_error_403_masked_as_5xx() { - let (status, error_type, message) = - classify_provider_error(403, "Forbidden: backend ACL denied".to_string()); + let (status, error_type, message) = classify_provider_error( + 403, + "Forbidden: backend ACL denied".to_string(), + ( + StatusCode::INTERNAL_SERVER_ERROR, + "Embeddings request failed. Please try again later.", + ), + ); assert_eq!(status, StatusCode::INTERNAL_SERVER_ERROR); assert_eq!(error_type, "server_error"); assert_eq!( @@ -3903,8 +3937,14 @@ mod tests { #[test] fn test_classify_provider_error_407_masked_as_5xx() { - let (status, error_type, _) = - classify_provider_error(407, "proxy auth required".to_string()); + let (status, error_type, _) = classify_provider_error( + 407, + "proxy auth required".to_string(), + ( + StatusCode::BAD_GATEWAY, + "Privacy redact request failed. Please try again later.", + ), + ); assert_eq!(status, StatusCode::INTERNAL_SERVER_ERROR); assert_eq!(error_type, "server_error"); } @@ -3915,6 +3955,10 @@ mod tests { let (status, error_type, message) = classify_provider_error( 502, "RuntimeError: traceback...internal-host:9000".to_string(), + ( + StatusCode::INTERNAL_SERVER_ERROR, + "Embeddings request failed. Please try again later.", + ), ); assert_eq!(status, StatusCode::INTERNAL_SERVER_ERROR); assert_eq!(error_type, "server_error"); @@ -3922,6 +3966,54 @@ mod tests { assert!(!message.contains("internal-host")); } + #[test] + fn test_classify_provider_error_non_4xx_uses_caller_fallback() { + // Given: a provider redirect and the privacy endpoint's generic contract. + let upstream_message = "redirected to internal-host".to_string(); + + // When: the shared classifier handles the non-4xx response. + let (status, error_type, message) = classify_provider_error( + 302, + upstream_message, + ( + StatusCode::BAD_GATEWAY, + "Privacy classify request failed. Please try again later.", + ), + ); + + // Then: the caller-selected gateway response masks upstream details. + assert_eq!(status, StatusCode::BAD_GATEWAY); + assert_eq!(error_type, "server_error"); + assert_eq!( + message, + "Privacy classify request failed. Please try again later." + ); + } + + #[test] + fn test_classify_provider_error_internal_server_error_stays_500() { + // Given: the service layer has already masked upstream credentials as 500. + let upstream_message = "The model is currently unavailable".to_string(); + + // When: a privacy caller otherwise uses 502 for generic failures. + let (status, error_type, message) = classify_provider_error( + 500, + upstream_message, + ( + StatusCode::BAD_GATEWAY, + "Privacy classify request failed. Please try again later.", + ), + ); + + // Then: the credential/configuration failure remains a generic 500. + assert_eq!(status, StatusCode::INTERNAL_SERVER_ERROR); + assert_eq!(error_type, "server_error"); + assert_eq!( + message, + "Privacy classify request failed. Please try again later." + ); + } + #[test] fn test_stream_chunk_serialization_preserves_field_order() { // Verify that StreamChunk::Chat serializes with struct field order @@ -5782,22 +5874,18 @@ struct EmbeddingsResponseDoc { /// - Other 4xx → preserve upstream status with `invalid_request_error`. The /// upstream message is surfaced so the user can see *why* their request was /// rejected (e.g. "dimensions is not supported for this model"). -/// - 5xx (and anything else) → 500 server_error with a generic message. The -/// upstream body may contain stack traces or internal details we don't want -/// to leak. +/// - 500 → generic 500, preserving service-layer masking for upstream auth faults. +/// - Other 5xx (and anything else) → the caller's generic server-error response. +/// The upstream body may contain stack traces or internal details we don't +/// want to leak. fn classify_provider_error( upstream_status: u16, upstream_message: String, + generic_response: (StatusCode, &str), ) -> (StatusCode, &'static str, String) { - let generic = || { - ( - StatusCode::INTERNAL_SERVER_ERROR, - "server_error", - "Embeddings request failed. Please try again later.".to_string(), - ) - }; + let generic = |status| (status, "server_error", generic_response.1.to_string()); match upstream_status { - 401 | 403 | 407 => generic(), + 401 | 403 | 407 => generic(StatusCode::INTERNAL_SERVER_ERROR), 404 => (StatusCode::NOT_FOUND, "not_found_error", upstream_message), 429 => ( StatusCode::TOO_MANY_REQUESTS, @@ -5808,7 +5896,8 @@ fn classify_provider_error( let http_status = StatusCode::from_u16(s).unwrap_or(StatusCode::BAD_REQUEST); (http_status, "invalid_request_error", upstream_message) } - _ => generic(), + 500 => generic(StatusCode::INTERNAL_SERVER_ERROR), + _ => generic(generic_response.0), } } @@ -6054,7 +6143,14 @@ pub async fn embeddings( // upstream URL on connection failures, where `status_code` is // synthetic) — without it the log can't distinguish a backend // 5xx from an unreachable SNI. - let classified = classify_provider_error(status_code, message.clone()); + let classified = classify_provider_error( + status_code, + message.clone(), + ( + StatusCode::INTERNAL_SERVER_ERROR, + "Embeddings request failed. Please try again later.", + ), + ); if classified.0.is_client_error() { tracing::warn!( model = %model_name, @@ -6384,18 +6480,29 @@ pub async fn privacy_classify( status_code, message, } => { - tracing::error!( - upstream_status = status_code, - detail = %message, - "Privacy classify provider error" + // The provider pool replaces every upstream body with a + // status-derived message before it reaches this handler, + // which makes forwarding classified client errors safe. + let classified = classify_provider_error( + status_code, + message, + ( + StatusCode::BAD_GATEWAY, + "Privacy classify request failed. Please try again later.", + ), ); - let http_status = StatusCode::from_u16(status_code) - .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR); - ( - http_status, - "server_error", - "Privacy classify request failed. Please try again later.".to_string(), - ) + if classified.0.is_client_error() { + tracing::warn!( + upstream_status = status_code, + "Privacy classify provider error" + ); + } else { + tracing::error!( + upstream_status = status_code, + "Privacy classify provider error" + ); + } + classified } services::completions::ports::CompletionError::InvalidModel(msg) => { tracing::warn!("Privacy classify model not found"); @@ -6686,18 +6793,29 @@ pub async fn privacy_redact( status_code, message, } => { - tracing::error!( - upstream_status = status_code, - detail = %message, - "Privacy redact provider error" + // The provider pool replaces every upstream body with a + // status-derived message before it reaches this handler, + // which makes forwarding classified client errors safe. + let classified = classify_provider_error( + status_code, + message, + ( + StatusCode::BAD_GATEWAY, + "Privacy redact request failed. Please try again later.", + ), ); - let http_status = StatusCode::from_u16(status_code) - .unwrap_or(StatusCode::INTERNAL_SERVER_ERROR); - ( - http_status, - "server_error", - "Privacy redact request failed. Please try again later.".to_string(), - ) + if classified.0.is_client_error() { + tracing::warn!( + upstream_status = status_code, + "Privacy redact provider error" + ); + } else { + tracing::error!( + upstream_status = status_code, + "Privacy redact provider error" + ); + } + classified } services::completions::ports::CompletionError::InvalidModel(msg) => { tracing::warn!("Privacy redact model not found"); @@ -6706,9 +6824,9 @@ pub async fn privacy_redact( services::completions::ports::CompletionError::ServiceOverloaded(_) => { tracing::warn!("Privacy redact service overloaded"); ( - StatusCode::SERVICE_UNAVAILABLE, + crate::routes::common::status_overloaded(), "service_overloaded", - "The service is temporarily overloaded. Please retry with exponential backoff.".to_string(), + "All inference backends are overloaded. Please retry with exponential backoff.".to_string(), ) } _ => { diff --git a/crates/api/tests/e2e_all/privacy_classify.rs b/crates/api/tests/e2e_all/privacy_classify.rs index c5fd4e8f..f9bda8c0 100644 --- a/crates/api/tests/e2e_all/privacy_classify.rs +++ b/crates/api/tests/e2e_all/privacy_classify.rs @@ -5,6 +5,7 @@ use crate::common::*; use api::models::{BatchUpdateModelApiRequest, ErrorResponse}; +use axum::http::header::RETRY_AFTER; async fn setup_privacy_filter_model(server: &axum_test::TestServer) -> String { let mut batch = BatchUpdateModelApiRequest::new(); @@ -243,3 +244,201 @@ async fn test_privacy_classify_costs_deducted() { "Privacy classify should bill 10 tokens × 1_000_000 = 10_000_000", ); } + +#[tokio::test] +async fn test_privacy_classify_upstream_429_returns_rate_limit_with_retry_after() { + // Given: an upstream privacy provider that rejects the request with a + // body containing data that must never leave the provider boundary. + const UPSTREAM_BODY_SENTINEL: &str = + "UPSTREAM_PRIVACY_BODY_SENTINEL::alice@example.com::123-45-6789"; + let (server, _pool, mock_provider, _db) = setup_test_server_with_pool().await; + setup_privacy_filter_model(&server).await; + let org = setup_org_with_credits(&server, 10_000_000_000i64).await; + let api_key = get_api_key_for_org(&server, org.id).await; + mock_provider + .set_privacy_classify_error_override(Some( + inference_providers::PrivacyClassifyError::HttpError { + status_code: 429, + message: UPSTREAM_BODY_SENTINEL.to_string(), + }, + )) + .await; + + // When: a client calls the real classify route. + let response = server + .post("/v1/privacy/classify") + .add_header("Authorization", format!("Bearer {api_key}")) + .add_header("User-Agent", MOCK_USER_AGENT) + .json(&serde_json::json!({ + "model": "openai/privacy-filter", + "input": "Classify this text" + })) + .await; + + // Then: retry semantics are explicit and the upstream body is absent. + assert_eq!(response.status_code(), 429); + assert_eq!( + response + .headers() + .get(RETRY_AFTER) + .and_then(|value| value.to_str().ok()), + Some("2") + ); + let error: ErrorResponse = response.json(); + assert_eq!(error.error.r#type, "rate_limit_error"); + assert!(!error.error.message.contains(UPSTREAM_BODY_SENTINEL)); +} + +#[tokio::test] +async fn test_privacy_classify_upstream_413_returns_invalid_request() { + // Given: an oversized upstream response whose body contains unsafe data. + const UPSTREAM_BODY_SENTINEL: &str = + "UPSTREAM_PRIVACY_BODY_SENTINEL::alice@example.com::123-45-6789"; + let (server, _pool, mock_provider, _db) = setup_test_server_with_pool().await; + setup_privacy_filter_model(&server).await; + let org = setup_org_with_credits(&server, 10_000_000_000i64).await; + let api_key = get_api_key_for_org(&server, org.id).await; + mock_provider + .set_privacy_classify_error_override(Some( + inference_providers::PrivacyClassifyError::HttpError { + status_code: 413, + message: UPSTREAM_BODY_SENTINEL.to_string(), + }, + )) + .await; + + // When: a client calls the real classify route. + let response = server + .post("/v1/privacy/classify") + .add_header("Authorization", format!("Bearer {api_key}")) + .add_header("User-Agent", MOCK_USER_AGENT) + .json(&serde_json::json!({ + "model": "openai/privacy-filter", + "input": "Classify this text" + })) + .await; + + // Then: the client receives the status-only invalid-request error. + assert_eq!(response.status_code(), 413); + let error: ErrorResponse = response.json(); + assert_eq!(error.error.r#type, "invalid_request_error"); + assert_eq!(error.error.message, "PII detector returned HTTP 413"); + assert!(!error.error.message.contains(UPSTREAM_BODY_SENTINEL)); +} + +#[tokio::test] +async fn test_privacy_classify_upstream_500_returns_502_without_upstream_body() { + // Given: an upstream server failure whose body contains unsafe data. + const UPSTREAM_BODY_SENTINEL: &str = + "UPSTREAM_PRIVACY_BODY_SENTINEL::alice@example.com::123-45-6789"; + let (server, _pool, mock_provider, _db) = setup_test_server_with_pool().await; + setup_privacy_filter_model(&server).await; + let org = setup_org_with_credits(&server, 10_000_000_000i64).await; + let api_key = get_api_key_for_org(&server, org.id).await; + mock_provider + .set_privacy_classify_error_override(Some( + inference_providers::PrivacyClassifyError::HttpError { + status_code: 500, + message: UPSTREAM_BODY_SENTINEL.to_string(), + }, + )) + .await; + + // When: a client calls the real classify route. + let response = server + .post("/v1/privacy/classify") + .add_header("Authorization", format!("Bearer {api_key}")) + .add_header("User-Agent", MOCK_USER_AGENT) + .json(&serde_json::json!({ + "model": "openai/privacy-filter", + "input": "Classify this text" + })) + .await; + + // Then: the server failure remains generic and the body stays private. + assert_eq!(response.status_code(), 502); + let error: ErrorResponse = response.json(); + assert_eq!(error.error.r#type, "server_error"); + assert_eq!( + error.error.message, + "Privacy classify request failed. Please try again later." + ); + assert!(!error.error.message.contains(UPSTREAM_BODY_SENTINEL)); +} + +#[tokio::test] +async fn test_privacy_classify_upstream_404_uses_not_found_mapping() { + // Given: an upstream routing failure whose body contains unsafe data. + const UPSTREAM_BODY_SENTINEL: &str = + "UPSTREAM_PRIVACY_BODY_SENTINEL::alice@example.com::123-45-6789"; + let (server, _pool, mock_provider, _db) = setup_test_server_with_pool().await; + setup_privacy_filter_model(&server).await; + let org = setup_org_with_credits(&server, 10_000_000_000i64).await; + let api_key = get_api_key_for_org(&server, org.id).await; + mock_provider + .set_privacy_classify_error_override(Some( + inference_providers::PrivacyClassifyError::HttpError { + status_code: 404, + message: UPSTREAM_BODY_SENTINEL.to_string(), + }, + )) + .await; + + // When: a client calls the real classify route. + let response = server + .post("/v1/privacy/classify") + .add_header("Authorization", format!("Bearer {api_key}")) + .add_header("User-Agent", MOCK_USER_AGENT) + .json(&serde_json::json!({ + "model": "openai/privacy-filter", + "input": "Classify this text" + })) + .await; + + // Then: the routing failure uses the shared not-found classification. + assert_eq!(response.status_code(), 404); + let error: ErrorResponse = response.json(); + assert_eq!(error.error.r#type, "not_found_error"); + assert_eq!(error.error.message, "PII detector returned HTTP 404"); + assert!(!error.error.message.contains(UPSTREAM_BODY_SENTINEL)); +} + +#[tokio::test] +async fn test_privacy_classify_upstream_302_returns_generic_bad_gateway() { + // Given: an invalid upstream redirect whose body contains unsafe data. + const UPSTREAM_BODY_SENTINEL: &str = + "UPSTREAM_PRIVACY_BODY_SENTINEL::alice@example.com::123-45-6789"; + let (server, _pool, mock_provider, _db) = setup_test_server_with_pool().await; + setup_privacy_filter_model(&server).await; + let org = setup_org_with_credits(&server, 10_000_000_000i64).await; + let api_key = get_api_key_for_org(&server, org.id).await; + mock_provider + .set_privacy_classify_error_override(Some( + inference_providers::PrivacyClassifyError::HttpError { + status_code: 302, + message: UPSTREAM_BODY_SENTINEL.to_string(), + }, + )) + .await; + + // When: a client calls the real classify route. + let response = server + .post("/v1/privacy/classify") + .add_header("Authorization", format!("Bearer {api_key}")) + .add_header("User-Agent", MOCK_USER_AGENT) + .json(&serde_json::json!({ + "model": "openai/privacy-filter", + "input": "Classify this text" + })) + .await; + + // Then: the redirect is masked as a generic gateway failure. + assert_eq!(response.status_code(), 502); + let error: ErrorResponse = response.json(); + assert_eq!(error.error.r#type, "server_error"); + assert_eq!( + error.error.message, + "Privacy classify request failed. Please try again later." + ); + assert!(!error.error.message.contains(UPSTREAM_BODY_SENTINEL)); +} diff --git a/crates/api/tests/e2e_all/privacy_redact.rs b/crates/api/tests/e2e_all/privacy_redact.rs index ce4daa14..e6122a82 100644 --- a/crates/api/tests/e2e_all/privacy_redact.rs +++ b/crates/api/tests/e2e_all/privacy_redact.rs @@ -6,6 +6,7 @@ use crate::common::*; use api::models::{BatchUpdateModelApiRequest, ErrorResponse}; +use axum::http::header::RETRY_AFTER; async fn setup_privacy_filter_model(server: &axum_test::TestServer) -> String { let mut batch = BatchUpdateModelApiRequest::new(); @@ -377,3 +378,211 @@ async fn test_privacy_redact_costs_deducted() { "Privacy redact should bill 10 tokens × 1_000_000 = 10_000_000", ); } + +#[tokio::test] +async fn test_privacy_redact_upstream_429_returns_rate_limit_with_retry_after() { + // Given: an upstream privacy provider that rejects the request with a + // body containing data that must never leave the provider boundary. + const UPSTREAM_BODY_SENTINEL: &str = + "UPSTREAM_PRIVACY_BODY_SENTINEL::alice@example.com::123-45-6789"; + let (server, _pool, mock_provider, _db) = setup_test_server_with_pool().await; + setup_privacy_filter_model(&server).await; + let org = setup_org_with_credits(&server, 10_000_000_000i64).await; + let api_key = get_api_key_for_org(&server, org.id).await; + mock_provider + .set_privacy_classify_error_override(Some( + inference_providers::PrivacyClassifyError::HttpError { + status_code: 429, + message: UPSTREAM_BODY_SENTINEL.to_string(), + }, + )) + .await; + + // When: a client calls the real redact route. + let response = server + .post("/v1/privacy/redact") + .add_header("Authorization", format!("Bearer {api_key}")) + .add_header("User-Agent", MOCK_USER_AGENT) + .json(&serde_json::json!({ + "model": "openai/privacy-filter", + "input": "Redact this text" + })) + .await; + + // Then: retry semantics are explicit and the upstream body is absent. + assert_eq!(response.status_code(), 429); + assert_eq!( + response + .headers() + .get(RETRY_AFTER) + .and_then(|value| value.to_str().ok()), + Some("2") + ); + let error: ErrorResponse = response.json(); + assert_eq!(error.error.r#type, "rate_limit_error"); + assert!(!error.error.message.contains(UPSTREAM_BODY_SENTINEL)); +} + +#[tokio::test] +async fn test_privacy_redact_upstream_413_returns_invalid_request() { + // Given: an oversized upstream response whose body contains unsafe data. + const UPSTREAM_BODY_SENTINEL: &str = + "UPSTREAM_PRIVACY_BODY_SENTINEL::alice@example.com::123-45-6789"; + let (server, _pool, mock_provider, _db) = setup_test_server_with_pool().await; + setup_privacy_filter_model(&server).await; + let org = setup_org_with_credits(&server, 10_000_000_000i64).await; + let api_key = get_api_key_for_org(&server, org.id).await; + mock_provider + .set_privacy_classify_error_override(Some( + inference_providers::PrivacyClassifyError::HttpError { + status_code: 413, + message: UPSTREAM_BODY_SENTINEL.to_string(), + }, + )) + .await; + + // When: a client calls the real redact route. + let response = server + .post("/v1/privacy/redact") + .add_header("Authorization", format!("Bearer {api_key}")) + .add_header("User-Agent", MOCK_USER_AGENT) + .json(&serde_json::json!({ + "model": "openai/privacy-filter", + "input": "Redact this text" + })) + .await; + + // Then: the client receives the status-only invalid-request error. + assert_eq!(response.status_code(), 413); + let error: ErrorResponse = response.json(); + assert_eq!(error.error.r#type, "invalid_request_error"); + assert_eq!(error.error.message, "PII detector returned HTTP 413"); + assert!(!error.error.message.contains(UPSTREAM_BODY_SENTINEL)); +} + +#[tokio::test] +async fn test_privacy_redact_upstream_500_returns_502_without_upstream_body() { + // Given: an upstream server failure whose body contains unsafe data. + const UPSTREAM_BODY_SENTINEL: &str = + "UPSTREAM_PRIVACY_BODY_SENTINEL::alice@example.com::123-45-6789"; + let (server, _pool, mock_provider, _db) = setup_test_server_with_pool().await; + setup_privacy_filter_model(&server).await; + let org = setup_org_with_credits(&server, 10_000_000_000i64).await; + let api_key = get_api_key_for_org(&server, org.id).await; + mock_provider + .set_privacy_classify_error_override(Some( + inference_providers::PrivacyClassifyError::HttpError { + status_code: 500, + message: UPSTREAM_BODY_SENTINEL.to_string(), + }, + )) + .await; + + // When: a client calls the real redact route. + let response = server + .post("/v1/privacy/redact") + .add_header("Authorization", format!("Bearer {api_key}")) + .add_header("User-Agent", MOCK_USER_AGENT) + .json(&serde_json::json!({ + "model": "openai/privacy-filter", + "input": "Redact this text" + })) + .await; + + // Then: the server failure remains generic and the body stays private. + assert_eq!(response.status_code(), 502); + let error: ErrorResponse = response.json(); + assert_eq!(error.error.r#type, "server_error"); + assert_eq!( + error.error.message, + "Privacy redact request failed. Please try again later." + ); + assert!(!error.error.message.contains(UPSTREAM_BODY_SENTINEL)); +} + +#[tokio::test] +async fn test_privacy_redact_upstream_407_returns_generic_server_error() { + // Given: an upstream proxy-auth failure whose body contains unsafe data. + const UPSTREAM_BODY_SENTINEL: &str = + "UPSTREAM_PRIVACY_BODY_SENTINEL::alice@example.com::123-45-6789"; + let (server, _pool, mock_provider, _db) = setup_test_server_with_pool().await; + setup_privacy_filter_model(&server).await; + let org = setup_org_with_credits(&server, 10_000_000_000i64).await; + let api_key = get_api_key_for_org(&server, org.id).await; + mock_provider + .set_privacy_classify_error_override(Some( + inference_providers::PrivacyClassifyError::HttpError { + status_code: 407, + message: UPSTREAM_BODY_SENTINEL.to_string(), + }, + )) + .await; + + // When: a client calls the real redact route. + let response = server + .post("/v1/privacy/redact") + .add_header("Authorization", format!("Bearer {api_key}")) + .add_header("User-Agent", MOCK_USER_AGENT) + .json(&serde_json::json!({ + "model": "openai/privacy-filter", + "input": "Redact this text" + })) + .await; + + // Then: the infrastructure fault is masked as a generic server error. + assert_eq!(response.status_code(), 500); + let error: ErrorResponse = response.json(); + assert_eq!(error.error.r#type, "server_error"); + assert_eq!( + error.error.message, + "Privacy redact request failed. Please try again later." + ); + assert!(!error.error.message.contains(UPSTREAM_BODY_SENTINEL)); +} + +#[tokio::test] +async fn test_privacy_redact_upstream_503_returns_overloaded_with_default_retry_after() { + // Given: an overloaded upstream whose body contains unsafe data. + const UPSTREAM_BODY_SENTINEL: &str = + "UPSTREAM_PRIVACY_BODY_SENTINEL::alice@example.com::123-45-6789"; + let (server, _pool, mock_provider, _db) = setup_test_server_with_pool().await; + setup_privacy_filter_model(&server).await; + let org = setup_org_with_credits(&server, 10_000_000_000i64).await; + let api_key = get_api_key_for_org(&server, org.id).await; + mock_provider + .set_privacy_classify_error_override(Some( + inference_providers::PrivacyClassifyError::HttpError { + status_code: 503, + message: UPSTREAM_BODY_SENTINEL.to_string(), + }, + )) + .await; + + // When: a client calls the real redact route. + let response = server + .post("/v1/privacy/redact") + .add_header("Authorization", format!("Bearer {api_key}")) + .add_header("User-Agent", MOCK_USER_AGENT) + .json(&serde_json::json!({ + "model": "openai/privacy-filter", + "input": "Redact this text" + })) + .await; + + // Then: overload uses the canonical 429 response and middleware backoff. + assert_eq!(response.status_code(), 429); + assert_eq!( + response + .headers() + .get(RETRY_AFTER) + .and_then(|value| value.to_str().ok()), + Some("2") + ); + let error: ErrorResponse = response.json(); + assert_eq!(error.error.r#type, "service_overloaded"); + assert_eq!( + error.error.message, + "All inference backends are overloaded. Please retry with exponential backoff." + ); + assert!(!error.error.message.contains(UPSTREAM_BODY_SENTINEL)); +} diff --git a/crates/inference_providers/src/mock.rs b/crates/inference_providers/src/mock.rs index 0aaf165d..b33f844e 100644 --- a/crates/inference_providers/src/mock.rs +++ b/crates/inference_providers/src/mock.rs @@ -657,6 +657,8 @@ struct MockConfig { embedding_error_override: Option, /// When set, all audio transcription calls return this error instead of a response. audio_transcription_error_override: Option, + /// When set, all privacy classify calls return this error instead of a response. + privacy_classify_error_override: Option, } /// Builder for configuring a single expectation @@ -688,6 +690,7 @@ pub struct MockProvider { last_chat_params: Arc>>, /// When true, get_attestation_report returns an error (simulates blocked/broken backend) fail_attestation: Arc, + privacy_classify_call_count: Arc, /// Trust tier reported by [`InferenceProvider::tier`]; defaults to /// `NonAttested`. Set via [`MockProvider::with_tier`] to exercise tiered /// provider selection (e.g. a `Near` primary with an `Attested3p` fallback). @@ -730,9 +733,11 @@ impl MockProvider { stream_error_override: None, embedding_error_override: None, audio_transcription_error_override: None, + privacy_classify_error_override: None, })), last_chat_params: Arc::new(Mutex::new(None)), fail_attestation: Arc::new(std::sync::atomic::AtomicBool::new(false)), + privacy_classify_call_count: Arc::new(std::sync::atomic::AtomicUsize::new(0)), tier: crate::ProviderTier::NonAttested, provider_source: crate::ProviderSource::External, supports_streaming: true, @@ -755,9 +760,11 @@ impl MockProvider { stream_error_override: None, embedding_error_override: None, audio_transcription_error_override: None, + privacy_classify_error_override: None, })), last_chat_params: Arc::new(Mutex::new(None)), fail_attestation: Arc::new(std::sync::atomic::AtomicBool::new(false)), + privacy_classify_call_count: Arc::new(std::sync::atomic::AtomicUsize::new(0)), tier: crate::ProviderTier::NonAttested, provider_source: crate::ProviderSource::External, supports_streaming: true, @@ -778,9 +785,11 @@ impl MockProvider { stream_error_override: None, embedding_error_override: None, audio_transcription_error_override: None, + privacy_classify_error_override: None, })), last_chat_params: Arc::new(Mutex::new(None)), fail_attestation: Arc::new(std::sync::atomic::AtomicBool::new(false)), + privacy_classify_call_count: Arc::new(std::sync::atomic::AtomicUsize::new(0)), tier: crate::ProviderTier::NonAttested, provider_source: crate::ProviderSource::External, supports_streaming: true, @@ -901,6 +910,18 @@ impl MockProvider { config.audio_transcription_error_override = error; } + /// Override the privacy classify response with an error. Pass `None` to clear. + pub async fn set_privacy_classify_error_override(&self, error: Option) { + let mut config = self.config.lock().await; + config.privacy_classify_error_override = error; + } + + /// Return how many privacy classification calls this mock received. + pub fn privacy_classify_call_count(&self) -> usize { + self.privacy_classify_call_count + .load(std::sync::atomic::Ordering::Relaxed) + } + /// Generate a completion ID fn generate_id(&self) -> String { use std::collections::hash_map::DefaultHasher; @@ -1450,6 +1471,15 @@ impl crate::InferenceProvider for MockProvider { body: bytes::Bytes, _extra: std::collections::HashMap, ) -> Result { + self.privacy_classify_call_count + .fetch_add(1, std::sync::atomic::Ordering::Relaxed); + { + let config = self.config.lock().await; + if let Some(ref error) = config.privacy_classify_error_override { + return Err(error.clone()); + } + } + // Echo the requested model so round-trip assertions in tests are meaningful. let parsed: serde_json::Value = serde_json::from_slice(&body).unwrap_or(serde_json::Value::Null); diff --git a/crates/inference_providers/src/models.rs b/crates/inference_providers/src/models.rs index 0d544063..dbe47fef 100644 --- a/crates/inference_providers/src/models.rs +++ b/crates/inference_providers/src/models.rs @@ -1883,7 +1883,7 @@ pub enum EmbeddingError { HttpError { status_code: u16, message: String }, } -#[derive(Debug, thiserror::Error)] +#[derive(Debug, Clone, thiserror::Error)] pub enum PrivacyClassifyError { #[error("Privacy classify request failed: {0}")] RequestFailed(String), diff --git a/crates/services/src/inference_provider_pool/mod.rs b/crates/services/src/inference_provider_pool/mod.rs index b8265055..c658b70c 100644 --- a/crates/services/src/inference_provider_pool/mod.rs +++ b/crates/services/src/inference_provider_pool/mod.rs @@ -4070,6 +4070,25 @@ impl InferenceProviderPool { return Ok(response); } Err(e) => { + let retryable_error = match e { + inference_providers::PrivacyClassifyError::HttpError { + status_code, + .. + } if (400..=499).contains(&status_code) + && !matches!(status_code, 408 | 429) => + { + return Err(inference_providers::PrivacyClassifyError::HttpError { + status_code, + message: format!("PII detector returned HTTP {status_code}"), + }); + } + error @ inference_providers::PrivacyClassifyError::HttpError { .. } => { + error + } + error @ inference_providers::PrivacyClassifyError::RequestFailed(_) => { + error + } + }; // Privacy-filter error messages may embed the upstream // response body (HttpError carries the verbatim text). // A misbehaving filter that echoes its input would @@ -4077,10 +4096,10 @@ impl InferenceProviderPool { // Log only the category + status code. tracing::warn!( model = %model, - error_category = %Self::privacy_classify_error_category(&e), + error_category = %Self::privacy_classify_error_category(&retryable_error), "Privacy classify failed with provider, trying next" ); - last_error = Some(e); + last_error = Some(retryable_error); } } } @@ -4088,21 +4107,22 @@ impl InferenceProviderPool { // Final user-facing error: only the status code escapes; no // upstream response body. (`sanitize_error_message` would still // include the body via Display, so we route around it.) - let error_msg = last_error - .as_ref() - .map(|e| match e { - inference_providers::PrivacyClassifyError::HttpError { status_code, .. } => { - format!("PII detector returned HTTP {status_code}") - } - inference_providers::PrivacyClassifyError::RequestFailed(_) => { - "PII detector unreachable".to_string() + Err(match last_error { + Some(inference_providers::PrivacyClassifyError::HttpError { status_code, .. }) => { + inference_providers::PrivacyClassifyError::HttpError { + status_code, + message: format!("PII detector returned HTTP {status_code}"), } - }) - .unwrap_or_else(|| "No providers available for privacy classify".to_string()); - - Err(inference_providers::PrivacyClassifyError::RequestFailed( - error_msg, - )) + } + Some(inference_providers::PrivacyClassifyError::RequestFailed(_)) => { + inference_providers::PrivacyClassifyError::RequestFailed( + "PII detector unreachable".to_string(), + ) + } + None => inference_providers::PrivacyClassifyError::RequestFailed( + "No providers available for privacy classify".to_string(), + ), + }) } pub async fn score( @@ -7879,6 +7899,109 @@ mod tests { (pool, model_id) } + #[tokio::test] + async fn privacy_classify_preserves_http_status_and_discards_response_body() { + // Given: a provider error whose message contains an unsafe upstream body. + const UPSTREAM_BODY_SENTINEL: &str = + "UPSTREAM_PRIVACY_BODY_SENTINEL::alice@example.com::123-45-6789"; + let pool = InferenceProviderPool::new(None, ExternalProvidersConfig::default()); + let mock_provider = Arc::new(inference_providers::mock::MockProvider::new()); + mock_provider + .set_privacy_classify_error_override(Some( + inference_providers::PrivacyClassifyError::HttpError { + status_code: 413, + message: UPSTREAM_BODY_SENTINEL.to_string(), + }, + )) + .await; + let model_id = "openai/privacy-filter"; + pool.register_provider(model_id.to_string(), mock_provider) + .await; + + // When: the pool exhausts its privacy providers. + let result = pool + .privacy_classify( + model_id, + bytes::Bytes::from_static(b"{}"), + std::collections::HashMap::new(), + ) + .await; + + // Then: only the typed status and a synthesized message escape. + match result { + Err(inference_providers::PrivacyClassifyError::HttpError { + status_code, + message, + }) => { + assert_eq!(status_code, 413); + assert_eq!(message, "PII detector returned HTTP 413"); + assert!(!message.contains(UPSTREAM_BODY_SENTINEL)); + } + Err(other) => panic!("Expected sanitized HttpError, got {other:?}"), + Ok(_) => panic!("Expected privacy classify to fail"), + } + } + + #[tokio::test] + async fn privacy_classify_non_retryable_client_error_short_circuits_later_providers() { + // Given: provider A rejects the request as too large, while provider B + // would replace that actionable status with a transport failure. + const FIRST_UPSTREAM_BODY_SENTINEL: &str = + "FIRST_UPSTREAM_PRIVACY_BODY_SENTINEL::alice@example.com"; + const SECOND_UPSTREAM_BODY_SENTINEL: &str = + "SECOND_UPSTREAM_PRIVACY_BODY_SENTINEL::123-45-6789"; + let pool = InferenceProviderPool::new(None, ExternalProvidersConfig::default()); + let first_provider = Arc::new(inference_providers::mock::MockProvider::new()); + first_provider + .set_privacy_classify_error_override(Some( + inference_providers::PrivacyClassifyError::HttpError { + status_code: 413, + message: FIRST_UPSTREAM_BODY_SENTINEL.to_string(), + }, + )) + .await; + let second_provider = Arc::new(inference_providers::mock::MockProvider::new()); + second_provider + .set_privacy_classify_error_override(Some( + inference_providers::PrivacyClassifyError::RequestFailed( + SECOND_UPSTREAM_BODY_SENTINEL.to_string(), + ), + )) + .await; + let model_id = "openai/privacy-filter"; + pool.register_providers(vec![ + (model_id.to_string(), first_provider.clone()), + (model_id.to_string(), second_provider.clone()), + ]) + .await; + + // When: privacy classification tries the ordered provider set. + let result = pool + .privacy_classify( + model_id, + bytes::Bytes::from_static(b"{}"), + std::collections::HashMap::new(), + ) + .await; + + // Then: the first non-retryable status wins without consulting provider B. + match result { + Err(inference_providers::PrivacyClassifyError::HttpError { + status_code, + message, + }) => { + assert_eq!(status_code, 413); + assert_eq!(message, "PII detector returned HTTP 413"); + assert!(!message.contains(FIRST_UPSTREAM_BODY_SENTINEL)); + assert!(!message.contains(SECOND_UPSTREAM_BODY_SENTINEL)); + } + Err(other) => panic!("Expected sanitized HttpError, got {other:?}"), + Ok(_) => panic!("Expected privacy classify to fail"), + } + assert_eq!(first_provider.privacy_classify_call_count(), 1); + assert_eq!(second_provider.privacy_classify_call_count(), 0); + } + #[tokio::test] async fn test_4xx_error_does_not_retry() { let (pool, model_id) = pool_with_mock_provider().await;