fix: enforce provider upstream stream policy

This commit is contained in:
stabey
2026-05-09 12:30:56 +08:00
parent 8958bf5e08
commit 7e5a08e09d
29 changed files with 1518 additions and 118 deletions
@@ -37,6 +37,52 @@ pub(crate) fn force_upstream_streaming_for_provider(
force_upstream_streaming_for_provider_impl(provider_type, provider_api_format)
}
pub(crate) fn resolve_upstream_is_stream_for_provider(
endpoint_config: Option<&serde_json::Value>,
provider_type: &str,
provider_api_format: &str,
client_is_stream: bool,
hard_requires_streaming: bool,
) -> bool {
let hard_requires_streaming = hard_requires_streaming
|| force_upstream_streaming_for_provider(provider_type, provider_api_format);
aether_ai_formats::resolve_upstream_is_stream_from_endpoint_config(
endpoint_config,
client_is_stream,
hard_requires_streaming,
)
}
pub(crate) fn endpoint_config_forces_body_stream_field(
endpoint_config: Option<&serde_json::Value>,
) -> bool {
aether_ai_formats::endpoint_config_forces_upstream_stream_policy(endpoint_config)
}
pub(crate) fn request_requires_body_stream_field(
body_json: &serde_json::Value,
force_body_stream_field: bool,
) -> bool {
force_body_stream_field
|| body_json
.as_object()
.is_some_and(|object| object.contains_key("stream"))
}
pub(crate) fn enforce_provider_body_stream_policy(
provider_request_body: &mut serde_json::Value,
provider_api_format: &str,
upstream_is_stream: bool,
require_body_stream_field: bool,
) {
aether_ai_formats::enforce_request_body_stream_field(
provider_request_body,
provider_api_format,
upstream_is_stream,
require_body_stream_field,
);
}
pub(crate) fn extract_standard_requested_model(body_json: &serde_json::Value) -> Option<String> {
aether_ai_serving::extract_ai_standard_requested_model(body_json)
}
@@ -56,8 +102,10 @@ pub(crate) fn extract_requested_model_from_request(
#[cfg(test)]
mod tests {
use super::{
endpoint_config_forces_body_stream_field, enforce_provider_body_stream_policy,
extract_requested_model_from_request, extract_standard_requested_model,
force_upstream_streaming_for_provider, RequestedModelFamily,
force_upstream_streaming_for_provider, resolve_upstream_is_stream_for_provider,
RequestedModelFamily,
};
use axum::http::Request;
use serde_json::json;
@@ -90,6 +138,67 @@ mod tests {
));
}
#[test]
fn resolves_endpoint_upstream_stream_policy_with_provider_hard_constraints() {
assert!(resolve_upstream_is_stream_for_provider(
Some(&json!({"upstream_stream_policy": "force_stream"})),
"openai",
"openai:chat",
false,
false,
));
assert!(!resolve_upstream_is_stream_for_provider(
Some(&json!({"upstream_stream_policy": "force_non_stream"})),
"openai",
"openai:chat",
true,
false,
));
assert!(resolve_upstream_is_stream_for_provider(
Some(&json!({"upstream_stream_policy": "auto"})),
"openai",
"openai:chat",
true,
false,
));
assert!(resolve_upstream_is_stream_for_provider(
Some(&json!({"upstream_stream_policy": "force_non_stream"})),
"codex",
"openai:responses",
true,
false,
));
}
#[test]
fn enforces_provider_body_stream_policy_for_body_and_streamless_formats() {
let mut openai_chat = json!({"stream": true});
enforce_provider_body_stream_policy(&mut openai_chat, "openai:chat", false, false);
assert_eq!(openai_chat.get("stream"), Some(&json!(false)));
let mut ordinary_sync = json!({"messages": []});
enforce_provider_body_stream_policy(&mut ordinary_sync, "openai:chat", false, false);
assert!(ordinary_sync.get("stream").is_none());
let mut compact = json!({"stream": true});
enforce_provider_body_stream_policy(&mut compact, "openai:responses:compact", true, true);
assert!(compact.get("stream").is_none());
}
#[test]
fn detects_endpoint_configs_that_force_body_stream_field() {
assert!(endpoint_config_forces_body_stream_field(Some(
&json!({"upstream_stream_policy": "force_stream"})
)));
assert!(endpoint_config_forces_body_stream_field(Some(
&json!({"upstream_stream_policy": "force_non_stream"})
)));
assert!(!endpoint_config_forces_body_stream_field(Some(
&json!({"upstream_stream_policy": "auto"})
)));
assert!(!endpoint_config_forces_body_stream_field(None));
}
#[test]
fn extracts_standard_requested_model_from_request_body() {
let requested_model =