Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
82 changes: 82 additions & 0 deletions crates/api/tests/e2e_all/first_stream_event.rs
Original file line number Diff line number Diff line change
@@ -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::<serde_json::Value>();
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::<serde_json::Value>(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());
}
1 change: 1 addition & 0 deletions crates/api/tests/e2e_all/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
112 changes: 99 additions & 13 deletions crates/services/src/inference_provider_pool/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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]
Expand Down
Loading