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
@@ -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()))
}
}
@@ -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()));
@@ -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,
)
}
@@ -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)
}
@@ -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()),
}
}
+26 -5
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 =
@@ -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 {
@@ -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(
@@ -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,
});
}
@@ -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,
@@ -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)?
@@ -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()))
})
}
@@ -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())
}
@@ -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(),
@@ -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(),
},
@@ -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>,
@@ -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(
@@ -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())
}
@@ -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,
},
},
}),
}
}
}
@@ -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(
@@ -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(),
},
))
@@ -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,
})
@@ -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,
@@ -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(),
},
))
@@ -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,
@@ -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(),
@@ -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(),
@@ -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(),
},
))
@@ -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),
})
}
@@ -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,
@@ -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(),
},
))
@@ -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),
})
}
@@ -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(),
@@ -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(),
},
))
@@ -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),
})
}
@@ -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(),
@@ -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!(
@@ -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!(
@@ -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,
@@ -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,