refactor: 优化调度候选排序与用量写入链路并改进 Fernet 缓存与前端批量列表

This commit is contained in:
fawney19
2026-04-21 16:19:07 +08:00
parent c5c56ff92f
commit 25a2b417be
77 changed files with 4801 additions and 2881 deletions

View File

@@ -11,15 +11,15 @@ enum ProviderPrivateStreamNormalizeMode {
KiroToClaudeCli(KiroToClaudeCliStreamState),
}
pub(crate) struct ProviderPrivateStreamNormalizer {
report_context: Value,
pub(crate) struct ProviderPrivateStreamNormalizer<'a> {
report_context: &'a Value,
buffered: Vec<u8>,
mode: ProviderPrivateStreamNormalizeMode,
}
pub(crate) fn maybe_build_provider_private_stream_normalizer(
report_context: Option<&Value>,
) -> Option<ProviderPrivateStreamNormalizer> {
pub(crate) fn maybe_build_provider_private_stream_normalizer<'a>(
report_context: Option<&'a Value>,
) -> Option<ProviderPrivateStreamNormalizer<'a>> {
let report_context = report_context?;
if !report_context
.get("has_envelope")
@@ -51,17 +51,17 @@ pub(crate) fn maybe_build_provider_private_stream_normalizer(
return None;
};
Some(ProviderPrivateStreamNormalizer {
report_context: report_context.clone(),
report_context,
buffered: Vec::new(),
mode,
})
}
impl ProviderPrivateStreamNormalizer {
impl ProviderPrivateStreamNormalizer<'_> {
pub(crate) fn push_chunk(&mut self, chunk: &[u8]) -> Result<Vec<u8>, GatewayError> {
match &mut self.mode {
ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => {
state.push_chunk(&self.report_context, chunk)
state.push_chunk(self.report_context, chunk)
}
ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => {
self.buffered.extend_from_slice(chunk);
@@ -69,7 +69,7 @@ impl ProviderPrivateStreamNormalizer {
while let Some(line_end) = self.buffered.iter().position(|byte| *byte == b'\n') {
let line = self.buffered.drain(..=line_end).collect::<Vec<_>>();
output.extend(
transform_provider_private_stream_line(&self.report_context, line)
transform_provider_private_stream_line(self.report_context, line)
.map_err(|err| GatewayError::Internal(err.to_string()))?,
);
}
@@ -81,14 +81,14 @@ impl ProviderPrivateStreamNormalizer {
pub(crate) fn finish(&mut self) -> Result<Vec<u8>, GatewayError> {
match &mut self.mode {
ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => {
state.finish(&self.report_context)
state.finish(self.report_context)
}
ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => {
if self.buffered.is_empty() {
return Ok(Vec::new());
}
let line = std::mem::take(&mut self.buffered);
transform_provider_private_stream_line(&self.report_context, line)
transform_provider_private_stream_line(self.report_context, line)
.map_err(|err| GatewayError::Internal(err.to_string()))
}
}

View File

@@ -60,7 +60,7 @@ fn normalize_provider_private_stream_bytes(
report_context: &Value,
body: &[u8],
) -> Result<Option<Vec<u8>>, GatewayError> {
let Some(mut normalizer): Option<ProviderPrivateStreamNormalizer> =
let Some(mut normalizer): Option<ProviderPrivateStreamNormalizer<'_>> =
maybe_build_provider_private_stream_normalizer(Some(report_context))
else {
return Ok(Some(body.to_vec()));

View File

@@ -24,21 +24,40 @@ pub(crate) struct LocalCoreSyncFinalizeOutcome {
pub(crate) background_report: Option<GatewaySyncReportRequest>,
}
fn build_local_success_response(
trace_id: &str,
decision: &GatewayControlDecision,
status_code: u16,
body_bytes: Vec<u8>,
headers: BTreeMap<String, String>,
) -> Result<Response<Body>, GatewayError> {
build_client_response_from_parts(
status_code,
&headers,
Body::from(body_bytes),
trace_id,
Some(decision),
)
}
pub(crate) fn build_local_success_outcome(
trace_id: &str,
decision: &GatewayControlDecision,
payload: &GatewaySyncReportRequest,
body_json: Value,
) -> Result<LocalCoreSyncFinalizeOutcome, GatewayError> {
let headers = payload.headers.clone();
let report_headers = payload.headers.clone();
let (body_bytes, response_headers) =
prepare_local_success_response_parts_impl(&payload.headers, &body_json)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let background_report =
build_local_success_background_report_impl(payload, body_json.clone(), headers.clone());
build_local_success_background_report_impl(payload, body_json, report_headers);
build_local_success_outcome_with_report(
trace_id,
decision,
payload.status_code,
body_json,
headers,
body_bytes,
response_headers,
background_report,
)
}
@@ -47,19 +66,12 @@ pub(crate) fn build_local_success_outcome_with_report(
trace_id: &str,
decision: &GatewayControlDecision,
status_code: u16,
body_json: Value,
body_bytes: Vec<u8>,
headers: BTreeMap<String, String>,
background_report: Option<GatewaySyncReportRequest>,
) -> Result<LocalCoreSyncFinalizeOutcome, GatewayError> {
let (body_bytes, headers) = prepare_local_success_response_parts_impl(&headers, &body_json)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let response = build_client_response_from_parts(
status_code,
&headers,
Body::from(body_bytes),
trace_id,
Some(decision),
)?;
let response =
build_local_success_response(trace_id, decision, status_code, body_bytes, headers)?;
Ok(LocalCoreSyncFinalizeOutcome {
response,
background_report,
@@ -73,9 +85,12 @@ pub(crate) fn build_local_success_outcome_with_conversion_report(
client_body_json: Value,
provider_body_json: Value,
) -> Result<LocalCoreSyncFinalizeOutcome, GatewayError> {
let (body_bytes, response_headers) =
prepare_local_success_response_parts_impl(&payload.headers, &client_body_json)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let report_payload = build_local_success_conversion_background_report_impl(
payload,
client_body_json.clone(),
client_body_json,
provider_body_json,
);
@@ -83,8 +98,8 @@ pub(crate) fn build_local_success_outcome_with_conversion_report(
trace_id,
decision,
payload.status_code,
client_body_json,
payload.headers.clone(),
body_bytes,
response_headers,
report_payload,
)
}

View File

@@ -34,6 +34,6 @@ pub(crate) fn maybe_compile_sync_finalize_response(
pub(crate) fn maybe_build_stream_response_rewriter(
report_context: Option<&Value>,
) -> Option<LocalStreamRewriter> {
) -> Option<LocalStreamRewriter<'_>> {
stream::maybe_build_local_stream_rewriter(report_context)
}

View File

@@ -12,15 +12,15 @@ enum RewriteMode {
KiroToClaudeCli(KiroToClaudeCliStreamState),
}
pub(crate) struct LocalStreamRewriter {
report_context: Value,
pub(crate) struct LocalStreamRewriter<'a> {
report_context: &'a Value,
buffered: Vec<u8>,
mode: RewriteMode,
}
pub(crate) fn maybe_build_local_stream_rewriter(
report_context: Option<&Value>,
) -> Option<LocalStreamRewriter> {
pub(crate) fn maybe_build_local_stream_rewriter<'a>(
report_context: Option<&'a Value>,
) -> Option<LocalStreamRewriter<'a>> {
let report_context = report_context?;
let mode = match resolve_finalize_stream_rewrite_mode(report_context)? {
FinalizeStreamRewriteMode::EnvelopeUnwrap => RewriteMode::EnvelopeUnwrap,
@@ -33,16 +33,16 @@ pub(crate) fn maybe_build_local_stream_rewriter(
};
Some(LocalStreamRewriter {
report_context: report_context.clone(),
report_context,
buffered: Vec::new(),
mode,
})
}
impl LocalStreamRewriter {
impl LocalStreamRewriter<'_> {
pub(crate) fn push_chunk(&mut self, chunk: &[u8]) -> Result<Vec<u8>, GatewayError> {
if let RewriteMode::KiroToClaudeCli(state) = &mut self.mode {
return state.push_chunk(&self.report_context, chunk);
return state.push_chunk(self.report_context, chunk);
}
self.buffered.extend_from_slice(chunk);
let mut output = Vec::new();
@@ -55,11 +55,11 @@ impl LocalStreamRewriter {
pub(crate) fn finish(&mut self) -> Result<Vec<u8>, GatewayError> {
if let RewriteMode::KiroToClaudeCli(state) = &mut self.mode {
return state.finish(&self.report_context);
return state.finish(self.report_context);
}
if self.buffered.is_empty() {
match &mut self.mode {
RewriteMode::Standard(state) => return state.finish(&self.report_context),
RewriteMode::Standard(state) => return state.finish(self.report_context),
RewriteMode::KiroToClaudeCli(_) => {}
RewriteMode::EnvelopeUnwrap => {}
}
@@ -69,7 +69,7 @@ impl LocalStreamRewriter {
let mut output = self.transform_line(line)?;
match &mut self.mode {
RewriteMode::Standard(state) => {
output.extend(state.finish(&self.report_context)?);
output.extend(state.finish(self.report_context)?);
}
RewriteMode::KiroToClaudeCli(_) => {}
RewriteMode::EnvelopeUnwrap => {}
@@ -79,9 +79,9 @@ impl LocalStreamRewriter {
fn transform_line(&mut self, line: Vec<u8>) -> Result<Vec<u8>, GatewayError> {
match &mut self.mode {
RewriteMode::EnvelopeUnwrap => transform_envelope_line(&self.report_context, line)
RewriteMode::EnvelopeUnwrap => transform_envelope_line(self.report_context, line)
.map_err(|err| GatewayError::Internal(err.to_string())),
RewriteMode::Standard(state) => state.transform_line(&self.report_context, line),
RewriteMode::Standard(state) => state.transform_line(self.report_context, line),
RewriteMode::KiroToClaudeCli(_) => Ok(Vec::new()),
}
}

View File

@@ -8,7 +8,7 @@ pub(crate) mod transport;
use axum::body::Body;
use axum::http::{Response, Uri};
use serde_json::Value;
use serde_json::{json, Value};
use crate::{usage::GatewaySyncReportRequest, AppState, GatewayError};
@@ -62,8 +62,18 @@ pub(crate) fn collect_control_headers(
crate::headers::collect_control_headers(headers)
}
pub(crate) fn build_report_context_original_request_echo(body_json: &Value) -> Option<Value> {
(!body_json.is_null()).then(|| body_json.clone())
pub(crate) fn build_report_context_original_request_echo(
body_json: Option<&Value>,
body_bytes_b64: Option<&str>,
) -> Option<Value> {
if let Some(body_bytes_b64) = body_bytes_b64
.map(str::trim)
.filter(|value| !value.is_empty())
{
return Some(json!({ "body_bytes_b64": body_bytes_b64 }));
}
body_json.filter(|body| !body.is_null()).cloned()
}
pub(crate) fn is_json_request(headers: &http::HeaderMap) -> bool {
@@ -138,12 +148,23 @@ mod tests {
"body_bytes_b64": "aGVsbG8=",
});
let echo =
build_report_context_original_request_echo(&body).expect("echo should be produced");
let echo = build_report_context_original_request_echo(Some(&body), None)
.expect("echo should be produced");
assert_eq!(echo, body);
}
#[test]
fn build_report_context_original_request_echo_prefers_binary_body_bytes() {
let echo = build_report_context_original_request_echo(
Some(&json!({"ignored": true})),
Some("aGVsbG8="),
)
.expect("echo should be produced");
assert_eq!(echo, json!({"body_bytes_b64": "aGVsbG8="}));
}
#[test]
fn extract_gemini_model_from_path_trims_method_suffix() {
let model =

View File

@@ -1,3 +1,5 @@
use std::collections::BTreeMap;
use tracing::warn;
use crate::ai_pipeline::{
@@ -8,8 +10,8 @@ use crate::scheduler::affinity::SCHEDULER_AFFINITY_TTL;
use crate::scheduler::config::{read_scheduler_ordering_config, SchedulerOrderingConfig};
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
use aether_scheduler_core::{
build_scheduler_affinity_cache_key_for_api_key_id, compare_candidates_by_priority_mode,
requested_capability_priority_for_candidate, SchedulerAffinityTarget, SchedulerPriorityMode,
build_scheduler_affinity_cache_key_for_api_key_id, requested_capability_priority_for_candidate,
SchedulerAffinityTarget, SchedulerPriorityMode,
};
use super::candidate_eligibility::{
@@ -18,6 +20,8 @@ use super::candidate_eligibility::{
const PLANNER_SCHEDULER_AFFINITY_MAX_ENTRIES: usize = 10_000;
type CandidateTransportIdentity<'a> = (&'a str, &'a str, &'a str);
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
enum TunnelOwnerAffinityBucket {
LocalTunnel = 0,
@@ -31,20 +35,36 @@ struct CandidateExecutionOrdering {
keep_priority_on_conversion: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct PlannerCandidateRankingState {
capability_priority: (u32, u32),
tunnel_bucket: TunnelOwnerAffinityBucket,
demote_cross_format: bool,
format_preference: (u8, u8),
original_index: usize,
}
pub(crate) async fn prefer_local_tunnel_owner_candidates(
state: PlannerAppState<'_>,
candidates: Vec<SchedulerMinimalCandidateSelectionCandidate>,
) -> Vec<SchedulerMinimalCandidateSelectionCandidate> {
let mut ranked = Vec::with_capacity(candidates.len());
for (original_index, candidate) in candidates.into_iter().enumerate() {
let bucket = resolve_candidate_tunnel_owner_affinity(state, &candidate).await;
ranked.push((bucket, original_index, candidate));
let mut candidates = candidates;
let mut rankings = Vec::with_capacity(candidates.len());
let mut tunnel_affinity_cache = BTreeMap::new();
for (original_index, candidate) in candidates.iter().enumerate() {
let bucket = resolve_cached_candidate_tunnel_owner_affinity(
state,
&mut tunnel_affinity_cache,
candidate,
)
.await;
rankings.push((bucket, original_index));
}
ranked.sort_by(|left, right| left.0.cmp(&right.0).then(left.1.cmp(&right.1)));
ranked
.into_iter()
.map(|(_, _, candidate)| candidate)
.collect()
let mut order = (0..candidates.len()).collect::<Vec<_>>();
order.sort_by(|left, right| rankings[*left].cmp(&rankings[*right]));
drop(tunnel_affinity_cache);
apply_order(&mut candidates, order);
candidates
}
#[cfg(test)]
@@ -56,11 +76,18 @@ async fn rank_local_execution_candidates(
) -> Vec<SchedulerMinimalCandidateSelectionCandidate> {
let normalized_client_api_format = client_api_format.trim().to_ascii_lowercase();
let ordering_config = read_scheduler_ordering_config_or_default(state).await;
let mut ranked = Vec::with_capacity(candidates.len());
let mut candidates = candidates;
let mut rankings = Vec::with_capacity(candidates.len());
let mut ordering_cache = BTreeMap::new();
for (original_index, candidate) in candidates.into_iter().enumerate() {
let ordering =
resolve_candidate_execution_ordering(state, &candidate, ordering_config).await;
for (original_index, candidate) in candidates.iter().enumerate() {
let ordering = resolve_cached_candidate_execution_ordering(
state,
&mut ordering_cache,
candidate,
ordering_config,
)
.await;
let is_same_format = candidate
.endpoint_api_format
.trim()
@@ -71,112 +98,82 @@ async fn rank_local_execution_candidates(
candidate.endpoint_api_format.as_str(),
);
let capability_priority =
requested_capability_priority_for_candidate(required_capabilities, &candidate);
ranked.push((
capability_priority.0,
capability_priority.1,
ordering.tunnel_bucket,
requested_capability_priority_for_candidate(required_capabilities, candidate);
rankings.push(PlannerCandidateRankingState {
capability_priority,
tunnel_bucket: ordering.tunnel_bucket,
demote_cross_format,
format_preference,
original_index,
candidate,
));
});
}
ranked.sort_by(|left, right| {
left.0
.cmp(&right.0)
.then(left.1.cmp(&right.1))
.then(left.2.cmp(&right.2))
.then(left.3.cmp(&right.3))
.then_with(|| {
compare_candidate_priority_slot(&left.6, &right.6, ordering_config.priority_mode)
})
.then(left.4.cmp(&right.4))
.then_with(|| {
compare_candidates_by_priority_mode(
&left.6,
&right.6,
ordering_config.priority_mode,
None,
)
})
.then(left.5.cmp(&right.5))
let mut order = (0..candidates.len()).collect::<Vec<_>>();
order.sort_by(|left, right| {
compare_planner_candidate_ranking(
&rankings[*left],
&candidates[*left],
&rankings[*right],
&candidates[*right],
ordering_config.priority_mode,
)
});
ranked
.into_iter()
.map(|(_, _, _, _, _, _, candidate)| candidate)
.collect()
drop(ordering_cache);
apply_order(&mut candidates, order);
candidates
}
pub(crate) async fn rank_eligible_local_execution_candidates(
state: PlannerAppState<'_>,
candidates: Vec<EligibleLocalExecutionCandidate>,
client_api_format: &str,
normalized_client_api_format: &str,
required_capabilities: Option<&serde_json::Value>,
) -> Vec<EligibleLocalExecutionCandidate> {
let normalized_client_api_format = client_api_format.trim().to_ascii_lowercase();
let ordering_config = read_scheduler_ordering_config_or_default(state).await;
let mut ranked = Vec::with_capacity(candidates.len());
let mut candidates = candidates;
let mut rankings = Vec::with_capacity(candidates.len());
let mut ordering_cache = BTreeMap::new();
for (original_index, eligible) in candidates.into_iter().enumerate() {
let ordering = resolve_candidate_execution_ordering_from_transport(
for (original_index, eligible) in candidates.iter().enumerate() {
let ordering = resolve_cached_eligible_candidate_execution_ordering(
state,
&eligible.transport,
&mut ordering_cache,
eligible,
ordering_config,
)
.await;
let is_same_format = eligible
.provider_api_format
.eq_ignore_ascii_case(normalized_client_api_format.as_str());
.eq_ignore_ascii_case(normalized_client_api_format);
let demote_cross_format = !is_same_format && !ordering.keep_priority_on_conversion;
let format_preference = candidate_api_format_preference(
normalized_client_api_format.as_str(),
normalized_client_api_format,
eligible.provider_api_format.as_str(),
);
let capability_priority =
requested_capability_priority_for_candidate(required_capabilities, &eligible.candidate);
ranked.push((
capability_priority.0,
capability_priority.1,
ordering.tunnel_bucket,
rankings.push(PlannerCandidateRankingState {
capability_priority,
tunnel_bucket: ordering.tunnel_bucket,
demote_cross_format,
format_preference,
original_index,
eligible,
));
});
}
ranked.sort_by(|left, right| {
left.0
.cmp(&right.0)
.then(left.1.cmp(&right.1))
.then(left.2.cmp(&right.2))
.then(left.3.cmp(&right.3))
.then_with(|| {
compare_candidate_priority_slot(
&left.6.candidate,
&right.6.candidate,
ordering_config.priority_mode,
)
})
.then(left.4.cmp(&right.4))
.then_with(|| {
compare_candidates_by_priority_mode(
&left.6.candidate,
&right.6.candidate,
ordering_config.priority_mode,
None,
)
})
.then(left.5.cmp(&right.5))
let mut order = (0..candidates.len()).collect::<Vec<_>>();
order.sort_by(|left, right| {
compare_planner_candidate_ranking(
&rankings[*left],
&candidates[*left].candidate,
&rankings[*right],
&candidates[*right].candidate,
ordering_config.priority_mode,
)
});
ranked
.into_iter()
.map(|(_, _, _, _, _, _, eligible)| eligible)
.collect()
drop(ordering_cache);
apply_order(&mut candidates, order);
candidates
}
pub(crate) fn remember_scheduler_affinity_for_candidate(
@@ -223,6 +220,21 @@ async fn resolve_candidate_tunnel_owner_affinity(
resolve_tunnel_owner_affinity_from_transport(state, &transport).await
}
async fn resolve_cached_candidate_tunnel_owner_affinity<'a>(
state: PlannerAppState<'_>,
cache: &mut BTreeMap<CandidateTransportIdentity<'a>, TunnelOwnerAffinityBucket>,
candidate: &'a SchedulerMinimalCandidateSelectionCandidate,
) -> TunnelOwnerAffinityBucket {
let identity = candidate_transport_identity(candidate);
if let Some(bucket) = cache.get(&identity).copied() {
return bucket;
}
let bucket = resolve_candidate_tunnel_owner_affinity(state, candidate).await;
cache.insert(identity, bucket);
bucket
}
async fn resolve_candidate_execution_ordering(
state: PlannerAppState<'_>,
candidate: &SchedulerMinimalCandidateSelectionCandidate,
@@ -238,6 +250,43 @@ async fn resolve_candidate_execution_ordering(
resolve_candidate_execution_ordering_from_transport(state, &transport, ordering_config).await
}
async fn resolve_cached_candidate_execution_ordering<'a>(
state: PlannerAppState<'_>,
cache: &mut BTreeMap<CandidateTransportIdentity<'a>, CandidateExecutionOrdering>,
candidate: &'a SchedulerMinimalCandidateSelectionCandidate,
ordering_config: SchedulerOrderingConfig,
) -> CandidateExecutionOrdering {
let identity = candidate_transport_identity(candidate);
if let Some(ordering) = cache.get(&identity).copied() {
return ordering;
}
let ordering = resolve_candidate_execution_ordering(state, candidate, ordering_config).await;
cache.insert(identity, ordering);
ordering
}
async fn resolve_cached_eligible_candidate_execution_ordering<'a>(
state: PlannerAppState<'_>,
cache: &mut BTreeMap<CandidateTransportIdentity<'a>, CandidateExecutionOrdering>,
eligible: &'a EligibleLocalExecutionCandidate,
ordering_config: SchedulerOrderingConfig,
) -> CandidateExecutionOrdering {
let identity = candidate_transport_identity(&eligible.candidate);
if let Some(ordering) = cache.get(&identity).copied() {
return ordering;
}
let ordering = resolve_candidate_execution_ordering_from_transport(
state,
&eligible.transport,
ordering_config,
)
.await;
cache.insert(identity, ordering);
ordering
}
async fn resolve_candidate_execution_ordering_from_transport(
state: PlannerAppState<'_>,
transport: &GatewayProviderTransportSnapshot,
@@ -301,6 +350,16 @@ async fn resolve_tunnel_owner_affinity_from_transport(
}
}
fn candidate_transport_identity(
candidate: &SchedulerMinimalCandidateSelectionCandidate,
) -> CandidateTransportIdentity<'_> {
(
candidate.provider_id.as_str(),
candidate.endpoint_id.as_str(),
candidate.key_id.as_str(),
)
}
fn candidate_api_format_preference(client_api_format: &str, provider_api_format: &str) -> (u8, u8) {
request_candidate_api_format_preference(client_api_format, provider_api_format)
.unwrap_or((u8::MAX, u8::MAX))
@@ -325,6 +384,67 @@ fn compare_candidate_priority_slot(
}
}
fn compare_candidate_identity(
left: &SchedulerMinimalCandidateSelectionCandidate,
right: &SchedulerMinimalCandidateSelectionCandidate,
) -> std::cmp::Ordering {
left.provider_id
.cmp(&right.provider_id)
.then(left.endpoint_id.cmp(&right.endpoint_id))
.then(left.key_id.cmp(&right.key_id))
.then(
left.selected_provider_model_name
.cmp(&right.selected_provider_model_name),
)
}
fn compare_planner_candidate_ranking(
left_state: &PlannerCandidateRankingState,
left_candidate: &SchedulerMinimalCandidateSelectionCandidate,
right_state: &PlannerCandidateRankingState,
right_candidate: &SchedulerMinimalCandidateSelectionCandidate,
priority_mode: SchedulerPriorityMode,
) -> std::cmp::Ordering {
left_state
.capability_priority
.cmp(&right_state.capability_priority)
.then(left_state.tunnel_bucket.cmp(&right_state.tunnel_bucket))
.then(
left_state
.demote_cross_format
.cmp(&right_state.demote_cross_format),
)
.then_with(|| {
compare_candidate_priority_slot(left_candidate, right_candidate, priority_mode)
})
.then(
left_state
.format_preference
.cmp(&right_state.format_preference),
)
.then_with(|| compare_candidate_identity(left_candidate, right_candidate))
.then(left_state.original_index.cmp(&right_state.original_index))
}
fn apply_order<T>(items: &mut [T], sorted_old_indices: Vec<usize>) {
if items.len() < 2 {
return;
}
let mut target_positions = vec![0usize; sorted_old_indices.len()];
for (new_position, old_position) in sorted_old_indices.into_iter().enumerate() {
target_positions[old_position] = new_position;
}
for index in 0..items.len() {
while target_positions[index] != index {
let target = target_positions[index];
items.swap(index, target);
target_positions.swap(index, target);
}
}
}
async fn read_scheduler_ordering_config_or_default(
state: PlannerAppState<'_>,
) -> SchedulerOrderingConfig {

View File

@@ -1,3 +1,5 @@
use std::sync::Arc;
use tracing::warn;
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
@@ -11,7 +13,7 @@ use super::pool_scheduler::apply_local_execution_pool_scheduler;
#[derive(Debug, Clone, PartialEq)]
pub(crate) struct EligibleLocalExecutionCandidate {
pub(crate) candidate: SchedulerMinimalCandidateSelectionCandidate,
pub(crate) transport: GatewayProviderTransportSnapshot,
pub(crate) transport: Arc<GatewayProviderTransportSnapshot>,
pub(crate) provider_api_format: String,
pub(crate) orchestration: LocalExecutionCandidateMetadata,
}
@@ -20,10 +22,16 @@ pub(crate) struct EligibleLocalExecutionCandidate {
pub(crate) struct SkippedLocalExecutionCandidate {
pub(crate) candidate: SchedulerMinimalCandidateSelectionCandidate,
pub(crate) skip_reason: &'static str,
pub(crate) transport: Option<GatewayProviderTransportSnapshot>,
pub(crate) transport: Option<Arc<GatewayProviderTransportSnapshot>>,
pub(crate) extra_data: Option<serde_json::Value>,
}
impl SkippedLocalExecutionCandidate {
pub(crate) fn transport_ref(&self) -> Option<&GatewayProviderTransportSnapshot> {
self.transport.as_deref()
}
}
pub(crate) async fn filter_and_rank_local_execution_candidates(
state: PlannerAppState<'_>,
candidates: Vec<SchedulerMinimalCandidateSelectionCandidate>,
@@ -35,17 +43,18 @@ pub(crate) async fn filter_and_rank_local_execution_candidates(
Vec<EligibleLocalExecutionCandidate>,
Vec<SkippedLocalExecutionCandidate>,
) {
let requested_model = requested_model.trim();
filter_and_rank_local_execution_candidates_with_gate(
state,
candidates,
client_api_format,
required_capabilities,
sticky_session_token,
|candidate, transport| {
|candidate, transport, normalized_client_api_format| {
current_local_execution_candidate_skip_reason_with_transport(
candidate,
transport,
client_api_format,
normalized_client_api_format,
requested_model,
)
},
@@ -64,13 +73,14 @@ pub(crate) async fn filter_and_rank_local_execution_candidates_without_transport
Vec<EligibleLocalExecutionCandidate>,
Vec<SkippedLocalExecutionCandidate>,
) {
let requested_model = requested_model.map(str::trim);
filter_and_rank_local_execution_candidates_with_gate(
state,
candidates,
client_api_format,
required_capabilities,
sticky_session_token,
|candidate, transport| {
|candidate, transport, _normalized_client_api_format| {
current_local_execution_candidate_common_skip_reason_with_transport(
candidate,
transport,
@@ -96,10 +106,12 @@ where
F: Fn(
&SchedulerMinimalCandidateSelectionCandidate,
&GatewayProviderTransportSnapshot,
&str,
) -> Option<&'static str>,
{
let normalized_client_api_format = client_api_format.trim().to_ascii_lowercase();
let mut selectable = Vec::with_capacity(candidates.len());
let mut skipped = Vec::new();
let mut skipped = Vec::with_capacity(candidates.len());
for candidate in candidates {
let Some(transport) = read_candidate_transport_snapshot(state, &candidate).await else {
@@ -111,11 +123,18 @@ where
});
continue;
};
if candidate_is_ineligible_due_to_disabled_format_conversion(&transport, client_api_format)
{
let transport = Arc::new(transport);
if candidate_is_ineligible_due_to_disabled_format_conversion(
transport.as_ref(),
normalized_client_api_format.as_str(),
) {
continue;
}
match runtime_skip_reason(&candidate, &transport) {
match runtime_skip_reason(
&candidate,
transport.as_ref(),
normalized_client_api_format.as_str(),
) {
Some(skip_reason) => skipped.push(SkippedLocalExecutionCandidate {
candidate,
skip_reason,
@@ -134,7 +153,7 @@ where
let ranked = rank_eligible_local_execution_candidates(
state,
selectable,
client_api_format,
normalized_client_api_format.as_str(),
required_capabilities,
)
.await;
@@ -146,28 +165,27 @@ where
}
pub(crate) fn extract_pool_sticky_session_token(body_json: &serde_json::Value) -> Option<String> {
fn non_empty_string(value: Option<&serde_json::Value>) -> Option<String> {
fn non_empty_str(value: Option<&serde_json::Value>) -> Option<&str> {
value
.and_then(serde_json::Value::as_str)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
}
let object = body_json.as_object()?;
non_empty_string(object.get("prompt_cache_key"))
.or_else(|| non_empty_string(object.get("conversation_id")))
.or_else(|| non_empty_string(object.get("conversationId")))
.or_else(|| non_empty_string(object.get("session_id")))
.or_else(|| non_empty_string(object.get("sessionId")))
non_empty_str(object.get("prompt_cache_key"))
.or_else(|| non_empty_str(object.get("conversation_id")))
.or_else(|| non_empty_str(object.get("conversationId")))
.or_else(|| non_empty_str(object.get("session_id")))
.or_else(|| non_empty_str(object.get("sessionId")))
.or_else(|| {
object
.get("metadata")
.and_then(serde_json::Value::as_object)
.and_then(|metadata| {
non_empty_string(metadata.get("session_id"))
.or_else(|| non_empty_string(metadata.get("conversation_id")))
non_empty_str(metadata.get("session_id"))
.or_else(|| non_empty_str(metadata.get("conversation_id")))
})
})
.or_else(|| {
@@ -175,10 +193,11 @@ pub(crate) fn extract_pool_sticky_session_token(body_json: &serde_json::Value) -
.get("conversationState")
.and_then(serde_json::Value::as_object)
.and_then(|state| {
non_empty_string(state.get("conversationId"))
.or_else(|| non_empty_string(state.get("sessionId")))
non_empty_str(state.get("conversationId"))
.or_else(|| non_empty_str(state.get("sessionId")))
})
})
.map(ToOwned::to_owned)
}
fn current_local_execution_candidate_common_skip_reason_with_transport(
@@ -198,13 +217,16 @@ fn current_local_execution_candidate_common_skip_reason_with_transport(
return Some("key_inactive");
}
let candidate_api_format = candidate.endpoint_api_format.trim().to_ascii_lowercase();
let endpoint_api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
if endpoint_api_format != candidate_api_format {
let endpoint_api_format = transport.endpoint.api_format.trim();
if !candidate
.endpoint_api_format
.trim()
.eq_ignore_ascii_case(endpoint_api_format)
{
return Some("endpoint_api_format_changed");
}
if !transport_key_supports_api_format(transport, endpoint_api_format.as_str()) {
if !transport_key_supports_api_format(transport, endpoint_api_format) {
return Some("key_api_format_disabled");
}
if !transport_key_allows_candidate_model(transport, requested_model, candidate) {
@@ -216,34 +238,33 @@ fn current_local_execution_candidate_common_skip_reason_with_transport(
fn candidate_is_ineligible_due_to_disabled_format_conversion(
transport: &GatewayProviderTransportSnapshot,
client_api_format: &str,
normalized_client_api_format: &str,
) -> bool {
let endpoint_api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
let client_api_format = client_api_format.trim().to_ascii_lowercase();
if client_api_format == endpoint_api_format {
let endpoint_api_format = transport.endpoint.api_format.trim();
if endpoint_api_format.eq_ignore_ascii_case(normalized_client_api_format) {
return false;
}
crate::ai_pipeline::conversion::request_conversion_kind(
client_api_format.as_str(),
endpoint_api_format.as_str(),
normalized_client_api_format,
endpoint_api_format,
)
.is_some()
&& crate::ai_pipeline::conversion::request_conversion_requires_enable_flag(
client_api_format.as_str(),
endpoint_api_format.as_str(),
normalized_client_api_format,
endpoint_api_format,
)
&& !crate::ai_pipeline::conversion::request_conversion_enabled_for_transport(
transport,
client_api_format.as_str(),
endpoint_api_format.as_str(),
normalized_client_api_format,
endpoint_api_format,
)
}
fn current_local_execution_candidate_skip_reason_with_transport(
candidate: &SchedulerMinimalCandidateSelectionCandidate,
transport: &GatewayProviderTransportSnapshot,
client_api_format: &str,
normalized_client_api_format: &str,
requested_model: &str,
) -> Option<&'static str> {
if let Some(skip_reason) = current_local_execution_candidate_common_skip_reason_with_transport(
@@ -254,16 +275,15 @@ fn current_local_execution_candidate_skip_reason_with_transport(
return Some(skip_reason);
}
let endpoint_api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
let client_api_format = client_api_format.trim().to_ascii_lowercase();
if client_api_format == endpoint_api_format {
let endpoint_api_format = transport.endpoint.api_format.trim();
if endpoint_api_format.eq_ignore_ascii_case(normalized_client_api_format) {
return None;
}
if !crate::ai_pipeline::conversion::request_pair_allowed_for_transport(
transport,
client_api_format.as_str(),
endpoint_api_format.as_str(),
normalized_client_api_format,
endpoint_api_format,
) {
return Some("transport_unsupported");
}
@@ -292,15 +312,6 @@ fn transport_key_allows_candidate_model(
return true;
};
let allowed_models = allowed_models
.iter()
.map(|value| value.trim())
.filter(|value| !value.is_empty())
.collect::<Vec<_>>();
if allowed_models.is_empty() {
return false;
}
let requested_model = requested_model.trim();
let global_model_name = candidate.global_model_name.trim();
let selected_provider_model_name = candidate.selected_provider_model_name.trim();
@@ -310,12 +321,20 @@ fn transport_key_allows_candidate_model(
.map(str::trim)
.filter(|value| !value.is_empty());
allowed_models.iter().any(|allowed_model| {
*allowed_model == requested_model
|| *allowed_model == global_model_name
|| *allowed_model == selected_provider_model_name
|| mapping_matched_model.is_some_and(|value| value == *allowed_model)
})
for allowed_model in allowed_models.iter().map(String::as_str).map(str::trim) {
if allowed_model.is_empty() {
continue;
}
if allowed_model == requested_model
|| allowed_model == global_model_name
|| allowed_model == selected_provider_model_name
|| mapping_matched_model.is_some_and(|value| value == allowed_model)
{
return true;
}
}
false
}
pub(crate) async fn read_candidate_transport_snapshot(

View File

@@ -9,7 +9,7 @@ use crate::ai_pipeline::planner::candidate_eligibility::{
use crate::ai_pipeline::planner::runtime_miss::record_local_runtime_candidate_skip_reason;
use crate::ai_pipeline::{GatewayAuthApiKeySnapshot, PlannerAppState};
use crate::clock::current_unix_ms;
use crate::orchestration::{build_local_attempt_identities, ExecutionAttemptIdentity};
use crate::orchestration::{local_attempt_slot_count, ExecutionAttemptIdentity};
use crate::AppState;
#[derive(Debug, Clone)]
@@ -17,15 +17,13 @@ pub(crate) struct LocalExecutionCandidateAttempt {
pub(crate) eligible: EligibleLocalExecutionCandidate,
pub(crate) candidate_index: u32,
pub(crate) retry_index: u32,
pub(crate) pool_key_index: Option<u32>,
pub(crate) candidate_group_id: Option<String>,
pub(crate) candidate_id: String,
}
impl LocalExecutionCandidateAttempt {
pub(crate) fn attempt_identity(&self) -> ExecutionAttemptIdentity {
ExecutionAttemptIdentity::new(self.candidate_index, self.retry_index)
.with_pool_key_index(self.pool_key_index)
.with_pool_key_index(self.eligible.orchestration.pool_key_index)
}
}
@@ -84,17 +82,25 @@ where
F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value>,
{
let created_at_unix_ms = current_unix_ms();
let mut materialized = Vec::new();
let total_attempts = candidates
.iter()
.map(|eligible| local_attempt_slot_count(&eligible.transport) as usize)
.sum();
let mut materialized = Vec::with_capacity(total_attempts);
for (candidate_index, eligible) in candidates.into_iter().enumerate() {
let candidate_index = candidate_index as u32;
let attempt_identities =
build_local_attempt_identities(candidate_index, &eligible.transport)
.into_iter()
.map(|identity| identity.with_pool_key_index(eligible.orchestration.pool_key_index))
.collect::<Vec<_>>();
let attempt_slots = local_attempt_slot_count(&eligible.transport);
let pool_key_index = eligible.orchestration.pool_key_index;
let extra_data = build_extra_data(&eligible);
let mut owned_eligible = Some(eligible);
for attempt_identity in attempt_identities {
for retry_index in 0..attempt_slots {
let eligible = owned_eligible
.as_ref()
.expect("eligible candidate should remain available until final retry");
let attempt_identity = ExecutionAttemptIdentity::new(candidate_index, retry_index)
.with_pool_key_index(pool_key_index);
let generated_candidate_id = Uuid::new_v4().to_string();
let candidate_id = state
.persist_available_local_candidate(
@@ -106,18 +112,23 @@ where
attempt_identity.retry_index,
&generated_candidate_id,
required_capabilities,
build_extra_data(&eligible),
extra_data.clone(),
created_at_unix_ms,
error_context,
)
.await;
let eligible = if retry_index + 1 == attempt_slots {
owned_eligible
.take()
.expect("final retry should consume owned eligible candidate")
} else {
eligible.clone()
};
materialized.push(LocalExecutionCandidateAttempt {
eligible: eligible.clone(),
eligible,
candidate_index: attempt_identity.candidate_index,
retry_index: attempt_identity.retry_index,
pool_key_index: attempt_identity.pool_key_index,
candidate_group_id: eligible.orchestration.candidate_group_id.clone(),
candidate_id,
});
}

View File

@@ -54,7 +54,7 @@ pub(crate) fn build_local_execution_candidate_metadata(
) -> Value {
build_local_execution_candidate_metadata_for_candidate(
&parts.eligible.candidate,
Some(&parts.eligible.transport),
Some(parts.eligible.transport.as_ref()),
parts.provider_api_format,
parts.client_api_format,
parts.extra_fields,
@@ -127,7 +127,7 @@ pub(crate) fn build_local_execution_candidate_contract_metadata(
append_execution_contract_fields_to_value(
build_local_execution_candidate_metadata_for_candidate(
&parts.eligible.candidate,
Some(&parts.eligible.transport),
Some(parts.eligible.transport.as_ref()),
parts.provider_api_format,
parts.client_api_format,
parts.extra_fields,

View File

@@ -81,9 +81,9 @@ fn build_sync_plan_payload_from_decision(
parts: &http::request::Parts,
body_json: &serde_json::Value,
plan_kind: &str,
payload: GatewayControlSyncDecisionResponse,
mut payload: GatewayControlSyncDecisionResponse,
) -> Result<Option<GatewayControlPlanResponse>, GatewayError> {
let auth_context = payload.auth_context.clone();
let auth_context = payload.auth_context.take();
let plan_and_report = match plan_kind {
OPENAI_CHAT_SYNC_PLAN_KIND => {
build_openai_chat_sync_plan_from_decision(parts, body_json, payload)?
@@ -121,9 +121,9 @@ fn build_stream_plan_payload_from_decision(
parts: &http::request::Parts,
body_json: &serde_json::Value,
plan_kind: &str,
payload: GatewayControlSyncDecisionResponse,
mut payload: GatewayControlSyncDecisionResponse,
) -> Result<Option<GatewayControlPlanResponse>, GatewayError> {
let auth_context = payload.auth_context.clone();
let auth_context = payload.auth_context.take();
let plan_and_report = match plan_kind {
OPENAI_CHAT_STREAM_PLAN_KIND => {
build_openai_chat_stream_plan_from_decision(parts, body_json, payload)?

View File

@@ -143,36 +143,69 @@ async fn maybe_build_local_video_task_follow_up_sync_decision_payload(
return Ok(None);
};
let auth_pair = extract_auth_header_pair(&follow_up.plan.headers);
let execution_strategy =
if follow_up.plan.provider_api_format == follow_up.plan.client_api_format {
ExecutionStrategy::LocalSameFormat
} else {
ExecutionStrategy::LocalCrossFormat
};
let conversion_mode = if follow_up.plan.provider_api_format == follow_up.plan.client_api_format
{
let aether_video_tasks_core::LocalVideoTaskFollowUpPlan {
plan,
report_kind,
report_context,
} = follow_up;
let aether_contracts::ExecutionPlan {
request_id: _request_id,
candidate_id,
provider_name,
provider_id,
endpoint_id,
key_id,
method,
url,
headers,
content_type,
content_encoding: _content_encoding,
body,
stream: _stream,
client_api_format,
provider_api_format,
model_name,
proxy,
tls_profile,
timeouts,
} = plan;
let auth_pair = extract_auth_header_pair(&headers);
let execution_strategy = if provider_api_format == client_api_format {
ExecutionStrategy::LocalSameFormat
} else {
ExecutionStrategy::LocalCrossFormat
};
let conversion_mode = if provider_api_format == client_api_format {
ConversionMode::None
} else {
ConversionMode::Bidirectional
};
let upstream_base_url = infer_upstream_base_url(&follow_up.plan.url);
let upstream_base_url = infer_upstream_base_url(&url);
let provider_contract = provider_api_format.clone();
let client_contract = client_api_format.clone();
let auth_header = auth_pair.map(|(name, _)| name.to_string());
let auth_value = auth_pair.map(|(_, value)| value.to_string());
let aether_contracts::RequestBody {
json_body,
body_bytes_b64,
body_ref: _body_ref,
} = body;
debug!(
event_name = "local_video_follow_up_sync_decision_payload_built",
log_type = "debug",
trace_id = %trace_id,
request_id = %trace_id,
candidate_id = ?follow_up.plan.candidate_id,
provider_id = %follow_up.plan.provider_id,
endpoint_id = %follow_up.plan.endpoint_id,
key_id = %follow_up.plan.key_id,
candidate_id = ?candidate_id,
provider_id = %provider_id,
endpoint_id = %endpoint_id,
key_id = %key_id,
plan_kind,
downstream_path = %parts.uri.path(),
provider_api_format = %follow_up.plan.provider_api_format,
client_api_format = %follow_up.plan.client_api_format,
provider_api_format = %provider_api_format,
client_api_format = %client_api_format,
upstream_base_url = ?upstream_base_url,
upstream_url = %follow_up.plan.url,
upstream_url = %url,
"gateway built local video follow-up sync decision payload"
);
@@ -182,39 +215,41 @@ async fn maybe_build_local_video_task_follow_up_sync_decision_payload(
execution_strategy: Some(execution_strategy.as_str().to_string()),
conversion_mode: Some(conversion_mode.as_str().to_string()),
request_id: Some(trace_id.to_string()),
candidate_id: follow_up.plan.candidate_id.clone(),
provider_name: follow_up.plan.provider_name.clone(),
provider_id: Some(follow_up.plan.provider_id.clone()),
endpoint_id: Some(follow_up.plan.endpoint_id.clone()),
key_id: Some(follow_up.plan.key_id.clone()),
candidate_id,
provider_name,
provider_id: Some(provider_id),
endpoint_id: Some(endpoint_id),
key_id: Some(key_id),
upstream_base_url,
upstream_url: Some(follow_up.plan.url.clone()),
provider_request_method: Some(follow_up.plan.method.clone()),
auth_header: auth_pair.as_ref().map(|(name, _)| name.clone()),
auth_value: auth_pair.as_ref().map(|(_, value)| value.clone()),
provider_api_format: Some(follow_up.plan.provider_api_format.clone()),
client_api_format: Some(follow_up.plan.client_api_format.clone()),
provider_contract: Some(follow_up.plan.provider_api_format.clone()),
client_contract: Some(follow_up.plan.client_api_format.clone()),
model_name: follow_up.plan.model_name.clone(),
upstream_url: Some(url),
provider_request_method: Some(method),
auth_header,
auth_value,
provider_api_format: Some(provider_api_format),
client_api_format: Some(client_api_format),
provider_contract: Some(provider_contract),
client_contract: Some(client_contract),
model_name,
mapped_model: None,
prompt_cache_key: None,
extra_headers: BTreeMap::new(),
provider_request_headers: follow_up.plan.headers.clone(),
provider_request_body: follow_up.plan.body.json_body.clone(),
provider_request_body_base64: follow_up.plan.body.body_bytes_b64.clone(),
content_type: follow_up.plan.content_type.clone(),
proxy: follow_up.plan.proxy.clone(),
tls_profile: follow_up.plan.tls_profile.clone(),
timeouts: follow_up.plan.timeouts.clone(),
provider_request_headers: headers,
provider_request_body: json_body,
provider_request_body_base64: body_bytes_b64,
content_type,
proxy,
tls_profile,
timeouts,
upstream_is_stream: false,
report_kind: follow_up.report_kind,
report_context: follow_up.report_context,
report_kind,
report_context,
auth_context: Some(build_execution_runtime_auth_context(&auth_context)),
}))
}
fn extract_auth_header_pair(headers: &BTreeMap<String, String>) -> Option<(String, String)> {
fn extract_auth_header_pair<'a>(
headers: &'a BTreeMap<String, String>,
) -> Option<(&'a str, &'a str)> {
[
"authorization",
"x-api-key",
@@ -227,7 +262,7 @@ fn extract_auth_header_pair(headers: &BTreeMap<String, String>) -> Option<(Strin
headers
.iter()
.find(|(header_name, _)| header_name.eq_ignore_ascii_case(name))
.map(|(header_name, value)| (header_name.clone(), value.clone()))
.map(|(header_name, value)| (header_name.as_str(), value.as_str()))
})
}

View File

@@ -1,103 +1,76 @@
use aether_contracts::{ExecutionPlan, RequestBody};
use super::{augment_sync_report_context, LocalStreamPlanAndReport, LocalSyncPlanAndReport};
use super::{
augment_sync_report_context, take_non_empty_string, LocalStreamPlanAndReport,
LocalSyncPlanAndReport,
};
use crate::{GatewayControlSyncDecisionResponse, GatewayError};
pub(crate) fn build_passthrough_sync_plan_from_decision(
parts: &http::request::Parts,
payload: GatewayControlSyncDecisionResponse,
) -> Result<Option<LocalSyncPlanAndReport>, GatewayError> {
let Some(request_id) = payload
.request_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let mut payload = payload;
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
return Ok(None);
};
let Some(provider_id) = payload
.provider_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
return Ok(None);
};
let Some(endpoint_id) = payload
.endpoint_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
return Ok(None);
};
let Some(key_id) = payload
.key_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
return Ok(None);
};
let Some(provider_api_format) = payload
.provider_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
return Ok(None);
};
let Some(client_api_format) = payload
.client_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
return Ok(None);
};
let Some(upstream_url) = payload
.upstream_url
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(upstream_url) = take_non_empty_string(&mut payload.upstream_url) else {
return Ok(None);
};
let (request_body, provider_request_body_for_report) = resolve_passthrough_sync_request_body(
payload.provider_request_body.clone(),
payload.provider_request_body_base64.clone(),
let provider_request_headers = std::mem::take(&mut payload.provider_request_headers);
let ignored_provider_request_body = serde_json::Value::Null;
let report_context = augment_sync_report_context(
payload.report_context.take(),
&provider_request_headers,
&ignored_provider_request_body,
)?;
let request_body = resolve_passthrough_sync_request_body(
payload.provider_request_body.take(),
payload.provider_request_body_base64.take(),
);
let provider_request_method = take_non_empty_string(&mut payload.provider_request_method);
let content_type = payload
.content_type
.take()
.or_else(|| provider_request_headers.get("content-type").cloned());
let plan = ExecutionPlan {
request_id,
candidate_id: payload.candidate_id.clone(),
provider_name: payload.provider_name.clone(),
candidate_id: payload.candidate_id.take(),
provider_name: payload.provider_name.take(),
provider_id,
endpoint_id,
key_id,
method: payload
.provider_request_method
.clone()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| parts.method.to_string()),
method: provider_request_method.unwrap_or_else(|| parts.method.to_string()),
url: upstream_url,
headers: payload.provider_request_headers.clone(),
content_type: payload.content_type.clone().or_else(|| {
payload
.provider_request_headers
.get("content-type")
.cloned()
}),
headers: provider_request_headers,
content_type,
content_encoding: None,
body: request_body,
stream: false,
client_api_format,
provider_api_format,
model_name: payload.model_name.clone(),
proxy: payload.proxy.clone(),
tls_profile: payload.tls_profile.clone(),
timeouts: payload.timeouts.clone(),
model_name: payload.model_name.take(),
proxy: payload.proxy.take(),
tls_profile: payload.tls_profile.take(),
timeouts: payload.timeouts.take(),
};
let report_context = augment_sync_report_context(
payload.report_context,
&plan.headers,
&provider_request_body_for_report,
)?;
Ok(Some(LocalSyncPlanAndReport {
plan,
report_kind: payload.report_kind,
@@ -109,71 +82,44 @@ pub(crate) fn build_passthrough_stream_plan_from_decision(
parts: &http::request::Parts,
payload: GatewayControlSyncDecisionResponse,
) -> Result<Option<LocalStreamPlanAndReport>, GatewayError> {
let Some(request_id) = payload
.request_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let mut payload = payload;
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
return Ok(None);
};
let Some(provider_id) = payload
.provider_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
return Ok(None);
};
let Some(endpoint_id) = payload
.endpoint_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
return Ok(None);
};
let Some(key_id) = payload
.key_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
return Ok(None);
};
let Some(provider_api_format) = payload
.provider_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
return Ok(None);
};
let Some(client_api_format) = payload
.client_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
return Ok(None);
};
let Some(upstream_url) = payload
.upstream_url
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(upstream_url) = take_non_empty_string(&mut payload.upstream_url) else {
return Ok(None);
};
let provider_request_headers = std::mem::take(&mut payload.provider_request_headers);
let content_type = payload
.content_type
.take()
.or_else(|| provider_request_headers.get("content-type").cloned());
let plan = ExecutionPlan {
request_id,
candidate_id: payload.candidate_id.clone(),
provider_name: payload.provider_name.clone(),
candidate_id: payload.candidate_id.take(),
provider_name: payload.provider_name.take(),
provider_id,
endpoint_id,
key_id,
method: parts.method.to_string(),
url: upstream_url,
headers: payload.provider_request_headers.clone(),
content_type: payload.content_type.clone().or_else(|| {
payload
.provider_request_headers
.get("content-type")
.cloned()
}),
headers: provider_request_headers,
content_type,
content_encoding: None,
body: RequestBody {
json_body: None,
@@ -183,10 +129,10 @@ pub(crate) fn build_passthrough_stream_plan_from_decision(
stream: true,
client_api_format,
provider_api_format,
model_name: payload.model_name.clone(),
proxy: payload.proxy.clone(),
tls_profile: payload.tls_profile.clone(),
timeouts: payload.timeouts.clone(),
model_name: payload.model_name.take(),
proxy: payload.proxy.take(),
tls_profile: payload.tls_profile.take(),
timeouts: payload.timeouts.take(),
};
Ok(Some(LocalStreamPlanAndReport {
@@ -199,35 +145,33 @@ pub(crate) fn build_passthrough_stream_plan_from_decision(
fn resolve_passthrough_sync_request_body(
provider_request_body: Option<serde_json::Value>,
provider_request_body_base64: Option<String>,
) -> (RequestBody, serde_json::Value) {
if let Some(body_bytes_b64) = provider_request_body_base64
.as_ref()
.map(|value| value.trim())
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
) -> RequestBody {
if let Some(body_bytes_b64) = provider_request_body_base64.and_then(trim_owned_non_empty_string)
{
return (
RequestBody {
json_body: None,
body_bytes_b64: Some(body_bytes_b64.clone()),
body_ref: None,
},
serde_json::json!({"body_bytes_b64": body_bytes_b64}),
);
return RequestBody {
json_body: None,
body_bytes_b64: Some(body_bytes_b64),
body_ref: None,
};
}
match provider_request_body.unwrap_or(serde_json::Value::Null) {
serde_json::Value::Null => (
RequestBody {
json_body: None,
body_bytes_b64: None,
body_ref: None,
},
serde_json::Value::Null,
),
other => {
let report_body = other.clone();
(RequestBody::from_json(other), report_body)
}
serde_json::Value::Null => RequestBody {
json_body: None,
body_bytes_b64: None,
body_ref: None,
},
other => RequestBody::from_json(other),
}
}
fn trim_owned_non_empty_string(value: String) -> Option<String> {
let trimmed = value.trim();
if trimmed.is_empty() {
return None;
}
if trimmed.len() == value.len() {
return Some(value);
}
Some(trimmed.to_owned())
}

View File

@@ -134,7 +134,7 @@ pub(crate) async fn materialize_local_same_format_provider_candidate_attempts(
skipped_candidate.extra_data = Some(
build_local_execution_candidate_contract_metadata_for_candidate(
&skipped_candidate.candidate,
skipped_candidate.transport.as_ref(),
skipped_candidate.transport_ref(),
provider_api_format.as_str(),
spec_metadata.api_format,
serde_json::Map::new(),

View File

@@ -41,7 +41,6 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
let LocalSameFormatProviderCandidateAttempt {
eligible,
candidate_index,
candidate_group_id,
candidate_id,
..
} = &attempt;
@@ -95,12 +94,13 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
provider_api_format: spec_metadata.api_format,
client_api_format: spec_metadata.api_format,
mapped_model: Some(&resolved.mapped_model),
candidate_group_id: candidate_group_id.as_deref(),
candidate_group_id: eligible.orchestration.candidate_group_id.as_deref(),
upstream_url: Some(&resolved.upstream_url),
provider_request_method: Some(serde_json::Value::Null),
provider_request_headers: Some(&resolved.provider_request_headers),
original_headers: &parts.headers,
original_request_body: body_json,
original_request_body_json: Some(body_json),
original_request_body_base64: None,
has_envelope: resolved.is_kiro || resolved.is_antigravity,
needs_conversion: false,
extra_fields,
@@ -112,6 +112,19 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
),
&resolved.transport,
);
let super::request::LocalSameFormatProviderCandidatePayloadParts {
transport,
is_antigravity: _,
is_kiro: _,
auth_header,
auth_value,
mapped_model,
report_kind,
upstream_is_stream,
upstream_url,
provider_request_headers,
provider_request_body,
} = resolved;
Some(build_local_execution_decision_response(
LocalExecutionDecisionResponseParts {
@@ -121,29 +134,29 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
conversion_mode: ConversionMode::None,
request_id: trace_id.to_string(),
candidate_id: candidate_id.to_string(),
provider_name: resolved.transport.provider.name.clone(),
provider_name: transport.provider.name.clone(),
provider_id: candidate.provider_id.clone(),
endpoint_id: candidate.endpoint_id.clone(),
key_id: candidate.key_id.clone(),
upstream_base_url: resolved.transport.endpoint.base_url.clone(),
upstream_url: resolved.upstream_url.clone(),
upstream_base_url: transport.endpoint.base_url.clone(),
upstream_url,
provider_request_method: None,
auth_header: resolved.auth_header.clone(),
auth_value: resolved.auth_value.clone(),
auth_header,
auth_value,
provider_api_format: spec_metadata.api_format.to_string(),
client_api_format: spec_metadata.api_format.to_string(),
model_name: input.requested_model.clone(),
mapped_model: resolved.mapped_model.clone(),
mapped_model,
prompt_cache_key,
provider_request_headers: resolved.provider_request_headers.clone(),
provider_request_body: Some(resolved.provider_request_body.clone()),
provider_request_headers,
provider_request_body: Some(provider_request_body),
provider_request_body_base64: None,
content_type: Some("application/json".to_string()),
proxy,
tls_profile,
timeouts: resolve_transport_execution_timeouts(&resolved.transport),
upstream_is_stream: resolved.upstream_is_stream,
report_kind: Some(resolved.report_kind.to_string()),
timeouts: resolve_transport_execution_timeouts(&transport),
upstream_is_stream,
report_kind: Some(report_kind.to_string()),
report_context: Some(report_context),
auth_context: input.auth_context.clone(),
},

View File

@@ -1,4 +1,5 @@
use std::collections::BTreeMap;
use std::sync::Arc;
use serde_json::Value;
@@ -72,7 +73,7 @@ pub(crate) fn resolve_same_format_provider_transport_unsupported_reason_for_trac
}
pub(crate) struct LocalSameFormatProviderCandidatePayloadParts {
pub(super) transport: GatewayProviderTransportSnapshot,
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
pub(super) is_antigravity: bool,
pub(super) is_kiro: bool,
pub(super) auth_header: Option<String>,

View File

@@ -1,3 +1,5 @@
use std::sync::Arc;
use crate::ai_pipeline::planner::candidate_eligibility::EligibleLocalExecutionCandidate;
use crate::ai_pipeline::planner::candidate_preparation::{
resolve_candidate_mapped_model, resolve_candidate_oauth_auth, OauthPreparationContext,
@@ -19,7 +21,7 @@ use super::policy::{
};
pub(super) struct PreparedSameFormatProviderCandidate {
pub(super) transport: GatewayProviderTransportSnapshot,
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
pub(super) is_antigravity: bool,
pub(super) is_claude_code: bool,
pub(super) is_vertex: bool,
@@ -44,7 +46,7 @@ pub(super) async fn prepare_local_same_format_provider_candidate(
let spec_metadata = local_same_format_provider_spec_metadata(spec);
let planner_state = PlannerAppState::new(state);
let candidate = &eligible.candidate;
let transport = eligible.transport.clone();
let transport = Arc::clone(&eligible.transport);
let behavior = classify_same_format_provider_request_behavior(&transport, spec_metadata);
if !same_format_provider_transport_supported(

View File

@@ -40,3 +40,7 @@ pub(super) fn augment_sync_report_context(
)
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(super) fn take_non_empty_string(value: &mut Option<String>) -> Option<String> {
value.take().filter(|value| !value.trim().is_empty())
}

View File

@@ -1,5 +1,5 @@
use std::cmp::Ordering;
use std::collections::{BTreeMap, BTreeSet};
use std::collections::{btree_map::Entry, BTreeMap, BTreeSet};
use std::hash::{Hash, Hasher};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
@@ -123,14 +123,11 @@ async fn read_pool_catalog_key_contexts_by_id(
if pool_config_for_candidate(candidate).is_none() {
continue;
}
if provider_type_by_key_id.contains_key(&candidate.candidate.key_id) {
continue;
let key_id = candidate.candidate.key_id.clone();
if let Entry::Vacant(entry) = provider_type_by_key_id.entry(key_id.clone()) {
entry.insert(candidate.transport.provider.provider_type.clone());
key_ids.push(key_id);
}
provider_type_by_key_id.insert(
candidate.candidate.key_id.clone(),
candidate.transport.provider.provider_type.clone(),
);
key_ids.push(candidate.candidate.key_id.clone());
}
if key_ids.is_empty() {
@@ -324,10 +321,15 @@ fn apply_local_execution_pool_scheduler_with_runtime_map(
for candidate in candidates {
let pool_enabled = pool_config_for_candidate(&candidate).is_some();
let group_key = pool_group_key(&candidate, pool_enabled);
if !groups.contains_key(&group_key) {
group_order.push(group_key.clone());
match groups.entry(group_key) {
Entry::Vacant(entry) => {
group_order.push(entry.key().clone());
entry.insert(vec![candidate]);
}
Entry::Occupied(mut entry) => {
entry.get_mut().push(candidate);
}
}
groups.entry(group_key).or_default().push(candidate);
}
let mut reordered = Vec::new();
@@ -1509,7 +1511,7 @@ mod tests {
},
provider_api_format: "openai:chat".to_string(),
orchestration: LocalExecutionCandidateMetadata::default(),
transport: crate::ai_pipeline::GatewayProviderTransportSnapshot {
transport: Arc::new(crate::ai_pipeline::GatewayProviderTransportSnapshot {
provider: GatewayProviderTransportProvider {
id: provider_id.to_string(),
name: provider_id.to_string(),
@@ -1558,7 +1560,7 @@ mod tests {
decrypted_api_key: "secret".to_string(),
decrypted_auth_config: None,
},
},
}),
}
}
}

View File

@@ -24,7 +24,8 @@ pub(crate) struct LocalExecutionReportContextParts<'a> {
pub(crate) provider_request_method: Option<Value>,
pub(crate) provider_request_headers: Option<&'a BTreeMap<String, String>>,
pub(crate) original_headers: &'a http::HeaderMap,
pub(crate) original_request_body: &'a Value,
pub(crate) original_request_body_json: Option<&'a Value>,
pub(crate) original_request_body_base64: Option<&'a str>,
pub(crate) has_envelope: bool,
pub(crate) needs_conversion: bool,
pub(crate) extra_fields: Map<String, Value>,
@@ -110,8 +111,11 @@ pub(crate) fn build_local_execution_report_context(
);
object.insert(
"original_request_body".to_string(),
crate::ai_pipeline::build_report_context_original_request_echo(parts.original_request_body)
.unwrap_or(Value::Null),
crate::ai_pipeline::build_report_context_original_request_echo(
parts.original_request_body_json,
parts.original_request_body_base64,
)
.unwrap_or(Value::Null),
);
object.insert("has_envelope".to_string(), Value::Bool(parts.has_envelope));
object.insert(

View File

@@ -49,7 +49,6 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
.await?;
let LocalGeminiFilesCandidateAttempt {
eligible,
candidate_group_id,
candidate_id,
..
} = attempt;
@@ -66,6 +65,41 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
}
extra_fields.insert("file_key_id".to_string(), json!(candidate.key_id));
extra_fields.insert("file_name".to_string(), json!(resolved.file_name));
let report_context = build_local_execution_report_context(LocalExecutionReportContextParts {
auth_context: &input.auth_context,
request_id: trace_id,
candidate_id: &candidate_id,
attempt_identity,
model: "gemini-files",
provider_name: &transport.provider.name,
provider_id: &candidate.provider_id,
endpoint_id: &candidate.endpoint_id,
key_id: &candidate.key_id,
key_name: None,
provider_api_format: GEMINI_FILES_CLIENT_API_FORMAT,
client_api_format: GEMINI_FILES_CLIENT_API_FORMAT,
mapped_model: None,
candidate_group_id: eligible.orchestration.candidate_group_id.as_deref(),
upstream_url: None,
provider_request_method: None,
provider_request_headers: None,
original_headers: &parts.headers,
original_request_body_json: Some(body_json),
original_request_body_base64: resolved.provider_request_body_base64.as_deref(),
has_envelope: false,
needs_conversion: false,
extra_fields,
});
let super::request::LocalGeminiFilesCandidatePayloadParts {
transport: _,
auth_header,
auth_value,
provider_request_headers,
provider_request_body,
provider_request_body_base64,
upstream_url,
file_name: _,
} = resolved;
Some(build_local_execution_decision_response(
LocalExecutionDecisionResponseParts {
@@ -80,18 +114,18 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
endpoint_id: candidate.endpoint_id.clone(),
key_id: candidate.key_id.clone(),
upstream_base_url: transport.endpoint.base_url.clone(),
upstream_url: resolved.upstream_url,
upstream_url,
provider_request_method: Some(parts.method.to_string()),
auth_header: Some(resolved.auth_header),
auth_value: Some(resolved.auth_value),
auth_header: Some(auth_header),
auth_value: Some(auth_value),
provider_api_format: GEMINI_FILES_CLIENT_API_FORMAT.to_string(),
client_api_format: GEMINI_FILES_CLIENT_API_FORMAT.to_string(),
model_name: "gemini-files".to_string(),
mapped_model: candidate.selected_provider_model_name.clone(),
prompt_cache_key: None,
provider_request_headers: resolved.provider_request_headers,
provider_request_body: resolved.provider_request_body,
provider_request_body_base64: resolved.provider_request_body_base64,
provider_request_headers,
provider_request_body,
provider_request_body_base64,
content_type: parts
.headers
.get(http::header::CONTENT_TYPE)
@@ -104,32 +138,7 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
timeouts: resolve_transport_execution_timeouts(&transport),
upstream_is_stream: spec_metadata.require_streaming,
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
report_context: Some(build_local_execution_report_context(
LocalExecutionReportContextParts {
auth_context: &input.auth_context,
request_id: trace_id,
candidate_id: &candidate_id,
attempt_identity,
model: "gemini-files",
provider_name: &transport.provider.name,
provider_id: &candidate.provider_id,
endpoint_id: &candidate.endpoint_id,
key_id: &candidate.key_id,
key_name: None,
provider_api_format: GEMINI_FILES_CLIENT_API_FORMAT,
client_api_format: GEMINI_FILES_CLIENT_API_FORMAT,
mapped_model: None,
candidate_group_id: candidate_group_id.as_deref(),
upstream_url: None,
provider_request_method: None,
provider_request_headers: None,
original_headers: &parts.headers,
original_request_body: &resolved.original_request_body,
has_envelope: false,
needs_conversion: false,
extra_fields,
},
)),
report_context: Some(report_context),
auth_context: input.auth_context.clone(),
},
))

View File

@@ -1,4 +1,5 @@
use std::collections::BTreeMap;
use std::sync::Arc;
use serde_json::json;
@@ -20,13 +21,12 @@ use super::support::{
use super::LocalGeminiFilesSpec;
pub(super) struct LocalGeminiFilesCandidatePayloadParts {
pub(super) transport: GatewayProviderTransportSnapshot,
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
pub(super) auth_header: String,
pub(super) auth_value: String,
pub(super) provider_request_headers: BTreeMap<String, String>,
pub(super) provider_request_body: Option<serde_json::Value>,
pub(super) provider_request_body_base64: Option<String>,
pub(super) original_request_body: serde_json::Value,
pub(super) upstream_url: String,
pub(super) file_name: String,
}
@@ -119,13 +119,6 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts(
} else {
None
};
let original_request_body = if let Some(body_bytes_b64) = provider_request_body_base64.clone() {
json!({"body_bytes_b64": body_bytes_b64})
} else if !body_is_empty {
body_json.clone()
} else {
serde_json::Value::Null
};
if provider_request_body_base64.is_some() && transport.endpoint.body_rules.is_some() {
mark_skipped_local_gemini_files_candidate(
state,
@@ -165,14 +158,22 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts(
&auth_value,
&BTreeMap::new(),
);
let null_original_request_body = serde_json::Value::Null;
let base64_original_request_body = provider_request_body_base64
.as_ref()
.map(|body_bytes_b64| json!({ "body_bytes_b64": body_bytes_b64 }));
let original_request_body = base64_original_request_body
.as_ref()
.or_else(|| (!body_is_empty).then_some(body_json))
.unwrap_or(&null_original_request_body);
if !apply_local_header_rules(
&mut provider_request_headers,
transport.endpoint.header_rules.as_ref(),
&[&auth_header, "content-type"],
provider_request_body
.as_ref()
.unwrap_or(&original_request_body),
Some(&original_request_body),
.unwrap_or(original_request_body),
Some(original_request_body),
) {
mark_skipped_local_gemini_files_candidate(
state,
@@ -195,13 +196,12 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts(
.to_string();
Some(LocalGeminiFilesCandidatePayloadParts {
transport: transport.clone(),
transport: Arc::clone(transport),
auth_header,
auth_value,
provider_request_headers,
provider_request_body,
provider_request_body_base64,
original_request_body,
upstream_url,
file_name,
})

View File

@@ -145,7 +145,7 @@ pub(super) async fn materialize_local_gemini_files_candidate_attempts(
skipped_candidate.extra_data =
Some(build_local_execution_candidate_metadata_for_candidate(
&skipped_candidate.candidate,
skipped_candidate.transport.as_ref(),
skipped_candidate.transport_ref(),
GEMINI_FILES_CLIENT_API_FORMAT,
GEMINI_FILES_CLIENT_API_FORMAT,
extra_fields,

View File

@@ -34,7 +34,6 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
.await?;
let LocalVideoCreateCandidateAttempt {
eligible,
candidate_group_id,
candidate_id,
..
} = attempt;
@@ -49,6 +48,40 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) {
extra_fields.insert("proxy".to_string(), proxy_value);
}
let report_context = build_local_execution_report_context(LocalExecutionReportContextParts {
auth_context: &input.auth_context,
request_id: trace_id,
candidate_id: &candidate_id,
attempt_identity,
model: &input.requested_model,
provider_name: &transport.provider.name,
provider_id: &candidate.provider_id,
endpoint_id: &candidate.endpoint_id,
key_id: &candidate.key_id,
key_name: None,
provider_api_format: spec_metadata.api_format,
client_api_format: spec_metadata.api_format,
mapped_model: Some(&resolved.mapped_model),
candidate_group_id: eligible.orchestration.candidate_group_id.as_deref(),
upstream_url: None,
provider_request_method: None,
provider_request_headers: None,
original_headers: &parts.headers,
original_request_body_json: Some(body_json),
original_request_body_base64: None,
has_envelope: false,
needs_conversion: false,
extra_fields,
});
let super::request::LocalVideoCreateCandidatePayloadParts {
transport: _,
auth_header,
auth_value,
mapped_model,
provider_request_headers,
provider_request_body,
upstream_url,
} = resolved;
Some(build_local_execution_decision_response(
LocalExecutionDecisionResponseParts {
@@ -63,17 +96,17 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
endpoint_id: candidate.endpoint_id.clone(),
key_id: candidate.key_id.clone(),
upstream_base_url: transport.endpoint.base_url.clone(),
upstream_url: resolved.upstream_url,
upstream_url,
provider_request_method: Some(parts.method.to_string()),
auth_header: Some(resolved.auth_header),
auth_value: Some(resolved.auth_value),
auth_header: Some(auth_header),
auth_value: Some(auth_value),
provider_api_format: spec_metadata.api_format.to_string(),
client_api_format: spec_metadata.api_format.to_string(),
model_name: input.requested_model.clone(),
mapped_model: resolved.mapped_model.clone(),
mapped_model,
prompt_cache_key: None,
provider_request_headers: resolved.provider_request_headers,
provider_request_body: Some(resolved.provider_request_body),
provider_request_headers,
provider_request_body: Some(provider_request_body),
provider_request_body_base64: None,
content_type: parts
.headers
@@ -87,32 +120,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
timeouts: resolve_transport_execution_timeouts(&transport),
upstream_is_stream: false,
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
report_context: Some(build_local_execution_report_context(
LocalExecutionReportContextParts {
auth_context: &input.auth_context,
request_id: trace_id,
candidate_id: &candidate_id,
attempt_identity,
model: &input.requested_model,
provider_name: &transport.provider.name,
provider_id: &candidate.provider_id,
endpoint_id: &candidate.endpoint_id,
key_id: &candidate.key_id,
key_name: None,
provider_api_format: spec_metadata.api_format,
client_api_format: spec_metadata.api_format,
mapped_model: Some(&resolved.mapped_model),
candidate_group_id: candidate_group_id.as_deref(),
upstream_url: None,
provider_request_method: None,
provider_request_headers: None,
original_headers: &parts.headers,
original_request_body: body_json,
has_envelope: false,
needs_conversion: false,
extra_fields,
},
)),
report_context: Some(report_context),
auth_context: input.auth_context.clone(),
},
))

View File

@@ -1,4 +1,5 @@
use std::collections::BTreeMap;
use std::sync::Arc;
use serde_json::Value;
@@ -26,7 +27,7 @@ use super::support::{
use super::{LocalVideoCreateFamily, LocalVideoCreateSpec};
pub(super) struct LocalVideoCreateCandidatePayloadParts {
pub(super) transport: GatewayProviderTransportSnapshot,
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
pub(super) auth_header: String,
pub(super) auth_value: String,
pub(super) mapped_model: String,
@@ -168,7 +169,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
}
Some(LocalVideoCreateCandidatePayloadParts {
transport: transport.clone(),
transport: Arc::clone(transport),
auth_header,
auth_value,
mapped_model,

View File

@@ -212,7 +212,7 @@ async fn materialize_local_video_create_candidate_attempts(
skipped_candidate.extra_data =
Some(build_local_execution_candidate_metadata_for_candidate(
&skipped_candidate.candidate,
skipped_candidate.transport.as_ref(),
skipped_candidate.transport_ref(),
api_format,
api_format,
serde_json::Map::new(),

View File

@@ -213,7 +213,7 @@ pub(super) async fn materialize_local_standard_candidate_attempts(
skipped_candidate.extra_data = Some(
build_local_execution_candidate_contract_metadata_for_candidate(
&skipped_candidate.candidate,
skipped_candidate.transport.as_ref(),
skipped_candidate.transport_ref(),
provider_api_format.as_str(),
spec_metadata.api_format,
serde_json::Map::new(),

View File

@@ -35,7 +35,6 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
let LocalStandardCandidateAttempt {
eligible,
candidate_index,
candidate_group_id,
candidate_id,
..
} = &attempt;
@@ -53,6 +52,53 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
{
extra_fields.insert("proxy".to_string(), proxy_value);
}
let report_context = append_local_failover_policy_to_value(
append_execution_contract_fields_to_value(
build_local_execution_report_context(LocalExecutionReportContextParts {
auth_context: &input.auth_context,
request_id: trace_id,
candidate_id,
attempt_identity: attempt.attempt_identity(),
model: &input.requested_model,
provider_name: &candidate.provider_name,
provider_id: &candidate.provider_id,
endpoint_id: &candidate.endpoint_id,
key_id: &candidate.key_id,
key_name: Some(&candidate.key_name),
provider_api_format: &resolved.provider_api_format,
client_api_format: spec_metadata.api_format,
mapped_model: Some(&resolved.mapped_model),
candidate_group_id: eligible.orchestration.candidate_group_id.as_deref(),
upstream_url: Some(&resolved.upstream_url),
provider_request_method: Some(serde_json::Value::Null),
provider_request_headers: Some(&resolved.provider_request_headers),
original_headers: &parts.headers,
original_request_body_json: Some(body_json),
original_request_body_base64: None,
has_envelope: false,
needs_conversion: true,
extra_fields,
}),
ExecutionStrategy::LocalCrossFormat,
ConversionMode::Bidirectional,
spec_metadata.api_format,
candidate.endpoint_api_format.as_str(),
),
&resolved.transport,
);
let tls_profile = resolve_transport_tls_profile(&resolved.transport);
let timeouts = resolve_transport_execution_timeouts(&resolved.transport);
let super::request::LocalStandardCandidatePayloadParts {
auth_header,
auth_value,
mapped_model,
provider_api_format,
provider_request_body,
provider_request_headers,
upstream_url,
upstream_is_stream,
transport,
} = resolved;
Some(build_local_execution_decision_response(
LocalExecutionDecisionResponseParts {
@@ -66,58 +112,26 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
provider_id: candidate.provider_id.clone(),
endpoint_id: candidate.endpoint_id.clone(),
key_id: candidate.key_id.clone(),
upstream_base_url: resolved.transport.endpoint.base_url.clone(),
upstream_url: resolved.upstream_url.clone(),
upstream_base_url: transport.endpoint.base_url.clone(),
upstream_url,
provider_request_method: None,
auth_header: Some(resolved.auth_header.clone()),
auth_value: Some(resolved.auth_value.clone()),
provider_api_format: resolved.provider_api_format.clone(),
auth_header: Some(auth_header),
auth_value: Some(auth_value),
provider_api_format,
client_api_format: spec_metadata.api_format.to_string(),
model_name: input.requested_model.clone(),
mapped_model: resolved.mapped_model.clone(),
mapped_model,
prompt_cache_key: None,
provider_request_headers: resolved.provider_request_headers.clone(),
provider_request_body: Some(resolved.provider_request_body.clone()),
provider_request_headers,
provider_request_body: Some(provider_request_body),
provider_request_body_base64: None,
content_type: Some("application/json".to_string()),
proxy,
tls_profile: resolve_transport_tls_profile(&resolved.transport),
timeouts: resolve_transport_execution_timeouts(&resolved.transport),
upstream_is_stream: resolved.upstream_is_stream,
tls_profile,
timeouts,
upstream_is_stream,
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
report_context: Some(append_local_failover_policy_to_value(
append_execution_contract_fields_to_value(
build_local_execution_report_context(LocalExecutionReportContextParts {
auth_context: &input.auth_context,
request_id: trace_id,
candidate_id,
attempt_identity: attempt.attempt_identity(),
model: &input.requested_model,
provider_name: &candidate.provider_name,
provider_id: &candidate.provider_id,
endpoint_id: &candidate.endpoint_id,
key_id: &candidate.key_id,
key_name: Some(&candidate.key_name),
provider_api_format: &resolved.provider_api_format,
client_api_format: spec_metadata.api_format,
mapped_model: Some(&resolved.mapped_model),
candidate_group_id: candidate_group_id.as_deref(),
upstream_url: Some(&resolved.upstream_url),
provider_request_method: Some(serde_json::Value::Null),
provider_request_headers: Some(&resolved.provider_request_headers),
original_headers: &parts.headers,
original_request_body: body_json,
has_envelope: false,
needs_conversion: true,
extra_fields,
}),
ExecutionStrategy::LocalCrossFormat,
ConversionMode::Bidirectional,
spec_metadata.api_format,
candidate.endpoint_api_format.as_str(),
),
&resolved.transport,
)),
report_context: Some(report_context),
auth_context: input.auth_context.clone(),
},
))

View File

@@ -1,4 +1,5 @@
use std::collections::BTreeMap;
use std::sync::Arc;
use serde_json::Value;
@@ -27,7 +28,7 @@ pub(crate) struct LocalStandardCandidatePayloadParts {
pub(super) provider_request_headers: BTreeMap<String, String>,
pub(super) upstream_url: String,
pub(super) upstream_is_stream: bool,
pub(super) transport: GatewayProviderTransportSnapshot,
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
}
pub(crate) async fn resolve_local_standard_candidate_payload_parts(
@@ -220,6 +221,6 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
provider_request_headers,
upstream_url,
upstream_is_stream,
transport: transport.clone(),
transport: Arc::clone(transport),
})
}

View File

@@ -2,7 +2,7 @@ use aether_contracts::{ExecutionPlan, RequestBody};
use super::{
augment_sync_report_context, generic_decision_missing_exact_provider_request,
LocalStreamPlanAndReport, LocalSyncPlanAndReport,
take_non_empty_string, LocalStreamPlanAndReport, LocalSyncPlanAndReport,
};
use crate::ai_pipeline::transport::ensure_upstream_auth_header;
use crate::{GatewayControlSyncDecisionResponse, GatewayError};
@@ -12,74 +12,41 @@ pub(crate) fn build_gemini_sync_plan_from_decision(
_body_json: &serde_json::Value,
payload: GatewayControlSyncDecisionResponse,
) -> Result<Option<LocalSyncPlanAndReport>, GatewayError> {
let mut payload = payload;
if generic_decision_missing_exact_provider_request(&payload) {
return Ok(None);
}
let Some(request_id) = payload
.request_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
return Ok(None);
};
let Some(provider_id) = payload
.provider_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
return Ok(None);
};
let Some(endpoint_id) = payload
.endpoint_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
return Ok(None);
};
let Some(key_id) = payload
.key_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
return Ok(None);
};
let Some(url) = payload
.upstream_url
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(url) = take_non_empty_string(&mut payload.upstream_url) else {
return Ok(None);
};
let auth_header = payload
.auth_header
.clone()
.filter(|value| !value.trim().is_empty());
let auth_value = payload
.auth_value
.clone()
.filter(|value| !value.trim().is_empty());
let auth_header = take_non_empty_string(&mut payload.auth_header);
let auth_value = take_non_empty_string(&mut payload.auth_value);
if auth_header.is_some() != auth_value.is_some() {
return Ok(None);
}
let Some(provider_api_format) = payload
.provider_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
return Ok(None);
};
let Some(client_api_format) = payload
.client_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
return Ok(None);
};
let Some(provider_request_body_value) = payload.provider_request_body.clone() else {
let Some(provider_request_body_value) = payload.provider_request_body.take() else {
return Ok(None);
};
let mut provider_request_headers = payload.provider_request_headers.clone();
let mut provider_request_headers = std::mem::take(&mut payload.provider_request_headers);
if let (Some(auth_header), Some(auth_value)) = (auth_header.as_deref(), auth_value.as_deref()) {
ensure_upstream_auth_header(&mut provider_request_headers, auth_header, auth_value);
}
@@ -88,37 +55,37 @@ pub(crate) fn build_gemini_sync_plan_from_decision(
.entry("accept".to_string())
.or_insert_with(|| "text/event-stream".to_string());
}
let content_type = payload
.content_type
.take()
.or_else(|| Some("application/json".to_string()));
let report_context = augment_sync_report_context(
payload.report_context.take(),
&provider_request_headers,
&provider_request_body_value,
)?;
let plan = ExecutionPlan {
request_id,
candidate_id: payload.candidate_id.clone(),
provider_name: payload.provider_name.clone(),
candidate_id: payload.candidate_id.take(),
provider_name: payload.provider_name.take(),
provider_id,
endpoint_id,
key_id,
method: "POST".to_string(),
url,
headers: std::mem::take(&mut provider_request_headers),
content_type: payload
.content_type
.clone()
.or_else(|| Some("application/json".to_string())),
content_type,
content_encoding: None,
body: RequestBody::from_json(provider_request_body_value.clone()),
body: RequestBody::from_json(provider_request_body_value),
stream: payload.upstream_is_stream,
client_api_format,
provider_api_format,
model_name: payload.model_name.clone(),
proxy: payload.proxy.clone(),
tls_profile: payload.tls_profile.clone(),
timeouts: payload.timeouts.clone(),
model_name: payload.model_name.take(),
proxy: payload.proxy.take(),
tls_profile: payload.tls_profile.take(),
timeouts: payload.timeouts.take(),
};
let report_context = augment_sync_report_context(
payload.report_context,
&plan.headers,
&provider_request_body_value,
)?;
Ok(Some(LocalSyncPlanAndReport {
plan,
report_kind: payload.report_kind,
@@ -131,109 +98,76 @@ pub(crate) fn build_gemini_stream_plan_from_decision(
_body_json: &serde_json::Value,
payload: GatewayControlSyncDecisionResponse,
) -> Result<Option<LocalStreamPlanAndReport>, GatewayError> {
let mut payload = payload;
if generic_decision_missing_exact_provider_request(&payload) {
return Ok(None);
}
let Some(request_id) = payload
.request_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
return Ok(None);
};
let Some(provider_id) = payload
.provider_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
return Ok(None);
};
let Some(endpoint_id) = payload
.endpoint_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
return Ok(None);
};
let Some(key_id) = payload
.key_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
return Ok(None);
};
let Some(url) = payload
.upstream_url
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(url) = take_non_empty_string(&mut payload.upstream_url) else {
return Ok(None);
};
let auth_header = payload
.auth_header
.clone()
.filter(|value| !value.trim().is_empty());
let auth_value = payload
.auth_value
.clone()
.filter(|value| !value.trim().is_empty());
let auth_header = take_non_empty_string(&mut payload.auth_header);
let auth_value = take_non_empty_string(&mut payload.auth_value);
if auth_header.is_some() != auth_value.is_some() {
return Ok(None);
}
let Some(provider_api_format) = payload
.provider_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
return Ok(None);
};
let Some(client_api_format) = payload
.client_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
return Ok(None);
};
let Some(provider_request_body_value) = payload.provider_request_body.clone() else {
let Some(provider_request_body_value) = payload.provider_request_body.take() else {
return Ok(None);
};
let mut provider_request_headers = payload.provider_request_headers.clone();
let mut provider_request_headers = std::mem::take(&mut payload.provider_request_headers);
if let (Some(auth_header), Some(auth_value)) = (auth_header.as_deref(), auth_value.as_deref()) {
ensure_upstream_auth_header(&mut provider_request_headers, auth_header, auth_value);
}
provider_request_headers.insert("accept".to_string(), "text/event-stream".to_string());
let content_type = payload
.content_type
.take()
.or_else(|| Some("application/json".to_string()));
let report_context = augment_sync_report_context(
payload.report_context.take(),
&provider_request_headers,
&provider_request_body_value,
)?;
let plan = ExecutionPlan {
request_id,
candidate_id: payload.candidate_id.clone(),
provider_name: payload.provider_name.clone(),
candidate_id: payload.candidate_id.take(),
provider_name: payload.provider_name.take(),
provider_id,
endpoint_id,
key_id,
method: "POST".to_string(),
url,
headers: std::mem::take(&mut provider_request_headers),
content_type: payload
.content_type
.clone()
.or_else(|| Some("application/json".to_string())),
content_type,
content_encoding: None,
body: RequestBody::from_json(provider_request_body_value.clone()),
body: RequestBody::from_json(provider_request_body_value),
stream: true,
client_api_format,
provider_api_format,
model_name: payload.model_name.clone(),
proxy: payload.proxy.clone(),
tls_profile: payload.tls_profile.clone(),
timeouts: payload.timeouts.clone(),
model_name: payload.model_name.take(),
proxy: payload.proxy.take(),
tls_profile: payload.tls_profile.take(),
timeouts: payload.timeouts.take(),
};
let report_context = augment_sync_report_context(
payload.report_context,
&plan.headers,
&provider_request_body_value,
)?;
Ok(Some(LocalStreamPlanAndReport {
plan,
report_kind: payload.report_kind,

View File

@@ -32,7 +32,6 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
let LocalOpenAiChatCandidateAttempt {
eligible,
candidate_index,
candidate_group_id,
candidate_id,
..
} = attempt;
@@ -70,74 +69,89 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
{
extra_fields.insert("proxy".to_string(), proxy_value);
}
let report_context = append_local_failover_policy_to_value(
append_execution_contract_fields_to_value(
build_local_execution_report_context(LocalExecutionReportContextParts {
auth_context: &input.auth_context,
request_id: trace_id,
candidate_id: &candidate_id,
attempt_identity,
model: &input.requested_model,
provider_name: &resolved.transport.provider.name,
provider_id: &candidate.provider_id,
endpoint_id: &candidate.endpoint_id,
key_id: &candidate.key_id,
key_name: Some(&candidate.key_name),
provider_api_format: &resolved.provider_api_format,
client_api_format: "openai:chat",
mapped_model: Some(&resolved.mapped_model),
candidate_group_id: eligible.orchestration.candidate_group_id.as_deref(),
upstream_url: Some(&resolved.upstream_url),
provider_request_method: Some(serde_json::Value::Null),
provider_request_headers: Some(&resolved.provider_request_headers),
original_headers: &parts.headers,
original_request_body_json: Some(body_json),
original_request_body_base64: None,
has_envelope: false,
needs_conversion: matches!(
resolved.conversion_mode,
crate::ai_pipeline::ConversionMode::Bidirectional
),
extra_fields,
}),
resolved.execution_strategy,
resolved.conversion_mode,
"openai:chat",
candidate.endpoint_api_format.as_str(),
),
&resolved.transport,
);
let super::request::LocalOpenAiChatCandidatePayloadParts {
auth_header,
auth_value,
mapped_model,
provider_api_format,
provider_request_body,
provider_request_headers,
upstream_url,
execution_strategy,
conversion_mode,
report_kind,
transport,
} = resolved;
Some(build_local_execution_decision_response(
LocalExecutionDecisionResponseParts {
decision_is_stream: upstream_is_stream,
decision_kind: decision_kind.to_string(),
execution_strategy: resolved.execution_strategy,
conversion_mode: resolved.conversion_mode,
execution_strategy,
conversion_mode,
request_id: trace_id.to_string(),
candidate_id: candidate_id.clone(),
provider_name: resolved.transport.provider.name.clone(),
provider_name: transport.provider.name.clone(),
provider_id: candidate.provider_id.clone(),
endpoint_id: candidate.endpoint_id.clone(),
key_id: candidate.key_id.clone(),
upstream_base_url: resolved.transport.endpoint.base_url.clone(),
upstream_url: resolved.upstream_url.clone(),
upstream_base_url: transport.endpoint.base_url.clone(),
upstream_url,
provider_request_method: None,
auth_header: Some(resolved.auth_header.clone()),
auth_value: Some(resolved.auth_value.clone()),
provider_api_format: resolved.provider_api_format.clone(),
auth_header: Some(auth_header),
auth_value: Some(auth_value),
provider_api_format,
client_api_format: "openai:chat".to_string(),
model_name: input.requested_model.clone(),
mapped_model: resolved.mapped_model.clone(),
mapped_model,
prompt_cache_key,
provider_request_headers: resolved.provider_request_headers.clone(),
provider_request_body: Some(resolved.provider_request_body.clone()),
provider_request_headers,
provider_request_body: Some(provider_request_body),
provider_request_body_base64: None,
content_type: Some("application/json".to_string()),
proxy,
tls_profile,
timeouts,
upstream_is_stream,
report_kind: Some(resolved.report_kind.clone()),
report_context: Some(append_local_failover_policy_to_value(
append_execution_contract_fields_to_value(
build_local_execution_report_context(LocalExecutionReportContextParts {
auth_context: &input.auth_context,
request_id: trace_id,
candidate_id: &candidate_id,
attempt_identity,
model: &input.requested_model,
provider_name: &resolved.transport.provider.name,
provider_id: &candidate.provider_id,
endpoint_id: &candidate.endpoint_id,
key_id: &candidate.key_id,
key_name: Some(&candidate.key_name),
provider_api_format: &resolved.provider_api_format,
client_api_format: "openai:chat",
mapped_model: Some(&resolved.mapped_model),
candidate_group_id: candidate_group_id.as_deref(),
upstream_url: Some(&resolved.upstream_url),
provider_request_method: Some(serde_json::Value::Null),
provider_request_headers: Some(&resolved.provider_request_headers),
original_headers: &parts.headers,
original_request_body: body_json,
has_envelope: false,
needs_conversion: matches!(
resolved.conversion_mode,
crate::ai_pipeline::ConversionMode::Bidirectional
),
extra_fields,
}),
resolved.execution_strategy,
resolved.conversion_mode,
"openai:chat",
candidate.endpoint_api_format.as_str(),
),
&resolved.transport,
)),
report_kind: Some(report_kind),
report_context: Some(report_context),
auth_context: input.auth_context.clone(),
},
))

View File

@@ -1,4 +1,5 @@
use std::collections::BTreeMap;
use std::sync::Arc;
use serde_json::Value;
@@ -36,7 +37,7 @@ pub(crate) struct LocalOpenAiChatCandidatePayloadParts {
pub(super) execution_strategy: ExecutionStrategy,
pub(super) conversion_mode: ConversionMode,
pub(super) report_kind: String,
pub(super) transport: GatewayProviderTransportSnapshot,
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
}
#[allow(clippy::too_many_arguments)]
@@ -192,7 +193,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
execution_strategy: ExecutionStrategy::LocalSameFormat,
conversion_mode: ConversionMode::None,
report_kind: report_kind.to_string(),
transport: transport.clone(),
transport: Arc::clone(transport),
});
}
@@ -374,6 +375,6 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
execution_strategy: ExecutionStrategy::LocalCrossFormat,
conversion_mode: ConversionMode::Bidirectional,
report_kind: resolved_report_kind,
transport: transport.clone(),
transport: Arc::clone(transport),
})
}

View File

@@ -103,7 +103,7 @@ pub(crate) async fn materialize_local_openai_chat_candidate_attempts(
skipped_candidate.extra_data = Some(
build_local_execution_candidate_contract_metadata_for_candidate(
&skipped_candidate.candidate,
skipped_candidate.transport.as_ref(),
skipped_candidate.transport_ref(),
provider_api_format.as_str(),
"openai:chat",
serde_json::Map::new(),

View File

@@ -35,7 +35,6 @@ pub(crate) async fn maybe_build_local_openai_cli_decision_payload_for_candidate(
let LocalOpenAiCliCandidateAttempt {
eligible,
candidate_index,
candidate_group_id,
candidate_id,
..
} = attempt;
@@ -74,6 +73,43 @@ pub(crate) async fn maybe_build_local_openai_cli_decision_payload_for_candidate(
if resolved.is_antigravity {
extra_fields.insert("envelope_name".to_string(), json!("antigravity:v1internal"));
}
let report_context = append_local_failover_policy_to_value(
append_execution_contract_fields_to_value(
build_local_execution_report_context(LocalExecutionReportContextParts {
auth_context: &input.auth_context,
request_id: trace_id,
candidate_id: &candidate_id,
attempt_identity,
model: &input.requested_model,
provider_name: &resolved.transport.provider.name,
provider_id: &candidate.provider_id,
endpoint_id: &candidate.endpoint_id,
key_id: &candidate.key_id,
key_name: Some(&candidate.key_name),
provider_api_format: &resolved.provider_api_format,
client_api_format: spec_metadata.api_format,
mapped_model: Some(&resolved.mapped_model),
candidate_group_id: eligible.orchestration.candidate_group_id.as_deref(),
upstream_url: Some(&resolved.upstream_url),
provider_request_method: Some(serde_json::Value::Null),
provider_request_headers: Some(&resolved.provider_request_headers),
original_headers: &parts.headers,
original_request_body_json: Some(body_json),
original_request_body_base64: None,
has_envelope: resolved.is_antigravity,
needs_conversion: matches!(
resolved.conversion_mode,
crate::ai_pipeline::ConversionMode::Bidirectional
),
extra_fields,
}),
resolved.execution_strategy,
resolved.conversion_mode,
spec_metadata.api_format,
candidate.endpoint_api_format.as_str(),
),
&resolved.transport,
);
debug!(
event_name = "local_openai_cli_decision_payload_built",
@@ -98,74 +134,53 @@ pub(crate) async fn maybe_build_local_openai_cli_decision_payload_for_candidate(
has_envelope = resolved.is_antigravity,
"gateway built local openai cli decision payload"
);
let super::request::LocalOpenAiCliCandidatePayloadParts {
auth_header,
auth_value,
mapped_model,
provider_api_format,
provider_request_body,
provider_request_headers,
upstream_url,
execution_strategy,
conversion_mode,
is_antigravity: _,
upstream_is_stream,
transport,
} = resolved;
Some(build_local_execution_decision_response(
LocalExecutionDecisionResponseParts {
decision_is_stream: spec_metadata.require_streaming,
decision_kind: spec_metadata.decision_kind.to_string(),
execution_strategy: resolved.execution_strategy,
conversion_mode: resolved.conversion_mode,
execution_strategy,
conversion_mode,
request_id: trace_id.to_string(),
candidate_id: candidate_id.clone(),
provider_name: resolved.transport.provider.name.clone(),
provider_name: transport.provider.name.clone(),
provider_id: candidate.provider_id.clone(),
endpoint_id: candidate.endpoint_id.clone(),
key_id: candidate.key_id.clone(),
upstream_base_url: resolved.transport.endpoint.base_url.clone(),
upstream_url: resolved.upstream_url.clone(),
upstream_base_url: transport.endpoint.base_url.clone(),
upstream_url,
provider_request_method: None,
auth_header: Some(resolved.auth_header.clone()),
auth_value: Some(resolved.auth_value.clone()),
provider_api_format: resolved.provider_api_format.clone(),
auth_header: Some(auth_header),
auth_value: Some(auth_value),
provider_api_format,
client_api_format: spec_metadata.api_format.to_string(),
model_name: input.requested_model.clone(),
mapped_model: resolved.mapped_model.clone(),
mapped_model,
prompt_cache_key,
provider_request_headers: resolved.provider_request_headers.clone(),
provider_request_body: Some(resolved.provider_request_body.clone()),
provider_request_headers,
provider_request_body: Some(provider_request_body),
provider_request_body_base64: None,
content_type: Some("application/json".to_string()),
proxy,
tls_profile,
timeouts,
upstream_is_stream: resolved.upstream_is_stream,
upstream_is_stream,
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
report_context: Some(append_local_failover_policy_to_value(
append_execution_contract_fields_to_value(
build_local_execution_report_context(LocalExecutionReportContextParts {
auth_context: &input.auth_context,
request_id: trace_id,
candidate_id: &candidate_id,
attempt_identity,
model: &input.requested_model,
provider_name: &resolved.transport.provider.name,
provider_id: &candidate.provider_id,
endpoint_id: &candidate.endpoint_id,
key_id: &candidate.key_id,
key_name: Some(&candidate.key_name),
provider_api_format: &resolved.provider_api_format,
client_api_format: spec_metadata.api_format,
mapped_model: Some(&resolved.mapped_model),
candidate_group_id: candidate_group_id.as_deref(),
upstream_url: Some(&resolved.upstream_url),
provider_request_method: Some(serde_json::Value::Null),
provider_request_headers: Some(&resolved.provider_request_headers),
original_headers: &parts.headers,
original_request_body: body_json,
has_envelope: resolved.is_antigravity,
needs_conversion: matches!(
resolved.conversion_mode,
crate::ai_pipeline::ConversionMode::Bidirectional
),
extra_fields,
}),
resolved.execution_strategy,
resolved.conversion_mode,
spec_metadata.api_format,
candidate.endpoint_api_format.as_str(),
),
&resolved.transport,
)),
report_context: Some(report_context),
auth_context: input.auth_context.clone(),
},
))

View File

@@ -1,4 +1,5 @@
use std::collections::BTreeMap;
use std::sync::Arc;
use serde_json::Value;
use tracing::debug;
@@ -48,7 +49,7 @@ pub(crate) struct LocalOpenAiCliCandidatePayloadParts {
pub(super) conversion_mode: ConversionMode,
pub(super) is_antigravity: bool,
pub(super) upstream_is_stream: bool,
pub(super) transport: GatewayProviderTransportSnapshot,
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
}
#[allow(clippy::too_many_arguments)]
@@ -379,6 +380,6 @@ pub(crate) async fn resolve_local_openai_cli_candidate_payload_parts(
is_antigravity: is_antigravity
|| antigravity_auth.is_some() && ANTIGRAVITY_ENVELOPE_NAME == "antigravity:v1internal",
upstream_is_stream,
transport: transport.clone(),
transport: Arc::clone(transport),
})
}

View File

@@ -210,7 +210,7 @@ pub(crate) async fn materialize_local_openai_cli_candidate_attempts(
skipped_candidate.extra_data = Some(
build_local_execution_candidate_contract_metadata_for_candidate(
&skipped_candidate.candidate,
skipped_candidate.transport.as_ref(),
skipped_candidate.transport_ref(),
provider_api_format.as_str(),
spec_metadata.api_format,
serde_json::Map::new(),

View File

@@ -3,7 +3,7 @@ use tracing::debug;
use super::super::{
augment_sync_report_context, generic_decision_missing_exact_provider_request,
LocalStreamPlanAndReport,
take_non_empty_string, LocalStreamPlanAndReport,
};
use crate::ai_pipeline::provider_adaptation_requires_eventstream_accept;
use crate::ai_pipeline::transport::auth::{
@@ -18,79 +18,40 @@ pub(crate) fn build_openai_chat_stream_plan_from_decision(
body_json: &serde_json::Value,
payload: GatewayControlSyncDecisionResponse,
) -> Result<Option<LocalStreamPlanAndReport>, GatewayError> {
let Some(request_id) = payload
.request_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let mut payload = payload;
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
return Ok(None);
};
let Some(provider_id) = payload
.provider_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
return Ok(None);
};
let Some(endpoint_id) = payload
.endpoint_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
return Ok(None);
};
let Some(key_id) = payload
.key_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
return Ok(None);
};
let Some(auth_header) = payload
.auth_header
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(auth_header) = take_non_empty_string(&mut payload.auth_header) else {
return Ok(None);
};
let Some(auth_value) = payload
.auth_value
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(auth_value) = take_non_empty_string(&mut payload.auth_value) else {
return Ok(None);
};
let Some(provider_api_format) = payload
.provider_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
return Ok(None);
};
let Some(client_api_format) = payload
.client_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
return Ok(None);
};
let url = if let Some(upstream_url) = payload
.upstream_url
.clone()
.filter(|value| !value.trim().is_empty())
{
let url = if let Some(upstream_url) = take_non_empty_string(&mut payload.upstream_url) {
upstream_url
} else {
let Some(upstream_base_url) = payload
.upstream_base_url
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(upstream_base_url) = take_non_empty_string(&mut payload.upstream_base_url) else {
return Ok(None);
};
build_openai_chat_url(&upstream_base_url, parts.uri.query())
};
let provider_request_body_value = if let Some(body) = payload.provider_request_body.clone() {
let provider_request_body_value = if let Some(body) = payload.provider_request_body.take() {
body
} else {
let Some(request_body_object) = body_json.as_object() else {
@@ -102,22 +63,12 @@ pub(crate) fn build_openai_chat_stream_plan_from_decision(
.iter()
.map(|(key, value)| (key.clone(), value.clone())),
);
if let Some(mapped_model) = payload
.mapped_model
.as_ref()
.filter(|value| !value.trim().is_empty())
{
provider_request_body.insert(
"model".to_string(),
serde_json::Value::String(mapped_model.clone()),
);
if let Some(mapped_model) = take_non_empty_string(&mut payload.mapped_model) {
provider_request_body
.insert("model".to_string(), serde_json::Value::String(mapped_model));
}
provider_request_body.insert("stream".to_string(), serde_json::Value::Bool(true));
if let Some(prompt_cache_key) = payload
.prompt_cache_key
.as_ref()
.filter(|value| !value.trim().is_empty())
{
if let Some(prompt_cache_key) = take_non_empty_string(&mut payload.prompt_cache_key) {
let existing = provider_request_body
.get("prompt_cache_key")
.and_then(|value| value.as_str())
@@ -126,20 +77,21 @@ pub(crate) fn build_openai_chat_stream_plan_from_decision(
if existing.is_empty() {
provider_request_body.insert(
"prompt_cache_key".to_string(),
serde_json::Value::String(prompt_cache_key.clone()),
serde_json::Value::String(prompt_cache_key),
);
}
}
serde_json::Value::Object(provider_request_body)
};
let mut provider_request_headers = if payload.provider_request_headers.is_empty() {
let existing_provider_request_headers = std::mem::take(&mut payload.provider_request_headers);
let extra_headers = std::mem::take(&mut payload.extra_headers);
let mut provider_request_headers = if existing_provider_request_headers.is_empty() {
if provider_api_format == client_api_format {
build_complete_passthrough_headers_with_auth(
&parts.headers,
&auth_header,
&auth_value,
&payload.extra_headers,
&extra_headers,
payload.content_type.as_deref(),
)
} else if provider_api_format.starts_with("claude:") {
@@ -147,7 +99,7 @@ pub(crate) fn build_openai_chat_stream_plan_from_decision(
&parts.headers,
&auth_header,
&auth_value,
&payload.extra_headers,
&extra_headers,
payload.content_type.as_deref(),
)
} else {
@@ -155,46 +107,46 @@ pub(crate) fn build_openai_chat_stream_plan_from_decision(
&parts.headers,
&auth_header,
&auth_value,
&payload.extra_headers,
&extra_headers,
payload.content_type.as_deref(),
)
}
} else {
payload.provider_request_headers.clone()
existing_provider_request_headers
};
ensure_upstream_auth_header(&mut provider_request_headers, &auth_header, &auth_value);
provider_request_headers.insert("accept".to_string(), "text/event-stream".to_string());
let content_type = payload
.content_type
.take()
.or_else(|| Some("application/json".to_string()));
let report_context = augment_sync_report_context(
payload.report_context.take(),
&provider_request_headers,
&provider_request_body_value,
)?;
let plan = ExecutionPlan {
request_id,
candidate_id: payload.candidate_id.clone(),
provider_name: payload.provider_name.clone(),
candidate_id: payload.candidate_id.take(),
provider_name: payload.provider_name.take(),
provider_id,
endpoint_id,
key_id,
method: "POST".to_string(),
url,
headers: std::mem::take(&mut provider_request_headers),
content_type: payload
.content_type
.clone()
.or_else(|| Some("application/json".to_string())),
content_type,
content_encoding: None,
body: RequestBody::from_json(provider_request_body_value.clone()),
body: RequestBody::from_json(provider_request_body_value),
stream: true,
client_api_format,
provider_api_format,
model_name: payload.model_name.clone(),
proxy: payload.proxy.clone(),
tls_profile: payload.tls_profile.clone(),
timeouts: payload.timeouts.clone(),
model_name: payload.model_name.take(),
proxy: payload.proxy.take(),
tls_profile: payload.tls_profile.take(),
timeouts: payload.timeouts.take(),
};
let report_context = augment_sync_report_context(
payload.report_context,
&plan.headers,
&provider_request_body_value,
)?;
Ok(Some(LocalStreamPlanAndReport {
plan,
report_kind: payload.report_kind,
@@ -208,74 +160,39 @@ pub(crate) fn build_openai_cli_stream_plan_from_decision(
payload: GatewayControlSyncDecisionResponse,
compact: bool,
) -> Result<Option<LocalStreamPlanAndReport>, GatewayError> {
let mut payload = payload;
if generic_decision_missing_exact_provider_request(&payload) {
return Ok(None);
}
let Some(request_id) = payload
.request_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
return Ok(None);
};
let Some(provider_id) = payload
.provider_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
return Ok(None);
};
let Some(endpoint_id) = payload
.endpoint_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
return Ok(None);
};
let Some(key_id) = payload
.key_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
return Ok(None);
};
let auth_header = payload
.auth_header
.clone()
.filter(|value| !value.trim().is_empty());
let auth_value = payload
.auth_value
.clone()
.filter(|value| !value.trim().is_empty());
let auth_header = take_non_empty_string(&mut payload.auth_header);
let auth_value = take_non_empty_string(&mut payload.auth_value);
if auth_header.is_some() != auth_value.is_some() {
return Ok(None);
}
let Some(provider_api_format) = payload
.provider_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
return Ok(None);
};
let Some(client_api_format) = payload
.client_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
return Ok(None);
};
let (url, url_source) = if let Some(upstream_url) = payload
.upstream_url
.clone()
.filter(|value| !value.trim().is_empty())
let (url, url_source) = if let Some(upstream_url) =
take_non_empty_string(&mut payload.upstream_url)
{
(upstream_url, "upstream_url")
} else {
let Some(upstream_base_url) = payload
.upstream_base_url
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(upstream_base_url) = take_non_empty_string(&mut payload.upstream_base_url) else {
return Ok(None);
};
(
@@ -283,7 +200,7 @@ pub(crate) fn build_openai_cli_stream_plan_from_decision(
"upstream_base_url",
)
};
let Some(provider_request_body_value) = payload.provider_request_body.clone() else {
let Some(provider_request_body_value) = payload.provider_request_body.take() else {
return Ok(None);
};
@@ -292,7 +209,7 @@ pub(crate) fn build_openai_cli_stream_plan_from_decision(
.as_ref()
.and_then(|context| context.get("envelope_name"))
.and_then(serde_json::Value::as_str);
let mut provider_request_headers = payload.provider_request_headers.clone();
let mut provider_request_headers = std::mem::take(&mut payload.provider_request_headers);
if let (Some(auth_header), Some(auth_value)) = (auth_header.as_deref(), auth_value.as_deref()) {
ensure_upstream_auth_header(&mut provider_request_headers, auth_header, auth_value);
}
@@ -304,34 +221,35 @@ pub(crate) fn build_openai_cli_stream_plan_from_decision(
} else {
provider_request_headers.insert("accept".to_string(), "text/event-stream".to_string());
}
let content_type = payload
.content_type
.take()
.or_else(|| Some("application/json".to_string()));
let report_context = augment_sync_report_context(
payload.report_context,
payload.report_context.take(),
&provider_request_headers,
&provider_request_body_value,
)?;
let plan = ExecutionPlan {
request_id,
candidate_id: payload.candidate_id.clone(),
provider_name: payload.provider_name.clone(),
candidate_id: payload.candidate_id.take(),
provider_name: payload.provider_name.take(),
provider_id,
endpoint_id,
key_id,
method: "POST".to_string(),
url,
headers: std::mem::take(&mut provider_request_headers),
content_type: payload
.content_type
.clone()
.or_else(|| Some("application/json".to_string())),
content_type,
content_encoding: None,
body: RequestBody::from_json(provider_request_body_value.clone()),
body: RequestBody::from_json(provider_request_body_value),
stream: true,
client_api_format,
provider_api_format,
model_name: payload.model_name.clone(),
proxy: payload.proxy.clone(),
tls_profile: payload.tls_profile.clone(),
timeouts: payload.timeouts.clone(),
model_name: payload.model_name.take(),
proxy: payload.proxy.take(),
tls_profile: payload.tls_profile.take(),
timeouts: payload.timeouts.take(),
};
debug!(

View File

@@ -3,7 +3,7 @@ use tracing::debug;
use super::super::{
augment_sync_report_context, generic_decision_missing_exact_provider_request,
LocalSyncPlanAndReport,
take_non_empty_string, LocalSyncPlanAndReport,
};
use crate::ai_pipeline::transport::auth::{
build_claude_passthrough_headers, build_complete_passthrough_headers_with_auth,
@@ -17,79 +17,40 @@ pub(crate) fn build_openai_chat_sync_plan_from_decision(
body_json: &serde_json::Value,
payload: GatewayControlSyncDecisionResponse,
) -> Result<Option<LocalSyncPlanAndReport>, GatewayError> {
let Some(request_id) = payload
.request_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let mut payload = payload;
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
return Ok(None);
};
let Some(provider_id) = payload
.provider_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
return Ok(None);
};
let Some(endpoint_id) = payload
.endpoint_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
return Ok(None);
};
let Some(key_id) = payload
.key_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
return Ok(None);
};
let Some(auth_header) = payload
.auth_header
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(auth_header) = take_non_empty_string(&mut payload.auth_header) else {
return Ok(None);
};
let Some(auth_value) = payload
.auth_value
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(auth_value) = take_non_empty_string(&mut payload.auth_value) else {
return Ok(None);
};
let Some(provider_api_format) = payload
.provider_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
return Ok(None);
};
let Some(client_api_format) = payload
.client_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
return Ok(None);
};
let url = if let Some(upstream_url) = payload
.upstream_url
.clone()
.filter(|value| !value.trim().is_empty())
{
let url = if let Some(upstream_url) = take_non_empty_string(&mut payload.upstream_url) {
upstream_url
} else {
let Some(upstream_base_url) = payload
.upstream_base_url
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(upstream_base_url) = take_non_empty_string(&mut payload.upstream_base_url) else {
return Ok(None);
};
build_openai_chat_url(&upstream_base_url, parts.uri.query())
};
let provider_request_body_value = if let Some(body) = payload.provider_request_body.clone() {
let provider_request_body_value = if let Some(body) = payload.provider_request_body.take() {
body
} else {
let Some(request_body_object) = body_json.as_object() else {
@@ -100,24 +61,14 @@ pub(crate) fn build_openai_chat_sync_plan_from_decision(
.iter()
.map(|(key, value)| (key.clone(), value.clone())),
);
if let Some(mapped_model) = payload
.mapped_model
.as_ref()
.filter(|value| !value.trim().is_empty())
{
provider_request_body.insert(
"model".to_string(),
serde_json::Value::String(mapped_model.clone()),
);
if let Some(mapped_model) = take_non_empty_string(&mut payload.mapped_model) {
provider_request_body
.insert("model".to_string(), serde_json::Value::String(mapped_model));
}
if payload.upstream_is_stream {
provider_request_body.insert("stream".to_string(), serde_json::Value::Bool(true));
}
if let Some(prompt_cache_key) = payload
.prompt_cache_key
.as_ref()
.filter(|value| !value.trim().is_empty())
{
if let Some(prompt_cache_key) = take_non_empty_string(&mut payload.prompt_cache_key) {
let existing = provider_request_body
.get("prompt_cache_key")
.and_then(|value| value.as_str())
@@ -126,20 +77,21 @@ pub(crate) fn build_openai_chat_sync_plan_from_decision(
if existing.is_empty() {
provider_request_body.insert(
"prompt_cache_key".to_string(),
serde_json::Value::String(prompt_cache_key.clone()),
serde_json::Value::String(prompt_cache_key),
);
}
}
serde_json::Value::Object(provider_request_body)
};
let mut provider_request_headers = if payload.provider_request_headers.is_empty() {
let existing_provider_request_headers = std::mem::take(&mut payload.provider_request_headers);
let extra_headers = std::mem::take(&mut payload.extra_headers);
let mut provider_request_headers = if existing_provider_request_headers.is_empty() {
if provider_api_format == client_api_format {
build_complete_passthrough_headers_with_auth(
&parts.headers,
&auth_header,
&auth_value,
&payload.extra_headers,
&extra_headers,
payload.content_type.as_deref(),
)
} else if provider_api_format.starts_with("claude:") {
@@ -147,7 +99,7 @@ pub(crate) fn build_openai_chat_sync_plan_from_decision(
&parts.headers,
&auth_header,
&auth_value,
&payload.extra_headers,
&extra_headers,
payload.content_type.as_deref(),
)
} else {
@@ -155,12 +107,12 @@ pub(crate) fn build_openai_chat_sync_plan_from_decision(
&parts.headers,
&auth_header,
&auth_value,
&payload.extra_headers,
&extra_headers,
payload.content_type.as_deref(),
)
}
} else {
payload.provider_request_headers.clone()
existing_provider_request_headers
};
ensure_upstream_auth_header(&mut provider_request_headers, &auth_header, &auth_value);
if payload.upstream_is_stream {
@@ -168,37 +120,37 @@ pub(crate) fn build_openai_chat_sync_plan_from_decision(
.entry("accept".to_string())
.or_insert_with(|| "text/event-stream".to_string());
}
let content_type = payload
.content_type
.take()
.or_else(|| Some("application/json".to_string()));
let report_context = augment_sync_report_context(
payload.report_context.take(),
&provider_request_headers,
&provider_request_body_value,
)?;
let plan = ExecutionPlan {
request_id,
candidate_id: payload.candidate_id.clone(),
provider_name: payload.provider_name.clone(),
candidate_id: payload.candidate_id.take(),
provider_name: payload.provider_name.take(),
provider_id,
endpoint_id,
key_id,
method: "POST".to_string(),
url,
headers: std::mem::take(&mut provider_request_headers),
content_type: payload
.content_type
.clone()
.or_else(|| Some("application/json".to_string())),
content_type,
content_encoding: None,
body: RequestBody::from_json(provider_request_body_value.clone()),
body: RequestBody::from_json(provider_request_body_value),
stream: payload.upstream_is_stream,
client_api_format,
provider_api_format,
model_name: payload.model_name.clone(),
proxy: payload.proxy.clone(),
tls_profile: payload.tls_profile.clone(),
timeouts: payload.timeouts.clone(),
model_name: payload.model_name.take(),
proxy: payload.proxy.take(),
tls_profile: payload.tls_profile.take(),
timeouts: payload.timeouts.take(),
};
let report_context = augment_sync_report_context(
payload.report_context,
&plan.headers,
&provider_request_body_value,
)?;
Ok(Some(LocalSyncPlanAndReport {
plan,
report_kind: payload.report_kind,
@@ -212,74 +164,39 @@ pub(crate) fn build_openai_cli_sync_plan_from_decision(
payload: GatewayControlSyncDecisionResponse,
compact: bool,
) -> Result<Option<LocalSyncPlanAndReport>, GatewayError> {
let mut payload = payload;
if generic_decision_missing_exact_provider_request(&payload) {
return Ok(None);
}
let Some(request_id) = payload
.request_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
return Ok(None);
};
let Some(provider_id) = payload
.provider_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
return Ok(None);
};
let Some(endpoint_id) = payload
.endpoint_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
return Ok(None);
};
let Some(key_id) = payload
.key_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
return Ok(None);
};
let auth_header = payload
.auth_header
.clone()
.filter(|value| !value.trim().is_empty());
let auth_value = payload
.auth_value
.clone()
.filter(|value| !value.trim().is_empty());
let auth_header = take_non_empty_string(&mut payload.auth_header);
let auth_value = take_non_empty_string(&mut payload.auth_value);
if auth_header.is_some() != auth_value.is_some() {
return Ok(None);
}
let Some(provider_api_format) = payload
.provider_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
return Ok(None);
};
let Some(client_api_format) = payload
.client_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
return Ok(None);
};
let (url, url_source) = if let Some(upstream_url) = payload
.upstream_url
.clone()
.filter(|value| !value.trim().is_empty())
let (url, url_source) = if let Some(upstream_url) =
take_non_empty_string(&mut payload.upstream_url)
{
(upstream_url, "upstream_url")
} else {
let Some(upstream_base_url) = payload
.upstream_base_url
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(upstream_base_url) = take_non_empty_string(&mut payload.upstream_base_url) else {
return Ok(None);
};
(
@@ -287,45 +204,45 @@ pub(crate) fn build_openai_cli_sync_plan_from_decision(
"upstream_base_url",
)
};
let Some(provider_request_body_value) = payload.provider_request_body.clone() else {
let Some(provider_request_body_value) = payload.provider_request_body.take() else {
return Ok(None);
};
let mut provider_request_headers = payload.provider_request_headers.clone();
let mut provider_request_headers = std::mem::take(&mut payload.provider_request_headers);
if let (Some(auth_header), Some(auth_value)) = (auth_header.as_deref(), auth_value.as_deref()) {
ensure_upstream_auth_header(&mut provider_request_headers, auth_header, auth_value);
}
if payload.upstream_is_stream && !provider_request_headers.contains_key("accept") {
provider_request_headers.insert("accept".to_string(), "text/event-stream".to_string());
}
let content_type = payload
.content_type
.take()
.or_else(|| Some("application/json".to_string()));
let report_context = augment_sync_report_context(
payload.report_context,
payload.report_context.take(),
&provider_request_headers,
&provider_request_body_value,
)?;
let plan = ExecutionPlan {
request_id,
candidate_id: payload.candidate_id.clone(),
provider_name: payload.provider_name.clone(),
candidate_id: payload.candidate_id.take(),
provider_name: payload.provider_name.take(),
provider_id,
endpoint_id,
key_id,
method: "POST".to_string(),
url,
headers: std::mem::take(&mut provider_request_headers),
content_type: payload
.content_type
.clone()
.or_else(|| Some("application/json".to_string())),
content_type,
content_encoding: None,
body: RequestBody::from_json(provider_request_body_value.clone()),
body: RequestBody::from_json(provider_request_body_value),
stream: payload.upstream_is_stream,
client_api_format,
provider_api_format,
model_name: payload.model_name.clone(),
proxy: payload.proxy.clone(),
tls_profile: payload.tls_profile.clone(),
timeouts: payload.timeouts.clone(),
model_name: payload.model_name.take(),
proxy: payload.proxy.take(),
tls_profile: payload.tls_profile.take(),
timeouts: payload.timeouts.take(),
};
debug!(

View File

@@ -1,6 +1,9 @@
use aether_contracts::{ExecutionPlan, RequestBody};
use super::{augment_sync_report_context, LocalStreamPlanAndReport, LocalSyncPlanAndReport};
use super::{
augment_sync_report_context, take_non_empty_string, LocalStreamPlanAndReport,
LocalSyncPlanAndReport,
};
use crate::ai_pipeline::contracts::generic_decision_missing_exact_provider_request;
use crate::ai_pipeline::provider_adaptation_requires_eventstream_accept;
use crate::ai_pipeline::transport::ensure_upstream_auth_header;
@@ -11,74 +14,40 @@ pub(crate) fn build_standard_sync_plan_from_decision(
_body_json: &serde_json::Value,
payload: GatewayControlSyncDecisionResponse,
) -> Result<Option<LocalSyncPlanAndReport>, GatewayError> {
let mut payload = payload;
if generic_decision_missing_exact_provider_request(&payload) {
return Ok(None);
}
let Some(request_id) = payload
.request_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
return Ok(None);
};
let Some(provider_id) = payload
.provider_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
return Ok(None);
};
let Some(endpoint_id) = payload
.endpoint_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
return Ok(None);
};
let Some(key_id) = payload
.key_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
return Ok(None);
};
let Some(url) = payload
.upstream_url
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(url) = take_non_empty_string(&mut payload.upstream_url) else {
return Ok(None);
};
let auth_header = payload
.auth_header
.clone()
.filter(|value| !value.trim().is_empty());
let auth_value = payload
.auth_value
.clone()
.filter(|value| !value.trim().is_empty());
let auth_header = take_non_empty_string(&mut payload.auth_header);
let auth_value = take_non_empty_string(&mut payload.auth_value);
if auth_header.is_some() != auth_value.is_some() {
return Ok(None);
}
let Some(provider_api_format) = payload
.provider_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
return Ok(None);
};
let Some(client_api_format) = payload
.client_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
return Ok(None);
};
let Some(provider_request_body_value) = payload.provider_request_body.clone() else {
let Some(provider_request_body_value) = payload.provider_request_body.take() else {
return Ok(None);
};
let mut provider_request_headers = payload.provider_request_headers.clone();
let mut provider_request_headers = std::mem::take(&mut payload.provider_request_headers);
if let (Some(auth_header), Some(auth_value)) = (auth_header.as_deref(), auth_value.as_deref()) {
ensure_upstream_auth_header(&mut provider_request_headers, auth_header, auth_value);
}
@@ -87,37 +56,37 @@ pub(crate) fn build_standard_sync_plan_from_decision(
.entry("accept".to_string())
.or_insert_with(|| "text/event-stream".to_string());
}
let content_type = payload
.content_type
.take()
.or_else(|| Some("application/json".to_string()));
let report_context = augment_sync_report_context(
payload.report_context.take(),
&provider_request_headers,
&provider_request_body_value,
)?;
let plan = ExecutionPlan {
request_id,
candidate_id: payload.candidate_id.clone(),
provider_name: payload.provider_name.clone(),
candidate_id: payload.candidate_id.take(),
provider_name: payload.provider_name.take(),
provider_id,
endpoint_id,
key_id,
method: "POST".to_string(),
url,
headers: std::mem::take(&mut provider_request_headers),
content_type: payload
.content_type
.clone()
.or_else(|| Some("application/json".to_string())),
content_type,
content_encoding: None,
body: RequestBody::from_json(provider_request_body_value.clone()),
body: RequestBody::from_json(provider_request_body_value),
stream: payload.upstream_is_stream,
client_api_format,
provider_api_format,
model_name: payload.model_name.clone(),
proxy: payload.proxy.clone(),
tls_profile: payload.tls_profile.clone(),
timeouts: payload.timeouts.clone(),
model_name: payload.model_name.take(),
proxy: payload.proxy.take(),
tls_profile: payload.tls_profile.take(),
timeouts: payload.timeouts.take(),
};
let report_context = augment_sync_report_context(
payload.report_context,
&plan.headers,
&provider_request_body_value,
)?;
Ok(Some(LocalSyncPlanAndReport {
plan,
report_kind: payload.report_kind,
@@ -131,70 +100,37 @@ pub(crate) fn build_standard_stream_plan_from_decision(
payload: GatewayControlSyncDecisionResponse,
_inject_stream_flag: bool,
) -> Result<Option<LocalStreamPlanAndReport>, GatewayError> {
let mut payload = payload;
if generic_decision_missing_exact_provider_request(&payload) {
return Ok(None);
}
let Some(request_id) = payload
.request_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
return Ok(None);
};
let Some(provider_id) = payload
.provider_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
return Ok(None);
};
let Some(endpoint_id) = payload
.endpoint_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
return Ok(None);
};
let Some(key_id) = payload
.key_id
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
return Ok(None);
};
let Some(url) = payload
.upstream_url
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(url) = take_non_empty_string(&mut payload.upstream_url) else {
return Ok(None);
};
let auth_header = payload
.auth_header
.clone()
.filter(|value| !value.trim().is_empty());
let auth_value = payload
.auth_value
.clone()
.filter(|value| !value.trim().is_empty());
let auth_header = take_non_empty_string(&mut payload.auth_header);
let auth_value = take_non_empty_string(&mut payload.auth_value);
if auth_header.is_some() != auth_value.is_some() {
return Ok(None);
}
let Some(provider_api_format) = payload
.provider_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
return Ok(None);
};
let Some(client_api_format) = payload
.client_api_format
.clone()
.filter(|value| !value.trim().is_empty())
else {
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
return Ok(None);
};
let Some(provider_request_body_value) = payload.provider_request_body.clone() else {
let Some(provider_request_body_value) = payload.provider_request_body.take() else {
return Ok(None);
};
@@ -203,7 +139,7 @@ pub(crate) fn build_standard_stream_plan_from_decision(
.as_ref()
.and_then(|context| context.get("envelope_name"))
.and_then(serde_json::Value::as_str);
let mut provider_request_headers = payload.provider_request_headers.clone();
let mut provider_request_headers = std::mem::take(&mut payload.provider_request_headers);
if let (Some(auth_header), Some(auth_value)) = (auth_header.as_deref(), auth_value.as_deref()) {
ensure_upstream_auth_header(&mut provider_request_headers, auth_header, auth_value);
}
@@ -215,37 +151,37 @@ pub(crate) fn build_standard_stream_plan_from_decision(
} else {
provider_request_headers.insert("accept".to_string(), "text/event-stream".to_string());
}
let content_type = payload
.content_type
.take()
.or_else(|| Some("application/json".to_string()));
let report_context = augment_sync_report_context(
payload.report_context.take(),
&provider_request_headers,
&provider_request_body_value,
)?;
let plan = ExecutionPlan {
request_id,
candidate_id: payload.candidate_id.clone(),
provider_name: payload.provider_name.clone(),
candidate_id: payload.candidate_id.take(),
provider_name: payload.provider_name.take(),
provider_id,
endpoint_id,
key_id,
method: "POST".to_string(),
url,
headers: std::mem::take(&mut provider_request_headers),
content_type: payload
.content_type
.clone()
.or_else(|| Some("application/json".to_string())),
content_type,
content_encoding: None,
body: RequestBody::from_json(provider_request_body_value.clone()),
body: RequestBody::from_json(provider_request_body_value),
stream: true,
client_api_format,
provider_api_format,
model_name: payload.model_name.clone(),
proxy: payload.proxy.clone(),
tls_profile: payload.tls_profile.clone(),
timeouts: payload.timeouts.clone(),
model_name: payload.model_name.take(),
proxy: payload.proxy.take(),
tls_profile: payload.tls_profile.take(),
timeouts: payload.timeouts.take(),
};
let report_context = augment_sync_report_context(
payload.report_context,
&plan.headers,
&provider_request_body_value,
)?;
Ok(Some(LocalStreamPlanAndReport {
plan,
report_kind: payload.report_kind,

View File

@@ -39,9 +39,9 @@ pub(crate) use aether_ai_pipeline::api::{
normalize_openai_cli_request_to_openai_chat_request, normalize_provider_private_report_context,
normalize_provider_private_response_value, normalize_standard_request_to_openai_chat_request,
parse_direct_request_body, parse_openai_stop_sequences, parse_openai_tool_result_content,
prepare_local_success_response_parts, provider_adaptation_allows_sync_finalize_envelope,
provider_adaptation_anchor_api_format, provider_adaptation_descriptor_for_envelope,
provider_adaptation_descriptor_for_provider_type,
prepare_local_success_response_parts, prepare_local_success_response_parts_owned,
provider_adaptation_allows_sync_finalize_envelope, provider_adaptation_anchor_api_format,
provider_adaptation_descriptor_for_envelope, provider_adaptation_descriptor_for_provider_type,
provider_adaptation_requires_eventstream_accept,
provider_adaptation_should_unwrap_stream_envelope,
provider_private_response_allows_sync_finalize, request_candidate_api_format_preference,

View File

@@ -84,6 +84,27 @@ pub(crate) fn build_client_response_from_parts(
trace_id: &str,
control_decision: Option<&GatewayControlDecision>,
) -> Result<Response<Body>, GatewayError> {
build_client_response_from_parts_with_mutator(
status_code,
upstream_headers,
body,
trace_id,
control_decision,
|_| Ok(()),
)
}
pub(crate) fn build_client_response_from_parts_with_mutator<F>(
status_code: u16,
upstream_headers: &BTreeMap<String, String>,
body: Body,
trace_id: &str,
control_decision: Option<&GatewayControlDecision>,
mutate_headers: F,
) -> Result<Response<Body>, GatewayError>
where
F: FnOnce(&mut http::HeaderMap) -> Result<(), GatewayError>,
{
let mut response = Response::builder()
.status(status_code)
.body(body)
@@ -99,6 +120,7 @@ pub(crate) fn build_client_response_from_parts(
HeaderValue::from_str(value).map_err(|err| GatewayError::Internal(err.to_string()))?;
response.headers_mut().insert(header_name, header_value);
}
mutate_headers(response.headers_mut())?;
apply_streaming_response_headers(response.headers_mut());
insert_header_if_missing(response.headers_mut(), TRACE_ID_HEADER, trace_id)?;
insert_header_if_missing(response.headers_mut(), GATEWAY_HEADER, "rust-phase3b")?;

View File

@@ -42,7 +42,8 @@ pub(crate) use sync::{
execute_execution_runtime_sync, maybe_build_local_sync_finalize_response,
maybe_build_local_video_error_response, maybe_build_local_video_success_outcome,
resolve_local_sync_error_background_report_kind,
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessOutcome,
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessBuild,
LocalVideoSyncSuccessOutcome,
};
pub(crate) use transport::{
execute_sync_plan as execute_execution_runtime_sync_plan, DirectSyncExecutionRuntime,

View File

@@ -221,7 +221,7 @@ async fn execute_sync(
let plan = parse_request_json::<ExecutionPlan>(request).await?;
let result = state
.execution_runtime
.execute_sync(plan)
.execute_sync(&plan)
.await
.map_err(|err| ExecutionRuntimeAppError(ExecutionRuntimeServerError::Transport(err)))?;
Ok(maybe_hold_axum_response_permit(
@@ -238,7 +238,7 @@ async fn execute_stream(
let plan = parse_request_json::<ExecutionPlan>(request).await?;
let execution = state
.execution_runtime
.execute_stream(plan)
.execute_stream(&plan)
.await
.map_err(|err| ExecutionRuntimeAppError(ExecutionRuntimeServerError::Transport(err)))?;

View File

@@ -74,6 +74,7 @@ use crate::orchestration::{
};
use crate::request_candidate_runtime::{
ensure_execution_request_candidate_slot, record_local_request_candidate_status,
record_local_request_candidate_status_snapshot, snapshot_local_request_candidate_status,
};
use crate::usage::submit_stream_report;
use crate::usage::{GatewayStreamReportRequest, GatewaySyncReportRequest};
@@ -89,7 +90,30 @@ fn record_sync_terminal_usage(
let payload_seed = build_sync_terminal_usage_payload_seed(payload);
state
.usage_runtime
.record_sync_terminal(state.data.as_ref(), &context_seed, &payload_seed);
.record_sync_terminal(state.data.as_ref(), context_seed, payload_seed);
}
fn build_stream_sync_payload(
trace_id: &str,
report_kind: String,
report_context: Option<Value>,
status_code: u16,
headers: BTreeMap<String, String>,
body_json: Option<Value>,
body_base64: Option<String>,
telemetry: Option<ExecutionTelemetry>,
) -> GatewaySyncReportRequest {
GatewaySyncReportRequest {
trace_id: trace_id.to_string(),
report_kind,
report_context,
status_code,
headers,
body_json,
client_body_json: None,
body_base64,
telemetry,
}
}
fn record_stream_terminal_usage(
@@ -103,12 +127,61 @@ fn record_stream_terminal_usage(
let payload_seed = build_stream_terminal_usage_payload_seed(payload);
state.usage_runtime.record_stream_terminal(
state.data.as_ref(),
&context_seed,
&payload_seed,
context_seed,
payload_seed,
cancelled,
);
}
fn build_stream_body_capture(
body: &[u8],
truncated: bool,
) -> (Option<String>, Option<UsageBodyCaptureState>) {
let body_base64 =
(!body.is_empty()).then(|| base64::engine::general_purpose::STANDARD.encode(body));
let body_state = Some(if truncated {
UsageBodyCaptureState::Truncated
} else if body.is_empty() {
UsageBodyCaptureState::None
} else {
UsageBodyCaptureState::Inline
});
(body_base64, body_state)
}
#[allow(clippy::too_many_arguments)] // stream report payload assembly mirrors runtime state
fn build_stream_usage_payload(
trace_id: String,
report_kind: String,
report_context: Option<Value>,
status_code: u16,
headers: BTreeMap<String, String>,
provider_body: &[u8],
provider_body_truncated: bool,
client_body: &[u8],
client_body_truncated: bool,
terminal_summary: Option<ExecutionStreamTerminalSummary>,
telemetry: Option<ExecutionTelemetry>,
) -> GatewayStreamReportRequest {
let (provider_body_base64, provider_body_state) =
build_stream_body_capture(provider_body, provider_body_truncated);
let (client_body_base64, client_body_state) =
build_stream_body_capture(client_body, client_body_truncated);
GatewayStreamReportRequest {
trace_id,
report_kind,
report_context,
status_code,
headers,
provider_body_base64,
provider_body_state,
client_body_base64,
client_body_state,
terminal_summary,
telemetry,
}
}
fn append_stream_capture_bytes(
buffer: &mut Vec<u8>,
chunk: &[u8],
@@ -138,9 +211,7 @@ async fn execute_in_process_stream(
return Ok(execution);
}
DirectSyncExecutionRuntime::new()
.execute_stream(plan.clone())
.await
DirectSyncExecutionRuntime::new().execute_stream(plan).await
}
#[allow(clippy::too_many_arguments)] // internal function, grouping would add unnecessary indirection
@@ -155,19 +226,18 @@ pub(crate) async fn execute_execution_runtime_stream(
) -> Result<Option<Response<Body>>, GatewayError> {
ensure_execution_request_candidate_slot(state, &mut plan, &mut report_context).await;
let lifecycle_seed = build_lifecycle_usage_seed(&plan, report_context.as_ref());
let request_candidate_status_snapshot =
snapshot_local_request_candidate_status(&plan, report_context.as_ref());
state
.usage_runtime
.record_pending(state.data.as_ref(), &lifecycle_seed);
.record_pending(state.data.as_ref(), lifecycle_seed.clone());
let candidate_started_unix_secs = current_request_candidate_unix_ms();
{
if let Some(snapshot) = request_candidate_status_snapshot.clone() {
let state_bg = state.clone();
let plan_bg = plan.clone();
let report_context_bg = report_context.clone();
tokio::spawn(async move {
record_local_request_candidate_status(
record_local_request_candidate_status_snapshot(
&state_bg,
&plan_bg,
report_context_bg.as_ref(),
&snapshot,
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Pending,
status_code: None,
@@ -412,8 +482,9 @@ fn response_headers_indicate_sse(headers: &BTreeMap<String, String>) -> bool {
}
fn encode_terminal_sse_error_event(failure: &StreamFailureReport) -> Result<Bytes, std::io::Error> {
let payload =
serde_json::to_string(&failure.body_json).map_err(|err| IoError::other(err.to_string()))?;
let payload = failure
.to_json_string()
.map_err(|err| IoError::other(err.to_string()))?;
let mut event = String::from("event: aether.error\n");
for line in payload.lines() {
event.push_str("data: ");
@@ -527,6 +598,8 @@ async fn execute_stream_from_frame_stream(
let provider_name = plan.provider_name.as_deref().unwrap_or("-");
let model_name = plan.model_name.as_deref().unwrap_or("-");
let lifecycle_seed = build_lifecycle_usage_seed(&plan, report_context.as_ref());
let request_candidate_status_snapshot =
snapshot_local_request_candidate_status(&plan, report_context.as_ref());
let candidate_index = parse_request_candidate_report_context(report_context.as_ref())
.and_then(|context| context.candidate_index)
.map(|value| value.to_string())
@@ -756,27 +829,26 @@ async fn execute_stream_from_frame_stream(
return Ok(None);
}
let usage_report_kind = stream_error_finalize_kind
.clone()
.or_else(|| report_kind.clone())
.unwrap_or_default();
let usage_payload = GatewaySyncReportRequest {
trace_id: trace_id.to_string(),
report_kind: usage_report_kind,
report_context: report_context.clone(),
let payload = build_stream_sync_payload(
trace_id,
stream_error_finalize_kind
.as_deref()
.or(report_kind.as_deref())
.unwrap_or_default()
.to_string(),
report_context,
status_code,
headers: headers.clone(),
body_json: body_json.clone(),
client_body_json: None,
body_base64: body_base64.clone(),
telemetry: None,
};
record_sync_terminal_usage(state, &plan, report_context.as_ref(), &usage_payload);
headers,
body_json,
body_base64,
None,
);
record_sync_terminal_usage(state, &plan, payload.report_context.as_ref(), &payload);
let terminal_unix_secs = current_request_candidate_unix_ms();
record_local_request_candidate_status(
state,
&plan,
report_context.as_ref(),
payload.report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed,
status_code: Some(status_code),
@@ -790,18 +862,7 @@ async fn execute_stream_from_frame_stream(
},
)
.await;
if let Some(report_kind) = stream_error_finalize_kind {
let payload = GatewaySyncReportRequest {
trace_id: trace_id.to_string(),
report_kind,
report_context,
status_code,
headers: headers.clone(),
body_json,
client_body_json: None,
body_base64,
telemetry: None,
};
if stream_error_finalize_kind.is_some() {
let response =
submit_local_core_error_or_sync_finalize(state, trace_id, decision, payload)
.await?;
@@ -817,7 +878,7 @@ async fn execute_stream_from_frame_stream(
decision,
plan_kind,
status_code,
headers,
payload.headers,
error_body,
)?,
Some(request_id),
@@ -891,12 +952,12 @@ async fn execute_stream_from_frame_stream(
trace_id,
decision,
&plan,
report_context.clone(),
report_context,
request_id,
candidate_id,
report_kind,
&headers,
prefetched_telemetry.clone(),
headers,
prefetched_telemetry,
&provider_prefetched_body,
failure,
)
@@ -924,12 +985,12 @@ async fn execute_stream_from_frame_stream(
trace_id,
decision,
&plan,
report_context.clone(),
report_context,
request_id,
candidate_id,
report_kind,
&headers,
prefetched_telemetry.clone(),
headers,
prefetched_telemetry,
&prefetched_body,
failure,
)
@@ -964,21 +1025,20 @@ async fn execute_stream_from_frame_stream(
provider_prefetched_body_bytes = provider_prefetched_body.len(),
"gateway detected embedded error while prefetching execution runtime stream"
);
let payload = GatewaySyncReportRequest {
trace_id: trace_id.to_string(),
report_kind: report_kind.clone(),
report_context: report_context.clone(),
let payload = build_stream_sync_payload(
trace_id,
report_kind.clone(),
report_context,
status_code,
headers: headers.clone(),
body_json: Some(body_json),
client_body_json: None,
body_base64: None,
telemetry: prefetched_telemetry.clone(),
};
headers,
Some(body_json),
None,
prefetched_telemetry,
);
record_sync_terminal_usage(
state,
&plan,
report_context.as_ref(),
payload.report_context.as_ref(),
&payload,
);
let response = submit_local_core_error_or_sync_finalize(
@@ -1013,12 +1073,12 @@ async fn execute_stream_from_frame_stream(
trace_id,
decision,
&plan,
report_context.clone(),
report_context,
request_id,
candidate_id,
report_kind,
&headers,
prefetched_telemetry.clone(),
headers,
prefetched_telemetry,
&provider_prefetched_body,
failure,
)
@@ -1044,12 +1104,12 @@ async fn execute_stream_from_frame_stream(
trace_id,
decision,
&plan,
report_context.clone(),
report_context,
request_id,
candidate_id,
report_kind,
&headers,
prefetched_telemetry.clone(),
headers,
prefetched_telemetry,
&provider_prefetched_body,
failure,
)
@@ -1093,12 +1153,12 @@ async fn execute_stream_from_frame_stream(
trace_id,
decision,
&plan,
report_context.clone(),
report_context,
request_id,
candidate_id,
report_kind,
&headers,
prefetched_telemetry.clone(),
headers,
prefetched_telemetry,
&provider_prefetched_body,
build_stream_failure_from_execution_error(&error),
)
@@ -1108,6 +1168,8 @@ async fn execute_stream_from_frame_stream(
}
}
}
drop(private_stream_normalizer);
drop(local_stream_rewriter);
state.usage_runtime.record_stream_started(
state.data.as_ref(),
@@ -1115,18 +1177,15 @@ async fn execute_stream_from_frame_stream(
status_code,
prefetched_telemetry.as_ref(),
);
{
if let Some(snapshot) = request_candidate_status_snapshot {
let state_bg = state.clone();
let plan_bg = plan.clone();
let report_context_bg = report_context.clone();
let latency_ms = prefetched_telemetry
.as_ref()
.and_then(|telemetry| telemetry.elapsed_ms);
tokio::spawn(async move {
record_local_request_candidate_status(
record_local_request_candidate_status_snapshot(
&state_bg,
&plan_bg,
report_context_bg.as_ref(),
&snapshot,
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Streaming,
status_code: Some(status_code),
@@ -1141,24 +1200,27 @@ async fn execute_stream_from_frame_stream(
});
}
let request_id = request_id.to_string();
let candidate_id = candidate_id.map(ToOwned::to_owned);
let (tx, mut rx) = mpsc::channel::<Result<Bytes, IoError>>(16);
let state_for_report = state.clone();
let plan_for_report = plan.clone();
let plan_for_report = plan;
let trace_id_owned = trace_id.to_string();
let headers_for_report = headers.clone();
let report_kind_owned = report_kind.clone();
let report_context_owned = report_context.clone();
let lifecycle_seed_for_report = lifecycle_seed.clone();
let report_kind_owned = report_kind;
let report_context_owned = report_context;
let normalized_stream_report_context_owned = normalized_stream_report_context;
let lifecycle_seed_for_report = lifecycle_seed;
let provider_prefetched_body_for_report = provider_prefetched_body;
let prefetched_body_for_report = prefetched_body;
let prefetched_chunks_for_body = prefetched_chunks;
let initial_telemetry = prefetched_telemetry.clone();
let initial_telemetry = prefetched_telemetry;
let initial_reached_eof = reached_eof;
let direct_stream_finalize_kind_owned = direct_stream_finalize_kind.clone();
let direct_stream_finalize_kind_owned = direct_stream_finalize_kind;
let candidate_started_unix_secs_for_report = candidate_started_unix_secs;
let request_id_for_report = request_id.to_string();
let request_id_for_report_log = short_request_id(request_id);
let candidate_id_for_report = candidate_id.map(ToOwned::to_owned);
let request_id_for_report = request_id.clone();
let request_id_for_report_log = short_request_id(&request_id);
let candidate_id_for_report = candidate_id.clone();
let emit_passthrough_sse_terminal_error =
skip_direct_finalize_prefetch && response_headers_indicate_sse(&headers);
let body_capture_policy = match UsageRuntimeAccess::body_capture_policy(state.data.as_ref())
@@ -1194,6 +1256,10 @@ async fn execute_stream_from_frame_stream(
let mut buffered_body = Vec::new();
let mut provider_body_truncated = false;
let mut client_body_truncated = false;
let mut private_stream_normalizer =
maybe_build_provider_private_stream_normalizer(report_context_owned.as_ref());
let mut local_stream_rewriter =
maybe_build_stream_response_rewriter(normalized_stream_report_context_owned.as_ref());
append_stream_capture_bytes(
&mut provider_buffered_body,
&provider_prefetched_body_for_report,
@@ -1206,13 +1272,68 @@ async fn execute_stream_from_frame_stream(
max_stream_body_buffer_bytes,
&mut client_body_truncated,
);
let mut telemetry: Option<ExecutionTelemetry> = initial_telemetry.clone();
let mut usage_stream_telemetry: Option<ExecutionTelemetry> = initial_telemetry;
let mut usage_stream_telemetry: Option<ExecutionTelemetry> = initial_telemetry.clone();
let mut telemetry: Option<ExecutionTelemetry> = initial_telemetry;
let reached_eof = initial_reached_eof;
let mut downstream_dropped = false;
let mut terminal_failure: Option<StreamFailureReport> = None;
if !provider_prefetched_body_for_report.is_empty() {
let normalized_prefetched_chunk = if let Some(normalizer) =
private_stream_normalizer.as_mut()
{
match normalizer.push_chunk(&provider_prefetched_body_for_report) {
Ok(normalized_chunk) => Some(normalized_chunk),
Err(err) => {
warn!(
event_name = "stream_execution_prefetch_normalize_restore_failed",
log_type = "ops",
trace_id = %trace_id_owned,
request_id = %request_id_for_report_log,
candidate_id = ?candidate_id_for_report.as_deref(),
error = ?err,
"gateway failed to restore private stream normalization state after prefetch"
);
terminal_failure = Some(build_stream_failure_report(
"execution_runtime_stream_rewrite_error",
format!(
"failed to restore private stream normalization state after prefetch: {err:?}"
),
502,
));
None
}
}
} else {
None
};
let replay_chunk = normalized_prefetched_chunk
.as_deref()
.unwrap_or(provider_prefetched_body_for_report.as_slice());
if terminal_failure.is_none() {
if let Some(rewriter) = local_stream_rewriter.as_mut() {
if let Err(err) = rewriter.push_chunk(replay_chunk) {
warn!(
event_name = "stream_execution_prefetch_rewrite_restore_failed",
log_type = "ops",
trace_id = %trace_id_owned,
request_id = %request_id_for_report_log,
candidate_id = ?candidate_id_for_report.as_deref(),
error = ?err,
"gateway failed to restore local stream rewrite state after prefetch"
);
terminal_failure = Some(build_stream_failure_report(
"execution_runtime_stream_rewrite_error",
format!(
"failed to restore local stream rewrite state after prefetch: {err:?}"
),
502,
));
}
}
}
}
if !reached_eof {
if terminal_failure.is_none() && !reached_eof {
loop {
let next_frame = match next_stream_frame(&mut buffered_frames, &mut lines).await {
Ok(frame) => frame,
@@ -1355,7 +1476,6 @@ async fn execute_stream_from_frame_stream(
usage_stream_telemetry.as_ref(),
&frame_telemetry,
);
telemetry = Some(frame_telemetry.clone());
if should_refresh_stream_usage {
state_for_report.usage_runtime.record_stream_started(
state_for_report.data.as_ref(),
@@ -1363,8 +1483,9 @@ async fn execute_stream_from_frame_stream(
status_code,
Some(&frame_telemetry),
);
usage_stream_telemetry = Some(frame_telemetry);
usage_stream_telemetry = Some(frame_telemetry.clone());
}
telemetry = Some(frame_telemetry);
}
StreamFramePayload::Eof { summary } => {
stream_terminal_summary = summary;
@@ -1564,50 +1685,39 @@ async fn execute_stream_from_frame_stream(
trace_id = %trace_id_owned,
"gateway skipped stream report because downstream disconnected before completion"
);
let usage_payload = build_stream_usage_payload(
trace_id_owned,
report_kind_owned.unwrap_or_default(),
report_context_owned,
499,
headers_for_report,
&provider_buffered_body,
provider_body_truncated,
&buffered_body,
client_body_truncated,
stream_terminal_summary,
telemetry,
);
record_stream_terminal_usage(
&state_for_report,
&plan_for_report,
report_context_owned.as_ref(),
&GatewayStreamReportRequest {
trace_id: trace_id_owned.clone(),
report_kind: report_kind_owned.clone().unwrap_or_default(),
report_context: report_context_owned.clone(),
status_code: 499,
headers: headers_for_report.clone(),
provider_body_base64: (!provider_buffered_body.is_empty()).then(|| {
base64::engine::general_purpose::STANDARD.encode(&provider_buffered_body)
}),
provider_body_state: Some(if provider_body_truncated {
UsageBodyCaptureState::Truncated
} else if provider_buffered_body.is_empty() {
UsageBodyCaptureState::None
} else {
UsageBodyCaptureState::Inline
}),
client_body_base64: (!buffered_body.is_empty())
.then(|| base64::engine::general_purpose::STANDARD.encode(&buffered_body)),
client_body_state: Some(if client_body_truncated {
UsageBodyCaptureState::Truncated
} else if buffered_body.is_empty() {
UsageBodyCaptureState::None
} else {
UsageBodyCaptureState::Inline
}),
terminal_summary: stream_terminal_summary.clone(),
telemetry: telemetry.clone(),
},
usage_payload.report_context.as_ref(),
&usage_payload,
true,
);
record_local_request_candidate_status(
&state_for_report,
&plan_for_report,
report_context_owned.as_ref(),
usage_payload.report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Cancelled,
status_code: Some(499),
error_type: Some("downstream_disconnect".to_string()),
error_message: Some("client disconnected before stream completion".to_string()),
latency_ms: telemetry.as_ref().and_then(|value| value.elapsed_ms),
latency_ms: usage_payload
.telemetry
.as_ref()
.and_then(|value| value.elapsed_ms),
started_at_unix_ms: Some(candidate_started_unix_secs_for_report),
finished_at_unix_ms: Some(current_request_candidate_unix_ms()),
},
@@ -1622,9 +1732,9 @@ async fn execute_stream_from_frame_stream(
&trace_id_owned,
&plan_for_report,
direct_stream_finalize_kind_owned.as_deref(),
report_context_owned.as_ref(),
&headers_for_report,
telemetry.clone(),
report_context_owned,
headers_for_report,
telemetry,
&provider_buffered_body,
candidate_started_unix_secs_for_report,
failure,
@@ -1633,38 +1743,25 @@ async fn execute_stream_from_frame_stream(
return;
}
let usage_payload = GatewayStreamReportRequest {
trace_id: trace_id_owned.clone(),
report_kind: report_kind_owned.clone().unwrap_or_default(),
report_context: report_context_owned.clone(),
let should_submit_report = report_kind_owned.is_some();
let usage_payload = build_stream_usage_payload(
trace_id_owned.clone(),
report_kind_owned.unwrap_or_default(),
report_context_owned,
status_code,
headers: headers_for_report.clone(),
provider_body_base64: (!provider_buffered_body.is_empty())
.then(|| base64::engine::general_purpose::STANDARD.encode(&provider_buffered_body)),
provider_body_state: Some(if provider_body_truncated {
UsageBodyCaptureState::Truncated
} else if provider_buffered_body.is_empty() {
UsageBodyCaptureState::None
} else {
UsageBodyCaptureState::Inline
}),
client_body_base64: (!buffered_body.is_empty())
.then(|| base64::engine::general_purpose::STANDARD.encode(&buffered_body)),
client_body_state: Some(if client_body_truncated {
UsageBodyCaptureState::Truncated
} else if buffered_body.is_empty() {
UsageBodyCaptureState::None
} else {
UsageBodyCaptureState::Inline
}),
terminal_summary: stream_terminal_summary,
telemetry: telemetry.clone(),
};
headers_for_report,
&provider_buffered_body,
provider_body_truncated,
&buffered_body,
client_body_truncated,
stream_terminal_summary,
telemetry,
);
apply_local_execution_effect(
&state_for_report,
LocalExecutionEffectContext {
plan: &plan_for_report,
report_context: report_context_owned.as_ref(),
report_context: usage_payload.report_context.as_ref(),
},
LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect),
)
@@ -1673,7 +1770,7 @@ async fn execute_stream_from_frame_stream(
&state_for_report,
LocalExecutionEffectContext {
plan: &plan_for_report,
report_context: report_context_owned.as_ref(),
report_context: usage_payload.report_context.as_ref(),
},
LocalExecutionEffect::AdaptiveSuccess(LocalAdaptiveSuccessEffect),
)
@@ -1682,7 +1779,7 @@ async fn execute_stream_from_frame_stream(
&state_for_report,
LocalExecutionEffectContext {
plan: &plan_for_report,
report_context: report_context_owned.as_ref(),
report_context: usage_payload.report_context.as_ref(),
},
LocalExecutionEffect::PoolSuccessStream {
payload: &usage_payload,
@@ -1692,31 +1789,31 @@ async fn execute_stream_from_frame_stream(
record_stream_terminal_usage(
&state_for_report,
&plan_for_report,
report_context_owned.as_ref(),
usage_payload.report_context.as_ref(),
&usage_payload,
false,
);
record_local_request_candidate_status(
&state_for_report,
&plan_for_report,
report_context_owned.as_ref(),
usage_payload.report_context.as_ref(),
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Success,
status_code: Some(status_code),
error_type: None,
error_message: None,
latency_ms: telemetry.as_ref().and_then(|value| value.elapsed_ms),
latency_ms: usage_payload
.telemetry
.as_ref()
.and_then(|value| value.elapsed_ms),
started_at_unix_ms: Some(candidate_started_unix_secs_for_report),
finished_at_unix_ms: Some(current_request_candidate_unix_ms()),
},
)
.await;
if let Some(report_kind) = report_kind_owned {
let mut report = usage_payload;
report.report_kind = report_kind;
if let Err(err) = submit_stream_report(&state_for_report, &trace_id_owned, report).await
{
if should_submit_report {
if let Err(err) = submit_stream_report(&state_for_report, usage_payload).await {
warn!(
event_name = "execution_report_submit_failed",
log_type = "ops",
@@ -1740,12 +1837,10 @@ async fn execute_stream_from_frame_stream(
}
};
headers.insert(
CONTROL_REQUEST_ID_HEADER.to_string(),
request_id.to_string(),
);
headers.insert(CONTROL_REQUEST_ID_HEADER.to_string(), request_id.clone());
if let Some(candidate_id) = candidate_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{

View File

@@ -7,6 +7,7 @@ use aether_usage_runtime::{
use axum::body::Body;
use axum::http::Response;
use base64::Engine as _;
use serde::Serialize;
use serde_json::{Map, Value};
use tracing::warn;
@@ -32,7 +33,51 @@ pub(super) struct StreamFailureReport {
pub(super) status_code: u16,
pub(super) error_type: String,
pub(super) error_message: String,
pub(super) body_json: Value,
extra_error_fields: Map<String, Value>,
}
#[derive(Serialize)]
struct StreamFailureBody<'a> {
error: StreamFailureBodyFields<'a>,
}
#[derive(Serialize)]
struct StreamFailureBodyFields<'a> {
#[serde(rename = "type")]
error_type: &'a str,
message: &'a str,
code: u16,
#[serde(flatten)]
extra_error_fields: &'a Map<String, Value>,
}
impl StreamFailureReport {
fn into_body_json(self) -> Value {
let Self {
status_code,
error_type,
error_message,
mut extra_error_fields,
} = self;
extra_error_fields.insert("type".to_string(), Value::String(error_type));
extra_error_fields.insert("message".to_string(), Value::String(error_message));
extra_error_fields.insert("code".to_string(), Value::from(status_code));
Value::Object(Map::from_iter([(
"error".to_string(),
Value::Object(extra_error_fields),
)]))
}
pub(super) fn to_json_string(&self) -> serde_json::Result<String> {
serde_json::to_string(&StreamFailureBody {
error: StreamFailureBodyFields {
error_type: self.error_type.as_str(),
message: self.error_message.as_str(),
code: self.status_code,
extra_error_fields: &self.extra_error_fields,
},
})
}
}
pub(super) fn build_stream_failure_report(
@@ -44,16 +89,9 @@ pub(super) fn build_stream_failure_report(
let error_message = error_message.into();
StreamFailureReport {
status_code,
body_json: Value::Object(Map::from_iter([(
"error".to_string(),
Value::Object(Map::from_iter([
("type".to_string(), Value::String(error_type.clone())),
("message".to_string(), Value::String(error_message.clone())),
("code".to_string(), Value::from(status_code)),
])),
)])),
error_type,
error_message,
extra_error_fields: Map::new(),
}
}
@@ -65,11 +103,9 @@ pub(super) fn build_stream_failure_from_execution_error(
.ok()
.and_then(|value| value.as_str().map(ToOwned::to_owned))
.unwrap_or_else(|| "internal".to_string());
let error_message = error.message.trim().to_string();
let phase = serde_json::to_value(&error.phase).unwrap_or(Value::Null);
let mut error_object = Map::from_iter([
("type".to_string(), Value::String(error_type.clone())),
("message".to_string(), Value::String(error.message.clone())),
("code".to_string(), Value::from(status_code)),
("phase".to_string(), phase),
("retryable".to_string(), Value::Bool(error.retryable)),
(
@@ -84,11 +120,8 @@ pub(super) fn build_stream_failure_from_execution_error(
StreamFailureReport {
status_code,
error_type,
error_message: error.message.trim().to_string(),
body_json: Value::Object(Map::from_iter([(
"error".to_string(),
Value::Object(error_object),
)])),
error_message,
extra_error_fields: error_object,
}
}
@@ -96,23 +129,23 @@ fn build_stream_failure_sync_payload(
trace_id: &str,
report_kind: String,
report_context: Option<Value>,
headers: &std::collections::BTreeMap<String, String>,
mut headers: std::collections::BTreeMap<String, String>,
telemetry: Option<ExecutionTelemetry>,
provider_buffered_body: &[u8],
failure: &StreamFailureReport,
failure: StreamFailureReport,
) -> GatewaySyncReportRequest {
let mut response_headers = headers.clone();
response_headers.remove("content-encoding");
response_headers.remove("content-length");
response_headers.insert("content-type".to_string(), "application/json".to_string());
let status_code = failure.status_code;
headers.remove("content-encoding");
headers.remove("content-length");
headers.insert("content-type".to_string(), "application/json".to_string());
GatewaySyncReportRequest {
trace_id: trace_id.to_string(),
report_kind,
report_context,
status_code: failure.status_code,
headers: response_headers,
body_json: Some(failure.body_json.clone()),
status_code,
headers,
body_json: Some(failure.into_body_json()),
client_body_json: None,
body_base64: (!provider_buffered_body.is_empty())
.then(|| base64::engine::general_purpose::STANDARD.encode(provider_buffered_body)),
@@ -120,27 +153,40 @@ fn build_stream_failure_sync_payload(
}
}
fn stream_failure_body_field<'a>(
payload: &'a GatewaySyncReportRequest,
field: &str,
) -> Option<&'a str> {
payload
.body_json
.as_ref()
.and_then(|body_json| body_json.get("error"))
.and_then(|value| value.get(field))
.and_then(Value::as_str)
}
async fn record_stream_sync_failure(
state: &AppState,
plan: &ExecutionPlan,
report_context: Option<&Value>,
payload: &GatewaySyncReportRequest,
failure: &StreamFailureReport,
started_at_unix_ms: Option<u64>,
) {
let error_body = serde_json::to_string(&failure.body_json).ok();
let error_type = stream_failure_body_field(payload, "type").unwrap_or("internal");
let error_message = stream_failure_body_field(payload, "message").unwrap_or_default();
let error_body = payload
.body_json
.as_ref()
.and_then(|body_json| serde_json::to_string(body_json).ok());
let failure_analysis = resolve_local_failover_analysis_for_attempt(
state,
plan,
report_context,
failure.status_code,
payload.status_code,
error_body.as_deref(),
)
.await;
if matches!(
failure.error_type.as_str(),
"first_byte_timeout" | "read_timeout"
) {
if matches!(error_type, "first_byte_timeout" | "read_timeout") {
apply_local_execution_effect(
state,
LocalExecutionEffectContext {
@@ -158,7 +204,7 @@ async fn record_stream_sync_failure(
report_context,
},
LocalExecutionEffect::AttemptFailure(LocalAttemptFailureEffect {
status_code: failure.status_code,
status_code: payload.status_code,
classification: failure_analysis.classification,
}),
)
@@ -170,7 +216,7 @@ async fn record_stream_sync_failure(
report_context,
},
LocalExecutionEffect::AdaptiveRateLimit(LocalAdaptiveRateLimitEffect {
status_code: failure.status_code,
status_code: payload.status_code,
classification: failure_analysis.classification,
headers: Some(&payload.headers),
}),
@@ -183,7 +229,7 @@ async fn record_stream_sync_failure(
report_context,
},
LocalExecutionEffect::HealthFailure(LocalHealthFailureEffect {
status_code: failure.status_code,
status_code: payload.status_code,
classification: failure_analysis.classification,
}),
)
@@ -195,7 +241,7 @@ async fn record_stream_sync_failure(
report_context,
},
LocalExecutionEffect::OauthInvalidation(LocalOAuthInvalidationEffect {
status_code: failure.status_code,
status_code: payload.status_code,
response_text: error_body.as_deref(),
}),
)
@@ -207,7 +253,7 @@ async fn record_stream_sync_failure(
report_context,
},
LocalExecutionEffect::PoolError(LocalPoolErrorEffect {
status_code: failure.status_code,
status_code: payload.status_code,
classification: failure_analysis.classification,
headers: &payload.headers,
error_body: error_body.as_deref(),
@@ -218,16 +264,16 @@ async fn record_stream_sync_failure(
let payload_seed = build_sync_terminal_usage_payload_seed(payload);
state
.usage_runtime
.record_sync_terminal(state.data.as_ref(), &context_seed, &payload_seed);
.record_sync_terminal(state.data.as_ref(), context_seed, payload_seed);
let terminal_unix_secs = current_request_candidate_unix_ms();
record_report_request_candidate_status(
state,
report_context,
SchedulerRequestCandidateStatusUpdate {
status: RequestCandidateStatus::Failed,
status_code: Some(failure.status_code),
error_type: Some(failure.error_type.clone()),
error_message: Some(failure.error_message.clone()),
status_code: Some(payload.status_code),
error_type: Some(error_type.to_string()),
error_message: Some(error_message.to_string()),
latency_ms: payload
.telemetry
.as_ref()
@@ -249,7 +295,7 @@ pub(super) async fn handle_prefetch_stream_failure(
request_id: &str,
candidate_id: Option<&str>,
report_kind: &str,
headers: &std::collections::BTreeMap<String, String>,
headers: std::collections::BTreeMap<String, String>,
telemetry: Option<ExecutionTelemetry>,
buffered_body: &[u8],
failure: StreamFailureReport,
@@ -257,21 +303,13 @@ pub(super) async fn handle_prefetch_stream_failure(
let payload = build_stream_failure_sync_payload(
trace_id,
report_kind.to_string(),
report_context.clone(),
report_context,
headers,
telemetry,
buffered_body,
&failure,
failure,
);
record_stream_sync_failure(
state,
plan,
report_context.as_ref(),
&payload,
&failure,
None,
)
.await;
record_stream_sync_failure(state, plan, payload.report_context.as_ref(), &payload, None).await;
let response =
submit_local_core_error_or_sync_finalize(state, trace_id, decision, payload).await?;
@@ -287,8 +325,8 @@ pub(super) async fn submit_midstream_stream_failure(
trace_id: &str,
plan: &ExecutionPlan,
direct_stream_finalize_kind: Option<&str>,
report_context: Option<&Value>,
headers: &std::collections::BTreeMap<String, String>,
report_context: Option<Value>,
headers: std::collections::BTreeMap<String, String>,
telemetry: Option<ExecutionTelemetry>,
buffered_body: &[u8],
started_at_unix_ms: u64,
@@ -303,22 +341,21 @@ pub(super) async fn submit_midstream_stream_failure(
let payload = build_stream_failure_sync_payload(
trace_id,
report_kind,
report_context.cloned(),
report_context,
headers,
telemetry,
buffered_body,
&failure,
failure,
);
record_stream_sync_failure(
state,
plan,
report_context,
payload.report_context.as_ref(),
&payload,
&failure,
Some(started_at_unix_ms),
)
.await;
if let Err(err) = submit_sync_report(state, trace_id, payload).await {
if let Err(err) = submit_sync_report(state, payload).await {
let request_id = short_request_id(plan.request_id.as_str());
warn!(
event_name = "execution_report_submit_failed",

View File

@@ -274,7 +274,7 @@ fn format_error_chain(err: &(dyn std::error::Error + 'static)) -> String {
fn observe_stream_chunk(
observer: &mut StreamingStandardTerminalObserver,
report_context: &Value,
private_stream_normalizer: Option<&mut crate::ai_pipeline::ProviderPrivateStreamNormalizer>,
private_stream_normalizer: Option<&mut crate::ai_pipeline::ProviderPrivateStreamNormalizer<'_>>,
observer_buffered: &mut Vec<u8>,
chunk: &[u8],
) {
@@ -298,7 +298,7 @@ fn observe_stream_chunk(
fn finalize_stream_terminal_summary(
observer: &mut StreamingStandardTerminalObserver,
report_context: &Value,
private_stream_normalizer: Option<&mut crate::ai_pipeline::ProviderPrivateStreamNormalizer>,
private_stream_normalizer: Option<&mut crate::ai_pipeline::ProviderPrivateStreamNormalizer<'_>>,
observer_buffered: &mut Vec<u8>,
) -> Option<ExecutionStreamTerminalSummary> {
if let Some(normalizer) = private_stream_normalizer {
@@ -412,7 +412,7 @@ mod tests {
});
let execution = DirectSyncExecutionRuntime::new()
.execute_stream(ExecutionPlan {
.execute_stream(&ExecutionPlan {
request_id: "req-stream-ttfb-1".into(),
candidate_id: Some("cand-stream-ttfb-1".into()),
provider_name: Some("openai".into()),
@@ -499,7 +499,7 @@ mod tests {
let runtime = DirectSyncExecutionRuntime::new();
let execution = runtime
.execute_stream(ExecutionPlan {
.execute_stream(&ExecutionPlan {
request_id: "req-telemetry-order".to_string(),
candidate_id: Some("cand-telemetry-order".to_string()),
provider_name: Some("OpenAI".to_string()),

View File

@@ -545,7 +545,7 @@ pub(crate) async fn submit_local_core_error_or_sync_finalize(
{
let mut report_payload = payload.clone();
report_payload.report_kind = error_report_kind;
spawn_sync_report(state.clone(), trace_id.to_string(), report_payload);
spawn_sync_report(state.clone(), report_payload);
} else {
warn!(
event_name = "local_core_finalize_missing_error_report_mapping",

View File

@@ -23,7 +23,6 @@ use crate::api::response::{
attach_control_metadata_headers, build_client_response, build_client_response_from_parts,
};
use crate::clock::current_unix_ms as current_request_candidate_unix_ms;
use crate::constants::{CONTROL_CANDIDATE_ID_HEADER, CONTROL_REQUEST_ID_HEADER};
use crate::control::GatewayControlDecision;
#[cfg(test)]
use crate::execution_runtime::remote_compat::post_sync_plan_to_remote_execution_runtime;
@@ -57,7 +56,8 @@ use policy::decode_execution_result_body;
pub(crate) use response::{
maybe_build_local_sync_finalize_response, maybe_build_local_video_error_response,
maybe_build_local_video_success_outcome, resolve_local_sync_error_background_report_kind,
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessOutcome,
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessBuild,
LocalVideoSyncSuccessOutcome,
};
struct ImplicitSyncFinalizeOutcome {
@@ -75,7 +75,65 @@ fn record_sync_terminal_usage(
let payload_seed = build_sync_terminal_usage_payload_seed(payload);
state
.usage_runtime
.record_sync_terminal(state.data.as_ref(), &context_seed, &payload_seed);
.record_sync_terminal(state.data.as_ref(), context_seed, payload_seed);
}
fn build_sync_report_payload(
trace_id: &str,
report_kind: String,
report_context: Option<serde_json::Value>,
status_code: u16,
headers: BTreeMap<String, String>,
body_json: Option<serde_json::Value>,
body_base64: Option<String>,
telemetry: Option<ExecutionTelemetry>,
) -> GatewaySyncReportRequest {
GatewaySyncReportRequest {
trace_id: trace_id.to_string(),
report_kind,
report_context,
status_code,
headers,
body_json,
client_body_json: None,
body_base64,
telemetry,
}
}
async fn apply_sync_success_effects(
state: &AppState,
plan: &ExecutionPlan,
report_context: Option<&serde_json::Value>,
payload: &GatewaySyncReportRequest,
) {
apply_local_execution_effect(
state,
LocalExecutionEffectContext {
plan,
report_context,
},
LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect),
)
.await;
apply_local_execution_effect(
state,
LocalExecutionEffectContext {
plan,
report_context,
},
LocalExecutionEffect::AdaptiveSuccess(LocalAdaptiveSuccessEffect),
)
.await;
apply_local_execution_effect(
state,
LocalExecutionEffectContext {
plan,
report_context,
},
LocalExecutionEffect::PoolSuccessSync { payload },
)
.await;
}
#[cfg(test)]
@@ -112,7 +170,7 @@ pub(crate) async fn execute_execution_runtime_sync(
let lifecycle_seed = build_lifecycle_usage_seed(&plan, report_context.as_ref());
state
.usage_runtime
.record_pending(state.data.as_ref(), &lifecycle_seed);
.record_pending(state.data.as_ref(), lifecycle_seed);
record_local_request_candidate_status(
state,
&plan,
@@ -129,11 +187,8 @@ pub(crate) async fn execute_execution_runtime_sync(
)
.await;
#[cfg(not(test))]
let result = {
match DirectSyncExecutionRuntime::new()
.execute_sync(plan.clone())
.await
{
let mut result = {
match DirectSyncExecutionRuntime::new().execute_sync(&plan).await {
Ok(result) => result,
Err(err) => {
warn!(
@@ -171,7 +226,7 @@ pub(crate) async fn execute_execution_runtime_sync(
}
};
#[cfg(test)]
let result = {
let mut result = {
if let Some(override_fn) = state.execution_runtime_sync_override.as_ref() {
match (override_fn.0)(&plan) {
Ok(result) => result,
@@ -215,10 +270,7 @@ pub(crate) async fn execute_execution_runtime_sync(
.trim()
.is_empty()
{
match DirectSyncExecutionRuntime::new()
.execute_sync(plan.clone())
.await
{
match DirectSyncExecutionRuntime::new().execute_sync(&plan).await {
Ok(result) => result,
Err(err) => {
warn!(
@@ -287,8 +339,9 @@ pub(crate) async fn execute_execution_runtime_sync(
.telemetry
.as_ref()
.and_then(|telemetry| telemetry.elapsed_ms);
let mut headers = result.headers.clone();
let (body_bytes, body_json, body_base64) = decode_execution_result_body(&result, &mut headers)?;
let mut headers = std::mem::take(&mut result.headers);
let (body_bytes, body_json, body_base64) =
decode_execution_result_body(result.body.take(), &mut headers)?;
let local_failover_response_text = local_failover_response_text(
body_json.as_ref(),
&body_bytes,
@@ -403,38 +456,26 @@ pub(crate) async fn execute_execution_runtime_sync(
);
return Ok(None);
}
let request_id = (!result.request_id.trim().is_empty())
.then_some(result.request_id.as_str())
.or(Some(plan_request_id));
let request_id_for_log = short_request_id(request_id.unwrap_or("-"));
let candidate_id = result.candidate_id.as_deref().or(plan_candidate_id);
let status_code = result.status_code;
let has_body_bytes = body_base64.is_some();
let explicit_finalize = should_finalize_sync_response(report_kind.as_deref());
let mapped_error_finalize_kind =
resolve_core_sync_error_finalize_report_kind(plan_kind, &result, body_json.as_ref());
let implicit_finalize = if explicit_finalize || mapped_error_finalize_kind.is_some() {
None
} else {
let implicit_finalize = if !explicit_finalize && mapped_error_finalize_kind.is_none() {
maybe_build_implicit_sync_finalize_outcome(
trace_id,
decision,
plan_kind,
report_context.clone(),
result.status_code,
headers.clone(),
body_json.clone(),
body_base64.clone(),
result.telemetry.clone(),
&report_context,
status_code,
&headers,
&body_json,
&body_base64,
&result.telemetry,
)?
};
let finalize_report_kind = if explicit_finalize {
report_kind.clone()
} else if let Some(implicit_finalize) = implicit_finalize.as_ref() {
Some(implicit_finalize.payload.report_kind.clone())
} else {
mapped_error_finalize_kind.clone()
None
};
if !matches!(
local_failover_analysis.decision,
LocalFailoverDecision::StopLocalFailover
@@ -486,91 +527,83 @@ pub(crate) async fn execute_execution_runtime_sync(
)
.await;
let base_usage_payload = GatewaySyncReportRequest {
trace_id: trace_id.to_string(),
report_kind: finalize_report_kind
.clone()
.or_else(|| report_kind.clone())
.unwrap_or_default(),
report_context: report_context.clone(),
status_code: result.status_code,
headers: headers.clone(),
body_json: body_json.clone(),
client_body_json: None,
body_base64: body_base64.clone(),
telemetry: result.telemetry.clone(),
};
if result.status_code < 400 {
apply_local_execution_effect(
let request_id_owned = result.request_id;
let candidate_id_owned = result.candidate_id;
let request_id = (!request_id_owned.trim().is_empty())
.then_some(request_id_owned.as_str())
.or(Some(plan_request_id));
let request_id_for_log = short_request_id(request_id.unwrap_or("-"));
let candidate_id = candidate_id_owned.as_deref().or(plan_candidate_id);
let report_context = report_context;
let headers = headers;
let body_json = body_json;
let telemetry = result.telemetry;
if let Some(implicit_finalize) = implicit_finalize {
let usage_payload = implicit_finalize
.outcome
.background_report
.as_ref()
.unwrap_or(&implicit_finalize.payload);
apply_sync_success_effects(
state,
LocalExecutionEffectContext {
plan: &plan,
report_context: report_context.as_ref(),
},
LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect),
&plan,
implicit_finalize.payload.report_context.as_ref(),
usage_payload,
)
.await;
apply_local_execution_effect(
record_sync_terminal_usage(
state,
LocalExecutionEffectContext {
plan: &plan,
report_context: report_context.as_ref(),
},
LocalExecutionEffect::AdaptiveSuccess(LocalAdaptiveSuccessEffect),
)
.await;
apply_local_execution_effect(
state,
LocalExecutionEffectContext {
plan: &plan,
report_context: report_context.as_ref(),
},
LocalExecutionEffect::PoolSuccessSync {
payload: &base_usage_payload,
},
)
.await;
&plan,
implicit_finalize.payload.report_context.as_ref(),
usage_payload,
);
if let Some(report_payload) = implicit_finalize.outcome.background_report {
spawn_sync_report(state.clone(), report_payload);
} else {
warn!(
event_name = "local_core_finalize_missing_success_report_mapping",
log_type = "event",
trace_id = %trace_id,
report_kind = %implicit_finalize.payload.report_kind,
"gateway implicit local core finalize produced response without background success report mapping"
);
}
return Ok(Some(attach_control_metadata_headers(
implicit_finalize.outcome.response,
request_id,
candidate_id,
)?));
}
if let Some(finalize_report_kind) = finalize_report_kind {
if let Some(implicit_finalize) = implicit_finalize {
let usage_payload = implicit_finalize
.outcome
.background_report
.as_ref()
.unwrap_or(&implicit_finalize.payload);
record_sync_terminal_usage(state, &plan, report_context.as_ref(), usage_payload);
if let Some(report_payload) = implicit_finalize.outcome.background_report {
spawn_sync_report(state.clone(), trace_id.to_string(), report_payload);
} else {
warn!(
event_name = "local_core_finalize_missing_success_report_mapping",
log_type = "event",
trace_id = %trace_id,
report_kind = %implicit_finalize.payload.report_kind,
"gateway implicit local core finalize produced response without background success report mapping"
);
}
return Ok(Some(attach_control_metadata_headers(
implicit_finalize.outcome.response,
request_id,
candidate_id,
)?));
}
let finalize_report_kind = if explicit_finalize {
report_kind.clone()
} else {
mapped_error_finalize_kind
};
let payload = GatewaySyncReportRequest {
trace_id: trace_id.to_string(),
report_kind: finalize_report_kind,
if let Some(finalize_report_kind) = finalize_report_kind {
let mut payload = build_sync_report_payload(
trace_id,
finalize_report_kind,
report_context,
status_code: result.status_code,
headers: headers.clone(),
body_json: body_json.clone(),
client_body_json: None,
body_base64: body_base64.clone(),
telemetry: result.telemetry.clone(),
};
status_code,
headers,
body_json,
body_base64,
telemetry,
);
if let Some(outcome) = maybe_build_sync_finalize_outcome(trace_id, decision, &payload)? {
let usage_payload = outcome.background_report.as_ref().unwrap_or(&payload);
if status_code < 400 {
apply_sync_success_effects(
state,
&plan,
payload.report_context.as_ref(),
usage_payload,
)
.await;
}
record_sync_terminal_usage(
state,
&plan,
@@ -578,7 +611,7 @@ pub(crate) async fn execute_execution_runtime_sync(
usage_payload,
);
if let Some(report_payload) = outcome.background_report {
spawn_sync_report(state.clone(), trace_id.to_string(), report_payload);
spawn_sync_report(state.clone(), report_payload);
} else {
warn!(
event_name = "local_core_finalize_missing_success_report_mapping",
@@ -594,55 +627,62 @@ pub(crate) async fn execute_execution_runtime_sync(
candidate_id,
)?));
}
if let Some(outcome) = maybe_build_local_video_success_outcome(
let mut payload = match maybe_build_local_video_success_outcome(
trace_id,
decision,
&payload,
payload,
&state.video_tasks,
&plan,
)? {
record_sync_terminal_usage(
state,
&plan,
payload.report_context.as_ref(),
&outcome.report_payload,
);
if let Some(snapshot) = outcome.local_task_snapshot.clone() {
state.video_tasks.record_snapshot(snapshot.clone());
let _ = state.upsert_video_task_snapshot(&snapshot).await?;
}
match outcome.report_mode {
VideoTaskSyncReportMode::InlineSync => {
submit_sync_report(state, trace_id, outcome.report_payload).await?;
LocalVideoSyncSuccessBuild::Handled(outcome) => {
let LocalVideoSyncSuccessOutcome {
response,
report_payload,
original_report_context,
report_mode,
local_task_snapshot,
} = outcome;
apply_sync_success_effects(
state,
&plan,
original_report_context.as_ref(),
&report_payload,
)
.await;
record_sync_terminal_usage(
state,
&plan,
original_report_context.as_ref(),
&report_payload,
);
if let Some(snapshot) = local_task_snapshot {
let _ = state.upsert_video_task_snapshot(&snapshot).await?;
state.video_tasks.record_snapshot(snapshot);
}
VideoTaskSyncReportMode::Background => {
spawn_sync_report(state.clone(), trace_id.to_string(), outcome.report_payload);
match report_mode {
VideoTaskSyncReportMode::InlineSync => {
submit_sync_report(state, report_payload).await?;
}
VideoTaskSyncReportMode::Background => {
spawn_sync_report(state.clone(), report_payload);
}
}
return Ok(Some(attach_control_metadata_headers(
response,
request_id,
candidate_id,
)?));
}
return Ok(Some(attach_control_metadata_headers(
outcome.response,
request_id,
candidate_id,
)?));
}
LocalVideoSyncSuccessBuild::NotHandled(payload) => payload,
};
if let Some(response) =
maybe_build_local_sync_finalize_response(trace_id, decision, &payload)?
{
let usage_payload = if let Some(success_report_kind) =
resolve_local_sync_success_background_report_kind(payload.report_kind.as_str())
{
let mut report_payload = payload.clone();
report_payload.report_kind = success_report_kind.to_string();
report_payload
} else {
payload.clone()
};
record_sync_terminal_usage(
state,
&plan,
payload.report_context.as_ref(),
&usage_payload,
);
let background_success_report_kind =
resolve_local_sync_success_background_report_kind(payload.report_kind.as_str());
apply_sync_success_effects(state, &plan, payload.report_context.as_ref(), &payload)
.await;
record_sync_terminal_usage(state, &plan, payload.report_context.as_ref(), &payload);
state
.video_tasks
.apply_finalize_mutation(request_path, payload.report_kind.as_str());
@@ -652,12 +692,11 @@ pub(crate) async fn execute_execution_runtime_sync(
{
let _ = state.upsert_video_task_snapshot(&snapshot).await?;
}
if let Some(success_report_kind) =
resolve_local_sync_success_background_report_kind(payload.report_kind.as_str())
{
let mut report_payload = usage_payload;
report_payload.report_kind = success_report_kind.to_string();
spawn_sync_report(state.clone(), trace_id.to_string(), report_payload);
if let Some(success_report_kind) = background_success_report_kind {
payload.report_kind = success_report_kind.to_string();
}
if background_success_report_kind.is_some() {
spawn_sync_report(state.clone(), payload);
} else {
warn!(
event_name = "local_video_finalize_missing_success_report_mapping",
@@ -678,27 +717,14 @@ pub(crate) async fn execute_execution_runtime_sync(
if let Some(response) =
maybe_build_local_video_error_response(trace_id, decision, &payload)?
{
let usage_payload = if let Some(error_report_kind) =
resolve_local_sync_error_background_report_kind(payload.report_kind.as_str())
{
let mut report_payload = payload.clone();
report_payload.report_kind = error_report_kind.to_string();
report_payload
} else {
payload.clone()
};
record_sync_terminal_usage(
state,
&plan,
payload.report_context.as_ref(),
&usage_payload,
);
if let Some(error_report_kind) =
resolve_local_sync_error_background_report_kind(payload.report_kind.as_str())
{
let mut report_payload = usage_payload;
report_payload.report_kind = error_report_kind.to_string();
spawn_sync_report(state.clone(), trace_id.to_string(), report_payload);
let background_error_report_kind =
resolve_local_sync_error_background_report_kind(payload.report_kind.as_str());
if let Some(error_report_kind) = background_error_report_kind {
payload.report_kind = error_report_kind.to_string();
}
record_sync_terminal_usage(state, &plan, payload.report_context.as_ref(), &payload);
if background_error_report_kind.is_some() {
spawn_sync_report(state.clone(), payload);
} else {
warn!(
event_name = "local_video_finalize_missing_error_report_mapping",
@@ -726,49 +752,47 @@ pub(crate) async fn execute_execution_runtime_sync(
)?));
}
record_sync_terminal_usage(state, &plan, report_context.as_ref(), &base_usage_payload);
if let Some(report_kind) = report_kind {
let report = GatewaySyncReportRequest {
trace_id: trace_id.to_string(),
report_kind,
report_context,
status_code: result.status_code,
headers: headers.clone(),
body_json: body_json.clone(),
client_body_json: None,
body_base64: body_base64.clone(),
telemetry: result.telemetry.clone(),
};
spawn_sync_report(state.clone(), trace_id.to_string(), report);
}
let request_id_header: Option<&str> = request_id
.map(str::trim)
.filter(|value: &&str| !value.is_empty());
if let Some(request_id) = request_id_header {
headers.insert(
CONTROL_REQUEST_ID_HEADER.to_string(),
request_id.to_string(),
);
}
let candidate_id_header: Option<&str> = candidate_id
.map(str::trim)
.filter(|value: &&str| !value.is_empty());
if let Some(candidate_id) = candidate_id_header {
headers.insert(
CONTROL_CANDIDATE_ID_HEADER.to_string(),
candidate_id.to_string(),
);
}
Ok(Some(build_client_response_from_parts(
result.status_code,
&headers,
Body::from(body_bytes),
let usage_payload = build_sync_report_payload(
trace_id,
Some(decision),
)?))
report_kind.unwrap_or_default(),
report_context,
status_code,
headers,
body_json,
body_base64,
telemetry,
);
if status_code < 400 {
apply_sync_success_effects(
state,
&plan,
usage_payload.report_context.as_ref(),
&usage_payload,
)
.await;
}
record_sync_terminal_usage(
state,
&plan,
usage_payload.report_context.as_ref(),
&usage_payload,
);
let response = attach_control_metadata_headers(
build_client_response_from_parts(
status_code,
&usage_payload.headers,
Body::from(body_bytes),
trace_id,
Some(decision),
)?,
request_id,
candidate_id,
)?;
if !usage_payload.report_kind.trim().is_empty() {
spawn_sync_report(state.clone(), usage_payload);
}
Ok(Some(response))
}
#[allow(clippy::too_many_arguments)] // mirrors sync execution context
@@ -776,12 +800,12 @@ fn maybe_build_implicit_sync_finalize_outcome(
trace_id: &str,
decision: &GatewayControlDecision,
plan_kind: &str,
report_context: Option<serde_json::Value>,
report_context: &Option<serde_json::Value>,
status_code: u16,
headers: BTreeMap<String, String>,
body_json: Option<serde_json::Value>,
body_base64: Option<String>,
telemetry: Option<ExecutionTelemetry>,
headers: &BTreeMap<String, String>,
body_json: &Option<serde_json::Value>,
body_base64: &Option<String>,
telemetry: &Option<ExecutionTelemetry>,
) -> Result<Option<ImplicitSyncFinalizeOutcome>, GatewayError> {
if status_code >= 400 || body_json.is_some() || body_base64.is_none() {
return Ok(None);
@@ -794,13 +818,13 @@ fn maybe_build_implicit_sync_finalize_outcome(
let payload = GatewaySyncReportRequest {
trace_id: trace_id.to_string(),
report_kind: report_kind.to_string(),
report_context,
report_context: report_context.clone(),
status_code,
headers,
body_json,
headers: headers.clone(),
body_json: body_json.clone(),
client_body_json: None,
body_base64,
telemetry,
body_base64: body_base64.clone(),
telemetry: telemetry.clone(),
};
let Some(outcome) = maybe_build_sync_finalize_outcome(trace_id, decision, &payload)? else {
return Ok(None);

View File

@@ -1,6 +1,6 @@
use std::collections::BTreeMap;
use aether_contracts::ExecutionResult;
use aether_contracts::ResponseBody;
use base64::Engine as _;
use crate::GatewayError;
@@ -8,14 +8,14 @@ use crate::GatewayError;
type DecodedBody = (Vec<u8>, Option<serde_json::Value>, Option<String>);
pub(super) fn decode_execution_result_body(
result: &ExecutionResult,
body: Option<ResponseBody>,
headers: &mut BTreeMap<String, String>,
) -> Result<DecodedBody, GatewayError> {
let Some(body) = result.body.as_ref() else {
let Some(body) = body else {
return Ok((Vec::new(), None, None));
};
if let Some(json_body) = body.json_body.clone() {
if let Some(json_body) = body.json_body {
headers
.entry("content-type".to_string())
.or_insert_with(|| "application/json".to_string());
@@ -25,7 +25,7 @@ pub(super) fn decode_execution_result_body(
return Ok((bytes, Some(json_body), None));
}
if let Some(body_bytes_b64) = body.body_bytes_b64.clone() {
if let Some(body_bytes_b64) = body.body_bytes_b64 {
let bytes = base64::engine::general_purpose::STANDARD
.decode(&body_bytes_b64)
.map_err(|err| GatewayError::Internal(err.to_string()))?;

View File

@@ -2,10 +2,13 @@ use std::collections::BTreeMap;
use aether_contracts::ExecutionPlan;
use axum::body::Body;
use axum::http::header::HeaderValue;
use axum::http::Response;
use serde_json::json;
use crate::api::response::build_client_response_from_parts;
use crate::api::response::{
build_client_response_from_parts, build_client_response_from_parts_with_mutator,
};
use crate::async_task::VideoTaskService;
use crate::control::GatewayControlDecision;
use crate::video_tasks::{
@@ -17,9 +20,15 @@ pub(crate) use crate::video_tasks::{
};
use crate::{usage::GatewaySyncReportRequest, GatewayError};
pub(crate) enum LocalVideoSyncSuccessBuild {
Handled(LocalVideoSyncSuccessOutcome),
NotHandled(GatewaySyncReportRequest),
}
pub(crate) struct LocalVideoSyncSuccessOutcome {
pub(crate) response: Response<Body>,
pub(crate) report_payload: GatewaySyncReportRequest,
pub(crate) original_report_context: Option<serde_json::Value>,
pub(crate) report_mode: VideoTaskSyncReportMode,
pub(crate) local_task_snapshot: Option<LocalVideoTaskSnapshot>,
}
@@ -29,8 +38,9 @@ fn cloned_report_context_object(
) -> serde_json::Map<String, serde_json::Value> {
payload
.report_context
.clone()
.and_then(|value| value.as_object().cloned())
.as_ref()
.and_then(serde_json::Value::as_object)
.cloned()
.unwrap_or_default()
}
@@ -56,54 +66,53 @@ fn build_local_video_success_response(
pub(crate) fn maybe_build_local_video_success_outcome(
trace_id: &str,
decision: &GatewayControlDecision,
payload: &GatewaySyncReportRequest,
mut payload: GatewaySyncReportRequest,
video_tasks: &VideoTaskService,
plan: &ExecutionPlan,
) -> Result<Option<LocalVideoSyncSuccessOutcome>, GatewayError> {
) -> Result<LocalVideoSyncSuccessBuild, GatewayError> {
if payload.status_code >= 400 {
return Ok(None);
return Ok(LocalVideoSyncSuccessBuild::NotHandled(payload));
}
let provider_body = match payload
.body_json
.as_ref()
.and_then(serde_json::Value::as_object)
{
Some(value) => value,
None => return Ok(None),
let mut report_context = cloned_report_context_object(&payload);
let prepared_plan = {
let provider_body = match payload
.body_json
.as_ref()
.and_then(serde_json::Value::as_object)
{
Some(value) => value,
None => return Ok(LocalVideoSyncSuccessBuild::NotHandled(payload)),
};
video_tasks.prepare_sync_success(
payload.report_kind.as_str(),
provider_body,
&report_context,
plan,
)
};
let mut report_context = cloned_report_context_object(payload);
let Some(plan) = video_tasks.prepare_sync_success(
payload.report_kind.as_str(),
provider_body,
&report_context,
plan,
) else {
return Ok(None);
let Some(plan) = prepared_plan else {
return Ok(LocalVideoSyncSuccessBuild::NotHandled(payload));
};
plan.apply_to_report_context(&mut report_context);
let client_body_json = plan.client_body_json();
let response = build_local_video_success_response(trace_id, decision, &client_body_json)?;
let report_payload = GatewaySyncReportRequest {
trace_id: payload.trace_id.clone(),
report_kind: plan.success_report_kind().to_string(),
report_context: Some(serde_json::Value::Object(report_context)),
status_code: payload.status_code,
headers: payload.headers.clone(),
body_json: payload.body_json.clone(),
client_body_json: Some(client_body_json),
body_base64: None,
telemetry: payload.telemetry.clone(),
};
let original_report_context = payload.report_context.take();
payload.report_kind = plan.success_report_kind().to_string();
payload.report_context = Some(serde_json::Value::Object(report_context));
payload.client_body_json = Some(client_body_json);
Ok(Some(LocalVideoSyncSuccessOutcome {
response,
report_payload,
report_mode: plan.report_mode(),
local_task_snapshot: matches!(plan.report_mode(), VideoTaskSyncReportMode::Background)
.then(|| plan.to_snapshot()),
}))
Ok(LocalVideoSyncSuccessBuild::Handled(
LocalVideoSyncSuccessOutcome {
response,
report_payload: payload,
original_report_context,
report_mode: plan.report_mode(),
local_task_snapshot: matches!(plan.report_mode(), VideoTaskSyncReportMode::Background)
.then(|| plan.to_snapshot()),
},
))
}
pub(crate) fn maybe_build_local_sync_finalize_response(
@@ -147,21 +156,114 @@ pub(crate) fn maybe_build_local_video_error_response(
return Ok(None);
}
let response_body = payload.body_json.clone().unwrap_or_else(|| json!({}));
let body_bytes = serde_json::to_vec(&response_body)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let empty_body = json!({});
let response_body = payload.body_json.as_ref().unwrap_or(&empty_body);
let body_bytes =
serde_json::to_vec(response_body).map_err(|err| GatewayError::Internal(err.to_string()))?;
let body_len = body_bytes.len().to_string();
let mut response_headers = payload.headers.clone();
response_headers.remove("content-encoding");
response_headers.remove("content-length");
response_headers.insert("content-type".to_string(), "application/json".to_string());
response_headers.insert("content-length".to_string(), body_bytes.len().to_string());
Ok(Some(build_client_response_from_parts(
Ok(Some(build_client_response_from_parts_with_mutator(
payload.status_code,
&response_headers,
&payload.headers,
Body::from(body_bytes),
trace_id,
Some(decision),
|headers| {
headers.remove(http::header::CONTENT_ENCODING);
headers.remove(http::header::CONTENT_LENGTH);
headers.insert(
http::header::CONTENT_TYPE,
HeaderValue::from_static("application/json"),
);
headers.insert(
http::header::CONTENT_LENGTH,
HeaderValue::from_str(body_len.as_str())
.map_err(|err| GatewayError::Internal(err.to_string()))?,
);
Ok(())
},
)?))
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::to_bytes;
use serde_json::json;
#[tokio::test]
async fn local_video_error_response_rewrites_headers_without_mutating_payload() {
let decision = GatewayControlDecision::synthetic(
"/v1/videos",
Some("ai_public".to_string()),
Some("openai".to_string()),
Some("video".to_string()),
Some("openai:video".to_string()),
)
.with_execution_runtime_candidate(true);
let payload = GatewaySyncReportRequest {
trace_id: "trace-payload".to_string(),
report_kind: "openai_video_create_sync_finalize".to_string(),
report_context: Some(json!({
"request_id": "req_123",
})),
status_code: http::StatusCode::BAD_GATEWAY.as_u16(),
headers: BTreeMap::from([
("content-encoding".to_string(), "gzip".to_string()),
("content-length".to_string(), "999".to_string()),
("x-upstream-id".to_string(), "video-123".to_string()),
]),
body_json: Some(json!({
"error": {
"type": "video_backend_error",
"message": "backend failed",
}
})),
client_body_json: None,
body_base64: None,
telemetry: None,
};
let response =
maybe_build_local_video_error_response("trace-response", &decision, &payload)
.expect("video error response should build")
.expect("video error response should match local video error kinds");
assert_eq!(response.status(), http::StatusCode::BAD_GATEWAY);
assert_eq!(
response
.headers()
.get(http::header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok()),
Some("application/json")
);
assert_eq!(response.headers().get(http::header::CONTENT_ENCODING), None);
assert_eq!(
response
.headers()
.get("x-upstream-id")
.and_then(|value| value.to_str().ok()),
Some("video-123")
);
assert_eq!(
payload.headers.get("content-encoding").map(String::as_str),
Some("gzip")
);
assert_eq!(
payload.headers.get("content-length").map(String::as_str),
Some("999")
);
let body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("response body should read");
assert_eq!(
serde_json::from_slice::<serde_json::Value>(&body).expect("response body should parse"),
payload
.body_json
.clone()
.expect("payload body should exist")
);
}
}

View File

@@ -6,5 +6,6 @@ pub(crate) use execution::execute_execution_runtime_sync;
pub(crate) use execution::{
maybe_build_local_sync_finalize_response, maybe_build_local_video_error_response,
maybe_build_local_video_success_outcome, resolve_local_sync_error_background_report_kind,
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessOutcome,
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessBuild,
LocalVideoSyncSuccessOutcome,
};

View File

@@ -161,12 +161,12 @@ impl DirectSyncExecutionRuntime {
pub(crate) async fn execute_sync(
&self,
plan: ExecutionPlan,
plan: &ExecutionPlan,
) -> Result<ExecutionResult, ExecutionRuntimeTransportError> {
let body_bytes = build_request_body(&plan)?;
let body_bytes = build_request_body(plan)?;
let started_at = Instant::now();
let response = send_request(&plan, body_bytes).await?;
let response = send_request(plan, body_bytes).await?;
let ttfb_ms = started_at.elapsed().as_millis() as u64;
let status_code = response.status().as_u16();
let headers = collect_response_headers(response.headers());
@@ -200,8 +200,8 @@ impl DirectSyncExecutionRuntime {
};
Ok(ExecutionResult {
request_id: plan.request_id,
candidate_id: plan.candidate_id,
request_id: plan.request_id.clone(),
candidate_id: plan.candidate_id.clone(),
status_code,
headers,
body,
@@ -216,24 +216,24 @@ impl DirectSyncExecutionRuntime {
pub(crate) async fn execute_stream(
&self,
plan: ExecutionPlan,
plan: &ExecutionPlan,
) -> Result<DirectUpstreamStreamExecution, ExecutionRuntimeTransportError> {
if !plan.stream {
return Err(ExecutionRuntimeTransportError::StreamUnsupported);
}
let body_bytes = build_request_body(&plan)?;
let body_bytes = build_request_body(plan)?;
let started_at = Instant::now();
let response = send_request(&plan, body_bytes).await?;
let response = send_request(plan, body_bytes).await?;
let status_code = response.status().as_u16();
let headers = collect_response_headers(response.headers());
let stream_summary_report_context = build_stream_summary_report_context(&plan);
let stream_summary_report_context = build_stream_summary_report_context(plan);
Ok(DirectUpstreamStreamExecution {
request_id: plan.request_id,
candidate_id: plan.candidate_id,
request_id: plan.request_id.clone(),
candidate_id: plan.candidate_id.clone(),
status_code,
headers,
provider_api_format: plan.provider_api_format.clone(),
@@ -274,7 +274,7 @@ pub(crate) async fn execute_sync_plan(
let _ = state;
let _ = trace_id;
DirectSyncExecutionRuntime::new()
.execute_sync(plan.clone())
.execute_sync(plan)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
@@ -1101,7 +1101,7 @@ mod tests {
let execution_runtime = DirectSyncExecutionRuntime::new();
let result = execution_runtime
.execute_sync(ExecutionPlan {
.execute_sync(&ExecutionPlan {
request_id: "req-1".into(),
candidate_id: Some("cand-1".into()),
provider_name: Some("openai".into()),
@@ -1168,7 +1168,7 @@ mod tests {
let execution_runtime = DirectSyncExecutionRuntime::new();
let result = execution_runtime
.execute_sync(ExecutionPlan {
.execute_sync(&ExecutionPlan {
request_id: "req-1".into(),
candidate_id: None,
provider_name: None,
@@ -1369,7 +1369,7 @@ mod tests {
let execution_runtime = DirectSyncExecutionRuntime::new();
let result = execution_runtime
.execute_sync(ExecutionPlan {
.execute_sync(&ExecutionPlan {
request_id: "req-redirect-1".into(),
candidate_id: None,
provider_name: Some("provider_ops".into()),
@@ -1442,7 +1442,7 @@ mod tests {
let execution_runtime = DirectSyncExecutionRuntime::new();
let result = execution_runtime
.execute_sync(ExecutionPlan {
.execute_sync(&ExecutionPlan {
request_id: "req-redirect-2".into(),
candidate_id: None,
provider_name: Some("provider_oauth".into()),
@@ -1512,7 +1512,7 @@ mod tests {
let execution_runtime = DirectSyncExecutionRuntime::new();
let result = execution_runtime
.execute_sync(ExecutionPlan {
.execute_sync(&ExecutionPlan {
request_id: "req-relay-http1-1".into(),
candidate_id: None,
provider_name: Some("provider_ops".into()),
@@ -1579,7 +1579,7 @@ mod tests {
let execution_runtime = DirectSyncExecutionRuntime::new();
let result = execution_runtime
.execute_sync(ExecutionPlan {
.execute_sync(&ExecutionPlan {
request_id: "req-tls-1".into(),
candidate_id: Some("cand-1".into()),
provider_name: Some("claude".into()),
@@ -1654,7 +1654,7 @@ mod tests {
let execution_runtime = DirectSyncExecutionRuntime::new();
let result = execution_runtime
.execute_sync(ExecutionPlan {
.execute_sync(&ExecutionPlan {
request_id: "req-gzip-1".into(),
candidate_id: Some("cand-1".into()),
provider_name: Some("openai".into()),
@@ -1715,7 +1715,7 @@ mod tests {
let execution_runtime = DirectSyncExecutionRuntime::new();
let result = execution_runtime
.execute_sync(ExecutionPlan {
.execute_sync(&ExecutionPlan {
request_id: "req-ttfb-1".into(),
candidate_id: Some("cand-1".into()),
provider_name: Some("openai".into()),

View File

@@ -206,18 +206,18 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
Json(build_internal_gateway_fallback_plan_payload(None)).into_response(),
));
};
if let Some(auth_context) = payload.auth_context.clone() {
let provided_auth_context = payload.auth_context.is_some();
if let Some(auth_context) = payload.auth_context {
resolved.auth_context = Some(auth_context);
resolved.local_auth_rejection = None;
}
let auth_context = resolved.auth_context.clone();
let auth_context = resolved.auth_context.as_ref();
if auth_context
.as_ref()
.map(|value| !value.access_allowed)
.unwrap_or(true)
{
let fallback_auth_context = if payload.auth_context.is_none() {
auth_context.as_ref()
let fallback_auth_context = if !provided_auth_context {
auth_context
} else {
None
};
@@ -239,8 +239,8 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
)
.await?
else {
let fallback_auth_context = if payload.auth_context.is_none() {
auth_context.as_ref()
let fallback_auth_context = if !provided_auth_context {
auth_context
} else {
None
};
@@ -251,7 +251,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
.into_response(),
));
};
if payload.auth_context.is_some() {
if provided_auth_context {
local_payload.auth_context = None;
}
return Ok(Some(Json(local_payload).into_response()));
@@ -308,18 +308,18 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
Json(build_internal_gateway_fallback_plan_payload(None)).into_response(),
));
};
if let Some(auth_context) = payload.auth_context.clone() {
let provided_auth_context = payload.auth_context.is_some();
if let Some(auth_context) = payload.auth_context {
resolved.auth_context = Some(auth_context);
resolved.local_auth_rejection = None;
}
let auth_context = resolved.auth_context.clone();
let auth_context = resolved.auth_context.as_ref();
if auth_context
.as_ref()
.map(|value| !value.access_allowed)
.unwrap_or(true)
{
let fallback_auth_context = if payload.auth_context.is_none() {
auth_context.as_ref()
let fallback_auth_context = if !provided_auth_context {
auth_context
} else {
None
};
@@ -339,8 +339,8 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
)
.await?
else {
let fallback_auth_context = if payload.auth_context.is_none() {
auth_context.as_ref()
let fallback_auth_context = if !provided_auth_context {
auth_context
} else {
None
};
@@ -351,7 +351,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
.into_response(),
));
};
if payload.auth_context.is_some() {
if provided_auth_context {
local_payload.auth_context = None;
}
return Ok(Some(Json(local_payload).into_response()));
@@ -406,7 +406,8 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
else {
return Ok(Some(build_internal_gateway_proxy_public_response()));
};
if let Some(auth_context) = payload.auth_context.clone() {
let provided_auth_context = payload.auth_context.is_some();
if let Some(auth_context) = payload.auth_context {
resolved.auth_context = Some(auth_context);
resolved.local_auth_rejection = None;
}
@@ -421,7 +422,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
)
.await?
{
if payload.auth_context.is_some() {
if provided_auth_context {
planned.auth_context = None;
}
return Ok(Some(Json(planned).into_response()));
@@ -472,7 +473,8 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
else {
return Ok(Some(build_internal_gateway_proxy_public_response()));
};
if let Some(auth_context) = payload.auth_context.clone() {
let provided_auth_context = payload.auth_context.is_some();
if let Some(auth_context) = payload.auth_context {
resolved.auth_context = Some(auth_context);
resolved.local_auth_rejection = None;
}
@@ -485,7 +487,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
)
.await?
{
if payload.auth_context.is_some() {
if provided_auth_context {
planned.auth_context = None;
}
return Ok(Some(Json(planned).into_response()));
@@ -542,7 +544,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
else {
return Ok(None);
};
if let Some(auth_context) = payload.auth_context.clone() {
if let Some(auth_context) = payload.auth_context {
resolved.auth_context = Some(auth_context);
resolved.local_auth_rejection = None;
}
@@ -626,7 +628,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
else {
return Ok(None);
};
if let Some(auth_context) = payload.auth_context.clone() {
if let Some(auth_context) = payload.auth_context {
resolved.auth_context = Some(auth_context);
resolved.local_auth_rejection = None;
}
@@ -683,8 +685,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
)));
}
};
let trace_id = payload.trace_id.clone();
crate::usage::submit_sync_report(state, &trace_id, payload).await?;
crate::usage::submit_sync_report(state, payload).await?;
return Ok(Some(Json(json!({ "ok": true })).into_response()));
}
Some("report_stream")
@@ -707,8 +708,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
)));
}
};
let trace_id = payload.trace_id.clone();
crate::usage::submit_stream_report(state, &trace_id, payload).await?;
crate::usage::submit_stream_report(state, payload).await?;
return Ok(Some(Json(json!({ "ok": true })).into_response()));
}
Some("finalize_sync")
@@ -739,16 +739,12 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
};
let trace_id = payload.trace_id.clone();
if let Some(outcome) = ai_pipeline_api::maybe_build_sync_finalize_outcome(
&trace_id,
trace_id.as_str(),
&synthetic_decision,
&payload,
)? {
if let Some(background_report) = outcome.background_report {
crate::usage::spawn_sync_report(
state.clone(),
trace_id.clone(),
background_report,
);
crate::usage::spawn_sync_report(state.clone(), background_report);
}
let mut response = outcome.response;
response.headers_mut().insert(
@@ -761,7 +757,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
state,
trace_id.as_str(),
&synthetic_decision,
&payload,
payload,
)
.await?
{

View File

@@ -11,7 +11,7 @@ use crate::control::GatewayControlDecision;
use crate::execution_runtime::{
maybe_build_local_sync_finalize_response, maybe_build_local_video_error_response,
maybe_build_local_video_success_outcome, resolve_local_sync_error_background_report_kind,
resolve_local_sync_success_background_report_kind,
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessBuild,
};
use crate::handlers::shared::{
unix_secs_to_rfc3339, InternalTunnelHeartbeatRequest, InternalTunnelNodeStatusRequest,
@@ -244,59 +244,66 @@ pub(crate) async fn maybe_build_internal_finalize_video_response(
state: &AppState,
trace_id: &str,
decision: &GatewayControlDecision,
payload: &crate::usage::GatewaySyncReportRequest,
payload: crate::usage::GatewaySyncReportRequest,
) -> Result<Option<Response<Body>>, GatewayError> {
let Some(plan) = infer_internal_finalize_signature(payload).and_then(|signature| {
build_internal_finalize_video_plan(
payload.trace_id.as_str(),
signature.as_str(),
payload.report_context.as_ref(),
)
}) else {
let Some((signature, plan)) =
infer_internal_finalize_signature(&payload).and_then(|signature| {
build_internal_finalize_video_plan(
payload.trace_id.as_str(),
signature.as_str(),
payload.report_context.as_ref(),
)
.map(|plan| (signature, plan))
})
else {
return Ok(None);
};
if let Some(outcome) = maybe_build_local_video_success_outcome(
let mut payload = match maybe_build_local_video_success_outcome(
trace_id,
decision,
payload,
&state.video_tasks,
&plan,
)? {
if let Some(snapshot) = outcome.local_task_snapshot.clone() {
state.video_tasks.record_snapshot(snapshot.clone());
let _ = state.upsert_video_task_snapshot(&snapshot).await?;
}
match outcome.report_mode {
crate::video_tasks::VideoTaskSyncReportMode::InlineSync => {
crate::usage::submit_sync_report(state, trace_id, outcome.report_payload).await?;
LocalVideoSyncSuccessBuild::Handled(outcome) => {
let crate::execution_runtime::LocalVideoSyncSuccessOutcome {
response,
report_payload,
original_report_context: _,
report_mode,
local_task_snapshot,
} = outcome;
if let Some(snapshot) = local_task_snapshot {
let _ = state.upsert_video_task_snapshot(&snapshot).await?;
state.video_tasks.record_snapshot(snapshot);
}
crate::video_tasks::VideoTaskSyncReportMode::Background => {
crate::usage::spawn_sync_report(
state.clone(),
trace_id.to_string(),
outcome.report_payload,
);
match report_mode {
crate::video_tasks::VideoTaskSyncReportMode::InlineSync => {
crate::usage::submit_sync_report(state, report_payload).await?;
}
crate::video_tasks::VideoTaskSyncReportMode::Background => {
crate::usage::spawn_sync_report(state.clone(), report_payload);
}
}
let mut response = response;
response.headers_mut().insert(
HeaderName::from_static(CONTROL_EXECUTED_HEADER),
HeaderValue::from_static("true"),
);
return Ok(Some(response));
}
let mut response = outcome.response;
response.headers_mut().insert(
HeaderName::from_static(CONTROL_EXECUTED_HEADER),
HeaderValue::from_static("true"),
);
return Ok(Some(response));
}
LocalVideoSyncSuccessBuild::NotHandled(payload) => payload,
};
if let Some(mut response) =
maybe_build_local_sync_finalize_response(trace_id, decision, payload)?
maybe_build_local_sync_finalize_response(trace_id, decision, &payload)?
{
let request_path = infer_internal_finalize_signature(payload).and_then(|signature| {
build_local_sync_finalize_request_path(
payload.report_kind.as_str(),
signature.as_str(),
payload.report_context.as_ref(),
)
});
let request_path = build_local_sync_finalize_request_path(
payload.report_kind.as_str(),
signature.as_str(),
payload.report_context.as_ref(),
);
if let Some(request_path) = request_path {
state
.video_tasks
@@ -311,9 +318,8 @@ pub(crate) async fn maybe_build_internal_finalize_video_response(
if let Some(success_report_kind) =
resolve_local_sync_success_background_report_kind(payload.report_kind.as_str())
{
let mut report_payload = payload.clone();
report_payload.report_kind = success_report_kind.to_string();
crate::usage::spawn_sync_report(state.clone(), trace_id.to_string(), report_payload);
payload.report_kind = success_report_kind.to_string();
crate::usage::spawn_sync_report(state.clone(), payload);
}
response.headers_mut().insert(
HeaderName::from_static(CONTROL_EXECUTED_HEADER),
@@ -322,14 +328,14 @@ pub(crate) async fn maybe_build_internal_finalize_video_response(
return Ok(Some(response));
}
if let Some(mut response) = maybe_build_local_video_error_response(trace_id, decision, payload)?
if let Some(mut response) =
maybe_build_local_video_error_response(trace_id, decision, &payload)?
{
if let Some(error_report_kind) =
resolve_local_sync_error_background_report_kind(payload.report_kind.as_str())
{
let mut report_payload = payload.clone();
report_payload.report_kind = error_report_kind.to_string();
crate::usage::spawn_sync_report(state.clone(), trace_id.to_string(), report_payload);
payload.report_kind = error_report_kind.to_string();
crate::usage::spawn_sync_report(state.clone(), payload);
}
response.headers_mut().insert(
HeaderName::from_static(CONTROL_EXECUTED_HEADER),

View File

@@ -122,6 +122,10 @@ pub(crate) async fn build_public_providers_payload(
.iter()
.map(|provider| provider.id.clone())
.collect::<Vec<_>>();
let provider_ids_set = provider_ids
.iter()
.map(String::as_str)
.collect::<BTreeSet<_>>();
let endpoints = if provider_ids.is_empty() {
Vec::new()
} else {
@@ -156,7 +160,7 @@ pub(crate) async fn build_public_providers_payload(
.ok()
.unwrap_or_default();
for row in rows {
if provider_ids.contains(&row.provider_id) {
if provider_ids_set.contains(row.provider_id.as_str()) {
models_by_provider
.entry(row.provider_id.clone())
.or_default()
@@ -332,7 +336,6 @@ pub(crate) async fn build_api_format_health_monitor_payload(
let mut endpoint_ids_by_format = BTreeMap::<String, Vec<String>>::new();
let mut endpoint_to_format = BTreeMap::<String, String>::new();
let mut provider_ids_by_format = BTreeMap::<String, BTreeSet<String>>::new();
let mut active_provider_formats = BTreeSet::<(String, String)>::new();
for endpoint in active_endpoints {
endpoint_to_format.insert(endpoint.id.clone(), endpoint.api_format.clone());
endpoint_ids_by_format
@@ -343,7 +346,6 @@ pub(crate) async fn build_api_format_health_monitor_payload(
.entry(endpoint.api_format.clone())
.or_default()
.insert(endpoint.provider_id.clone());
active_provider_formats.insert((endpoint.provider_id, endpoint.api_format));
}
let all_endpoint_ids = endpoint_to_format.keys().cloned().collect::<Vec<_>>();
@@ -356,7 +358,9 @@ pub(crate) async fn build_api_format_health_monitor_payload(
.unwrap_or_default();
for key in keys.into_iter().filter(|key| key.is_active) {
for api_format in provider_key_api_formats(&key) {
if active_provider_formats.contains(&(key.provider_id.clone(), api_format.clone()))
if provider_ids_by_format
.get(&api_format)
.is_some_and(|provider_ids| provider_ids.contains(key.provider_id.as_str()))
{
*key_counts_by_format.entry(api_format).or_default() += 1;
}

View File

@@ -8,6 +8,7 @@ use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
use aether_crypto::{decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use serde_json::{json, Map, Value};
use std::borrow::Cow;
use std::time::{SystemTime, UNIX_EPOCH};
const OAUTH_ACCOUNT_BLOCK_PREFIX: &str = "[ACCOUNT_BLOCK] ";
@@ -63,23 +64,27 @@ pub(crate) fn decrypt_catalog_secret_with_fallbacks(
None
}
pub(crate) fn effective_catalog_encryption_key(state: &AppState) -> Option<String> {
pub(crate) fn effective_catalog_encryption_key(state: &AppState) -> Option<Cow<'_, str>> {
let encryption_key = state.encryption_key().map(str::trim).unwrap_or("");
if !encryption_key.is_empty() {
return Some(encryption_key.to_string());
return Some(Cow::Borrowed(encryption_key));
}
for env_key in ["AETHER_GATEWAY_DATA_ENCRYPTION_KEY", "ENCRYPTION_KEY"] {
let Ok(candidate) = std::env::var(env_key) else {
continue;
};
let candidate = candidate.trim();
if !candidate.is_empty() {
return Some(candidate.to_string());
let trimmed = candidate.trim();
if !trimmed.is_empty() {
return Some(if trimmed.len() == candidate.len() {
Cow::Owned(candidate)
} else {
Cow::Owned(trimmed.to_string())
});
}
}
#[cfg(test)]
{
return Some(DEVELOPMENT_ENCRYPTION_KEY.to_string());
return Some(Cow::Borrowed(DEVELOPMENT_ENCRYPTION_KEY));
}
#[allow(unreachable_code)]
None
@@ -90,7 +95,7 @@ pub(crate) fn encrypt_catalog_secret_with_fallbacks(
plaintext: &str,
) -> Option<String> {
let encryption_key = effective_catalog_encryption_key(state)?;
encrypt_python_fernet_plaintext(&encryption_key, plaintext).ok()
encrypt_python_fernet_plaintext(encryption_key.as_ref(), plaintext).ok()
}
pub(crate) fn masked_catalog_api_key(state: &AppState, key: &StoredProviderCatalogKey) -> String {

View File

@@ -64,13 +64,13 @@ pub(crate) fn build_local_attempt_identities(
candidate_index: u32,
transport: &GatewayProviderTransportSnapshot,
) -> Vec<ExecutionAttemptIdentity> {
let attempt_slots = resolve_local_attempt_slot_count(transport);
let attempt_slots = local_attempt_slot_count(transport);
(0..attempt_slots)
.map(|retry_index| ExecutionAttemptIdentity::new(candidate_index, retry_index))
.collect()
}
fn resolve_local_attempt_slot_count(transport: &GatewayProviderTransportSnapshot) -> u32 {
pub(crate) fn local_attempt_slot_count(transport: &GatewayProviderTransportSnapshot) -> u32 {
local_attempt_slots_from_transport(transport).unwrap_or(1)
}

View File

@@ -16,7 +16,7 @@ pub(crate) use self::adaptive::{
LocalAdaptiveRateLimitProjection, LocalAdaptiveSuccessProjection,
};
pub(crate) use self::attempt::{
attempt_identity_from_report_context, build_local_attempt_identities,
attempt_identity_from_report_context, build_local_attempt_identities, local_attempt_slot_count,
local_execution_candidate_metadata_from_report_context, ExecutionAttemptIdentity,
LocalExecutionCandidateMetadata,
};

View File

@@ -21,6 +21,19 @@ use crate::clock::current_unix_ms;
use crate::log_ids::short_request_id;
use crate::GatewayError;
#[derive(Debug, Clone)]
pub(crate) struct LocalRequestCandidateStatusSnapshot {
candidate_id: String,
request_id: String,
user_id: Option<String>,
api_key_id: Option<String>,
candidate_index: u32,
retry_index: u32,
provider_id: String,
endpoint_id: String,
key_id: String,
}
#[async_trait]
pub(crate) trait RequestCandidateRuntimeReader {
async fn read_request_candidates_by_request_id(
@@ -149,23 +162,37 @@ fn request_candidate_status_label(status: RequestCandidateStatus) -> &'static st
}
}
pub(crate) async fn record_local_request_candidate_status(
state: &(impl RequestCandidateRuntimeWriter + ?Sized),
pub(crate) fn snapshot_local_request_candidate_status(
plan: &ExecutionPlan,
report_context: Option<&Value>,
status_update: SchedulerRequestCandidateStatusUpdate,
) -> Option<LocalRequestCandidateStatusSnapshot> {
let candidate_id = plan
.candidate_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())?;
let metadata = parse_request_candidate_report_context(report_context)?;
let candidate_index = metadata.candidate_index?;
Some(LocalRequestCandidateStatusSnapshot {
candidate_id: candidate_id.to_string(),
request_id: plan.request_id.clone(),
user_id: metadata.user_id,
api_key_id: metadata.api_key_id,
candidate_index,
retry_index: metadata.retry_index,
provider_id: plan.provider_id.clone(),
endpoint_id: plan.endpoint_id.clone(),
key_id: plan.key_id.clone(),
})
}
async fn persist_local_request_candidate_status_record(
state: &(impl RequestCandidateRuntimeWriter + ?Sized),
record: UpsertRequestCandidateRecord,
) {
let Some(record) =
build_local_request_candidate_status_record(LocalRequestCandidateStatusRecordInput {
plan,
report_context,
status_update,
})
else {
return;
};
let candidate_id = record.id.clone();
let request_id = short_request_id(plan.request_id.as_str());
let request_id = short_request_id(record.request_id.as_str());
let candidate_index = record.candidate_index;
let retry_index = record.retry_index;
let status = record.status;
@@ -210,6 +237,67 @@ pub(crate) async fn record_local_request_candidate_status(
}
}
pub(crate) async fn record_local_request_candidate_status(
state: &(impl RequestCandidateRuntimeWriter + ?Sized),
plan: &ExecutionPlan,
report_context: Option<&Value>,
status_update: SchedulerRequestCandidateStatusUpdate,
) {
let Some(record) =
build_local_request_candidate_status_record(LocalRequestCandidateStatusRecordInput {
plan,
report_context,
status_update,
})
else {
return;
};
persist_local_request_candidate_status_record(state, record).await;
}
pub(crate) async fn record_local_request_candidate_status_snapshot(
state: &(impl RequestCandidateRuntimeWriter + ?Sized),
snapshot: &LocalRequestCandidateStatusSnapshot,
status_update: SchedulerRequestCandidateStatusUpdate,
) {
let SchedulerRequestCandidateStatusUpdate {
status,
status_code,
error_type,
error_message,
latency_ms,
started_at_unix_ms,
finished_at_unix_ms,
} = status_update;
let record = UpsertRequestCandidateRecord {
id: snapshot.candidate_id.clone(),
request_id: snapshot.request_id.clone(),
user_id: snapshot.user_id.clone(),
api_key_id: snapshot.api_key_id.clone(),
username: None,
api_key_name: None,
candidate_index: snapshot.candidate_index,
retry_index: snapshot.retry_index,
provider_id: Some(snapshot.provider_id.clone()),
endpoint_id: Some(snapshot.endpoint_id.clone()),
key_id: Some(snapshot.key_id.clone()),
status,
skip_reason: None,
is_cached: None,
status_code,
error_type,
error_message,
latency_ms,
concurrent_requests: None,
extra_data: None,
required_capabilities: None,
created_at_unix_ms: None,
started_at_unix_ms,
finished_at_unix_ms,
};
persist_local_request_candidate_status_record(state, record).await;
}
pub(crate) async fn record_report_request_candidate_status(
state: &(impl RequestCandidateRuntimeReader + RequestCandidateRuntimeWriter + ?Sized),
report_context: Option<&Value>,

View File

@@ -389,3 +389,109 @@ async fn gateway_stops_execution_runtime_stream_when_client_disconnects() {
execution_runtime_handle.abort();
upstream_handle.abort();
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn gateway_returns_error_body_when_prefetch_detects_embedded_stream_error() {
let public_hits = Arc::new(Mutex::new(0usize));
let public_hits_clone = Arc::clone(&public_hits);
let upstream = Router::new().route(
"/v1/chat/completions",
any(move |_request: Request| {
let public_hits_inner = Arc::clone(&public_hits_clone);
async move {
*public_hits_inner.lock().expect("mutex should lock") += 1;
(StatusCode::IM_A_TEAPOT, Body::from("public-route-hit"))
}
}),
);
let execution_runtime = Router::new().route(
"/v1/execute/stream",
any(|_request: Request| async move {
let body_stream = async_stream::stream! {
yield Ok::<Bytes, Infallible>(Bytes::from_static(
b"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n"
));
yield Ok::<Bytes, Infallible>(Bytes::from_static(
b"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"error\\\":{\\\"message\\\":\\\"slow down\\\",\\\"type\\\":\\\"rate_limit_error\\\",\\\"code\\\":\\\"rate_limit\\\"}}\\n\\n\"}}\n"
));
};
let mut response = Response::builder()
.status(StatusCode::OK)
.body(Body::from_stream(body_stream))
.expect("response should build");
response.headers_mut().insert(
http::header::CONTENT_TYPE,
HeaderValue::from_static("application/x-ndjson"),
);
response
}),
);
let (upstream_url, upstream_handle) = start_server(upstream).await;
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
Some(hash_api_key("sk-client-openai-stream-prefetch-error")),
sample_local_openai_auth_snapshot(
"api-key-openai-lifecycle-local-1",
"user-openai-lifecycle-local-1",
),
)]));
let candidate_selection_repository =
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
sample_local_openai_candidate_row(),
]));
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_local_openai_provider()],
vec![sample_local_openai_endpoint()],
vec![sample_local_openai_key()],
));
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
let gateway = build_router_with_state(
build_state_with_execution_runtime_override(execution_runtime_url)
.with_data_state_for_tests(
GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
auth_repository,
candidate_selection_repository,
provider_catalog_repository,
request_candidate_repository,
DEVELOPMENT_ENCRYPTION_KEY,
),
),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.post(format!("{gateway_url}/v1/chat/completions"))
.header(http::header::CONTENT_TYPE, "application/json")
.header(
http::header::AUTHORIZATION,
"Bearer sk-client-openai-stream-prefetch-error",
)
.header(
TRACE_ID_HEADER,
"trace-openai-chat-stream-prefetch-error-123",
)
.body("{\"model\":\"gpt-5\",\"messages\":[],\"stream\":true}")
.send()
.await
.expect("request should succeed");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response
.headers()
.get(http::header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok()),
Some("text/event-stream")
);
let body_text = response.text().await.expect("response body should read");
assert!(body_text.contains("\"rate_limit_error\""));
assert!(body_text.contains("\"slow down\""));
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort();
execution_runtime_handle.abort();
upstream_handle.abort();
}

View File

@@ -61,28 +61,28 @@ fn log_dropped_report(
pub(crate) async fn submit_sync_report(
state: &AppState,
trace_id: &str,
payload: GatewaySyncReportRequest,
mut payload: GatewaySyncReportRequest,
) -> Result<(), GatewayError> {
let original_report_context = payload.report_context.take();
if let Some(report_context) =
resolve_locally_actionable_report_context(state, payload.report_context.as_ref()).await
resolve_locally_actionable_report_context(state, original_report_context.as_ref()).await
{
let mut local_payload = payload.clone();
local_payload.report_context = Some(report_context);
payload.report_context = Some(report_context);
if should_handle_local_sync_report(
local_payload.report_context.as_ref(),
local_payload.report_kind.as_str(),
payload.report_context.as_ref(),
payload.report_kind.as_str(),
) {
handle_local_sync_report(state, &local_payload).await;
handle_local_sync_report(state, &payload).await;
log_local_report_handled(
trace_id,
&local_payload.report_kind,
payload.trace_id.as_str(),
&payload.report_kind,
"sync",
local_payload.report_context.as_ref(),
payload.report_context.as_ref(),
);
return Ok(());
}
}
payload.report_context = original_report_context;
if should_handle_local_sync_report(
payload.report_context.as_ref(),
@@ -90,7 +90,7 @@ pub(crate) async fn submit_sync_report(
) {
handle_local_sync_report(state, &payload).await;
log_local_report_handled(
trace_id,
payload.trace_id.as_str(),
&payload.report_kind,
"sync",
payload.report_context.as_ref(),
@@ -99,7 +99,7 @@ pub(crate) async fn submit_sync_report(
}
log_dropped_report(
trace_id,
payload.trace_id.as_str(),
&payload.report_kind,
"sync",
payload.report_context.as_ref(),
@@ -107,15 +107,12 @@ pub(crate) async fn submit_sync_report(
Ok(())
}
pub(crate) fn spawn_sync_report(
state: AppState,
trace_id: String,
payload: GatewaySyncReportRequest,
) {
pub(crate) fn spawn_sync_report(state: AppState, payload: GatewaySyncReportRequest) {
let report_request_id_for_log =
short_request_id(report_request_id(payload.report_context.as_ref()));
tokio::spawn(async move {
if let Err(err) = submit_sync_report(&state, &trace_id, payload).await {
let trace_id = payload.trace_id.clone();
if let Err(err) = submit_sync_report(&state, payload).await {
warn!(
event_name = "execution_report_submit_failed",
log_type = "ops",
@@ -131,39 +128,28 @@ pub(crate) fn spawn_sync_report(
pub(crate) async fn submit_stream_report(
state: &AppState,
trace_id: &str,
payload: GatewayStreamReportRequest,
mut payload: GatewayStreamReportRequest,
) -> Result<(), GatewayError> {
let original_report_context = payload.report_context.take();
if let Some(report_context) =
resolve_locally_actionable_report_context(state, payload.report_context.as_ref()).await
resolve_locally_actionable_report_context(state, original_report_context.as_ref()).await
{
let local_payload = GatewayStreamReportRequest {
trace_id: payload.trace_id.clone(),
report_kind: payload.report_kind.clone(),
report_context: Some(report_context),
status_code: payload.status_code,
headers: payload.headers.clone(),
provider_body_base64: payload.provider_body_base64.clone(),
provider_body_state: payload.provider_body_state,
client_body_base64: payload.client_body_base64.clone(),
client_body_state: payload.client_body_state,
terminal_summary: payload.terminal_summary.clone(),
telemetry: payload.telemetry.clone(),
};
payload.report_context = Some(report_context);
if should_handle_local_stream_report(
local_payload.report_context.as_ref(),
local_payload.report_kind.as_str(),
payload.report_context.as_ref(),
payload.report_kind.as_str(),
) {
handle_local_stream_report(state, &local_payload).await;
handle_local_stream_report(state, &payload).await;
log_local_report_handled(
trace_id,
&local_payload.report_kind,
payload.trace_id.as_str(),
&payload.report_kind,
"stream",
local_payload.report_context.as_ref(),
payload.report_context.as_ref(),
);
return Ok(());
}
}
payload.report_context = original_report_context;
if should_handle_local_stream_report(
payload.report_context.as_ref(),
@@ -171,7 +157,7 @@ pub(crate) async fn submit_stream_report(
) {
handle_local_stream_report(state, &payload).await;
log_local_report_handled(
trace_id,
payload.trace_id.as_str(),
&payload.report_kind,
"stream",
payload.report_context.as_ref(),
@@ -180,7 +166,7 @@ pub(crate) async fn submit_stream_report(
}
log_dropped_report(
trace_id,
payload.trace_id.as_str(),
&payload.report_kind,
"stream",
payload.report_context.as_ref(),
@@ -545,7 +531,6 @@ mod tests {
submit_sync_report(
&state,
"trace-reporting-sync-123",
GatewaySyncReportRequest {
trace_id: "trace-reporting-sync-123".to_string(),
report_kind: "openai_chat_sync_success".to_string(),
@@ -586,7 +571,6 @@ mod tests {
submit_sync_report(
&state,
"trace-reporting-sync-null-1",
GatewaySyncReportRequest {
trace_id: "trace-reporting-sync-null-1".to_string(),
report_kind: "claude_cli_sync_success".to_string(),
@@ -629,7 +613,6 @@ mod tests {
submit_stream_report(
&state,
"trace-reporting-stream-123",
GatewayStreamReportRequest {
trace_id: "trace-reporting-stream-123".to_string(),
report_kind: "openai_chat_stream_success".to_string(),
@@ -682,7 +665,6 @@ mod tests {
submit_sync_report(
&state,
"trace-codex-reporting-sync",
GatewaySyncReportRequest {
trace_id: "trace-codex-reporting-sync".to_string(),
report_kind: "openai_cli_sync_success".to_string(),
@@ -746,7 +728,6 @@ mod tests {
submit_stream_report(
&state,
"trace-codex-reporting-stream",
GatewayStreamReportRequest {
trace_id: "trace-codex-reporting-stream".to_string(),
report_kind: "openai_cli_stream_success".to_string(),
@@ -812,7 +793,6 @@ mod tests {
submit_sync_report(
&state,
"trace-gemini-files-store-123",
GatewaySyncReportRequest {
trace_id: "trace-gemini-files-store-123".to_string(),
report_kind: "gemini_files_store_mapping".to_string(),
@@ -883,7 +863,6 @@ mod tests {
submit_sync_report(
&state,
"trace-gemini-files-delete-123",
GatewaySyncReportRequest {
trace_id: "trace-gemini-files-delete-123".to_string(),
report_kind: "gemini_files_delete_mapping".to_string(),
@@ -926,7 +905,6 @@ mod tests {
submit_sync_report(
&state,
"trace-reporting-video-delete-123",
GatewaySyncReportRequest {
trace_id: "trace-reporting-video-delete-123".to_string(),
report_kind: "openai_video_delete_sync_success".to_string(),
@@ -991,7 +969,6 @@ mod tests {
submit_sync_report(
&state,
"trace-openai-video-reporting-123",
GatewaySyncReportRequest {
trace_id: "trace-openai-video-reporting-123".to_string(),
report_kind: "openai_video_create_sync_success".to_string(),
@@ -1053,7 +1030,6 @@ mod tests {
submit_sync_report(
&state,
"trace-gemini-video-reporting-123",
GatewaySyncReportRequest {
trace_id: "trace-gemini-video-reporting-123".to_string(),
report_kind: "gemini_video_create_sync_success".to_string(),
@@ -1103,7 +1079,6 @@ mod tests {
submit_sync_report(
&state,
"trace-openai-video-task-id-123",
GatewaySyncReportRequest {
trace_id: "trace-openai-video-task-id-123".to_string(),
report_kind: "openai_video_cancel_sync_success".to_string(),
@@ -1206,7 +1181,6 @@ mod tests {
submit_sync_report(
&state,
"trace-gemini-video-external-id-123",
GatewaySyncReportRequest {
trace_id: "trace-gemini-video-external-id-123".to_string(),
report_kind: "gemini_video_cancel_sync_success".to_string(),