mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 18:59:50 +08:00
feat(gateway): harden provider request execution
Preserve exact request payloads and model client surface and API operation explicitly. Add Anthropic compatibility profiles, bounded stream commitment, and scoped OAuth retry behavior across provider transports.
This commit is contained in:
@@ -202,6 +202,16 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalStandardSyncAttemptSour
|
||||
Ok(drained)
|
||||
}
|
||||
|
||||
async fn skip_credential(&mut self, key_id: &str) -> Result<(), GatewayError> {
|
||||
self.candidates.skip_credential(key_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn skip_endpoint(&mut self, endpoint_id: &str) -> Result<(), GatewayError> {
|
||||
self.candidates.skip_endpoint(endpoint_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
|
||||
self.candidates.skip_provider(provider_id);
|
||||
Ok(())
|
||||
@@ -235,6 +245,16 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalStandardStreamAttempt
|
||||
Ok(drained)
|
||||
}
|
||||
|
||||
async fn skip_credential(&mut self, key_id: &str) -> Result<(), GatewayError> {
|
||||
self.candidates.skip_credential(key_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn skip_endpoint(&mut self, endpoint_id: &str) -> Result<(), GatewayError> {
|
||||
self.candidates.skip_endpoint(endpoint_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
|
||||
self.candidates.skip_provider(provider_id);
|
||||
Ok(())
|
||||
|
||||
@@ -373,6 +373,8 @@ mod tests {
|
||||
auth_snapshot: sample_auth_snapshot(),
|
||||
required_capabilities: None,
|
||||
request_auth_channel: None,
|
||||
client_surface: None,
|
||||
gateway_credential_carrier: None,
|
||||
client_session_affinity: None,
|
||||
routing_policy: None,
|
||||
routing_trace_seed: None,
|
||||
|
||||
@@ -1,12 +1,10 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use aether_contracts::RequestBody;
|
||||
|
||||
use super::{
|
||||
augment_sync_report_context, build_ai_execution_plan_from_decision,
|
||||
generic_decision_missing_exact_provider_request, take_ai_decision_plan_core,
|
||||
take_ai_upstream_auth_pair, take_non_empty_string, AiExecutionPlanFromDecisionParts,
|
||||
AiStreamAttempt, AiSyncAttempt,
|
||||
generic_decision_missing_exact_provider_request, resolve_ai_passthrough_sync_request_body,
|
||||
take_ai_decision_plan_core, take_ai_upstream_auth_pair, take_non_empty_string,
|
||||
AiExecutionPlanFromDecisionParts, AiStreamAttempt, AiSyncAttempt,
|
||||
};
|
||||
use crate::ai_serving::transport::{
|
||||
build_standard_plan_fallback_headers, StandardPlanFallbackAcceptPolicy,
|
||||
@@ -61,6 +59,10 @@ pub(crate) fn build_gemini_sync_plan_from_decision(
|
||||
&provider_request_headers,
|
||||
&provider_request_body_value,
|
||||
)?;
|
||||
let request_body = resolve_ai_passthrough_sync_request_body(
|
||||
Some(provider_request_body_value),
|
||||
payload.provider_request_body_base64.take(),
|
||||
);
|
||||
let stream = payload.upstream_is_stream;
|
||||
let plan = build_ai_execution_plan_from_decision(
|
||||
&mut payload,
|
||||
@@ -70,7 +72,7 @@ pub(crate) fn build_gemini_sync_plan_from_decision(
|
||||
url,
|
||||
headers: std::mem::take(&mut provider_request_headers),
|
||||
content_type,
|
||||
body: RequestBody::from_json(provider_request_body_value),
|
||||
body: request_body,
|
||||
stream,
|
||||
},
|
||||
);
|
||||
@@ -129,6 +131,10 @@ pub(crate) fn build_gemini_stream_plan_from_decision(
|
||||
&provider_request_headers,
|
||||
&provider_request_body_value,
|
||||
)?;
|
||||
let request_body = resolve_ai_passthrough_sync_request_body(
|
||||
Some(provider_request_body_value),
|
||||
payload.provider_request_body_base64.take(),
|
||||
);
|
||||
let plan = build_ai_execution_plan_from_decision(
|
||||
&mut payload,
|
||||
AiExecutionPlanFromDecisionParts {
|
||||
@@ -137,7 +143,7 @@ pub(crate) fn build_gemini_stream_plan_from_decision(
|
||||
url,
|
||||
headers: std::mem::take(&mut provider_request_headers),
|
||||
content_type,
|
||||
body: RequestBody::from_json(provider_request_body_value),
|
||||
body: request_body,
|
||||
stream: true,
|
||||
},
|
||||
);
|
||||
|
||||
@@ -84,6 +84,7 @@ pub(crate) fn build_standard_upstream_url(
|
||||
upstream_is_stream,
|
||||
parts.uri.query(),
|
||||
None,
|
||||
None,
|
||||
provider_request_body,
|
||||
)
|
||||
}
|
||||
|
||||
@@ -2199,6 +2199,8 @@ mod tests {
|
||||
auth_snapshot: sample_auth_snapshot(),
|
||||
required_capabilities: None,
|
||||
request_auth_channel: None,
|
||||
client_surface: None,
|
||||
gateway_credential_carrier: None,
|
||||
client_session_affinity: None,
|
||||
routing_policy: None,
|
||||
routing_trace_seed: None,
|
||||
|
||||
@@ -139,6 +139,20 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiChatStreamAttem
|
||||
Ok(drained)
|
||||
}
|
||||
|
||||
async fn skip_credential(&mut self, key_id: &str) -> Result<(), GatewayError> {
|
||||
self.prefetched_attempts
|
||||
.retain(|attempt| attempt.eligible.candidate.key_id != key_id);
|
||||
self.candidates.skip_credential(key_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn skip_endpoint(&mut self, endpoint_id: &str) -> Result<(), GatewayError> {
|
||||
self.prefetched_attempts
|
||||
.retain(|attempt| attempt.eligible.candidate.endpoint_id != endpoint_id);
|
||||
self.candidates.skip_endpoint(endpoint_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
|
||||
self.prefetched_attempts
|
||||
.retain(|attempt| attempt.eligible.candidate.provider_id != provider_id);
|
||||
|
||||
@@ -117,6 +117,16 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiChatSyncAttemptSo
|
||||
Ok(drained)
|
||||
}
|
||||
|
||||
async fn skip_credential(&mut self, key_id: &str) -> Result<(), GatewayError> {
|
||||
self.candidates.skip_credential(key_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn skip_endpoint(&mut self, endpoint_id: &str) -> Result<(), GatewayError> {
|
||||
self.candidates.skip_endpoint(endpoint_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
|
||||
self.candidates.skip_provider(provider_id);
|
||||
Ok(())
|
||||
|
||||
@@ -186,6 +186,16 @@ impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiResponsesSyncAtte
|
||||
Ok(drained)
|
||||
}
|
||||
|
||||
async fn skip_credential(&mut self, key_id: &str) -> Result<(), GatewayError> {
|
||||
self.candidates.skip_credential(key_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn skip_endpoint(&mut self, endpoint_id: &str) -> Result<(), GatewayError> {
|
||||
self.candidates.skip_endpoint(endpoint_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
|
||||
self.candidates.skip_provider(provider_id);
|
||||
Ok(())
|
||||
@@ -219,6 +229,16 @@ impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiResponsesStream
|
||||
Ok(drained)
|
||||
}
|
||||
|
||||
async fn skip_credential(&mut self, key_id: &str) -> Result<(), GatewayError> {
|
||||
self.candidates.skip_credential(key_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn skip_endpoint(&mut self, endpoint_id: &str) -> Result<(), GatewayError> {
|
||||
self.candidates.skip_endpoint(endpoint_id);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn skip_provider(&mut self, provider_id: &str) -> Result<(), GatewayError> {
|
||||
self.candidates.skip_provider(provider_id);
|
||||
Ok(())
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use aether_contracts::RequestBody;
|
||||
|
||||
use super::{
|
||||
augment_sync_report_context, build_ai_execution_plan_from_decision, take_ai_decision_plan_core,
|
||||
augment_sync_report_context, build_ai_execution_plan_from_decision,
|
||||
resolve_ai_passthrough_sync_request_body, take_ai_decision_plan_core,
|
||||
take_ai_upstream_auth_pair, take_non_empty_string, AiExecutionPlanFromDecisionParts,
|
||||
AiStreamAttempt, AiSyncAttempt,
|
||||
};
|
||||
@@ -63,6 +62,10 @@ pub(crate) fn build_standard_sync_plan_from_decision(
|
||||
&provider_request_headers,
|
||||
&provider_request_body_value,
|
||||
)?;
|
||||
let request_body = resolve_ai_passthrough_sync_request_body(
|
||||
Some(provider_request_body_value),
|
||||
payload.provider_request_body_base64.take(),
|
||||
);
|
||||
let stream = payload.upstream_is_stream;
|
||||
let plan = build_ai_execution_plan_from_decision(
|
||||
&mut payload,
|
||||
@@ -72,7 +75,7 @@ pub(crate) fn build_standard_sync_plan_from_decision(
|
||||
url,
|
||||
headers: std::mem::take(&mut provider_request_headers),
|
||||
content_type,
|
||||
body: RequestBody::from_json(provider_request_body_value),
|
||||
body: request_body,
|
||||
stream,
|
||||
},
|
||||
);
|
||||
@@ -146,6 +149,10 @@ pub(crate) fn build_standard_stream_plan_from_decision(
|
||||
&provider_request_headers,
|
||||
&provider_request_body_value,
|
||||
)?;
|
||||
let request_body = resolve_ai_passthrough_sync_request_body(
|
||||
Some(provider_request_body_value),
|
||||
payload.provider_request_body_base64.take(),
|
||||
);
|
||||
let stream = payload.upstream_is_stream;
|
||||
let plan = build_ai_execution_plan_from_decision(
|
||||
&mut payload,
|
||||
@@ -155,7 +162,7 @@ pub(crate) fn build_standard_stream_plan_from_decision(
|
||||
url,
|
||||
headers: std::mem::take(&mut provider_request_headers),
|
||||
content_type,
|
||||
body: RequestBody::from_json(provider_request_body_value),
|
||||
body: request_body,
|
||||
stream,
|
||||
},
|
||||
);
|
||||
@@ -166,3 +173,88 @@ pub(crate) fn build_standard_stream_plan_from_decision(
|
||||
report_context,
|
||||
}))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use aether_contracts::{ExecutionResponseBodyMode, EXECUTION_RESPONSE_BODY_MODE_HEADER};
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
build_standard_stream_plan_from_decision, build_standard_sync_plan_from_decision,
|
||||
AiExecutionDecision,
|
||||
};
|
||||
|
||||
fn decision_with_raw_body(upstream_is_stream: bool) -> AiExecutionDecision {
|
||||
serde_json::from_value(json!({
|
||||
"action": if upstream_is_stream { "stream" } else { "sync" },
|
||||
"request_id": "req-raw",
|
||||
"provider_id": "provider-raw",
|
||||
"endpoint_id": "endpoint-raw",
|
||||
"key_id": "key-raw",
|
||||
"upstream_url": "https://api.anthropic.test/v1/messages",
|
||||
"provider_api_format": "claude:messages",
|
||||
"client_api_format": "claude:messages",
|
||||
"provider_request_headers": {
|
||||
"content-type": "application/json",
|
||||
(EXECUTION_RESPONSE_BODY_MODE_HEADER): ExecutionResponseBodyMode::PreserveBytes.as_str()
|
||||
},
|
||||
"provider_request_body": {
|
||||
"model": "claude-sonnet-4",
|
||||
"messages": []
|
||||
},
|
||||
"provider_request_body_base64": "eyAibW9kZWwiOiAiY2xhdWRlLXNvbm5ldC00IiwgIm1lc3NhZ2VzIjogW10gfQ==",
|
||||
"content_type": "application/json",
|
||||
"upstream_is_stream": upstream_is_stream
|
||||
}))
|
||||
.expect("decision should deserialize")
|
||||
}
|
||||
|
||||
fn request_parts() -> http::request::Parts {
|
||||
http::Request::builder()
|
||||
.uri("http://localhost/v1/messages")
|
||||
.body(())
|
||||
.expect("request should build")
|
||||
.into_parts()
|
||||
.0
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn standard_sync_plan_prefers_exact_request_body_bytes() {
|
||||
let built = build_standard_sync_plan_from_decision(
|
||||
&request_parts(),
|
||||
&json!({}),
|
||||
decision_with_raw_body(false),
|
||||
)
|
||||
.expect("plan should build")
|
||||
.expect("plan should exist");
|
||||
|
||||
assert!(built.plan.body.json_body.is_none());
|
||||
assert_eq!(
|
||||
built.plan.body.body_bytes_b64.as_deref(),
|
||||
Some("eyAibW9kZWwiOiAiY2xhdWRlLXNvbm5ldC00IiwgIm1lc3NhZ2VzIjogW10gfQ==")
|
||||
);
|
||||
assert_eq!(
|
||||
built
|
||||
.plan
|
||||
.headers
|
||||
.get(EXECUTION_RESPONSE_BODY_MODE_HEADER)
|
||||
.map(String::as_str),
|
||||
Some(ExecutionResponseBodyMode::PreserveBytes.as_str())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn standard_stream_plan_prefers_exact_request_body_bytes() {
|
||||
let built = build_standard_stream_plan_from_decision(
|
||||
&request_parts(),
|
||||
&json!({}),
|
||||
decision_with_raw_body(true),
|
||||
false,
|
||||
)
|
||||
.expect("plan should build")
|
||||
.expect("plan should exist");
|
||||
|
||||
assert!(built.plan.body.json_body.is_none());
|
||||
assert!(built.plan.body.body_bytes_b64.is_some());
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user