diff --git a/crates/api/tests/e2e_all/first_stream_event.rs b/crates/api/tests/e2e_all/first_stream_event.rs new file mode 100644 index 000000000..25d975194 --- /dev/null +++ b/crates/api/tests/e2e_all/first_stream_event.rs @@ -0,0 +1,82 @@ +use crate::common::*; +use inference_providers::mock::{RequestMatcher, ResponseTemplate}; + +#[tokio::test] +async fn first_upstream_error_is_returned_as_http_error_before_sse_starts() { + // Given + let (server, _pool, mock, _db) = setup_test_server_with_pool().await; + let model = setup_qwen_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.set_stream_error_override(Some(inference_providers::CompletionError::HttpError { + status_code: 400, + message: "Grammar error: Unimplemented keys: [\"uniqueItems\"]".to_string(), + is_external: false, + })) + .await; + + // When + let response = server + .post("/v1/chat/completions") + .add_header("Authorization", format!("Bearer {api_key}")) + .json(&serde_json::json!({ + "model": model, + "messages": [{"role": "user", "content": "Return JSON."}], + "stream": true + })) + .await; + + // Then + assert_eq!(response.status_code(), 400, "{}", response.text()); + assert_eq!( + response + .headers() + .get("content-type") + .and_then(|value| value.to_str().ok()), + Some("application/json") + ); + let body = response.json::(); + assert_eq!(body["error"]["type"], "invalid_request_error"); + assert!(body["error"]["message"] + .as_str() + .is_some_and(|message| message.contains("uniqueItems"))); +} + +#[tokio::test] +async fn normal_first_upstream_chunk_remains_first_and_unmodified() { + // Given + let (server, _pool, mock, _db) = setup_test_server_with_pool().await; + let model = setup_qwen_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.when(RequestMatcher::Any) + .respond_with(ResponseTemplate::new("first second")) + .await; + + // When + let response = server + .post("/v1/chat/completions") + .add_header("Authorization", format!("Bearer {api_key}")) + .json(&serde_json::json!({ + "model": model, + "messages": [{"role": "user", "content": "Stream two words."}], + "stream": true, + "stream_options": {"continuous_usage_stats": true} + })) + .await; + + // Then + assert_eq!(response.status_code(), 200, "{}", response.text()); + let response_text = response.text(); + let first = response_text + .lines() + .find_map(|line| line.strip_prefix("data: ")) + .expect("stream should contain a first data event"); + let chunk = serde_json::from_str::(first).expect("valid first chunk"); + assert_eq!(chunk["choices"][0]["delta"]["content"], "first"); + assert_eq!( + chunk["mock_upstream_only_field"], "dropped-by-typed-parse", + "the provider's raw first chunk must bypass typed re-serialization" + ); + assert!(chunk.get("error").is_none()); +} diff --git a/crates/api/tests/e2e_all/main.rs b/crates/api/tests/e2e_all/main.rs index 0f5e08ad6..83b38793c 100644 --- a/crates/api/tests/e2e_all/main.rs +++ b/crates/api/tests/e2e_all/main.rs @@ -41,6 +41,7 @@ mod error_msg; mod external_providers; mod feature_requests; mod files; +mod first_stream_event; mod function_tools; mod general; mod glm52_tier_routing; diff --git a/crates/services/src/inference_provider_pool/mod.rs b/crates/services/src/inference_provider_pool/mod.rs index 132725e63..6f44f22cc 100644 --- a/crates/services/src/inference_provider_pool/mod.rs +++ b/crates/services/src/inference_provider_pool/mod.rs @@ -3302,25 +3302,33 @@ impl InferenceProviderPool { } } - if let Some(Ok(event)) = peekable.peek().await { - if let Some(inference_providers::StreamChunk::Chat(chat_chunk)) = &event.chunk { - let chat_id = chat_chunk.id.clone(); - tracing::info!( - chat_id = %chat_id, - "Storing chat_id mapping for streaming completion" - ); - // Pin the dedicated TLS connection so signature fetches - // reuse the same connection that served this completion. - provider.pin_chat_connection(&request_hash, &chat_id); - pinned = true; - self.store_chat_id_mapping(chat_id, provider.clone()).await; + let first_error = match peekable.peek().await { + Some(Ok(event)) => { + if let Some(inference_providers::StreamChunk::Chat(chat_chunk)) = &event.chunk { + let chat_id = chat_chunk.id.clone(); + tracing::info!( + chat_id = %chat_id, + "Storing chat_id mapping for streaming completion" + ); + // Pin the dedicated TLS connection so signature fetches + // reuse the same connection that served this completion. + provider.pin_chat_connection(&request_hash, &chat_id); + pinned = true; + self.store_chat_id_mapping(chat_id, provider.clone()).await; + } + None } - } + Some(Err(error)) => Some(error.clone()), + None => None, + }; if !pinned { // Clean up orphaned pending client when peek fails or yields no chat_id provider.pin_chat_connection(&request_hash, ""); provider.unpin_chat_connection(""); } + if let Some(error) = first_error { + return Err(error); + } let stream: StreamingResult = if leading_control.is_empty() { Box::pin(peekable) } else { @@ -6130,6 +6138,84 @@ mod tests { assert!(pool.get_provider_by_chat_id(&chat_id).await.is_some()); } + #[tokio::test] + async fn test_first_stream_error_is_returned_before_stream() { + use inference_providers::mock::MockProvider; + + // Given + let pool = InferenceProviderPool::new(None, ExternalProvidersConfig::default()); + let mock_provider = Arc::new(MockProvider::new()); + let model_id = "Qwen/Qwen3-30B-A3B-Instruct-2507".to_string(); + mock_provider + .set_stream_error_override(Some(CompletionError::HttpError { + status_code: 400, + message: "Grammar error: unsupported schema keyword".to_string(), + is_external: true, + })) + .await; + pool.register_provider(model_id.clone(), mock_provider.clone()) + .await; + let params = inference_providers::ChatCompletionParams { + model: model_id, + messages: vec![inference_providers::ChatMessage { + role: inference_providers::MessageRole::User, + content: Some(serde_json::Value::String("Hello".to_string())), + name: None, + tool_call_id: None, + tool_calls: None, + }], + max_tokens: None, + temperature: None, + top_p: None, + stop: None, + stream: Some(true), + tools: None, + max_completion_tokens: None, + n: None, + frequency_penalty: None, + presence_penalty: None, + logit_bias: None, + logprobs: None, + top_logprobs: None, + user: None, + seed: None, + tool_choice: None, + parallel_tool_calls: None, + metadata: None, + store: None, + stream_options: None, + service_tier: None, + modalities: None, + original_request: None, + extra: std::collections::HashMap::new(), + }; + + // When + let result = pool + .chat_completion_stream( + params, + "test-request-hash".to_string(), + ChatRoutingHints::default(), + ) + .await; + + // Then + match result { + Err(CompletionError::HttpError { + status_code, + message, + is_external, + }) => { + assert_eq!(status_code, 400); + assert_eq!(message, "Grammar error: unsupported schema keyword"); + assert!(is_external); + } + Err(other) => panic!("Expected HttpError, got {other:?}"), + Ok(_) => panic!("Expected the pool to return the first stream error"), + } + assert_eq!(mock_provider.unpinned_chat_ids(), vec![String::new()]); + } + // ==================== Provider Tests ==================== #[tokio::test]