mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-10 03:09:50 +08:00
refactor: 优化调度候选排序与用量写入链路并改进 Fernet 缓存与前端批量列表
This commit is contained in:
@@ -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()),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
+1
-1
@@ -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(),
|
||||
|
||||
+27
-14
@@ -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>,
|
||||
|
||||
+4
-2
@@ -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,
|
||||
|
||||
+63
-49
@@ -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(),
|
||||
},
|
||||
))
|
||||
|
||||
+4
-3
@@ -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),
|
||||
})
|
||||
}
|
||||
|
||||
+1
-1
@@ -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(),
|
||||
|
||||
+64
-49
@@ -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(),
|
||||
|
||||
+69
-151
@@ -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!(
|
||||
|
||||
+69
-152
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user