mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 09:50:21 +08:00
refactor: 优化调度候选排序与用量写入链路并改进 Fernet 缓存与前端批量列表
This commit is contained in:
@@ -11,15 +11,15 @@ enum ProviderPrivateStreamNormalizeMode {
|
|||||||
KiroToClaudeCli(KiroToClaudeCliStreamState),
|
KiroToClaudeCli(KiroToClaudeCliStreamState),
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) struct ProviderPrivateStreamNormalizer {
|
pub(crate) struct ProviderPrivateStreamNormalizer<'a> {
|
||||||
report_context: Value,
|
report_context: &'a Value,
|
||||||
buffered: Vec<u8>,
|
buffered: Vec<u8>,
|
||||||
mode: ProviderPrivateStreamNormalizeMode,
|
mode: ProviderPrivateStreamNormalizeMode,
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn maybe_build_provider_private_stream_normalizer(
|
pub(crate) fn maybe_build_provider_private_stream_normalizer<'a>(
|
||||||
report_context: Option<&Value>,
|
report_context: Option<&'a Value>,
|
||||||
) -> Option<ProviderPrivateStreamNormalizer> {
|
) -> Option<ProviderPrivateStreamNormalizer<'a>> {
|
||||||
let report_context = report_context?;
|
let report_context = report_context?;
|
||||||
if !report_context
|
if !report_context
|
||||||
.get("has_envelope")
|
.get("has_envelope")
|
||||||
@@ -51,17 +51,17 @@ pub(crate) fn maybe_build_provider_private_stream_normalizer(
|
|||||||
return None;
|
return None;
|
||||||
};
|
};
|
||||||
Some(ProviderPrivateStreamNormalizer {
|
Some(ProviderPrivateStreamNormalizer {
|
||||||
report_context: report_context.clone(),
|
report_context,
|
||||||
buffered: Vec::new(),
|
buffered: Vec::new(),
|
||||||
mode,
|
mode,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
impl ProviderPrivateStreamNormalizer {
|
impl ProviderPrivateStreamNormalizer<'_> {
|
||||||
pub(crate) fn push_chunk(&mut self, chunk: &[u8]) -> Result<Vec<u8>, GatewayError> {
|
pub(crate) fn push_chunk(&mut self, chunk: &[u8]) -> Result<Vec<u8>, GatewayError> {
|
||||||
match &mut self.mode {
|
match &mut self.mode {
|
||||||
ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => {
|
ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => {
|
||||||
state.push_chunk(&self.report_context, chunk)
|
state.push_chunk(self.report_context, chunk)
|
||||||
}
|
}
|
||||||
ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => {
|
ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => {
|
||||||
self.buffered.extend_from_slice(chunk);
|
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') {
|
while let Some(line_end) = self.buffered.iter().position(|byte| *byte == b'\n') {
|
||||||
let line = self.buffered.drain(..=line_end).collect::<Vec<_>>();
|
let line = self.buffered.drain(..=line_end).collect::<Vec<_>>();
|
||||||
output.extend(
|
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()))?,
|
.map_err(|err| GatewayError::Internal(err.to_string()))?,
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
@@ -81,14 +81,14 @@ impl ProviderPrivateStreamNormalizer {
|
|||||||
pub(crate) fn finish(&mut self) -> Result<Vec<u8>, GatewayError> {
|
pub(crate) fn finish(&mut self) -> Result<Vec<u8>, GatewayError> {
|
||||||
match &mut self.mode {
|
match &mut self.mode {
|
||||||
ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => {
|
ProviderPrivateStreamNormalizeMode::KiroToClaudeCli(state) => {
|
||||||
state.finish(&self.report_context)
|
state.finish(self.report_context)
|
||||||
}
|
}
|
||||||
ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => {
|
ProviderPrivateStreamNormalizeMode::EnvelopeUnwrap => {
|
||||||
if self.buffered.is_empty() {
|
if self.buffered.is_empty() {
|
||||||
return Ok(Vec::new());
|
return Ok(Vec::new());
|
||||||
}
|
}
|
||||||
let line = std::mem::take(&mut self.buffered);
|
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()))
|
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -60,7 +60,7 @@ fn normalize_provider_private_stream_bytes(
|
|||||||
report_context: &Value,
|
report_context: &Value,
|
||||||
body: &[u8],
|
body: &[u8],
|
||||||
) -> Result<Option<Vec<u8>>, GatewayError> {
|
) -> 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))
|
maybe_build_provider_private_stream_normalizer(Some(report_context))
|
||||||
else {
|
else {
|
||||||
return Ok(Some(body.to_vec()));
|
return Ok(Some(body.to_vec()));
|
||||||
|
|||||||
@@ -24,21 +24,40 @@ pub(crate) struct LocalCoreSyncFinalizeOutcome {
|
|||||||
pub(crate) background_report: Option<GatewaySyncReportRequest>,
|
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(
|
pub(crate) fn build_local_success_outcome(
|
||||||
trace_id: &str,
|
trace_id: &str,
|
||||||
decision: &GatewayControlDecision,
|
decision: &GatewayControlDecision,
|
||||||
payload: &GatewaySyncReportRequest,
|
payload: &GatewaySyncReportRequest,
|
||||||
body_json: Value,
|
body_json: Value,
|
||||||
) -> Result<LocalCoreSyncFinalizeOutcome, GatewayError> {
|
) -> 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 =
|
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(
|
build_local_success_outcome_with_report(
|
||||||
trace_id,
|
trace_id,
|
||||||
decision,
|
decision,
|
||||||
payload.status_code,
|
payload.status_code,
|
||||||
body_json,
|
body_bytes,
|
||||||
headers,
|
response_headers,
|
||||||
background_report,
|
background_report,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
@@ -47,19 +66,12 @@ pub(crate) fn build_local_success_outcome_with_report(
|
|||||||
trace_id: &str,
|
trace_id: &str,
|
||||||
decision: &GatewayControlDecision,
|
decision: &GatewayControlDecision,
|
||||||
status_code: u16,
|
status_code: u16,
|
||||||
body_json: Value,
|
body_bytes: Vec<u8>,
|
||||||
headers: BTreeMap<String, String>,
|
headers: BTreeMap<String, String>,
|
||||||
background_report: Option<GatewaySyncReportRequest>,
|
background_report: Option<GatewaySyncReportRequest>,
|
||||||
) -> Result<LocalCoreSyncFinalizeOutcome, GatewayError> {
|
) -> Result<LocalCoreSyncFinalizeOutcome, GatewayError> {
|
||||||
let (body_bytes, headers) = prepare_local_success_response_parts_impl(&headers, &body_json)
|
let response =
|
||||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
build_local_success_response(trace_id, decision, status_code, body_bytes, headers)?;
|
||||||
let response = build_client_response_from_parts(
|
|
||||||
status_code,
|
|
||||||
&headers,
|
|
||||||
Body::from(body_bytes),
|
|
||||||
trace_id,
|
|
||||||
Some(decision),
|
|
||||||
)?;
|
|
||||||
Ok(LocalCoreSyncFinalizeOutcome {
|
Ok(LocalCoreSyncFinalizeOutcome {
|
||||||
response,
|
response,
|
||||||
background_report,
|
background_report,
|
||||||
@@ -73,9 +85,12 @@ pub(crate) fn build_local_success_outcome_with_conversion_report(
|
|||||||
client_body_json: Value,
|
client_body_json: Value,
|
||||||
provider_body_json: Value,
|
provider_body_json: Value,
|
||||||
) -> Result<LocalCoreSyncFinalizeOutcome, GatewayError> {
|
) -> 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(
|
let report_payload = build_local_success_conversion_background_report_impl(
|
||||||
payload,
|
payload,
|
||||||
client_body_json.clone(),
|
client_body_json,
|
||||||
provider_body_json,
|
provider_body_json,
|
||||||
);
|
);
|
||||||
|
|
||||||
@@ -83,8 +98,8 @@ pub(crate) fn build_local_success_outcome_with_conversion_report(
|
|||||||
trace_id,
|
trace_id,
|
||||||
decision,
|
decision,
|
||||||
payload.status_code,
|
payload.status_code,
|
||||||
client_body_json,
|
body_bytes,
|
||||||
payload.headers.clone(),
|
response_headers,
|
||||||
report_payload,
|
report_payload,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -34,6 +34,6 @@ pub(crate) fn maybe_compile_sync_finalize_response(
|
|||||||
|
|
||||||
pub(crate) fn maybe_build_stream_response_rewriter(
|
pub(crate) fn maybe_build_stream_response_rewriter(
|
||||||
report_context: Option<&Value>,
|
report_context: Option<&Value>,
|
||||||
) -> Option<LocalStreamRewriter> {
|
) -> Option<LocalStreamRewriter<'_>> {
|
||||||
stream::maybe_build_local_stream_rewriter(report_context)
|
stream::maybe_build_local_stream_rewriter(report_context)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -12,15 +12,15 @@ enum RewriteMode {
|
|||||||
KiroToClaudeCli(KiroToClaudeCliStreamState),
|
KiroToClaudeCli(KiroToClaudeCliStreamState),
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) struct LocalStreamRewriter {
|
pub(crate) struct LocalStreamRewriter<'a> {
|
||||||
report_context: Value,
|
report_context: &'a Value,
|
||||||
buffered: Vec<u8>,
|
buffered: Vec<u8>,
|
||||||
mode: RewriteMode,
|
mode: RewriteMode,
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn maybe_build_local_stream_rewriter(
|
pub(crate) fn maybe_build_local_stream_rewriter<'a>(
|
||||||
report_context: Option<&Value>,
|
report_context: Option<&'a Value>,
|
||||||
) -> Option<LocalStreamRewriter> {
|
) -> Option<LocalStreamRewriter<'a>> {
|
||||||
let report_context = report_context?;
|
let report_context = report_context?;
|
||||||
let mode = match resolve_finalize_stream_rewrite_mode(report_context)? {
|
let mode = match resolve_finalize_stream_rewrite_mode(report_context)? {
|
||||||
FinalizeStreamRewriteMode::EnvelopeUnwrap => RewriteMode::EnvelopeUnwrap,
|
FinalizeStreamRewriteMode::EnvelopeUnwrap => RewriteMode::EnvelopeUnwrap,
|
||||||
@@ -33,16 +33,16 @@ pub(crate) fn maybe_build_local_stream_rewriter(
|
|||||||
};
|
};
|
||||||
|
|
||||||
Some(LocalStreamRewriter {
|
Some(LocalStreamRewriter {
|
||||||
report_context: report_context.clone(),
|
report_context,
|
||||||
buffered: Vec::new(),
|
buffered: Vec::new(),
|
||||||
mode,
|
mode,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
impl LocalStreamRewriter {
|
impl LocalStreamRewriter<'_> {
|
||||||
pub(crate) fn push_chunk(&mut self, chunk: &[u8]) -> Result<Vec<u8>, GatewayError> {
|
pub(crate) fn push_chunk(&mut self, chunk: &[u8]) -> Result<Vec<u8>, GatewayError> {
|
||||||
if let RewriteMode::KiroToClaudeCli(state) = &mut self.mode {
|
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);
|
self.buffered.extend_from_slice(chunk);
|
||||||
let mut output = Vec::new();
|
let mut output = Vec::new();
|
||||||
@@ -55,11 +55,11 @@ impl LocalStreamRewriter {
|
|||||||
|
|
||||||
pub(crate) fn finish(&mut self) -> Result<Vec<u8>, GatewayError> {
|
pub(crate) fn finish(&mut self) -> Result<Vec<u8>, GatewayError> {
|
||||||
if let RewriteMode::KiroToClaudeCli(state) = &mut self.mode {
|
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() {
|
if self.buffered.is_empty() {
|
||||||
match &mut self.mode {
|
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::KiroToClaudeCli(_) => {}
|
||||||
RewriteMode::EnvelopeUnwrap => {}
|
RewriteMode::EnvelopeUnwrap => {}
|
||||||
}
|
}
|
||||||
@@ -69,7 +69,7 @@ impl LocalStreamRewriter {
|
|||||||
let mut output = self.transform_line(line)?;
|
let mut output = self.transform_line(line)?;
|
||||||
match &mut self.mode {
|
match &mut self.mode {
|
||||||
RewriteMode::Standard(state) => {
|
RewriteMode::Standard(state) => {
|
||||||
output.extend(state.finish(&self.report_context)?);
|
output.extend(state.finish(self.report_context)?);
|
||||||
}
|
}
|
||||||
RewriteMode::KiroToClaudeCli(_) => {}
|
RewriteMode::KiroToClaudeCli(_) => {}
|
||||||
RewriteMode::EnvelopeUnwrap => {}
|
RewriteMode::EnvelopeUnwrap => {}
|
||||||
@@ -79,9 +79,9 @@ impl LocalStreamRewriter {
|
|||||||
|
|
||||||
fn transform_line(&mut self, line: Vec<u8>) -> Result<Vec<u8>, GatewayError> {
|
fn transform_line(&mut self, line: Vec<u8>) -> Result<Vec<u8>, GatewayError> {
|
||||||
match &mut self.mode {
|
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())),
|
.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()),
|
RewriteMode::KiroToClaudeCli(_) => Ok(Vec::new()),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ pub(crate) mod transport;
|
|||||||
|
|
||||||
use axum::body::Body;
|
use axum::body::Body;
|
||||||
use axum::http::{Response, Uri};
|
use axum::http::{Response, Uri};
|
||||||
use serde_json::Value;
|
use serde_json::{json, Value};
|
||||||
|
|
||||||
use crate::{usage::GatewaySyncReportRequest, AppState, GatewayError};
|
use crate::{usage::GatewaySyncReportRequest, AppState, GatewayError};
|
||||||
|
|
||||||
@@ -62,8 +62,18 @@ pub(crate) fn collect_control_headers(
|
|||||||
crate::headers::collect_control_headers(headers)
|
crate::headers::collect_control_headers(headers)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn build_report_context_original_request_echo(body_json: &Value) -> Option<Value> {
|
pub(crate) fn build_report_context_original_request_echo(
|
||||||
(!body_json.is_null()).then(|| body_json.clone())
|
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 {
|
pub(crate) fn is_json_request(headers: &http::HeaderMap) -> bool {
|
||||||
@@ -138,12 +148,23 @@ mod tests {
|
|||||||
"body_bytes_b64": "aGVsbG8=",
|
"body_bytes_b64": "aGVsbG8=",
|
||||||
});
|
});
|
||||||
|
|
||||||
let echo =
|
let echo = build_report_context_original_request_echo(Some(&body), None)
|
||||||
build_report_context_original_request_echo(&body).expect("echo should be produced");
|
.expect("echo should be produced");
|
||||||
|
|
||||||
assert_eq!(echo, body);
|
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]
|
#[test]
|
||||||
fn extract_gemini_model_from_path_trims_method_suffix() {
|
fn extract_gemini_model_from_path_trims_method_suffix() {
|
||||||
let model =
|
let model =
|
||||||
|
|||||||
@@ -1,3 +1,5 @@
|
|||||||
|
use std::collections::BTreeMap;
|
||||||
|
|
||||||
use tracing::warn;
|
use tracing::warn;
|
||||||
|
|
||||||
use crate::ai_pipeline::{
|
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 crate::scheduler::config::{read_scheduler_ordering_config, SchedulerOrderingConfig};
|
||||||
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
|
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
|
||||||
use aether_scheduler_core::{
|
use aether_scheduler_core::{
|
||||||
build_scheduler_affinity_cache_key_for_api_key_id, compare_candidates_by_priority_mode,
|
build_scheduler_affinity_cache_key_for_api_key_id, requested_capability_priority_for_candidate,
|
||||||
requested_capability_priority_for_candidate, SchedulerAffinityTarget, SchedulerPriorityMode,
|
SchedulerAffinityTarget, SchedulerPriorityMode,
|
||||||
};
|
};
|
||||||
|
|
||||||
use super::candidate_eligibility::{
|
use super::candidate_eligibility::{
|
||||||
@@ -18,6 +20,8 @@ use super::candidate_eligibility::{
|
|||||||
|
|
||||||
const PLANNER_SCHEDULER_AFFINITY_MAX_ENTRIES: usize = 10_000;
|
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)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
|
||||||
enum TunnelOwnerAffinityBucket {
|
enum TunnelOwnerAffinityBucket {
|
||||||
LocalTunnel = 0,
|
LocalTunnel = 0,
|
||||||
@@ -31,20 +35,36 @@ struct CandidateExecutionOrdering {
|
|||||||
keep_priority_on_conversion: bool,
|
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(
|
pub(crate) async fn prefer_local_tunnel_owner_candidates(
|
||||||
state: PlannerAppState<'_>,
|
state: PlannerAppState<'_>,
|
||||||
candidates: Vec<SchedulerMinimalCandidateSelectionCandidate>,
|
candidates: Vec<SchedulerMinimalCandidateSelectionCandidate>,
|
||||||
) -> Vec<SchedulerMinimalCandidateSelectionCandidate> {
|
) -> Vec<SchedulerMinimalCandidateSelectionCandidate> {
|
||||||
let mut ranked = Vec::with_capacity(candidates.len());
|
let mut candidates = candidates;
|
||||||
for (original_index, candidate) in candidates.into_iter().enumerate() {
|
let mut rankings = Vec::with_capacity(candidates.len());
|
||||||
let bucket = resolve_candidate_tunnel_owner_affinity(state, &candidate).await;
|
let mut tunnel_affinity_cache = BTreeMap::new();
|
||||||
ranked.push((bucket, original_index, candidate));
|
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)));
|
let mut order = (0..candidates.len()).collect::<Vec<_>>();
|
||||||
ranked
|
order.sort_by(|left, right| rankings[*left].cmp(&rankings[*right]));
|
||||||
.into_iter()
|
drop(tunnel_affinity_cache);
|
||||||
.map(|(_, _, candidate)| candidate)
|
apply_order(&mut candidates, order);
|
||||||
.collect()
|
candidates
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
@@ -56,11 +76,18 @@ async fn rank_local_execution_candidates(
|
|||||||
) -> Vec<SchedulerMinimalCandidateSelectionCandidate> {
|
) -> Vec<SchedulerMinimalCandidateSelectionCandidate> {
|
||||||
let normalized_client_api_format = client_api_format.trim().to_ascii_lowercase();
|
let normalized_client_api_format = client_api_format.trim().to_ascii_lowercase();
|
||||||
let ordering_config = read_scheduler_ordering_config_or_default(state).await;
|
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() {
|
for (original_index, candidate) in candidates.iter().enumerate() {
|
||||||
let ordering =
|
let ordering = resolve_cached_candidate_execution_ordering(
|
||||||
resolve_candidate_execution_ordering(state, &candidate, ordering_config).await;
|
state,
|
||||||
|
&mut ordering_cache,
|
||||||
|
candidate,
|
||||||
|
ordering_config,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
let is_same_format = candidate
|
let is_same_format = candidate
|
||||||
.endpoint_api_format
|
.endpoint_api_format
|
||||||
.trim()
|
.trim()
|
||||||
@@ -71,112 +98,82 @@ async fn rank_local_execution_candidates(
|
|||||||
candidate.endpoint_api_format.as_str(),
|
candidate.endpoint_api_format.as_str(),
|
||||||
);
|
);
|
||||||
let capability_priority =
|
let capability_priority =
|
||||||
requested_capability_priority_for_candidate(required_capabilities, &candidate);
|
requested_capability_priority_for_candidate(required_capabilities, candidate);
|
||||||
ranked.push((
|
rankings.push(PlannerCandidateRankingState {
|
||||||
capability_priority.0,
|
capability_priority,
|
||||||
capability_priority.1,
|
tunnel_bucket: ordering.tunnel_bucket,
|
||||||
ordering.tunnel_bucket,
|
|
||||||
demote_cross_format,
|
demote_cross_format,
|
||||||
format_preference,
|
format_preference,
|
||||||
original_index,
|
original_index,
|
||||||
candidate,
|
});
|
||||||
));
|
|
||||||
}
|
}
|
||||||
|
|
||||||
ranked.sort_by(|left, right| {
|
let mut order = (0..candidates.len()).collect::<Vec<_>>();
|
||||||
left.0
|
order.sort_by(|left, right| {
|
||||||
.cmp(&right.0)
|
compare_planner_candidate_ranking(
|
||||||
.then(left.1.cmp(&right.1))
|
&rankings[*left],
|
||||||
.then(left.2.cmp(&right.2))
|
&candidates[*left],
|
||||||
.then(left.3.cmp(&right.3))
|
&rankings[*right],
|
||||||
.then_with(|| {
|
&candidates[*right],
|
||||||
compare_candidate_priority_slot(&left.6, &right.6, ordering_config.priority_mode)
|
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))
|
|
||||||
});
|
});
|
||||||
|
drop(ordering_cache);
|
||||||
ranked
|
apply_order(&mut candidates, order);
|
||||||
.into_iter()
|
candidates
|
||||||
.map(|(_, _, _, _, _, _, candidate)| candidate)
|
|
||||||
.collect()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) async fn rank_eligible_local_execution_candidates(
|
pub(crate) async fn rank_eligible_local_execution_candidates(
|
||||||
state: PlannerAppState<'_>,
|
state: PlannerAppState<'_>,
|
||||||
candidates: Vec<EligibleLocalExecutionCandidate>,
|
candidates: Vec<EligibleLocalExecutionCandidate>,
|
||||||
client_api_format: &str,
|
normalized_client_api_format: &str,
|
||||||
required_capabilities: Option<&serde_json::Value>,
|
required_capabilities: Option<&serde_json::Value>,
|
||||||
) -> Vec<EligibleLocalExecutionCandidate> {
|
) -> 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 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() {
|
for (original_index, eligible) in candidates.iter().enumerate() {
|
||||||
let ordering = resolve_candidate_execution_ordering_from_transport(
|
let ordering = resolve_cached_eligible_candidate_execution_ordering(
|
||||||
state,
|
state,
|
||||||
&eligible.transport,
|
&mut ordering_cache,
|
||||||
|
eligible,
|
||||||
ordering_config,
|
ordering_config,
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
let is_same_format = eligible
|
let is_same_format = eligible
|
||||||
.provider_api_format
|
.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 demote_cross_format = !is_same_format && !ordering.keep_priority_on_conversion;
|
||||||
let format_preference = candidate_api_format_preference(
|
let format_preference = candidate_api_format_preference(
|
||||||
normalized_client_api_format.as_str(),
|
normalized_client_api_format,
|
||||||
eligible.provider_api_format.as_str(),
|
eligible.provider_api_format.as_str(),
|
||||||
);
|
);
|
||||||
let capability_priority =
|
let capability_priority =
|
||||||
requested_capability_priority_for_candidate(required_capabilities, &eligible.candidate);
|
requested_capability_priority_for_candidate(required_capabilities, &eligible.candidate);
|
||||||
ranked.push((
|
rankings.push(PlannerCandidateRankingState {
|
||||||
capability_priority.0,
|
capability_priority,
|
||||||
capability_priority.1,
|
tunnel_bucket: ordering.tunnel_bucket,
|
||||||
ordering.tunnel_bucket,
|
|
||||||
demote_cross_format,
|
demote_cross_format,
|
||||||
format_preference,
|
format_preference,
|
||||||
original_index,
|
original_index,
|
||||||
eligible,
|
});
|
||||||
));
|
|
||||||
}
|
}
|
||||||
|
|
||||||
ranked.sort_by(|left, right| {
|
let mut order = (0..candidates.len()).collect::<Vec<_>>();
|
||||||
left.0
|
order.sort_by(|left, right| {
|
||||||
.cmp(&right.0)
|
compare_planner_candidate_ranking(
|
||||||
.then(left.1.cmp(&right.1))
|
&rankings[*left],
|
||||||
.then(left.2.cmp(&right.2))
|
&candidates[*left].candidate,
|
||||||
.then(left.3.cmp(&right.3))
|
&rankings[*right],
|
||||||
.then_with(|| {
|
&candidates[*right].candidate,
|
||||||
compare_candidate_priority_slot(
|
ordering_config.priority_mode,
|
||||||
&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))
|
|
||||||
});
|
});
|
||||||
|
drop(ordering_cache);
|
||||||
ranked
|
apply_order(&mut candidates, order);
|
||||||
.into_iter()
|
candidates
|
||||||
.map(|(_, _, _, _, _, _, eligible)| eligible)
|
|
||||||
.collect()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn remember_scheduler_affinity_for_candidate(
|
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
|
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(
|
async fn resolve_candidate_execution_ordering(
|
||||||
state: PlannerAppState<'_>,
|
state: PlannerAppState<'_>,
|
||||||
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
||||||
@@ -238,6 +250,43 @@ async fn resolve_candidate_execution_ordering(
|
|||||||
resolve_candidate_execution_ordering_from_transport(state, &transport, ordering_config).await
|
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(
|
async fn resolve_candidate_execution_ordering_from_transport(
|
||||||
state: PlannerAppState<'_>,
|
state: PlannerAppState<'_>,
|
||||||
transport: &GatewayProviderTransportSnapshot,
|
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) {
|
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)
|
request_candidate_api_format_preference(client_api_format, provider_api_format)
|
||||||
.unwrap_or((u8::MAX, u8::MAX))
|
.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(
|
async fn read_scheduler_ordering_config_or_default(
|
||||||
state: PlannerAppState<'_>,
|
state: PlannerAppState<'_>,
|
||||||
) -> SchedulerOrderingConfig {
|
) -> SchedulerOrderingConfig {
|
||||||
|
|||||||
@@ -1,3 +1,5 @@
|
|||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
use tracing::warn;
|
use tracing::warn;
|
||||||
|
|
||||||
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
|
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
|
||||||
@@ -11,7 +13,7 @@ use super::pool_scheduler::apply_local_execution_pool_scheduler;
|
|||||||
#[derive(Debug, Clone, PartialEq)]
|
#[derive(Debug, Clone, PartialEq)]
|
||||||
pub(crate) struct EligibleLocalExecutionCandidate {
|
pub(crate) struct EligibleLocalExecutionCandidate {
|
||||||
pub(crate) candidate: SchedulerMinimalCandidateSelectionCandidate,
|
pub(crate) candidate: SchedulerMinimalCandidateSelectionCandidate,
|
||||||
pub(crate) transport: GatewayProviderTransportSnapshot,
|
pub(crate) transport: Arc<GatewayProviderTransportSnapshot>,
|
||||||
pub(crate) provider_api_format: String,
|
pub(crate) provider_api_format: String,
|
||||||
pub(crate) orchestration: LocalExecutionCandidateMetadata,
|
pub(crate) orchestration: LocalExecutionCandidateMetadata,
|
||||||
}
|
}
|
||||||
@@ -20,10 +22,16 @@ pub(crate) struct EligibleLocalExecutionCandidate {
|
|||||||
pub(crate) struct SkippedLocalExecutionCandidate {
|
pub(crate) struct SkippedLocalExecutionCandidate {
|
||||||
pub(crate) candidate: SchedulerMinimalCandidateSelectionCandidate,
|
pub(crate) candidate: SchedulerMinimalCandidateSelectionCandidate,
|
||||||
pub(crate) skip_reason: &'static str,
|
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>,
|
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(
|
pub(crate) async fn filter_and_rank_local_execution_candidates(
|
||||||
state: PlannerAppState<'_>,
|
state: PlannerAppState<'_>,
|
||||||
candidates: Vec<SchedulerMinimalCandidateSelectionCandidate>,
|
candidates: Vec<SchedulerMinimalCandidateSelectionCandidate>,
|
||||||
@@ -35,17 +43,18 @@ pub(crate) async fn filter_and_rank_local_execution_candidates(
|
|||||||
Vec<EligibleLocalExecutionCandidate>,
|
Vec<EligibleLocalExecutionCandidate>,
|
||||||
Vec<SkippedLocalExecutionCandidate>,
|
Vec<SkippedLocalExecutionCandidate>,
|
||||||
) {
|
) {
|
||||||
|
let requested_model = requested_model.trim();
|
||||||
filter_and_rank_local_execution_candidates_with_gate(
|
filter_and_rank_local_execution_candidates_with_gate(
|
||||||
state,
|
state,
|
||||||
candidates,
|
candidates,
|
||||||
client_api_format,
|
client_api_format,
|
||||||
required_capabilities,
|
required_capabilities,
|
||||||
sticky_session_token,
|
sticky_session_token,
|
||||||
|candidate, transport| {
|
|candidate, transport, normalized_client_api_format| {
|
||||||
current_local_execution_candidate_skip_reason_with_transport(
|
current_local_execution_candidate_skip_reason_with_transport(
|
||||||
candidate,
|
candidate,
|
||||||
transport,
|
transport,
|
||||||
client_api_format,
|
normalized_client_api_format,
|
||||||
requested_model,
|
requested_model,
|
||||||
)
|
)
|
||||||
},
|
},
|
||||||
@@ -64,13 +73,14 @@ pub(crate) async fn filter_and_rank_local_execution_candidates_without_transport
|
|||||||
Vec<EligibleLocalExecutionCandidate>,
|
Vec<EligibleLocalExecutionCandidate>,
|
||||||
Vec<SkippedLocalExecutionCandidate>,
|
Vec<SkippedLocalExecutionCandidate>,
|
||||||
) {
|
) {
|
||||||
|
let requested_model = requested_model.map(str::trim);
|
||||||
filter_and_rank_local_execution_candidates_with_gate(
|
filter_and_rank_local_execution_candidates_with_gate(
|
||||||
state,
|
state,
|
||||||
candidates,
|
candidates,
|
||||||
client_api_format,
|
client_api_format,
|
||||||
required_capabilities,
|
required_capabilities,
|
||||||
sticky_session_token,
|
sticky_session_token,
|
||||||
|candidate, transport| {
|
|candidate, transport, _normalized_client_api_format| {
|
||||||
current_local_execution_candidate_common_skip_reason_with_transport(
|
current_local_execution_candidate_common_skip_reason_with_transport(
|
||||||
candidate,
|
candidate,
|
||||||
transport,
|
transport,
|
||||||
@@ -96,10 +106,12 @@ where
|
|||||||
F: Fn(
|
F: Fn(
|
||||||
&SchedulerMinimalCandidateSelectionCandidate,
|
&SchedulerMinimalCandidateSelectionCandidate,
|
||||||
&GatewayProviderTransportSnapshot,
|
&GatewayProviderTransportSnapshot,
|
||||||
|
&str,
|
||||||
) -> Option<&'static 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 selectable = Vec::with_capacity(candidates.len());
|
||||||
let mut skipped = Vec::new();
|
let mut skipped = Vec::with_capacity(candidates.len());
|
||||||
|
|
||||||
for candidate in candidates {
|
for candidate in candidates {
|
||||||
let Some(transport) = read_candidate_transport_snapshot(state, &candidate).await else {
|
let Some(transport) = read_candidate_transport_snapshot(state, &candidate).await else {
|
||||||
@@ -111,11 +123,18 @@ where
|
|||||||
});
|
});
|
||||||
continue;
|
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;
|
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 {
|
Some(skip_reason) => skipped.push(SkippedLocalExecutionCandidate {
|
||||||
candidate,
|
candidate,
|
||||||
skip_reason,
|
skip_reason,
|
||||||
@@ -134,7 +153,7 @@ where
|
|||||||
let ranked = rank_eligible_local_execution_candidates(
|
let ranked = rank_eligible_local_execution_candidates(
|
||||||
state,
|
state,
|
||||||
selectable,
|
selectable,
|
||||||
client_api_format,
|
normalized_client_api_format.as_str(),
|
||||||
required_capabilities,
|
required_capabilities,
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
@@ -146,28 +165,27 @@ where
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn extract_pool_sticky_session_token(body_json: &serde_json::Value) -> Option<String> {
|
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
|
value
|
||||||
.and_then(serde_json::Value::as_str)
|
.and_then(serde_json::Value::as_str)
|
||||||
.map(str::trim)
|
.map(str::trim)
|
||||||
.filter(|value| !value.is_empty())
|
.filter(|value| !value.is_empty())
|
||||||
.map(ToOwned::to_owned)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
let object = body_json.as_object()?;
|
let object = body_json.as_object()?;
|
||||||
|
|
||||||
non_empty_string(object.get("prompt_cache_key"))
|
non_empty_str(object.get("prompt_cache_key"))
|
||||||
.or_else(|| non_empty_string(object.get("conversation_id")))
|
.or_else(|| non_empty_str(object.get("conversation_id")))
|
||||||
.or_else(|| non_empty_string(object.get("conversationId")))
|
.or_else(|| non_empty_str(object.get("conversationId")))
|
||||||
.or_else(|| non_empty_string(object.get("session_id")))
|
.or_else(|| non_empty_str(object.get("session_id")))
|
||||||
.or_else(|| non_empty_string(object.get("sessionId")))
|
.or_else(|| non_empty_str(object.get("sessionId")))
|
||||||
.or_else(|| {
|
.or_else(|| {
|
||||||
object
|
object
|
||||||
.get("metadata")
|
.get("metadata")
|
||||||
.and_then(serde_json::Value::as_object)
|
.and_then(serde_json::Value::as_object)
|
||||||
.and_then(|metadata| {
|
.and_then(|metadata| {
|
||||||
non_empty_string(metadata.get("session_id"))
|
non_empty_str(metadata.get("session_id"))
|
||||||
.or_else(|| non_empty_string(metadata.get("conversation_id")))
|
.or_else(|| non_empty_str(metadata.get("conversation_id")))
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
.or_else(|| {
|
.or_else(|| {
|
||||||
@@ -175,10 +193,11 @@ pub(crate) fn extract_pool_sticky_session_token(body_json: &serde_json::Value) -
|
|||||||
.get("conversationState")
|
.get("conversationState")
|
||||||
.and_then(serde_json::Value::as_object)
|
.and_then(serde_json::Value::as_object)
|
||||||
.and_then(|state| {
|
.and_then(|state| {
|
||||||
non_empty_string(state.get("conversationId"))
|
non_empty_str(state.get("conversationId"))
|
||||||
.or_else(|| non_empty_string(state.get("sessionId")))
|
.or_else(|| non_empty_str(state.get("sessionId")))
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
|
.map(ToOwned::to_owned)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn current_local_execution_candidate_common_skip_reason_with_transport(
|
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");
|
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();
|
||||||
let endpoint_api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
|
if !candidate
|
||||||
if endpoint_api_format != candidate_api_format {
|
.endpoint_api_format
|
||||||
|
.trim()
|
||||||
|
.eq_ignore_ascii_case(endpoint_api_format)
|
||||||
|
{
|
||||||
return Some("endpoint_api_format_changed");
|
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");
|
return Some("key_api_format_disabled");
|
||||||
}
|
}
|
||||||
if !transport_key_allows_candidate_model(transport, requested_model, candidate) {
|
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(
|
fn candidate_is_ineligible_due_to_disabled_format_conversion(
|
||||||
transport: &GatewayProviderTransportSnapshot,
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
client_api_format: &str,
|
normalized_client_api_format: &str,
|
||||||
) -> bool {
|
) -> bool {
|
||||||
let endpoint_api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
|
let endpoint_api_format = transport.endpoint.api_format.trim();
|
||||||
let client_api_format = client_api_format.trim().to_ascii_lowercase();
|
if endpoint_api_format.eq_ignore_ascii_case(normalized_client_api_format) {
|
||||||
if client_api_format == endpoint_api_format {
|
|
||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
|
|
||||||
crate::ai_pipeline::conversion::request_conversion_kind(
|
crate::ai_pipeline::conversion::request_conversion_kind(
|
||||||
client_api_format.as_str(),
|
normalized_client_api_format,
|
||||||
endpoint_api_format.as_str(),
|
endpoint_api_format,
|
||||||
)
|
)
|
||||||
.is_some()
|
.is_some()
|
||||||
&& crate::ai_pipeline::conversion::request_conversion_requires_enable_flag(
|
&& crate::ai_pipeline::conversion::request_conversion_requires_enable_flag(
|
||||||
client_api_format.as_str(),
|
normalized_client_api_format,
|
||||||
endpoint_api_format.as_str(),
|
endpoint_api_format,
|
||||||
)
|
)
|
||||||
&& !crate::ai_pipeline::conversion::request_conversion_enabled_for_transport(
|
&& !crate::ai_pipeline::conversion::request_conversion_enabled_for_transport(
|
||||||
transport,
|
transport,
|
||||||
client_api_format.as_str(),
|
normalized_client_api_format,
|
||||||
endpoint_api_format.as_str(),
|
endpoint_api_format,
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn current_local_execution_candidate_skip_reason_with_transport(
|
fn current_local_execution_candidate_skip_reason_with_transport(
|
||||||
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
||||||
transport: &GatewayProviderTransportSnapshot,
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
client_api_format: &str,
|
normalized_client_api_format: &str,
|
||||||
requested_model: &str,
|
requested_model: &str,
|
||||||
) -> Option<&'static str> {
|
) -> Option<&'static str> {
|
||||||
if let Some(skip_reason) = current_local_execution_candidate_common_skip_reason_with_transport(
|
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);
|
return Some(skip_reason);
|
||||||
}
|
}
|
||||||
|
|
||||||
let endpoint_api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
|
let endpoint_api_format = transport.endpoint.api_format.trim();
|
||||||
let client_api_format = client_api_format.trim().to_ascii_lowercase();
|
if endpoint_api_format.eq_ignore_ascii_case(normalized_client_api_format) {
|
||||||
if client_api_format == endpoint_api_format {
|
|
||||||
return None;
|
return None;
|
||||||
}
|
}
|
||||||
|
|
||||||
if !crate::ai_pipeline::conversion::request_pair_allowed_for_transport(
|
if !crate::ai_pipeline::conversion::request_pair_allowed_for_transport(
|
||||||
transport,
|
transport,
|
||||||
client_api_format.as_str(),
|
normalized_client_api_format,
|
||||||
endpoint_api_format.as_str(),
|
endpoint_api_format,
|
||||||
) {
|
) {
|
||||||
return Some("transport_unsupported");
|
return Some("transport_unsupported");
|
||||||
}
|
}
|
||||||
@@ -292,15 +312,6 @@ fn transport_key_allows_candidate_model(
|
|||||||
return true;
|
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 requested_model = requested_model.trim();
|
||||||
let global_model_name = candidate.global_model_name.trim();
|
let global_model_name = candidate.global_model_name.trim();
|
||||||
let selected_provider_model_name = candidate.selected_provider_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)
|
.map(str::trim)
|
||||||
.filter(|value| !value.is_empty());
|
.filter(|value| !value.is_empty());
|
||||||
|
|
||||||
allowed_models.iter().any(|allowed_model| {
|
for allowed_model in allowed_models.iter().map(String::as_str).map(str::trim) {
|
||||||
*allowed_model == requested_model
|
if allowed_model.is_empty() {
|
||||||
|| *allowed_model == global_model_name
|
continue;
|
||||||
|| *allowed_model == selected_provider_model_name
|
}
|
||||||
|| mapping_matched_model.is_some_and(|value| value == *allowed_model)
|
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(
|
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::planner::runtime_miss::record_local_runtime_candidate_skip_reason;
|
||||||
use crate::ai_pipeline::{GatewayAuthApiKeySnapshot, PlannerAppState};
|
use crate::ai_pipeline::{GatewayAuthApiKeySnapshot, PlannerAppState};
|
||||||
use crate::clock::current_unix_ms;
|
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;
|
use crate::AppState;
|
||||||
|
|
||||||
#[derive(Debug, Clone)]
|
#[derive(Debug, Clone)]
|
||||||
@@ -17,15 +17,13 @@ pub(crate) struct LocalExecutionCandidateAttempt {
|
|||||||
pub(crate) eligible: EligibleLocalExecutionCandidate,
|
pub(crate) eligible: EligibleLocalExecutionCandidate,
|
||||||
pub(crate) candidate_index: u32,
|
pub(crate) candidate_index: u32,
|
||||||
pub(crate) retry_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,
|
pub(crate) candidate_id: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl LocalExecutionCandidateAttempt {
|
impl LocalExecutionCandidateAttempt {
|
||||||
pub(crate) fn attempt_identity(&self) -> ExecutionAttemptIdentity {
|
pub(crate) fn attempt_identity(&self) -> ExecutionAttemptIdentity {
|
||||||
ExecutionAttemptIdentity::new(self.candidate_index, self.retry_index)
|
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>,
|
F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value>,
|
||||||
{
|
{
|
||||||
let created_at_unix_ms = current_unix_ms();
|
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() {
|
for (candidate_index, eligible) in candidates.into_iter().enumerate() {
|
||||||
let candidate_index = candidate_index as u32;
|
let candidate_index = candidate_index as u32;
|
||||||
let attempt_identities =
|
let attempt_slots = local_attempt_slot_count(&eligible.transport);
|
||||||
build_local_attempt_identities(candidate_index, &eligible.transport)
|
let pool_key_index = eligible.orchestration.pool_key_index;
|
||||||
.into_iter()
|
let extra_data = build_extra_data(&eligible);
|
||||||
.map(|identity| identity.with_pool_key_index(eligible.orchestration.pool_key_index))
|
let mut owned_eligible = Some(eligible);
|
||||||
.collect::<Vec<_>>();
|
|
||||||
|
|
||||||
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 generated_candidate_id = Uuid::new_v4().to_string();
|
||||||
let candidate_id = state
|
let candidate_id = state
|
||||||
.persist_available_local_candidate(
|
.persist_available_local_candidate(
|
||||||
@@ -106,18 +112,23 @@ where
|
|||||||
attempt_identity.retry_index,
|
attempt_identity.retry_index,
|
||||||
&generated_candidate_id,
|
&generated_candidate_id,
|
||||||
required_capabilities,
|
required_capabilities,
|
||||||
build_extra_data(&eligible),
|
extra_data.clone(),
|
||||||
created_at_unix_ms,
|
created_at_unix_ms,
|
||||||
error_context,
|
error_context,
|
||||||
)
|
)
|
||||||
.await;
|
.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 {
|
materialized.push(LocalExecutionCandidateAttempt {
|
||||||
eligible: eligible.clone(),
|
eligible,
|
||||||
candidate_index: attempt_identity.candidate_index,
|
candidate_index: attempt_identity.candidate_index,
|
||||||
retry_index: attempt_identity.retry_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,
|
candidate_id,
|
||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -54,7 +54,7 @@ pub(crate) fn build_local_execution_candidate_metadata(
|
|||||||
) -> Value {
|
) -> Value {
|
||||||
build_local_execution_candidate_metadata_for_candidate(
|
build_local_execution_candidate_metadata_for_candidate(
|
||||||
&parts.eligible.candidate,
|
&parts.eligible.candidate,
|
||||||
Some(&parts.eligible.transport),
|
Some(parts.eligible.transport.as_ref()),
|
||||||
parts.provider_api_format,
|
parts.provider_api_format,
|
||||||
parts.client_api_format,
|
parts.client_api_format,
|
||||||
parts.extra_fields,
|
parts.extra_fields,
|
||||||
@@ -127,7 +127,7 @@ pub(crate) fn build_local_execution_candidate_contract_metadata(
|
|||||||
append_execution_contract_fields_to_value(
|
append_execution_contract_fields_to_value(
|
||||||
build_local_execution_candidate_metadata_for_candidate(
|
build_local_execution_candidate_metadata_for_candidate(
|
||||||
&parts.eligible.candidate,
|
&parts.eligible.candidate,
|
||||||
Some(&parts.eligible.transport),
|
Some(parts.eligible.transport.as_ref()),
|
||||||
parts.provider_api_format,
|
parts.provider_api_format,
|
||||||
parts.client_api_format,
|
parts.client_api_format,
|
||||||
parts.extra_fields,
|
parts.extra_fields,
|
||||||
|
|||||||
@@ -81,9 +81,9 @@ fn build_sync_plan_payload_from_decision(
|
|||||||
parts: &http::request::Parts,
|
parts: &http::request::Parts,
|
||||||
body_json: &serde_json::Value,
|
body_json: &serde_json::Value,
|
||||||
plan_kind: &str,
|
plan_kind: &str,
|
||||||
payload: GatewayControlSyncDecisionResponse,
|
mut payload: GatewayControlSyncDecisionResponse,
|
||||||
) -> Result<Option<GatewayControlPlanResponse>, GatewayError> {
|
) -> 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 {
|
let plan_and_report = match plan_kind {
|
||||||
OPENAI_CHAT_SYNC_PLAN_KIND => {
|
OPENAI_CHAT_SYNC_PLAN_KIND => {
|
||||||
build_openai_chat_sync_plan_from_decision(parts, body_json, payload)?
|
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,
|
parts: &http::request::Parts,
|
||||||
body_json: &serde_json::Value,
|
body_json: &serde_json::Value,
|
||||||
plan_kind: &str,
|
plan_kind: &str,
|
||||||
payload: GatewayControlSyncDecisionResponse,
|
mut payload: GatewayControlSyncDecisionResponse,
|
||||||
) -> Result<Option<GatewayControlPlanResponse>, GatewayError> {
|
) -> 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 {
|
let plan_and_report = match plan_kind {
|
||||||
OPENAI_CHAT_STREAM_PLAN_KIND => {
|
OPENAI_CHAT_STREAM_PLAN_KIND => {
|
||||||
build_openai_chat_stream_plan_from_decision(parts, body_json, payload)?
|
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);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
|
|
||||||
let auth_pair = extract_auth_header_pair(&follow_up.plan.headers);
|
let aether_video_tasks_core::LocalVideoTaskFollowUpPlan {
|
||||||
let execution_strategy =
|
plan,
|
||||||
if follow_up.plan.provider_api_format == follow_up.plan.client_api_format {
|
report_kind,
|
||||||
ExecutionStrategy::LocalSameFormat
|
report_context,
|
||||||
} else {
|
} = follow_up;
|
||||||
ExecutionStrategy::LocalCrossFormat
|
let aether_contracts::ExecutionPlan {
|
||||||
};
|
request_id: _request_id,
|
||||||
let conversion_mode = if follow_up.plan.provider_api_format == follow_up.plan.client_api_format
|
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
|
ConversionMode::None
|
||||||
} else {
|
} else {
|
||||||
ConversionMode::Bidirectional
|
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!(
|
debug!(
|
||||||
event_name = "local_video_follow_up_sync_decision_payload_built",
|
event_name = "local_video_follow_up_sync_decision_payload_built",
|
||||||
log_type = "debug",
|
log_type = "debug",
|
||||||
trace_id = %trace_id,
|
trace_id = %trace_id,
|
||||||
request_id = %trace_id,
|
request_id = %trace_id,
|
||||||
candidate_id = ?follow_up.plan.candidate_id,
|
candidate_id = ?candidate_id,
|
||||||
provider_id = %follow_up.plan.provider_id,
|
provider_id = %provider_id,
|
||||||
endpoint_id = %follow_up.plan.endpoint_id,
|
endpoint_id = %endpoint_id,
|
||||||
key_id = %follow_up.plan.key_id,
|
key_id = %key_id,
|
||||||
plan_kind,
|
plan_kind,
|
||||||
downstream_path = %parts.uri.path(),
|
downstream_path = %parts.uri.path(),
|
||||||
provider_api_format = %follow_up.plan.provider_api_format,
|
provider_api_format = %provider_api_format,
|
||||||
client_api_format = %follow_up.plan.client_api_format,
|
client_api_format = %client_api_format,
|
||||||
upstream_base_url = ?upstream_base_url,
|
upstream_base_url = ?upstream_base_url,
|
||||||
upstream_url = %follow_up.plan.url,
|
upstream_url = %url,
|
||||||
"gateway built local video follow-up sync decision payload"
|
"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()),
|
execution_strategy: Some(execution_strategy.as_str().to_string()),
|
||||||
conversion_mode: Some(conversion_mode.as_str().to_string()),
|
conversion_mode: Some(conversion_mode.as_str().to_string()),
|
||||||
request_id: Some(trace_id.to_string()),
|
request_id: Some(trace_id.to_string()),
|
||||||
candidate_id: follow_up.plan.candidate_id.clone(),
|
candidate_id,
|
||||||
provider_name: follow_up.plan.provider_name.clone(),
|
provider_name,
|
||||||
provider_id: Some(follow_up.plan.provider_id.clone()),
|
provider_id: Some(provider_id),
|
||||||
endpoint_id: Some(follow_up.plan.endpoint_id.clone()),
|
endpoint_id: Some(endpoint_id),
|
||||||
key_id: Some(follow_up.plan.key_id.clone()),
|
key_id: Some(key_id),
|
||||||
upstream_base_url,
|
upstream_base_url,
|
||||||
upstream_url: Some(follow_up.plan.url.clone()),
|
upstream_url: Some(url),
|
||||||
provider_request_method: Some(follow_up.plan.method.clone()),
|
provider_request_method: Some(method),
|
||||||
auth_header: auth_pair.as_ref().map(|(name, _)| name.clone()),
|
auth_header,
|
||||||
auth_value: auth_pair.as_ref().map(|(_, value)| value.clone()),
|
auth_value,
|
||||||
provider_api_format: Some(follow_up.plan.provider_api_format.clone()),
|
provider_api_format: Some(provider_api_format),
|
||||||
client_api_format: Some(follow_up.plan.client_api_format.clone()),
|
client_api_format: Some(client_api_format),
|
||||||
provider_contract: Some(follow_up.plan.provider_api_format.clone()),
|
provider_contract: Some(provider_contract),
|
||||||
client_contract: Some(follow_up.plan.client_api_format.clone()),
|
client_contract: Some(client_contract),
|
||||||
model_name: follow_up.plan.model_name.clone(),
|
model_name,
|
||||||
mapped_model: None,
|
mapped_model: None,
|
||||||
prompt_cache_key: None,
|
prompt_cache_key: None,
|
||||||
extra_headers: BTreeMap::new(),
|
extra_headers: BTreeMap::new(),
|
||||||
provider_request_headers: follow_up.plan.headers.clone(),
|
provider_request_headers: headers,
|
||||||
provider_request_body: follow_up.plan.body.json_body.clone(),
|
provider_request_body: json_body,
|
||||||
provider_request_body_base64: follow_up.plan.body.body_bytes_b64.clone(),
|
provider_request_body_base64: body_bytes_b64,
|
||||||
content_type: follow_up.plan.content_type.clone(),
|
content_type,
|
||||||
proxy: follow_up.plan.proxy.clone(),
|
proxy,
|
||||||
tls_profile: follow_up.plan.tls_profile.clone(),
|
tls_profile,
|
||||||
timeouts: follow_up.plan.timeouts.clone(),
|
timeouts,
|
||||||
upstream_is_stream: false,
|
upstream_is_stream: false,
|
||||||
report_kind: follow_up.report_kind,
|
report_kind,
|
||||||
report_context: follow_up.report_context,
|
report_context,
|
||||||
auth_context: Some(build_execution_runtime_auth_context(&auth_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",
|
"authorization",
|
||||||
"x-api-key",
|
"x-api-key",
|
||||||
@@ -227,7 +262,7 @@ fn extract_auth_header_pair(headers: &BTreeMap<String, String>) -> Option<(Strin
|
|||||||
headers
|
headers
|
||||||
.iter()
|
.iter()
|
||||||
.find(|(header_name, _)| header_name.eq_ignore_ascii_case(name))
|
.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 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};
|
use crate::{GatewayControlSyncDecisionResponse, GatewayError};
|
||||||
|
|
||||||
pub(crate) fn build_passthrough_sync_plan_from_decision(
|
pub(crate) fn build_passthrough_sync_plan_from_decision(
|
||||||
parts: &http::request::Parts,
|
parts: &http::request::Parts,
|
||||||
payload: GatewayControlSyncDecisionResponse,
|
payload: GatewayControlSyncDecisionResponse,
|
||||||
) -> Result<Option<LocalSyncPlanAndReport>, GatewayError> {
|
) -> Result<Option<LocalSyncPlanAndReport>, GatewayError> {
|
||||||
let Some(request_id) = payload
|
let mut payload = payload;
|
||||||
.request_id
|
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(provider_id) = payload
|
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
|
||||||
.provider_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(endpoint_id) = payload
|
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
|
||||||
.endpoint_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(key_id) = payload
|
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
|
||||||
.key_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(provider_api_format) = payload
|
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
|
||||||
.provider_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(client_api_format) = payload
|
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
|
||||||
.client_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(upstream_url) = payload
|
let Some(upstream_url) = take_non_empty_string(&mut payload.upstream_url) else {
|
||||||
.upstream_url
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let (request_body, provider_request_body_for_report) = resolve_passthrough_sync_request_body(
|
let provider_request_headers = std::mem::take(&mut payload.provider_request_headers);
|
||||||
payload.provider_request_body.clone(),
|
let ignored_provider_request_body = serde_json::Value::Null;
|
||||||
payload.provider_request_body_base64.clone(),
|
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 {
|
let plan = ExecutionPlan {
|
||||||
request_id,
|
request_id,
|
||||||
candidate_id: payload.candidate_id.clone(),
|
candidate_id: payload.candidate_id.take(),
|
||||||
provider_name: payload.provider_name.clone(),
|
provider_name: payload.provider_name.take(),
|
||||||
provider_id,
|
provider_id,
|
||||||
endpoint_id,
|
endpoint_id,
|
||||||
key_id,
|
key_id,
|
||||||
method: payload
|
method: provider_request_method.unwrap_or_else(|| parts.method.to_string()),
|
||||||
.provider_request_method
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
.unwrap_or_else(|| parts.method.to_string()),
|
|
||||||
url: upstream_url,
|
url: upstream_url,
|
||||||
headers: payload.provider_request_headers.clone(),
|
headers: provider_request_headers,
|
||||||
content_type: payload.content_type.clone().or_else(|| {
|
content_type,
|
||||||
payload
|
|
||||||
.provider_request_headers
|
|
||||||
.get("content-type")
|
|
||||||
.cloned()
|
|
||||||
}),
|
|
||||||
content_encoding: None,
|
content_encoding: None,
|
||||||
body: request_body,
|
body: request_body,
|
||||||
stream: false,
|
stream: false,
|
||||||
client_api_format,
|
client_api_format,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
model_name: payload.model_name.clone(),
|
model_name: payload.model_name.take(),
|
||||||
proxy: payload.proxy.clone(),
|
proxy: payload.proxy.take(),
|
||||||
tls_profile: payload.tls_profile.clone(),
|
tls_profile: payload.tls_profile.take(),
|
||||||
timeouts: payload.timeouts.clone(),
|
timeouts: payload.timeouts.take(),
|
||||||
};
|
};
|
||||||
|
|
||||||
let report_context = augment_sync_report_context(
|
|
||||||
payload.report_context,
|
|
||||||
&plan.headers,
|
|
||||||
&provider_request_body_for_report,
|
|
||||||
)?;
|
|
||||||
|
|
||||||
Ok(Some(LocalSyncPlanAndReport {
|
Ok(Some(LocalSyncPlanAndReport {
|
||||||
plan,
|
plan,
|
||||||
report_kind: payload.report_kind,
|
report_kind: payload.report_kind,
|
||||||
@@ -109,71 +82,44 @@ pub(crate) fn build_passthrough_stream_plan_from_decision(
|
|||||||
parts: &http::request::Parts,
|
parts: &http::request::Parts,
|
||||||
payload: GatewayControlSyncDecisionResponse,
|
payload: GatewayControlSyncDecisionResponse,
|
||||||
) -> Result<Option<LocalStreamPlanAndReport>, GatewayError> {
|
) -> Result<Option<LocalStreamPlanAndReport>, GatewayError> {
|
||||||
let Some(request_id) = payload
|
let mut payload = payload;
|
||||||
.request_id
|
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(provider_id) = payload
|
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
|
||||||
.provider_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(endpoint_id) = payload
|
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
|
||||||
.endpoint_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(key_id) = payload
|
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
|
||||||
.key_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(provider_api_format) = payload
|
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
|
||||||
.provider_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(client_api_format) = payload
|
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
|
||||||
.client_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(upstream_url) = payload
|
let Some(upstream_url) = take_non_empty_string(&mut payload.upstream_url) else {
|
||||||
.upstream_url
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
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 {
|
let plan = ExecutionPlan {
|
||||||
request_id,
|
request_id,
|
||||||
candidate_id: payload.candidate_id.clone(),
|
candidate_id: payload.candidate_id.take(),
|
||||||
provider_name: payload.provider_name.clone(),
|
provider_name: payload.provider_name.take(),
|
||||||
provider_id,
|
provider_id,
|
||||||
endpoint_id,
|
endpoint_id,
|
||||||
key_id,
|
key_id,
|
||||||
method: parts.method.to_string(),
|
method: parts.method.to_string(),
|
||||||
url: upstream_url,
|
url: upstream_url,
|
||||||
headers: payload.provider_request_headers.clone(),
|
headers: provider_request_headers,
|
||||||
content_type: payload.content_type.clone().or_else(|| {
|
content_type,
|
||||||
payload
|
|
||||||
.provider_request_headers
|
|
||||||
.get("content-type")
|
|
||||||
.cloned()
|
|
||||||
}),
|
|
||||||
content_encoding: None,
|
content_encoding: None,
|
||||||
body: RequestBody {
|
body: RequestBody {
|
||||||
json_body: None,
|
json_body: None,
|
||||||
@@ -183,10 +129,10 @@ pub(crate) fn build_passthrough_stream_plan_from_decision(
|
|||||||
stream: true,
|
stream: true,
|
||||||
client_api_format,
|
client_api_format,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
model_name: payload.model_name.clone(),
|
model_name: payload.model_name.take(),
|
||||||
proxy: payload.proxy.clone(),
|
proxy: payload.proxy.take(),
|
||||||
tls_profile: payload.tls_profile.clone(),
|
tls_profile: payload.tls_profile.take(),
|
||||||
timeouts: payload.timeouts.clone(),
|
timeouts: payload.timeouts.take(),
|
||||||
};
|
};
|
||||||
|
|
||||||
Ok(Some(LocalStreamPlanAndReport {
|
Ok(Some(LocalStreamPlanAndReport {
|
||||||
@@ -199,35 +145,33 @@ pub(crate) fn build_passthrough_stream_plan_from_decision(
|
|||||||
fn resolve_passthrough_sync_request_body(
|
fn resolve_passthrough_sync_request_body(
|
||||||
provider_request_body: Option<serde_json::Value>,
|
provider_request_body: Option<serde_json::Value>,
|
||||||
provider_request_body_base64: Option<String>,
|
provider_request_body_base64: Option<String>,
|
||||||
) -> (RequestBody, serde_json::Value) {
|
) -> RequestBody {
|
||||||
if let Some(body_bytes_b64) = provider_request_body_base64
|
if let Some(body_bytes_b64) = provider_request_body_base64.and_then(trim_owned_non_empty_string)
|
||||||
.as_ref()
|
|
||||||
.map(|value| value.trim())
|
|
||||||
.filter(|value| !value.is_empty())
|
|
||||||
.map(ToOwned::to_owned)
|
|
||||||
{
|
{
|
||||||
return (
|
return RequestBody {
|
||||||
RequestBody {
|
json_body: None,
|
||||||
json_body: None,
|
body_bytes_b64: Some(body_bytes_b64),
|
||||||
body_bytes_b64: Some(body_bytes_b64.clone()),
|
body_ref: None,
|
||||||
body_ref: None,
|
};
|
||||||
},
|
|
||||||
serde_json::json!({"body_bytes_b64": body_bytes_b64}),
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
match provider_request_body.unwrap_or(serde_json::Value::Null) {
|
match provider_request_body.unwrap_or(serde_json::Value::Null) {
|
||||||
serde_json::Value::Null => (
|
serde_json::Value::Null => RequestBody {
|
||||||
RequestBody {
|
json_body: None,
|
||||||
json_body: None,
|
body_bytes_b64: None,
|
||||||
body_bytes_b64: None,
|
body_ref: None,
|
||||||
body_ref: None,
|
},
|
||||||
},
|
other => RequestBody::from_json(other),
|
||||||
serde_json::Value::Null,
|
|
||||||
),
|
|
||||||
other => {
|
|
||||||
let report_body = other.clone();
|
|
||||||
(RequestBody::from_json(other), report_body)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn trim_owned_non_empty_string(value: String) -> Option<String> {
|
||||||
|
let trimmed = value.trim();
|
||||||
|
if trimmed.is_empty() {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
if trimmed.len() == value.len() {
|
||||||
|
return Some(value);
|
||||||
|
}
|
||||||
|
Some(trimmed.to_owned())
|
||||||
|
}
|
||||||
|
|||||||
@@ -134,7 +134,7 @@ pub(crate) async fn materialize_local_same_format_provider_candidate_attempts(
|
|||||||
skipped_candidate.extra_data = Some(
|
skipped_candidate.extra_data = Some(
|
||||||
build_local_execution_candidate_contract_metadata_for_candidate(
|
build_local_execution_candidate_contract_metadata_for_candidate(
|
||||||
&skipped_candidate.candidate,
|
&skipped_candidate.candidate,
|
||||||
skipped_candidate.transport.as_ref(),
|
skipped_candidate.transport_ref(),
|
||||||
provider_api_format.as_str(),
|
provider_api_format.as_str(),
|
||||||
spec_metadata.api_format,
|
spec_metadata.api_format,
|
||||||
serde_json::Map::new(),
|
serde_json::Map::new(),
|
||||||
|
|||||||
@@ -41,7 +41,6 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
|||||||
let LocalSameFormatProviderCandidateAttempt {
|
let LocalSameFormatProviderCandidateAttempt {
|
||||||
eligible,
|
eligible,
|
||||||
candidate_index,
|
candidate_index,
|
||||||
candidate_group_id,
|
|
||||||
candidate_id,
|
candidate_id,
|
||||||
..
|
..
|
||||||
} = &attempt;
|
} = &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,
|
provider_api_format: spec_metadata.api_format,
|
||||||
client_api_format: spec_metadata.api_format,
|
client_api_format: spec_metadata.api_format,
|
||||||
mapped_model: Some(&resolved.mapped_model),
|
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),
|
upstream_url: Some(&resolved.upstream_url),
|
||||||
provider_request_method: Some(serde_json::Value::Null),
|
provider_request_method: Some(serde_json::Value::Null),
|
||||||
provider_request_headers: Some(&resolved.provider_request_headers),
|
provider_request_headers: Some(&resolved.provider_request_headers),
|
||||||
original_headers: &parts.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,
|
has_envelope: resolved.is_kiro || resolved.is_antigravity,
|
||||||
needs_conversion: false,
|
needs_conversion: false,
|
||||||
extra_fields,
|
extra_fields,
|
||||||
@@ -112,6 +112,19 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
|||||||
),
|
),
|
||||||
&resolved.transport,
|
&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(
|
Some(build_local_execution_decision_response(
|
||||||
LocalExecutionDecisionResponseParts {
|
LocalExecutionDecisionResponseParts {
|
||||||
@@ -121,29 +134,29 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
|||||||
conversion_mode: ConversionMode::None,
|
conversion_mode: ConversionMode::None,
|
||||||
request_id: trace_id.to_string(),
|
request_id: trace_id.to_string(),
|
||||||
candidate_id: candidate_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(),
|
provider_id: candidate.provider_id.clone(),
|
||||||
endpoint_id: candidate.endpoint_id.clone(),
|
endpoint_id: candidate.endpoint_id.clone(),
|
||||||
key_id: candidate.key_id.clone(),
|
key_id: candidate.key_id.clone(),
|
||||||
upstream_base_url: resolved.transport.endpoint.base_url.clone(),
|
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||||
upstream_url: resolved.upstream_url.clone(),
|
upstream_url,
|
||||||
provider_request_method: None,
|
provider_request_method: None,
|
||||||
auth_header: resolved.auth_header.clone(),
|
auth_header,
|
||||||
auth_value: resolved.auth_value.clone(),
|
auth_value,
|
||||||
provider_api_format: spec_metadata.api_format.to_string(),
|
provider_api_format: spec_metadata.api_format.to_string(),
|
||||||
client_api_format: spec_metadata.api_format.to_string(),
|
client_api_format: spec_metadata.api_format.to_string(),
|
||||||
model_name: input.requested_model.clone(),
|
model_name: input.requested_model.clone(),
|
||||||
mapped_model: resolved.mapped_model.clone(),
|
mapped_model,
|
||||||
prompt_cache_key,
|
prompt_cache_key,
|
||||||
provider_request_headers: resolved.provider_request_headers.clone(),
|
provider_request_headers,
|
||||||
provider_request_body: Some(resolved.provider_request_body.clone()),
|
provider_request_body: Some(provider_request_body),
|
||||||
provider_request_body_base64: None,
|
provider_request_body_base64: None,
|
||||||
content_type: Some("application/json".to_string()),
|
content_type: Some("application/json".to_string()),
|
||||||
proxy,
|
proxy,
|
||||||
tls_profile,
|
tls_profile,
|
||||||
timeouts: resolve_transport_execution_timeouts(&resolved.transport),
|
timeouts: resolve_transport_execution_timeouts(&transport),
|
||||||
upstream_is_stream: resolved.upstream_is_stream,
|
upstream_is_stream,
|
||||||
report_kind: Some(resolved.report_kind.to_string()),
|
report_kind: Some(report_kind.to_string()),
|
||||||
report_context: Some(report_context),
|
report_context: Some(report_context),
|
||||||
auth_context: input.auth_context.clone(),
|
auth_context: input.auth_context.clone(),
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
use serde_json::Value;
|
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(crate) struct LocalSameFormatProviderCandidatePayloadParts {
|
||||||
pub(super) transport: GatewayProviderTransportSnapshot,
|
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
|
||||||
pub(super) is_antigravity: bool,
|
pub(super) is_antigravity: bool,
|
||||||
pub(super) is_kiro: bool,
|
pub(super) is_kiro: bool,
|
||||||
pub(super) auth_header: Option<String>,
|
pub(super) auth_header: Option<String>,
|
||||||
|
|||||||
@@ -1,3 +1,5 @@
|
|||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
use crate::ai_pipeline::planner::candidate_eligibility::EligibleLocalExecutionCandidate;
|
use crate::ai_pipeline::planner::candidate_eligibility::EligibleLocalExecutionCandidate;
|
||||||
use crate::ai_pipeline::planner::candidate_preparation::{
|
use crate::ai_pipeline::planner::candidate_preparation::{
|
||||||
resolve_candidate_mapped_model, resolve_candidate_oauth_auth, OauthPreparationContext,
|
resolve_candidate_mapped_model, resolve_candidate_oauth_auth, OauthPreparationContext,
|
||||||
@@ -19,7 +21,7 @@ use super::policy::{
|
|||||||
};
|
};
|
||||||
|
|
||||||
pub(super) struct PreparedSameFormatProviderCandidate {
|
pub(super) struct PreparedSameFormatProviderCandidate {
|
||||||
pub(super) transport: GatewayProviderTransportSnapshot,
|
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
|
||||||
pub(super) is_antigravity: bool,
|
pub(super) is_antigravity: bool,
|
||||||
pub(super) is_claude_code: bool,
|
pub(super) is_claude_code: bool,
|
||||||
pub(super) is_vertex: 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 spec_metadata = local_same_format_provider_spec_metadata(spec);
|
||||||
let planner_state = PlannerAppState::new(state);
|
let planner_state = PlannerAppState::new(state);
|
||||||
let candidate = &eligible.candidate;
|
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);
|
let behavior = classify_same_format_provider_request_behavior(&transport, spec_metadata);
|
||||||
|
|
||||||
if !same_format_provider_transport_supported(
|
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()))
|
.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::cmp::Ordering;
|
||||||
use std::collections::{BTreeMap, BTreeSet};
|
use std::collections::{btree_map::Entry, BTreeMap, BTreeSet};
|
||||||
use std::hash::{Hash, Hasher};
|
use std::hash::{Hash, Hasher};
|
||||||
|
|
||||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
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() {
|
if pool_config_for_candidate(candidate).is_none() {
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
if provider_type_by_key_id.contains_key(&candidate.candidate.key_id) {
|
let key_id = candidate.candidate.key_id.clone();
|
||||||
continue;
|
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() {
|
if key_ids.is_empty() {
|
||||||
@@ -324,10 +321,15 @@ fn apply_local_execution_pool_scheduler_with_runtime_map(
|
|||||||
for candidate in candidates {
|
for candidate in candidates {
|
||||||
let pool_enabled = pool_config_for_candidate(&candidate).is_some();
|
let pool_enabled = pool_config_for_candidate(&candidate).is_some();
|
||||||
let group_key = pool_group_key(&candidate, pool_enabled);
|
let group_key = pool_group_key(&candidate, pool_enabled);
|
||||||
if !groups.contains_key(&group_key) {
|
match groups.entry(group_key) {
|
||||||
group_order.push(group_key.clone());
|
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();
|
let mut reordered = Vec::new();
|
||||||
@@ -1509,7 +1511,7 @@ mod tests {
|
|||||||
},
|
},
|
||||||
provider_api_format: "openai:chat".to_string(),
|
provider_api_format: "openai:chat".to_string(),
|
||||||
orchestration: LocalExecutionCandidateMetadata::default(),
|
orchestration: LocalExecutionCandidateMetadata::default(),
|
||||||
transport: crate::ai_pipeline::GatewayProviderTransportSnapshot {
|
transport: Arc::new(crate::ai_pipeline::GatewayProviderTransportSnapshot {
|
||||||
provider: GatewayProviderTransportProvider {
|
provider: GatewayProviderTransportProvider {
|
||||||
id: provider_id.to_string(),
|
id: provider_id.to_string(),
|
||||||
name: provider_id.to_string(),
|
name: provider_id.to_string(),
|
||||||
@@ -1558,7 +1560,7 @@ mod tests {
|
|||||||
decrypted_api_key: "secret".to_string(),
|
decrypted_api_key: "secret".to_string(),
|
||||||
decrypted_auth_config: None,
|
decrypted_auth_config: None,
|
||||||
},
|
},
|
||||||
},
|
}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -24,7 +24,8 @@ pub(crate) struct LocalExecutionReportContextParts<'a> {
|
|||||||
pub(crate) provider_request_method: Option<Value>,
|
pub(crate) provider_request_method: Option<Value>,
|
||||||
pub(crate) provider_request_headers: Option<&'a BTreeMap<String, String>>,
|
pub(crate) provider_request_headers: Option<&'a BTreeMap<String, String>>,
|
||||||
pub(crate) original_headers: &'a http::HeaderMap,
|
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) has_envelope: bool,
|
||||||
pub(crate) needs_conversion: bool,
|
pub(crate) needs_conversion: bool,
|
||||||
pub(crate) extra_fields: Map<String, Value>,
|
pub(crate) extra_fields: Map<String, Value>,
|
||||||
@@ -110,8 +111,11 @@ pub(crate) fn build_local_execution_report_context(
|
|||||||
);
|
);
|
||||||
object.insert(
|
object.insert(
|
||||||
"original_request_body".to_string(),
|
"original_request_body".to_string(),
|
||||||
crate::ai_pipeline::build_report_context_original_request_echo(parts.original_request_body)
|
crate::ai_pipeline::build_report_context_original_request_echo(
|
||||||
.unwrap_or(Value::Null),
|
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("has_envelope".to_string(), Value::Bool(parts.has_envelope));
|
||||||
object.insert(
|
object.insert(
|
||||||
|
|||||||
@@ -49,7 +49,6 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
|
|||||||
.await?;
|
.await?;
|
||||||
let LocalGeminiFilesCandidateAttempt {
|
let LocalGeminiFilesCandidateAttempt {
|
||||||
eligible,
|
eligible,
|
||||||
candidate_group_id,
|
|
||||||
candidate_id,
|
candidate_id,
|
||||||
..
|
..
|
||||||
} = attempt;
|
} = 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_key_id".to_string(), json!(candidate.key_id));
|
||||||
extra_fields.insert("file_name".to_string(), json!(resolved.file_name));
|
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(
|
Some(build_local_execution_decision_response(
|
||||||
LocalExecutionDecisionResponseParts {
|
LocalExecutionDecisionResponseParts {
|
||||||
@@ -80,18 +114,18 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
|
|||||||
endpoint_id: candidate.endpoint_id.clone(),
|
endpoint_id: candidate.endpoint_id.clone(),
|
||||||
key_id: candidate.key_id.clone(),
|
key_id: candidate.key_id.clone(),
|
||||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||||
upstream_url: resolved.upstream_url,
|
upstream_url,
|
||||||
provider_request_method: Some(parts.method.to_string()),
|
provider_request_method: Some(parts.method.to_string()),
|
||||||
auth_header: Some(resolved.auth_header),
|
auth_header: Some(auth_header),
|
||||||
auth_value: Some(resolved.auth_value),
|
auth_value: Some(auth_value),
|
||||||
provider_api_format: GEMINI_FILES_CLIENT_API_FORMAT.to_string(),
|
provider_api_format: GEMINI_FILES_CLIENT_API_FORMAT.to_string(),
|
||||||
client_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(),
|
model_name: "gemini-files".to_string(),
|
||||||
mapped_model: candidate.selected_provider_model_name.clone(),
|
mapped_model: candidate.selected_provider_model_name.clone(),
|
||||||
prompt_cache_key: None,
|
prompt_cache_key: None,
|
||||||
provider_request_headers: resolved.provider_request_headers,
|
provider_request_headers,
|
||||||
provider_request_body: resolved.provider_request_body,
|
provider_request_body,
|
||||||
provider_request_body_base64: resolved.provider_request_body_base64,
|
provider_request_body_base64,
|
||||||
content_type: parts
|
content_type: parts
|
||||||
.headers
|
.headers
|
||||||
.get(http::header::CONTENT_TYPE)
|
.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),
|
timeouts: resolve_transport_execution_timeouts(&transport),
|
||||||
upstream_is_stream: spec_metadata.require_streaming,
|
upstream_is_stream: spec_metadata.require_streaming,
|
||||||
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
||||||
report_context: Some(build_local_execution_report_context(
|
report_context: Some(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,
|
|
||||||
},
|
|
||||||
)),
|
|
||||||
auth_context: input.auth_context.clone(),
|
auth_context: input.auth_context.clone(),
|
||||||
},
|
},
|
||||||
))
|
))
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
use serde_json::json;
|
use serde_json::json;
|
||||||
|
|
||||||
@@ -20,13 +21,12 @@ use super::support::{
|
|||||||
use super::LocalGeminiFilesSpec;
|
use super::LocalGeminiFilesSpec;
|
||||||
|
|
||||||
pub(super) struct LocalGeminiFilesCandidatePayloadParts {
|
pub(super) struct LocalGeminiFilesCandidatePayloadParts {
|
||||||
pub(super) transport: GatewayProviderTransportSnapshot,
|
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
|
||||||
pub(super) auth_header: String,
|
pub(super) auth_header: String,
|
||||||
pub(super) auth_value: String,
|
pub(super) auth_value: String,
|
||||||
pub(super) provider_request_headers: BTreeMap<String, String>,
|
pub(super) provider_request_headers: BTreeMap<String, String>,
|
||||||
pub(super) provider_request_body: Option<serde_json::Value>,
|
pub(super) provider_request_body: Option<serde_json::Value>,
|
||||||
pub(super) provider_request_body_base64: Option<String>,
|
pub(super) provider_request_body_base64: Option<String>,
|
||||||
pub(super) original_request_body: serde_json::Value,
|
|
||||||
pub(super) upstream_url: String,
|
pub(super) upstream_url: String,
|
||||||
pub(super) file_name: String,
|
pub(super) file_name: String,
|
||||||
}
|
}
|
||||||
@@ -119,13 +119,6 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts(
|
|||||||
} else {
|
} else {
|
||||||
None
|
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() {
|
if provider_request_body_base64.is_some() && transport.endpoint.body_rules.is_some() {
|
||||||
mark_skipped_local_gemini_files_candidate(
|
mark_skipped_local_gemini_files_candidate(
|
||||||
state,
|
state,
|
||||||
@@ -165,14 +158,22 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts(
|
|||||||
&auth_value,
|
&auth_value,
|
||||||
&BTreeMap::new(),
|
&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(
|
if !apply_local_header_rules(
|
||||||
&mut provider_request_headers,
|
&mut provider_request_headers,
|
||||||
transport.endpoint.header_rules.as_ref(),
|
transport.endpoint.header_rules.as_ref(),
|
||||||
&[&auth_header, "content-type"],
|
&[&auth_header, "content-type"],
|
||||||
provider_request_body
|
provider_request_body
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.unwrap_or(&original_request_body),
|
.unwrap_or(original_request_body),
|
||||||
Some(&original_request_body),
|
Some(original_request_body),
|
||||||
) {
|
) {
|
||||||
mark_skipped_local_gemini_files_candidate(
|
mark_skipped_local_gemini_files_candidate(
|
||||||
state,
|
state,
|
||||||
@@ -195,13 +196,12 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts(
|
|||||||
.to_string();
|
.to_string();
|
||||||
|
|
||||||
Some(LocalGeminiFilesCandidatePayloadParts {
|
Some(LocalGeminiFilesCandidatePayloadParts {
|
||||||
transport: transport.clone(),
|
transport: Arc::clone(transport),
|
||||||
auth_header,
|
auth_header,
|
||||||
auth_value,
|
auth_value,
|
||||||
provider_request_headers,
|
provider_request_headers,
|
||||||
provider_request_body,
|
provider_request_body,
|
||||||
provider_request_body_base64,
|
provider_request_body_base64,
|
||||||
original_request_body,
|
|
||||||
upstream_url,
|
upstream_url,
|
||||||
file_name,
|
file_name,
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -145,7 +145,7 @@ pub(super) async fn materialize_local_gemini_files_candidate_attempts(
|
|||||||
skipped_candidate.extra_data =
|
skipped_candidate.extra_data =
|
||||||
Some(build_local_execution_candidate_metadata_for_candidate(
|
Some(build_local_execution_candidate_metadata_for_candidate(
|
||||||
&skipped_candidate.candidate,
|
&skipped_candidate.candidate,
|
||||||
skipped_candidate.transport.as_ref(),
|
skipped_candidate.transport_ref(),
|
||||||
GEMINI_FILES_CLIENT_API_FORMAT,
|
GEMINI_FILES_CLIENT_API_FORMAT,
|
||||||
GEMINI_FILES_CLIENT_API_FORMAT,
|
GEMINI_FILES_CLIENT_API_FORMAT,
|
||||||
extra_fields,
|
extra_fields,
|
||||||
|
|||||||
@@ -34,7 +34,6 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
|
|||||||
.await?;
|
.await?;
|
||||||
let LocalVideoCreateCandidateAttempt {
|
let LocalVideoCreateCandidateAttempt {
|
||||||
eligible,
|
eligible,
|
||||||
candidate_group_id,
|
|
||||||
candidate_id,
|
candidate_id,
|
||||||
..
|
..
|
||||||
} = attempt;
|
} = 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()) {
|
if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) {
|
||||||
extra_fields.insert("proxy".to_string(), proxy_value);
|
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(
|
Some(build_local_execution_decision_response(
|
||||||
LocalExecutionDecisionResponseParts {
|
LocalExecutionDecisionResponseParts {
|
||||||
@@ -63,17 +96,17 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
|
|||||||
endpoint_id: candidate.endpoint_id.clone(),
|
endpoint_id: candidate.endpoint_id.clone(),
|
||||||
key_id: candidate.key_id.clone(),
|
key_id: candidate.key_id.clone(),
|
||||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||||
upstream_url: resolved.upstream_url,
|
upstream_url,
|
||||||
provider_request_method: Some(parts.method.to_string()),
|
provider_request_method: Some(parts.method.to_string()),
|
||||||
auth_header: Some(resolved.auth_header),
|
auth_header: Some(auth_header),
|
||||||
auth_value: Some(resolved.auth_value),
|
auth_value: Some(auth_value),
|
||||||
provider_api_format: spec_metadata.api_format.to_string(),
|
provider_api_format: spec_metadata.api_format.to_string(),
|
||||||
client_api_format: spec_metadata.api_format.to_string(),
|
client_api_format: spec_metadata.api_format.to_string(),
|
||||||
model_name: input.requested_model.clone(),
|
model_name: input.requested_model.clone(),
|
||||||
mapped_model: resolved.mapped_model.clone(),
|
mapped_model,
|
||||||
prompt_cache_key: None,
|
prompt_cache_key: None,
|
||||||
provider_request_headers: resolved.provider_request_headers,
|
provider_request_headers,
|
||||||
provider_request_body: Some(resolved.provider_request_body),
|
provider_request_body: Some(provider_request_body),
|
||||||
provider_request_body_base64: None,
|
provider_request_body_base64: None,
|
||||||
content_type: parts
|
content_type: parts
|
||||||
.headers
|
.headers
|
||||||
@@ -87,32 +120,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
|
|||||||
timeouts: resolve_transport_execution_timeouts(&transport),
|
timeouts: resolve_transport_execution_timeouts(&transport),
|
||||||
upstream_is_stream: false,
|
upstream_is_stream: false,
|
||||||
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
||||||
report_context: Some(build_local_execution_report_context(
|
report_context: Some(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,
|
|
||||||
},
|
|
||||||
)),
|
|
||||||
auth_context: input.auth_context.clone(),
|
auth_context: input.auth_context.clone(),
|
||||||
},
|
},
|
||||||
))
|
))
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
|
|
||||||
@@ -26,7 +27,7 @@ use super::support::{
|
|||||||
use super::{LocalVideoCreateFamily, LocalVideoCreateSpec};
|
use super::{LocalVideoCreateFamily, LocalVideoCreateSpec};
|
||||||
|
|
||||||
pub(super) struct LocalVideoCreateCandidatePayloadParts {
|
pub(super) struct LocalVideoCreateCandidatePayloadParts {
|
||||||
pub(super) transport: GatewayProviderTransportSnapshot,
|
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
|
||||||
pub(super) auth_header: String,
|
pub(super) auth_header: String,
|
||||||
pub(super) auth_value: String,
|
pub(super) auth_value: String,
|
||||||
pub(super) mapped_model: String,
|
pub(super) mapped_model: String,
|
||||||
@@ -168,7 +169,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
|||||||
}
|
}
|
||||||
|
|
||||||
Some(LocalVideoCreateCandidatePayloadParts {
|
Some(LocalVideoCreateCandidatePayloadParts {
|
||||||
transport: transport.clone(),
|
transport: Arc::clone(transport),
|
||||||
auth_header,
|
auth_header,
|
||||||
auth_value,
|
auth_value,
|
||||||
mapped_model,
|
mapped_model,
|
||||||
|
|||||||
@@ -212,7 +212,7 @@ async fn materialize_local_video_create_candidate_attempts(
|
|||||||
skipped_candidate.extra_data =
|
skipped_candidate.extra_data =
|
||||||
Some(build_local_execution_candidate_metadata_for_candidate(
|
Some(build_local_execution_candidate_metadata_for_candidate(
|
||||||
&skipped_candidate.candidate,
|
&skipped_candidate.candidate,
|
||||||
skipped_candidate.transport.as_ref(),
|
skipped_candidate.transport_ref(),
|
||||||
api_format,
|
api_format,
|
||||||
api_format,
|
api_format,
|
||||||
serde_json::Map::new(),
|
serde_json::Map::new(),
|
||||||
|
|||||||
@@ -213,7 +213,7 @@ pub(super) async fn materialize_local_standard_candidate_attempts(
|
|||||||
skipped_candidate.extra_data = Some(
|
skipped_candidate.extra_data = Some(
|
||||||
build_local_execution_candidate_contract_metadata_for_candidate(
|
build_local_execution_candidate_contract_metadata_for_candidate(
|
||||||
&skipped_candidate.candidate,
|
&skipped_candidate.candidate,
|
||||||
skipped_candidate.transport.as_ref(),
|
skipped_candidate.transport_ref(),
|
||||||
provider_api_format.as_str(),
|
provider_api_format.as_str(),
|
||||||
spec_metadata.api_format,
|
spec_metadata.api_format,
|
||||||
serde_json::Map::new(),
|
serde_json::Map::new(),
|
||||||
|
|||||||
@@ -35,7 +35,6 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
|
|||||||
let LocalStandardCandidateAttempt {
|
let LocalStandardCandidateAttempt {
|
||||||
eligible,
|
eligible,
|
||||||
candidate_index,
|
candidate_index,
|
||||||
candidate_group_id,
|
|
||||||
candidate_id,
|
candidate_id,
|
||||||
..
|
..
|
||||||
} = &attempt;
|
} = &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);
|
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(
|
Some(build_local_execution_decision_response(
|
||||||
LocalExecutionDecisionResponseParts {
|
LocalExecutionDecisionResponseParts {
|
||||||
@@ -66,58 +112,26 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
|
|||||||
provider_id: candidate.provider_id.clone(),
|
provider_id: candidate.provider_id.clone(),
|
||||||
endpoint_id: candidate.endpoint_id.clone(),
|
endpoint_id: candidate.endpoint_id.clone(),
|
||||||
key_id: candidate.key_id.clone(),
|
key_id: candidate.key_id.clone(),
|
||||||
upstream_base_url: resolved.transport.endpoint.base_url.clone(),
|
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||||
upstream_url: resolved.upstream_url.clone(),
|
upstream_url,
|
||||||
provider_request_method: None,
|
provider_request_method: None,
|
||||||
auth_header: Some(resolved.auth_header.clone()),
|
auth_header: Some(auth_header),
|
||||||
auth_value: Some(resolved.auth_value.clone()),
|
auth_value: Some(auth_value),
|
||||||
provider_api_format: resolved.provider_api_format.clone(),
|
provider_api_format,
|
||||||
client_api_format: spec_metadata.api_format.to_string(),
|
client_api_format: spec_metadata.api_format.to_string(),
|
||||||
model_name: input.requested_model.clone(),
|
model_name: input.requested_model.clone(),
|
||||||
mapped_model: resolved.mapped_model.clone(),
|
mapped_model,
|
||||||
prompt_cache_key: None,
|
prompt_cache_key: None,
|
||||||
provider_request_headers: resolved.provider_request_headers.clone(),
|
provider_request_headers,
|
||||||
provider_request_body: Some(resolved.provider_request_body.clone()),
|
provider_request_body: Some(provider_request_body),
|
||||||
provider_request_body_base64: None,
|
provider_request_body_base64: None,
|
||||||
content_type: Some("application/json".to_string()),
|
content_type: Some("application/json".to_string()),
|
||||||
proxy,
|
proxy,
|
||||||
tls_profile: resolve_transport_tls_profile(&resolved.transport),
|
tls_profile,
|
||||||
timeouts: resolve_transport_execution_timeouts(&resolved.transport),
|
timeouts,
|
||||||
upstream_is_stream: resolved.upstream_is_stream,
|
upstream_is_stream,
|
||||||
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
||||||
report_context: Some(append_local_failover_policy_to_value(
|
report_context: Some(report_context),
|
||||||
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,
|
|
||||||
)),
|
|
||||||
auth_context: input.auth_context.clone(),
|
auth_context: input.auth_context.clone(),
|
||||||
},
|
},
|
||||||
))
|
))
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
|
|
||||||
@@ -27,7 +28,7 @@ pub(crate) struct LocalStandardCandidatePayloadParts {
|
|||||||
pub(super) provider_request_headers: BTreeMap<String, String>,
|
pub(super) provider_request_headers: BTreeMap<String, String>,
|
||||||
pub(super) upstream_url: String,
|
pub(super) upstream_url: String,
|
||||||
pub(super) upstream_is_stream: bool,
|
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(
|
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,
|
provider_request_headers,
|
||||||
upstream_url,
|
upstream_url,
|
||||||
upstream_is_stream,
|
upstream_is_stream,
|
||||||
transport: transport.clone(),
|
transport: Arc::clone(transport),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ use aether_contracts::{ExecutionPlan, RequestBody};
|
|||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
augment_sync_report_context, generic_decision_missing_exact_provider_request,
|
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::ai_pipeline::transport::ensure_upstream_auth_header;
|
||||||
use crate::{GatewayControlSyncDecisionResponse, GatewayError};
|
use crate::{GatewayControlSyncDecisionResponse, GatewayError};
|
||||||
@@ -12,74 +12,41 @@ pub(crate) fn build_gemini_sync_plan_from_decision(
|
|||||||
_body_json: &serde_json::Value,
|
_body_json: &serde_json::Value,
|
||||||
payload: GatewayControlSyncDecisionResponse,
|
payload: GatewayControlSyncDecisionResponse,
|
||||||
) -> Result<Option<LocalSyncPlanAndReport>, GatewayError> {
|
) -> Result<Option<LocalSyncPlanAndReport>, GatewayError> {
|
||||||
|
let mut payload = payload;
|
||||||
if generic_decision_missing_exact_provider_request(&payload) {
|
if generic_decision_missing_exact_provider_request(&payload) {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
let Some(request_id) = payload
|
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
|
||||||
.request_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(provider_id) = payload
|
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
|
||||||
.provider_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(endpoint_id) = payload
|
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
|
||||||
.endpoint_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(key_id) = payload
|
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
|
||||||
.key_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(url) = payload
|
let Some(url) = take_non_empty_string(&mut payload.upstream_url) else {
|
||||||
.upstream_url
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let auth_header = payload
|
let auth_header = take_non_empty_string(&mut payload.auth_header);
|
||||||
.auth_header
|
let auth_value = take_non_empty_string(&mut payload.auth_value);
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty());
|
|
||||||
let auth_value = payload
|
|
||||||
.auth_value
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty());
|
|
||||||
if auth_header.is_some() != auth_value.is_some() {
|
if auth_header.is_some() != auth_value.is_some() {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
let Some(provider_api_format) = payload
|
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
|
||||||
.provider_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(client_api_format) = payload
|
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
|
||||||
.client_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
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);
|
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()) {
|
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);
|
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())
|
.entry("accept".to_string())
|
||||||
.or_insert_with(|| "text/event-stream".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 {
|
let plan = ExecutionPlan {
|
||||||
request_id,
|
request_id,
|
||||||
candidate_id: payload.candidate_id.clone(),
|
candidate_id: payload.candidate_id.take(),
|
||||||
provider_name: payload.provider_name.clone(),
|
provider_name: payload.provider_name.take(),
|
||||||
provider_id,
|
provider_id,
|
||||||
endpoint_id,
|
endpoint_id,
|
||||||
key_id,
|
key_id,
|
||||||
method: "POST".to_string(),
|
method: "POST".to_string(),
|
||||||
url,
|
url,
|
||||||
headers: std::mem::take(&mut provider_request_headers),
|
headers: std::mem::take(&mut provider_request_headers),
|
||||||
content_type: payload
|
content_type,
|
||||||
.content_type
|
|
||||||
.clone()
|
|
||||||
.or_else(|| Some("application/json".to_string())),
|
|
||||||
content_encoding: None,
|
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,
|
stream: payload.upstream_is_stream,
|
||||||
client_api_format,
|
client_api_format,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
model_name: payload.model_name.clone(),
|
model_name: payload.model_name.take(),
|
||||||
proxy: payload.proxy.clone(),
|
proxy: payload.proxy.take(),
|
||||||
tls_profile: payload.tls_profile.clone(),
|
tls_profile: payload.tls_profile.take(),
|
||||||
timeouts: payload.timeouts.clone(),
|
timeouts: payload.timeouts.take(),
|
||||||
};
|
};
|
||||||
|
|
||||||
let report_context = augment_sync_report_context(
|
|
||||||
payload.report_context,
|
|
||||||
&plan.headers,
|
|
||||||
&provider_request_body_value,
|
|
||||||
)?;
|
|
||||||
|
|
||||||
Ok(Some(LocalSyncPlanAndReport {
|
Ok(Some(LocalSyncPlanAndReport {
|
||||||
plan,
|
plan,
|
||||||
report_kind: payload.report_kind,
|
report_kind: payload.report_kind,
|
||||||
@@ -131,109 +98,76 @@ pub(crate) fn build_gemini_stream_plan_from_decision(
|
|||||||
_body_json: &serde_json::Value,
|
_body_json: &serde_json::Value,
|
||||||
payload: GatewayControlSyncDecisionResponse,
|
payload: GatewayControlSyncDecisionResponse,
|
||||||
) -> Result<Option<LocalStreamPlanAndReport>, GatewayError> {
|
) -> Result<Option<LocalStreamPlanAndReport>, GatewayError> {
|
||||||
|
let mut payload = payload;
|
||||||
if generic_decision_missing_exact_provider_request(&payload) {
|
if generic_decision_missing_exact_provider_request(&payload) {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
let Some(request_id) = payload
|
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
|
||||||
.request_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(provider_id) = payload
|
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
|
||||||
.provider_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(endpoint_id) = payload
|
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
|
||||||
.endpoint_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(key_id) = payload
|
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
|
||||||
.key_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(url) = payload
|
let Some(url) = take_non_empty_string(&mut payload.upstream_url) else {
|
||||||
.upstream_url
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let auth_header = payload
|
let auth_header = take_non_empty_string(&mut payload.auth_header);
|
||||||
.auth_header
|
let auth_value = take_non_empty_string(&mut payload.auth_value);
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty());
|
|
||||||
let auth_value = payload
|
|
||||||
.auth_value
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty());
|
|
||||||
if auth_header.is_some() != auth_value.is_some() {
|
if auth_header.is_some() != auth_value.is_some() {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
let Some(provider_api_format) = payload
|
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
|
||||||
.provider_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(client_api_format) = payload
|
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
|
||||||
.client_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
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);
|
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()) {
|
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);
|
ensure_upstream_auth_header(&mut provider_request_headers, auth_header, auth_value);
|
||||||
}
|
}
|
||||||
provider_request_headers.insert("accept".to_string(), "text/event-stream".to_string());
|
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 {
|
let plan = ExecutionPlan {
|
||||||
request_id,
|
request_id,
|
||||||
candidate_id: payload.candidate_id.clone(),
|
candidate_id: payload.candidate_id.take(),
|
||||||
provider_name: payload.provider_name.clone(),
|
provider_name: payload.provider_name.take(),
|
||||||
provider_id,
|
provider_id,
|
||||||
endpoint_id,
|
endpoint_id,
|
||||||
key_id,
|
key_id,
|
||||||
method: "POST".to_string(),
|
method: "POST".to_string(),
|
||||||
url,
|
url,
|
||||||
headers: std::mem::take(&mut provider_request_headers),
|
headers: std::mem::take(&mut provider_request_headers),
|
||||||
content_type: payload
|
content_type,
|
||||||
.content_type
|
|
||||||
.clone()
|
|
||||||
.or_else(|| Some("application/json".to_string())),
|
|
||||||
content_encoding: None,
|
content_encoding: None,
|
||||||
body: RequestBody::from_json(provider_request_body_value.clone()),
|
body: RequestBody::from_json(provider_request_body_value),
|
||||||
stream: true,
|
stream: true,
|
||||||
client_api_format,
|
client_api_format,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
model_name: payload.model_name.clone(),
|
model_name: payload.model_name.take(),
|
||||||
proxy: payload.proxy.clone(),
|
proxy: payload.proxy.take(),
|
||||||
tls_profile: payload.tls_profile.clone(),
|
tls_profile: payload.tls_profile.take(),
|
||||||
timeouts: payload.timeouts.clone(),
|
timeouts: payload.timeouts.take(),
|
||||||
};
|
};
|
||||||
|
|
||||||
let report_context = augment_sync_report_context(
|
|
||||||
payload.report_context,
|
|
||||||
&plan.headers,
|
|
||||||
&provider_request_body_value,
|
|
||||||
)?;
|
|
||||||
|
|
||||||
Ok(Some(LocalStreamPlanAndReport {
|
Ok(Some(LocalStreamPlanAndReport {
|
||||||
plan,
|
plan,
|
||||||
report_kind: payload.report_kind,
|
report_kind: payload.report_kind,
|
||||||
|
|||||||
@@ -32,7 +32,6 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
|
|||||||
let LocalOpenAiChatCandidateAttempt {
|
let LocalOpenAiChatCandidateAttempt {
|
||||||
eligible,
|
eligible,
|
||||||
candidate_index,
|
candidate_index,
|
||||||
candidate_group_id,
|
|
||||||
candidate_id,
|
candidate_id,
|
||||||
..
|
..
|
||||||
} = attempt;
|
} = 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);
|
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(
|
Some(build_local_execution_decision_response(
|
||||||
LocalExecutionDecisionResponseParts {
|
LocalExecutionDecisionResponseParts {
|
||||||
decision_is_stream: upstream_is_stream,
|
decision_is_stream: upstream_is_stream,
|
||||||
decision_kind: decision_kind.to_string(),
|
decision_kind: decision_kind.to_string(),
|
||||||
execution_strategy: resolved.execution_strategy,
|
execution_strategy,
|
||||||
conversion_mode: resolved.conversion_mode,
|
conversion_mode,
|
||||||
request_id: trace_id.to_string(),
|
request_id: trace_id.to_string(),
|
||||||
candidate_id: candidate_id.clone(),
|
candidate_id: candidate_id.clone(),
|
||||||
provider_name: resolved.transport.provider.name.clone(),
|
provider_name: transport.provider.name.clone(),
|
||||||
provider_id: candidate.provider_id.clone(),
|
provider_id: candidate.provider_id.clone(),
|
||||||
endpoint_id: candidate.endpoint_id.clone(),
|
endpoint_id: candidate.endpoint_id.clone(),
|
||||||
key_id: candidate.key_id.clone(),
|
key_id: candidate.key_id.clone(),
|
||||||
upstream_base_url: resolved.transport.endpoint.base_url.clone(),
|
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||||
upstream_url: resolved.upstream_url.clone(),
|
upstream_url,
|
||||||
provider_request_method: None,
|
provider_request_method: None,
|
||||||
auth_header: Some(resolved.auth_header.clone()),
|
auth_header: Some(auth_header),
|
||||||
auth_value: Some(resolved.auth_value.clone()),
|
auth_value: Some(auth_value),
|
||||||
provider_api_format: resolved.provider_api_format.clone(),
|
provider_api_format,
|
||||||
client_api_format: "openai:chat".to_string(),
|
client_api_format: "openai:chat".to_string(),
|
||||||
model_name: input.requested_model.clone(),
|
model_name: input.requested_model.clone(),
|
||||||
mapped_model: resolved.mapped_model.clone(),
|
mapped_model,
|
||||||
prompt_cache_key,
|
prompt_cache_key,
|
||||||
provider_request_headers: resolved.provider_request_headers.clone(),
|
provider_request_headers,
|
||||||
provider_request_body: Some(resolved.provider_request_body.clone()),
|
provider_request_body: Some(provider_request_body),
|
||||||
provider_request_body_base64: None,
|
provider_request_body_base64: None,
|
||||||
content_type: Some("application/json".to_string()),
|
content_type: Some("application/json".to_string()),
|
||||||
proxy,
|
proxy,
|
||||||
tls_profile,
|
tls_profile,
|
||||||
timeouts,
|
timeouts,
|
||||||
upstream_is_stream,
|
upstream_is_stream,
|
||||||
report_kind: Some(resolved.report_kind.clone()),
|
report_kind: Some(report_kind),
|
||||||
report_context: Some(append_local_failover_policy_to_value(
|
report_context: Some(report_context),
|
||||||
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,
|
|
||||||
)),
|
|
||||||
auth_context: input.auth_context.clone(),
|
auth_context: input.auth_context.clone(),
|
||||||
},
|
},
|
||||||
))
|
))
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
|
|
||||||
@@ -36,7 +37,7 @@ pub(crate) struct LocalOpenAiChatCandidatePayloadParts {
|
|||||||
pub(super) execution_strategy: ExecutionStrategy,
|
pub(super) execution_strategy: ExecutionStrategy,
|
||||||
pub(super) conversion_mode: ConversionMode,
|
pub(super) conversion_mode: ConversionMode,
|
||||||
pub(super) report_kind: String,
|
pub(super) report_kind: String,
|
||||||
pub(super) transport: GatewayProviderTransportSnapshot,
|
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[allow(clippy::too_many_arguments)]
|
#[allow(clippy::too_many_arguments)]
|
||||||
@@ -192,7 +193,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
|||||||
execution_strategy: ExecutionStrategy::LocalSameFormat,
|
execution_strategy: ExecutionStrategy::LocalSameFormat,
|
||||||
conversion_mode: ConversionMode::None,
|
conversion_mode: ConversionMode::None,
|
||||||
report_kind: report_kind.to_string(),
|
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,
|
execution_strategy: ExecutionStrategy::LocalCrossFormat,
|
||||||
conversion_mode: ConversionMode::Bidirectional,
|
conversion_mode: ConversionMode::Bidirectional,
|
||||||
report_kind: resolved_report_kind,
|
report_kind: resolved_report_kind,
|
||||||
transport: transport.clone(),
|
transport: Arc::clone(transport),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -103,7 +103,7 @@ pub(crate) async fn materialize_local_openai_chat_candidate_attempts(
|
|||||||
skipped_candidate.extra_data = Some(
|
skipped_candidate.extra_data = Some(
|
||||||
build_local_execution_candidate_contract_metadata_for_candidate(
|
build_local_execution_candidate_contract_metadata_for_candidate(
|
||||||
&skipped_candidate.candidate,
|
&skipped_candidate.candidate,
|
||||||
skipped_candidate.transport.as_ref(),
|
skipped_candidate.transport_ref(),
|
||||||
provider_api_format.as_str(),
|
provider_api_format.as_str(),
|
||||||
"openai:chat",
|
"openai:chat",
|
||||||
serde_json::Map::new(),
|
serde_json::Map::new(),
|
||||||
|
|||||||
@@ -35,7 +35,6 @@ pub(crate) async fn maybe_build_local_openai_cli_decision_payload_for_candidate(
|
|||||||
let LocalOpenAiCliCandidateAttempt {
|
let LocalOpenAiCliCandidateAttempt {
|
||||||
eligible,
|
eligible,
|
||||||
candidate_index,
|
candidate_index,
|
||||||
candidate_group_id,
|
|
||||||
candidate_id,
|
candidate_id,
|
||||||
..
|
..
|
||||||
} = attempt;
|
} = attempt;
|
||||||
@@ -74,6 +73,43 @@ pub(crate) async fn maybe_build_local_openai_cli_decision_payload_for_candidate(
|
|||||||
if resolved.is_antigravity {
|
if resolved.is_antigravity {
|
||||||
extra_fields.insert("envelope_name".to_string(), json!("antigravity:v1internal"));
|
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!(
|
debug!(
|
||||||
event_name = "local_openai_cli_decision_payload_built",
|
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,
|
has_envelope = resolved.is_antigravity,
|
||||||
"gateway built local openai cli decision payload"
|
"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(
|
Some(build_local_execution_decision_response(
|
||||||
LocalExecutionDecisionResponseParts {
|
LocalExecutionDecisionResponseParts {
|
||||||
decision_is_stream: spec_metadata.require_streaming,
|
decision_is_stream: spec_metadata.require_streaming,
|
||||||
decision_kind: spec_metadata.decision_kind.to_string(),
|
decision_kind: spec_metadata.decision_kind.to_string(),
|
||||||
execution_strategy: resolved.execution_strategy,
|
execution_strategy,
|
||||||
conversion_mode: resolved.conversion_mode,
|
conversion_mode,
|
||||||
request_id: trace_id.to_string(),
|
request_id: trace_id.to_string(),
|
||||||
candidate_id: candidate_id.clone(),
|
candidate_id: candidate_id.clone(),
|
||||||
provider_name: resolved.transport.provider.name.clone(),
|
provider_name: transport.provider.name.clone(),
|
||||||
provider_id: candidate.provider_id.clone(),
|
provider_id: candidate.provider_id.clone(),
|
||||||
endpoint_id: candidate.endpoint_id.clone(),
|
endpoint_id: candidate.endpoint_id.clone(),
|
||||||
key_id: candidate.key_id.clone(),
|
key_id: candidate.key_id.clone(),
|
||||||
upstream_base_url: resolved.transport.endpoint.base_url.clone(),
|
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||||
upstream_url: resolved.upstream_url.clone(),
|
upstream_url,
|
||||||
provider_request_method: None,
|
provider_request_method: None,
|
||||||
auth_header: Some(resolved.auth_header.clone()),
|
auth_header: Some(auth_header),
|
||||||
auth_value: Some(resolved.auth_value.clone()),
|
auth_value: Some(auth_value),
|
||||||
provider_api_format: resolved.provider_api_format.clone(),
|
provider_api_format,
|
||||||
client_api_format: spec_metadata.api_format.to_string(),
|
client_api_format: spec_metadata.api_format.to_string(),
|
||||||
model_name: input.requested_model.clone(),
|
model_name: input.requested_model.clone(),
|
||||||
mapped_model: resolved.mapped_model.clone(),
|
mapped_model,
|
||||||
prompt_cache_key,
|
prompt_cache_key,
|
||||||
provider_request_headers: resolved.provider_request_headers.clone(),
|
provider_request_headers,
|
||||||
provider_request_body: Some(resolved.provider_request_body.clone()),
|
provider_request_body: Some(provider_request_body),
|
||||||
provider_request_body_base64: None,
|
provider_request_body_base64: None,
|
||||||
content_type: Some("application/json".to_string()),
|
content_type: Some("application/json".to_string()),
|
||||||
proxy,
|
proxy,
|
||||||
tls_profile,
|
tls_profile,
|
||||||
timeouts,
|
timeouts,
|
||||||
upstream_is_stream: resolved.upstream_is_stream,
|
upstream_is_stream,
|
||||||
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
||||||
report_context: Some(append_local_failover_policy_to_value(
|
report_context: Some(report_context),
|
||||||
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,
|
|
||||||
)),
|
|
||||||
auth_context: input.auth_context.clone(),
|
auth_context: input.auth_context.clone(),
|
||||||
},
|
},
|
||||||
))
|
))
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
|
use std::sync::Arc;
|
||||||
|
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
use tracing::debug;
|
use tracing::debug;
|
||||||
@@ -48,7 +49,7 @@ pub(crate) struct LocalOpenAiCliCandidatePayloadParts {
|
|||||||
pub(super) conversion_mode: ConversionMode,
|
pub(super) conversion_mode: ConversionMode,
|
||||||
pub(super) is_antigravity: bool,
|
pub(super) is_antigravity: bool,
|
||||||
pub(super) upstream_is_stream: bool,
|
pub(super) upstream_is_stream: bool,
|
||||||
pub(super) transport: GatewayProviderTransportSnapshot,
|
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[allow(clippy::too_many_arguments)]
|
#[allow(clippy::too_many_arguments)]
|
||||||
@@ -379,6 +380,6 @@ pub(crate) async fn resolve_local_openai_cli_candidate_payload_parts(
|
|||||||
is_antigravity: is_antigravity
|
is_antigravity: is_antigravity
|
||||||
|| antigravity_auth.is_some() && ANTIGRAVITY_ENVELOPE_NAME == "antigravity:v1internal",
|
|| antigravity_auth.is_some() && ANTIGRAVITY_ENVELOPE_NAME == "antigravity:v1internal",
|
||||||
upstream_is_stream,
|
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(
|
skipped_candidate.extra_data = Some(
|
||||||
build_local_execution_candidate_contract_metadata_for_candidate(
|
build_local_execution_candidate_contract_metadata_for_candidate(
|
||||||
&skipped_candidate.candidate,
|
&skipped_candidate.candidate,
|
||||||
skipped_candidate.transport.as_ref(),
|
skipped_candidate.transport_ref(),
|
||||||
provider_api_format.as_str(),
|
provider_api_format.as_str(),
|
||||||
spec_metadata.api_format,
|
spec_metadata.api_format,
|
||||||
serde_json::Map::new(),
|
serde_json::Map::new(),
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ use tracing::debug;
|
|||||||
|
|
||||||
use super::super::{
|
use super::super::{
|
||||||
augment_sync_report_context, generic_decision_missing_exact_provider_request,
|
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::provider_adaptation_requires_eventstream_accept;
|
||||||
use crate::ai_pipeline::transport::auth::{
|
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,
|
body_json: &serde_json::Value,
|
||||||
payload: GatewayControlSyncDecisionResponse,
|
payload: GatewayControlSyncDecisionResponse,
|
||||||
) -> Result<Option<LocalStreamPlanAndReport>, GatewayError> {
|
) -> Result<Option<LocalStreamPlanAndReport>, GatewayError> {
|
||||||
let Some(request_id) = payload
|
let mut payload = payload;
|
||||||
.request_id
|
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(provider_id) = payload
|
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
|
||||||
.provider_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(endpoint_id) = payload
|
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
|
||||||
.endpoint_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(key_id) = payload
|
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
|
||||||
.key_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(auth_header) = payload
|
let Some(auth_header) = take_non_empty_string(&mut payload.auth_header) else {
|
||||||
.auth_header
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(auth_value) = payload
|
let Some(auth_value) = take_non_empty_string(&mut payload.auth_value) else {
|
||||||
.auth_value
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(provider_api_format) = payload
|
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
|
||||||
.provider_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(client_api_format) = payload
|
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
|
||||||
.client_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let url = if let Some(upstream_url) = payload
|
let url = if let Some(upstream_url) = take_non_empty_string(&mut payload.upstream_url) {
|
||||||
.upstream_url
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
{
|
|
||||||
upstream_url
|
upstream_url
|
||||||
} else {
|
} else {
|
||||||
let Some(upstream_base_url) = payload
|
let Some(upstream_base_url) = take_non_empty_string(&mut payload.upstream_base_url) else {
|
||||||
.upstream_base_url
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
build_openai_chat_url(&upstream_base_url, parts.uri.query())
|
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
|
body
|
||||||
} else {
|
} else {
|
||||||
let Some(request_body_object) = body_json.as_object() 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()
|
.iter()
|
||||||
.map(|(key, value)| (key.clone(), value.clone())),
|
.map(|(key, value)| (key.clone(), value.clone())),
|
||||||
);
|
);
|
||||||
if let Some(mapped_model) = payload
|
if let Some(mapped_model) = take_non_empty_string(&mut payload.mapped_model) {
|
||||||
.mapped_model
|
provider_request_body
|
||||||
.as_ref()
|
.insert("model".to_string(), serde_json::Value::String(mapped_model));
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
{
|
|
||||||
provider_request_body.insert(
|
|
||||||
"model".to_string(),
|
|
||||||
serde_json::Value::String(mapped_model.clone()),
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
provider_request_body.insert("stream".to_string(), serde_json::Value::Bool(true));
|
provider_request_body.insert("stream".to_string(), serde_json::Value::Bool(true));
|
||||||
if let Some(prompt_cache_key) = payload
|
if let Some(prompt_cache_key) = take_non_empty_string(&mut payload.prompt_cache_key) {
|
||||||
.prompt_cache_key
|
|
||||||
.as_ref()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
{
|
|
||||||
let existing = provider_request_body
|
let existing = provider_request_body
|
||||||
.get("prompt_cache_key")
|
.get("prompt_cache_key")
|
||||||
.and_then(|value| value.as_str())
|
.and_then(|value| value.as_str())
|
||||||
@@ -126,20 +77,21 @@ pub(crate) fn build_openai_chat_stream_plan_from_decision(
|
|||||||
if existing.is_empty() {
|
if existing.is_empty() {
|
||||||
provider_request_body.insert(
|
provider_request_body.insert(
|
||||||
"prompt_cache_key".to_string(),
|
"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)
|
serde_json::Value::Object(provider_request_body)
|
||||||
};
|
};
|
||||||
|
let existing_provider_request_headers = std::mem::take(&mut payload.provider_request_headers);
|
||||||
let mut provider_request_headers = if payload.provider_request_headers.is_empty() {
|
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 {
|
if provider_api_format == client_api_format {
|
||||||
build_complete_passthrough_headers_with_auth(
|
build_complete_passthrough_headers_with_auth(
|
||||||
&parts.headers,
|
&parts.headers,
|
||||||
&auth_header,
|
&auth_header,
|
||||||
&auth_value,
|
&auth_value,
|
||||||
&payload.extra_headers,
|
&extra_headers,
|
||||||
payload.content_type.as_deref(),
|
payload.content_type.as_deref(),
|
||||||
)
|
)
|
||||||
} else if provider_api_format.starts_with("claude:") {
|
} else if provider_api_format.starts_with("claude:") {
|
||||||
@@ -147,7 +99,7 @@ pub(crate) fn build_openai_chat_stream_plan_from_decision(
|
|||||||
&parts.headers,
|
&parts.headers,
|
||||||
&auth_header,
|
&auth_header,
|
||||||
&auth_value,
|
&auth_value,
|
||||||
&payload.extra_headers,
|
&extra_headers,
|
||||||
payload.content_type.as_deref(),
|
payload.content_type.as_deref(),
|
||||||
)
|
)
|
||||||
} else {
|
} else {
|
||||||
@@ -155,46 +107,46 @@ pub(crate) fn build_openai_chat_stream_plan_from_decision(
|
|||||||
&parts.headers,
|
&parts.headers,
|
||||||
&auth_header,
|
&auth_header,
|
||||||
&auth_value,
|
&auth_value,
|
||||||
&payload.extra_headers,
|
&extra_headers,
|
||||||
payload.content_type.as_deref(),
|
payload.content_type.as_deref(),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
payload.provider_request_headers.clone()
|
existing_provider_request_headers
|
||||||
};
|
};
|
||||||
ensure_upstream_auth_header(&mut provider_request_headers, &auth_header, &auth_value);
|
ensure_upstream_auth_header(&mut provider_request_headers, &auth_header, &auth_value);
|
||||||
provider_request_headers.insert("accept".to_string(), "text/event-stream".to_string());
|
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 {
|
let plan = ExecutionPlan {
|
||||||
request_id,
|
request_id,
|
||||||
candidate_id: payload.candidate_id.clone(),
|
candidate_id: payload.candidate_id.take(),
|
||||||
provider_name: payload.provider_name.clone(),
|
provider_name: payload.provider_name.take(),
|
||||||
provider_id,
|
provider_id,
|
||||||
endpoint_id,
|
endpoint_id,
|
||||||
key_id,
|
key_id,
|
||||||
method: "POST".to_string(),
|
method: "POST".to_string(),
|
||||||
url,
|
url,
|
||||||
headers: std::mem::take(&mut provider_request_headers),
|
headers: std::mem::take(&mut provider_request_headers),
|
||||||
content_type: payload
|
content_type,
|
||||||
.content_type
|
|
||||||
.clone()
|
|
||||||
.or_else(|| Some("application/json".to_string())),
|
|
||||||
content_encoding: None,
|
content_encoding: None,
|
||||||
body: RequestBody::from_json(provider_request_body_value.clone()),
|
body: RequestBody::from_json(provider_request_body_value),
|
||||||
stream: true,
|
stream: true,
|
||||||
client_api_format,
|
client_api_format,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
model_name: payload.model_name.clone(),
|
model_name: payload.model_name.take(),
|
||||||
proxy: payload.proxy.clone(),
|
proxy: payload.proxy.take(),
|
||||||
tls_profile: payload.tls_profile.clone(),
|
tls_profile: payload.tls_profile.take(),
|
||||||
timeouts: payload.timeouts.clone(),
|
timeouts: payload.timeouts.take(),
|
||||||
};
|
};
|
||||||
|
|
||||||
let report_context = augment_sync_report_context(
|
|
||||||
payload.report_context,
|
|
||||||
&plan.headers,
|
|
||||||
&provider_request_body_value,
|
|
||||||
)?;
|
|
||||||
|
|
||||||
Ok(Some(LocalStreamPlanAndReport {
|
Ok(Some(LocalStreamPlanAndReport {
|
||||||
plan,
|
plan,
|
||||||
report_kind: payload.report_kind,
|
report_kind: payload.report_kind,
|
||||||
@@ -208,74 +160,39 @@ pub(crate) fn build_openai_cli_stream_plan_from_decision(
|
|||||||
payload: GatewayControlSyncDecisionResponse,
|
payload: GatewayControlSyncDecisionResponse,
|
||||||
compact: bool,
|
compact: bool,
|
||||||
) -> Result<Option<LocalStreamPlanAndReport>, GatewayError> {
|
) -> Result<Option<LocalStreamPlanAndReport>, GatewayError> {
|
||||||
|
let mut payload = payload;
|
||||||
if generic_decision_missing_exact_provider_request(&payload) {
|
if generic_decision_missing_exact_provider_request(&payload) {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
let Some(request_id) = payload
|
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
|
||||||
.request_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(provider_id) = payload
|
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
|
||||||
.provider_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(endpoint_id) = payload
|
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
|
||||||
.endpoint_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(key_id) = payload
|
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
|
||||||
.key_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let auth_header = payload
|
let auth_header = take_non_empty_string(&mut payload.auth_header);
|
||||||
.auth_header
|
let auth_value = take_non_empty_string(&mut payload.auth_value);
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty());
|
|
||||||
let auth_value = payload
|
|
||||||
.auth_value
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty());
|
|
||||||
if auth_header.is_some() != auth_value.is_some() {
|
if auth_header.is_some() != auth_value.is_some() {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
let Some(provider_api_format) = payload
|
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
|
||||||
.provider_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(client_api_format) = payload
|
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
|
||||||
.client_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let (url, url_source) = if let Some(upstream_url) = payload
|
let (url, url_source) = if let Some(upstream_url) =
|
||||||
.upstream_url
|
take_non_empty_string(&mut payload.upstream_url)
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
{
|
{
|
||||||
(upstream_url, "upstream_url")
|
(upstream_url, "upstream_url")
|
||||||
} else {
|
} else {
|
||||||
let Some(upstream_base_url) = payload
|
let Some(upstream_base_url) = take_non_empty_string(&mut payload.upstream_base_url) else {
|
||||||
.upstream_base_url
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
(
|
(
|
||||||
@@ -283,7 +200,7 @@ pub(crate) fn build_openai_cli_stream_plan_from_decision(
|
|||||||
"upstream_base_url",
|
"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);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -292,7 +209,7 @@ pub(crate) fn build_openai_cli_stream_plan_from_decision(
|
|||||||
.as_ref()
|
.as_ref()
|
||||||
.and_then(|context| context.get("envelope_name"))
|
.and_then(|context| context.get("envelope_name"))
|
||||||
.and_then(serde_json::Value::as_str);
|
.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()) {
|
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);
|
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 {
|
} else {
|
||||||
provider_request_headers.insert("accept".to_string(), "text/event-stream".to_string());
|
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(
|
let report_context = augment_sync_report_context(
|
||||||
payload.report_context,
|
payload.report_context.take(),
|
||||||
&provider_request_headers,
|
&provider_request_headers,
|
||||||
&provider_request_body_value,
|
&provider_request_body_value,
|
||||||
)?;
|
)?;
|
||||||
let plan = ExecutionPlan {
|
let plan = ExecutionPlan {
|
||||||
request_id,
|
request_id,
|
||||||
candidate_id: payload.candidate_id.clone(),
|
candidate_id: payload.candidate_id.take(),
|
||||||
provider_name: payload.provider_name.clone(),
|
provider_name: payload.provider_name.take(),
|
||||||
provider_id,
|
provider_id,
|
||||||
endpoint_id,
|
endpoint_id,
|
||||||
key_id,
|
key_id,
|
||||||
method: "POST".to_string(),
|
method: "POST".to_string(),
|
||||||
url,
|
url,
|
||||||
headers: std::mem::take(&mut provider_request_headers),
|
headers: std::mem::take(&mut provider_request_headers),
|
||||||
content_type: payload
|
content_type,
|
||||||
.content_type
|
|
||||||
.clone()
|
|
||||||
.or_else(|| Some("application/json".to_string())),
|
|
||||||
content_encoding: None,
|
content_encoding: None,
|
||||||
body: RequestBody::from_json(provider_request_body_value.clone()),
|
body: RequestBody::from_json(provider_request_body_value),
|
||||||
stream: true,
|
stream: true,
|
||||||
client_api_format,
|
client_api_format,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
model_name: payload.model_name.clone(),
|
model_name: payload.model_name.take(),
|
||||||
proxy: payload.proxy.clone(),
|
proxy: payload.proxy.take(),
|
||||||
tls_profile: payload.tls_profile.clone(),
|
tls_profile: payload.tls_profile.take(),
|
||||||
timeouts: payload.timeouts.clone(),
|
timeouts: payload.timeouts.take(),
|
||||||
};
|
};
|
||||||
|
|
||||||
debug!(
|
debug!(
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ use tracing::debug;
|
|||||||
|
|
||||||
use super::super::{
|
use super::super::{
|
||||||
augment_sync_report_context, generic_decision_missing_exact_provider_request,
|
augment_sync_report_context, generic_decision_missing_exact_provider_request,
|
||||||
LocalSyncPlanAndReport,
|
take_non_empty_string, LocalSyncPlanAndReport,
|
||||||
};
|
};
|
||||||
use crate::ai_pipeline::transport::auth::{
|
use crate::ai_pipeline::transport::auth::{
|
||||||
build_claude_passthrough_headers, build_complete_passthrough_headers_with_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,
|
body_json: &serde_json::Value,
|
||||||
payload: GatewayControlSyncDecisionResponse,
|
payload: GatewayControlSyncDecisionResponse,
|
||||||
) -> Result<Option<LocalSyncPlanAndReport>, GatewayError> {
|
) -> Result<Option<LocalSyncPlanAndReport>, GatewayError> {
|
||||||
let Some(request_id) = payload
|
let mut payload = payload;
|
||||||
.request_id
|
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(provider_id) = payload
|
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
|
||||||
.provider_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(endpoint_id) = payload
|
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
|
||||||
.endpoint_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(key_id) = payload
|
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
|
||||||
.key_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(auth_header) = payload
|
let Some(auth_header) = take_non_empty_string(&mut payload.auth_header) else {
|
||||||
.auth_header
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(auth_value) = payload
|
let Some(auth_value) = take_non_empty_string(&mut payload.auth_value) else {
|
||||||
.auth_value
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(provider_api_format) = payload
|
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
|
||||||
.provider_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(client_api_format) = payload
|
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
|
||||||
.client_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let url = if let Some(upstream_url) = payload
|
let url = if let Some(upstream_url) = take_non_empty_string(&mut payload.upstream_url) {
|
||||||
.upstream_url
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
{
|
|
||||||
upstream_url
|
upstream_url
|
||||||
} else {
|
} else {
|
||||||
let Some(upstream_base_url) = payload
|
let Some(upstream_base_url) = take_non_empty_string(&mut payload.upstream_base_url) else {
|
||||||
.upstream_base_url
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
build_openai_chat_url(&upstream_base_url, parts.uri.query())
|
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
|
body
|
||||||
} else {
|
} else {
|
||||||
let Some(request_body_object) = body_json.as_object() 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()
|
.iter()
|
||||||
.map(|(key, value)| (key.clone(), value.clone())),
|
.map(|(key, value)| (key.clone(), value.clone())),
|
||||||
);
|
);
|
||||||
if let Some(mapped_model) = payload
|
if let Some(mapped_model) = take_non_empty_string(&mut payload.mapped_model) {
|
||||||
.mapped_model
|
provider_request_body
|
||||||
.as_ref()
|
.insert("model".to_string(), serde_json::Value::String(mapped_model));
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
{
|
|
||||||
provider_request_body.insert(
|
|
||||||
"model".to_string(),
|
|
||||||
serde_json::Value::String(mapped_model.clone()),
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
if payload.upstream_is_stream {
|
if payload.upstream_is_stream {
|
||||||
provider_request_body.insert("stream".to_string(), serde_json::Value::Bool(true));
|
provider_request_body.insert("stream".to_string(), serde_json::Value::Bool(true));
|
||||||
}
|
}
|
||||||
if let Some(prompt_cache_key) = payload
|
if let Some(prompt_cache_key) = take_non_empty_string(&mut payload.prompt_cache_key) {
|
||||||
.prompt_cache_key
|
|
||||||
.as_ref()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
{
|
|
||||||
let existing = provider_request_body
|
let existing = provider_request_body
|
||||||
.get("prompt_cache_key")
|
.get("prompt_cache_key")
|
||||||
.and_then(|value| value.as_str())
|
.and_then(|value| value.as_str())
|
||||||
@@ -126,20 +77,21 @@ pub(crate) fn build_openai_chat_sync_plan_from_decision(
|
|||||||
if existing.is_empty() {
|
if existing.is_empty() {
|
||||||
provider_request_body.insert(
|
provider_request_body.insert(
|
||||||
"prompt_cache_key".to_string(),
|
"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)
|
serde_json::Value::Object(provider_request_body)
|
||||||
};
|
};
|
||||||
|
let existing_provider_request_headers = std::mem::take(&mut payload.provider_request_headers);
|
||||||
let mut provider_request_headers = if payload.provider_request_headers.is_empty() {
|
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 {
|
if provider_api_format == client_api_format {
|
||||||
build_complete_passthrough_headers_with_auth(
|
build_complete_passthrough_headers_with_auth(
|
||||||
&parts.headers,
|
&parts.headers,
|
||||||
&auth_header,
|
&auth_header,
|
||||||
&auth_value,
|
&auth_value,
|
||||||
&payload.extra_headers,
|
&extra_headers,
|
||||||
payload.content_type.as_deref(),
|
payload.content_type.as_deref(),
|
||||||
)
|
)
|
||||||
} else if provider_api_format.starts_with("claude:") {
|
} else if provider_api_format.starts_with("claude:") {
|
||||||
@@ -147,7 +99,7 @@ pub(crate) fn build_openai_chat_sync_plan_from_decision(
|
|||||||
&parts.headers,
|
&parts.headers,
|
||||||
&auth_header,
|
&auth_header,
|
||||||
&auth_value,
|
&auth_value,
|
||||||
&payload.extra_headers,
|
&extra_headers,
|
||||||
payload.content_type.as_deref(),
|
payload.content_type.as_deref(),
|
||||||
)
|
)
|
||||||
} else {
|
} else {
|
||||||
@@ -155,12 +107,12 @@ pub(crate) fn build_openai_chat_sync_plan_from_decision(
|
|||||||
&parts.headers,
|
&parts.headers,
|
||||||
&auth_header,
|
&auth_header,
|
||||||
&auth_value,
|
&auth_value,
|
||||||
&payload.extra_headers,
|
&extra_headers,
|
||||||
payload.content_type.as_deref(),
|
payload.content_type.as_deref(),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
payload.provider_request_headers.clone()
|
existing_provider_request_headers
|
||||||
};
|
};
|
||||||
ensure_upstream_auth_header(&mut provider_request_headers, &auth_header, &auth_value);
|
ensure_upstream_auth_header(&mut provider_request_headers, &auth_header, &auth_value);
|
||||||
if payload.upstream_is_stream {
|
if payload.upstream_is_stream {
|
||||||
@@ -168,37 +120,37 @@ pub(crate) fn build_openai_chat_sync_plan_from_decision(
|
|||||||
.entry("accept".to_string())
|
.entry("accept".to_string())
|
||||||
.or_insert_with(|| "text/event-stream".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 {
|
let plan = ExecutionPlan {
|
||||||
request_id,
|
request_id,
|
||||||
candidate_id: payload.candidate_id.clone(),
|
candidate_id: payload.candidate_id.take(),
|
||||||
provider_name: payload.provider_name.clone(),
|
provider_name: payload.provider_name.take(),
|
||||||
provider_id,
|
provider_id,
|
||||||
endpoint_id,
|
endpoint_id,
|
||||||
key_id,
|
key_id,
|
||||||
method: "POST".to_string(),
|
method: "POST".to_string(),
|
||||||
url,
|
url,
|
||||||
headers: std::mem::take(&mut provider_request_headers),
|
headers: std::mem::take(&mut provider_request_headers),
|
||||||
content_type: payload
|
content_type,
|
||||||
.content_type
|
|
||||||
.clone()
|
|
||||||
.or_else(|| Some("application/json".to_string())),
|
|
||||||
content_encoding: None,
|
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,
|
stream: payload.upstream_is_stream,
|
||||||
client_api_format,
|
client_api_format,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
model_name: payload.model_name.clone(),
|
model_name: payload.model_name.take(),
|
||||||
proxy: payload.proxy.clone(),
|
proxy: payload.proxy.take(),
|
||||||
tls_profile: payload.tls_profile.clone(),
|
tls_profile: payload.tls_profile.take(),
|
||||||
timeouts: payload.timeouts.clone(),
|
timeouts: payload.timeouts.take(),
|
||||||
};
|
};
|
||||||
|
|
||||||
let report_context = augment_sync_report_context(
|
|
||||||
payload.report_context,
|
|
||||||
&plan.headers,
|
|
||||||
&provider_request_body_value,
|
|
||||||
)?;
|
|
||||||
|
|
||||||
Ok(Some(LocalSyncPlanAndReport {
|
Ok(Some(LocalSyncPlanAndReport {
|
||||||
plan,
|
plan,
|
||||||
report_kind: payload.report_kind,
|
report_kind: payload.report_kind,
|
||||||
@@ -212,74 +164,39 @@ pub(crate) fn build_openai_cli_sync_plan_from_decision(
|
|||||||
payload: GatewayControlSyncDecisionResponse,
|
payload: GatewayControlSyncDecisionResponse,
|
||||||
compact: bool,
|
compact: bool,
|
||||||
) -> Result<Option<LocalSyncPlanAndReport>, GatewayError> {
|
) -> Result<Option<LocalSyncPlanAndReport>, GatewayError> {
|
||||||
|
let mut payload = payload;
|
||||||
if generic_decision_missing_exact_provider_request(&payload) {
|
if generic_decision_missing_exact_provider_request(&payload) {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
let Some(request_id) = payload
|
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
|
||||||
.request_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(provider_id) = payload
|
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
|
||||||
.provider_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(endpoint_id) = payload
|
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
|
||||||
.endpoint_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(key_id) = payload
|
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
|
||||||
.key_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let auth_header = payload
|
let auth_header = take_non_empty_string(&mut payload.auth_header);
|
||||||
.auth_header
|
let auth_value = take_non_empty_string(&mut payload.auth_value);
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty());
|
|
||||||
let auth_value = payload
|
|
||||||
.auth_value
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty());
|
|
||||||
if auth_header.is_some() != auth_value.is_some() {
|
if auth_header.is_some() != auth_value.is_some() {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
let Some(provider_api_format) = payload
|
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
|
||||||
.provider_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(client_api_format) = payload
|
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
|
||||||
.client_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let (url, url_source) = if let Some(upstream_url) = payload
|
let (url, url_source) = if let Some(upstream_url) =
|
||||||
.upstream_url
|
take_non_empty_string(&mut payload.upstream_url)
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
{
|
{
|
||||||
(upstream_url, "upstream_url")
|
(upstream_url, "upstream_url")
|
||||||
} else {
|
} else {
|
||||||
let Some(upstream_base_url) = payload
|
let Some(upstream_base_url) = take_non_empty_string(&mut payload.upstream_base_url) else {
|
||||||
.upstream_base_url
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
(
|
(
|
||||||
@@ -287,45 +204,45 @@ pub(crate) fn build_openai_cli_sync_plan_from_decision(
|
|||||||
"upstream_base_url",
|
"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);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
|
let mut provider_request_headers = std::mem::take(&mut payload.provider_request_headers);
|
||||||
let mut provider_request_headers = payload.provider_request_headers.clone();
|
|
||||||
if let (Some(auth_header), Some(auth_value)) = (auth_header.as_deref(), auth_value.as_deref()) {
|
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);
|
ensure_upstream_auth_header(&mut provider_request_headers, auth_header, auth_value);
|
||||||
}
|
}
|
||||||
if payload.upstream_is_stream && !provider_request_headers.contains_key("accept") {
|
if payload.upstream_is_stream && !provider_request_headers.contains_key("accept") {
|
||||||
provider_request_headers.insert("accept".to_string(), "text/event-stream".to_string());
|
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(
|
let report_context = augment_sync_report_context(
|
||||||
payload.report_context,
|
payload.report_context.take(),
|
||||||
&provider_request_headers,
|
&provider_request_headers,
|
||||||
&provider_request_body_value,
|
&provider_request_body_value,
|
||||||
)?;
|
)?;
|
||||||
let plan = ExecutionPlan {
|
let plan = ExecutionPlan {
|
||||||
request_id,
|
request_id,
|
||||||
candidate_id: payload.candidate_id.clone(),
|
candidate_id: payload.candidate_id.take(),
|
||||||
provider_name: payload.provider_name.clone(),
|
provider_name: payload.provider_name.take(),
|
||||||
provider_id,
|
provider_id,
|
||||||
endpoint_id,
|
endpoint_id,
|
||||||
key_id,
|
key_id,
|
||||||
method: "POST".to_string(),
|
method: "POST".to_string(),
|
||||||
url,
|
url,
|
||||||
headers: std::mem::take(&mut provider_request_headers),
|
headers: std::mem::take(&mut provider_request_headers),
|
||||||
content_type: payload
|
content_type,
|
||||||
.content_type
|
|
||||||
.clone()
|
|
||||||
.or_else(|| Some("application/json".to_string())),
|
|
||||||
content_encoding: None,
|
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,
|
stream: payload.upstream_is_stream,
|
||||||
client_api_format,
|
client_api_format,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
model_name: payload.model_name.clone(),
|
model_name: payload.model_name.take(),
|
||||||
proxy: payload.proxy.clone(),
|
proxy: payload.proxy.take(),
|
||||||
tls_profile: payload.tls_profile.clone(),
|
tls_profile: payload.tls_profile.take(),
|
||||||
timeouts: payload.timeouts.clone(),
|
timeouts: payload.timeouts.take(),
|
||||||
};
|
};
|
||||||
|
|
||||||
debug!(
|
debug!(
|
||||||
|
|||||||
@@ -1,6 +1,9 @@
|
|||||||
use aether_contracts::{ExecutionPlan, RequestBody};
|
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::contracts::generic_decision_missing_exact_provider_request;
|
||||||
use crate::ai_pipeline::provider_adaptation_requires_eventstream_accept;
|
use crate::ai_pipeline::provider_adaptation_requires_eventstream_accept;
|
||||||
use crate::ai_pipeline::transport::ensure_upstream_auth_header;
|
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,
|
_body_json: &serde_json::Value,
|
||||||
payload: GatewayControlSyncDecisionResponse,
|
payload: GatewayControlSyncDecisionResponse,
|
||||||
) -> Result<Option<LocalSyncPlanAndReport>, GatewayError> {
|
) -> Result<Option<LocalSyncPlanAndReport>, GatewayError> {
|
||||||
|
let mut payload = payload;
|
||||||
if generic_decision_missing_exact_provider_request(&payload) {
|
if generic_decision_missing_exact_provider_request(&payload) {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
let Some(request_id) = payload
|
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
|
||||||
.request_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(provider_id) = payload
|
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
|
||||||
.provider_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(endpoint_id) = payload
|
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
|
||||||
.endpoint_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(key_id) = payload
|
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
|
||||||
.key_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(url) = payload
|
let Some(url) = take_non_empty_string(&mut payload.upstream_url) else {
|
||||||
.upstream_url
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let auth_header = payload
|
let auth_header = take_non_empty_string(&mut payload.auth_header);
|
||||||
.auth_header
|
let auth_value = take_non_empty_string(&mut payload.auth_value);
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty());
|
|
||||||
let auth_value = payload
|
|
||||||
.auth_value
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty());
|
|
||||||
if auth_header.is_some() != auth_value.is_some() {
|
if auth_header.is_some() != auth_value.is_some() {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
let Some(provider_api_format) = payload
|
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
|
||||||
.provider_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(client_api_format) = payload
|
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
|
||||||
.client_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
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);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
|
let mut provider_request_headers = std::mem::take(&mut payload.provider_request_headers);
|
||||||
let mut provider_request_headers = payload.provider_request_headers.clone();
|
|
||||||
if let (Some(auth_header), Some(auth_value)) = (auth_header.as_deref(), auth_value.as_deref()) {
|
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);
|
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())
|
.entry("accept".to_string())
|
||||||
.or_insert_with(|| "text/event-stream".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 {
|
let plan = ExecutionPlan {
|
||||||
request_id,
|
request_id,
|
||||||
candidate_id: payload.candidate_id.clone(),
|
candidate_id: payload.candidate_id.take(),
|
||||||
provider_name: payload.provider_name.clone(),
|
provider_name: payload.provider_name.take(),
|
||||||
provider_id,
|
provider_id,
|
||||||
endpoint_id,
|
endpoint_id,
|
||||||
key_id,
|
key_id,
|
||||||
method: "POST".to_string(),
|
method: "POST".to_string(),
|
||||||
url,
|
url,
|
||||||
headers: std::mem::take(&mut provider_request_headers),
|
headers: std::mem::take(&mut provider_request_headers),
|
||||||
content_type: payload
|
content_type,
|
||||||
.content_type
|
|
||||||
.clone()
|
|
||||||
.or_else(|| Some("application/json".to_string())),
|
|
||||||
content_encoding: None,
|
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,
|
stream: payload.upstream_is_stream,
|
||||||
client_api_format,
|
client_api_format,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
model_name: payload.model_name.clone(),
|
model_name: payload.model_name.take(),
|
||||||
proxy: payload.proxy.clone(),
|
proxy: payload.proxy.take(),
|
||||||
tls_profile: payload.tls_profile.clone(),
|
tls_profile: payload.tls_profile.take(),
|
||||||
timeouts: payload.timeouts.clone(),
|
timeouts: payload.timeouts.take(),
|
||||||
};
|
};
|
||||||
|
|
||||||
let report_context = augment_sync_report_context(
|
|
||||||
payload.report_context,
|
|
||||||
&plan.headers,
|
|
||||||
&provider_request_body_value,
|
|
||||||
)?;
|
|
||||||
|
|
||||||
Ok(Some(LocalSyncPlanAndReport {
|
Ok(Some(LocalSyncPlanAndReport {
|
||||||
plan,
|
plan,
|
||||||
report_kind: payload.report_kind,
|
report_kind: payload.report_kind,
|
||||||
@@ -131,70 +100,37 @@ pub(crate) fn build_standard_stream_plan_from_decision(
|
|||||||
payload: GatewayControlSyncDecisionResponse,
|
payload: GatewayControlSyncDecisionResponse,
|
||||||
_inject_stream_flag: bool,
|
_inject_stream_flag: bool,
|
||||||
) -> Result<Option<LocalStreamPlanAndReport>, GatewayError> {
|
) -> Result<Option<LocalStreamPlanAndReport>, GatewayError> {
|
||||||
|
let mut payload = payload;
|
||||||
if generic_decision_missing_exact_provider_request(&payload) {
|
if generic_decision_missing_exact_provider_request(&payload) {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
let Some(request_id) = payload
|
let Some(request_id) = take_non_empty_string(&mut payload.request_id) else {
|
||||||
.request_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(provider_id) = payload
|
let Some(provider_id) = take_non_empty_string(&mut payload.provider_id) else {
|
||||||
.provider_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(endpoint_id) = payload
|
let Some(endpoint_id) = take_non_empty_string(&mut payload.endpoint_id) else {
|
||||||
.endpoint_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(key_id) = payload
|
let Some(key_id) = take_non_empty_string(&mut payload.key_id) else {
|
||||||
.key_id
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(url) = payload
|
let Some(url) = take_non_empty_string(&mut payload.upstream_url) else {
|
||||||
.upstream_url
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let auth_header = payload
|
let auth_header = take_non_empty_string(&mut payload.auth_header);
|
||||||
.auth_header
|
let auth_value = take_non_empty_string(&mut payload.auth_value);
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty());
|
|
||||||
let auth_value = payload
|
|
||||||
.auth_value
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty());
|
|
||||||
if auth_header.is_some() != auth_value.is_some() {
|
if auth_header.is_some() != auth_value.is_some() {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
let Some(provider_api_format) = payload
|
let Some(provider_api_format) = take_non_empty_string(&mut payload.provider_api_format) else {
|
||||||
.provider_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
let Some(client_api_format) = payload
|
let Some(client_api_format) = take_non_empty_string(&mut payload.client_api_format) else {
|
||||||
.client_api_format
|
|
||||||
.clone()
|
|
||||||
.filter(|value| !value.trim().is_empty())
|
|
||||||
else {
|
|
||||||
return Ok(None);
|
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);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -203,7 +139,7 @@ pub(crate) fn build_standard_stream_plan_from_decision(
|
|||||||
.as_ref()
|
.as_ref()
|
||||||
.and_then(|context| context.get("envelope_name"))
|
.and_then(|context| context.get("envelope_name"))
|
||||||
.and_then(serde_json::Value::as_str);
|
.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()) {
|
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);
|
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 {
|
} else {
|
||||||
provider_request_headers.insert("accept".to_string(), "text/event-stream".to_string());
|
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 {
|
let plan = ExecutionPlan {
|
||||||
request_id,
|
request_id,
|
||||||
candidate_id: payload.candidate_id.clone(),
|
candidate_id: payload.candidate_id.take(),
|
||||||
provider_name: payload.provider_name.clone(),
|
provider_name: payload.provider_name.take(),
|
||||||
provider_id,
|
provider_id,
|
||||||
endpoint_id,
|
endpoint_id,
|
||||||
key_id,
|
key_id,
|
||||||
method: "POST".to_string(),
|
method: "POST".to_string(),
|
||||||
url,
|
url,
|
||||||
headers: std::mem::take(&mut provider_request_headers),
|
headers: std::mem::take(&mut provider_request_headers),
|
||||||
content_type: payload
|
content_type,
|
||||||
.content_type
|
|
||||||
.clone()
|
|
||||||
.or_else(|| Some("application/json".to_string())),
|
|
||||||
content_encoding: None,
|
content_encoding: None,
|
||||||
body: RequestBody::from_json(provider_request_body_value.clone()),
|
body: RequestBody::from_json(provider_request_body_value),
|
||||||
stream: true,
|
stream: true,
|
||||||
client_api_format,
|
client_api_format,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
model_name: payload.model_name.clone(),
|
model_name: payload.model_name.take(),
|
||||||
proxy: payload.proxy.clone(),
|
proxy: payload.proxy.take(),
|
||||||
tls_profile: payload.tls_profile.clone(),
|
tls_profile: payload.tls_profile.take(),
|
||||||
timeouts: payload.timeouts.clone(),
|
timeouts: payload.timeouts.take(),
|
||||||
};
|
};
|
||||||
|
|
||||||
let report_context = augment_sync_report_context(
|
|
||||||
payload.report_context,
|
|
||||||
&plan.headers,
|
|
||||||
&provider_request_body_value,
|
|
||||||
)?;
|
|
||||||
|
|
||||||
Ok(Some(LocalStreamPlanAndReport {
|
Ok(Some(LocalStreamPlanAndReport {
|
||||||
plan,
|
plan,
|
||||||
report_kind: payload.report_kind,
|
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_openai_cli_request_to_openai_chat_request, normalize_provider_private_report_context,
|
||||||
normalize_provider_private_response_value, normalize_standard_request_to_openai_chat_request,
|
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,
|
parse_direct_request_body, parse_openai_stop_sequences, parse_openai_tool_result_content,
|
||||||
prepare_local_success_response_parts, provider_adaptation_allows_sync_finalize_envelope,
|
prepare_local_success_response_parts, prepare_local_success_response_parts_owned,
|
||||||
provider_adaptation_anchor_api_format, provider_adaptation_descriptor_for_envelope,
|
provider_adaptation_allows_sync_finalize_envelope, provider_adaptation_anchor_api_format,
|
||||||
provider_adaptation_descriptor_for_provider_type,
|
provider_adaptation_descriptor_for_envelope, provider_adaptation_descriptor_for_provider_type,
|
||||||
provider_adaptation_requires_eventstream_accept,
|
provider_adaptation_requires_eventstream_accept,
|
||||||
provider_adaptation_should_unwrap_stream_envelope,
|
provider_adaptation_should_unwrap_stream_envelope,
|
||||||
provider_private_response_allows_sync_finalize, request_candidate_api_format_preference,
|
provider_private_response_allows_sync_finalize, request_candidate_api_format_preference,
|
||||||
|
|||||||
@@ -84,6 +84,27 @@ pub(crate) fn build_client_response_from_parts(
|
|||||||
trace_id: &str,
|
trace_id: &str,
|
||||||
control_decision: Option<&GatewayControlDecision>,
|
control_decision: Option<&GatewayControlDecision>,
|
||||||
) -> Result<Response<Body>, GatewayError> {
|
) -> Result<Response<Body>, GatewayError> {
|
||||||
|
build_client_response_from_parts_with_mutator(
|
||||||
|
status_code,
|
||||||
|
upstream_headers,
|
||||||
|
body,
|
||||||
|
trace_id,
|
||||||
|
control_decision,
|
||||||
|
|_| Ok(()),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(crate) fn build_client_response_from_parts_with_mutator<F>(
|
||||||
|
status_code: u16,
|
||||||
|
upstream_headers: &BTreeMap<String, String>,
|
||||||
|
body: Body,
|
||||||
|
trace_id: &str,
|
||||||
|
control_decision: Option<&GatewayControlDecision>,
|
||||||
|
mutate_headers: F,
|
||||||
|
) -> Result<Response<Body>, GatewayError>
|
||||||
|
where
|
||||||
|
F: FnOnce(&mut http::HeaderMap) -> Result<(), GatewayError>,
|
||||||
|
{
|
||||||
let mut response = Response::builder()
|
let mut response = Response::builder()
|
||||||
.status(status_code)
|
.status(status_code)
|
||||||
.body(body)
|
.body(body)
|
||||||
@@ -99,6 +120,7 @@ pub(crate) fn build_client_response_from_parts(
|
|||||||
HeaderValue::from_str(value).map_err(|err| GatewayError::Internal(err.to_string()))?;
|
HeaderValue::from_str(value).map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||||
response.headers_mut().insert(header_name, header_value);
|
response.headers_mut().insert(header_name, header_value);
|
||||||
}
|
}
|
||||||
|
mutate_headers(response.headers_mut())?;
|
||||||
apply_streaming_response_headers(response.headers_mut());
|
apply_streaming_response_headers(response.headers_mut());
|
||||||
insert_header_if_missing(response.headers_mut(), TRACE_ID_HEADER, trace_id)?;
|
insert_header_if_missing(response.headers_mut(), TRACE_ID_HEADER, trace_id)?;
|
||||||
insert_header_if_missing(response.headers_mut(), GATEWAY_HEADER, "rust-phase3b")?;
|
insert_header_if_missing(response.headers_mut(), GATEWAY_HEADER, "rust-phase3b")?;
|
||||||
|
|||||||
@@ -42,7 +42,8 @@ pub(crate) use sync::{
|
|||||||
execute_execution_runtime_sync, maybe_build_local_sync_finalize_response,
|
execute_execution_runtime_sync, maybe_build_local_sync_finalize_response,
|
||||||
maybe_build_local_video_error_response, maybe_build_local_video_success_outcome,
|
maybe_build_local_video_error_response, maybe_build_local_video_success_outcome,
|
||||||
resolve_local_sync_error_background_report_kind,
|
resolve_local_sync_error_background_report_kind,
|
||||||
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessOutcome,
|
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessBuild,
|
||||||
|
LocalVideoSyncSuccessOutcome,
|
||||||
};
|
};
|
||||||
pub(crate) use transport::{
|
pub(crate) use transport::{
|
||||||
execute_sync_plan as execute_execution_runtime_sync_plan, DirectSyncExecutionRuntime,
|
execute_sync_plan as execute_execution_runtime_sync_plan, DirectSyncExecutionRuntime,
|
||||||
|
|||||||
@@ -221,7 +221,7 @@ async fn execute_sync(
|
|||||||
let plan = parse_request_json::<ExecutionPlan>(request).await?;
|
let plan = parse_request_json::<ExecutionPlan>(request).await?;
|
||||||
let result = state
|
let result = state
|
||||||
.execution_runtime
|
.execution_runtime
|
||||||
.execute_sync(plan)
|
.execute_sync(&plan)
|
||||||
.await
|
.await
|
||||||
.map_err(|err| ExecutionRuntimeAppError(ExecutionRuntimeServerError::Transport(err)))?;
|
.map_err(|err| ExecutionRuntimeAppError(ExecutionRuntimeServerError::Transport(err)))?;
|
||||||
Ok(maybe_hold_axum_response_permit(
|
Ok(maybe_hold_axum_response_permit(
|
||||||
@@ -238,7 +238,7 @@ async fn execute_stream(
|
|||||||
let plan = parse_request_json::<ExecutionPlan>(request).await?;
|
let plan = parse_request_json::<ExecutionPlan>(request).await?;
|
||||||
let execution = state
|
let execution = state
|
||||||
.execution_runtime
|
.execution_runtime
|
||||||
.execute_stream(plan)
|
.execute_stream(&plan)
|
||||||
.await
|
.await
|
||||||
.map_err(|err| ExecutionRuntimeAppError(ExecutionRuntimeServerError::Transport(err)))?;
|
.map_err(|err| ExecutionRuntimeAppError(ExecutionRuntimeServerError::Transport(err)))?;
|
||||||
|
|
||||||
|
|||||||
@@ -74,6 +74,7 @@ use crate::orchestration::{
|
|||||||
};
|
};
|
||||||
use crate::request_candidate_runtime::{
|
use crate::request_candidate_runtime::{
|
||||||
ensure_execution_request_candidate_slot, record_local_request_candidate_status,
|
ensure_execution_request_candidate_slot, record_local_request_candidate_status,
|
||||||
|
record_local_request_candidate_status_snapshot, snapshot_local_request_candidate_status,
|
||||||
};
|
};
|
||||||
use crate::usage::submit_stream_report;
|
use crate::usage::submit_stream_report;
|
||||||
use crate::usage::{GatewayStreamReportRequest, GatewaySyncReportRequest};
|
use crate::usage::{GatewayStreamReportRequest, GatewaySyncReportRequest};
|
||||||
@@ -89,7 +90,30 @@ fn record_sync_terminal_usage(
|
|||||||
let payload_seed = build_sync_terminal_usage_payload_seed(payload);
|
let payload_seed = build_sync_terminal_usage_payload_seed(payload);
|
||||||
state
|
state
|
||||||
.usage_runtime
|
.usage_runtime
|
||||||
.record_sync_terminal(state.data.as_ref(), &context_seed, &payload_seed);
|
.record_sync_terminal(state.data.as_ref(), context_seed, payload_seed);
|
||||||
|
}
|
||||||
|
|
||||||
|
fn build_stream_sync_payload(
|
||||||
|
trace_id: &str,
|
||||||
|
report_kind: String,
|
||||||
|
report_context: Option<Value>,
|
||||||
|
status_code: u16,
|
||||||
|
headers: BTreeMap<String, String>,
|
||||||
|
body_json: Option<Value>,
|
||||||
|
body_base64: Option<String>,
|
||||||
|
telemetry: Option<ExecutionTelemetry>,
|
||||||
|
) -> GatewaySyncReportRequest {
|
||||||
|
GatewaySyncReportRequest {
|
||||||
|
trace_id: trace_id.to_string(),
|
||||||
|
report_kind,
|
||||||
|
report_context,
|
||||||
|
status_code,
|
||||||
|
headers,
|
||||||
|
body_json,
|
||||||
|
client_body_json: None,
|
||||||
|
body_base64,
|
||||||
|
telemetry,
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn record_stream_terminal_usage(
|
fn record_stream_terminal_usage(
|
||||||
@@ -103,12 +127,61 @@ fn record_stream_terminal_usage(
|
|||||||
let payload_seed = build_stream_terminal_usage_payload_seed(payload);
|
let payload_seed = build_stream_terminal_usage_payload_seed(payload);
|
||||||
state.usage_runtime.record_stream_terminal(
|
state.usage_runtime.record_stream_terminal(
|
||||||
state.data.as_ref(),
|
state.data.as_ref(),
|
||||||
&context_seed,
|
context_seed,
|
||||||
&payload_seed,
|
payload_seed,
|
||||||
cancelled,
|
cancelled,
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn build_stream_body_capture(
|
||||||
|
body: &[u8],
|
||||||
|
truncated: bool,
|
||||||
|
) -> (Option<String>, Option<UsageBodyCaptureState>) {
|
||||||
|
let body_base64 =
|
||||||
|
(!body.is_empty()).then(|| base64::engine::general_purpose::STANDARD.encode(body));
|
||||||
|
let body_state = Some(if truncated {
|
||||||
|
UsageBodyCaptureState::Truncated
|
||||||
|
} else if body.is_empty() {
|
||||||
|
UsageBodyCaptureState::None
|
||||||
|
} else {
|
||||||
|
UsageBodyCaptureState::Inline
|
||||||
|
});
|
||||||
|
(body_base64, body_state)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[allow(clippy::too_many_arguments)] // stream report payload assembly mirrors runtime state
|
||||||
|
fn build_stream_usage_payload(
|
||||||
|
trace_id: String,
|
||||||
|
report_kind: String,
|
||||||
|
report_context: Option<Value>,
|
||||||
|
status_code: u16,
|
||||||
|
headers: BTreeMap<String, String>,
|
||||||
|
provider_body: &[u8],
|
||||||
|
provider_body_truncated: bool,
|
||||||
|
client_body: &[u8],
|
||||||
|
client_body_truncated: bool,
|
||||||
|
terminal_summary: Option<ExecutionStreamTerminalSummary>,
|
||||||
|
telemetry: Option<ExecutionTelemetry>,
|
||||||
|
) -> GatewayStreamReportRequest {
|
||||||
|
let (provider_body_base64, provider_body_state) =
|
||||||
|
build_stream_body_capture(provider_body, provider_body_truncated);
|
||||||
|
let (client_body_base64, client_body_state) =
|
||||||
|
build_stream_body_capture(client_body, client_body_truncated);
|
||||||
|
GatewayStreamReportRequest {
|
||||||
|
trace_id,
|
||||||
|
report_kind,
|
||||||
|
report_context,
|
||||||
|
status_code,
|
||||||
|
headers,
|
||||||
|
provider_body_base64,
|
||||||
|
provider_body_state,
|
||||||
|
client_body_base64,
|
||||||
|
client_body_state,
|
||||||
|
terminal_summary,
|
||||||
|
telemetry,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
fn append_stream_capture_bytes(
|
fn append_stream_capture_bytes(
|
||||||
buffer: &mut Vec<u8>,
|
buffer: &mut Vec<u8>,
|
||||||
chunk: &[u8],
|
chunk: &[u8],
|
||||||
@@ -138,9 +211,7 @@ async fn execute_in_process_stream(
|
|||||||
return Ok(execution);
|
return Ok(execution);
|
||||||
}
|
}
|
||||||
|
|
||||||
DirectSyncExecutionRuntime::new()
|
DirectSyncExecutionRuntime::new().execute_stream(plan).await
|
||||||
.execute_stream(plan.clone())
|
|
||||||
.await
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[allow(clippy::too_many_arguments)] // internal function, grouping would add unnecessary indirection
|
#[allow(clippy::too_many_arguments)] // internal function, grouping would add unnecessary indirection
|
||||||
@@ -155,19 +226,18 @@ pub(crate) async fn execute_execution_runtime_stream(
|
|||||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||||
ensure_execution_request_candidate_slot(state, &mut plan, &mut report_context).await;
|
ensure_execution_request_candidate_slot(state, &mut plan, &mut report_context).await;
|
||||||
let lifecycle_seed = build_lifecycle_usage_seed(&plan, report_context.as_ref());
|
let lifecycle_seed = build_lifecycle_usage_seed(&plan, report_context.as_ref());
|
||||||
|
let request_candidate_status_snapshot =
|
||||||
|
snapshot_local_request_candidate_status(&plan, report_context.as_ref());
|
||||||
state
|
state
|
||||||
.usage_runtime
|
.usage_runtime
|
||||||
.record_pending(state.data.as_ref(), &lifecycle_seed);
|
.record_pending(state.data.as_ref(), lifecycle_seed.clone());
|
||||||
let candidate_started_unix_secs = current_request_candidate_unix_ms();
|
let candidate_started_unix_secs = current_request_candidate_unix_ms();
|
||||||
{
|
if let Some(snapshot) = request_candidate_status_snapshot.clone() {
|
||||||
let state_bg = state.clone();
|
let state_bg = state.clone();
|
||||||
let plan_bg = plan.clone();
|
|
||||||
let report_context_bg = report_context.clone();
|
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
record_local_request_candidate_status(
|
record_local_request_candidate_status_snapshot(
|
||||||
&state_bg,
|
&state_bg,
|
||||||
&plan_bg,
|
&snapshot,
|
||||||
report_context_bg.as_ref(),
|
|
||||||
SchedulerRequestCandidateStatusUpdate {
|
SchedulerRequestCandidateStatusUpdate {
|
||||||
status: RequestCandidateStatus::Pending,
|
status: RequestCandidateStatus::Pending,
|
||||||
status_code: None,
|
status_code: None,
|
||||||
@@ -412,8 +482,9 @@ fn response_headers_indicate_sse(headers: &BTreeMap<String, String>) -> bool {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn encode_terminal_sse_error_event(failure: &StreamFailureReport) -> Result<Bytes, std::io::Error> {
|
fn encode_terminal_sse_error_event(failure: &StreamFailureReport) -> Result<Bytes, std::io::Error> {
|
||||||
let payload =
|
let payload = failure
|
||||||
serde_json::to_string(&failure.body_json).map_err(|err| IoError::other(err.to_string()))?;
|
.to_json_string()
|
||||||
|
.map_err(|err| IoError::other(err.to_string()))?;
|
||||||
let mut event = String::from("event: aether.error\n");
|
let mut event = String::from("event: aether.error\n");
|
||||||
for line in payload.lines() {
|
for line in payload.lines() {
|
||||||
event.push_str("data: ");
|
event.push_str("data: ");
|
||||||
@@ -527,6 +598,8 @@ async fn execute_stream_from_frame_stream(
|
|||||||
let provider_name = plan.provider_name.as_deref().unwrap_or("-");
|
let provider_name = plan.provider_name.as_deref().unwrap_or("-");
|
||||||
let model_name = plan.model_name.as_deref().unwrap_or("-");
|
let model_name = plan.model_name.as_deref().unwrap_or("-");
|
||||||
let lifecycle_seed = build_lifecycle_usage_seed(&plan, report_context.as_ref());
|
let lifecycle_seed = build_lifecycle_usage_seed(&plan, report_context.as_ref());
|
||||||
|
let request_candidate_status_snapshot =
|
||||||
|
snapshot_local_request_candidate_status(&plan, report_context.as_ref());
|
||||||
let candidate_index = parse_request_candidate_report_context(report_context.as_ref())
|
let candidate_index = parse_request_candidate_report_context(report_context.as_ref())
|
||||||
.and_then(|context| context.candidate_index)
|
.and_then(|context| context.candidate_index)
|
||||||
.map(|value| value.to_string())
|
.map(|value| value.to_string())
|
||||||
@@ -756,27 +829,26 @@ async fn execute_stream_from_frame_stream(
|
|||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
|
|
||||||
let usage_report_kind = stream_error_finalize_kind
|
let payload = build_stream_sync_payload(
|
||||||
.clone()
|
trace_id,
|
||||||
.or_else(|| report_kind.clone())
|
stream_error_finalize_kind
|
||||||
.unwrap_or_default();
|
.as_deref()
|
||||||
let usage_payload = GatewaySyncReportRequest {
|
.or(report_kind.as_deref())
|
||||||
trace_id: trace_id.to_string(),
|
.unwrap_or_default()
|
||||||
report_kind: usage_report_kind,
|
.to_string(),
|
||||||
report_context: report_context.clone(),
|
report_context,
|
||||||
status_code,
|
status_code,
|
||||||
headers: headers.clone(),
|
headers,
|
||||||
body_json: body_json.clone(),
|
body_json,
|
||||||
client_body_json: None,
|
body_base64,
|
||||||
body_base64: body_base64.clone(),
|
None,
|
||||||
telemetry: None,
|
);
|
||||||
};
|
record_sync_terminal_usage(state, &plan, payload.report_context.as_ref(), &payload);
|
||||||
record_sync_terminal_usage(state, &plan, report_context.as_ref(), &usage_payload);
|
|
||||||
let terminal_unix_secs = current_request_candidate_unix_ms();
|
let terminal_unix_secs = current_request_candidate_unix_ms();
|
||||||
record_local_request_candidate_status(
|
record_local_request_candidate_status(
|
||||||
state,
|
state,
|
||||||
&plan,
|
&plan,
|
||||||
report_context.as_ref(),
|
payload.report_context.as_ref(),
|
||||||
SchedulerRequestCandidateStatusUpdate {
|
SchedulerRequestCandidateStatusUpdate {
|
||||||
status: RequestCandidateStatus::Failed,
|
status: RequestCandidateStatus::Failed,
|
||||||
status_code: Some(status_code),
|
status_code: Some(status_code),
|
||||||
@@ -790,18 +862,7 @@ async fn execute_stream_from_frame_stream(
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
if let Some(report_kind) = stream_error_finalize_kind {
|
if stream_error_finalize_kind.is_some() {
|
||||||
let payload = GatewaySyncReportRequest {
|
|
||||||
trace_id: trace_id.to_string(),
|
|
||||||
report_kind,
|
|
||||||
report_context,
|
|
||||||
status_code,
|
|
||||||
headers: headers.clone(),
|
|
||||||
body_json,
|
|
||||||
client_body_json: None,
|
|
||||||
body_base64,
|
|
||||||
telemetry: None,
|
|
||||||
};
|
|
||||||
let response =
|
let response =
|
||||||
submit_local_core_error_or_sync_finalize(state, trace_id, decision, payload)
|
submit_local_core_error_or_sync_finalize(state, trace_id, decision, payload)
|
||||||
.await?;
|
.await?;
|
||||||
@@ -817,7 +878,7 @@ async fn execute_stream_from_frame_stream(
|
|||||||
decision,
|
decision,
|
||||||
plan_kind,
|
plan_kind,
|
||||||
status_code,
|
status_code,
|
||||||
headers,
|
payload.headers,
|
||||||
error_body,
|
error_body,
|
||||||
)?,
|
)?,
|
||||||
Some(request_id),
|
Some(request_id),
|
||||||
@@ -891,12 +952,12 @@ async fn execute_stream_from_frame_stream(
|
|||||||
trace_id,
|
trace_id,
|
||||||
decision,
|
decision,
|
||||||
&plan,
|
&plan,
|
||||||
report_context.clone(),
|
report_context,
|
||||||
request_id,
|
request_id,
|
||||||
candidate_id,
|
candidate_id,
|
||||||
report_kind,
|
report_kind,
|
||||||
&headers,
|
headers,
|
||||||
prefetched_telemetry.clone(),
|
prefetched_telemetry,
|
||||||
&provider_prefetched_body,
|
&provider_prefetched_body,
|
||||||
failure,
|
failure,
|
||||||
)
|
)
|
||||||
@@ -924,12 +985,12 @@ async fn execute_stream_from_frame_stream(
|
|||||||
trace_id,
|
trace_id,
|
||||||
decision,
|
decision,
|
||||||
&plan,
|
&plan,
|
||||||
report_context.clone(),
|
report_context,
|
||||||
request_id,
|
request_id,
|
||||||
candidate_id,
|
candidate_id,
|
||||||
report_kind,
|
report_kind,
|
||||||
&headers,
|
headers,
|
||||||
prefetched_telemetry.clone(),
|
prefetched_telemetry,
|
||||||
&prefetched_body,
|
&prefetched_body,
|
||||||
failure,
|
failure,
|
||||||
)
|
)
|
||||||
@@ -964,21 +1025,20 @@ async fn execute_stream_from_frame_stream(
|
|||||||
provider_prefetched_body_bytes = provider_prefetched_body.len(),
|
provider_prefetched_body_bytes = provider_prefetched_body.len(),
|
||||||
"gateway detected embedded error while prefetching execution runtime stream"
|
"gateway detected embedded error while prefetching execution runtime stream"
|
||||||
);
|
);
|
||||||
let payload = GatewaySyncReportRequest {
|
let payload = build_stream_sync_payload(
|
||||||
trace_id: trace_id.to_string(),
|
trace_id,
|
||||||
report_kind: report_kind.clone(),
|
report_kind.clone(),
|
||||||
report_context: report_context.clone(),
|
report_context,
|
||||||
status_code,
|
status_code,
|
||||||
headers: headers.clone(),
|
headers,
|
||||||
body_json: Some(body_json),
|
Some(body_json),
|
||||||
client_body_json: None,
|
None,
|
||||||
body_base64: None,
|
prefetched_telemetry,
|
||||||
telemetry: prefetched_telemetry.clone(),
|
);
|
||||||
};
|
|
||||||
record_sync_terminal_usage(
|
record_sync_terminal_usage(
|
||||||
state,
|
state,
|
||||||
&plan,
|
&plan,
|
||||||
report_context.as_ref(),
|
payload.report_context.as_ref(),
|
||||||
&payload,
|
&payload,
|
||||||
);
|
);
|
||||||
let response = submit_local_core_error_or_sync_finalize(
|
let response = submit_local_core_error_or_sync_finalize(
|
||||||
@@ -1013,12 +1073,12 @@ async fn execute_stream_from_frame_stream(
|
|||||||
trace_id,
|
trace_id,
|
||||||
decision,
|
decision,
|
||||||
&plan,
|
&plan,
|
||||||
report_context.clone(),
|
report_context,
|
||||||
request_id,
|
request_id,
|
||||||
candidate_id,
|
candidate_id,
|
||||||
report_kind,
|
report_kind,
|
||||||
&headers,
|
headers,
|
||||||
prefetched_telemetry.clone(),
|
prefetched_telemetry,
|
||||||
&provider_prefetched_body,
|
&provider_prefetched_body,
|
||||||
failure,
|
failure,
|
||||||
)
|
)
|
||||||
@@ -1044,12 +1104,12 @@ async fn execute_stream_from_frame_stream(
|
|||||||
trace_id,
|
trace_id,
|
||||||
decision,
|
decision,
|
||||||
&plan,
|
&plan,
|
||||||
report_context.clone(),
|
report_context,
|
||||||
request_id,
|
request_id,
|
||||||
candidate_id,
|
candidate_id,
|
||||||
report_kind,
|
report_kind,
|
||||||
&headers,
|
headers,
|
||||||
prefetched_telemetry.clone(),
|
prefetched_telemetry,
|
||||||
&provider_prefetched_body,
|
&provider_prefetched_body,
|
||||||
failure,
|
failure,
|
||||||
)
|
)
|
||||||
@@ -1093,12 +1153,12 @@ async fn execute_stream_from_frame_stream(
|
|||||||
trace_id,
|
trace_id,
|
||||||
decision,
|
decision,
|
||||||
&plan,
|
&plan,
|
||||||
report_context.clone(),
|
report_context,
|
||||||
request_id,
|
request_id,
|
||||||
candidate_id,
|
candidate_id,
|
||||||
report_kind,
|
report_kind,
|
||||||
&headers,
|
headers,
|
||||||
prefetched_telemetry.clone(),
|
prefetched_telemetry,
|
||||||
&provider_prefetched_body,
|
&provider_prefetched_body,
|
||||||
build_stream_failure_from_execution_error(&error),
|
build_stream_failure_from_execution_error(&error),
|
||||||
)
|
)
|
||||||
@@ -1108,6 +1168,8 @@ async fn execute_stream_from_frame_stream(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
drop(private_stream_normalizer);
|
||||||
|
drop(local_stream_rewriter);
|
||||||
|
|
||||||
state.usage_runtime.record_stream_started(
|
state.usage_runtime.record_stream_started(
|
||||||
state.data.as_ref(),
|
state.data.as_ref(),
|
||||||
@@ -1115,18 +1177,15 @@ async fn execute_stream_from_frame_stream(
|
|||||||
status_code,
|
status_code,
|
||||||
prefetched_telemetry.as_ref(),
|
prefetched_telemetry.as_ref(),
|
||||||
);
|
);
|
||||||
{
|
if let Some(snapshot) = request_candidate_status_snapshot {
|
||||||
let state_bg = state.clone();
|
let state_bg = state.clone();
|
||||||
let plan_bg = plan.clone();
|
|
||||||
let report_context_bg = report_context.clone();
|
|
||||||
let latency_ms = prefetched_telemetry
|
let latency_ms = prefetched_telemetry
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.and_then(|telemetry| telemetry.elapsed_ms);
|
.and_then(|telemetry| telemetry.elapsed_ms);
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
record_local_request_candidate_status(
|
record_local_request_candidate_status_snapshot(
|
||||||
&state_bg,
|
&state_bg,
|
||||||
&plan_bg,
|
&snapshot,
|
||||||
report_context_bg.as_ref(),
|
|
||||||
SchedulerRequestCandidateStatusUpdate {
|
SchedulerRequestCandidateStatusUpdate {
|
||||||
status: RequestCandidateStatus::Streaming,
|
status: RequestCandidateStatus::Streaming,
|
||||||
status_code: Some(status_code),
|
status_code: Some(status_code),
|
||||||
@@ -1141,24 +1200,27 @@ async fn execute_stream_from_frame_stream(
|
|||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
|
let request_id = request_id.to_string();
|
||||||
|
let candidate_id = candidate_id.map(ToOwned::to_owned);
|
||||||
let (tx, mut rx) = mpsc::channel::<Result<Bytes, IoError>>(16);
|
let (tx, mut rx) = mpsc::channel::<Result<Bytes, IoError>>(16);
|
||||||
let state_for_report = state.clone();
|
let state_for_report = state.clone();
|
||||||
let plan_for_report = plan.clone();
|
let plan_for_report = plan;
|
||||||
let trace_id_owned = trace_id.to_string();
|
let trace_id_owned = trace_id.to_string();
|
||||||
let headers_for_report = headers.clone();
|
let headers_for_report = headers.clone();
|
||||||
let report_kind_owned = report_kind.clone();
|
let report_kind_owned = report_kind;
|
||||||
let report_context_owned = report_context.clone();
|
let report_context_owned = report_context;
|
||||||
let lifecycle_seed_for_report = lifecycle_seed.clone();
|
let normalized_stream_report_context_owned = normalized_stream_report_context;
|
||||||
|
let lifecycle_seed_for_report = lifecycle_seed;
|
||||||
let provider_prefetched_body_for_report = provider_prefetched_body;
|
let provider_prefetched_body_for_report = provider_prefetched_body;
|
||||||
let prefetched_body_for_report = prefetched_body;
|
let prefetched_body_for_report = prefetched_body;
|
||||||
let prefetched_chunks_for_body = prefetched_chunks;
|
let prefetched_chunks_for_body = prefetched_chunks;
|
||||||
let initial_telemetry = prefetched_telemetry.clone();
|
let initial_telemetry = prefetched_telemetry;
|
||||||
let initial_reached_eof = reached_eof;
|
let initial_reached_eof = reached_eof;
|
||||||
let direct_stream_finalize_kind_owned = direct_stream_finalize_kind.clone();
|
let direct_stream_finalize_kind_owned = direct_stream_finalize_kind;
|
||||||
let candidate_started_unix_secs_for_report = candidate_started_unix_secs;
|
let candidate_started_unix_secs_for_report = candidate_started_unix_secs;
|
||||||
let request_id_for_report = request_id.to_string();
|
let request_id_for_report = request_id.clone();
|
||||||
let request_id_for_report_log = short_request_id(request_id);
|
let request_id_for_report_log = short_request_id(&request_id);
|
||||||
let candidate_id_for_report = candidate_id.map(ToOwned::to_owned);
|
let candidate_id_for_report = candidate_id.clone();
|
||||||
let emit_passthrough_sse_terminal_error =
|
let emit_passthrough_sse_terminal_error =
|
||||||
skip_direct_finalize_prefetch && response_headers_indicate_sse(&headers);
|
skip_direct_finalize_prefetch && response_headers_indicate_sse(&headers);
|
||||||
let body_capture_policy = match UsageRuntimeAccess::body_capture_policy(state.data.as_ref())
|
let body_capture_policy = match UsageRuntimeAccess::body_capture_policy(state.data.as_ref())
|
||||||
@@ -1194,6 +1256,10 @@ async fn execute_stream_from_frame_stream(
|
|||||||
let mut buffered_body = Vec::new();
|
let mut buffered_body = Vec::new();
|
||||||
let mut provider_body_truncated = false;
|
let mut provider_body_truncated = false;
|
||||||
let mut client_body_truncated = false;
|
let mut client_body_truncated = false;
|
||||||
|
let mut private_stream_normalizer =
|
||||||
|
maybe_build_provider_private_stream_normalizer(report_context_owned.as_ref());
|
||||||
|
let mut local_stream_rewriter =
|
||||||
|
maybe_build_stream_response_rewriter(normalized_stream_report_context_owned.as_ref());
|
||||||
append_stream_capture_bytes(
|
append_stream_capture_bytes(
|
||||||
&mut provider_buffered_body,
|
&mut provider_buffered_body,
|
||||||
&provider_prefetched_body_for_report,
|
&provider_prefetched_body_for_report,
|
||||||
@@ -1206,13 +1272,68 @@ async fn execute_stream_from_frame_stream(
|
|||||||
max_stream_body_buffer_bytes,
|
max_stream_body_buffer_bytes,
|
||||||
&mut client_body_truncated,
|
&mut client_body_truncated,
|
||||||
);
|
);
|
||||||
let mut telemetry: Option<ExecutionTelemetry> = initial_telemetry.clone();
|
let mut usage_stream_telemetry: Option<ExecutionTelemetry> = initial_telemetry.clone();
|
||||||
let mut usage_stream_telemetry: Option<ExecutionTelemetry> = initial_telemetry;
|
let mut telemetry: Option<ExecutionTelemetry> = initial_telemetry;
|
||||||
let reached_eof = initial_reached_eof;
|
let reached_eof = initial_reached_eof;
|
||||||
let mut downstream_dropped = false;
|
let mut downstream_dropped = false;
|
||||||
let mut terminal_failure: Option<StreamFailureReport> = None;
|
let mut terminal_failure: Option<StreamFailureReport> = None;
|
||||||
|
if !provider_prefetched_body_for_report.is_empty() {
|
||||||
|
let normalized_prefetched_chunk = if let Some(normalizer) =
|
||||||
|
private_stream_normalizer.as_mut()
|
||||||
|
{
|
||||||
|
match normalizer.push_chunk(&provider_prefetched_body_for_report) {
|
||||||
|
Ok(normalized_chunk) => Some(normalized_chunk),
|
||||||
|
Err(err) => {
|
||||||
|
warn!(
|
||||||
|
event_name = "stream_execution_prefetch_normalize_restore_failed",
|
||||||
|
log_type = "ops",
|
||||||
|
trace_id = %trace_id_owned,
|
||||||
|
request_id = %request_id_for_report_log,
|
||||||
|
candidate_id = ?candidate_id_for_report.as_deref(),
|
||||||
|
error = ?err,
|
||||||
|
"gateway failed to restore private stream normalization state after prefetch"
|
||||||
|
);
|
||||||
|
terminal_failure = Some(build_stream_failure_report(
|
||||||
|
"execution_runtime_stream_rewrite_error",
|
||||||
|
format!(
|
||||||
|
"failed to restore private stream normalization state after prefetch: {err:?}"
|
||||||
|
),
|
||||||
|
502,
|
||||||
|
));
|
||||||
|
None
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
None
|
||||||
|
};
|
||||||
|
let replay_chunk = normalized_prefetched_chunk
|
||||||
|
.as_deref()
|
||||||
|
.unwrap_or(provider_prefetched_body_for_report.as_slice());
|
||||||
|
if terminal_failure.is_none() {
|
||||||
|
if let Some(rewriter) = local_stream_rewriter.as_mut() {
|
||||||
|
if let Err(err) = rewriter.push_chunk(replay_chunk) {
|
||||||
|
warn!(
|
||||||
|
event_name = "stream_execution_prefetch_rewrite_restore_failed",
|
||||||
|
log_type = "ops",
|
||||||
|
trace_id = %trace_id_owned,
|
||||||
|
request_id = %request_id_for_report_log,
|
||||||
|
candidate_id = ?candidate_id_for_report.as_deref(),
|
||||||
|
error = ?err,
|
||||||
|
"gateway failed to restore local stream rewrite state after prefetch"
|
||||||
|
);
|
||||||
|
terminal_failure = Some(build_stream_failure_report(
|
||||||
|
"execution_runtime_stream_rewrite_error",
|
||||||
|
format!(
|
||||||
|
"failed to restore local stream rewrite state after prefetch: {err:?}"
|
||||||
|
),
|
||||||
|
502,
|
||||||
|
));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
if !reached_eof {
|
if terminal_failure.is_none() && !reached_eof {
|
||||||
loop {
|
loop {
|
||||||
let next_frame = match next_stream_frame(&mut buffered_frames, &mut lines).await {
|
let next_frame = match next_stream_frame(&mut buffered_frames, &mut lines).await {
|
||||||
Ok(frame) => frame,
|
Ok(frame) => frame,
|
||||||
@@ -1355,7 +1476,6 @@ async fn execute_stream_from_frame_stream(
|
|||||||
usage_stream_telemetry.as_ref(),
|
usage_stream_telemetry.as_ref(),
|
||||||
&frame_telemetry,
|
&frame_telemetry,
|
||||||
);
|
);
|
||||||
telemetry = Some(frame_telemetry.clone());
|
|
||||||
if should_refresh_stream_usage {
|
if should_refresh_stream_usage {
|
||||||
state_for_report.usage_runtime.record_stream_started(
|
state_for_report.usage_runtime.record_stream_started(
|
||||||
state_for_report.data.as_ref(),
|
state_for_report.data.as_ref(),
|
||||||
@@ -1363,8 +1483,9 @@ async fn execute_stream_from_frame_stream(
|
|||||||
status_code,
|
status_code,
|
||||||
Some(&frame_telemetry),
|
Some(&frame_telemetry),
|
||||||
);
|
);
|
||||||
usage_stream_telemetry = Some(frame_telemetry);
|
usage_stream_telemetry = Some(frame_telemetry.clone());
|
||||||
}
|
}
|
||||||
|
telemetry = Some(frame_telemetry);
|
||||||
}
|
}
|
||||||
StreamFramePayload::Eof { summary } => {
|
StreamFramePayload::Eof { summary } => {
|
||||||
stream_terminal_summary = summary;
|
stream_terminal_summary = summary;
|
||||||
@@ -1564,50 +1685,39 @@ async fn execute_stream_from_frame_stream(
|
|||||||
trace_id = %trace_id_owned,
|
trace_id = %trace_id_owned,
|
||||||
"gateway skipped stream report because downstream disconnected before completion"
|
"gateway skipped stream report because downstream disconnected before completion"
|
||||||
);
|
);
|
||||||
|
let usage_payload = build_stream_usage_payload(
|
||||||
|
trace_id_owned,
|
||||||
|
report_kind_owned.unwrap_or_default(),
|
||||||
|
report_context_owned,
|
||||||
|
499,
|
||||||
|
headers_for_report,
|
||||||
|
&provider_buffered_body,
|
||||||
|
provider_body_truncated,
|
||||||
|
&buffered_body,
|
||||||
|
client_body_truncated,
|
||||||
|
stream_terminal_summary,
|
||||||
|
telemetry,
|
||||||
|
);
|
||||||
record_stream_terminal_usage(
|
record_stream_terminal_usage(
|
||||||
&state_for_report,
|
&state_for_report,
|
||||||
&plan_for_report,
|
&plan_for_report,
|
||||||
report_context_owned.as_ref(),
|
usage_payload.report_context.as_ref(),
|
||||||
&GatewayStreamReportRequest {
|
&usage_payload,
|
||||||
trace_id: trace_id_owned.clone(),
|
|
||||||
report_kind: report_kind_owned.clone().unwrap_or_default(),
|
|
||||||
report_context: report_context_owned.clone(),
|
|
||||||
status_code: 499,
|
|
||||||
headers: headers_for_report.clone(),
|
|
||||||
provider_body_base64: (!provider_buffered_body.is_empty()).then(|| {
|
|
||||||
base64::engine::general_purpose::STANDARD.encode(&provider_buffered_body)
|
|
||||||
}),
|
|
||||||
provider_body_state: Some(if provider_body_truncated {
|
|
||||||
UsageBodyCaptureState::Truncated
|
|
||||||
} else if provider_buffered_body.is_empty() {
|
|
||||||
UsageBodyCaptureState::None
|
|
||||||
} else {
|
|
||||||
UsageBodyCaptureState::Inline
|
|
||||||
}),
|
|
||||||
client_body_base64: (!buffered_body.is_empty())
|
|
||||||
.then(|| base64::engine::general_purpose::STANDARD.encode(&buffered_body)),
|
|
||||||
client_body_state: Some(if client_body_truncated {
|
|
||||||
UsageBodyCaptureState::Truncated
|
|
||||||
} else if buffered_body.is_empty() {
|
|
||||||
UsageBodyCaptureState::None
|
|
||||||
} else {
|
|
||||||
UsageBodyCaptureState::Inline
|
|
||||||
}),
|
|
||||||
terminal_summary: stream_terminal_summary.clone(),
|
|
||||||
telemetry: telemetry.clone(),
|
|
||||||
},
|
|
||||||
true,
|
true,
|
||||||
);
|
);
|
||||||
record_local_request_candidate_status(
|
record_local_request_candidate_status(
|
||||||
&state_for_report,
|
&state_for_report,
|
||||||
&plan_for_report,
|
&plan_for_report,
|
||||||
report_context_owned.as_ref(),
|
usage_payload.report_context.as_ref(),
|
||||||
SchedulerRequestCandidateStatusUpdate {
|
SchedulerRequestCandidateStatusUpdate {
|
||||||
status: RequestCandidateStatus::Cancelled,
|
status: RequestCandidateStatus::Cancelled,
|
||||||
status_code: Some(499),
|
status_code: Some(499),
|
||||||
error_type: Some("downstream_disconnect".to_string()),
|
error_type: Some("downstream_disconnect".to_string()),
|
||||||
error_message: Some("client disconnected before stream completion".to_string()),
|
error_message: Some("client disconnected before stream completion".to_string()),
|
||||||
latency_ms: telemetry.as_ref().and_then(|value| value.elapsed_ms),
|
latency_ms: usage_payload
|
||||||
|
.telemetry
|
||||||
|
.as_ref()
|
||||||
|
.and_then(|value| value.elapsed_ms),
|
||||||
started_at_unix_ms: Some(candidate_started_unix_secs_for_report),
|
started_at_unix_ms: Some(candidate_started_unix_secs_for_report),
|
||||||
finished_at_unix_ms: Some(current_request_candidate_unix_ms()),
|
finished_at_unix_ms: Some(current_request_candidate_unix_ms()),
|
||||||
},
|
},
|
||||||
@@ -1622,9 +1732,9 @@ async fn execute_stream_from_frame_stream(
|
|||||||
&trace_id_owned,
|
&trace_id_owned,
|
||||||
&plan_for_report,
|
&plan_for_report,
|
||||||
direct_stream_finalize_kind_owned.as_deref(),
|
direct_stream_finalize_kind_owned.as_deref(),
|
||||||
report_context_owned.as_ref(),
|
report_context_owned,
|
||||||
&headers_for_report,
|
headers_for_report,
|
||||||
telemetry.clone(),
|
telemetry,
|
||||||
&provider_buffered_body,
|
&provider_buffered_body,
|
||||||
candidate_started_unix_secs_for_report,
|
candidate_started_unix_secs_for_report,
|
||||||
failure,
|
failure,
|
||||||
@@ -1633,38 +1743,25 @@ async fn execute_stream_from_frame_stream(
|
|||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
let usage_payload = GatewayStreamReportRequest {
|
let should_submit_report = report_kind_owned.is_some();
|
||||||
trace_id: trace_id_owned.clone(),
|
let usage_payload = build_stream_usage_payload(
|
||||||
report_kind: report_kind_owned.clone().unwrap_or_default(),
|
trace_id_owned.clone(),
|
||||||
report_context: report_context_owned.clone(),
|
report_kind_owned.unwrap_or_default(),
|
||||||
|
report_context_owned,
|
||||||
status_code,
|
status_code,
|
||||||
headers: headers_for_report.clone(),
|
headers_for_report,
|
||||||
provider_body_base64: (!provider_buffered_body.is_empty())
|
&provider_buffered_body,
|
||||||
.then(|| base64::engine::general_purpose::STANDARD.encode(&provider_buffered_body)),
|
provider_body_truncated,
|
||||||
provider_body_state: Some(if provider_body_truncated {
|
&buffered_body,
|
||||||
UsageBodyCaptureState::Truncated
|
client_body_truncated,
|
||||||
} else if provider_buffered_body.is_empty() {
|
stream_terminal_summary,
|
||||||
UsageBodyCaptureState::None
|
telemetry,
|
||||||
} else {
|
);
|
||||||
UsageBodyCaptureState::Inline
|
|
||||||
}),
|
|
||||||
client_body_base64: (!buffered_body.is_empty())
|
|
||||||
.then(|| base64::engine::general_purpose::STANDARD.encode(&buffered_body)),
|
|
||||||
client_body_state: Some(if client_body_truncated {
|
|
||||||
UsageBodyCaptureState::Truncated
|
|
||||||
} else if buffered_body.is_empty() {
|
|
||||||
UsageBodyCaptureState::None
|
|
||||||
} else {
|
|
||||||
UsageBodyCaptureState::Inline
|
|
||||||
}),
|
|
||||||
terminal_summary: stream_terminal_summary,
|
|
||||||
telemetry: telemetry.clone(),
|
|
||||||
};
|
|
||||||
apply_local_execution_effect(
|
apply_local_execution_effect(
|
||||||
&state_for_report,
|
&state_for_report,
|
||||||
LocalExecutionEffectContext {
|
LocalExecutionEffectContext {
|
||||||
plan: &plan_for_report,
|
plan: &plan_for_report,
|
||||||
report_context: report_context_owned.as_ref(),
|
report_context: usage_payload.report_context.as_ref(),
|
||||||
},
|
},
|
||||||
LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect),
|
LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect),
|
||||||
)
|
)
|
||||||
@@ -1673,7 +1770,7 @@ async fn execute_stream_from_frame_stream(
|
|||||||
&state_for_report,
|
&state_for_report,
|
||||||
LocalExecutionEffectContext {
|
LocalExecutionEffectContext {
|
||||||
plan: &plan_for_report,
|
plan: &plan_for_report,
|
||||||
report_context: report_context_owned.as_ref(),
|
report_context: usage_payload.report_context.as_ref(),
|
||||||
},
|
},
|
||||||
LocalExecutionEffect::AdaptiveSuccess(LocalAdaptiveSuccessEffect),
|
LocalExecutionEffect::AdaptiveSuccess(LocalAdaptiveSuccessEffect),
|
||||||
)
|
)
|
||||||
@@ -1682,7 +1779,7 @@ async fn execute_stream_from_frame_stream(
|
|||||||
&state_for_report,
|
&state_for_report,
|
||||||
LocalExecutionEffectContext {
|
LocalExecutionEffectContext {
|
||||||
plan: &plan_for_report,
|
plan: &plan_for_report,
|
||||||
report_context: report_context_owned.as_ref(),
|
report_context: usage_payload.report_context.as_ref(),
|
||||||
},
|
},
|
||||||
LocalExecutionEffect::PoolSuccessStream {
|
LocalExecutionEffect::PoolSuccessStream {
|
||||||
payload: &usage_payload,
|
payload: &usage_payload,
|
||||||
@@ -1692,31 +1789,31 @@ async fn execute_stream_from_frame_stream(
|
|||||||
record_stream_terminal_usage(
|
record_stream_terminal_usage(
|
||||||
&state_for_report,
|
&state_for_report,
|
||||||
&plan_for_report,
|
&plan_for_report,
|
||||||
report_context_owned.as_ref(),
|
usage_payload.report_context.as_ref(),
|
||||||
&usage_payload,
|
&usage_payload,
|
||||||
false,
|
false,
|
||||||
);
|
);
|
||||||
record_local_request_candidate_status(
|
record_local_request_candidate_status(
|
||||||
&state_for_report,
|
&state_for_report,
|
||||||
&plan_for_report,
|
&plan_for_report,
|
||||||
report_context_owned.as_ref(),
|
usage_payload.report_context.as_ref(),
|
||||||
SchedulerRequestCandidateStatusUpdate {
|
SchedulerRequestCandidateStatusUpdate {
|
||||||
status: RequestCandidateStatus::Success,
|
status: RequestCandidateStatus::Success,
|
||||||
status_code: Some(status_code),
|
status_code: Some(status_code),
|
||||||
error_type: None,
|
error_type: None,
|
||||||
error_message: None,
|
error_message: None,
|
||||||
latency_ms: telemetry.as_ref().and_then(|value| value.elapsed_ms),
|
latency_ms: usage_payload
|
||||||
|
.telemetry
|
||||||
|
.as_ref()
|
||||||
|
.and_then(|value| value.elapsed_ms),
|
||||||
started_at_unix_ms: Some(candidate_started_unix_secs_for_report),
|
started_at_unix_ms: Some(candidate_started_unix_secs_for_report),
|
||||||
finished_at_unix_ms: Some(current_request_candidate_unix_ms()),
|
finished_at_unix_ms: Some(current_request_candidate_unix_ms()),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
if let Some(report_kind) = report_kind_owned {
|
if should_submit_report {
|
||||||
let mut report = usage_payload;
|
if let Err(err) = submit_stream_report(&state_for_report, usage_payload).await {
|
||||||
report.report_kind = report_kind;
|
|
||||||
if let Err(err) = submit_stream_report(&state_for_report, &trace_id_owned, report).await
|
|
||||||
{
|
|
||||||
warn!(
|
warn!(
|
||||||
event_name = "execution_report_submit_failed",
|
event_name = "execution_report_submit_failed",
|
||||||
log_type = "ops",
|
log_type = "ops",
|
||||||
@@ -1740,12 +1837,10 @@ async fn execute_stream_from_frame_stream(
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
headers.insert(
|
headers.insert(CONTROL_REQUEST_ID_HEADER.to_string(), request_id.clone());
|
||||||
CONTROL_REQUEST_ID_HEADER.to_string(),
|
|
||||||
request_id.to_string(),
|
|
||||||
);
|
|
||||||
|
|
||||||
if let Some(candidate_id) = candidate_id
|
if let Some(candidate_id) = candidate_id
|
||||||
|
.as_deref()
|
||||||
.map(str::trim)
|
.map(str::trim)
|
||||||
.filter(|value| !value.is_empty())
|
.filter(|value| !value.is_empty())
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ use aether_usage_runtime::{
|
|||||||
use axum::body::Body;
|
use axum::body::Body;
|
||||||
use axum::http::Response;
|
use axum::http::Response;
|
||||||
use base64::Engine as _;
|
use base64::Engine as _;
|
||||||
|
use serde::Serialize;
|
||||||
use serde_json::{Map, Value};
|
use serde_json::{Map, Value};
|
||||||
use tracing::warn;
|
use tracing::warn;
|
||||||
|
|
||||||
@@ -32,7 +33,51 @@ pub(super) struct StreamFailureReport {
|
|||||||
pub(super) status_code: u16,
|
pub(super) status_code: u16,
|
||||||
pub(super) error_type: String,
|
pub(super) error_type: String,
|
||||||
pub(super) error_message: String,
|
pub(super) error_message: String,
|
||||||
pub(super) body_json: Value,
|
extra_error_fields: Map<String, Value>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Serialize)]
|
||||||
|
struct StreamFailureBody<'a> {
|
||||||
|
error: StreamFailureBodyFields<'a>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Serialize)]
|
||||||
|
struct StreamFailureBodyFields<'a> {
|
||||||
|
#[serde(rename = "type")]
|
||||||
|
error_type: &'a str,
|
||||||
|
message: &'a str,
|
||||||
|
code: u16,
|
||||||
|
#[serde(flatten)]
|
||||||
|
extra_error_fields: &'a Map<String, Value>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl StreamFailureReport {
|
||||||
|
fn into_body_json(self) -> Value {
|
||||||
|
let Self {
|
||||||
|
status_code,
|
||||||
|
error_type,
|
||||||
|
error_message,
|
||||||
|
mut extra_error_fields,
|
||||||
|
} = self;
|
||||||
|
extra_error_fields.insert("type".to_string(), Value::String(error_type));
|
||||||
|
extra_error_fields.insert("message".to_string(), Value::String(error_message));
|
||||||
|
extra_error_fields.insert("code".to_string(), Value::from(status_code));
|
||||||
|
Value::Object(Map::from_iter([(
|
||||||
|
"error".to_string(),
|
||||||
|
Value::Object(extra_error_fields),
|
||||||
|
)]))
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(super) fn to_json_string(&self) -> serde_json::Result<String> {
|
||||||
|
serde_json::to_string(&StreamFailureBody {
|
||||||
|
error: StreamFailureBodyFields {
|
||||||
|
error_type: self.error_type.as_str(),
|
||||||
|
message: self.error_message.as_str(),
|
||||||
|
code: self.status_code,
|
||||||
|
extra_error_fields: &self.extra_error_fields,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(super) fn build_stream_failure_report(
|
pub(super) fn build_stream_failure_report(
|
||||||
@@ -44,16 +89,9 @@ pub(super) fn build_stream_failure_report(
|
|||||||
let error_message = error_message.into();
|
let error_message = error_message.into();
|
||||||
StreamFailureReport {
|
StreamFailureReport {
|
||||||
status_code,
|
status_code,
|
||||||
body_json: Value::Object(Map::from_iter([(
|
|
||||||
"error".to_string(),
|
|
||||||
Value::Object(Map::from_iter([
|
|
||||||
("type".to_string(), Value::String(error_type.clone())),
|
|
||||||
("message".to_string(), Value::String(error_message.clone())),
|
|
||||||
("code".to_string(), Value::from(status_code)),
|
|
||||||
])),
|
|
||||||
)])),
|
|
||||||
error_type,
|
error_type,
|
||||||
error_message,
|
error_message,
|
||||||
|
extra_error_fields: Map::new(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -65,11 +103,9 @@ pub(super) fn build_stream_failure_from_execution_error(
|
|||||||
.ok()
|
.ok()
|
||||||
.and_then(|value| value.as_str().map(ToOwned::to_owned))
|
.and_then(|value| value.as_str().map(ToOwned::to_owned))
|
||||||
.unwrap_or_else(|| "internal".to_string());
|
.unwrap_or_else(|| "internal".to_string());
|
||||||
|
let error_message = error.message.trim().to_string();
|
||||||
let phase = serde_json::to_value(&error.phase).unwrap_or(Value::Null);
|
let phase = serde_json::to_value(&error.phase).unwrap_or(Value::Null);
|
||||||
let mut error_object = Map::from_iter([
|
let mut error_object = Map::from_iter([
|
||||||
("type".to_string(), Value::String(error_type.clone())),
|
|
||||||
("message".to_string(), Value::String(error.message.clone())),
|
|
||||||
("code".to_string(), Value::from(status_code)),
|
|
||||||
("phase".to_string(), phase),
|
("phase".to_string(), phase),
|
||||||
("retryable".to_string(), Value::Bool(error.retryable)),
|
("retryable".to_string(), Value::Bool(error.retryable)),
|
||||||
(
|
(
|
||||||
@@ -84,11 +120,8 @@ pub(super) fn build_stream_failure_from_execution_error(
|
|||||||
StreamFailureReport {
|
StreamFailureReport {
|
||||||
status_code,
|
status_code,
|
||||||
error_type,
|
error_type,
|
||||||
error_message: error.message.trim().to_string(),
|
error_message,
|
||||||
body_json: Value::Object(Map::from_iter([(
|
extra_error_fields: error_object,
|
||||||
"error".to_string(),
|
|
||||||
Value::Object(error_object),
|
|
||||||
)])),
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -96,23 +129,23 @@ fn build_stream_failure_sync_payload(
|
|||||||
trace_id: &str,
|
trace_id: &str,
|
||||||
report_kind: String,
|
report_kind: String,
|
||||||
report_context: Option<Value>,
|
report_context: Option<Value>,
|
||||||
headers: &std::collections::BTreeMap<String, String>,
|
mut headers: std::collections::BTreeMap<String, String>,
|
||||||
telemetry: Option<ExecutionTelemetry>,
|
telemetry: Option<ExecutionTelemetry>,
|
||||||
provider_buffered_body: &[u8],
|
provider_buffered_body: &[u8],
|
||||||
failure: &StreamFailureReport,
|
failure: StreamFailureReport,
|
||||||
) -> GatewaySyncReportRequest {
|
) -> GatewaySyncReportRequest {
|
||||||
let mut response_headers = headers.clone();
|
let status_code = failure.status_code;
|
||||||
response_headers.remove("content-encoding");
|
headers.remove("content-encoding");
|
||||||
response_headers.remove("content-length");
|
headers.remove("content-length");
|
||||||
response_headers.insert("content-type".to_string(), "application/json".to_string());
|
headers.insert("content-type".to_string(), "application/json".to_string());
|
||||||
|
|
||||||
GatewaySyncReportRequest {
|
GatewaySyncReportRequest {
|
||||||
trace_id: trace_id.to_string(),
|
trace_id: trace_id.to_string(),
|
||||||
report_kind,
|
report_kind,
|
||||||
report_context,
|
report_context,
|
||||||
status_code: failure.status_code,
|
status_code,
|
||||||
headers: response_headers,
|
headers,
|
||||||
body_json: Some(failure.body_json.clone()),
|
body_json: Some(failure.into_body_json()),
|
||||||
client_body_json: None,
|
client_body_json: None,
|
||||||
body_base64: (!provider_buffered_body.is_empty())
|
body_base64: (!provider_buffered_body.is_empty())
|
||||||
.then(|| base64::engine::general_purpose::STANDARD.encode(provider_buffered_body)),
|
.then(|| base64::engine::general_purpose::STANDARD.encode(provider_buffered_body)),
|
||||||
@@ -120,27 +153,40 @@ fn build_stream_failure_sync_payload(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn stream_failure_body_field<'a>(
|
||||||
|
payload: &'a GatewaySyncReportRequest,
|
||||||
|
field: &str,
|
||||||
|
) -> Option<&'a str> {
|
||||||
|
payload
|
||||||
|
.body_json
|
||||||
|
.as_ref()
|
||||||
|
.and_then(|body_json| body_json.get("error"))
|
||||||
|
.and_then(|value| value.get(field))
|
||||||
|
.and_then(Value::as_str)
|
||||||
|
}
|
||||||
|
|
||||||
async fn record_stream_sync_failure(
|
async fn record_stream_sync_failure(
|
||||||
state: &AppState,
|
state: &AppState,
|
||||||
plan: &ExecutionPlan,
|
plan: &ExecutionPlan,
|
||||||
report_context: Option<&Value>,
|
report_context: Option<&Value>,
|
||||||
payload: &GatewaySyncReportRequest,
|
payload: &GatewaySyncReportRequest,
|
||||||
failure: &StreamFailureReport,
|
|
||||||
started_at_unix_ms: Option<u64>,
|
started_at_unix_ms: Option<u64>,
|
||||||
) {
|
) {
|
||||||
let error_body = serde_json::to_string(&failure.body_json).ok();
|
let error_type = stream_failure_body_field(payload, "type").unwrap_or("internal");
|
||||||
|
let error_message = stream_failure_body_field(payload, "message").unwrap_or_default();
|
||||||
|
let error_body = payload
|
||||||
|
.body_json
|
||||||
|
.as_ref()
|
||||||
|
.and_then(|body_json| serde_json::to_string(body_json).ok());
|
||||||
let failure_analysis = resolve_local_failover_analysis_for_attempt(
|
let failure_analysis = resolve_local_failover_analysis_for_attempt(
|
||||||
state,
|
state,
|
||||||
plan,
|
plan,
|
||||||
report_context,
|
report_context,
|
||||||
failure.status_code,
|
payload.status_code,
|
||||||
error_body.as_deref(),
|
error_body.as_deref(),
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
if matches!(
|
if matches!(error_type, "first_byte_timeout" | "read_timeout") {
|
||||||
failure.error_type.as_str(),
|
|
||||||
"first_byte_timeout" | "read_timeout"
|
|
||||||
) {
|
|
||||||
apply_local_execution_effect(
|
apply_local_execution_effect(
|
||||||
state,
|
state,
|
||||||
LocalExecutionEffectContext {
|
LocalExecutionEffectContext {
|
||||||
@@ -158,7 +204,7 @@ async fn record_stream_sync_failure(
|
|||||||
report_context,
|
report_context,
|
||||||
},
|
},
|
||||||
LocalExecutionEffect::AttemptFailure(LocalAttemptFailureEffect {
|
LocalExecutionEffect::AttemptFailure(LocalAttemptFailureEffect {
|
||||||
status_code: failure.status_code,
|
status_code: payload.status_code,
|
||||||
classification: failure_analysis.classification,
|
classification: failure_analysis.classification,
|
||||||
}),
|
}),
|
||||||
)
|
)
|
||||||
@@ -170,7 +216,7 @@ async fn record_stream_sync_failure(
|
|||||||
report_context,
|
report_context,
|
||||||
},
|
},
|
||||||
LocalExecutionEffect::AdaptiveRateLimit(LocalAdaptiveRateLimitEffect {
|
LocalExecutionEffect::AdaptiveRateLimit(LocalAdaptiveRateLimitEffect {
|
||||||
status_code: failure.status_code,
|
status_code: payload.status_code,
|
||||||
classification: failure_analysis.classification,
|
classification: failure_analysis.classification,
|
||||||
headers: Some(&payload.headers),
|
headers: Some(&payload.headers),
|
||||||
}),
|
}),
|
||||||
@@ -183,7 +229,7 @@ async fn record_stream_sync_failure(
|
|||||||
report_context,
|
report_context,
|
||||||
},
|
},
|
||||||
LocalExecutionEffect::HealthFailure(LocalHealthFailureEffect {
|
LocalExecutionEffect::HealthFailure(LocalHealthFailureEffect {
|
||||||
status_code: failure.status_code,
|
status_code: payload.status_code,
|
||||||
classification: failure_analysis.classification,
|
classification: failure_analysis.classification,
|
||||||
}),
|
}),
|
||||||
)
|
)
|
||||||
@@ -195,7 +241,7 @@ async fn record_stream_sync_failure(
|
|||||||
report_context,
|
report_context,
|
||||||
},
|
},
|
||||||
LocalExecutionEffect::OauthInvalidation(LocalOAuthInvalidationEffect {
|
LocalExecutionEffect::OauthInvalidation(LocalOAuthInvalidationEffect {
|
||||||
status_code: failure.status_code,
|
status_code: payload.status_code,
|
||||||
response_text: error_body.as_deref(),
|
response_text: error_body.as_deref(),
|
||||||
}),
|
}),
|
||||||
)
|
)
|
||||||
@@ -207,7 +253,7 @@ async fn record_stream_sync_failure(
|
|||||||
report_context,
|
report_context,
|
||||||
},
|
},
|
||||||
LocalExecutionEffect::PoolError(LocalPoolErrorEffect {
|
LocalExecutionEffect::PoolError(LocalPoolErrorEffect {
|
||||||
status_code: failure.status_code,
|
status_code: payload.status_code,
|
||||||
classification: failure_analysis.classification,
|
classification: failure_analysis.classification,
|
||||||
headers: &payload.headers,
|
headers: &payload.headers,
|
||||||
error_body: error_body.as_deref(),
|
error_body: error_body.as_deref(),
|
||||||
@@ -218,16 +264,16 @@ async fn record_stream_sync_failure(
|
|||||||
let payload_seed = build_sync_terminal_usage_payload_seed(payload);
|
let payload_seed = build_sync_terminal_usage_payload_seed(payload);
|
||||||
state
|
state
|
||||||
.usage_runtime
|
.usage_runtime
|
||||||
.record_sync_terminal(state.data.as_ref(), &context_seed, &payload_seed);
|
.record_sync_terminal(state.data.as_ref(), context_seed, payload_seed);
|
||||||
let terminal_unix_secs = current_request_candidate_unix_ms();
|
let terminal_unix_secs = current_request_candidate_unix_ms();
|
||||||
record_report_request_candidate_status(
|
record_report_request_candidate_status(
|
||||||
state,
|
state,
|
||||||
report_context,
|
report_context,
|
||||||
SchedulerRequestCandidateStatusUpdate {
|
SchedulerRequestCandidateStatusUpdate {
|
||||||
status: RequestCandidateStatus::Failed,
|
status: RequestCandidateStatus::Failed,
|
||||||
status_code: Some(failure.status_code),
|
status_code: Some(payload.status_code),
|
||||||
error_type: Some(failure.error_type.clone()),
|
error_type: Some(error_type.to_string()),
|
||||||
error_message: Some(failure.error_message.clone()),
|
error_message: Some(error_message.to_string()),
|
||||||
latency_ms: payload
|
latency_ms: payload
|
||||||
.telemetry
|
.telemetry
|
||||||
.as_ref()
|
.as_ref()
|
||||||
@@ -249,7 +295,7 @@ pub(super) async fn handle_prefetch_stream_failure(
|
|||||||
request_id: &str,
|
request_id: &str,
|
||||||
candidate_id: Option<&str>,
|
candidate_id: Option<&str>,
|
||||||
report_kind: &str,
|
report_kind: &str,
|
||||||
headers: &std::collections::BTreeMap<String, String>,
|
headers: std::collections::BTreeMap<String, String>,
|
||||||
telemetry: Option<ExecutionTelemetry>,
|
telemetry: Option<ExecutionTelemetry>,
|
||||||
buffered_body: &[u8],
|
buffered_body: &[u8],
|
||||||
failure: StreamFailureReport,
|
failure: StreamFailureReport,
|
||||||
@@ -257,21 +303,13 @@ pub(super) async fn handle_prefetch_stream_failure(
|
|||||||
let payload = build_stream_failure_sync_payload(
|
let payload = build_stream_failure_sync_payload(
|
||||||
trace_id,
|
trace_id,
|
||||||
report_kind.to_string(),
|
report_kind.to_string(),
|
||||||
report_context.clone(),
|
report_context,
|
||||||
headers,
|
headers,
|
||||||
telemetry,
|
telemetry,
|
||||||
buffered_body,
|
buffered_body,
|
||||||
&failure,
|
failure,
|
||||||
);
|
);
|
||||||
record_stream_sync_failure(
|
record_stream_sync_failure(state, plan, payload.report_context.as_ref(), &payload, None).await;
|
||||||
state,
|
|
||||||
plan,
|
|
||||||
report_context.as_ref(),
|
|
||||||
&payload,
|
|
||||||
&failure,
|
|
||||||
None,
|
|
||||||
)
|
|
||||||
.await;
|
|
||||||
|
|
||||||
let response =
|
let response =
|
||||||
submit_local_core_error_or_sync_finalize(state, trace_id, decision, payload).await?;
|
submit_local_core_error_or_sync_finalize(state, trace_id, decision, payload).await?;
|
||||||
@@ -287,8 +325,8 @@ pub(super) async fn submit_midstream_stream_failure(
|
|||||||
trace_id: &str,
|
trace_id: &str,
|
||||||
plan: &ExecutionPlan,
|
plan: &ExecutionPlan,
|
||||||
direct_stream_finalize_kind: Option<&str>,
|
direct_stream_finalize_kind: Option<&str>,
|
||||||
report_context: Option<&Value>,
|
report_context: Option<Value>,
|
||||||
headers: &std::collections::BTreeMap<String, String>,
|
headers: std::collections::BTreeMap<String, String>,
|
||||||
telemetry: Option<ExecutionTelemetry>,
|
telemetry: Option<ExecutionTelemetry>,
|
||||||
buffered_body: &[u8],
|
buffered_body: &[u8],
|
||||||
started_at_unix_ms: u64,
|
started_at_unix_ms: u64,
|
||||||
@@ -303,22 +341,21 @@ pub(super) async fn submit_midstream_stream_failure(
|
|||||||
let payload = build_stream_failure_sync_payload(
|
let payload = build_stream_failure_sync_payload(
|
||||||
trace_id,
|
trace_id,
|
||||||
report_kind,
|
report_kind,
|
||||||
report_context.cloned(),
|
report_context,
|
||||||
headers,
|
headers,
|
||||||
telemetry,
|
telemetry,
|
||||||
buffered_body,
|
buffered_body,
|
||||||
&failure,
|
failure,
|
||||||
);
|
);
|
||||||
record_stream_sync_failure(
|
record_stream_sync_failure(
|
||||||
state,
|
state,
|
||||||
plan,
|
plan,
|
||||||
report_context,
|
payload.report_context.as_ref(),
|
||||||
&payload,
|
&payload,
|
||||||
&failure,
|
|
||||||
Some(started_at_unix_ms),
|
Some(started_at_unix_ms),
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
if let Err(err) = submit_sync_report(state, trace_id, payload).await {
|
if let Err(err) = submit_sync_report(state, payload).await {
|
||||||
let request_id = short_request_id(plan.request_id.as_str());
|
let request_id = short_request_id(plan.request_id.as_str());
|
||||||
warn!(
|
warn!(
|
||||||
event_name = "execution_report_submit_failed",
|
event_name = "execution_report_submit_failed",
|
||||||
|
|||||||
@@ -274,7 +274,7 @@ fn format_error_chain(err: &(dyn std::error::Error + 'static)) -> String {
|
|||||||
fn observe_stream_chunk(
|
fn observe_stream_chunk(
|
||||||
observer: &mut StreamingStandardTerminalObserver,
|
observer: &mut StreamingStandardTerminalObserver,
|
||||||
report_context: &Value,
|
report_context: &Value,
|
||||||
private_stream_normalizer: Option<&mut crate::ai_pipeline::ProviderPrivateStreamNormalizer>,
|
private_stream_normalizer: Option<&mut crate::ai_pipeline::ProviderPrivateStreamNormalizer<'_>>,
|
||||||
observer_buffered: &mut Vec<u8>,
|
observer_buffered: &mut Vec<u8>,
|
||||||
chunk: &[u8],
|
chunk: &[u8],
|
||||||
) {
|
) {
|
||||||
@@ -298,7 +298,7 @@ fn observe_stream_chunk(
|
|||||||
fn finalize_stream_terminal_summary(
|
fn finalize_stream_terminal_summary(
|
||||||
observer: &mut StreamingStandardTerminalObserver,
|
observer: &mut StreamingStandardTerminalObserver,
|
||||||
report_context: &Value,
|
report_context: &Value,
|
||||||
private_stream_normalizer: Option<&mut crate::ai_pipeline::ProviderPrivateStreamNormalizer>,
|
private_stream_normalizer: Option<&mut crate::ai_pipeline::ProviderPrivateStreamNormalizer<'_>>,
|
||||||
observer_buffered: &mut Vec<u8>,
|
observer_buffered: &mut Vec<u8>,
|
||||||
) -> Option<ExecutionStreamTerminalSummary> {
|
) -> Option<ExecutionStreamTerminalSummary> {
|
||||||
if let Some(normalizer) = private_stream_normalizer {
|
if let Some(normalizer) = private_stream_normalizer {
|
||||||
@@ -412,7 +412,7 @@ mod tests {
|
|||||||
});
|
});
|
||||||
|
|
||||||
let execution = DirectSyncExecutionRuntime::new()
|
let execution = DirectSyncExecutionRuntime::new()
|
||||||
.execute_stream(ExecutionPlan {
|
.execute_stream(&ExecutionPlan {
|
||||||
request_id: "req-stream-ttfb-1".into(),
|
request_id: "req-stream-ttfb-1".into(),
|
||||||
candidate_id: Some("cand-stream-ttfb-1".into()),
|
candidate_id: Some("cand-stream-ttfb-1".into()),
|
||||||
provider_name: Some("openai".into()),
|
provider_name: Some("openai".into()),
|
||||||
@@ -499,7 +499,7 @@ mod tests {
|
|||||||
|
|
||||||
let runtime = DirectSyncExecutionRuntime::new();
|
let runtime = DirectSyncExecutionRuntime::new();
|
||||||
let execution = runtime
|
let execution = runtime
|
||||||
.execute_stream(ExecutionPlan {
|
.execute_stream(&ExecutionPlan {
|
||||||
request_id: "req-telemetry-order".to_string(),
|
request_id: "req-telemetry-order".to_string(),
|
||||||
candidate_id: Some("cand-telemetry-order".to_string()),
|
candidate_id: Some("cand-telemetry-order".to_string()),
|
||||||
provider_name: Some("OpenAI".to_string()),
|
provider_name: Some("OpenAI".to_string()),
|
||||||
|
|||||||
@@ -545,7 +545,7 @@ pub(crate) async fn submit_local_core_error_or_sync_finalize(
|
|||||||
{
|
{
|
||||||
let mut report_payload = payload.clone();
|
let mut report_payload = payload.clone();
|
||||||
report_payload.report_kind = error_report_kind;
|
report_payload.report_kind = error_report_kind;
|
||||||
spawn_sync_report(state.clone(), trace_id.to_string(), report_payload);
|
spawn_sync_report(state.clone(), report_payload);
|
||||||
} else {
|
} else {
|
||||||
warn!(
|
warn!(
|
||||||
event_name = "local_core_finalize_missing_error_report_mapping",
|
event_name = "local_core_finalize_missing_error_report_mapping",
|
||||||
|
|||||||
@@ -23,7 +23,6 @@ use crate::api::response::{
|
|||||||
attach_control_metadata_headers, build_client_response, build_client_response_from_parts,
|
attach_control_metadata_headers, build_client_response, build_client_response_from_parts,
|
||||||
};
|
};
|
||||||
use crate::clock::current_unix_ms as current_request_candidate_unix_ms;
|
use crate::clock::current_unix_ms as current_request_candidate_unix_ms;
|
||||||
use crate::constants::{CONTROL_CANDIDATE_ID_HEADER, CONTROL_REQUEST_ID_HEADER};
|
|
||||||
use crate::control::GatewayControlDecision;
|
use crate::control::GatewayControlDecision;
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
use crate::execution_runtime::remote_compat::post_sync_plan_to_remote_execution_runtime;
|
use crate::execution_runtime::remote_compat::post_sync_plan_to_remote_execution_runtime;
|
||||||
@@ -57,7 +56,8 @@ use policy::decode_execution_result_body;
|
|||||||
pub(crate) use response::{
|
pub(crate) use response::{
|
||||||
maybe_build_local_sync_finalize_response, maybe_build_local_video_error_response,
|
maybe_build_local_sync_finalize_response, maybe_build_local_video_error_response,
|
||||||
maybe_build_local_video_success_outcome, resolve_local_sync_error_background_report_kind,
|
maybe_build_local_video_success_outcome, resolve_local_sync_error_background_report_kind,
|
||||||
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessOutcome,
|
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessBuild,
|
||||||
|
LocalVideoSyncSuccessOutcome,
|
||||||
};
|
};
|
||||||
|
|
||||||
struct ImplicitSyncFinalizeOutcome {
|
struct ImplicitSyncFinalizeOutcome {
|
||||||
@@ -75,7 +75,65 @@ fn record_sync_terminal_usage(
|
|||||||
let payload_seed = build_sync_terminal_usage_payload_seed(payload);
|
let payload_seed = build_sync_terminal_usage_payload_seed(payload);
|
||||||
state
|
state
|
||||||
.usage_runtime
|
.usage_runtime
|
||||||
.record_sync_terminal(state.data.as_ref(), &context_seed, &payload_seed);
|
.record_sync_terminal(state.data.as_ref(), context_seed, payload_seed);
|
||||||
|
}
|
||||||
|
|
||||||
|
fn build_sync_report_payload(
|
||||||
|
trace_id: &str,
|
||||||
|
report_kind: String,
|
||||||
|
report_context: Option<serde_json::Value>,
|
||||||
|
status_code: u16,
|
||||||
|
headers: BTreeMap<String, String>,
|
||||||
|
body_json: Option<serde_json::Value>,
|
||||||
|
body_base64: Option<String>,
|
||||||
|
telemetry: Option<ExecutionTelemetry>,
|
||||||
|
) -> GatewaySyncReportRequest {
|
||||||
|
GatewaySyncReportRequest {
|
||||||
|
trace_id: trace_id.to_string(),
|
||||||
|
report_kind,
|
||||||
|
report_context,
|
||||||
|
status_code,
|
||||||
|
headers,
|
||||||
|
body_json,
|
||||||
|
client_body_json: None,
|
||||||
|
body_base64,
|
||||||
|
telemetry,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn apply_sync_success_effects(
|
||||||
|
state: &AppState,
|
||||||
|
plan: &ExecutionPlan,
|
||||||
|
report_context: Option<&serde_json::Value>,
|
||||||
|
payload: &GatewaySyncReportRequest,
|
||||||
|
) {
|
||||||
|
apply_local_execution_effect(
|
||||||
|
state,
|
||||||
|
LocalExecutionEffectContext {
|
||||||
|
plan,
|
||||||
|
report_context,
|
||||||
|
},
|
||||||
|
LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect),
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
apply_local_execution_effect(
|
||||||
|
state,
|
||||||
|
LocalExecutionEffectContext {
|
||||||
|
plan,
|
||||||
|
report_context,
|
||||||
|
},
|
||||||
|
LocalExecutionEffect::AdaptiveSuccess(LocalAdaptiveSuccessEffect),
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
apply_local_execution_effect(
|
||||||
|
state,
|
||||||
|
LocalExecutionEffectContext {
|
||||||
|
plan,
|
||||||
|
report_context,
|
||||||
|
},
|
||||||
|
LocalExecutionEffect::PoolSuccessSync { payload },
|
||||||
|
)
|
||||||
|
.await;
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
@@ -112,7 +170,7 @@ pub(crate) async fn execute_execution_runtime_sync(
|
|||||||
let lifecycle_seed = build_lifecycle_usage_seed(&plan, report_context.as_ref());
|
let lifecycle_seed = build_lifecycle_usage_seed(&plan, report_context.as_ref());
|
||||||
state
|
state
|
||||||
.usage_runtime
|
.usage_runtime
|
||||||
.record_pending(state.data.as_ref(), &lifecycle_seed);
|
.record_pending(state.data.as_ref(), lifecycle_seed);
|
||||||
record_local_request_candidate_status(
|
record_local_request_candidate_status(
|
||||||
state,
|
state,
|
||||||
&plan,
|
&plan,
|
||||||
@@ -129,11 +187,8 @@ pub(crate) async fn execute_execution_runtime_sync(
|
|||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
#[cfg(not(test))]
|
#[cfg(not(test))]
|
||||||
let result = {
|
let mut result = {
|
||||||
match DirectSyncExecutionRuntime::new()
|
match DirectSyncExecutionRuntime::new().execute_sync(&plan).await {
|
||||||
.execute_sync(plan.clone())
|
|
||||||
.await
|
|
||||||
{
|
|
||||||
Ok(result) => result,
|
Ok(result) => result,
|
||||||
Err(err) => {
|
Err(err) => {
|
||||||
warn!(
|
warn!(
|
||||||
@@ -171,7 +226,7 @@ pub(crate) async fn execute_execution_runtime_sync(
|
|||||||
}
|
}
|
||||||
};
|
};
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
let result = {
|
let mut result = {
|
||||||
if let Some(override_fn) = state.execution_runtime_sync_override.as_ref() {
|
if let Some(override_fn) = state.execution_runtime_sync_override.as_ref() {
|
||||||
match (override_fn.0)(&plan) {
|
match (override_fn.0)(&plan) {
|
||||||
Ok(result) => result,
|
Ok(result) => result,
|
||||||
@@ -215,10 +270,7 @@ pub(crate) async fn execute_execution_runtime_sync(
|
|||||||
.trim()
|
.trim()
|
||||||
.is_empty()
|
.is_empty()
|
||||||
{
|
{
|
||||||
match DirectSyncExecutionRuntime::new()
|
match DirectSyncExecutionRuntime::new().execute_sync(&plan).await {
|
||||||
.execute_sync(plan.clone())
|
|
||||||
.await
|
|
||||||
{
|
|
||||||
Ok(result) => result,
|
Ok(result) => result,
|
||||||
Err(err) => {
|
Err(err) => {
|
||||||
warn!(
|
warn!(
|
||||||
@@ -287,8 +339,9 @@ pub(crate) async fn execute_execution_runtime_sync(
|
|||||||
.telemetry
|
.telemetry
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.and_then(|telemetry| telemetry.elapsed_ms);
|
.and_then(|telemetry| telemetry.elapsed_ms);
|
||||||
let mut headers = result.headers.clone();
|
let mut headers = std::mem::take(&mut result.headers);
|
||||||
let (body_bytes, body_json, body_base64) = decode_execution_result_body(&result, &mut headers)?;
|
let (body_bytes, body_json, body_base64) =
|
||||||
|
decode_execution_result_body(result.body.take(), &mut headers)?;
|
||||||
let local_failover_response_text = local_failover_response_text(
|
let local_failover_response_text = local_failover_response_text(
|
||||||
body_json.as_ref(),
|
body_json.as_ref(),
|
||||||
&body_bytes,
|
&body_bytes,
|
||||||
@@ -403,38 +456,26 @@ pub(crate) async fn execute_execution_runtime_sync(
|
|||||||
);
|
);
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
let request_id = (!result.request_id.trim().is_empty())
|
let status_code = result.status_code;
|
||||||
.then_some(result.request_id.as_str())
|
|
||||||
.or(Some(plan_request_id));
|
|
||||||
let request_id_for_log = short_request_id(request_id.unwrap_or("-"));
|
|
||||||
let candidate_id = result.candidate_id.as_deref().or(plan_candidate_id);
|
|
||||||
let has_body_bytes = body_base64.is_some();
|
let has_body_bytes = body_base64.is_some();
|
||||||
let explicit_finalize = should_finalize_sync_response(report_kind.as_deref());
|
let explicit_finalize = should_finalize_sync_response(report_kind.as_deref());
|
||||||
let mapped_error_finalize_kind =
|
let mapped_error_finalize_kind =
|
||||||
resolve_core_sync_error_finalize_report_kind(plan_kind, &result, body_json.as_ref());
|
resolve_core_sync_error_finalize_report_kind(plan_kind, &result, body_json.as_ref());
|
||||||
let implicit_finalize = if explicit_finalize || mapped_error_finalize_kind.is_some() {
|
let implicit_finalize = if !explicit_finalize && mapped_error_finalize_kind.is_none() {
|
||||||
None
|
|
||||||
} else {
|
|
||||||
maybe_build_implicit_sync_finalize_outcome(
|
maybe_build_implicit_sync_finalize_outcome(
|
||||||
trace_id,
|
trace_id,
|
||||||
decision,
|
decision,
|
||||||
plan_kind,
|
plan_kind,
|
||||||
report_context.clone(),
|
&report_context,
|
||||||
result.status_code,
|
status_code,
|
||||||
headers.clone(),
|
&headers,
|
||||||
body_json.clone(),
|
&body_json,
|
||||||
body_base64.clone(),
|
&body_base64,
|
||||||
result.telemetry.clone(),
|
&result.telemetry,
|
||||||
)?
|
)?
|
||||||
};
|
|
||||||
let finalize_report_kind = if explicit_finalize {
|
|
||||||
report_kind.clone()
|
|
||||||
} else if let Some(implicit_finalize) = implicit_finalize.as_ref() {
|
|
||||||
Some(implicit_finalize.payload.report_kind.clone())
|
|
||||||
} else {
|
} else {
|
||||||
mapped_error_finalize_kind.clone()
|
None
|
||||||
};
|
};
|
||||||
|
|
||||||
if !matches!(
|
if !matches!(
|
||||||
local_failover_analysis.decision,
|
local_failover_analysis.decision,
|
||||||
LocalFailoverDecision::StopLocalFailover
|
LocalFailoverDecision::StopLocalFailover
|
||||||
@@ -486,91 +527,83 @@ pub(crate) async fn execute_execution_runtime_sync(
|
|||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
|
|
||||||
let base_usage_payload = GatewaySyncReportRequest {
|
let request_id_owned = result.request_id;
|
||||||
trace_id: trace_id.to_string(),
|
let candidate_id_owned = result.candidate_id;
|
||||||
report_kind: finalize_report_kind
|
let request_id = (!request_id_owned.trim().is_empty())
|
||||||
.clone()
|
.then_some(request_id_owned.as_str())
|
||||||
.or_else(|| report_kind.clone())
|
.or(Some(plan_request_id));
|
||||||
.unwrap_or_default(),
|
let request_id_for_log = short_request_id(request_id.unwrap_or("-"));
|
||||||
report_context: report_context.clone(),
|
let candidate_id = candidate_id_owned.as_deref().or(plan_candidate_id);
|
||||||
status_code: result.status_code,
|
let report_context = report_context;
|
||||||
headers: headers.clone(),
|
let headers = headers;
|
||||||
body_json: body_json.clone(),
|
let body_json = body_json;
|
||||||
client_body_json: None,
|
let telemetry = result.telemetry;
|
||||||
body_base64: body_base64.clone(),
|
|
||||||
telemetry: result.telemetry.clone(),
|
if let Some(implicit_finalize) = implicit_finalize {
|
||||||
};
|
let usage_payload = implicit_finalize
|
||||||
if result.status_code < 400 {
|
.outcome
|
||||||
apply_local_execution_effect(
|
.background_report
|
||||||
|
.as_ref()
|
||||||
|
.unwrap_or(&implicit_finalize.payload);
|
||||||
|
apply_sync_success_effects(
|
||||||
state,
|
state,
|
||||||
LocalExecutionEffectContext {
|
&plan,
|
||||||
plan: &plan,
|
implicit_finalize.payload.report_context.as_ref(),
|
||||||
report_context: report_context.as_ref(),
|
usage_payload,
|
||||||
},
|
|
||||||
LocalExecutionEffect::HealthSuccess(LocalHealthSuccessEffect),
|
|
||||||
)
|
)
|
||||||
.await;
|
.await;
|
||||||
apply_local_execution_effect(
|
record_sync_terminal_usage(
|
||||||
state,
|
state,
|
||||||
LocalExecutionEffectContext {
|
&plan,
|
||||||
plan: &plan,
|
implicit_finalize.payload.report_context.as_ref(),
|
||||||
report_context: report_context.as_ref(),
|
usage_payload,
|
||||||
},
|
);
|
||||||
LocalExecutionEffect::AdaptiveSuccess(LocalAdaptiveSuccessEffect),
|
if let Some(report_payload) = implicit_finalize.outcome.background_report {
|
||||||
)
|
spawn_sync_report(state.clone(), report_payload);
|
||||||
.await;
|
} else {
|
||||||
apply_local_execution_effect(
|
warn!(
|
||||||
state,
|
event_name = "local_core_finalize_missing_success_report_mapping",
|
||||||
LocalExecutionEffectContext {
|
log_type = "event",
|
||||||
plan: &plan,
|
trace_id = %trace_id,
|
||||||
report_context: report_context.as_ref(),
|
report_kind = %implicit_finalize.payload.report_kind,
|
||||||
},
|
"gateway implicit local core finalize produced response without background success report mapping"
|
||||||
LocalExecutionEffect::PoolSuccessSync {
|
);
|
||||||
payload: &base_usage_payload,
|
}
|
||||||
},
|
return Ok(Some(attach_control_metadata_headers(
|
||||||
)
|
implicit_finalize.outcome.response,
|
||||||
.await;
|
request_id,
|
||||||
|
candidate_id,
|
||||||
|
)?));
|
||||||
}
|
}
|
||||||
|
|
||||||
if let Some(finalize_report_kind) = finalize_report_kind {
|
let finalize_report_kind = if explicit_finalize {
|
||||||
if let Some(implicit_finalize) = implicit_finalize {
|
report_kind.clone()
|
||||||
let usage_payload = implicit_finalize
|
} else {
|
||||||
.outcome
|
mapped_error_finalize_kind
|
||||||
.background_report
|
};
|
||||||
.as_ref()
|
|
||||||
.unwrap_or(&implicit_finalize.payload);
|
|
||||||
record_sync_terminal_usage(state, &plan, report_context.as_ref(), usage_payload);
|
|
||||||
if let Some(report_payload) = implicit_finalize.outcome.background_report {
|
|
||||||
spawn_sync_report(state.clone(), trace_id.to_string(), report_payload);
|
|
||||||
} else {
|
|
||||||
warn!(
|
|
||||||
event_name = "local_core_finalize_missing_success_report_mapping",
|
|
||||||
log_type = "event",
|
|
||||||
trace_id = %trace_id,
|
|
||||||
report_kind = %implicit_finalize.payload.report_kind,
|
|
||||||
"gateway implicit local core finalize produced response without background success report mapping"
|
|
||||||
);
|
|
||||||
}
|
|
||||||
return Ok(Some(attach_control_metadata_headers(
|
|
||||||
implicit_finalize.outcome.response,
|
|
||||||
request_id,
|
|
||||||
candidate_id,
|
|
||||||
)?));
|
|
||||||
}
|
|
||||||
|
|
||||||
let payload = GatewaySyncReportRequest {
|
if let Some(finalize_report_kind) = finalize_report_kind {
|
||||||
trace_id: trace_id.to_string(),
|
let mut payload = build_sync_report_payload(
|
||||||
report_kind: finalize_report_kind,
|
trace_id,
|
||||||
|
finalize_report_kind,
|
||||||
report_context,
|
report_context,
|
||||||
status_code: result.status_code,
|
status_code,
|
||||||
headers: headers.clone(),
|
headers,
|
||||||
body_json: body_json.clone(),
|
body_json,
|
||||||
client_body_json: None,
|
body_base64,
|
||||||
body_base64: body_base64.clone(),
|
telemetry,
|
||||||
telemetry: result.telemetry.clone(),
|
);
|
||||||
};
|
|
||||||
if let Some(outcome) = maybe_build_sync_finalize_outcome(trace_id, decision, &payload)? {
|
if let Some(outcome) = maybe_build_sync_finalize_outcome(trace_id, decision, &payload)? {
|
||||||
let usage_payload = outcome.background_report.as_ref().unwrap_or(&payload);
|
let usage_payload = outcome.background_report.as_ref().unwrap_or(&payload);
|
||||||
|
if status_code < 400 {
|
||||||
|
apply_sync_success_effects(
|
||||||
|
state,
|
||||||
|
&plan,
|
||||||
|
payload.report_context.as_ref(),
|
||||||
|
usage_payload,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
}
|
||||||
record_sync_terminal_usage(
|
record_sync_terminal_usage(
|
||||||
state,
|
state,
|
||||||
&plan,
|
&plan,
|
||||||
@@ -578,7 +611,7 @@ pub(crate) async fn execute_execution_runtime_sync(
|
|||||||
usage_payload,
|
usage_payload,
|
||||||
);
|
);
|
||||||
if let Some(report_payload) = outcome.background_report {
|
if let Some(report_payload) = outcome.background_report {
|
||||||
spawn_sync_report(state.clone(), trace_id.to_string(), report_payload);
|
spawn_sync_report(state.clone(), report_payload);
|
||||||
} else {
|
} else {
|
||||||
warn!(
|
warn!(
|
||||||
event_name = "local_core_finalize_missing_success_report_mapping",
|
event_name = "local_core_finalize_missing_success_report_mapping",
|
||||||
@@ -594,55 +627,62 @@ pub(crate) async fn execute_execution_runtime_sync(
|
|||||||
candidate_id,
|
candidate_id,
|
||||||
)?));
|
)?));
|
||||||
}
|
}
|
||||||
if let Some(outcome) = maybe_build_local_video_success_outcome(
|
let mut payload = match maybe_build_local_video_success_outcome(
|
||||||
trace_id,
|
trace_id,
|
||||||
decision,
|
decision,
|
||||||
&payload,
|
payload,
|
||||||
&state.video_tasks,
|
&state.video_tasks,
|
||||||
&plan,
|
&plan,
|
||||||
)? {
|
)? {
|
||||||
record_sync_terminal_usage(
|
LocalVideoSyncSuccessBuild::Handled(outcome) => {
|
||||||
state,
|
let LocalVideoSyncSuccessOutcome {
|
||||||
&plan,
|
response,
|
||||||
payload.report_context.as_ref(),
|
report_payload,
|
||||||
&outcome.report_payload,
|
original_report_context,
|
||||||
);
|
report_mode,
|
||||||
if let Some(snapshot) = outcome.local_task_snapshot.clone() {
|
local_task_snapshot,
|
||||||
state.video_tasks.record_snapshot(snapshot.clone());
|
} = outcome;
|
||||||
let _ = state.upsert_video_task_snapshot(&snapshot).await?;
|
apply_sync_success_effects(
|
||||||
}
|
state,
|
||||||
match outcome.report_mode {
|
&plan,
|
||||||
VideoTaskSyncReportMode::InlineSync => {
|
original_report_context.as_ref(),
|
||||||
submit_sync_report(state, trace_id, outcome.report_payload).await?;
|
&report_payload,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
record_sync_terminal_usage(
|
||||||
|
state,
|
||||||
|
&plan,
|
||||||
|
original_report_context.as_ref(),
|
||||||
|
&report_payload,
|
||||||
|
);
|
||||||
|
if let Some(snapshot) = local_task_snapshot {
|
||||||
|
let _ = state.upsert_video_task_snapshot(&snapshot).await?;
|
||||||
|
state.video_tasks.record_snapshot(snapshot);
|
||||||
}
|
}
|
||||||
VideoTaskSyncReportMode::Background => {
|
match report_mode {
|
||||||
spawn_sync_report(state.clone(), trace_id.to_string(), outcome.report_payload);
|
VideoTaskSyncReportMode::InlineSync => {
|
||||||
|
submit_sync_report(state, report_payload).await?;
|
||||||
|
}
|
||||||
|
VideoTaskSyncReportMode::Background => {
|
||||||
|
spawn_sync_report(state.clone(), report_payload);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
return Ok(Some(attach_control_metadata_headers(
|
||||||
|
response,
|
||||||
|
request_id,
|
||||||
|
candidate_id,
|
||||||
|
)?));
|
||||||
}
|
}
|
||||||
return Ok(Some(attach_control_metadata_headers(
|
LocalVideoSyncSuccessBuild::NotHandled(payload) => payload,
|
||||||
outcome.response,
|
};
|
||||||
request_id,
|
|
||||||
candidate_id,
|
|
||||||
)?));
|
|
||||||
}
|
|
||||||
if let Some(response) =
|
if let Some(response) =
|
||||||
maybe_build_local_sync_finalize_response(trace_id, decision, &payload)?
|
maybe_build_local_sync_finalize_response(trace_id, decision, &payload)?
|
||||||
{
|
{
|
||||||
let usage_payload = if let Some(success_report_kind) =
|
let background_success_report_kind =
|
||||||
resolve_local_sync_success_background_report_kind(payload.report_kind.as_str())
|
resolve_local_sync_success_background_report_kind(payload.report_kind.as_str());
|
||||||
{
|
apply_sync_success_effects(state, &plan, payload.report_context.as_ref(), &payload)
|
||||||
let mut report_payload = payload.clone();
|
.await;
|
||||||
report_payload.report_kind = success_report_kind.to_string();
|
record_sync_terminal_usage(state, &plan, payload.report_context.as_ref(), &payload);
|
||||||
report_payload
|
|
||||||
} else {
|
|
||||||
payload.clone()
|
|
||||||
};
|
|
||||||
record_sync_terminal_usage(
|
|
||||||
state,
|
|
||||||
&plan,
|
|
||||||
payload.report_context.as_ref(),
|
|
||||||
&usage_payload,
|
|
||||||
);
|
|
||||||
state
|
state
|
||||||
.video_tasks
|
.video_tasks
|
||||||
.apply_finalize_mutation(request_path, payload.report_kind.as_str());
|
.apply_finalize_mutation(request_path, payload.report_kind.as_str());
|
||||||
@@ -652,12 +692,11 @@ pub(crate) async fn execute_execution_runtime_sync(
|
|||||||
{
|
{
|
||||||
let _ = state.upsert_video_task_snapshot(&snapshot).await?;
|
let _ = state.upsert_video_task_snapshot(&snapshot).await?;
|
||||||
}
|
}
|
||||||
if let Some(success_report_kind) =
|
if let Some(success_report_kind) = background_success_report_kind {
|
||||||
resolve_local_sync_success_background_report_kind(payload.report_kind.as_str())
|
payload.report_kind = success_report_kind.to_string();
|
||||||
{
|
}
|
||||||
let mut report_payload = usage_payload;
|
if background_success_report_kind.is_some() {
|
||||||
report_payload.report_kind = success_report_kind.to_string();
|
spawn_sync_report(state.clone(), payload);
|
||||||
spawn_sync_report(state.clone(), trace_id.to_string(), report_payload);
|
|
||||||
} else {
|
} else {
|
||||||
warn!(
|
warn!(
|
||||||
event_name = "local_video_finalize_missing_success_report_mapping",
|
event_name = "local_video_finalize_missing_success_report_mapping",
|
||||||
@@ -678,27 +717,14 @@ pub(crate) async fn execute_execution_runtime_sync(
|
|||||||
if let Some(response) =
|
if let Some(response) =
|
||||||
maybe_build_local_video_error_response(trace_id, decision, &payload)?
|
maybe_build_local_video_error_response(trace_id, decision, &payload)?
|
||||||
{
|
{
|
||||||
let usage_payload = if let Some(error_report_kind) =
|
let background_error_report_kind =
|
||||||
resolve_local_sync_error_background_report_kind(payload.report_kind.as_str())
|
resolve_local_sync_error_background_report_kind(payload.report_kind.as_str());
|
||||||
{
|
if let Some(error_report_kind) = background_error_report_kind {
|
||||||
let mut report_payload = payload.clone();
|
payload.report_kind = error_report_kind.to_string();
|
||||||
report_payload.report_kind = error_report_kind.to_string();
|
}
|
||||||
report_payload
|
record_sync_terminal_usage(state, &plan, payload.report_context.as_ref(), &payload);
|
||||||
} else {
|
if background_error_report_kind.is_some() {
|
||||||
payload.clone()
|
spawn_sync_report(state.clone(), payload);
|
||||||
};
|
|
||||||
record_sync_terminal_usage(
|
|
||||||
state,
|
|
||||||
&plan,
|
|
||||||
payload.report_context.as_ref(),
|
|
||||||
&usage_payload,
|
|
||||||
);
|
|
||||||
if let Some(error_report_kind) =
|
|
||||||
resolve_local_sync_error_background_report_kind(payload.report_kind.as_str())
|
|
||||||
{
|
|
||||||
let mut report_payload = usage_payload;
|
|
||||||
report_payload.report_kind = error_report_kind.to_string();
|
|
||||||
spawn_sync_report(state.clone(), trace_id.to_string(), report_payload);
|
|
||||||
} else {
|
} else {
|
||||||
warn!(
|
warn!(
|
||||||
event_name = "local_video_finalize_missing_error_report_mapping",
|
event_name = "local_video_finalize_missing_error_report_mapping",
|
||||||
@@ -726,49 +752,47 @@ pub(crate) async fn execute_execution_runtime_sync(
|
|||||||
)?));
|
)?));
|
||||||
}
|
}
|
||||||
|
|
||||||
record_sync_terminal_usage(state, &plan, report_context.as_ref(), &base_usage_payload);
|
let usage_payload = build_sync_report_payload(
|
||||||
if let Some(report_kind) = report_kind {
|
|
||||||
let report = GatewaySyncReportRequest {
|
|
||||||
trace_id: trace_id.to_string(),
|
|
||||||
report_kind,
|
|
||||||
report_context,
|
|
||||||
status_code: result.status_code,
|
|
||||||
headers: headers.clone(),
|
|
||||||
body_json: body_json.clone(),
|
|
||||||
client_body_json: None,
|
|
||||||
body_base64: body_base64.clone(),
|
|
||||||
telemetry: result.telemetry.clone(),
|
|
||||||
};
|
|
||||||
spawn_sync_report(state.clone(), trace_id.to_string(), report);
|
|
||||||
}
|
|
||||||
|
|
||||||
let request_id_header: Option<&str> = request_id
|
|
||||||
.map(str::trim)
|
|
||||||
.filter(|value: &&str| !value.is_empty());
|
|
||||||
if let Some(request_id) = request_id_header {
|
|
||||||
headers.insert(
|
|
||||||
CONTROL_REQUEST_ID_HEADER.to_string(),
|
|
||||||
request_id.to_string(),
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
let candidate_id_header: Option<&str> = candidate_id
|
|
||||||
.map(str::trim)
|
|
||||||
.filter(|value: &&str| !value.is_empty());
|
|
||||||
if let Some(candidate_id) = candidate_id_header {
|
|
||||||
headers.insert(
|
|
||||||
CONTROL_CANDIDATE_ID_HEADER.to_string(),
|
|
||||||
candidate_id.to_string(),
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
Ok(Some(build_client_response_from_parts(
|
|
||||||
result.status_code,
|
|
||||||
&headers,
|
|
||||||
Body::from(body_bytes),
|
|
||||||
trace_id,
|
trace_id,
|
||||||
Some(decision),
|
report_kind.unwrap_or_default(),
|
||||||
)?))
|
report_context,
|
||||||
|
status_code,
|
||||||
|
headers,
|
||||||
|
body_json,
|
||||||
|
body_base64,
|
||||||
|
telemetry,
|
||||||
|
);
|
||||||
|
if status_code < 400 {
|
||||||
|
apply_sync_success_effects(
|
||||||
|
state,
|
||||||
|
&plan,
|
||||||
|
usage_payload.report_context.as_ref(),
|
||||||
|
&usage_payload,
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
record_sync_terminal_usage(
|
||||||
|
state,
|
||||||
|
&plan,
|
||||||
|
usage_payload.report_context.as_ref(),
|
||||||
|
&usage_payload,
|
||||||
|
);
|
||||||
|
let response = attach_control_metadata_headers(
|
||||||
|
build_client_response_from_parts(
|
||||||
|
status_code,
|
||||||
|
&usage_payload.headers,
|
||||||
|
Body::from(body_bytes),
|
||||||
|
trace_id,
|
||||||
|
Some(decision),
|
||||||
|
)?,
|
||||||
|
request_id,
|
||||||
|
candidate_id,
|
||||||
|
)?;
|
||||||
|
if !usage_payload.report_kind.trim().is_empty() {
|
||||||
|
spawn_sync_report(state.clone(), usage_payload);
|
||||||
|
}
|
||||||
|
|
||||||
|
Ok(Some(response))
|
||||||
}
|
}
|
||||||
|
|
||||||
#[allow(clippy::too_many_arguments)] // mirrors sync execution context
|
#[allow(clippy::too_many_arguments)] // mirrors sync execution context
|
||||||
@@ -776,12 +800,12 @@ fn maybe_build_implicit_sync_finalize_outcome(
|
|||||||
trace_id: &str,
|
trace_id: &str,
|
||||||
decision: &GatewayControlDecision,
|
decision: &GatewayControlDecision,
|
||||||
plan_kind: &str,
|
plan_kind: &str,
|
||||||
report_context: Option<serde_json::Value>,
|
report_context: &Option<serde_json::Value>,
|
||||||
status_code: u16,
|
status_code: u16,
|
||||||
headers: BTreeMap<String, String>,
|
headers: &BTreeMap<String, String>,
|
||||||
body_json: Option<serde_json::Value>,
|
body_json: &Option<serde_json::Value>,
|
||||||
body_base64: Option<String>,
|
body_base64: &Option<String>,
|
||||||
telemetry: Option<ExecutionTelemetry>,
|
telemetry: &Option<ExecutionTelemetry>,
|
||||||
) -> Result<Option<ImplicitSyncFinalizeOutcome>, GatewayError> {
|
) -> Result<Option<ImplicitSyncFinalizeOutcome>, GatewayError> {
|
||||||
if status_code >= 400 || body_json.is_some() || body_base64.is_none() {
|
if status_code >= 400 || body_json.is_some() || body_base64.is_none() {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
@@ -794,13 +818,13 @@ fn maybe_build_implicit_sync_finalize_outcome(
|
|||||||
let payload = GatewaySyncReportRequest {
|
let payload = GatewaySyncReportRequest {
|
||||||
trace_id: trace_id.to_string(),
|
trace_id: trace_id.to_string(),
|
||||||
report_kind: report_kind.to_string(),
|
report_kind: report_kind.to_string(),
|
||||||
report_context,
|
report_context: report_context.clone(),
|
||||||
status_code,
|
status_code,
|
||||||
headers,
|
headers: headers.clone(),
|
||||||
body_json,
|
body_json: body_json.clone(),
|
||||||
client_body_json: None,
|
client_body_json: None,
|
||||||
body_base64,
|
body_base64: body_base64.clone(),
|
||||||
telemetry,
|
telemetry: telemetry.clone(),
|
||||||
};
|
};
|
||||||
let Some(outcome) = maybe_build_sync_finalize_outcome(trace_id, decision, &payload)? else {
|
let Some(outcome) = maybe_build_sync_finalize_outcome(trace_id, decision, &payload)? else {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
|
|
||||||
use aether_contracts::ExecutionResult;
|
use aether_contracts::ResponseBody;
|
||||||
use base64::Engine as _;
|
use base64::Engine as _;
|
||||||
|
|
||||||
use crate::GatewayError;
|
use crate::GatewayError;
|
||||||
@@ -8,14 +8,14 @@ use crate::GatewayError;
|
|||||||
type DecodedBody = (Vec<u8>, Option<serde_json::Value>, Option<String>);
|
type DecodedBody = (Vec<u8>, Option<serde_json::Value>, Option<String>);
|
||||||
|
|
||||||
pub(super) fn decode_execution_result_body(
|
pub(super) fn decode_execution_result_body(
|
||||||
result: &ExecutionResult,
|
body: Option<ResponseBody>,
|
||||||
headers: &mut BTreeMap<String, String>,
|
headers: &mut BTreeMap<String, String>,
|
||||||
) -> Result<DecodedBody, GatewayError> {
|
) -> Result<DecodedBody, GatewayError> {
|
||||||
let Some(body) = result.body.as_ref() else {
|
let Some(body) = body else {
|
||||||
return Ok((Vec::new(), None, None));
|
return Ok((Vec::new(), None, None));
|
||||||
};
|
};
|
||||||
|
|
||||||
if let Some(json_body) = body.json_body.clone() {
|
if let Some(json_body) = body.json_body {
|
||||||
headers
|
headers
|
||||||
.entry("content-type".to_string())
|
.entry("content-type".to_string())
|
||||||
.or_insert_with(|| "application/json".to_string());
|
.or_insert_with(|| "application/json".to_string());
|
||||||
@@ -25,7 +25,7 @@ pub(super) fn decode_execution_result_body(
|
|||||||
return Ok((bytes, Some(json_body), None));
|
return Ok((bytes, Some(json_body), None));
|
||||||
}
|
}
|
||||||
|
|
||||||
if let Some(body_bytes_b64) = body.body_bytes_b64.clone() {
|
if let Some(body_bytes_b64) = body.body_bytes_b64 {
|
||||||
let bytes = base64::engine::general_purpose::STANDARD
|
let bytes = base64::engine::general_purpose::STANDARD
|
||||||
.decode(&body_bytes_b64)
|
.decode(&body_bytes_b64)
|
||||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||||
|
|||||||
@@ -2,10 +2,13 @@ use std::collections::BTreeMap;
|
|||||||
|
|
||||||
use aether_contracts::ExecutionPlan;
|
use aether_contracts::ExecutionPlan;
|
||||||
use axum::body::Body;
|
use axum::body::Body;
|
||||||
|
use axum::http::header::HeaderValue;
|
||||||
use axum::http::Response;
|
use axum::http::Response;
|
||||||
use serde_json::json;
|
use serde_json::json;
|
||||||
|
|
||||||
use crate::api::response::build_client_response_from_parts;
|
use crate::api::response::{
|
||||||
|
build_client_response_from_parts, build_client_response_from_parts_with_mutator,
|
||||||
|
};
|
||||||
use crate::async_task::VideoTaskService;
|
use crate::async_task::VideoTaskService;
|
||||||
use crate::control::GatewayControlDecision;
|
use crate::control::GatewayControlDecision;
|
||||||
use crate::video_tasks::{
|
use crate::video_tasks::{
|
||||||
@@ -17,9 +20,15 @@ pub(crate) use crate::video_tasks::{
|
|||||||
};
|
};
|
||||||
use crate::{usage::GatewaySyncReportRequest, GatewayError};
|
use crate::{usage::GatewaySyncReportRequest, GatewayError};
|
||||||
|
|
||||||
|
pub(crate) enum LocalVideoSyncSuccessBuild {
|
||||||
|
Handled(LocalVideoSyncSuccessOutcome),
|
||||||
|
NotHandled(GatewaySyncReportRequest),
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) struct LocalVideoSyncSuccessOutcome {
|
pub(crate) struct LocalVideoSyncSuccessOutcome {
|
||||||
pub(crate) response: Response<Body>,
|
pub(crate) response: Response<Body>,
|
||||||
pub(crate) report_payload: GatewaySyncReportRequest,
|
pub(crate) report_payload: GatewaySyncReportRequest,
|
||||||
|
pub(crate) original_report_context: Option<serde_json::Value>,
|
||||||
pub(crate) report_mode: VideoTaskSyncReportMode,
|
pub(crate) report_mode: VideoTaskSyncReportMode,
|
||||||
pub(crate) local_task_snapshot: Option<LocalVideoTaskSnapshot>,
|
pub(crate) local_task_snapshot: Option<LocalVideoTaskSnapshot>,
|
||||||
}
|
}
|
||||||
@@ -29,8 +38,9 @@ fn cloned_report_context_object(
|
|||||||
) -> serde_json::Map<String, serde_json::Value> {
|
) -> serde_json::Map<String, serde_json::Value> {
|
||||||
payload
|
payload
|
||||||
.report_context
|
.report_context
|
||||||
.clone()
|
.as_ref()
|
||||||
.and_then(|value| value.as_object().cloned())
|
.and_then(serde_json::Value::as_object)
|
||||||
|
.cloned()
|
||||||
.unwrap_or_default()
|
.unwrap_or_default()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -56,54 +66,53 @@ fn build_local_video_success_response(
|
|||||||
pub(crate) fn maybe_build_local_video_success_outcome(
|
pub(crate) fn maybe_build_local_video_success_outcome(
|
||||||
trace_id: &str,
|
trace_id: &str,
|
||||||
decision: &GatewayControlDecision,
|
decision: &GatewayControlDecision,
|
||||||
payload: &GatewaySyncReportRequest,
|
mut payload: GatewaySyncReportRequest,
|
||||||
video_tasks: &VideoTaskService,
|
video_tasks: &VideoTaskService,
|
||||||
plan: &ExecutionPlan,
|
plan: &ExecutionPlan,
|
||||||
) -> Result<Option<LocalVideoSyncSuccessOutcome>, GatewayError> {
|
) -> Result<LocalVideoSyncSuccessBuild, GatewayError> {
|
||||||
if payload.status_code >= 400 {
|
if payload.status_code >= 400 {
|
||||||
return Ok(None);
|
return Ok(LocalVideoSyncSuccessBuild::NotHandled(payload));
|
||||||
}
|
}
|
||||||
|
|
||||||
let provider_body = match payload
|
let mut report_context = cloned_report_context_object(&payload);
|
||||||
.body_json
|
let prepared_plan = {
|
||||||
.as_ref()
|
let provider_body = match payload
|
||||||
.and_then(serde_json::Value::as_object)
|
.body_json
|
||||||
{
|
.as_ref()
|
||||||
Some(value) => value,
|
.and_then(serde_json::Value::as_object)
|
||||||
None => return Ok(None),
|
{
|
||||||
|
Some(value) => value,
|
||||||
|
None => return Ok(LocalVideoSyncSuccessBuild::NotHandled(payload)),
|
||||||
|
};
|
||||||
|
video_tasks.prepare_sync_success(
|
||||||
|
payload.report_kind.as_str(),
|
||||||
|
provider_body,
|
||||||
|
&report_context,
|
||||||
|
plan,
|
||||||
|
)
|
||||||
};
|
};
|
||||||
let mut report_context = cloned_report_context_object(payload);
|
let Some(plan) = prepared_plan else {
|
||||||
let Some(plan) = video_tasks.prepare_sync_success(
|
return Ok(LocalVideoSyncSuccessBuild::NotHandled(payload));
|
||||||
payload.report_kind.as_str(),
|
|
||||||
provider_body,
|
|
||||||
&report_context,
|
|
||||||
plan,
|
|
||||||
) else {
|
|
||||||
return Ok(None);
|
|
||||||
};
|
};
|
||||||
plan.apply_to_report_context(&mut report_context);
|
plan.apply_to_report_context(&mut report_context);
|
||||||
let client_body_json = plan.client_body_json();
|
let client_body_json = plan.client_body_json();
|
||||||
|
|
||||||
let response = build_local_video_success_response(trace_id, decision, &client_body_json)?;
|
let response = build_local_video_success_response(trace_id, decision, &client_body_json)?;
|
||||||
let report_payload = GatewaySyncReportRequest {
|
let original_report_context = payload.report_context.take();
|
||||||
trace_id: payload.trace_id.clone(),
|
payload.report_kind = plan.success_report_kind().to_string();
|
||||||
report_kind: plan.success_report_kind().to_string(),
|
payload.report_context = Some(serde_json::Value::Object(report_context));
|
||||||
report_context: Some(serde_json::Value::Object(report_context)),
|
payload.client_body_json = Some(client_body_json);
|
||||||
status_code: payload.status_code,
|
|
||||||
headers: payload.headers.clone(),
|
|
||||||
body_json: payload.body_json.clone(),
|
|
||||||
client_body_json: Some(client_body_json),
|
|
||||||
body_base64: None,
|
|
||||||
telemetry: payload.telemetry.clone(),
|
|
||||||
};
|
|
||||||
|
|
||||||
Ok(Some(LocalVideoSyncSuccessOutcome {
|
Ok(LocalVideoSyncSuccessBuild::Handled(
|
||||||
response,
|
LocalVideoSyncSuccessOutcome {
|
||||||
report_payload,
|
response,
|
||||||
report_mode: plan.report_mode(),
|
report_payload: payload,
|
||||||
local_task_snapshot: matches!(plan.report_mode(), VideoTaskSyncReportMode::Background)
|
original_report_context,
|
||||||
.then(|| plan.to_snapshot()),
|
report_mode: plan.report_mode(),
|
||||||
}))
|
local_task_snapshot: matches!(plan.report_mode(), VideoTaskSyncReportMode::Background)
|
||||||
|
.then(|| plan.to_snapshot()),
|
||||||
|
},
|
||||||
|
))
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn maybe_build_local_sync_finalize_response(
|
pub(crate) fn maybe_build_local_sync_finalize_response(
|
||||||
@@ -147,21 +156,114 @@ pub(crate) fn maybe_build_local_video_error_response(
|
|||||||
return Ok(None);
|
return Ok(None);
|
||||||
}
|
}
|
||||||
|
|
||||||
let response_body = payload.body_json.clone().unwrap_or_else(|| json!({}));
|
let empty_body = json!({});
|
||||||
let body_bytes = serde_json::to_vec(&response_body)
|
let response_body = payload.body_json.as_ref().unwrap_or(&empty_body);
|
||||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
let body_bytes =
|
||||||
|
serde_json::to_vec(response_body).map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||||
|
let body_len = body_bytes.len().to_string();
|
||||||
|
|
||||||
let mut response_headers = payload.headers.clone();
|
Ok(Some(build_client_response_from_parts_with_mutator(
|
||||||
response_headers.remove("content-encoding");
|
|
||||||
response_headers.remove("content-length");
|
|
||||||
response_headers.insert("content-type".to_string(), "application/json".to_string());
|
|
||||||
response_headers.insert("content-length".to_string(), body_bytes.len().to_string());
|
|
||||||
|
|
||||||
Ok(Some(build_client_response_from_parts(
|
|
||||||
payload.status_code,
|
payload.status_code,
|
||||||
&response_headers,
|
&payload.headers,
|
||||||
Body::from(body_bytes),
|
Body::from(body_bytes),
|
||||||
trace_id,
|
trace_id,
|
||||||
Some(decision),
|
Some(decision),
|
||||||
|
|headers| {
|
||||||
|
headers.remove(http::header::CONTENT_ENCODING);
|
||||||
|
headers.remove(http::header::CONTENT_LENGTH);
|
||||||
|
headers.insert(
|
||||||
|
http::header::CONTENT_TYPE,
|
||||||
|
HeaderValue::from_static("application/json"),
|
||||||
|
);
|
||||||
|
headers.insert(
|
||||||
|
http::header::CONTENT_LENGTH,
|
||||||
|
HeaderValue::from_str(body_len.as_str())
|
||||||
|
.map_err(|err| GatewayError::Internal(err.to_string()))?,
|
||||||
|
);
|
||||||
|
Ok(())
|
||||||
|
},
|
||||||
)?))
|
)?))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::*;
|
||||||
|
|
||||||
|
use axum::body::to_bytes;
|
||||||
|
use serde_json::json;
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn local_video_error_response_rewrites_headers_without_mutating_payload() {
|
||||||
|
let decision = GatewayControlDecision::synthetic(
|
||||||
|
"/v1/videos",
|
||||||
|
Some("ai_public".to_string()),
|
||||||
|
Some("openai".to_string()),
|
||||||
|
Some("video".to_string()),
|
||||||
|
Some("openai:video".to_string()),
|
||||||
|
)
|
||||||
|
.with_execution_runtime_candidate(true);
|
||||||
|
let payload = GatewaySyncReportRequest {
|
||||||
|
trace_id: "trace-payload".to_string(),
|
||||||
|
report_kind: "openai_video_create_sync_finalize".to_string(),
|
||||||
|
report_context: Some(json!({
|
||||||
|
"request_id": "req_123",
|
||||||
|
})),
|
||||||
|
status_code: http::StatusCode::BAD_GATEWAY.as_u16(),
|
||||||
|
headers: BTreeMap::from([
|
||||||
|
("content-encoding".to_string(), "gzip".to_string()),
|
||||||
|
("content-length".to_string(), "999".to_string()),
|
||||||
|
("x-upstream-id".to_string(), "video-123".to_string()),
|
||||||
|
]),
|
||||||
|
body_json: Some(json!({
|
||||||
|
"error": {
|
||||||
|
"type": "video_backend_error",
|
||||||
|
"message": "backend failed",
|
||||||
|
}
|
||||||
|
})),
|
||||||
|
client_body_json: None,
|
||||||
|
body_base64: None,
|
||||||
|
telemetry: None,
|
||||||
|
};
|
||||||
|
|
||||||
|
let response =
|
||||||
|
maybe_build_local_video_error_response("trace-response", &decision, &payload)
|
||||||
|
.expect("video error response should build")
|
||||||
|
.expect("video error response should match local video error kinds");
|
||||||
|
|
||||||
|
assert_eq!(response.status(), http::StatusCode::BAD_GATEWAY);
|
||||||
|
assert_eq!(
|
||||||
|
response
|
||||||
|
.headers()
|
||||||
|
.get(http::header::CONTENT_TYPE)
|
||||||
|
.and_then(|value| value.to_str().ok()),
|
||||||
|
Some("application/json")
|
||||||
|
);
|
||||||
|
assert_eq!(response.headers().get(http::header::CONTENT_ENCODING), None);
|
||||||
|
assert_eq!(
|
||||||
|
response
|
||||||
|
.headers()
|
||||||
|
.get("x-upstream-id")
|
||||||
|
.and_then(|value| value.to_str().ok()),
|
||||||
|
Some("video-123")
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
payload.headers.get("content-encoding").map(String::as_str),
|
||||||
|
Some("gzip")
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
payload.headers.get("content-length").map(String::as_str),
|
||||||
|
Some("999")
|
||||||
|
);
|
||||||
|
|
||||||
|
let body = to_bytes(response.into_body(), usize::MAX)
|
||||||
|
.await
|
||||||
|
.expect("response body should read");
|
||||||
|
assert_eq!(
|
||||||
|
serde_json::from_slice::<serde_json::Value>(&body).expect("response body should parse"),
|
||||||
|
payload
|
||||||
|
.body_json
|
||||||
|
.clone()
|
||||||
|
.expect("payload body should exist")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -6,5 +6,6 @@ pub(crate) use execution::execute_execution_runtime_sync;
|
|||||||
pub(crate) use execution::{
|
pub(crate) use execution::{
|
||||||
maybe_build_local_sync_finalize_response, maybe_build_local_video_error_response,
|
maybe_build_local_sync_finalize_response, maybe_build_local_video_error_response,
|
||||||
maybe_build_local_video_success_outcome, resolve_local_sync_error_background_report_kind,
|
maybe_build_local_video_success_outcome, resolve_local_sync_error_background_report_kind,
|
||||||
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessOutcome,
|
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessBuild,
|
||||||
|
LocalVideoSyncSuccessOutcome,
|
||||||
};
|
};
|
||||||
|
|||||||
@@ -161,12 +161,12 @@ impl DirectSyncExecutionRuntime {
|
|||||||
|
|
||||||
pub(crate) async fn execute_sync(
|
pub(crate) async fn execute_sync(
|
||||||
&self,
|
&self,
|
||||||
plan: ExecutionPlan,
|
plan: &ExecutionPlan,
|
||||||
) -> Result<ExecutionResult, ExecutionRuntimeTransportError> {
|
) -> Result<ExecutionResult, ExecutionRuntimeTransportError> {
|
||||||
let body_bytes = build_request_body(&plan)?;
|
let body_bytes = build_request_body(plan)?;
|
||||||
|
|
||||||
let started_at = Instant::now();
|
let started_at = Instant::now();
|
||||||
let response = send_request(&plan, body_bytes).await?;
|
let response = send_request(plan, body_bytes).await?;
|
||||||
let ttfb_ms = started_at.elapsed().as_millis() as u64;
|
let ttfb_ms = started_at.elapsed().as_millis() as u64;
|
||||||
let status_code = response.status().as_u16();
|
let status_code = response.status().as_u16();
|
||||||
let headers = collect_response_headers(response.headers());
|
let headers = collect_response_headers(response.headers());
|
||||||
@@ -200,8 +200,8 @@ impl DirectSyncExecutionRuntime {
|
|||||||
};
|
};
|
||||||
|
|
||||||
Ok(ExecutionResult {
|
Ok(ExecutionResult {
|
||||||
request_id: plan.request_id,
|
request_id: plan.request_id.clone(),
|
||||||
candidate_id: plan.candidate_id,
|
candidate_id: plan.candidate_id.clone(),
|
||||||
status_code,
|
status_code,
|
||||||
headers,
|
headers,
|
||||||
body,
|
body,
|
||||||
@@ -216,24 +216,24 @@ impl DirectSyncExecutionRuntime {
|
|||||||
|
|
||||||
pub(crate) async fn execute_stream(
|
pub(crate) async fn execute_stream(
|
||||||
&self,
|
&self,
|
||||||
plan: ExecutionPlan,
|
plan: &ExecutionPlan,
|
||||||
) -> Result<DirectUpstreamStreamExecution, ExecutionRuntimeTransportError> {
|
) -> Result<DirectUpstreamStreamExecution, ExecutionRuntimeTransportError> {
|
||||||
if !plan.stream {
|
if !plan.stream {
|
||||||
return Err(ExecutionRuntimeTransportError::StreamUnsupported);
|
return Err(ExecutionRuntimeTransportError::StreamUnsupported);
|
||||||
}
|
}
|
||||||
|
|
||||||
let body_bytes = build_request_body(&plan)?;
|
let body_bytes = build_request_body(plan)?;
|
||||||
|
|
||||||
let started_at = Instant::now();
|
let started_at = Instant::now();
|
||||||
let response = send_request(&plan, body_bytes).await?;
|
let response = send_request(plan, body_bytes).await?;
|
||||||
let status_code = response.status().as_u16();
|
let status_code = response.status().as_u16();
|
||||||
let headers = collect_response_headers(response.headers());
|
let headers = collect_response_headers(response.headers());
|
||||||
|
|
||||||
let stream_summary_report_context = build_stream_summary_report_context(&plan);
|
let stream_summary_report_context = build_stream_summary_report_context(plan);
|
||||||
|
|
||||||
Ok(DirectUpstreamStreamExecution {
|
Ok(DirectUpstreamStreamExecution {
|
||||||
request_id: plan.request_id,
|
request_id: plan.request_id.clone(),
|
||||||
candidate_id: plan.candidate_id,
|
candidate_id: plan.candidate_id.clone(),
|
||||||
status_code,
|
status_code,
|
||||||
headers,
|
headers,
|
||||||
provider_api_format: plan.provider_api_format.clone(),
|
provider_api_format: plan.provider_api_format.clone(),
|
||||||
@@ -274,7 +274,7 @@ pub(crate) async fn execute_sync_plan(
|
|||||||
let _ = state;
|
let _ = state;
|
||||||
let _ = trace_id;
|
let _ = trace_id;
|
||||||
DirectSyncExecutionRuntime::new()
|
DirectSyncExecutionRuntime::new()
|
||||||
.execute_sync(plan.clone())
|
.execute_sync(plan)
|
||||||
.await
|
.await
|
||||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||||
}
|
}
|
||||||
@@ -1101,7 +1101,7 @@ mod tests {
|
|||||||
|
|
||||||
let execution_runtime = DirectSyncExecutionRuntime::new();
|
let execution_runtime = DirectSyncExecutionRuntime::new();
|
||||||
let result = execution_runtime
|
let result = execution_runtime
|
||||||
.execute_sync(ExecutionPlan {
|
.execute_sync(&ExecutionPlan {
|
||||||
request_id: "req-1".into(),
|
request_id: "req-1".into(),
|
||||||
candidate_id: Some("cand-1".into()),
|
candidate_id: Some("cand-1".into()),
|
||||||
provider_name: Some("openai".into()),
|
provider_name: Some("openai".into()),
|
||||||
@@ -1168,7 +1168,7 @@ mod tests {
|
|||||||
|
|
||||||
let execution_runtime = DirectSyncExecutionRuntime::new();
|
let execution_runtime = DirectSyncExecutionRuntime::new();
|
||||||
let result = execution_runtime
|
let result = execution_runtime
|
||||||
.execute_sync(ExecutionPlan {
|
.execute_sync(&ExecutionPlan {
|
||||||
request_id: "req-1".into(),
|
request_id: "req-1".into(),
|
||||||
candidate_id: None,
|
candidate_id: None,
|
||||||
provider_name: None,
|
provider_name: None,
|
||||||
@@ -1369,7 +1369,7 @@ mod tests {
|
|||||||
|
|
||||||
let execution_runtime = DirectSyncExecutionRuntime::new();
|
let execution_runtime = DirectSyncExecutionRuntime::new();
|
||||||
let result = execution_runtime
|
let result = execution_runtime
|
||||||
.execute_sync(ExecutionPlan {
|
.execute_sync(&ExecutionPlan {
|
||||||
request_id: "req-redirect-1".into(),
|
request_id: "req-redirect-1".into(),
|
||||||
candidate_id: None,
|
candidate_id: None,
|
||||||
provider_name: Some("provider_ops".into()),
|
provider_name: Some("provider_ops".into()),
|
||||||
@@ -1442,7 +1442,7 @@ mod tests {
|
|||||||
|
|
||||||
let execution_runtime = DirectSyncExecutionRuntime::new();
|
let execution_runtime = DirectSyncExecutionRuntime::new();
|
||||||
let result = execution_runtime
|
let result = execution_runtime
|
||||||
.execute_sync(ExecutionPlan {
|
.execute_sync(&ExecutionPlan {
|
||||||
request_id: "req-redirect-2".into(),
|
request_id: "req-redirect-2".into(),
|
||||||
candidate_id: None,
|
candidate_id: None,
|
||||||
provider_name: Some("provider_oauth".into()),
|
provider_name: Some("provider_oauth".into()),
|
||||||
@@ -1512,7 +1512,7 @@ mod tests {
|
|||||||
|
|
||||||
let execution_runtime = DirectSyncExecutionRuntime::new();
|
let execution_runtime = DirectSyncExecutionRuntime::new();
|
||||||
let result = execution_runtime
|
let result = execution_runtime
|
||||||
.execute_sync(ExecutionPlan {
|
.execute_sync(&ExecutionPlan {
|
||||||
request_id: "req-relay-http1-1".into(),
|
request_id: "req-relay-http1-1".into(),
|
||||||
candidate_id: None,
|
candidate_id: None,
|
||||||
provider_name: Some("provider_ops".into()),
|
provider_name: Some("provider_ops".into()),
|
||||||
@@ -1579,7 +1579,7 @@ mod tests {
|
|||||||
|
|
||||||
let execution_runtime = DirectSyncExecutionRuntime::new();
|
let execution_runtime = DirectSyncExecutionRuntime::new();
|
||||||
let result = execution_runtime
|
let result = execution_runtime
|
||||||
.execute_sync(ExecutionPlan {
|
.execute_sync(&ExecutionPlan {
|
||||||
request_id: "req-tls-1".into(),
|
request_id: "req-tls-1".into(),
|
||||||
candidate_id: Some("cand-1".into()),
|
candidate_id: Some("cand-1".into()),
|
||||||
provider_name: Some("claude".into()),
|
provider_name: Some("claude".into()),
|
||||||
@@ -1654,7 +1654,7 @@ mod tests {
|
|||||||
|
|
||||||
let execution_runtime = DirectSyncExecutionRuntime::new();
|
let execution_runtime = DirectSyncExecutionRuntime::new();
|
||||||
let result = execution_runtime
|
let result = execution_runtime
|
||||||
.execute_sync(ExecutionPlan {
|
.execute_sync(&ExecutionPlan {
|
||||||
request_id: "req-gzip-1".into(),
|
request_id: "req-gzip-1".into(),
|
||||||
candidate_id: Some("cand-1".into()),
|
candidate_id: Some("cand-1".into()),
|
||||||
provider_name: Some("openai".into()),
|
provider_name: Some("openai".into()),
|
||||||
@@ -1715,7 +1715,7 @@ mod tests {
|
|||||||
|
|
||||||
let execution_runtime = DirectSyncExecutionRuntime::new();
|
let execution_runtime = DirectSyncExecutionRuntime::new();
|
||||||
let result = execution_runtime
|
let result = execution_runtime
|
||||||
.execute_sync(ExecutionPlan {
|
.execute_sync(&ExecutionPlan {
|
||||||
request_id: "req-ttfb-1".into(),
|
request_id: "req-ttfb-1".into(),
|
||||||
candidate_id: Some("cand-1".into()),
|
candidate_id: Some("cand-1".into()),
|
||||||
provider_name: Some("openai".into()),
|
provider_name: Some("openai".into()),
|
||||||
|
|||||||
@@ -206,18 +206,18 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
Json(build_internal_gateway_fallback_plan_payload(None)).into_response(),
|
Json(build_internal_gateway_fallback_plan_payload(None)).into_response(),
|
||||||
));
|
));
|
||||||
};
|
};
|
||||||
if let Some(auth_context) = payload.auth_context.clone() {
|
let provided_auth_context = payload.auth_context.is_some();
|
||||||
|
if let Some(auth_context) = payload.auth_context {
|
||||||
resolved.auth_context = Some(auth_context);
|
resolved.auth_context = Some(auth_context);
|
||||||
resolved.local_auth_rejection = None;
|
resolved.local_auth_rejection = None;
|
||||||
}
|
}
|
||||||
let auth_context = resolved.auth_context.clone();
|
let auth_context = resolved.auth_context.as_ref();
|
||||||
if auth_context
|
if auth_context
|
||||||
.as_ref()
|
|
||||||
.map(|value| !value.access_allowed)
|
.map(|value| !value.access_allowed)
|
||||||
.unwrap_or(true)
|
.unwrap_or(true)
|
||||||
{
|
{
|
||||||
let fallback_auth_context = if payload.auth_context.is_none() {
|
let fallback_auth_context = if !provided_auth_context {
|
||||||
auth_context.as_ref()
|
auth_context
|
||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
};
|
};
|
||||||
@@ -239,8 +239,8 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
)
|
)
|
||||||
.await?
|
.await?
|
||||||
else {
|
else {
|
||||||
let fallback_auth_context = if payload.auth_context.is_none() {
|
let fallback_auth_context = if !provided_auth_context {
|
||||||
auth_context.as_ref()
|
auth_context
|
||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
};
|
};
|
||||||
@@ -251,7 +251,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
.into_response(),
|
.into_response(),
|
||||||
));
|
));
|
||||||
};
|
};
|
||||||
if payload.auth_context.is_some() {
|
if provided_auth_context {
|
||||||
local_payload.auth_context = None;
|
local_payload.auth_context = None;
|
||||||
}
|
}
|
||||||
return Ok(Some(Json(local_payload).into_response()));
|
return Ok(Some(Json(local_payload).into_response()));
|
||||||
@@ -308,18 +308,18 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
Json(build_internal_gateway_fallback_plan_payload(None)).into_response(),
|
Json(build_internal_gateway_fallback_plan_payload(None)).into_response(),
|
||||||
));
|
));
|
||||||
};
|
};
|
||||||
if let Some(auth_context) = payload.auth_context.clone() {
|
let provided_auth_context = payload.auth_context.is_some();
|
||||||
|
if let Some(auth_context) = payload.auth_context {
|
||||||
resolved.auth_context = Some(auth_context);
|
resolved.auth_context = Some(auth_context);
|
||||||
resolved.local_auth_rejection = None;
|
resolved.local_auth_rejection = None;
|
||||||
}
|
}
|
||||||
let auth_context = resolved.auth_context.clone();
|
let auth_context = resolved.auth_context.as_ref();
|
||||||
if auth_context
|
if auth_context
|
||||||
.as_ref()
|
|
||||||
.map(|value| !value.access_allowed)
|
.map(|value| !value.access_allowed)
|
||||||
.unwrap_or(true)
|
.unwrap_or(true)
|
||||||
{
|
{
|
||||||
let fallback_auth_context = if payload.auth_context.is_none() {
|
let fallback_auth_context = if !provided_auth_context {
|
||||||
auth_context.as_ref()
|
auth_context
|
||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
};
|
};
|
||||||
@@ -339,8 +339,8 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
)
|
)
|
||||||
.await?
|
.await?
|
||||||
else {
|
else {
|
||||||
let fallback_auth_context = if payload.auth_context.is_none() {
|
let fallback_auth_context = if !provided_auth_context {
|
||||||
auth_context.as_ref()
|
auth_context
|
||||||
} else {
|
} else {
|
||||||
None
|
None
|
||||||
};
|
};
|
||||||
@@ -351,7 +351,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
.into_response(),
|
.into_response(),
|
||||||
));
|
));
|
||||||
};
|
};
|
||||||
if payload.auth_context.is_some() {
|
if provided_auth_context {
|
||||||
local_payload.auth_context = None;
|
local_payload.auth_context = None;
|
||||||
}
|
}
|
||||||
return Ok(Some(Json(local_payload).into_response()));
|
return Ok(Some(Json(local_payload).into_response()));
|
||||||
@@ -406,7 +406,8 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
else {
|
else {
|
||||||
return Ok(Some(build_internal_gateway_proxy_public_response()));
|
return Ok(Some(build_internal_gateway_proxy_public_response()));
|
||||||
};
|
};
|
||||||
if let Some(auth_context) = payload.auth_context.clone() {
|
let provided_auth_context = payload.auth_context.is_some();
|
||||||
|
if let Some(auth_context) = payload.auth_context {
|
||||||
resolved.auth_context = Some(auth_context);
|
resolved.auth_context = Some(auth_context);
|
||||||
resolved.local_auth_rejection = None;
|
resolved.local_auth_rejection = None;
|
||||||
}
|
}
|
||||||
@@ -421,7 +422,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
)
|
)
|
||||||
.await?
|
.await?
|
||||||
{
|
{
|
||||||
if payload.auth_context.is_some() {
|
if provided_auth_context {
|
||||||
planned.auth_context = None;
|
planned.auth_context = None;
|
||||||
}
|
}
|
||||||
return Ok(Some(Json(planned).into_response()));
|
return Ok(Some(Json(planned).into_response()));
|
||||||
@@ -472,7 +473,8 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
else {
|
else {
|
||||||
return Ok(Some(build_internal_gateway_proxy_public_response()));
|
return Ok(Some(build_internal_gateway_proxy_public_response()));
|
||||||
};
|
};
|
||||||
if let Some(auth_context) = payload.auth_context.clone() {
|
let provided_auth_context = payload.auth_context.is_some();
|
||||||
|
if let Some(auth_context) = payload.auth_context {
|
||||||
resolved.auth_context = Some(auth_context);
|
resolved.auth_context = Some(auth_context);
|
||||||
resolved.local_auth_rejection = None;
|
resolved.local_auth_rejection = None;
|
||||||
}
|
}
|
||||||
@@ -485,7 +487,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
)
|
)
|
||||||
.await?
|
.await?
|
||||||
{
|
{
|
||||||
if payload.auth_context.is_some() {
|
if provided_auth_context {
|
||||||
planned.auth_context = None;
|
planned.auth_context = None;
|
||||||
}
|
}
|
||||||
return Ok(Some(Json(planned).into_response()));
|
return Ok(Some(Json(planned).into_response()));
|
||||||
@@ -542,7 +544,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
else {
|
else {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
if let Some(auth_context) = payload.auth_context.clone() {
|
if let Some(auth_context) = payload.auth_context {
|
||||||
resolved.auth_context = Some(auth_context);
|
resolved.auth_context = Some(auth_context);
|
||||||
resolved.local_auth_rejection = None;
|
resolved.local_auth_rejection = None;
|
||||||
}
|
}
|
||||||
@@ -626,7 +628,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
else {
|
else {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
if let Some(auth_context) = payload.auth_context.clone() {
|
if let Some(auth_context) = payload.auth_context {
|
||||||
resolved.auth_context = Some(auth_context);
|
resolved.auth_context = Some(auth_context);
|
||||||
resolved.local_auth_rejection = None;
|
resolved.local_auth_rejection = None;
|
||||||
}
|
}
|
||||||
@@ -683,8 +685,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
)));
|
)));
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
let trace_id = payload.trace_id.clone();
|
crate::usage::submit_sync_report(state, payload).await?;
|
||||||
crate::usage::submit_sync_report(state, &trace_id, payload).await?;
|
|
||||||
return Ok(Some(Json(json!({ "ok": true })).into_response()));
|
return Ok(Some(Json(json!({ "ok": true })).into_response()));
|
||||||
}
|
}
|
||||||
Some("report_stream")
|
Some("report_stream")
|
||||||
@@ -707,8 +708,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
)));
|
)));
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
let trace_id = payload.trace_id.clone();
|
crate::usage::submit_stream_report(state, payload).await?;
|
||||||
crate::usage::submit_stream_report(state, &trace_id, payload).await?;
|
|
||||||
return Ok(Some(Json(json!({ "ok": true })).into_response()));
|
return Ok(Some(Json(json!({ "ok": true })).into_response()));
|
||||||
}
|
}
|
||||||
Some("finalize_sync")
|
Some("finalize_sync")
|
||||||
@@ -739,16 +739,12 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
};
|
};
|
||||||
let trace_id = payload.trace_id.clone();
|
let trace_id = payload.trace_id.clone();
|
||||||
if let Some(outcome) = ai_pipeline_api::maybe_build_sync_finalize_outcome(
|
if let Some(outcome) = ai_pipeline_api::maybe_build_sync_finalize_outcome(
|
||||||
&trace_id,
|
trace_id.as_str(),
|
||||||
&synthetic_decision,
|
&synthetic_decision,
|
||||||
&payload,
|
&payload,
|
||||||
)? {
|
)? {
|
||||||
if let Some(background_report) = outcome.background_report {
|
if let Some(background_report) = outcome.background_report {
|
||||||
crate::usage::spawn_sync_report(
|
crate::usage::spawn_sync_report(state.clone(), background_report);
|
||||||
state.clone(),
|
|
||||||
trace_id.clone(),
|
|
||||||
background_report,
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
let mut response = outcome.response;
|
let mut response = outcome.response;
|
||||||
response.headers_mut().insert(
|
response.headers_mut().insert(
|
||||||
@@ -761,7 +757,7 @@ pub(crate) async fn maybe_build_local_internal_proxy_response_impl(
|
|||||||
state,
|
state,
|
||||||
trace_id.as_str(),
|
trace_id.as_str(),
|
||||||
&synthetic_decision,
|
&synthetic_decision,
|
||||||
&payload,
|
payload,
|
||||||
)
|
)
|
||||||
.await?
|
.await?
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ use crate::control::GatewayControlDecision;
|
|||||||
use crate::execution_runtime::{
|
use crate::execution_runtime::{
|
||||||
maybe_build_local_sync_finalize_response, maybe_build_local_video_error_response,
|
maybe_build_local_sync_finalize_response, maybe_build_local_video_error_response,
|
||||||
maybe_build_local_video_success_outcome, resolve_local_sync_error_background_report_kind,
|
maybe_build_local_video_success_outcome, resolve_local_sync_error_background_report_kind,
|
||||||
resolve_local_sync_success_background_report_kind,
|
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessBuild,
|
||||||
};
|
};
|
||||||
use crate::handlers::shared::{
|
use crate::handlers::shared::{
|
||||||
unix_secs_to_rfc3339, InternalTunnelHeartbeatRequest, InternalTunnelNodeStatusRequest,
|
unix_secs_to_rfc3339, InternalTunnelHeartbeatRequest, InternalTunnelNodeStatusRequest,
|
||||||
@@ -244,59 +244,66 @@ pub(crate) async fn maybe_build_internal_finalize_video_response(
|
|||||||
state: &AppState,
|
state: &AppState,
|
||||||
trace_id: &str,
|
trace_id: &str,
|
||||||
decision: &GatewayControlDecision,
|
decision: &GatewayControlDecision,
|
||||||
payload: &crate::usage::GatewaySyncReportRequest,
|
payload: crate::usage::GatewaySyncReportRequest,
|
||||||
) -> Result<Option<Response<Body>>, GatewayError> {
|
) -> Result<Option<Response<Body>>, GatewayError> {
|
||||||
let Some(plan) = infer_internal_finalize_signature(payload).and_then(|signature| {
|
let Some((signature, plan)) =
|
||||||
build_internal_finalize_video_plan(
|
infer_internal_finalize_signature(&payload).and_then(|signature| {
|
||||||
payload.trace_id.as_str(),
|
build_internal_finalize_video_plan(
|
||||||
signature.as_str(),
|
payload.trace_id.as_str(),
|
||||||
payload.report_context.as_ref(),
|
signature.as_str(),
|
||||||
)
|
payload.report_context.as_ref(),
|
||||||
}) else {
|
)
|
||||||
|
.map(|plan| (signature, plan))
|
||||||
|
})
|
||||||
|
else {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
};
|
};
|
||||||
|
|
||||||
if let Some(outcome) = maybe_build_local_video_success_outcome(
|
let mut payload = match maybe_build_local_video_success_outcome(
|
||||||
trace_id,
|
trace_id,
|
||||||
decision,
|
decision,
|
||||||
payload,
|
payload,
|
||||||
&state.video_tasks,
|
&state.video_tasks,
|
||||||
&plan,
|
&plan,
|
||||||
)? {
|
)? {
|
||||||
if let Some(snapshot) = outcome.local_task_snapshot.clone() {
|
LocalVideoSyncSuccessBuild::Handled(outcome) => {
|
||||||
state.video_tasks.record_snapshot(snapshot.clone());
|
let crate::execution_runtime::LocalVideoSyncSuccessOutcome {
|
||||||
let _ = state.upsert_video_task_snapshot(&snapshot).await?;
|
response,
|
||||||
}
|
report_payload,
|
||||||
match outcome.report_mode {
|
original_report_context: _,
|
||||||
crate::video_tasks::VideoTaskSyncReportMode::InlineSync => {
|
report_mode,
|
||||||
crate::usage::submit_sync_report(state, trace_id, outcome.report_payload).await?;
|
local_task_snapshot,
|
||||||
|
} = outcome;
|
||||||
|
if let Some(snapshot) = local_task_snapshot {
|
||||||
|
let _ = state.upsert_video_task_snapshot(&snapshot).await?;
|
||||||
|
state.video_tasks.record_snapshot(snapshot);
|
||||||
}
|
}
|
||||||
crate::video_tasks::VideoTaskSyncReportMode::Background => {
|
match report_mode {
|
||||||
crate::usage::spawn_sync_report(
|
crate::video_tasks::VideoTaskSyncReportMode::InlineSync => {
|
||||||
state.clone(),
|
crate::usage::submit_sync_report(state, report_payload).await?;
|
||||||
trace_id.to_string(),
|
}
|
||||||
outcome.report_payload,
|
crate::video_tasks::VideoTaskSyncReportMode::Background => {
|
||||||
);
|
crate::usage::spawn_sync_report(state.clone(), report_payload);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
let mut response = response;
|
||||||
|
response.headers_mut().insert(
|
||||||
|
HeaderName::from_static(CONTROL_EXECUTED_HEADER),
|
||||||
|
HeaderValue::from_static("true"),
|
||||||
|
);
|
||||||
|
return Ok(Some(response));
|
||||||
}
|
}
|
||||||
let mut response = outcome.response;
|
LocalVideoSyncSuccessBuild::NotHandled(payload) => payload,
|
||||||
response.headers_mut().insert(
|
};
|
||||||
HeaderName::from_static(CONTROL_EXECUTED_HEADER),
|
|
||||||
HeaderValue::from_static("true"),
|
|
||||||
);
|
|
||||||
return Ok(Some(response));
|
|
||||||
}
|
|
||||||
|
|
||||||
if let Some(mut response) =
|
if let Some(mut response) =
|
||||||
maybe_build_local_sync_finalize_response(trace_id, decision, payload)?
|
maybe_build_local_sync_finalize_response(trace_id, decision, &payload)?
|
||||||
{
|
{
|
||||||
let request_path = infer_internal_finalize_signature(payload).and_then(|signature| {
|
let request_path = build_local_sync_finalize_request_path(
|
||||||
build_local_sync_finalize_request_path(
|
payload.report_kind.as_str(),
|
||||||
payload.report_kind.as_str(),
|
signature.as_str(),
|
||||||
signature.as_str(),
|
payload.report_context.as_ref(),
|
||||||
payload.report_context.as_ref(),
|
);
|
||||||
)
|
|
||||||
});
|
|
||||||
if let Some(request_path) = request_path {
|
if let Some(request_path) = request_path {
|
||||||
state
|
state
|
||||||
.video_tasks
|
.video_tasks
|
||||||
@@ -311,9 +318,8 @@ pub(crate) async fn maybe_build_internal_finalize_video_response(
|
|||||||
if let Some(success_report_kind) =
|
if let Some(success_report_kind) =
|
||||||
resolve_local_sync_success_background_report_kind(payload.report_kind.as_str())
|
resolve_local_sync_success_background_report_kind(payload.report_kind.as_str())
|
||||||
{
|
{
|
||||||
let mut report_payload = payload.clone();
|
payload.report_kind = success_report_kind.to_string();
|
||||||
report_payload.report_kind = success_report_kind.to_string();
|
crate::usage::spawn_sync_report(state.clone(), payload);
|
||||||
crate::usage::spawn_sync_report(state.clone(), trace_id.to_string(), report_payload);
|
|
||||||
}
|
}
|
||||||
response.headers_mut().insert(
|
response.headers_mut().insert(
|
||||||
HeaderName::from_static(CONTROL_EXECUTED_HEADER),
|
HeaderName::from_static(CONTROL_EXECUTED_HEADER),
|
||||||
@@ -322,14 +328,14 @@ pub(crate) async fn maybe_build_internal_finalize_video_response(
|
|||||||
return Ok(Some(response));
|
return Ok(Some(response));
|
||||||
}
|
}
|
||||||
|
|
||||||
if let Some(mut response) = maybe_build_local_video_error_response(trace_id, decision, payload)?
|
if let Some(mut response) =
|
||||||
|
maybe_build_local_video_error_response(trace_id, decision, &payload)?
|
||||||
{
|
{
|
||||||
if let Some(error_report_kind) =
|
if let Some(error_report_kind) =
|
||||||
resolve_local_sync_error_background_report_kind(payload.report_kind.as_str())
|
resolve_local_sync_error_background_report_kind(payload.report_kind.as_str())
|
||||||
{
|
{
|
||||||
let mut report_payload = payload.clone();
|
payload.report_kind = error_report_kind.to_string();
|
||||||
report_payload.report_kind = error_report_kind.to_string();
|
crate::usage::spawn_sync_report(state.clone(), payload);
|
||||||
crate::usage::spawn_sync_report(state.clone(), trace_id.to_string(), report_payload);
|
|
||||||
}
|
}
|
||||||
response.headers_mut().insert(
|
response.headers_mut().insert(
|
||||||
HeaderName::from_static(CONTROL_EXECUTED_HEADER),
|
HeaderName::from_static(CONTROL_EXECUTED_HEADER),
|
||||||
|
|||||||
@@ -122,6 +122,10 @@ pub(crate) async fn build_public_providers_payload(
|
|||||||
.iter()
|
.iter()
|
||||||
.map(|provider| provider.id.clone())
|
.map(|provider| provider.id.clone())
|
||||||
.collect::<Vec<_>>();
|
.collect::<Vec<_>>();
|
||||||
|
let provider_ids_set = provider_ids
|
||||||
|
.iter()
|
||||||
|
.map(String::as_str)
|
||||||
|
.collect::<BTreeSet<_>>();
|
||||||
let endpoints = if provider_ids.is_empty() {
|
let endpoints = if provider_ids.is_empty() {
|
||||||
Vec::new()
|
Vec::new()
|
||||||
} else {
|
} else {
|
||||||
@@ -156,7 +160,7 @@ pub(crate) async fn build_public_providers_payload(
|
|||||||
.ok()
|
.ok()
|
||||||
.unwrap_or_default();
|
.unwrap_or_default();
|
||||||
for row in rows {
|
for row in rows {
|
||||||
if provider_ids.contains(&row.provider_id) {
|
if provider_ids_set.contains(row.provider_id.as_str()) {
|
||||||
models_by_provider
|
models_by_provider
|
||||||
.entry(row.provider_id.clone())
|
.entry(row.provider_id.clone())
|
||||||
.or_default()
|
.or_default()
|
||||||
@@ -332,7 +336,6 @@ pub(crate) async fn build_api_format_health_monitor_payload(
|
|||||||
let mut endpoint_ids_by_format = BTreeMap::<String, Vec<String>>::new();
|
let mut endpoint_ids_by_format = BTreeMap::<String, Vec<String>>::new();
|
||||||
let mut endpoint_to_format = BTreeMap::<String, String>::new();
|
let mut endpoint_to_format = BTreeMap::<String, String>::new();
|
||||||
let mut provider_ids_by_format = BTreeMap::<String, BTreeSet<String>>::new();
|
let mut provider_ids_by_format = BTreeMap::<String, BTreeSet<String>>::new();
|
||||||
let mut active_provider_formats = BTreeSet::<(String, String)>::new();
|
|
||||||
for endpoint in active_endpoints {
|
for endpoint in active_endpoints {
|
||||||
endpoint_to_format.insert(endpoint.id.clone(), endpoint.api_format.clone());
|
endpoint_to_format.insert(endpoint.id.clone(), endpoint.api_format.clone());
|
||||||
endpoint_ids_by_format
|
endpoint_ids_by_format
|
||||||
@@ -343,7 +346,6 @@ pub(crate) async fn build_api_format_health_monitor_payload(
|
|||||||
.entry(endpoint.api_format.clone())
|
.entry(endpoint.api_format.clone())
|
||||||
.or_default()
|
.or_default()
|
||||||
.insert(endpoint.provider_id.clone());
|
.insert(endpoint.provider_id.clone());
|
||||||
active_provider_formats.insert((endpoint.provider_id, endpoint.api_format));
|
|
||||||
}
|
}
|
||||||
let all_endpoint_ids = endpoint_to_format.keys().cloned().collect::<Vec<_>>();
|
let all_endpoint_ids = endpoint_to_format.keys().cloned().collect::<Vec<_>>();
|
||||||
|
|
||||||
@@ -356,7 +358,9 @@ pub(crate) async fn build_api_format_health_monitor_payload(
|
|||||||
.unwrap_or_default();
|
.unwrap_or_default();
|
||||||
for key in keys.into_iter().filter(|key| key.is_active) {
|
for key in keys.into_iter().filter(|key| key.is_active) {
|
||||||
for api_format in provider_key_api_formats(&key) {
|
for api_format in provider_key_api_formats(&key) {
|
||||||
if active_provider_formats.contains(&(key.provider_id.clone(), api_format.clone()))
|
if provider_ids_by_format
|
||||||
|
.get(&api_format)
|
||||||
|
.is_some_and(|provider_ids| provider_ids.contains(key.provider_id.as_str()))
|
||||||
{
|
{
|
||||||
*key_counts_by_format.entry(api_format).or_default() += 1;
|
*key_counts_by_format.entry(api_format).or_default() += 1;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ use aether_crypto::DEVELOPMENT_ENCRYPTION_KEY;
|
|||||||
use aether_crypto::{decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext};
|
use aether_crypto::{decrypt_python_fernet_ciphertext, encrypt_python_fernet_plaintext};
|
||||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
||||||
use serde_json::{json, Map, Value};
|
use serde_json::{json, Map, Value};
|
||||||
|
use std::borrow::Cow;
|
||||||
use std::time::{SystemTime, UNIX_EPOCH};
|
use std::time::{SystemTime, UNIX_EPOCH};
|
||||||
|
|
||||||
const OAUTH_ACCOUNT_BLOCK_PREFIX: &str = "[ACCOUNT_BLOCK] ";
|
const OAUTH_ACCOUNT_BLOCK_PREFIX: &str = "[ACCOUNT_BLOCK] ";
|
||||||
@@ -63,23 +64,27 @@ pub(crate) fn decrypt_catalog_secret_with_fallbacks(
|
|||||||
None
|
None
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn effective_catalog_encryption_key(state: &AppState) -> Option<String> {
|
pub(crate) fn effective_catalog_encryption_key(state: &AppState) -> Option<Cow<'_, str>> {
|
||||||
let encryption_key = state.encryption_key().map(str::trim).unwrap_or("");
|
let encryption_key = state.encryption_key().map(str::trim).unwrap_or("");
|
||||||
if !encryption_key.is_empty() {
|
if !encryption_key.is_empty() {
|
||||||
return Some(encryption_key.to_string());
|
return Some(Cow::Borrowed(encryption_key));
|
||||||
}
|
}
|
||||||
for env_key in ["AETHER_GATEWAY_DATA_ENCRYPTION_KEY", "ENCRYPTION_KEY"] {
|
for env_key in ["AETHER_GATEWAY_DATA_ENCRYPTION_KEY", "ENCRYPTION_KEY"] {
|
||||||
let Ok(candidate) = std::env::var(env_key) else {
|
let Ok(candidate) = std::env::var(env_key) else {
|
||||||
continue;
|
continue;
|
||||||
};
|
};
|
||||||
let candidate = candidate.trim();
|
let trimmed = candidate.trim();
|
||||||
if !candidate.is_empty() {
|
if !trimmed.is_empty() {
|
||||||
return Some(candidate.to_string());
|
return Some(if trimmed.len() == candidate.len() {
|
||||||
|
Cow::Owned(candidate)
|
||||||
|
} else {
|
||||||
|
Cow::Owned(trimmed.to_string())
|
||||||
|
});
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
{
|
{
|
||||||
return Some(DEVELOPMENT_ENCRYPTION_KEY.to_string());
|
return Some(Cow::Borrowed(DEVELOPMENT_ENCRYPTION_KEY));
|
||||||
}
|
}
|
||||||
#[allow(unreachable_code)]
|
#[allow(unreachable_code)]
|
||||||
None
|
None
|
||||||
@@ -90,7 +95,7 @@ pub(crate) fn encrypt_catalog_secret_with_fallbacks(
|
|||||||
plaintext: &str,
|
plaintext: &str,
|
||||||
) -> Option<String> {
|
) -> Option<String> {
|
||||||
let encryption_key = effective_catalog_encryption_key(state)?;
|
let encryption_key = effective_catalog_encryption_key(state)?;
|
||||||
encrypt_python_fernet_plaintext(&encryption_key, plaintext).ok()
|
encrypt_python_fernet_plaintext(encryption_key.as_ref(), plaintext).ok()
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn masked_catalog_api_key(state: &AppState, key: &StoredProviderCatalogKey) -> String {
|
pub(crate) fn masked_catalog_api_key(state: &AppState, key: &StoredProviderCatalogKey) -> String {
|
||||||
|
|||||||
@@ -64,13 +64,13 @@ pub(crate) fn build_local_attempt_identities(
|
|||||||
candidate_index: u32,
|
candidate_index: u32,
|
||||||
transport: &GatewayProviderTransportSnapshot,
|
transport: &GatewayProviderTransportSnapshot,
|
||||||
) -> Vec<ExecutionAttemptIdentity> {
|
) -> Vec<ExecutionAttemptIdentity> {
|
||||||
let attempt_slots = resolve_local_attempt_slot_count(transport);
|
let attempt_slots = local_attempt_slot_count(transport);
|
||||||
(0..attempt_slots)
|
(0..attempt_slots)
|
||||||
.map(|retry_index| ExecutionAttemptIdentity::new(candidate_index, retry_index))
|
.map(|retry_index| ExecutionAttemptIdentity::new(candidate_index, retry_index))
|
||||||
.collect()
|
.collect()
|
||||||
}
|
}
|
||||||
|
|
||||||
fn resolve_local_attempt_slot_count(transport: &GatewayProviderTransportSnapshot) -> u32 {
|
pub(crate) fn local_attempt_slot_count(transport: &GatewayProviderTransportSnapshot) -> u32 {
|
||||||
local_attempt_slots_from_transport(transport).unwrap_or(1)
|
local_attempt_slots_from_transport(transport).unwrap_or(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ pub(crate) use self::adaptive::{
|
|||||||
LocalAdaptiveRateLimitProjection, LocalAdaptiveSuccessProjection,
|
LocalAdaptiveRateLimitProjection, LocalAdaptiveSuccessProjection,
|
||||||
};
|
};
|
||||||
pub(crate) use self::attempt::{
|
pub(crate) use self::attempt::{
|
||||||
attempt_identity_from_report_context, build_local_attempt_identities,
|
attempt_identity_from_report_context, build_local_attempt_identities, local_attempt_slot_count,
|
||||||
local_execution_candidate_metadata_from_report_context, ExecutionAttemptIdentity,
|
local_execution_candidate_metadata_from_report_context, ExecutionAttemptIdentity,
|
||||||
LocalExecutionCandidateMetadata,
|
LocalExecutionCandidateMetadata,
|
||||||
};
|
};
|
||||||
|
|||||||
@@ -21,6 +21,19 @@ use crate::clock::current_unix_ms;
|
|||||||
use crate::log_ids::short_request_id;
|
use crate::log_ids::short_request_id;
|
||||||
use crate::GatewayError;
|
use crate::GatewayError;
|
||||||
|
|
||||||
|
#[derive(Debug, Clone)]
|
||||||
|
pub(crate) struct LocalRequestCandidateStatusSnapshot {
|
||||||
|
candidate_id: String,
|
||||||
|
request_id: String,
|
||||||
|
user_id: Option<String>,
|
||||||
|
api_key_id: Option<String>,
|
||||||
|
candidate_index: u32,
|
||||||
|
retry_index: u32,
|
||||||
|
provider_id: String,
|
||||||
|
endpoint_id: String,
|
||||||
|
key_id: String,
|
||||||
|
}
|
||||||
|
|
||||||
#[async_trait]
|
#[async_trait]
|
||||||
pub(crate) trait RequestCandidateRuntimeReader {
|
pub(crate) trait RequestCandidateRuntimeReader {
|
||||||
async fn read_request_candidates_by_request_id(
|
async fn read_request_candidates_by_request_id(
|
||||||
@@ -149,23 +162,37 @@ fn request_candidate_status_label(status: RequestCandidateStatus) -> &'static st
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) async fn record_local_request_candidate_status(
|
pub(crate) fn snapshot_local_request_candidate_status(
|
||||||
state: &(impl RequestCandidateRuntimeWriter + ?Sized),
|
|
||||||
plan: &ExecutionPlan,
|
plan: &ExecutionPlan,
|
||||||
report_context: Option<&Value>,
|
report_context: Option<&Value>,
|
||||||
status_update: SchedulerRequestCandidateStatusUpdate,
|
) -> Option<LocalRequestCandidateStatusSnapshot> {
|
||||||
|
let candidate_id = plan
|
||||||
|
.candidate_id
|
||||||
|
.as_deref()
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|value| !value.is_empty())?;
|
||||||
|
let metadata = parse_request_candidate_report_context(report_context)?;
|
||||||
|
let candidate_index = metadata.candidate_index?;
|
||||||
|
|
||||||
|
Some(LocalRequestCandidateStatusSnapshot {
|
||||||
|
candidate_id: candidate_id.to_string(),
|
||||||
|
request_id: plan.request_id.clone(),
|
||||||
|
user_id: metadata.user_id,
|
||||||
|
api_key_id: metadata.api_key_id,
|
||||||
|
candidate_index,
|
||||||
|
retry_index: metadata.retry_index,
|
||||||
|
provider_id: plan.provider_id.clone(),
|
||||||
|
endpoint_id: plan.endpoint_id.clone(),
|
||||||
|
key_id: plan.key_id.clone(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn persist_local_request_candidate_status_record(
|
||||||
|
state: &(impl RequestCandidateRuntimeWriter + ?Sized),
|
||||||
|
record: UpsertRequestCandidateRecord,
|
||||||
) {
|
) {
|
||||||
let Some(record) =
|
|
||||||
build_local_request_candidate_status_record(LocalRequestCandidateStatusRecordInput {
|
|
||||||
plan,
|
|
||||||
report_context,
|
|
||||||
status_update,
|
|
||||||
})
|
|
||||||
else {
|
|
||||||
return;
|
|
||||||
};
|
|
||||||
let candidate_id = record.id.clone();
|
let candidate_id = record.id.clone();
|
||||||
let request_id = short_request_id(plan.request_id.as_str());
|
let request_id = short_request_id(record.request_id.as_str());
|
||||||
let candidate_index = record.candidate_index;
|
let candidate_index = record.candidate_index;
|
||||||
let retry_index = record.retry_index;
|
let retry_index = record.retry_index;
|
||||||
let status = record.status;
|
let status = record.status;
|
||||||
@@ -210,6 +237,67 @@ pub(crate) async fn record_local_request_candidate_status(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub(crate) async fn record_local_request_candidate_status(
|
||||||
|
state: &(impl RequestCandidateRuntimeWriter + ?Sized),
|
||||||
|
plan: &ExecutionPlan,
|
||||||
|
report_context: Option<&Value>,
|
||||||
|
status_update: SchedulerRequestCandidateStatusUpdate,
|
||||||
|
) {
|
||||||
|
let Some(record) =
|
||||||
|
build_local_request_candidate_status_record(LocalRequestCandidateStatusRecordInput {
|
||||||
|
plan,
|
||||||
|
report_context,
|
||||||
|
status_update,
|
||||||
|
})
|
||||||
|
else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
persist_local_request_candidate_status_record(state, record).await;
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(crate) async fn record_local_request_candidate_status_snapshot(
|
||||||
|
state: &(impl RequestCandidateRuntimeWriter + ?Sized),
|
||||||
|
snapshot: &LocalRequestCandidateStatusSnapshot,
|
||||||
|
status_update: SchedulerRequestCandidateStatusUpdate,
|
||||||
|
) {
|
||||||
|
let SchedulerRequestCandidateStatusUpdate {
|
||||||
|
status,
|
||||||
|
status_code,
|
||||||
|
error_type,
|
||||||
|
error_message,
|
||||||
|
latency_ms,
|
||||||
|
started_at_unix_ms,
|
||||||
|
finished_at_unix_ms,
|
||||||
|
} = status_update;
|
||||||
|
let record = UpsertRequestCandidateRecord {
|
||||||
|
id: snapshot.candidate_id.clone(),
|
||||||
|
request_id: snapshot.request_id.clone(),
|
||||||
|
user_id: snapshot.user_id.clone(),
|
||||||
|
api_key_id: snapshot.api_key_id.clone(),
|
||||||
|
username: None,
|
||||||
|
api_key_name: None,
|
||||||
|
candidate_index: snapshot.candidate_index,
|
||||||
|
retry_index: snapshot.retry_index,
|
||||||
|
provider_id: Some(snapshot.provider_id.clone()),
|
||||||
|
endpoint_id: Some(snapshot.endpoint_id.clone()),
|
||||||
|
key_id: Some(snapshot.key_id.clone()),
|
||||||
|
status,
|
||||||
|
skip_reason: None,
|
||||||
|
is_cached: None,
|
||||||
|
status_code,
|
||||||
|
error_type,
|
||||||
|
error_message,
|
||||||
|
latency_ms,
|
||||||
|
concurrent_requests: None,
|
||||||
|
extra_data: None,
|
||||||
|
required_capabilities: None,
|
||||||
|
created_at_unix_ms: None,
|
||||||
|
started_at_unix_ms,
|
||||||
|
finished_at_unix_ms,
|
||||||
|
};
|
||||||
|
persist_local_request_candidate_status_record(state, record).await;
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) async fn record_report_request_candidate_status(
|
pub(crate) async fn record_report_request_candidate_status(
|
||||||
state: &(impl RequestCandidateRuntimeReader + RequestCandidateRuntimeWriter + ?Sized),
|
state: &(impl RequestCandidateRuntimeReader + RequestCandidateRuntimeWriter + ?Sized),
|
||||||
report_context: Option<&Value>,
|
report_context: Option<&Value>,
|
||||||
|
|||||||
@@ -389,3 +389,109 @@ async fn gateway_stops_execution_runtime_stream_when_client_disconnects() {
|
|||||||
execution_runtime_handle.abort();
|
execution_runtime_handle.abort();
|
||||||
upstream_handle.abort();
|
upstream_handle.abort();
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
|
||||||
|
async fn gateway_returns_error_body_when_prefetch_detects_embedded_stream_error() {
|
||||||
|
let public_hits = Arc::new(Mutex::new(0usize));
|
||||||
|
let public_hits_clone = Arc::clone(&public_hits);
|
||||||
|
|
||||||
|
let upstream = Router::new().route(
|
||||||
|
"/v1/chat/completions",
|
||||||
|
any(move |_request: Request| {
|
||||||
|
let public_hits_inner = Arc::clone(&public_hits_clone);
|
||||||
|
async move {
|
||||||
|
*public_hits_inner.lock().expect("mutex should lock") += 1;
|
||||||
|
(StatusCode::IM_A_TEAPOT, Body::from("public-route-hit"))
|
||||||
|
}
|
||||||
|
}),
|
||||||
|
);
|
||||||
|
|
||||||
|
let execution_runtime = Router::new().route(
|
||||||
|
"/v1/execute/stream",
|
||||||
|
any(|_request: Request| async move {
|
||||||
|
let body_stream = async_stream::stream! {
|
||||||
|
yield Ok::<Bytes, Infallible>(Bytes::from_static(
|
||||||
|
b"{\"type\":\"headers\",\"payload\":{\"kind\":\"headers\",\"status_code\":200,\"headers\":{\"content-type\":\"text/event-stream\"}}}\n"
|
||||||
|
));
|
||||||
|
yield Ok::<Bytes, Infallible>(Bytes::from_static(
|
||||||
|
b"{\"type\":\"data\",\"payload\":{\"kind\":\"data\",\"text\":\"data: {\\\"error\\\":{\\\"message\\\":\\\"slow down\\\",\\\"type\\\":\\\"rate_limit_error\\\",\\\"code\\\":\\\"rate_limit\\\"}}\\n\\n\"}}\n"
|
||||||
|
));
|
||||||
|
};
|
||||||
|
let mut response = Response::builder()
|
||||||
|
.status(StatusCode::OK)
|
||||||
|
.body(Body::from_stream(body_stream))
|
||||||
|
.expect("response should build");
|
||||||
|
response.headers_mut().insert(
|
||||||
|
http::header::CONTENT_TYPE,
|
||||||
|
HeaderValue::from_static("application/x-ndjson"),
|
||||||
|
);
|
||||||
|
response
|
||||||
|
}),
|
||||||
|
);
|
||||||
|
|
||||||
|
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||||
|
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
|
||||||
|
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
|
||||||
|
Some(hash_api_key("sk-client-openai-stream-prefetch-error")),
|
||||||
|
sample_local_openai_auth_snapshot(
|
||||||
|
"api-key-openai-lifecycle-local-1",
|
||||||
|
"user-openai-lifecycle-local-1",
|
||||||
|
),
|
||||||
|
)]));
|
||||||
|
let candidate_selection_repository =
|
||||||
|
Arc::new(InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||||
|
sample_local_openai_candidate_row(),
|
||||||
|
]));
|
||||||
|
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
|
||||||
|
vec![sample_local_openai_provider()],
|
||||||
|
vec![sample_local_openai_endpoint()],
|
||||||
|
vec![sample_local_openai_key()],
|
||||||
|
));
|
||||||
|
let request_candidate_repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
||||||
|
let gateway = build_router_with_state(
|
||||||
|
build_state_with_execution_runtime_override(execution_runtime_url)
|
||||||
|
.with_data_state_for_tests(
|
||||||
|
GatewayDataState::with_auth_candidate_selection_provider_catalog_and_request_candidate_repository_for_tests(
|
||||||
|
auth_repository,
|
||||||
|
candidate_selection_repository,
|
||||||
|
provider_catalog_repository,
|
||||||
|
request_candidate_repository,
|
||||||
|
DEVELOPMENT_ENCRYPTION_KEY,
|
||||||
|
),
|
||||||
|
),
|
||||||
|
);
|
||||||
|
let (gateway_url, gateway_handle) = start_server(gateway).await;
|
||||||
|
|
||||||
|
let response = reqwest::Client::new()
|
||||||
|
.post(format!("{gateway_url}/v1/chat/completions"))
|
||||||
|
.header(http::header::CONTENT_TYPE, "application/json")
|
||||||
|
.header(
|
||||||
|
http::header::AUTHORIZATION,
|
||||||
|
"Bearer sk-client-openai-stream-prefetch-error",
|
||||||
|
)
|
||||||
|
.header(
|
||||||
|
TRACE_ID_HEADER,
|
||||||
|
"trace-openai-chat-stream-prefetch-error-123",
|
||||||
|
)
|
||||||
|
.body("{\"model\":\"gpt-5\",\"messages\":[],\"stream\":true}")
|
||||||
|
.send()
|
||||||
|
.await
|
||||||
|
.expect("request should succeed");
|
||||||
|
|
||||||
|
assert_eq!(response.status(), StatusCode::OK);
|
||||||
|
assert_eq!(
|
||||||
|
response
|
||||||
|
.headers()
|
||||||
|
.get(http::header::CONTENT_TYPE)
|
||||||
|
.and_then(|value| value.to_str().ok()),
|
||||||
|
Some("text/event-stream")
|
||||||
|
);
|
||||||
|
let body_text = response.text().await.expect("response body should read");
|
||||||
|
assert!(body_text.contains("\"rate_limit_error\""));
|
||||||
|
assert!(body_text.contains("\"slow down\""));
|
||||||
|
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
|
||||||
|
|
||||||
|
gateway_handle.abort();
|
||||||
|
execution_runtime_handle.abort();
|
||||||
|
upstream_handle.abort();
|
||||||
|
}
|
||||||
|
|||||||
@@ -61,28 +61,28 @@ fn log_dropped_report(
|
|||||||
|
|
||||||
pub(crate) async fn submit_sync_report(
|
pub(crate) async fn submit_sync_report(
|
||||||
state: &AppState,
|
state: &AppState,
|
||||||
trace_id: &str,
|
mut payload: GatewaySyncReportRequest,
|
||||||
payload: GatewaySyncReportRequest,
|
|
||||||
) -> Result<(), GatewayError> {
|
) -> Result<(), GatewayError> {
|
||||||
|
let original_report_context = payload.report_context.take();
|
||||||
if let Some(report_context) =
|
if let Some(report_context) =
|
||||||
resolve_locally_actionable_report_context(state, payload.report_context.as_ref()).await
|
resolve_locally_actionable_report_context(state, original_report_context.as_ref()).await
|
||||||
{
|
{
|
||||||
let mut local_payload = payload.clone();
|
payload.report_context = Some(report_context);
|
||||||
local_payload.report_context = Some(report_context);
|
|
||||||
if should_handle_local_sync_report(
|
if should_handle_local_sync_report(
|
||||||
local_payload.report_context.as_ref(),
|
payload.report_context.as_ref(),
|
||||||
local_payload.report_kind.as_str(),
|
payload.report_kind.as_str(),
|
||||||
) {
|
) {
|
||||||
handle_local_sync_report(state, &local_payload).await;
|
handle_local_sync_report(state, &payload).await;
|
||||||
log_local_report_handled(
|
log_local_report_handled(
|
||||||
trace_id,
|
payload.trace_id.as_str(),
|
||||||
&local_payload.report_kind,
|
&payload.report_kind,
|
||||||
"sync",
|
"sync",
|
||||||
local_payload.report_context.as_ref(),
|
payload.report_context.as_ref(),
|
||||||
);
|
);
|
||||||
return Ok(());
|
return Ok(());
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
payload.report_context = original_report_context;
|
||||||
|
|
||||||
if should_handle_local_sync_report(
|
if should_handle_local_sync_report(
|
||||||
payload.report_context.as_ref(),
|
payload.report_context.as_ref(),
|
||||||
@@ -90,7 +90,7 @@ pub(crate) async fn submit_sync_report(
|
|||||||
) {
|
) {
|
||||||
handle_local_sync_report(state, &payload).await;
|
handle_local_sync_report(state, &payload).await;
|
||||||
log_local_report_handled(
|
log_local_report_handled(
|
||||||
trace_id,
|
payload.trace_id.as_str(),
|
||||||
&payload.report_kind,
|
&payload.report_kind,
|
||||||
"sync",
|
"sync",
|
||||||
payload.report_context.as_ref(),
|
payload.report_context.as_ref(),
|
||||||
@@ -99,7 +99,7 @@ pub(crate) async fn submit_sync_report(
|
|||||||
}
|
}
|
||||||
|
|
||||||
log_dropped_report(
|
log_dropped_report(
|
||||||
trace_id,
|
payload.trace_id.as_str(),
|
||||||
&payload.report_kind,
|
&payload.report_kind,
|
||||||
"sync",
|
"sync",
|
||||||
payload.report_context.as_ref(),
|
payload.report_context.as_ref(),
|
||||||
@@ -107,15 +107,12 @@ pub(crate) async fn submit_sync_report(
|
|||||||
Ok(())
|
Ok(())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn spawn_sync_report(
|
pub(crate) fn spawn_sync_report(state: AppState, payload: GatewaySyncReportRequest) {
|
||||||
state: AppState,
|
|
||||||
trace_id: String,
|
|
||||||
payload: GatewaySyncReportRequest,
|
|
||||||
) {
|
|
||||||
let report_request_id_for_log =
|
let report_request_id_for_log =
|
||||||
short_request_id(report_request_id(payload.report_context.as_ref()));
|
short_request_id(report_request_id(payload.report_context.as_ref()));
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
if let Err(err) = submit_sync_report(&state, &trace_id, payload).await {
|
let trace_id = payload.trace_id.clone();
|
||||||
|
if let Err(err) = submit_sync_report(&state, payload).await {
|
||||||
warn!(
|
warn!(
|
||||||
event_name = "execution_report_submit_failed",
|
event_name = "execution_report_submit_failed",
|
||||||
log_type = "ops",
|
log_type = "ops",
|
||||||
@@ -131,39 +128,28 @@ pub(crate) fn spawn_sync_report(
|
|||||||
|
|
||||||
pub(crate) async fn submit_stream_report(
|
pub(crate) async fn submit_stream_report(
|
||||||
state: &AppState,
|
state: &AppState,
|
||||||
trace_id: &str,
|
mut payload: GatewayStreamReportRequest,
|
||||||
payload: GatewayStreamReportRequest,
|
|
||||||
) -> Result<(), GatewayError> {
|
) -> Result<(), GatewayError> {
|
||||||
|
let original_report_context = payload.report_context.take();
|
||||||
if let Some(report_context) =
|
if let Some(report_context) =
|
||||||
resolve_locally_actionable_report_context(state, payload.report_context.as_ref()).await
|
resolve_locally_actionable_report_context(state, original_report_context.as_ref()).await
|
||||||
{
|
{
|
||||||
let local_payload = GatewayStreamReportRequest {
|
payload.report_context = Some(report_context);
|
||||||
trace_id: payload.trace_id.clone(),
|
|
||||||
report_kind: payload.report_kind.clone(),
|
|
||||||
report_context: Some(report_context),
|
|
||||||
status_code: payload.status_code,
|
|
||||||
headers: payload.headers.clone(),
|
|
||||||
provider_body_base64: payload.provider_body_base64.clone(),
|
|
||||||
provider_body_state: payload.provider_body_state,
|
|
||||||
client_body_base64: payload.client_body_base64.clone(),
|
|
||||||
client_body_state: payload.client_body_state,
|
|
||||||
terminal_summary: payload.terminal_summary.clone(),
|
|
||||||
telemetry: payload.telemetry.clone(),
|
|
||||||
};
|
|
||||||
if should_handle_local_stream_report(
|
if should_handle_local_stream_report(
|
||||||
local_payload.report_context.as_ref(),
|
payload.report_context.as_ref(),
|
||||||
local_payload.report_kind.as_str(),
|
payload.report_kind.as_str(),
|
||||||
) {
|
) {
|
||||||
handle_local_stream_report(state, &local_payload).await;
|
handle_local_stream_report(state, &payload).await;
|
||||||
log_local_report_handled(
|
log_local_report_handled(
|
||||||
trace_id,
|
payload.trace_id.as_str(),
|
||||||
&local_payload.report_kind,
|
&payload.report_kind,
|
||||||
"stream",
|
"stream",
|
||||||
local_payload.report_context.as_ref(),
|
payload.report_context.as_ref(),
|
||||||
);
|
);
|
||||||
return Ok(());
|
return Ok(());
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
payload.report_context = original_report_context;
|
||||||
|
|
||||||
if should_handle_local_stream_report(
|
if should_handle_local_stream_report(
|
||||||
payload.report_context.as_ref(),
|
payload.report_context.as_ref(),
|
||||||
@@ -171,7 +157,7 @@ pub(crate) async fn submit_stream_report(
|
|||||||
) {
|
) {
|
||||||
handle_local_stream_report(state, &payload).await;
|
handle_local_stream_report(state, &payload).await;
|
||||||
log_local_report_handled(
|
log_local_report_handled(
|
||||||
trace_id,
|
payload.trace_id.as_str(),
|
||||||
&payload.report_kind,
|
&payload.report_kind,
|
||||||
"stream",
|
"stream",
|
||||||
payload.report_context.as_ref(),
|
payload.report_context.as_ref(),
|
||||||
@@ -180,7 +166,7 @@ pub(crate) async fn submit_stream_report(
|
|||||||
}
|
}
|
||||||
|
|
||||||
log_dropped_report(
|
log_dropped_report(
|
||||||
trace_id,
|
payload.trace_id.as_str(),
|
||||||
&payload.report_kind,
|
&payload.report_kind,
|
||||||
"stream",
|
"stream",
|
||||||
payload.report_context.as_ref(),
|
payload.report_context.as_ref(),
|
||||||
@@ -545,7 +531,6 @@ mod tests {
|
|||||||
|
|
||||||
submit_sync_report(
|
submit_sync_report(
|
||||||
&state,
|
&state,
|
||||||
"trace-reporting-sync-123",
|
|
||||||
GatewaySyncReportRequest {
|
GatewaySyncReportRequest {
|
||||||
trace_id: "trace-reporting-sync-123".to_string(),
|
trace_id: "trace-reporting-sync-123".to_string(),
|
||||||
report_kind: "openai_chat_sync_success".to_string(),
|
report_kind: "openai_chat_sync_success".to_string(),
|
||||||
@@ -586,7 +571,6 @@ mod tests {
|
|||||||
|
|
||||||
submit_sync_report(
|
submit_sync_report(
|
||||||
&state,
|
&state,
|
||||||
"trace-reporting-sync-null-1",
|
|
||||||
GatewaySyncReportRequest {
|
GatewaySyncReportRequest {
|
||||||
trace_id: "trace-reporting-sync-null-1".to_string(),
|
trace_id: "trace-reporting-sync-null-1".to_string(),
|
||||||
report_kind: "claude_cli_sync_success".to_string(),
|
report_kind: "claude_cli_sync_success".to_string(),
|
||||||
@@ -629,7 +613,6 @@ mod tests {
|
|||||||
|
|
||||||
submit_stream_report(
|
submit_stream_report(
|
||||||
&state,
|
&state,
|
||||||
"trace-reporting-stream-123",
|
|
||||||
GatewayStreamReportRequest {
|
GatewayStreamReportRequest {
|
||||||
trace_id: "trace-reporting-stream-123".to_string(),
|
trace_id: "trace-reporting-stream-123".to_string(),
|
||||||
report_kind: "openai_chat_stream_success".to_string(),
|
report_kind: "openai_chat_stream_success".to_string(),
|
||||||
@@ -682,7 +665,6 @@ mod tests {
|
|||||||
|
|
||||||
submit_sync_report(
|
submit_sync_report(
|
||||||
&state,
|
&state,
|
||||||
"trace-codex-reporting-sync",
|
|
||||||
GatewaySyncReportRequest {
|
GatewaySyncReportRequest {
|
||||||
trace_id: "trace-codex-reporting-sync".to_string(),
|
trace_id: "trace-codex-reporting-sync".to_string(),
|
||||||
report_kind: "openai_cli_sync_success".to_string(),
|
report_kind: "openai_cli_sync_success".to_string(),
|
||||||
@@ -746,7 +728,6 @@ mod tests {
|
|||||||
|
|
||||||
submit_stream_report(
|
submit_stream_report(
|
||||||
&state,
|
&state,
|
||||||
"trace-codex-reporting-stream",
|
|
||||||
GatewayStreamReportRequest {
|
GatewayStreamReportRequest {
|
||||||
trace_id: "trace-codex-reporting-stream".to_string(),
|
trace_id: "trace-codex-reporting-stream".to_string(),
|
||||||
report_kind: "openai_cli_stream_success".to_string(),
|
report_kind: "openai_cli_stream_success".to_string(),
|
||||||
@@ -812,7 +793,6 @@ mod tests {
|
|||||||
|
|
||||||
submit_sync_report(
|
submit_sync_report(
|
||||||
&state,
|
&state,
|
||||||
"trace-gemini-files-store-123",
|
|
||||||
GatewaySyncReportRequest {
|
GatewaySyncReportRequest {
|
||||||
trace_id: "trace-gemini-files-store-123".to_string(),
|
trace_id: "trace-gemini-files-store-123".to_string(),
|
||||||
report_kind: "gemini_files_store_mapping".to_string(),
|
report_kind: "gemini_files_store_mapping".to_string(),
|
||||||
@@ -883,7 +863,6 @@ mod tests {
|
|||||||
|
|
||||||
submit_sync_report(
|
submit_sync_report(
|
||||||
&state,
|
&state,
|
||||||
"trace-gemini-files-delete-123",
|
|
||||||
GatewaySyncReportRequest {
|
GatewaySyncReportRequest {
|
||||||
trace_id: "trace-gemini-files-delete-123".to_string(),
|
trace_id: "trace-gemini-files-delete-123".to_string(),
|
||||||
report_kind: "gemini_files_delete_mapping".to_string(),
|
report_kind: "gemini_files_delete_mapping".to_string(),
|
||||||
@@ -926,7 +905,6 @@ mod tests {
|
|||||||
|
|
||||||
submit_sync_report(
|
submit_sync_report(
|
||||||
&state,
|
&state,
|
||||||
"trace-reporting-video-delete-123",
|
|
||||||
GatewaySyncReportRequest {
|
GatewaySyncReportRequest {
|
||||||
trace_id: "trace-reporting-video-delete-123".to_string(),
|
trace_id: "trace-reporting-video-delete-123".to_string(),
|
||||||
report_kind: "openai_video_delete_sync_success".to_string(),
|
report_kind: "openai_video_delete_sync_success".to_string(),
|
||||||
@@ -991,7 +969,6 @@ mod tests {
|
|||||||
|
|
||||||
submit_sync_report(
|
submit_sync_report(
|
||||||
&state,
|
&state,
|
||||||
"trace-openai-video-reporting-123",
|
|
||||||
GatewaySyncReportRequest {
|
GatewaySyncReportRequest {
|
||||||
trace_id: "trace-openai-video-reporting-123".to_string(),
|
trace_id: "trace-openai-video-reporting-123".to_string(),
|
||||||
report_kind: "openai_video_create_sync_success".to_string(),
|
report_kind: "openai_video_create_sync_success".to_string(),
|
||||||
@@ -1053,7 +1030,6 @@ mod tests {
|
|||||||
|
|
||||||
submit_sync_report(
|
submit_sync_report(
|
||||||
&state,
|
&state,
|
||||||
"trace-gemini-video-reporting-123",
|
|
||||||
GatewaySyncReportRequest {
|
GatewaySyncReportRequest {
|
||||||
trace_id: "trace-gemini-video-reporting-123".to_string(),
|
trace_id: "trace-gemini-video-reporting-123".to_string(),
|
||||||
report_kind: "gemini_video_create_sync_success".to_string(),
|
report_kind: "gemini_video_create_sync_success".to_string(),
|
||||||
@@ -1103,7 +1079,6 @@ mod tests {
|
|||||||
|
|
||||||
submit_sync_report(
|
submit_sync_report(
|
||||||
&state,
|
&state,
|
||||||
"trace-openai-video-task-id-123",
|
|
||||||
GatewaySyncReportRequest {
|
GatewaySyncReportRequest {
|
||||||
trace_id: "trace-openai-video-task-id-123".to_string(),
|
trace_id: "trace-openai-video-task-id-123".to_string(),
|
||||||
report_kind: "openai_video_cancel_sync_success".to_string(),
|
report_kind: "openai_video_cancel_sync_success".to_string(),
|
||||||
@@ -1206,7 +1181,6 @@ mod tests {
|
|||||||
|
|
||||||
submit_sync_report(
|
submit_sync_report(
|
||||||
&state,
|
&state,
|
||||||
"trace-gemini-video-external-id-123",
|
|
||||||
GatewaySyncReportRequest {
|
GatewaySyncReportRequest {
|
||||||
trace_id: "trace-gemini-video-external-id-123".to_string(),
|
trace_id: "trace-gemini-video-external-id-123".to_string(),
|
||||||
report_kind: "gemini_video_cancel_sync_success".to_string(),
|
report_kind: "gemini_video_cancel_sync_success".to_string(),
|
||||||
|
|||||||
@@ -80,7 +80,7 @@ pub use crate::conversion::{
|
|||||||
pub use crate::finalize::common::{
|
pub use crate::finalize::common::{
|
||||||
build_generated_tool_call_id, build_local_success_background_report,
|
build_generated_tool_call_id, build_local_success_background_report,
|
||||||
build_local_success_conversion_background_report, canonicalize_tool_arguments,
|
build_local_success_conversion_background_report, canonicalize_tool_arguments,
|
||||||
prepare_local_success_response_parts,
|
prepare_local_success_response_parts, prepare_local_success_response_parts_owned,
|
||||||
};
|
};
|
||||||
pub use crate::finalize::sse::{encode_done_sse, encode_json_sse, map_claude_stop_reason};
|
pub use crate::finalize::sse::{encode_done_sse, encode_json_sse, map_claude_stop_reason};
|
||||||
pub use crate::finalize::standard::claude::stream::{ClaudeClientEmitter, ClaudeProviderState};
|
pub use crate::finalize::standard::claude::stream::{ClaudeClientEmitter, ClaudeProviderState};
|
||||||
|
|||||||
@@ -21,7 +21,13 @@ pub fn prepare_local_success_response_parts(
|
|||||||
headers: &BTreeMap<String, String>,
|
headers: &BTreeMap<String, String>,
|
||||||
body_json: &Value,
|
body_json: &Value,
|
||||||
) -> serde_json::Result<(Vec<u8>, BTreeMap<String, String>)> {
|
) -> serde_json::Result<(Vec<u8>, BTreeMap<String, String>)> {
|
||||||
let mut headers = headers.clone();
|
prepare_local_success_response_parts_owned(headers.clone(), body_json)
|
||||||
|
}
|
||||||
|
|
||||||
|
pub fn prepare_local_success_response_parts_owned(
|
||||||
|
mut headers: BTreeMap<String, String>,
|
||||||
|
body_json: &Value,
|
||||||
|
) -> serde_json::Result<(Vec<u8>, BTreeMap<String, String>)> {
|
||||||
headers.remove("content-encoding");
|
headers.remove("content-encoding");
|
||||||
headers.remove("content-length");
|
headers.remove("content-length");
|
||||||
headers.insert("content-type".to_string(), "application/json".to_string());
|
headers.insert("content-type".to_string(), "application/json".to_string());
|
||||||
@@ -77,7 +83,7 @@ mod tests {
|
|||||||
use super::{
|
use super::{
|
||||||
build_generated_tool_call_id, build_local_success_background_report,
|
build_generated_tool_call_id, build_local_success_background_report,
|
||||||
build_local_success_conversion_background_report, canonicalize_tool_arguments,
|
build_local_success_conversion_background_report, canonicalize_tool_arguments,
|
||||||
prepare_local_success_response_parts,
|
prepare_local_success_response_parts, prepare_local_success_response_parts_owned,
|
||||||
};
|
};
|
||||||
use aether_usage_runtime::GatewaySyncReportRequest;
|
use aether_usage_runtime::GatewaySyncReportRequest;
|
||||||
use std::collections::BTreeMap;
|
use std::collections::BTreeMap;
|
||||||
@@ -127,6 +133,37 @@ mod tests {
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn prepare_local_success_response_parts_owned_normalizes_headers() {
|
||||||
|
let headers = BTreeMap::from([
|
||||||
|
("content-encoding".to_string(), "gzip".to_string()),
|
||||||
|
("content-length".to_string(), "999".to_string()),
|
||||||
|
("x-test".to_string(), "1".to_string()),
|
||||||
|
]);
|
||||||
|
let (body_bytes, normalized_headers) =
|
||||||
|
prepare_local_success_response_parts_owned(headers, &serde_json::json!({"ok": true}))
|
||||||
|
.expect("response parts should serialize");
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
serde_json::from_slice::<Value>(&body_bytes).expect("json body"),
|
||||||
|
serde_json::json!({"ok": true})
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
normalized_headers.get("content-type").map(String::as_str),
|
||||||
|
Some("application/json")
|
||||||
|
);
|
||||||
|
assert!(!normalized_headers.contains_key("content-encoding"));
|
||||||
|
let expected_length = body_bytes.len().to_string();
|
||||||
|
assert_eq!(
|
||||||
|
normalized_headers.get("content-length").map(String::as_str),
|
||||||
|
Some(expected_length.as_str())
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
normalized_headers.get("x-test").map(String::as_str),
|
||||||
|
Some("1")
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
#[test]
|
#[test]
|
||||||
fn build_local_success_background_report_maps_finalize_kind() {
|
fn build_local_success_background_report_maps_finalize_kind() {
|
||||||
let payload = GatewaySyncReportRequest {
|
let payload = GatewaySyncReportRequest {
|
||||||
|
|||||||
@@ -1,3 +1,5 @@
|
|||||||
|
use std::borrow::Cow;
|
||||||
|
|
||||||
use aether_provider_transport::url::{
|
use aether_provider_transport::url::{
|
||||||
build_claude_messages_url, build_gemini_content_url, build_openai_chat_url,
|
build_claude_messages_url, build_gemini_content_url, build_openai_chat_url,
|
||||||
build_openai_cli_url, build_passthrough_path_url,
|
build_openai_cli_url, build_passthrough_path_url,
|
||||||
@@ -30,13 +32,13 @@ pub fn build_standard_request_body(
|
|||||||
body_rules: Option<&Value>,
|
body_rules: Option<&Value>,
|
||||||
user_api_key_id: Option<&str>,
|
user_api_key_id: Option<&str>,
|
||||||
) -> Option<Value> {
|
) -> Option<Value> {
|
||||||
let canonical_request = normalize_standard_request_to_openai_chat_request(
|
let canonical_request = normalize_standard_request_to_openai_chat_request_cow(
|
||||||
body_json,
|
body_json,
|
||||||
client_api_format,
|
client_api_format,
|
||||||
request_path,
|
request_path,
|
||||||
)?;
|
)?;
|
||||||
let mut provider_request_body = build_standard_request_body_from_canonical(
|
let mut provider_request_body = build_standard_request_body_from_canonical(
|
||||||
&canonical_request,
|
canonical_request.as_ref(),
|
||||||
mapped_model,
|
mapped_model,
|
||||||
provider_api_format,
|
provider_api_format,
|
||||||
upstream_is_stream,
|
upstream_is_stream,
|
||||||
@@ -99,14 +101,29 @@ pub fn normalize_standard_request_to_openai_chat_request(
|
|||||||
client_api_format: &str,
|
client_api_format: &str,
|
||||||
request_path: &str,
|
request_path: &str,
|
||||||
) -> Option<Value> {
|
) -> Option<Value> {
|
||||||
|
normalize_standard_request_to_openai_chat_request_cow(
|
||||||
|
body_json,
|
||||||
|
client_api_format,
|
||||||
|
request_path,
|
||||||
|
)
|
||||||
|
.map(Cow::into_owned)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn normalize_standard_request_to_openai_chat_request_cow<'a>(
|
||||||
|
body_json: &'a Value,
|
||||||
|
client_api_format: &str,
|
||||||
|
request_path: &str,
|
||||||
|
) -> Option<Cow<'a, Value>> {
|
||||||
match client_api_format.trim().to_ascii_lowercase().as_str() {
|
match client_api_format.trim().to_ascii_lowercase().as_str() {
|
||||||
"openai:chat" => Some(body_json.clone()),
|
"openai:chat" => Some(Cow::Borrowed(body_json)),
|
||||||
"openai:cli" | "openai:compact" => {
|
"openai:cli" | "openai:compact" => {
|
||||||
normalize_openai_cli_request_to_openai_chat_request(body_json)
|
normalize_openai_cli_request_to_openai_chat_request(body_json).map(Cow::Owned)
|
||||||
|
}
|
||||||
|
"claude:chat" | "claude:cli" => {
|
||||||
|
normalize_claude_request_to_openai_chat_request(body_json).map(Cow::Owned)
|
||||||
}
|
}
|
||||||
"claude:chat" | "claude:cli" => normalize_claude_request_to_openai_chat_request(body_json),
|
|
||||||
"gemini:chat" | "gemini:cli" => {
|
"gemini:chat" | "gemini:cli" => {
|
||||||
normalize_gemini_request_to_openai_chat_request(body_json, request_path)
|
normalize_gemini_request_to_openai_chat_request(body_json, request_path).map(Cow::Owned)
|
||||||
}
|
}
|
||||||
_ => None,
|
_ => None,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,8 +1,9 @@
|
|||||||
use std::collections::HashMap;
|
use std::collections::{HashMap, VecDeque};
|
||||||
use std::sync::{LazyLock, Mutex};
|
use std::sync::{Arc, LazyLock, RwLock};
|
||||||
use std::time::{SystemTime, UNIX_EPOCH};
|
use std::time::{SystemTime, UNIX_EPOCH};
|
||||||
|
|
||||||
use aes::cipher::{block_padding::Pkcs7, BlockDecryptMut, BlockEncryptMut, KeyIvInit};
|
use aes::cipher::{block_padding::Pkcs7, BlockDecryptMut, BlockEncryptMut, KeyIvInit};
|
||||||
|
use base64::decoded_len_estimate;
|
||||||
use base64::engine::general_purpose::{STANDARD, URL_SAFE, URL_SAFE_NO_PAD};
|
use base64::engine::general_purpose::{STANDARD, URL_SAFE, URL_SAFE_NO_PAD};
|
||||||
use base64::Engine as _;
|
use base64::Engine as _;
|
||||||
use cbc::{Decryptor, Encryptor};
|
use cbc::{Decryptor, Encryptor};
|
||||||
@@ -16,7 +17,8 @@ const HMAC_SIZE: usize = 32;
|
|||||||
const IV_SIZE: usize = 16;
|
const IV_SIZE: usize = 16;
|
||||||
const SIGNING_KEY_SIZE: usize = 16;
|
const SIGNING_KEY_SIZE: usize = 16;
|
||||||
const ENCRYPTION_KEY_SIZE: usize = 16;
|
const ENCRYPTION_KEY_SIZE: usize = 16;
|
||||||
const MIN_TOKEN_SIZE: usize = 1 + 8 + IV_SIZE + HMAC_SIZE;
|
const MIN_CIPHERTEXT_SIZE: usize = 16;
|
||||||
|
const MIN_TOKEN_SIZE: usize = 1 + 8 + IV_SIZE + MIN_CIPHERTEXT_SIZE + HMAC_SIZE;
|
||||||
const PBKDF2_ITERATIONS: u32 = 100_000;
|
const PBKDF2_ITERATIONS: u32 = 100_000;
|
||||||
const MAX_CACHED_DERIVED_KEYS: usize = 16;
|
const MAX_CACHED_DERIVED_KEYS: usize = 16;
|
||||||
|
|
||||||
@@ -24,13 +26,63 @@ pub const APP_SALT_SEED: &[u8] = b"aether-v1";
|
|||||||
pub const APP_SALT_HEX: &str = "8797080a7a4b45b4810e934d1af36261";
|
pub const APP_SALT_HEX: &str = "8797080a7a4b45b4810e934d1af36261";
|
||||||
pub const DEVELOPMENT_ENCRYPTION_KEY: &str = "dev-encryption-key-do-not-use-in-production";
|
pub const DEVELOPMENT_ENCRYPTION_KEY: &str = "dev-encryption-key-do-not-use-in-production";
|
||||||
|
|
||||||
static RAW_FERNET_KEY_CACHE: LazyLock<Mutex<HashMap<Box<str>, [u8; 32]>>> =
|
static RAW_FERNET_KEY_CACHE: LazyLock<RwLock<RawFernetKeyCache>> =
|
||||||
LazyLock::new(|| Mutex::new(HashMap::new()));
|
LazyLock::new(|| RwLock::new(RawFernetKeyCache::default()));
|
||||||
|
static APP_SALT: LazyLock<[u8; 16]> = LazyLock::new(|| {
|
||||||
|
let mut salt = [0u8; 16];
|
||||||
|
salt.copy_from_slice(&Sha256::digest(APP_SALT_SEED)[..16]);
|
||||||
|
salt
|
||||||
|
});
|
||||||
|
|
||||||
type Aes128CbcDec = Decryptor<aes::Aes128>;
|
type Aes128CbcDec = Decryptor<aes::Aes128>;
|
||||||
type Aes128CbcEnc = Encryptor<aes::Aes128>;
|
type Aes128CbcEnc = Encryptor<aes::Aes128>;
|
||||||
type HmacSha256 = Hmac<Sha256>;
|
type HmacSha256 = Hmac<Sha256>;
|
||||||
|
|
||||||
|
fn base64_encoded_len(input_len: usize) -> usize {
|
||||||
|
input_len.div_ceil(3) * 4
|
||||||
|
}
|
||||||
|
|
||||||
|
fn base64_unpadded_len(input_len: usize) -> usize {
|
||||||
|
let full_chunks = (input_len / 3) * 4;
|
||||||
|
match input_len % 3 {
|
||||||
|
0 => full_chunks,
|
||||||
|
1 => full_chunks + 2,
|
||||||
|
_ => full_chunks + 3,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn minimum_wrapped_token_len() -> usize {
|
||||||
|
base64_unpadded_len(base64_unpadded_len(MIN_TOKEN_SIZE))
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Default)]
|
||||||
|
struct RawFernetKeyCache {
|
||||||
|
entries: HashMap<Arc<str>, [u8; 32]>,
|
||||||
|
insertion_order: VecDeque<Arc<str>>,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl RawFernetKeyCache {
|
||||||
|
fn get(&self, secret: &str) -> Option<[u8; 32]> {
|
||||||
|
self.entries.get(secret).copied()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn insert(&mut self, secret: &str, raw_key: [u8; 32]) {
|
||||||
|
if self.entries.contains_key(secret) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
if self.entries.len() >= MAX_CACHED_DERIVED_KEYS {
|
||||||
|
if let Some(oldest) = self.insertion_order.pop_front() {
|
||||||
|
self.entries.remove(oldest.as_ref());
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let secret: Arc<str> = Arc::from(secret);
|
||||||
|
self.insertion_order.push_back(secret.clone());
|
||||||
|
self.entries.insert(secret, raw_key);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Debug, thiserror::Error)]
|
#[derive(Debug, thiserror::Error)]
|
||||||
pub enum PythonFernetError {
|
pub enum PythonFernetError {
|
||||||
#[error("invalid Python Fernet outer base64 payload")]
|
#[error("invalid Python Fernet outer base64 payload")]
|
||||||
@@ -68,9 +120,9 @@ impl PythonFernetCompat {
|
|||||||
|
|
||||||
let outer =
|
let outer =
|
||||||
decode_urlsafe(ciphertext).map_err(|_| PythonFernetError::InvalidOuterBase64)?;
|
decode_urlsafe(ciphertext).map_err(|_| PythonFernetError::InvalidOuterBase64)?;
|
||||||
let inner =
|
let token =
|
||||||
decode_urlsafe_bytes(&outer).map_err(|_| PythonFernetError::InvalidInnerBase64)?;
|
decode_urlsafe_bytes(&outer).map_err(|_| PythonFernetError::InvalidInnerBase64)?;
|
||||||
let plaintext = self.decrypt_token_bytes(&inner)?;
|
let plaintext = self.decrypt_token_bytes(token)?;
|
||||||
String::from_utf8(plaintext).map_err(PythonFernetError::InvalidUtf8)
|
String::from_utf8(plaintext).map_err(PythonFernetError::InvalidUtf8)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -93,7 +145,7 @@ impl PythonFernetCompat {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn decrypt_token_bytes(&self, token: &[u8]) -> Result<Vec<u8>, PythonFernetError> {
|
fn decrypt_token_bytes(&self, mut token: Vec<u8>) -> Result<Vec<u8>, PythonFernetError> {
|
||||||
if token.len() < MIN_TOKEN_SIZE {
|
if token.len() < MIN_TOKEN_SIZE {
|
||||||
return Err(PythonFernetError::InvalidTokenStructure);
|
return Err(PythonFernetError::InvalidTokenStructure);
|
||||||
}
|
}
|
||||||
@@ -102,22 +154,30 @@ impl PythonFernetCompat {
|
|||||||
}
|
}
|
||||||
|
|
||||||
let signed_len = token.len() - HMAC_SIZE;
|
let signed_len = token.len() - HMAC_SIZE;
|
||||||
let (signed, signature) = token.split_at(signed_len);
|
{
|
||||||
|
let (signed, signature) = token.split_at(signed_len);
|
||||||
let mut mac = HmacSha256::new_from_slice(&self.signing_key)
|
let mut mac = HmacSha256::new_from_slice(&self.signing_key)
|
||||||
.map_err(|_| PythonFernetError::InvalidTokenSignature)?;
|
.map_err(|_| PythonFernetError::InvalidTokenSignature)?;
|
||||||
mac.update(signed);
|
mac.update(signed);
|
||||||
mac.verify_slice(signature)
|
mac.verify_slice(signature)
|
||||||
.map_err(|_| PythonFernetError::InvalidTokenSignature)?;
|
.map_err(|_| PythonFernetError::InvalidTokenSignature)?;
|
||||||
|
}
|
||||||
|
|
||||||
let iv_offset = 1 + 8;
|
let iv_offset = 1 + 8;
|
||||||
let ciphertext_offset = iv_offset + IV_SIZE;
|
let ciphertext_offset = iv_offset + IV_SIZE;
|
||||||
let iv = &token[iv_offset..ciphertext_offset];
|
let plaintext_len = {
|
||||||
let mut ciphertext = token[ciphertext_offset..signed_len].to_vec();
|
let (_, payload) = token.split_at_mut(iv_offset);
|
||||||
let plaintext = Aes128CbcDec::new((&self.encryption_key).into(), iv.into())
|
let (iv, ciphertext_and_signature) = payload.split_at_mut(IV_SIZE);
|
||||||
.decrypt_padded_mut::<Pkcs7>(&mut ciphertext)
|
let ciphertext = &mut ciphertext_and_signature[..signed_len - ciphertext_offset];
|
||||||
.map_err(|_| PythonFernetError::InvalidPadding)?;
|
Aes128CbcDec::new((&self.encryption_key).into(), (&iv[..]).into())
|
||||||
Ok(plaintext.to_vec())
|
.decrypt_padded_mut::<Pkcs7>(ciphertext)
|
||||||
|
.map_err(|_| PythonFernetError::InvalidPadding)?
|
||||||
|
.len()
|
||||||
|
};
|
||||||
|
|
||||||
|
token.copy_within(ciphertext_offset..ciphertext_offset + plaintext_len, 0);
|
||||||
|
token.truncate(plaintext_len);
|
||||||
|
Ok(token)
|
||||||
}
|
}
|
||||||
|
|
||||||
fn encrypt_token(
|
fn encrypt_token(
|
||||||
@@ -131,14 +191,13 @@ impl PythonFernetCompat {
|
|||||||
padded[..plaintext.len()].copy_from_slice(plaintext);
|
padded[..plaintext.len()].copy_from_slice(plaintext);
|
||||||
let ciphertext = Aes128CbcEnc::new((&self.encryption_key).into(), (&iv).into())
|
let ciphertext = Aes128CbcEnc::new((&self.encryption_key).into(), (&iv).into())
|
||||||
.encrypt_padded_mut::<Pkcs7>(&mut padded, plaintext.len())
|
.encrypt_padded_mut::<Pkcs7>(&mut padded, plaintext.len())
|
||||||
.map_err(|_| PythonFernetError::InvalidPadding)?
|
.map_err(|_| PythonFernetError::InvalidPadding)?;
|
||||||
.to_vec();
|
|
||||||
|
|
||||||
let mut signed = Vec::with_capacity(1 + 8 + IV_SIZE + ciphertext.len() + HMAC_SIZE);
|
let mut signed = Vec::with_capacity(1 + 8 + IV_SIZE + ciphertext.len() + HMAC_SIZE);
|
||||||
signed.push(FERNET_VERSION);
|
signed.push(FERNET_VERSION);
|
||||||
signed.extend_from_slice(×tamp.to_be_bytes());
|
signed.extend_from_slice(×tamp.to_be_bytes());
|
||||||
signed.extend_from_slice(&iv);
|
signed.extend_from_slice(&iv);
|
||||||
signed.extend_from_slice(&ciphertext);
|
signed.extend_from_slice(ciphertext);
|
||||||
|
|
||||||
let mut mac = HmacSha256::new_from_slice(&self.signing_key)
|
let mut mac = HmacSha256::new_from_slice(&self.signing_key)
|
||||||
.map_err(|_| PythonFernetError::InvalidTokenSignature)?;
|
.map_err(|_| PythonFernetError::InvalidTokenSignature)?;
|
||||||
@@ -146,8 +205,12 @@ impl PythonFernetCompat {
|
|||||||
let signature = mac.finalize().into_bytes();
|
let signature = mac.finalize().into_bytes();
|
||||||
signed.extend_from_slice(&signature);
|
signed.extend_from_slice(&signature);
|
||||||
|
|
||||||
let inner = URL_SAFE.encode(signed);
|
let mut inner = String::with_capacity(base64_encoded_len(signed.len()));
|
||||||
Ok(URL_SAFE.encode(inner.as_bytes()))
|
URL_SAFE.encode_string(&signed, &mut inner);
|
||||||
|
|
||||||
|
let mut outer = String::with_capacity(base64_encoded_len(inner.len()));
|
||||||
|
URL_SAFE.encode_string(inner.as_bytes(), &mut outer);
|
||||||
|
Ok(outer)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -164,7 +227,7 @@ pub fn decrypt_python_fernet_ciphertext(
|
|||||||
|
|
||||||
pub fn looks_like_python_fernet_ciphertext(ciphertext: &str) -> bool {
|
pub fn looks_like_python_fernet_ciphertext(ciphertext: &str) -> bool {
|
||||||
let ciphertext = ciphertext.trim();
|
let ciphertext = ciphertext.trim();
|
||||||
if ciphertext.is_empty() {
|
if ciphertext.is_empty() || ciphertext.len() < minimum_wrapped_token_len() {
|
||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -190,32 +253,36 @@ pub fn warm_python_fernet_secret(secret: &str) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn raw_fernet_key(secret: &str) -> [u8; 32] {
|
fn raw_fernet_key(secret: &str) -> [u8; 32] {
|
||||||
if let Ok(raw_key) = decode_direct_fernet_key(secret) {
|
|
||||||
return raw_key;
|
|
||||||
}
|
|
||||||
|
|
||||||
if let Some(raw_key) = RAW_FERNET_KEY_CACHE
|
if let Some(raw_key) = RAW_FERNET_KEY_CACHE
|
||||||
.lock()
|
.read()
|
||||||
.expect("raw fernet key cache should lock")
|
.expect("raw fernet key cache should lock")
|
||||||
.get(secret)
|
.get(secret)
|
||||||
.copied()
|
|
||||||
{
|
{
|
||||||
return raw_key;
|
return raw_key;
|
||||||
}
|
}
|
||||||
|
|
||||||
let mut salt = [0u8; 16];
|
|
||||||
salt.copy_from_slice(&Sha256::digest(APP_SALT_SEED)[..16]);
|
|
||||||
|
|
||||||
let mut raw_key = [0u8; 32];
|
|
||||||
pbkdf2_hmac::<Sha256>(secret.as_bytes(), &salt, PBKDF2_ITERATIONS, &mut raw_key);
|
|
||||||
|
|
||||||
let mut cache = RAW_FERNET_KEY_CACHE
|
let mut cache = RAW_FERNET_KEY_CACHE
|
||||||
.lock()
|
.write()
|
||||||
.expect("raw fernet key cache should lock");
|
.expect("raw fernet key cache should lock");
|
||||||
if cache.len() >= MAX_CACHED_DERIVED_KEYS && !cache.contains_key(secret) {
|
if let Some(raw_key) = cache.get(secret) {
|
||||||
cache.clear();
|
return raw_key;
|
||||||
}
|
}
|
||||||
cache.insert(secret.into(), raw_key);
|
|
||||||
|
let raw_key = match decode_direct_fernet_key(secret) {
|
||||||
|
Ok(raw_key) => raw_key,
|
||||||
|
Err(_) => {
|
||||||
|
let mut raw_key = [0u8; 32];
|
||||||
|
pbkdf2_hmac::<Sha256>(
|
||||||
|
secret.as_bytes(),
|
||||||
|
&*APP_SALT,
|
||||||
|
PBKDF2_ITERATIONS,
|
||||||
|
&mut raw_key,
|
||||||
|
);
|
||||||
|
raw_key
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
cache.insert(secret, raw_key);
|
||||||
raw_key
|
raw_key
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -232,15 +299,23 @@ fn decode_direct_fernet_key(secret: &str) -> Result<[u8; 32], PythonFernetError>
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn decode_urlsafe(value: &str) -> Result<Vec<u8>, base64::DecodeError> {
|
fn decode_urlsafe(value: &str) -> Result<Vec<u8>, base64::DecodeError> {
|
||||||
URL_SAFE
|
decode_with_engine_fallback(value.as_bytes())
|
||||||
.decode(value)
|
|
||||||
.or_else(|_| URL_SAFE_NO_PAD.decode(value))
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn decode_urlsafe_bytes(value: &[u8]) -> Result<Vec<u8>, base64::DecodeError> {
|
fn decode_urlsafe_bytes(value: &[u8]) -> Result<Vec<u8>, base64::DecodeError> {
|
||||||
URL_SAFE
|
decode_with_engine_fallback(value)
|
||||||
.decode(value)
|
}
|
||||||
.or_else(|_| URL_SAFE_NO_PAD.decode(value))
|
|
||||||
|
fn decode_with_engine_fallback(value: &[u8]) -> Result<Vec<u8>, base64::DecodeError> {
|
||||||
|
let mut decoded = Vec::with_capacity(decoded_len_estimate(value.len()));
|
||||||
|
match URL_SAFE.decode_vec(value, &mut decoded) {
|
||||||
|
Ok(()) => Ok(decoded),
|
||||||
|
Err(_) => {
|
||||||
|
decoded.clear();
|
||||||
|
URL_SAFE_NO_PAD.decode_vec(value, &mut decoded)?;
|
||||||
|
Ok(decoded)
|
||||||
|
}
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
|
|||||||
@@ -113,11 +113,12 @@ fn decrypt_secret(
|
|||||||
ciphertext: &str,
|
ciphertext: &str,
|
||||||
field_name: &str,
|
field_name: &str,
|
||||||
) -> Result<String, DataLayerError> {
|
) -> Result<String, DataLayerError> {
|
||||||
|
if should_use_plaintext_secret(ciphertext, field_name) {
|
||||||
|
return Ok(ciphertext.trim().to_string());
|
||||||
|
}
|
||||||
|
|
||||||
match decrypt_python_fernet_ciphertext(encryption_key, ciphertext) {
|
match decrypt_python_fernet_ciphertext(encryption_key, ciphertext) {
|
||||||
Ok(value) => Ok(value),
|
Ok(value) => Ok(value),
|
||||||
Err(_error) if should_use_plaintext_secret(ciphertext, field_name) => {
|
|
||||||
Ok(ciphertext.trim().to_string())
|
|
||||||
}
|
|
||||||
Err(error) => {
|
Err(error) => {
|
||||||
for fallback_encryption_key in fallback_encryption_keys {
|
for fallback_encryption_key in fallback_encryption_keys {
|
||||||
if let Ok(value) =
|
if let Ok(value) =
|
||||||
@@ -156,14 +157,19 @@ fn should_use_plaintext_secret(ciphertext: &str, field_name: &str) -> bool {
|
|||||||
if ciphertext.is_empty() {
|
if ciphertext.is_empty() {
|
||||||
return false;
|
return false;
|
||||||
}
|
}
|
||||||
if looks_like_python_fernet_ciphertext(ciphertext) {
|
|
||||||
return false;
|
|
||||||
}
|
|
||||||
|
|
||||||
match field_name {
|
match field_name {
|
||||||
"provider_api_keys.api_key" => !ciphertext.starts_with('{') && !ciphertext.starts_with('['),
|
"provider_api_keys.api_key" => {
|
||||||
|
if ciphertext.starts_with('{') || ciphertext.starts_with('[') {
|
||||||
|
return false;
|
||||||
|
}
|
||||||
|
!looks_like_python_fernet_ciphertext(ciphertext)
|
||||||
|
}
|
||||||
"provider_api_keys.auth_config" => {
|
"provider_api_keys.auth_config" => {
|
||||||
ciphertext.starts_with('{') || ciphertext.starts_with('[')
|
if ciphertext.starts_with('{') || ciphertext.starts_with('[') {
|
||||||
|
return true;
|
||||||
|
}
|
||||||
|
false
|
||||||
}
|
}
|
||||||
_ => false,
|
_ => false,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -45,6 +45,20 @@ pub struct BuildMinimalCandidateSelectionInput<'a> {
|
|||||||
pub priority_mode: SchedulerPriorityMode,
|
pub priority_mode: SchedulerPriorityMode,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Clone, Copy)]
|
||||||
|
struct RequiredCapabilityDescriptor<'a> {
|
||||||
|
name: &'a str,
|
||||||
|
compatible: bool,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Clone, Copy)]
|
||||||
|
struct CandidateOrderingState {
|
||||||
|
capability_priority: (u32, u32),
|
||||||
|
affinity_hash: Option<u64>,
|
||||||
|
health_bucket: Option<crate::ProviderKeyHealthBucket>,
|
||||||
|
health_score: f64,
|
||||||
|
}
|
||||||
|
|
||||||
pub fn candidate_supports_required_capability(
|
pub fn candidate_supports_required_capability(
|
||||||
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
||||||
required_capability: &str,
|
required_capability: &str,
|
||||||
@@ -91,23 +105,17 @@ pub fn requested_capability_priority_for_candidate(
|
|||||||
return (0, 0);
|
return (0, 0);
|
||||||
};
|
};
|
||||||
|
|
||||||
let mut exclusive_misses = 0u32;
|
requested_capability_priority_for_candidate_descriptors(
|
||||||
let mut compatible_misses = 0u32;
|
required_capabilities
|
||||||
for (capability, value) in required_capabilities {
|
.iter()
|
||||||
if !requested_capability_is_enabled(value) {
|
.filter_map(|(capability, value)| {
|
||||||
continue;
|
requested_capability_is_enabled(value).then_some(RequiredCapabilityDescriptor {
|
||||||
}
|
name: capability.as_str(),
|
||||||
if candidate_supports_required_capability(candidate, capability) {
|
compatible: requested_capability_is_compatible(capability),
|
||||||
continue;
|
})
|
||||||
}
|
}),
|
||||||
if requested_capability_is_compatible(capability) {
|
candidate,
|
||||||
compatible_misses += 1;
|
)
|
||||||
} else {
|
|
||||||
exclusive_misses += 1;
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
(exclusive_misses, compatible_misses)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn auth_api_key_concurrency_limit_reached(
|
pub fn auth_api_key_concurrency_limit_reached(
|
||||||
@@ -153,7 +161,8 @@ pub fn build_minimal_candidate_selection(
|
|||||||
return Ok(Vec::new());
|
return Ok(Vec::new());
|
||||||
}
|
}
|
||||||
|
|
||||||
let mut candidates = Vec::new();
|
let required_capabilities = enabled_required_capabilities(required_capabilities);
|
||||||
|
let mut candidates = Vec::with_capacity(rows.len());
|
||||||
for row in rows {
|
for row in rows {
|
||||||
if !crate::auth_constraints_allow_provider(
|
if !crate::auth_constraints_allow_provider(
|
||||||
auth_constraints,
|
auth_constraints,
|
||||||
@@ -196,16 +205,9 @@ pub fn build_minimal_candidate_selection(
|
|||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
candidates.sort_by(|left, right| {
|
let ordering_states =
|
||||||
requested_capability_priority_for_candidate(required_capabilities, left)
|
build_candidate_ordering_states(&candidates, &required_capabilities, affinity_key, None);
|
||||||
.cmp(&requested_capability_priority_for_candidate(
|
sort_candidates_by_ordering_state(&mut candidates, &ordering_states, priority_mode, false);
|
||||||
required_capabilities,
|
|
||||||
right,
|
|
||||||
))
|
|
||||||
.then_with(|| {
|
|
||||||
compare_candidates_by_priority_mode(left, right, priority_mode, affinity_key)
|
|
||||||
})
|
|
||||||
});
|
|
||||||
|
|
||||||
Ok(candidates)
|
Ok(candidates)
|
||||||
}
|
}
|
||||||
@@ -298,28 +300,27 @@ pub fn collect_selectable_candidates_from_keys(
|
|||||||
selectable_keys: &BTreeSet<(String, String, String)>,
|
selectable_keys: &BTreeSet<(String, String, String)>,
|
||||||
cached_affinity_target: Option<&crate::SchedulerAffinityTarget>,
|
cached_affinity_target: Option<&crate::SchedulerAffinityTarget>,
|
||||||
) -> Vec<SchedulerMinimalCandidateSelectionCandidate> {
|
) -> Vec<SchedulerMinimalCandidateSelectionCandidate> {
|
||||||
let mut selected = Vec::new();
|
let mut promoted = None;
|
||||||
|
let mut selected = Vec::with_capacity(candidates.len());
|
||||||
let mut emitted_keys = BTreeSet::new();
|
let mut emitted_keys = BTreeSet::new();
|
||||||
|
|
||||||
if let Some(target) = cached_affinity_target {
|
|
||||||
if let Some(candidate) = candidates
|
|
||||||
.iter()
|
|
||||||
.find(|candidate| crate::matches_affinity_target(candidate, target))
|
|
||||||
.cloned()
|
|
||||||
{
|
|
||||||
let key = crate::candidate_key(&candidate);
|
|
||||||
if selectable_keys.contains(&key) && emitted_keys.insert(key) {
|
|
||||||
selected.push(candidate);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
for candidate in candidates {
|
for candidate in candidates {
|
||||||
let key = crate::candidate_key(&candidate);
|
let key = crate::candidate_key(&candidate);
|
||||||
if !selectable_keys.contains(&key) || !emitted_keys.insert(key) {
|
if !selectable_keys.contains(&key) || !emitted_keys.insert(key) {
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
selected.push(candidate);
|
if promoted.is_none()
|
||||||
|
&& cached_affinity_target
|
||||||
|
.is_some_and(|target| crate::matches_affinity_target(&candidate, target))
|
||||||
|
{
|
||||||
|
promoted = Some(candidate);
|
||||||
|
} else {
|
||||||
|
selected.push(candidate);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if let Some(candidate) = promoted {
|
||||||
|
selected.insert(0, candidate);
|
||||||
}
|
}
|
||||||
|
|
||||||
selected
|
selected
|
||||||
@@ -332,35 +333,14 @@ pub fn reorder_candidates_by_scheduler_health(
|
|||||||
affinity_key: Option<&str>,
|
affinity_key: Option<&str>,
|
||||||
priority_mode: SchedulerPriorityMode,
|
priority_mode: SchedulerPriorityMode,
|
||||||
) {
|
) {
|
||||||
candidates.sort_by(|left, right| {
|
let required_capabilities = enabled_required_capabilities(required_capabilities);
|
||||||
requested_capability_priority_for_candidate(required_capabilities, left)
|
let ordering_states = build_candidate_ordering_states(
|
||||||
.cmp(&requested_capability_priority_for_candidate(
|
candidates,
|
||||||
required_capabilities,
|
&required_capabilities,
|
||||||
right,
|
affinity_key,
|
||||||
))
|
Some(provider_key_rpm_states),
|
||||||
.then_with(|| match priority_mode {
|
);
|
||||||
SchedulerPriorityMode::Provider => left
|
sort_candidates_by_ordering_state(candidates, &ordering_states, priority_mode, true);
|
||||||
.provider_priority
|
|
||||||
.cmp(&right.provider_priority)
|
|
||||||
.then(left.key_internal_priority.cmp(&right.key_internal_priority))
|
|
||||||
.then_with(|| {
|
|
||||||
compare_provider_key_health_order(left, right, provider_key_rpm_states)
|
|
||||||
})
|
|
||||||
.then_with(|| crate::compare_affinity_order(left, right, affinity_key))
|
|
||||||
.then_with(|| compare_candidate_identity(left, right)),
|
|
||||||
SchedulerPriorityMode::GlobalKey => left
|
|
||||||
.key_global_priority_for_format
|
|
||||||
.unwrap_or(i32::MAX)
|
|
||||||
.cmp(&right.key_global_priority_for_format.unwrap_or(i32::MAX))
|
|
||||||
.then_with(|| {
|
|
||||||
compare_provider_key_health_order(left, right, provider_key_rpm_states)
|
|
||||||
})
|
|
||||||
.then_with(|| crate::compare_affinity_order(left, right, affinity_key))
|
|
||||||
.then(left.provider_priority.cmp(&right.provider_priority))
|
|
||||||
.then(left.key_internal_priority.cmp(&right.key_internal_priority))
|
|
||||||
.then_with(|| compare_candidate_identity(left, right)),
|
|
||||||
})
|
|
||||||
});
|
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Clone, Copy, Debug)]
|
#[derive(Clone, Copy, Debug)]
|
||||||
@@ -456,33 +436,6 @@ pub fn candidate_runtime_skip_reason_with_state(
|
|||||||
None
|
None
|
||||||
}
|
}
|
||||||
|
|
||||||
fn compare_provider_key_health_order(
|
|
||||||
left: &SchedulerMinimalCandidateSelectionCandidate,
|
|
||||||
right: &SchedulerMinimalCandidateSelectionCandidate,
|
|
||||||
provider_key_rpm_states: &BTreeMap<String, StoredProviderCatalogKey>,
|
|
||||||
) -> std::cmp::Ordering {
|
|
||||||
let left_bucket = candidate_provider_key_health_bucket(left, provider_key_rpm_states);
|
|
||||||
let right_bucket = candidate_provider_key_health_bucket(right, provider_key_rpm_states);
|
|
||||||
right_bucket.cmp(&left_bucket).then_with(|| {
|
|
||||||
let left_score = candidate_provider_key_health_score(left, provider_key_rpm_states);
|
|
||||||
let right_score = candidate_provider_key_health_score(right, provider_key_rpm_states);
|
|
||||||
right_score
|
|
||||||
.partial_cmp(&left_score)
|
|
||||||
.unwrap_or(std::cmp::Ordering::Equal)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
fn candidate_provider_key_health_bucket(
|
|
||||||
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
|
||||||
provider_key_rpm_states: &BTreeMap<String, StoredProviderCatalogKey>,
|
|
||||||
) -> Option<crate::ProviderKeyHealthBucket> {
|
|
||||||
provider_key_rpm_states
|
|
||||||
.get(&candidate.key_id)
|
|
||||||
.and_then(|key| {
|
|
||||||
crate::provider_key_health_bucket(key, candidate.endpoint_api_format.as_str())
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
fn compare_candidate_identity(
|
fn compare_candidate_identity(
|
||||||
left: &SchedulerMinimalCandidateSelectionCandidate,
|
left: &SchedulerMinimalCandidateSelectionCandidate,
|
||||||
right: &SchedulerMinimalCandidateSelectionCandidate,
|
right: &SchedulerMinimalCandidateSelectionCandidate,
|
||||||
@@ -497,12 +450,190 @@ fn compare_candidate_identity(
|
|||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn enabled_required_capabilities(
|
||||||
|
required_capabilities: Option<&serde_json::Value>,
|
||||||
|
) -> Vec<RequiredCapabilityDescriptor<'_>> {
|
||||||
|
let Some(required_capabilities) = required_capabilities.and_then(serde_json::Value::as_object)
|
||||||
|
else {
|
||||||
|
return Vec::new();
|
||||||
|
};
|
||||||
|
|
||||||
|
required_capabilities
|
||||||
|
.iter()
|
||||||
|
.filter_map(|(capability, value)| {
|
||||||
|
requested_capability_is_enabled(value).then_some(RequiredCapabilityDescriptor {
|
||||||
|
name: capability.as_str(),
|
||||||
|
compatible: requested_capability_is_compatible(capability),
|
||||||
|
})
|
||||||
|
})
|
||||||
|
.collect()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn requested_capability_priority_for_candidate_descriptors<'a, I>(
|
||||||
|
required_capabilities: I,
|
||||||
|
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
||||||
|
) -> (u32, u32)
|
||||||
|
where
|
||||||
|
I: IntoIterator<Item = RequiredCapabilityDescriptor<'a>>,
|
||||||
|
{
|
||||||
|
let mut exclusive_misses = 0u32;
|
||||||
|
let mut compatible_misses = 0u32;
|
||||||
|
for capability in required_capabilities {
|
||||||
|
if candidate_supports_required_capability(candidate, capability.name) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
if capability.compatible {
|
||||||
|
compatible_misses += 1;
|
||||||
|
} else {
|
||||||
|
exclusive_misses += 1;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
(exclusive_misses, compatible_misses)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn build_candidate_ordering_states(
|
||||||
|
candidates: &[SchedulerMinimalCandidateSelectionCandidate],
|
||||||
|
required_capabilities: &[RequiredCapabilityDescriptor<'_>],
|
||||||
|
affinity_key: Option<&str>,
|
||||||
|
provider_key_rpm_states: Option<&BTreeMap<String, StoredProviderCatalogKey>>,
|
||||||
|
) -> Vec<CandidateOrderingState> {
|
||||||
|
candidates
|
||||||
|
.iter()
|
||||||
|
.map(|candidate| CandidateOrderingState {
|
||||||
|
capability_priority: requested_capability_priority_for_candidate_descriptors(
|
||||||
|
required_capabilities.iter().copied(),
|
||||||
|
candidate,
|
||||||
|
),
|
||||||
|
affinity_hash: affinity_key.map(|key| crate::candidate_affinity_hash(key, candidate)),
|
||||||
|
health_bucket: provider_key_rpm_states.and_then(|states| {
|
||||||
|
states.get(&candidate.key_id).and_then(|key| {
|
||||||
|
crate::provider_key_health_bucket(key, candidate.endpoint_api_format.as_str())
|
||||||
|
})
|
||||||
|
}),
|
||||||
|
health_score: candidate_provider_key_health_score(candidate, provider_key_rpm_states),
|
||||||
|
})
|
||||||
|
.collect()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn sort_candidates_by_ordering_state(
|
||||||
|
candidates: &mut [SchedulerMinimalCandidateSelectionCandidate],
|
||||||
|
ordering_states: &[CandidateOrderingState],
|
||||||
|
priority_mode: SchedulerPriorityMode,
|
||||||
|
include_health: bool,
|
||||||
|
) {
|
||||||
|
if candidates.len() < 2 {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut order = (0..candidates.len()).collect::<Vec<_>>();
|
||||||
|
order.sort_by(|left, right| {
|
||||||
|
compare_candidates_with_ordering_state(
|
||||||
|
&ordering_states[*left],
|
||||||
|
&candidates[*left],
|
||||||
|
&ordering_states[*right],
|
||||||
|
&candidates[*right],
|
||||||
|
priority_mode,
|
||||||
|
include_health,
|
||||||
|
)
|
||||||
|
});
|
||||||
|
apply_candidate_order(candidates, order);
|
||||||
|
}
|
||||||
|
|
||||||
|
fn compare_candidates_with_ordering_state(
|
||||||
|
left_state: &CandidateOrderingState,
|
||||||
|
left_candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
||||||
|
right_state: &CandidateOrderingState,
|
||||||
|
right_candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
||||||
|
priority_mode: SchedulerPriorityMode,
|
||||||
|
include_health: bool,
|
||||||
|
) -> std::cmp::Ordering {
|
||||||
|
left_state
|
||||||
|
.capability_priority
|
||||||
|
.cmp(&right_state.capability_priority)
|
||||||
|
.then_with(|| {
|
||||||
|
compare_priority_before_health(left_candidate, right_candidate, priority_mode)
|
||||||
|
})
|
||||||
|
.then_with(|| {
|
||||||
|
if include_health {
|
||||||
|
compare_provider_key_health_state(left_state, right_state)
|
||||||
|
} else {
|
||||||
|
std::cmp::Ordering::Equal
|
||||||
|
}
|
||||||
|
})
|
||||||
|
.then_with(|| left_state.affinity_hash.cmp(&right_state.affinity_hash))
|
||||||
|
.then_with(|| {
|
||||||
|
compare_priority_after_affinity(left_candidate, right_candidate, priority_mode)
|
||||||
|
})
|
||||||
|
.then_with(|| compare_candidate_identity(left_candidate, right_candidate))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn compare_priority_before_health(
|
||||||
|
left: &SchedulerMinimalCandidateSelectionCandidate,
|
||||||
|
right: &SchedulerMinimalCandidateSelectionCandidate,
|
||||||
|
priority_mode: SchedulerPriorityMode,
|
||||||
|
) -> std::cmp::Ordering {
|
||||||
|
match priority_mode {
|
||||||
|
SchedulerPriorityMode::Provider => left
|
||||||
|
.provider_priority
|
||||||
|
.cmp(&right.provider_priority)
|
||||||
|
.then(left.key_internal_priority.cmp(&right.key_internal_priority)),
|
||||||
|
SchedulerPriorityMode::GlobalKey => left
|
||||||
|
.key_global_priority_for_format
|
||||||
|
.unwrap_or(i32::MAX)
|
||||||
|
.cmp(&right.key_global_priority_for_format.unwrap_or(i32::MAX)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn compare_priority_after_affinity(
|
||||||
|
left: &SchedulerMinimalCandidateSelectionCandidate,
|
||||||
|
right: &SchedulerMinimalCandidateSelectionCandidate,
|
||||||
|
priority_mode: SchedulerPriorityMode,
|
||||||
|
) -> std::cmp::Ordering {
|
||||||
|
match priority_mode {
|
||||||
|
SchedulerPriorityMode::Provider => std::cmp::Ordering::Equal,
|
||||||
|
SchedulerPriorityMode::GlobalKey => left
|
||||||
|
.provider_priority
|
||||||
|
.cmp(&right.provider_priority)
|
||||||
|
.then(left.key_internal_priority.cmp(&right.key_internal_priority)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn compare_provider_key_health_state(
|
||||||
|
left: &CandidateOrderingState,
|
||||||
|
right: &CandidateOrderingState,
|
||||||
|
) -> std::cmp::Ordering {
|
||||||
|
right
|
||||||
|
.health_bucket
|
||||||
|
.cmp(&left.health_bucket)
|
||||||
|
.then_with(|| right.health_score.total_cmp(&left.health_score))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn apply_candidate_order(
|
||||||
|
candidates: &mut [SchedulerMinimalCandidateSelectionCandidate],
|
||||||
|
sorted_old_indices: Vec<usize>,
|
||||||
|
) {
|
||||||
|
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..candidates.len() {
|
||||||
|
let current = index;
|
||||||
|
while target_positions[current] != current {
|
||||||
|
let target = target_positions[current];
|
||||||
|
candidates.swap(current, target);
|
||||||
|
target_positions.swap(current, target);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
fn candidate_provider_key_health_score(
|
fn candidate_provider_key_health_score(
|
||||||
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
||||||
provider_key_rpm_states: &BTreeMap<String, StoredProviderCatalogKey>,
|
provider_key_rpm_states: Option<&BTreeMap<String, StoredProviderCatalogKey>>,
|
||||||
) -> f64 {
|
) -> f64 {
|
||||||
provider_key_rpm_states
|
provider_key_rpm_states
|
||||||
.get(&candidate.key_id)
|
.and_then(|states| states.get(&candidate.key_id))
|
||||||
.and_then(|key| {
|
.and_then(|key| {
|
||||||
crate::effective_provider_key_health_score(key, candidate.endpoint_api_format.as_str())
|
crate::effective_provider_key_health_score(key, candidate.endpoint_api_format.as_str())
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -44,15 +44,17 @@ fn resolve_global_model_name_by<F>(
|
|||||||
where
|
where
|
||||||
F: Fn(&StoredMinimalCandidateSelectionRow) -> bool,
|
F: Fn(&StoredMinimalCandidateSelectionRow) -> bool,
|
||||||
{
|
{
|
||||||
let mut matches = rows
|
let mut best_match = None::<&str>;
|
||||||
.iter()
|
for row in rows.iter().filter(|row| matches(row)) {
|
||||||
.filter(|row| matches(row))
|
let candidate = row.global_model_name.trim();
|
||||||
.map(|row| row.global_model_name.trim())
|
if candidate.is_empty() {
|
||||||
.filter(|value| !value.is_empty())
|
continue;
|
||||||
.map(ToOwned::to_owned)
|
}
|
||||||
.collect::<BTreeSet<_>>()
|
if best_match.is_none_or(|current| candidate < current) {
|
||||||
.into_iter();
|
best_match = Some(candidate);
|
||||||
matches.next()
|
}
|
||||||
|
}
|
||||||
|
best_match.map(ToOwned::to_owned)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn resolve_provider_model_name(
|
pub fn resolve_provider_model_name(
|
||||||
@@ -75,25 +77,26 @@ pub fn resolve_provider_model_name(
|
|||||||
return Some((selected_provider_model_name, None));
|
return Some((selected_provider_model_name, None));
|
||||||
}
|
}
|
||||||
|
|
||||||
let candidate_models = candidate_model_names(row, api_format);
|
|
||||||
let mut sorted_allowed_models = key_allowed_models
|
let mut sorted_allowed_models = key_allowed_models
|
||||||
.iter()
|
.iter()
|
||||||
.map(|value| value.trim())
|
.map(String::as_str)
|
||||||
|
.map(str::trim)
|
||||||
.filter(|value| !value.is_empty())
|
.filter(|value| !value.is_empty())
|
||||||
.map(ToOwned::to_owned)
|
|
||||||
.collect::<Vec<_>>();
|
.collect::<Vec<_>>();
|
||||||
sorted_allowed_models.sort();
|
sorted_allowed_models.sort_unstable();
|
||||||
|
|
||||||
for allowed_model in &sorted_allowed_models {
|
for &allowed_model in &sorted_allowed_models {
|
||||||
if candidate_models.contains(allowed_model.as_str()) {
|
if row_has_candidate_model_name(row, api_format, allowed_model) {
|
||||||
return Some((allowed_model.clone(), Some(allowed_model.clone())));
|
let allowed_model = allowed_model.to_owned();
|
||||||
|
return Some((allowed_model.clone(), Some(allowed_model)));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
let global_model_mappings = row.global_model_mappings.as_ref()?;
|
let global_model_mappings = row.global_model_mappings.as_ref()?;
|
||||||
for allowed_model in sorted_allowed_models {
|
for &allowed_model in &sorted_allowed_models {
|
||||||
for pattern in global_model_mappings {
|
for pattern in global_model_mappings {
|
||||||
if matches_model_mapping(pattern, &allowed_model) {
|
if matches_model_mapping(pattern, allowed_model) {
|
||||||
|
let allowed_model = allowed_model.to_owned();
|
||||||
return Some((allowed_model.clone(), Some(allowed_model)));
|
return Some((allowed_model.clone(), Some(allowed_model)));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -110,23 +113,14 @@ pub fn select_provider_model_name(
|
|||||||
return row.model_provider_model_name.clone();
|
return row.model_provider_model_name.clone();
|
||||||
};
|
};
|
||||||
|
|
||||||
let mut scoped = mappings
|
mappings
|
||||||
.iter()
|
.iter()
|
||||||
.filter(|mapping| mapping_scope_matches(mapping, api_format))
|
.filter(|mapping| mapping_scope_matches(mapping, api_format))
|
||||||
.collect::<Vec<_>>();
|
.min_by(|left, right| {
|
||||||
if scoped.is_empty() {
|
left.priority
|
||||||
return row.model_provider_model_name.clone();
|
.cmp(&right.priority)
|
||||||
}
|
.then(left.name.cmp(&right.name))
|
||||||
|
})
|
||||||
scoped.sort_by(|left, right| {
|
|
||||||
left.priority
|
|
||||||
.cmp(&right.priority)
|
|
||||||
.then(left.name.cmp(&right.name))
|
|
||||||
});
|
|
||||||
let top_priority = scoped[0].priority;
|
|
||||||
scoped
|
|
||||||
.into_iter()
|
|
||||||
.find(|mapping| mapping.priority == top_priority)
|
|
||||||
.map(|mapping| mapping.name.clone())
|
.map(|mapping| mapping.name.clone())
|
||||||
.unwrap_or_else(|| row.model_provider_model_name.clone())
|
.unwrap_or_else(|| row.model_provider_model_name.clone())
|
||||||
}
|
}
|
||||||
@@ -153,7 +147,7 @@ fn mapping_scope_matches(mapping: &StoredProviderModelMapping, api_format: &str)
|
|||||||
|
|
||||||
api_formats
|
api_formats
|
||||||
.iter()
|
.iter()
|
||||||
.any(|value| normalize_api_format(value) == api_format)
|
.any(|value| api_format_matches(value, api_format))
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn row_supports_required_capability(
|
pub fn row_supports_required_capability(
|
||||||
@@ -230,7 +224,7 @@ pub fn extract_global_priority_for_format(
|
|||||||
|
|
||||||
let Some(value) = object
|
let Some(value) = object
|
||||||
.iter()
|
.iter()
|
||||||
.find(|(key, _)| normalize_api_format(key) == api_format)
|
.find(|(key, _)| api_format_matches(key, api_format))
|
||||||
.map(|(_, value)| value)
|
.map(|(_, value)| value)
|
||||||
else {
|
else {
|
||||||
return Ok(None);
|
return Ok(None);
|
||||||
@@ -262,6 +256,26 @@ pub fn normalize_api_format(value: &str) -> String {
|
|||||||
value.trim().to_ascii_lowercase()
|
value.trim().to_ascii_lowercase()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn row_has_candidate_model_name(
|
||||||
|
row: &StoredMinimalCandidateSelectionRow,
|
||||||
|
api_format: &str,
|
||||||
|
model_name: &str,
|
||||||
|
) -> bool {
|
||||||
|
row.model_provider_model_name == model_name
|
||||||
|
|| row
|
||||||
|
.model_provider_model_mappings
|
||||||
|
.as_ref()
|
||||||
|
.is_some_and(|mappings| {
|
||||||
|
mappings.iter().any(|mapping| {
|
||||||
|
mapping_scope_matches(mapping, api_format) && mapping.name == model_name
|
||||||
|
})
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
fn api_format_matches(left: &str, right: &str) -> bool {
|
||||||
|
left.trim().eq_ignore_ascii_case(right.trim())
|
||||||
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::matches_model_mapping;
|
use super::matches_model_mapping;
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ use aether_contracts::{ExecutionError, ExecutionPlan};
|
|||||||
use aether_data_contracts::repository::candidates::{
|
use aether_data_contracts::repository::candidates::{
|
||||||
RequestCandidateStatus, StoredRequestCandidate, UpsertRequestCandidateRecord,
|
RequestCandidateStatus, StoredRequestCandidate, UpsertRequestCandidateRecord,
|
||||||
};
|
};
|
||||||
use serde_json::Value;
|
use serde_json::{Map, Value};
|
||||||
|
|
||||||
#[derive(Debug, Clone, PartialEq)]
|
#[derive(Debug, Clone, PartialEq)]
|
||||||
pub struct SchedulerRequestCandidateReportContext {
|
pub struct SchedulerRequestCandidateReportContext {
|
||||||
@@ -91,20 +91,13 @@ pub fn parse_request_candidate_report_context(
|
|||||||
report_context: Option<&Value>,
|
report_context: Option<&Value>,
|
||||||
) -> Option<SchedulerRequestCandidateReportContext> {
|
) -> Option<SchedulerRequestCandidateReportContext> {
|
||||||
let report_context = report_context?;
|
let report_context = report_context?;
|
||||||
let retry_index = report_context
|
|
||||||
.get("retry_index")
|
|
||||||
.and_then(Value::as_u64)
|
|
||||||
.unwrap_or_default();
|
|
||||||
Some(SchedulerRequestCandidateReportContext {
|
Some(SchedulerRequestCandidateReportContext {
|
||||||
request_id: string_field(report_context, "request_id"),
|
request_id: string_field(report_context, "request_id"),
|
||||||
candidate_id: string_field(report_context, "candidate_id"),
|
candidate_id: string_field(report_context, "candidate_id"),
|
||||||
user_id: string_field(report_context, "user_id"),
|
user_id: string_field(report_context, "user_id"),
|
||||||
api_key_id: string_field(report_context, "api_key_id"),
|
api_key_id: string_field(report_context, "api_key_id"),
|
||||||
candidate_index: report_context
|
candidate_index: u32_field(report_context, "candidate_index"),
|
||||||
.get("candidate_index")
|
retry_index: u32_field(report_context, "retry_index").unwrap_or_default(),
|
||||||
.and_then(Value::as_u64)
|
|
||||||
.and_then(|value| u32::try_from(value).ok()),
|
|
||||||
retry_index: u32::try_from(retry_index).unwrap_or(u32::MAX),
|
|
||||||
provider_id: string_field(report_context, "provider_id"),
|
provider_id: string_field(report_context, "provider_id"),
|
||||||
endpoint_id: string_field(report_context, "endpoint_id"),
|
endpoint_id: string_field(report_context, "endpoint_id"),
|
||||||
key_id: string_field(report_context, "key_id"),
|
key_id: string_field(report_context, "key_id"),
|
||||||
@@ -123,9 +116,24 @@ pub fn resolve_report_request_candidate_slot(
|
|||||||
now_unix_ms: u64,
|
now_unix_ms: u64,
|
||||||
generated_candidate_id: String,
|
generated_candidate_id: String,
|
||||||
) -> Option<SchedulerResolvedReportRequestCandidateSlot> {
|
) -> Option<SchedulerResolvedReportRequestCandidateSlot> {
|
||||||
let request_id = metadata.request_id.clone()?;
|
|
||||||
let matched_candidate = match_existing_report_candidate(existing_candidates, &metadata);
|
let matched_candidate = match_existing_report_candidate(existing_candidates, &metadata);
|
||||||
let synthesized_extra_data = build_report_candidate_extra_data(&metadata);
|
let SchedulerRequestCandidateReportContext {
|
||||||
|
request_id,
|
||||||
|
candidate_id,
|
||||||
|
user_id,
|
||||||
|
api_key_id,
|
||||||
|
candidate_index: metadata_candidate_index,
|
||||||
|
retry_index,
|
||||||
|
provider_id,
|
||||||
|
endpoint_id,
|
||||||
|
key_id,
|
||||||
|
client_api_format,
|
||||||
|
provider_api_format,
|
||||||
|
proxy,
|
||||||
|
} = metadata;
|
||||||
|
let request_id = request_id?;
|
||||||
|
let synthesized_extra_data =
|
||||||
|
build_report_candidate_extra_data(client_api_format, provider_api_format, proxy);
|
||||||
let created_at_unix_ms = matched_candidate
|
let created_at_unix_ms = matched_candidate
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.map(|candidate| candidate.created_at_unix_ms)
|
.map(|candidate| candidate.created_at_unix_ms)
|
||||||
@@ -133,42 +141,42 @@ pub fn resolve_report_request_candidate_slot(
|
|||||||
let candidate_index = matched_candidate
|
let candidate_index = matched_candidate
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.map(|candidate| candidate.candidate_index)
|
.map(|candidate| candidate.candidate_index)
|
||||||
.or(metadata.candidate_index)
|
.or(metadata_candidate_index)
|
||||||
.unwrap_or_else(|| next_candidate_index(existing_candidates));
|
.unwrap_or_else(|| next_candidate_index(existing_candidates));
|
||||||
let retry_index = matched_candidate
|
let retry_index = matched_candidate
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.map(|candidate| candidate.retry_index)
|
.map(|candidate| candidate.retry_index)
|
||||||
.unwrap_or(metadata.retry_index);
|
.unwrap_or(retry_index);
|
||||||
|
|
||||||
Some(SchedulerResolvedReportRequestCandidateSlot {
|
Some(SchedulerResolvedReportRequestCandidateSlot {
|
||||||
id: matched_candidate
|
id: matched_candidate
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.map(|candidate| candidate.id.clone())
|
.map(|candidate| candidate.id.clone())
|
||||||
.or(metadata.candidate_id)
|
.or(candidate_id)
|
||||||
.unwrap_or(generated_candidate_id),
|
.unwrap_or(generated_candidate_id),
|
||||||
request_id,
|
request_id,
|
||||||
user_id: matched_candidate
|
user_id: matched_candidate
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.and_then(|candidate| candidate.user_id.clone())
|
.and_then(|candidate| candidate.user_id.clone())
|
||||||
.or(metadata.user_id),
|
.or(user_id),
|
||||||
api_key_id: matched_candidate
|
api_key_id: matched_candidate
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.and_then(|candidate| candidate.api_key_id.clone())
|
.and_then(|candidate| candidate.api_key_id.clone())
|
||||||
.or(metadata.api_key_id),
|
.or(api_key_id),
|
||||||
candidate_index,
|
candidate_index,
|
||||||
retry_index,
|
retry_index,
|
||||||
provider_id: matched_candidate
|
provider_id: matched_candidate
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.and_then(|candidate| candidate.provider_id.clone())
|
.and_then(|candidate| candidate.provider_id.clone())
|
||||||
.or(metadata.provider_id),
|
.or(provider_id),
|
||||||
endpoint_id: matched_candidate
|
endpoint_id: matched_candidate
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.and_then(|candidate| candidate.endpoint_id.clone())
|
.and_then(|candidate| candidate.endpoint_id.clone())
|
||||||
.or(metadata.endpoint_id),
|
.or(endpoint_id),
|
||||||
key_id: matched_candidate
|
key_id: matched_candidate
|
||||||
.as_ref()
|
.as_ref()
|
||||||
.and_then(|candidate| candidate.key_id.clone())
|
.and_then(|candidate| candidate.key_id.clone())
|
||||||
.or(metadata.key_id),
|
.or(key_id),
|
||||||
extra_data: merge_request_candidate_extra_data(
|
extra_data: merge_request_candidate_extra_data(
|
||||||
matched_candidate
|
matched_candidate
|
||||||
.as_ref()
|
.as_ref()
|
||||||
@@ -195,22 +203,14 @@ pub fn build_execution_request_candidate_seed(
|
|||||||
.and_then(Value::as_object)
|
.and_then(Value::as_object)
|
||||||
.cloned()
|
.cloned()
|
||||||
.unwrap_or_default();
|
.unwrap_or_default();
|
||||||
let request_id = string_field(&Value::Object(context.clone()), "request_id")
|
let request_id =
|
||||||
.unwrap_or_else(|| plan.request_id.clone());
|
string_field_from_object(&context, "request_id").unwrap_or_else(|| plan.request_id.clone());
|
||||||
let candidate_index = context
|
let candidate_index = u32_field_from_object(&context, "candidate_index").unwrap_or(0);
|
||||||
.get("candidate_index")
|
let retry_index = u32_field_from_object(&context, "retry_index").unwrap_or(0);
|
||||||
.and_then(Value::as_u64)
|
let candidate_id =
|
||||||
.and_then(|value| u32::try_from(value).ok())
|
string_field_from_object(&context, "candidate_id").unwrap_or(generated_candidate_id);
|
||||||
.unwrap_or(0);
|
let user_id = string_field_from_object(&context, "user_id");
|
||||||
let retry_index = context
|
let api_key_id = string_field_from_object(&context, "api_key_id");
|
||||||
.get("retry_index")
|
|
||||||
.and_then(Value::as_u64)
|
|
||||||
.and_then(|value| u32::try_from(value).ok())
|
|
||||||
.unwrap_or(0);
|
|
||||||
let candidate_id = string_field(&Value::Object(context.clone()), "candidate_id")
|
|
||||||
.unwrap_or(generated_candidate_id);
|
|
||||||
let user_id = string_field(&Value::Object(context.clone()), "user_id");
|
|
||||||
let api_key_id = string_field(&Value::Object(context.clone()), "api_key_id");
|
|
||||||
|
|
||||||
context.insert("request_id".to_string(), Value::String(request_id.clone()));
|
context.insert("request_id".to_string(), Value::String(request_id.clone()));
|
||||||
context.insert(
|
context.insert(
|
||||||
@@ -374,7 +374,10 @@ pub fn finalize_execution_request_candidate_report_context(
|
|||||||
report_context: Value,
|
report_context: Value,
|
||||||
candidate_id: &str,
|
candidate_id: &str,
|
||||||
) -> Value {
|
) -> Value {
|
||||||
let mut context = report_context.as_object().cloned().unwrap_or_default();
|
let mut context = match report_context {
|
||||||
|
Value::Object(context) => context,
|
||||||
|
_ => Map::new(),
|
||||||
|
};
|
||||||
let candidate_id = candidate_id.trim();
|
let candidate_id = candidate_id.trim();
|
||||||
if !candidate_id.is_empty() {
|
if !candidate_id.is_empty() {
|
||||||
context.insert(
|
context.insert(
|
||||||
@@ -410,6 +413,12 @@ fn extract_error_message(body_json: &Value) -> Option<&str> {
|
|||||||
|
|
||||||
fn string_field(value: &Value, key: &str) -> Option<String> {
|
fn string_field(value: &Value, key: &str) -> Option<String> {
|
||||||
value
|
value
|
||||||
|
.as_object()
|
||||||
|
.and_then(|object| string_field_from_object(object, key))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn string_field_from_object(object: &Map<String, Value>, key: &str) -> Option<String> {
|
||||||
|
object
|
||||||
.get(key)
|
.get(key)
|
||||||
.and_then(Value::as_str)
|
.and_then(Value::as_str)
|
||||||
.map(str::trim)
|
.map(str::trim)
|
||||||
@@ -417,6 +426,19 @@ fn string_field(value: &Value, key: &str) -> Option<String> {
|
|||||||
.map(ToOwned::to_owned)
|
.map(ToOwned::to_owned)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn u32_field(value: &Value, key: &str) -> Option<u32> {
|
||||||
|
value
|
||||||
|
.as_object()
|
||||||
|
.and_then(|object| u32_field_from_object(object, key))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn u32_field_from_object(object: &Map<String, Value>, key: &str) -> Option<u32> {
|
||||||
|
object
|
||||||
|
.get(key)
|
||||||
|
.and_then(Value::as_u64)
|
||||||
|
.and_then(|value| u32::try_from(value).ok())
|
||||||
|
}
|
||||||
|
|
||||||
fn match_existing_report_candidate<'a>(
|
fn match_existing_report_candidate<'a>(
|
||||||
candidates: &'a [StoredRequestCandidate],
|
candidates: &'a [StoredRequestCandidate],
|
||||||
metadata: &SchedulerRequestCandidateReportContext,
|
metadata: &SchedulerRequestCandidateReportContext,
|
||||||
@@ -465,24 +487,26 @@ fn next_candidate_index(candidates: &[StoredRequestCandidate]) -> u32 {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn build_report_candidate_extra_data(
|
fn build_report_candidate_extra_data(
|
||||||
metadata: &SchedulerRequestCandidateReportContext,
|
client_api_format: Option<String>,
|
||||||
|
provider_api_format: Option<String>,
|
||||||
|
proxy: Option<Value>,
|
||||||
) -> Option<Value> {
|
) -> Option<Value> {
|
||||||
let mut extra_data = serde_json::Map::new();
|
let mut extra_data = Map::with_capacity(5);
|
||||||
extra_data.insert("gateway_execution_runtime".to_string(), Value::Bool(true));
|
extra_data.insert("gateway_execution_runtime".to_string(), Value::Bool(true));
|
||||||
extra_data.insert("phase".to_string(), Value::String("3c_trial".to_string()));
|
extra_data.insert("phase".to_string(), Value::String("3c_trial".to_string()));
|
||||||
if let Some(client_api_format) = metadata.client_api_format.clone() {
|
if let Some(client_api_format) = client_api_format {
|
||||||
extra_data.insert(
|
extra_data.insert(
|
||||||
"client_api_format".to_string(),
|
"client_api_format".to_string(),
|
||||||
Value::String(client_api_format),
|
Value::String(client_api_format),
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
if let Some(provider_api_format) = metadata.provider_api_format.clone() {
|
if let Some(provider_api_format) = provider_api_format {
|
||||||
extra_data.insert(
|
extra_data.insert(
|
||||||
"provider_api_format".to_string(),
|
"provider_api_format".to_string(),
|
||||||
Value::String(provider_api_format),
|
Value::String(provider_api_format),
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
if let Some(proxy) = metadata.proxy.clone() {
|
if let Some(proxy) = proxy {
|
||||||
extra_data.insert("proxy".to_string(), proxy);
|
extra_data.insert("proxy".to_string(), proxy);
|
||||||
}
|
}
|
||||||
(!extra_data.is_empty()).then_some(Value::Object(extra_data))
|
(!extra_data.is_empty()).then_some(Value::Object(extra_data))
|
||||||
|
|||||||
@@ -1,6 +1,9 @@
|
|||||||
|
use std::io::{self, Write};
|
||||||
|
|
||||||
use aether_data_contracts::repository::usage::{
|
use aether_data_contracts::repository::usage::{
|
||||||
UpsertUsageRecord, UsageBodyCaptureState, UsageBodyField,
|
UpsertUsageRecord, UsageBodyCaptureState, UsageBodyField,
|
||||||
};
|
};
|
||||||
|
use serde::Serialize;
|
||||||
use serde_json::{json, Map, Value};
|
use serde_json::{json, Map, Value};
|
||||||
|
|
||||||
use crate::event::UsageEvent;
|
use crate::event::UsageEvent;
|
||||||
@@ -76,6 +79,22 @@ pub struct UsageBodyCaptureEngine {
|
|||||||
policy: UsageBodyCapturePolicy,
|
policy: UsageBodyCapturePolicy,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(Default)]
|
||||||
|
struct CountingWriter {
|
||||||
|
bytes: u64,
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Write for CountingWriter {
|
||||||
|
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
|
||||||
|
self.bytes = self.bytes.saturating_add(buf.len() as u64);
|
||||||
|
Ok(buf.len())
|
||||||
|
}
|
||||||
|
|
||||||
|
fn flush(&mut self) -> io::Result<()> {
|
||||||
|
Ok(())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||||
pub(crate) struct RuntimeBodyCaptureStates {
|
pub(crate) struct RuntimeBodyCaptureStates {
|
||||||
pub request: UsageBodyCaptureState,
|
pub request: UsageBodyCaptureState,
|
||||||
@@ -286,9 +305,7 @@ fn limit_usage_body_capture_value(
|
|||||||
value: Value,
|
value: Value,
|
||||||
max_bytes: Option<usize>,
|
max_bytes: Option<usize>,
|
||||||
) -> LimitedUsageBodyCapture {
|
) -> LimitedUsageBodyCapture {
|
||||||
let source_bytes = serde_json::to_vec(&value)
|
let source_bytes = json_serialized_len(&value);
|
||||||
.ok()
|
|
||||||
.map(|bytes| bytes.len() as u64);
|
|
||||||
let Some(limit) = max_bytes.filter(|value| *value > 0) else {
|
let Some(limit) = max_bytes.filter(|value| *value > 0) else {
|
||||||
return LimitedUsageBodyCapture {
|
return LimitedUsageBodyCapture {
|
||||||
stored_bytes: source_bytes,
|
stored_bytes: source_bytes,
|
||||||
@@ -327,9 +344,7 @@ fn limit_usage_body_capture_value(
|
|||||||
"value_kind": usage_value_kind(&other),
|
"value_kind": usage_value_kind(&other),
|
||||||
}),
|
}),
|
||||||
};
|
};
|
||||||
let stored_bytes = serde_json::to_vec(&truncated_value)
|
let stored_bytes = json_serialized_len(&truncated_value);
|
||||||
.ok()
|
|
||||||
.map(|bytes| bytes.len() as u64);
|
|
||||||
LimitedUsageBodyCapture {
|
LimitedUsageBodyCapture {
|
||||||
value: truncated_value,
|
value: truncated_value,
|
||||||
source_bytes: Some(source_len),
|
source_bytes: Some(source_len),
|
||||||
@@ -340,13 +355,6 @@ fn limit_usage_body_capture_value(
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn truncate_usage_body_string(value: &str, max_bytes: usize) -> String {
|
fn truncate_usage_body_string(value: &str, max_bytes: usize) -> String {
|
||||||
if serde_json::to_vec(&Value::String(value.to_string()))
|
|
||||||
.ok()
|
|
||||||
.is_some_and(|bytes| bytes.len() <= max_bytes)
|
|
||||||
{
|
|
||||||
return value.to_string();
|
|
||||||
}
|
|
||||||
|
|
||||||
let mut end = value.len();
|
let mut end = value.len();
|
||||||
while end > 0 {
|
while end > 0 {
|
||||||
while end > 0 && !value.is_char_boundary(end) {
|
while end > 0 && !value.is_char_boundary(end) {
|
||||||
@@ -354,10 +362,7 @@ fn truncate_usage_body_string(value: &str, max_bytes: usize) -> String {
|
|||||||
}
|
}
|
||||||
let mut candidate = value[..end].to_string();
|
let mut candidate = value[..end].to_string();
|
||||||
candidate.push_str(TRUNCATED_BODY_STRING_SUFFIX);
|
candidate.push_str(TRUNCATED_BODY_STRING_SUFFIX);
|
||||||
if serde_json::to_vec(&Value::String(candidate.clone()))
|
if json_serialized_len(&candidate).is_some_and(|bytes| bytes <= max_bytes as u64) {
|
||||||
.ok()
|
|
||||||
.is_some_and(|bytes| bytes.len() <= max_bytes)
|
|
||||||
{
|
|
||||||
return candidate;
|
return candidate;
|
||||||
}
|
}
|
||||||
end = value[..end]
|
end = value[..end]
|
||||||
@@ -379,27 +384,45 @@ fn truncate_usage_body_string(value: &str, max_bytes: usize) -> String {
|
|||||||
.to_string()
|
.to_string()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn json_serialized_len<T: Serialize>(value: &T) -> Option<u64> {
|
||||||
|
let mut writer = CountingWriter::default();
|
||||||
|
serde_json::to_writer(&mut writer, value).ok()?;
|
||||||
|
Some(writer.bytes)
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) fn sync_usage_body_ref_metadata(
|
pub(crate) fn sync_usage_body_ref_metadata(
|
||||||
metadata: &mut Option<Value>,
|
metadata: &mut Option<Value>,
|
||||||
field: UsageBodyField,
|
field: UsageBodyField,
|
||||||
body_ref: Option<&str>,
|
body_ref: Option<&str>,
|
||||||
) {
|
) {
|
||||||
|
let key = field.as_ref_key();
|
||||||
let Some(body_ref) = body_ref.map(str::trim).filter(|value| !value.is_empty()) else {
|
let Some(body_ref) = body_ref.map(str::trim).filter(|value| !value.is_empty()) else {
|
||||||
if let Some(object) = metadata.as_mut().and_then(Value::as_object_mut) {
|
let clear_metadata = match metadata.as_mut() {
|
||||||
object.remove(field.as_ref_key());
|
Some(Value::Object(object)) => {
|
||||||
|
object.remove(key);
|
||||||
|
object.is_empty()
|
||||||
|
}
|
||||||
|
_ => false,
|
||||||
|
};
|
||||||
|
if clear_metadata {
|
||||||
|
*metadata = None;
|
||||||
}
|
}
|
||||||
return;
|
return;
|
||||||
};
|
};
|
||||||
|
if let Some(Value::Object(object)) = metadata.as_mut() {
|
||||||
|
if object.get(key).and_then(Value::as_str) == Some(body_ref) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
object.insert(key.to_owned(), Value::String(body_ref.to_owned()));
|
||||||
|
return;
|
||||||
|
}
|
||||||
let object = metadata
|
let object = metadata
|
||||||
.get_or_insert_with(|| Value::Object(Map::new()))
|
.get_or_insert_with(|| Value::Object(Map::new()))
|
||||||
.as_object_mut();
|
.as_object_mut();
|
||||||
let Some(object) = object else {
|
let Some(object) = object else {
|
||||||
return;
|
return;
|
||||||
};
|
};
|
||||||
object.insert(
|
object.insert(key.to_owned(), Value::String(body_ref.to_owned()));
|
||||||
field.as_ref_key().to_string(),
|
|
||||||
Value::String(body_ref.to_string()),
|
|
||||||
);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn build_payload_body_capture_metadata(
|
pub(crate) fn build_payload_body_capture_metadata(
|
||||||
@@ -408,36 +431,44 @@ pub(crate) fn build_payload_body_capture_metadata(
|
|||||||
provider_body_state: Option<UsageBodyCaptureState>,
|
provider_body_state: Option<UsageBodyCaptureState>,
|
||||||
client_body_state: Option<UsageBodyCaptureState>,
|
client_body_state: Option<UsageBodyCaptureState>,
|
||||||
) -> Option<Value> {
|
) -> Option<Value> {
|
||||||
let mut metadata = Map::new();
|
let provider_decoded_len = provider_body_base64.and_then(decoded_base64_len_hint);
|
||||||
if let Some(decoded_len) = provider_body_base64.and_then(decoded_base64_len_hint) {
|
let client_decoded_len = client_body_base64.and_then(decoded_base64_len_hint);
|
||||||
|
let body_capture_capacity =
|
||||||
|
usize::from(provider_body_state.is_some()) + usize::from(client_body_state.is_some());
|
||||||
|
let mut metadata = Map::with_capacity(
|
||||||
|
usize::from(provider_decoded_len.is_some())
|
||||||
|
+ usize::from(client_decoded_len.is_some())
|
||||||
|
+ usize::from(body_capture_capacity > 0),
|
||||||
|
);
|
||||||
|
if let Some(decoded_len) = provider_decoded_len {
|
||||||
metadata.insert(
|
metadata.insert(
|
||||||
"provider_response_body_base64_bytes".to_string(),
|
"provider_response_body_base64_bytes".to_string(),
|
||||||
Value::Number(decoded_len.into()),
|
Value::Number(decoded_len.into()),
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
if let Some(decoded_len) = client_body_base64.and_then(decoded_base64_len_hint) {
|
if let Some(decoded_len) = client_decoded_len {
|
||||||
metadata.insert(
|
metadata.insert(
|
||||||
"client_response_body_base64_bytes".to_string(),
|
"client_response_body_base64_bytes".to_string(),
|
||||||
Value::Number(decoded_len.into()),
|
Value::Number(decoded_len.into()),
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
let mut body_capture = Map::new();
|
if body_capture_capacity > 0 {
|
||||||
append_body_capture_metadata_entry(
|
let mut body_capture = Map::with_capacity(body_capture_capacity);
|
||||||
&mut body_capture,
|
append_body_capture_metadata_entry(
|
||||||
"response",
|
&mut body_capture,
|
||||||
provider_body_state,
|
"response",
|
||||||
provider_body_base64.and_then(decoded_base64_len_hint),
|
provider_body_state,
|
||||||
provider_body_base64.and_then(decoded_base64_len_hint),
|
provider_decoded_len,
|
||||||
);
|
provider_decoded_len,
|
||||||
append_body_capture_metadata_entry(
|
);
|
||||||
&mut body_capture,
|
append_body_capture_metadata_entry(
|
||||||
"client_response",
|
&mut body_capture,
|
||||||
client_body_state,
|
"client_response",
|
||||||
client_body_base64.and_then(decoded_base64_len_hint),
|
client_body_state,
|
||||||
client_body_base64.and_then(decoded_base64_len_hint),
|
client_decoded_len,
|
||||||
);
|
client_decoded_len,
|
||||||
if !body_capture.is_empty() {
|
);
|
||||||
metadata.insert("body_capture".to_string(), Value::Object(body_capture));
|
metadata.insert("body_capture".to_string(), Value::Object(body_capture));
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -476,21 +507,29 @@ pub(crate) fn append_runtime_body_capture_metadata(
|
|||||||
input.provider_request_body_ref,
|
input.provider_request_body_ref,
|
||||||
input.provider_request_unavailable,
|
input.provider_request_unavailable,
|
||||||
);
|
);
|
||||||
upsert_body_capture_metadata_entry(metadata, "request", Some(states.request), None, None, None);
|
let Some(body_capture_object) = body_capture_object_mut(metadata, 2) else {
|
||||||
upsert_body_capture_metadata_entry(
|
return;
|
||||||
metadata,
|
};
|
||||||
"provider_request",
|
body_capture_object.insert(
|
||||||
Some(states.provider_request),
|
"request".to_string(),
|
||||||
input.provider_request_source_bytes,
|
build_body_capture_metadata_entry(states.request, None, None, None),
|
||||||
input.provider_request_source_bytes,
|
);
|
||||||
input.provider_request_unavailable_reason,
|
body_capture_object.insert(
|
||||||
|
"provider_request".to_string(),
|
||||||
|
build_body_capture_metadata_entry(
|
||||||
|
states.provider_request,
|
||||||
|
input.provider_request_source_bytes,
|
||||||
|
input.provider_request_source_bytes,
|
||||||
|
input.provider_request_unavailable_reason,
|
||||||
|
),
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn build_plan_body_capture_metadata(
|
pub(crate) fn build_plan_body_capture_metadata(
|
||||||
provider_request_body_base64: Option<&str>,
|
provider_request_body_base64: Option<&str>,
|
||||||
) -> Option<Value> {
|
) -> Option<Value> {
|
||||||
let mut metadata = Map::new();
|
provider_request_body_base64?;
|
||||||
|
let mut metadata = Map::with_capacity(2);
|
||||||
append_plan_body_capture_metadata(&mut metadata, provider_request_body_base64);
|
append_plan_body_capture_metadata(&mut metadata, provider_request_body_base64);
|
||||||
(!metadata.is_empty()).then_some(Value::Object(metadata))
|
(!metadata.is_empty()).then_some(Value::Object(metadata))
|
||||||
}
|
}
|
||||||
@@ -507,13 +546,17 @@ pub(crate) fn append_plan_body_capture_metadata(
|
|||||||
Value::Number(decoded_len.into()),
|
Value::Number(decoded_len.into()),
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
upsert_body_capture_metadata_entry(
|
let Some(body_capture_object) = body_capture_object_mut(metadata, 1) else {
|
||||||
metadata,
|
return;
|
||||||
"provider_request",
|
};
|
||||||
Some(UsageBodyCaptureState::Unavailable),
|
body_capture_object.insert(
|
||||||
decoded_len,
|
"provider_request".to_string(),
|
||||||
decoded_len,
|
build_body_capture_metadata_entry(
|
||||||
Some("body_bytes_base64_only"),
|
UsageBodyCaptureState::Unavailable,
|
||||||
|
decoded_len,
|
||||||
|
decoded_len,
|
||||||
|
Some("body_bytes_base64_only"),
|
||||||
|
),
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -528,58 +571,16 @@ fn append_body_capture_metadata_entry(
|
|||||||
let Some(state) = state else {
|
let Some(state) = state else {
|
||||||
return;
|
return;
|
||||||
};
|
};
|
||||||
let mut entry = Map::new();
|
target.insert(
|
||||||
entry.insert(
|
key.to_string(),
|
||||||
"state".to_string(),
|
build_body_capture_metadata_entry(
|
||||||
Value::String(state.as_str().to_string()),
|
state,
|
||||||
|
stored_bytes,
|
||||||
|
source_bytes,
|
||||||
|
matches!(state, UsageBodyCaptureState::Truncated)
|
||||||
|
.then_some("body_capture_limit_exceeded"),
|
||||||
|
),
|
||||||
);
|
);
|
||||||
if let Some(stored_bytes) = stored_bytes {
|
|
||||||
entry.insert("stored_bytes".to_string(), json!(stored_bytes));
|
|
||||||
}
|
|
||||||
if let Some(source_bytes) = source_bytes {
|
|
||||||
entry.insert("source_bytes".to_string(), json!(source_bytes));
|
|
||||||
}
|
|
||||||
if matches!(state, UsageBodyCaptureState::Truncated) {
|
|
||||||
entry.insert(
|
|
||||||
"reason".to_string(),
|
|
||||||
Value::String("body_capture_limit_exceeded".to_string()),
|
|
||||||
);
|
|
||||||
}
|
|
||||||
target.insert(key.to_string(), Value::Object(entry));
|
|
||||||
}
|
|
||||||
|
|
||||||
pub(crate) fn upsert_body_capture_metadata_entry(
|
|
||||||
metadata: &mut Map<String, Value>,
|
|
||||||
key: &str,
|
|
||||||
state: Option<UsageBodyCaptureState>,
|
|
||||||
stored_bytes: Option<u64>,
|
|
||||||
source_bytes: Option<u64>,
|
|
||||||
reason: Option<&str>,
|
|
||||||
) {
|
|
||||||
let body_capture = metadata
|
|
||||||
.entry("body_capture".to_string())
|
|
||||||
.or_insert_with(|| Value::Object(Map::new()));
|
|
||||||
let Some(body_capture_object) = body_capture.as_object_mut() else {
|
|
||||||
return;
|
|
||||||
};
|
|
||||||
let Some(state) = state else {
|
|
||||||
return;
|
|
||||||
};
|
|
||||||
let mut entry = Map::new();
|
|
||||||
entry.insert(
|
|
||||||
"state".to_string(),
|
|
||||||
Value::String(state.as_str().to_string()),
|
|
||||||
);
|
|
||||||
if let Some(bytes) = stored_bytes {
|
|
||||||
entry.insert("stored_bytes".to_string(), json!(bytes));
|
|
||||||
}
|
|
||||||
if let Some(bytes) = source_bytes {
|
|
||||||
entry.insert("source_bytes".to_string(), json!(bytes));
|
|
||||||
}
|
|
||||||
if let Some(reason) = reason {
|
|
||||||
entry.insert("reason".to_string(), Value::String(reason.to_string()));
|
|
||||||
}
|
|
||||||
body_capture_object.insert(key.to_string(), Value::Object(entry));
|
|
||||||
}
|
}
|
||||||
|
|
||||||
fn upsert_body_capture_metadata_value_entry(
|
fn upsert_body_capture_metadata_value_entry(
|
||||||
@@ -593,22 +594,63 @@ fn upsert_body_capture_metadata_value_entry(
|
|||||||
let Some(state) = state else {
|
let Some(state) = state else {
|
||||||
return;
|
return;
|
||||||
};
|
};
|
||||||
let metadata_object = metadata
|
let Some(body_capture_object) = body_capture_value_object_mut(metadata, 1) else {
|
||||||
.get_or_insert_with(|| Value::Object(Map::new()))
|
|
||||||
.as_object_mut();
|
|
||||||
let Some(metadata_object) = metadata_object else {
|
|
||||||
return;
|
return;
|
||||||
};
|
};
|
||||||
upsert_body_capture_metadata_entry(
|
body_capture_object.insert(
|
||||||
metadata_object,
|
key.to_string(),
|
||||||
key,
|
build_body_capture_metadata_entry(state, stored_bytes, source_bytes, reason),
|
||||||
Some(state),
|
|
||||||
stored_bytes,
|
|
||||||
source_bytes,
|
|
||||||
reason,
|
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn body_capture_object_mut(
|
||||||
|
metadata: &mut Map<String, Value>,
|
||||||
|
capacity: usize,
|
||||||
|
) -> Option<&mut Map<String, Value>> {
|
||||||
|
let body_capture = metadata
|
||||||
|
.entry("body_capture".to_string())
|
||||||
|
.or_insert_with(|| Value::Object(Map::with_capacity(capacity)));
|
||||||
|
body_capture.as_object_mut()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn body_capture_value_object_mut(
|
||||||
|
metadata: &mut Option<Value>,
|
||||||
|
capacity: usize,
|
||||||
|
) -> Option<&mut Map<String, Value>> {
|
||||||
|
let metadata_object = metadata
|
||||||
|
.get_or_insert_with(|| Value::Object(Map::with_capacity(1)))
|
||||||
|
.as_object_mut();
|
||||||
|
let metadata_object = metadata_object?;
|
||||||
|
body_capture_object_mut(metadata_object, capacity)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn build_body_capture_metadata_entry(
|
||||||
|
state: UsageBodyCaptureState,
|
||||||
|
stored_bytes: Option<u64>,
|
||||||
|
source_bytes: Option<u64>,
|
||||||
|
reason: Option<&str>,
|
||||||
|
) -> Value {
|
||||||
|
let mut entry = Map::with_capacity(
|
||||||
|
1 + usize::from(stored_bytes.is_some())
|
||||||
|
+ usize::from(source_bytes.is_some())
|
||||||
|
+ usize::from(reason.is_some()),
|
||||||
|
);
|
||||||
|
entry.insert(
|
||||||
|
"state".to_string(),
|
||||||
|
Value::String(state.as_str().to_owned()),
|
||||||
|
);
|
||||||
|
if let Some(bytes) = stored_bytes {
|
||||||
|
entry.insert("stored_bytes".to_string(), json!(bytes));
|
||||||
|
}
|
||||||
|
if let Some(bytes) = source_bytes {
|
||||||
|
entry.insert("source_bytes".to_string(), json!(bytes));
|
||||||
|
}
|
||||||
|
if let Some(reason) = reason {
|
||||||
|
entry.insert("reason".to_string(), Value::String(reason.to_owned()));
|
||||||
|
}
|
||||||
|
Value::Object(entry)
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) fn decoded_base64_len_hint(body_base64: &str) -> Option<u64> {
|
pub(crate) fn decoded_base64_len_hint(body_base64: &str) -> Option<u64> {
|
||||||
let body_base64 = body_base64.trim();
|
let body_base64 = body_base64.trim();
|
||||||
if body_base64.is_empty() {
|
if body_base64.is_empty() {
|
||||||
@@ -642,9 +684,18 @@ pub(crate) fn decoded_base64_len_hint(body_base64: &str) -> Option<u64> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
fn sanitize_usage_body_ref(value: Option<String>) -> Option<String> {
|
fn sanitize_usage_body_ref(value: Option<String>) -> Option<String> {
|
||||||
value
|
value.and_then(trim_owned_non_empty_string)
|
||||||
.map(|value| value.trim().to_string())
|
}
|
||||||
.filter(|value| !value.is_empty())
|
|
||||||
|
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_string())
|
||||||
}
|
}
|
||||||
|
|
||||||
fn usage_value_kind(value: &Value) -> &'static str {
|
fn usage_value_kind(value: &Value) -> &'static str {
|
||||||
@@ -657,3 +708,122 @@ fn usage_value_kind(value: &Value) -> &'static str {
|
|||||||
Value::Object(_) => "object",
|
Value::Object(_) => "object",
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use super::{
|
||||||
|
build_plan_body_capture_metadata, sync_usage_body_ref_metadata,
|
||||||
|
trim_owned_non_empty_string, truncate_usage_body_string,
|
||||||
|
upsert_body_capture_metadata_value_entry,
|
||||||
|
};
|
||||||
|
use aether_data_contracts::repository::usage::UsageBodyCaptureState;
|
||||||
|
use aether_data_contracts::repository::usage::UsageBodyField;
|
||||||
|
use serde_json::{Map, Value};
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn build_plan_body_capture_metadata_returns_none_without_base64_body() {
|
||||||
|
assert!(build_plan_body_capture_metadata(None).is_none());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn trim_owned_non_empty_string_preserves_clean_values_and_drops_blank_ones() {
|
||||||
|
assert_eq!(
|
||||||
|
trim_owned_non_empty_string("blob://body-ref-1".to_string()),
|
||||||
|
Some("blob://body-ref-1".to_string()),
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
trim_owned_non_empty_string(" blob://body-ref-1 ".to_string()),
|
||||||
|
Some("blob://body-ref-1".to_string()),
|
||||||
|
);
|
||||||
|
assert_eq!(trim_owned_non_empty_string(" ".to_string()), None);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn upsert_body_capture_metadata_value_entry_ignores_none_state() {
|
||||||
|
let mut metadata = Some(Value::Object(Map::<String, Value>::new()));
|
||||||
|
upsert_body_capture_metadata_value_entry(&mut metadata, "response", None, None, None, None);
|
||||||
|
assert_eq!(metadata, Some(Value::Object(Map::new())));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn upsert_body_capture_metadata_value_entry_preserves_existing_metadata_fields() {
|
||||||
|
let mut metadata = Some(Value::Object(Map::from_iter([(
|
||||||
|
"request_body_ref".to_string(),
|
||||||
|
Value::String("blob://body-ref-1".to_string()),
|
||||||
|
)])));
|
||||||
|
|
||||||
|
upsert_body_capture_metadata_value_entry(
|
||||||
|
&mut metadata,
|
||||||
|
"response",
|
||||||
|
Some(UsageBodyCaptureState::Reference),
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
);
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
metadata,
|
||||||
|
Some(Value::Object(Map::from_iter([
|
||||||
|
(
|
||||||
|
"request_body_ref".to_string(),
|
||||||
|
Value::String("blob://body-ref-1".to_string()),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"body_capture".to_string(),
|
||||||
|
Value::Object(Map::from_iter([(
|
||||||
|
"response".to_string(),
|
||||||
|
Value::Object(Map::from_iter([(
|
||||||
|
"state".to_string(),
|
||||||
|
Value::String("reference".to_string()),
|
||||||
|
)])),
|
||||||
|
)])),
|
||||||
|
),
|
||||||
|
]))),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn sync_usage_body_ref_metadata_clears_empty_metadata_object() {
|
||||||
|
let mut metadata = Some(Value::Object(Map::from_iter([(
|
||||||
|
"request_body_ref".to_string(),
|
||||||
|
Value::String("blob://body-ref-1".to_string()),
|
||||||
|
)])));
|
||||||
|
|
||||||
|
sync_usage_body_ref_metadata(&mut metadata, UsageBodyField::RequestBody, None);
|
||||||
|
|
||||||
|
assert!(metadata.is_none());
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn sync_usage_body_ref_metadata_preserves_existing_ref_value() {
|
||||||
|
let mut metadata = Some(Value::Object(Map::from_iter([(
|
||||||
|
"request_body_ref".to_string(),
|
||||||
|
Value::String("blob://body-ref-1".to_string()),
|
||||||
|
)])));
|
||||||
|
|
||||||
|
sync_usage_body_ref_metadata(
|
||||||
|
&mut metadata,
|
||||||
|
UsageBodyField::RequestBody,
|
||||||
|
Some("blob://body-ref-1"),
|
||||||
|
);
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
metadata,
|
||||||
|
Some(Value::Object(Map::from_iter([(
|
||||||
|
"request_body_ref".to_string(),
|
||||||
|
Value::String("blob://body-ref-1".to_string()),
|
||||||
|
)]))),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn truncate_usage_body_string_respects_json_byte_limit() {
|
||||||
|
let limit = 32usize;
|
||||||
|
let truncated = truncate_usage_body_string("x".repeat(256).as_str(), limit);
|
||||||
|
|
||||||
|
assert!(truncated.ends_with("...[truncated]"));
|
||||||
|
assert!(serde_json::to_vec(&truncated)
|
||||||
|
.ok()
|
||||||
|
.is_some_and(|bytes| bytes.len() <= limit));
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -31,13 +31,36 @@ pub(crate) fn merge_usage_request_metadata(
|
|||||||
(!metadata.is_empty()).then_some(Value::Object(metadata))
|
(!metadata.is_empty()).then_some(Value::Object(metadata))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub(crate) fn merge_usage_request_metadata_owned(
|
||||||
|
base: Option<Value>,
|
||||||
|
override_value: Option<Value>,
|
||||||
|
) -> Option<Value> {
|
||||||
|
let mut metadata = match base {
|
||||||
|
Some(Value::Object(base)) => base,
|
||||||
|
_ => Map::new(),
|
||||||
|
};
|
||||||
|
if let Some(Value::Object(override_object)) = override_value {
|
||||||
|
move_allowed_metadata_fields(override_object, &mut metadata);
|
||||||
|
}
|
||||||
|
(!metadata.is_empty()).then_some(Value::Object(metadata))
|
||||||
|
}
|
||||||
|
|
||||||
pub(crate) fn sanitize_usage_request_metadata(value: Option<Value>) -> Option<Value> {
|
pub(crate) fn sanitize_usage_request_metadata(value: Option<Value>) -> Option<Value> {
|
||||||
let Value::Object(object) = value? else {
|
let Value::Object(object) = value? else {
|
||||||
return None;
|
return None;
|
||||||
};
|
};
|
||||||
|
|
||||||
let mut filtered = Map::new();
|
let mut filtered = Map::new();
|
||||||
copy_allowed_metadata_fields(&object, &mut filtered);
|
move_allowed_metadata_fields(object, &mut filtered);
|
||||||
|
|
||||||
|
(!filtered.is_empty()).then_some(Value::Object(filtered))
|
||||||
|
}
|
||||||
|
|
||||||
|
pub(crate) fn sanitize_usage_request_metadata_ref(value: Option<&Value>) -> Option<Value> {
|
||||||
|
let object = value.and_then(Value::as_object)?;
|
||||||
|
|
||||||
|
let mut filtered = Map::new();
|
||||||
|
copy_allowed_metadata_fields(object, &mut filtered);
|
||||||
|
|
||||||
(!filtered.is_empty()).then_some(Value::Object(filtered))
|
(!filtered.is_empty()).then_some(Value::Object(filtered))
|
||||||
}
|
}
|
||||||
@@ -62,6 +85,26 @@ fn copy_allowed_metadata_fields(source: &Map<String, Value>, target: &mut Map<St
|
|||||||
copy_number(source, target, "price_per_request");
|
copy_number(source, target, "price_per_request");
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn move_allowed_metadata_fields(mut source: Map<String, Value>, target: &mut Map<String, Value>) {
|
||||||
|
remove_non_empty_string(&mut source, target, "trace_id");
|
||||||
|
remove_number(&mut source, target, "provider_request_body_base64_bytes");
|
||||||
|
remove_number(&mut source, target, "provider_response_body_base64_bytes");
|
||||||
|
remove_number(&mut source, target, "client_response_body_base64_bytes");
|
||||||
|
remove_non_null_value(&mut source, target, "billing_snapshot");
|
||||||
|
remove_non_empty_string(&mut source, target, "billing_snapshot_schema_version");
|
||||||
|
remove_non_empty_string(&mut source, target, "billing_snapshot_status");
|
||||||
|
remove_non_null_value(&mut source, target, "dimensions");
|
||||||
|
remove_non_null_value(&mut source, target, "billing_rule_snapshot");
|
||||||
|
remove_non_null_value(&mut source, target, "scheduling_audit");
|
||||||
|
remove_number(&mut source, target, "rate_multiplier");
|
||||||
|
remove_bool(&mut source, target, "is_free_tier");
|
||||||
|
remove_number(&mut source, target, "input_price_per_1m");
|
||||||
|
remove_number(&mut source, target, "output_price_per_1m");
|
||||||
|
remove_number(&mut source, target, "cache_creation_price_per_1m");
|
||||||
|
remove_number(&mut source, target, "cache_read_price_per_1m");
|
||||||
|
remove_number(&mut source, target, "price_per_request");
|
||||||
|
}
|
||||||
|
|
||||||
fn copy_non_empty_string(source: &Map<String, Value>, target: &mut Map<String, Value>, key: &str) {
|
fn copy_non_empty_string(source: &Map<String, Value>, target: &mut Map<String, Value>, key: &str) {
|
||||||
let Some(value) = source
|
let Some(value) = source
|
||||||
.get(key)
|
.get(key)
|
||||||
@@ -77,6 +120,20 @@ fn copy_non_empty_string(source: &Map<String, Value>, target: &mut Map<String, V
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn remove_non_empty_string(
|
||||||
|
source: &mut Map<String, Value>,
|
||||||
|
target: &mut Map<String, Value>,
|
||||||
|
key: &str,
|
||||||
|
) {
|
||||||
|
let Some(Value::String(value)) = source.remove(key) else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
let Some(value) = trim_and_truncate_usage_request_metadata_string_owned(value) else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
target.insert(key.to_string(), Value::String(value));
|
||||||
|
}
|
||||||
|
|
||||||
fn copy_number(source: &Map<String, Value>, target: &mut Map<String, Value>, key: &str) {
|
fn copy_number(source: &Map<String, Value>, target: &mut Map<String, Value>, key: &str) {
|
||||||
let Some(value) = source.get(key).filter(|value| value.is_number()) else {
|
let Some(value) = source.get(key).filter(|value| value.is_number()) else {
|
||||||
return;
|
return;
|
||||||
@@ -84,6 +141,13 @@ fn copy_number(source: &Map<String, Value>, target: &mut Map<String, Value>, key
|
|||||||
target.insert(key.to_string(), value.clone());
|
target.insert(key.to_string(), value.clone());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn remove_number(source: &mut Map<String, Value>, target: &mut Map<String, Value>, key: &str) {
|
||||||
|
let Some(value) = source.remove(key).filter(|value| value.is_number()) else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
target.insert(key.to_string(), value);
|
||||||
|
}
|
||||||
|
|
||||||
fn copy_bool(source: &Map<String, Value>, target: &mut Map<String, Value>, key: &str) {
|
fn copy_bool(source: &Map<String, Value>, target: &mut Map<String, Value>, key: &str) {
|
||||||
let Some(value) = source.get(key).filter(|value| value.is_boolean()) else {
|
let Some(value) = source.get(key).filter(|value| value.is_boolean()) else {
|
||||||
return;
|
return;
|
||||||
@@ -91,6 +155,13 @@ fn copy_bool(source: &Map<String, Value>, target: &mut Map<String, Value>, key:
|
|||||||
target.insert(key.to_string(), value.clone());
|
target.insert(key.to_string(), value.clone());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn remove_bool(source: &mut Map<String, Value>, target: &mut Map<String, Value>, key: &str) {
|
||||||
|
let Some(value) = source.remove(key).filter(|value| value.is_boolean()) else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
target.insert(key.to_string(), value);
|
||||||
|
}
|
||||||
|
|
||||||
fn copy_non_null_value(source: &Map<String, Value>, target: &mut Map<String, Value>, key: &str) {
|
fn copy_non_null_value(source: &Map<String, Value>, target: &mut Map<String, Value>, key: &str) {
|
||||||
let Some(value) = source.get(key).filter(|value| !value.is_null()) else {
|
let Some(value) = source.get(key).filter(|value| !value.is_null()) else {
|
||||||
return;
|
return;
|
||||||
@@ -101,6 +172,20 @@ fn copy_non_null_value(source: &Map<String, Value>, target: &mut Map<String, Val
|
|||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn remove_non_null_value(
|
||||||
|
source: &mut Map<String, Value>,
|
||||||
|
target: &mut Map<String, Value>,
|
||||||
|
key: &str,
|
||||||
|
) {
|
||||||
|
let Some(value) = source.remove(key).filter(|value| !value.is_null()) else {
|
||||||
|
return;
|
||||||
|
};
|
||||||
|
target.insert(
|
||||||
|
key.to_string(),
|
||||||
|
sanitize_usage_request_metadata_value_owned(value),
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
fn sanitize_usage_request_metadata_value(value: &Value) -> Value {
|
fn sanitize_usage_request_metadata_value(value: &Value) -> Value {
|
||||||
match value {
|
match value {
|
||||||
Value::String(text) => Value::String(truncate_usage_request_metadata_string(text)),
|
Value::String(text) => Value::String(truncate_usage_request_metadata_string(text)),
|
||||||
@@ -109,6 +194,14 @@ fn sanitize_usage_request_metadata_value(value: &Value) -> Value {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn sanitize_usage_request_metadata_value_owned(value: Value) -> Value {
|
||||||
|
match value {
|
||||||
|
Value::String(text) => Value::String(truncate_usage_request_metadata_string_owned(text)),
|
||||||
|
_ if usage_request_metadata_within_limits(&value) => value,
|
||||||
|
_ => truncated_usage_request_metadata_value(&value),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
fn truncate_usage_request_metadata_string(value: &str) -> String {
|
fn truncate_usage_request_metadata_string(value: &str) -> String {
|
||||||
const TRUNCATED_SUFFIX: &str = "...[truncated]";
|
const TRUNCATED_SUFFIX: &str = "...[truncated]";
|
||||||
|
|
||||||
@@ -134,6 +227,24 @@ fn truncate_usage_request_metadata_string(value: &str) -> String {
|
|||||||
format!("{}{TRUNCATED_SUFFIX}", &value[..end])
|
format!("{}{TRUNCATED_SUFFIX}", &value[..end])
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn trim_and_truncate_usage_request_metadata_string_owned(value: String) -> Option<String> {
|
||||||
|
let trimmed = value.trim();
|
||||||
|
if trimmed.is_empty() {
|
||||||
|
return None;
|
||||||
|
}
|
||||||
|
if trimmed.len() == value.len() {
|
||||||
|
return Some(truncate_usage_request_metadata_string_owned(value));
|
||||||
|
}
|
||||||
|
Some(truncate_usage_request_metadata_string(trimmed))
|
||||||
|
}
|
||||||
|
|
||||||
|
fn truncate_usage_request_metadata_string_owned(value: String) -> String {
|
||||||
|
if value.len() <= MAX_USAGE_REQUEST_METADATA_STRING_BYTES {
|
||||||
|
return value;
|
||||||
|
}
|
||||||
|
truncate_usage_request_metadata_string(value.as_str())
|
||||||
|
}
|
||||||
|
|
||||||
fn truncated_usage_request_metadata_value(value: &Value) -> Value {
|
fn truncated_usage_request_metadata_value(value: &Value) -> Value {
|
||||||
json!({
|
json!({
|
||||||
"truncated": true,
|
"truncated": true,
|
||||||
@@ -217,7 +328,8 @@ mod tests {
|
|||||||
|
|
||||||
use super::{
|
use super::{
|
||||||
build_usage_request_metadata_seed, merge_usage_request_metadata,
|
build_usage_request_metadata_seed, merge_usage_request_metadata,
|
||||||
sanitize_usage_request_metadata, MAX_USAGE_REQUEST_METADATA_BYTES,
|
merge_usage_request_metadata_owned, sanitize_usage_request_metadata,
|
||||||
|
sanitize_usage_request_metadata_ref, MAX_USAGE_REQUEST_METADATA_BYTES,
|
||||||
MAX_USAGE_REQUEST_METADATA_DEPTH, MAX_USAGE_REQUEST_METADATA_NODES,
|
MAX_USAGE_REQUEST_METADATA_DEPTH, MAX_USAGE_REQUEST_METADATA_NODES,
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -363,4 +475,35 @@ mod tests {
|
|||||||
|
|
||||||
assert_eq!(metadata, None);
|
assert_eq!(metadata, None);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn owned_merge_matches_filtered_merge_for_trusted_objects() {
|
||||||
|
let base = Some(json!({
|
||||||
|
"trace_id": "trace-1",
|
||||||
|
"provider_request_body_base64_bytes": 128
|
||||||
|
}));
|
||||||
|
let override_value = Some(json!({
|
||||||
|
"billing_snapshot_status": "complete",
|
||||||
|
"trace_id": "trace-2"
|
||||||
|
}));
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
merge_usage_request_metadata_owned(base.clone(), override_value.clone()),
|
||||||
|
merge_usage_request_metadata(base, override_value)
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn borrowed_sanitize_matches_owned_sanitize() {
|
||||||
|
let value = json!({
|
||||||
|
"trace_id": "trace-1",
|
||||||
|
"billing_snapshot": {"status": "complete"},
|
||||||
|
"provider_name": "OpenAI"
|
||||||
|
});
|
||||||
|
|
||||||
|
assert_eq!(
|
||||||
|
sanitize_usage_request_metadata_ref(Some(&value)),
|
||||||
|
sanitize_usage_request_metadata(Some(value))
|
||||||
|
);
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -12,8 +12,7 @@ use tracing::warn;
|
|||||||
use crate::executor::spawn_on_usage_background_runtime;
|
use crate::executor::spawn_on_usage_background_runtime;
|
||||||
use crate::{
|
use crate::{
|
||||||
apply_usage_body_capture_policy_to_event, apply_usage_body_capture_policy_to_record,
|
apply_usage_body_capture_policy_to_event, apply_usage_body_capture_policy_to_record,
|
||||||
build_pending_usage_record_from_seed, build_stream_terminal_usage_seed,
|
build_stream_terminal_usage_seed, build_sync_terminal_usage_seed,
|
||||||
build_streaming_usage_record_from_seed, build_sync_terminal_usage_seed,
|
|
||||||
build_terminal_usage_event_from_seed, build_upsert_usage_record_from_event,
|
build_terminal_usage_event_from_seed, build_upsert_usage_record_from_event,
|
||||||
build_usage_queue_worker, settle_usage_if_needed, LifecycleUsageSeed,
|
build_usage_queue_worker, settle_usage_if_needed, LifecycleUsageSeed,
|
||||||
StreamTerminalUsagePayloadSeed, SyncTerminalUsagePayloadSeed, TerminalUsageContextSeed,
|
StreamTerminalUsagePayloadSeed, SyncTerminalUsagePayloadSeed, TerminalUsageContextSeed,
|
||||||
@@ -74,17 +73,6 @@ pub struct UsageRuntime {
|
|||||||
config: UsageRuntimeConfig,
|
config: UsageRuntimeConfig,
|
||||||
}
|
}
|
||||||
|
|
||||||
struct SyncTerminalUsageTaskInput {
|
|
||||||
context_seed: TerminalUsageContextSeed,
|
|
||||||
payload_seed: SyncTerminalUsagePayloadSeed,
|
|
||||||
}
|
|
||||||
|
|
||||||
struct StreamTerminalUsageTaskInput {
|
|
||||||
context_seed: TerminalUsageContextSeed,
|
|
||||||
payload_seed: StreamTerminalUsagePayloadSeed,
|
|
||||||
cancelled: bool,
|
|
||||||
}
|
|
||||||
|
|
||||||
impl Default for UsageRuntime {
|
impl Default for UsageRuntime {
|
||||||
fn default() -> Self {
|
fn default() -> Self {
|
||||||
Self::disabled()
|
Self::disabled()
|
||||||
@@ -126,7 +114,7 @@ impl UsageRuntime {
|
|||||||
Some(worker.spawn())
|
Some(worker.spawn())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub fn record_pending<T>(&self, data: &T, seed: &LifecycleUsageSeed)
|
pub fn record_pending<T>(&self, data: &T, seed: LifecycleUsageSeed)
|
||||||
where
|
where
|
||||||
T: UsageRuntimeAccess + Clone + 'static,
|
T: UsageRuntimeAccess + Clone + 'static,
|
||||||
{
|
{
|
||||||
@@ -134,11 +122,10 @@ impl UsageRuntime {
|
|||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
let data = T::clone(data);
|
let data = T::clone(data);
|
||||||
let seed = seed.clone();
|
|
||||||
let request_id = seed.request_id.clone();
|
let request_id = seed.request_id.clone();
|
||||||
spawn_on_usage_background_runtime(boxed_usage_task(async move {
|
spawn_on_usage_background_runtime(boxed_usage_task(async move {
|
||||||
let now_unix_secs = now_unix_secs();
|
let now_unix_secs = now_unix_secs();
|
||||||
match build_pending_usage_record_offthread(&seed, now_unix_secs).await {
|
match build_pending_usage_record_offthread(seed, now_unix_secs).await {
|
||||||
Ok(mut record) => {
|
Ok(mut record) => {
|
||||||
apply_body_capture_policy_to_record_from_data(&data, &mut record).await;
|
apply_body_capture_policy_to_record_from_data(&data, &mut record).await;
|
||||||
if let Err(err) = data.upsert_usage_record(record).await {
|
if let Err(err) = data.upsert_usage_record(record).await {
|
||||||
@@ -183,9 +170,9 @@ impl UsageRuntime {
|
|||||||
spawn_on_usage_background_runtime(boxed_usage_task(async move {
|
spawn_on_usage_background_runtime(boxed_usage_task(async move {
|
||||||
let now_unix_secs = now_unix_secs();
|
let now_unix_secs = now_unix_secs();
|
||||||
match build_streaming_usage_record_offthread(
|
match build_streaming_usage_record_offthread(
|
||||||
&seed,
|
seed,
|
||||||
status_code,
|
status_code,
|
||||||
telemetry.as_ref(),
|
telemetry,
|
||||||
now_unix_secs,
|
now_unix_secs,
|
||||||
)
|
)
|
||||||
.await
|
.await
|
||||||
@@ -218,8 +205,8 @@ impl UsageRuntime {
|
|||||||
pub fn record_sync_terminal<T>(
|
pub fn record_sync_terminal<T>(
|
||||||
&self,
|
&self,
|
||||||
data: &T,
|
data: &T,
|
||||||
context_seed: &TerminalUsageContextSeed,
|
context_seed: TerminalUsageContextSeed,
|
||||||
payload_seed: &SyncTerminalUsagePayloadSeed,
|
payload_seed: SyncTerminalUsagePayloadSeed,
|
||||||
) where
|
) where
|
||||||
T: UsageRuntimeAccess + Clone + 'static,
|
T: UsageRuntimeAccess + Clone + 'static,
|
||||||
{
|
{
|
||||||
@@ -229,12 +216,8 @@ impl UsageRuntime {
|
|||||||
let runtime = self.clone();
|
let runtime = self.clone();
|
||||||
let data = T::clone(data);
|
let data = T::clone(data);
|
||||||
let request_id = context_seed.request_id.clone();
|
let request_id = context_seed.request_id.clone();
|
||||||
let input = Box::new(SyncTerminalUsageTaskInput {
|
|
||||||
context_seed: context_seed.clone(),
|
|
||||||
payload_seed: payload_seed.clone(),
|
|
||||||
});
|
|
||||||
spawn_on_usage_background_runtime(boxed_usage_task(async move {
|
spawn_on_usage_background_runtime(boxed_usage_task(async move {
|
||||||
match build_sync_terminal_usage_event_offthread(input).await {
|
match build_sync_terminal_usage_event_offthread(context_seed, payload_seed).await {
|
||||||
Ok(mut event) => {
|
Ok(mut event) => {
|
||||||
apply_body_capture_policy_from_data(&data, &mut event).await;
|
apply_body_capture_policy_from_data(&data, &mut event).await;
|
||||||
if let Err(err) = data.enrich_usage_event(&mut event).await {
|
if let Err(err) = data.enrich_usage_event(&mut event).await {
|
||||||
@@ -264,8 +247,8 @@ impl UsageRuntime {
|
|||||||
pub fn record_stream_terminal<T>(
|
pub fn record_stream_terminal<T>(
|
||||||
&self,
|
&self,
|
||||||
data: &T,
|
data: &T,
|
||||||
context_seed: &TerminalUsageContextSeed,
|
context_seed: TerminalUsageContextSeed,
|
||||||
payload_seed: &StreamTerminalUsagePayloadSeed,
|
payload_seed: StreamTerminalUsagePayloadSeed,
|
||||||
cancelled: bool,
|
cancelled: bool,
|
||||||
) where
|
) where
|
||||||
T: UsageRuntimeAccess + Clone + 'static,
|
T: UsageRuntimeAccess + Clone + 'static,
|
||||||
@@ -276,13 +259,10 @@ impl UsageRuntime {
|
|||||||
let runtime = self.clone();
|
let runtime = self.clone();
|
||||||
let data = T::clone(data);
|
let data = T::clone(data);
|
||||||
let request_id = context_seed.request_id.clone();
|
let request_id = context_seed.request_id.clone();
|
||||||
let input = Box::new(StreamTerminalUsageTaskInput {
|
|
||||||
context_seed: context_seed.clone(),
|
|
||||||
payload_seed: payload_seed.clone(),
|
|
||||||
cancelled,
|
|
||||||
});
|
|
||||||
spawn_on_usage_background_runtime(boxed_usage_task(async move {
|
spawn_on_usage_background_runtime(boxed_usage_task(async move {
|
||||||
match build_stream_terminal_usage_event_offthread(input).await {
|
match build_stream_terminal_usage_event_offthread(context_seed, payload_seed, cancelled)
|
||||||
|
.await
|
||||||
|
{
|
||||||
Ok(mut event) => {
|
Ok(mut event) => {
|
||||||
apply_body_capture_policy_from_data(&data, &mut event).await;
|
apply_body_capture_policy_from_data(&data, &mut event).await;
|
||||||
if let Err(err) = data.enrich_usage_event(&mut event).await {
|
if let Err(err) = data.enrich_usage_event(&mut event).await {
|
||||||
@@ -413,28 +393,27 @@ impl UsageRuntime {
|
|||||||
}
|
}
|
||||||
|
|
||||||
async fn build_pending_usage_record_offthread(
|
async fn build_pending_usage_record_offthread(
|
||||||
seed: &LifecycleUsageSeed,
|
seed: LifecycleUsageSeed,
|
||||||
now_unix_secs: u64,
|
now_unix_secs: u64,
|
||||||
) -> Result<UpsertUsageRecord, DataLayerError> {
|
) -> Result<UpsertUsageRecord, DataLayerError> {
|
||||||
let seed = seed.clone();
|
tokio::task::spawn_blocking(move || {
|
||||||
tokio::task::spawn_blocking(move || build_pending_usage_record_from_seed(&seed, now_unix_secs))
|
crate::write::build_pending_usage_record_from_owned_seed(seed, now_unix_secs)
|
||||||
.await
|
})
|
||||||
.map_err(join_error_to_data_layer)?
|
.await
|
||||||
|
.map_err(join_error_to_data_layer)?
|
||||||
}
|
}
|
||||||
|
|
||||||
async fn build_streaming_usage_record_offthread(
|
async fn build_streaming_usage_record_offthread(
|
||||||
seed: &LifecycleUsageSeed,
|
seed: LifecycleUsageSeed,
|
||||||
status_code: u16,
|
status_code: u16,
|
||||||
telemetry: Option<&ExecutionTelemetry>,
|
telemetry: Option<ExecutionTelemetry>,
|
||||||
now_unix_secs: u64,
|
now_unix_secs: u64,
|
||||||
) -> Result<UpsertUsageRecord, DataLayerError> {
|
) -> Result<UpsertUsageRecord, DataLayerError> {
|
||||||
let seed = seed.clone();
|
|
||||||
let telemetry = telemetry.cloned();
|
|
||||||
tokio::task::spawn_blocking(move || {
|
tokio::task::spawn_blocking(move || {
|
||||||
build_streaming_usage_record_from_seed(
|
crate::write::build_streaming_usage_record_from_owned_seed(
|
||||||
&seed,
|
seed,
|
||||||
status_code,
|
status_code,
|
||||||
telemetry.as_ref(),
|
telemetry,
|
||||||
now_unix_secs,
|
now_unix_secs,
|
||||||
)
|
)
|
||||||
})
|
})
|
||||||
@@ -443,12 +422,13 @@ async fn build_streaming_usage_record_offthread(
|
|||||||
}
|
}
|
||||||
|
|
||||||
async fn build_sync_terminal_usage_event_offthread(
|
async fn build_sync_terminal_usage_event_offthread(
|
||||||
input: Box<SyncTerminalUsageTaskInput>,
|
context_seed: TerminalUsageContextSeed,
|
||||||
|
payload_seed: SyncTerminalUsagePayloadSeed,
|
||||||
) -> Result<UsageEvent, DataLayerError> {
|
) -> Result<UsageEvent, DataLayerError> {
|
||||||
tokio::task::spawn_blocking(move || {
|
tokio::task::spawn_blocking(move || {
|
||||||
build_terminal_usage_event_from_seed(build_sync_terminal_usage_seed(
|
build_terminal_usage_event_from_seed(build_sync_terminal_usage_seed(
|
||||||
input.context_seed,
|
context_seed,
|
||||||
input.payload_seed,
|
payload_seed,
|
||||||
))
|
))
|
||||||
})
|
})
|
||||||
.await
|
.await
|
||||||
@@ -456,13 +436,15 @@ async fn build_sync_terminal_usage_event_offthread(
|
|||||||
}
|
}
|
||||||
|
|
||||||
async fn build_stream_terminal_usage_event_offthread(
|
async fn build_stream_terminal_usage_event_offthread(
|
||||||
input: Box<StreamTerminalUsageTaskInput>,
|
context_seed: TerminalUsageContextSeed,
|
||||||
|
payload_seed: StreamTerminalUsagePayloadSeed,
|
||||||
|
cancelled: bool,
|
||||||
) -> Result<UsageEvent, DataLayerError> {
|
) -> Result<UsageEvent, DataLayerError> {
|
||||||
tokio::task::spawn_blocking(move || {
|
tokio::task::spawn_blocking(move || {
|
||||||
build_terminal_usage_event_from_seed(build_stream_terminal_usage_seed(
|
build_terminal_usage_event_from_seed(build_stream_terminal_usage_seed(
|
||||||
input.context_seed,
|
context_seed,
|
||||||
input.payload_seed,
|
payload_seed,
|
||||||
input.cancelled,
|
cancelled,
|
||||||
))
|
))
|
||||||
})
|
})
|
||||||
.await
|
.await
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -61,7 +61,7 @@
|
|||||||
|
|
||||||
<div class="flex flex-col gap-2 text-xs lg:flex-row lg:items-center lg:justify-between">
|
<div class="flex flex-col gap-2 text-xs lg:flex-row lg:items-center lg:justify-between">
|
||||||
<div class="text-muted-foreground">
|
<div class="text-muted-foreground">
|
||||||
共 {{ filteredTotal }} 个匹配账号,当前页 {{ pageKeys.length }} 个,已选 {{ selectedCount }} 个
|
共 {{ filteredTotal }} 个匹配账号,当前页 {{ pageKeyRows.length }} 个,已选 {{ selectedCount }} 个
|
||||||
</div>
|
</div>
|
||||||
<div class="flex flex-wrap items-center gap-1">
|
<div class="flex flex-wrap items-center gap-1">
|
||||||
<div class="mr-1 flex items-center gap-2">
|
<div class="mr-1 flex items-center gap-2">
|
||||||
@@ -77,7 +77,7 @@
|
|||||||
variant="ghost"
|
variant="ghost"
|
||||||
size="sm"
|
size="sm"
|
||||||
class="h-7 px-2 text-[11px]"
|
class="h-7 px-2 text-[11px]"
|
||||||
:disabled="pageKeys.length === 0 || loading || executing || selectAllFiltered"
|
:disabled="pageKeyRows.length === 0 || loading || executing || selectAllFiltered"
|
||||||
@click="toggleSelectCurrentPage"
|
@click="toggleSelectCurrentPage"
|
||||||
>
|
>
|
||||||
{{ isCurrentPageFullySelected ? '取消本页全选' : '本页全选' }}
|
{{ isCurrentPageFullySelected ? '取消本页全选' : '本页全选' }}
|
||||||
@@ -106,54 +106,54 @@
|
|||||||
正在加载账号列表...
|
正在加载账号列表...
|
||||||
</div>
|
</div>
|
||||||
<div
|
<div
|
||||||
v-else-if="pageKeys.length === 0"
|
v-else-if="pageKeyRows.length === 0"
|
||||||
class="py-10 text-center text-sm text-muted-foreground"
|
class="py-10 text-center text-sm text-muted-foreground"
|
||||||
>
|
>
|
||||||
无匹配账号
|
无匹配账号
|
||||||
</div>
|
</div>
|
||||||
<label
|
<label
|
||||||
v-for="key in pageKeys"
|
v-for="row in pageKeyRows"
|
||||||
:key="key.key_id"
|
:key="row.key.key_id"
|
||||||
class="flex items-center gap-2.5 px-3 py-2 border-b last:border-b-0 cursor-pointer hover:bg-muted/30"
|
class="flex items-center gap-2.5 px-3 py-2 border-b last:border-b-0 cursor-pointer hover:bg-muted/30"
|
||||||
>
|
>
|
||||||
<Checkbox
|
<Checkbox
|
||||||
:checked="selectAllFiltered || selectedIdSet.has(key.key_id)"
|
:checked="selectAllFiltered || selectedIdSet.has(row.key.key_id)"
|
||||||
:disabled="executing || selectAllFiltered"
|
:disabled="executing || selectAllFiltered"
|
||||||
@update:checked="(checked) => toggleOne(key.key_id, checked === true)"
|
@update:checked="(checked) => toggleOne(row.key.key_id, checked === true)"
|
||||||
/>
|
/>
|
||||||
<div class="min-w-0 flex-1">
|
<div class="min-w-0 flex-1">
|
||||||
<div class="flex items-center gap-1.5">
|
<div class="flex items-center gap-1.5">
|
||||||
<span class="text-xs font-medium truncate">{{ key.key_name || '未命名' }}</span>
|
<span class="text-xs font-medium truncate">{{ row.key.key_name || '未命名' }}</span>
|
||||||
<Badge
|
<Badge
|
||||||
variant="outline"
|
variant="outline"
|
||||||
class="text-[10px] px-1 py-0 h-4 shrink-0"
|
class="text-[10px] px-1 py-0 h-4 shrink-0"
|
||||||
>{{ normalizeAuthTypeLabel(key) }}</Badge>
|
>{{ row.authTypeLabel }}</Badge>
|
||||||
<Badge
|
<Badge
|
||||||
v-if="getStatusBadgeLabel(key)"
|
v-if="row.statusBadgeLabel"
|
||||||
variant="destructive"
|
variant="destructive"
|
||||||
class="text-[10px] px-1 py-0 h-4 shrink-0"
|
class="text-[10px] px-1 py-0 h-4 shrink-0"
|
||||||
:title="getStatusBadgeTitle(key)"
|
:title="row.statusBadgeTitle"
|
||||||
>{{ getStatusBadgeLabel(key) }}</Badge>
|
>{{ row.statusBadgeLabel }}</Badge>
|
||||||
<Badge
|
<Badge
|
||||||
v-if="key.oauth_plan_type"
|
v-if="row.key.oauth_plan_type"
|
||||||
variant="outline"
|
variant="outline"
|
||||||
class="text-[10px] px-1 py-0 h-4 shrink-0"
|
class="text-[10px] px-1 py-0 h-4 shrink-0"
|
||||||
>{{ key.oauth_plan_type }}</Badge>
|
>{{ row.key.oauth_plan_type }}</Badge>
|
||||||
<Badge
|
<Badge
|
||||||
v-if="getOAuthOrgBadge(key)"
|
v-if="row.oauthOrgBadge"
|
||||||
variant="secondary"
|
variant="secondary"
|
||||||
class="text-[10px] px-1 py-0 h-4 shrink-0"
|
class="text-[10px] px-1 py-0 h-4 shrink-0"
|
||||||
:title="getOAuthOrgBadge(key)?.title"
|
:title="row.oauthOrgBadge.title"
|
||||||
>{{ getOAuthOrgBadge(key)?.label }}</Badge>
|
>{{ row.oauthOrgBadge.label }}</Badge>
|
||||||
</div>
|
</div>
|
||||||
<div class="flex items-center gap-1.5 mt-0.5 text-[11px] text-muted-foreground flex-wrap">
|
<div class="flex items-center gap-1.5 mt-0.5 text-[11px] text-muted-foreground flex-wrap">
|
||||||
<span :class="key.is_active ? '' : 'text-destructive'">{{ key.is_active ? '启用' : '禁用' }}</span>
|
<span :class="row.key.is_active ? '' : 'text-destructive'">{{ row.key.is_active ? '启用' : '禁用' }}</span>
|
||||||
<span v-if="getQuotaText(key)">{{ shortenQuota(getQuotaText(key) || '') }}</span>
|
<span v-if="row.quotaText">{{ row.quotaTextShort }}</span>
|
||||||
<span v-if="key.proxy?.node_id">独立代理</span>
|
<span v-if="row.key.proxy?.node_id">独立代理</span>
|
||||||
<span
|
<span
|
||||||
v-if="key.last_used_at"
|
v-if="row.lastUsedRelative"
|
||||||
class="ml-auto shrink-0"
|
class="ml-auto shrink-0"
|
||||||
>{{ formatRelativeTime(key.last_used_at) }}</span>
|
>{{ row.lastUsedRelative }}</span>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
</label>
|
</label>
|
||||||
@@ -331,6 +331,17 @@ type BatchActionOption = {
|
|||||||
destructive?: boolean
|
destructive?: boolean
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type PageKeyRow = {
|
||||||
|
key: PoolKeyDetail
|
||||||
|
authTypeLabel: string
|
||||||
|
statusBadgeLabel: string | null
|
||||||
|
statusBadgeTitle: string
|
||||||
|
oauthOrgBadge: ReturnType<typeof getOAuthOrgBadge>
|
||||||
|
quotaText: string | null
|
||||||
|
quotaTextShort: string
|
||||||
|
lastUsedRelative: string
|
||||||
|
}
|
||||||
|
|
||||||
const props = defineProps<{
|
const props = defineProps<{
|
||||||
modelValue: boolean
|
modelValue: boolean
|
||||||
providerId: string
|
providerId: string
|
||||||
@@ -407,17 +418,32 @@ const totalPages = computed(() => Math.max(1, Math.ceil(filteredTotal.value / PA
|
|||||||
const isAllFilteredSelected = computed(() => selectAllFiltered.value && filteredTotal.value > 0)
|
const isAllFilteredSelected = computed(() => selectAllFiltered.value && filteredTotal.value > 0)
|
||||||
const isPartiallyFilteredSelected = computed(() => !selectAllFiltered.value && selectedKeyIds.value.length > 0)
|
const isPartiallyFilteredSelected = computed(() => !selectAllFiltered.value && selectedKeyIds.value.length > 0)
|
||||||
const hasActiveFilters = computed(() => searchText.value.trim().length > 0 || activeQuickSelectors.value.length > 0)
|
const hasActiveFilters = computed(() => searchText.value.trim().length > 0 || activeQuickSelectors.value.length > 0)
|
||||||
|
const pageKeyRows = computed<PageKeyRow[]>(() => pageKeys.value.map((key) => {
|
||||||
|
const statusBadgeLabel = getStatusBadgeLabel(key)
|
||||||
|
const quotaText = getQuotaText(key)
|
||||||
|
|
||||||
|
return {
|
||||||
|
key,
|
||||||
|
authTypeLabel: normalizeAuthTypeLabel(key),
|
||||||
|
statusBadgeLabel,
|
||||||
|
statusBadgeTitle: statusBadgeLabel ? getStatusBadgeTitle(key) : '',
|
||||||
|
oauthOrgBadge: getOAuthOrgBadge(key),
|
||||||
|
quotaText,
|
||||||
|
quotaTextShort: quotaText ? shortenQuota(quotaText) : '',
|
||||||
|
lastUsedRelative: key.last_used_at ? formatRelativeTime(key.last_used_at) : '',
|
||||||
|
}
|
||||||
|
}))
|
||||||
const selectedOnCurrentPageCount = computed(() => {
|
const selectedOnCurrentPageCount = computed(() => {
|
||||||
if (selectAllFiltered.value) return pageKeys.value.length
|
if (selectAllFiltered.value) return pageKeyRows.value.length
|
||||||
let count = 0
|
let count = 0
|
||||||
for (const key of pageKeys.value) {
|
for (const row of pageKeyRows.value) {
|
||||||
if (selectedIdSet.value.has(key.key_id)) count += 1
|
if (selectedIdSet.value.has(row.key.key_id)) count += 1
|
||||||
}
|
}
|
||||||
return count
|
return count
|
||||||
})
|
})
|
||||||
const isCurrentPageFullySelected = computed(() => {
|
const isCurrentPageFullySelected = computed(() => {
|
||||||
if (selectAllFiltered.value || pageKeys.value.length === 0) return false
|
if (selectAllFiltered.value || pageKeyRows.value.length === 0) return false
|
||||||
return selectedOnCurrentPageCount.value === pageKeys.value.length
|
return selectedOnCurrentPageCount.value === pageKeyRows.value.length
|
||||||
})
|
})
|
||||||
const canClearSelection = computed(() => selectAllFiltered.value || selectedKeyIds.value.length > 0)
|
const canClearSelection = computed(() => selectAllFiltered.value || selectedKeyIds.value.length > 0)
|
||||||
const activeQuickSelectorSet = computed(() => new Set(activeQuickSelectors.value))
|
const activeQuickSelectorSet = computed(() => new Set(activeQuickSelectors.value))
|
||||||
|
|||||||
@@ -40,4 +40,43 @@ describe('getOAuthOrgBadge', () => {
|
|||||||
|
|
||||||
expect(badge).toBeNull()
|
expect(badge).toBeNull()
|
||||||
})
|
})
|
||||||
|
|
||||||
|
it('reuses cached badge objects when identity fields are unchanged', () => {
|
||||||
|
const identity = {
|
||||||
|
oauth_account_id: 'acct-demo-001',
|
||||||
|
oauth_account_name: 'Workspace Alpha',
|
||||||
|
oauth_account_user_id: 'user-1__acct-demo-001',
|
||||||
|
oauth_organizations: [
|
||||||
|
{ id: 'org-personal-1234', title: 'Personal', is_default: true },
|
||||||
|
],
|
||||||
|
}
|
||||||
|
|
||||||
|
const first = getOAuthOrgBadge(identity)
|
||||||
|
const second = getOAuthOrgBadge(identity)
|
||||||
|
|
||||||
|
expect(first).toBe(second)
|
||||||
|
})
|
||||||
|
|
||||||
|
it('invalidates the cached badge when the selected organization changes in place', () => {
|
||||||
|
const organizations = [
|
||||||
|
{ id: 'org-personal-1234', title: 'Personal', is_default: true },
|
||||||
|
{ id: 'org-team-5678', title: 'Team', is_default: false },
|
||||||
|
]
|
||||||
|
const identity = {
|
||||||
|
oauth_account_id: 'acct-demo-001',
|
||||||
|
oauth_organizations: organizations,
|
||||||
|
}
|
||||||
|
|
||||||
|
const first = getOAuthOrgBadge(identity)
|
||||||
|
organizations[0].is_default = false
|
||||||
|
organizations[1].is_default = true
|
||||||
|
const second = getOAuthOrgBadge(identity)
|
||||||
|
|
||||||
|
expect(first).not.toBe(second)
|
||||||
|
expect(second).toEqual({
|
||||||
|
id: 'org-team-5678',
|
||||||
|
label: 'org:team-5678',
|
||||||
|
title: 'account_id: acct-demo-001 | org_id: org-team-5678 | org_title: Team',
|
||||||
|
})
|
||||||
|
})
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -7,23 +7,65 @@ type OAuthIdentityDisplayValue = {
|
|||||||
oauth_organizations?: OAuthOrganizationInfo[] | null
|
oauth_organizations?: OAuthOrganizationInfo[] | null
|
||||||
} | null | undefined
|
} | null | undefined
|
||||||
|
|
||||||
function formatOAuthIdentityShort(
|
type OAuthOrgBadge = {
|
||||||
value: string | null | undefined,
|
id: string
|
||||||
head = 8,
|
label: string
|
||||||
tail = 6,
|
title: string
|
||||||
): string {
|
|
||||||
const normalized = String(value || '').trim()
|
|
||||||
if (!normalized) return ''
|
|
||||||
if (normalized.length <= head + tail + 3) return normalized
|
|
||||||
return `${normalized.slice(0, head)}...${normalized.slice(-tail)}`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type SelectedOAuthOrganization = {
|
||||||
|
id: string
|
||||||
|
title: string
|
||||||
|
rawId: unknown
|
||||||
|
rawTitle: unknown
|
||||||
|
rawIsDefault: boolean | null | undefined
|
||||||
|
index: number
|
||||||
|
}
|
||||||
|
|
||||||
|
type OAuthOrgBadgeCacheEntry = {
|
||||||
|
rawAccountId: unknown
|
||||||
|
rawAccountName: unknown
|
||||||
|
rawAccountUserId: unknown
|
||||||
|
rawOrganizations: OAuthOrganizationInfo[] | null
|
||||||
|
organizationCount: number
|
||||||
|
selectedOrganizationIndex: number
|
||||||
|
selectedOrganizationRawId: unknown
|
||||||
|
selectedOrganizationRawTitle: unknown
|
||||||
|
selectedOrganizationRawIsDefault: boolean | null | undefined
|
||||||
|
result: OAuthOrgBadge | null
|
||||||
|
}
|
||||||
|
|
||||||
|
const oauthOrgBadgeCache = new WeakMap<object, OAuthOrgBadgeCacheEntry>()
|
||||||
|
|
||||||
function readStr(raw: unknown): string {
|
function readStr(raw: unknown): string {
|
||||||
return typeof raw === 'string' ? raw.trim() : ''
|
return typeof raw === 'string' ? raw.trim() : ''
|
||||||
}
|
}
|
||||||
|
|
||||||
|
function formatOAuthIdentityShort(value: string, head = 8, tail = 6): string {
|
||||||
|
if (!value) return ''
|
||||||
|
if (value.length <= head + tail + 3) return value
|
||||||
|
return `${value.slice(0, head)}...${value.slice(-tail)}`
|
||||||
|
}
|
||||||
|
|
||||||
function stripOAuthOrganizationPrefix(orgId: string): string {
|
function stripOAuthOrganizationPrefix(orgId: string): string {
|
||||||
return orgId.replace(/^org[-_:]+/i, '').trim()
|
if (orgId.length < 4) return orgId
|
||||||
|
|
||||||
|
const first = orgId.charCodeAt(0) | 32
|
||||||
|
const second = orgId.charCodeAt(1) | 32
|
||||||
|
const third = orgId.charCodeAt(2) | 32
|
||||||
|
|
||||||
|
if (first !== 111 || second !== 114 || third !== 103) {
|
||||||
|
return orgId
|
||||||
|
}
|
||||||
|
|
||||||
|
let index = 3
|
||||||
|
while (index < orgId.length) {
|
||||||
|
const code = orgId.charCodeAt(index)
|
||||||
|
if (code !== 45 && code !== 95 && code !== 58) break
|
||||||
|
index += 1
|
||||||
|
}
|
||||||
|
|
||||||
|
return index === 3 ? orgId : orgId.slice(index)
|
||||||
}
|
}
|
||||||
|
|
||||||
function formatOAuthOrganizationBadge(orgId: string): string {
|
function formatOAuthOrganizationBadge(orgId: string): string {
|
||||||
@@ -45,60 +87,154 @@ function formatOAuthAccountUserBadge(accountUserId: string): string {
|
|||||||
|
|
||||||
function getPrimaryOAuthOrganization(
|
function getPrimaryOAuthOrganization(
|
||||||
value: OAuthIdentityDisplayValue,
|
value: OAuthIdentityDisplayValue,
|
||||||
): { id: string; title: string } | null {
|
): SelectedOAuthOrganization | null {
|
||||||
const organizations: OAuthOrganizationInfo[] = Array.isArray(value?.oauth_organizations)
|
const organizations = Array.isArray(value?.oauth_organizations)
|
||||||
? value.oauth_organizations
|
? value.oauth_organizations
|
||||||
: []
|
: null
|
||||||
let firstWithId: OAuthOrganizationInfo | null = null
|
if (!organizations?.length) return null
|
||||||
|
|
||||||
|
let firstWithId: SelectedOAuthOrganization | null = null
|
||||||
|
|
||||||
for (let index = 0; index < organizations.length; index += 1) {
|
for (let index = 0; index < organizations.length; index += 1) {
|
||||||
const org = organizations[index]
|
const org = organizations[index]
|
||||||
if (typeof org?.id !== 'string' || !org.id.trim()) continue
|
const rawId = org?.id
|
||||||
if (!firstWithId) firstWithId = org
|
if (typeof rawId !== 'string') continue
|
||||||
|
|
||||||
|
const id = rawId.trim()
|
||||||
|
if (!id) continue
|
||||||
|
|
||||||
|
const selectedOrganization: SelectedOAuthOrganization = {
|
||||||
|
id,
|
||||||
|
title: readStr(org?.title),
|
||||||
|
rawId,
|
||||||
|
rawTitle: org?.title,
|
||||||
|
rawIsDefault: org?.is_default,
|
||||||
|
index,
|
||||||
|
}
|
||||||
|
|
||||||
|
if (!firstWithId) firstWithId = selectedOrganization
|
||||||
if (org.is_default) {
|
if (org.is_default) {
|
||||||
firstWithId = org
|
firstWithId = selectedOrganization
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if (!firstWithId?.id) return null
|
return firstWithId
|
||||||
|
}
|
||||||
|
|
||||||
return {
|
function appendTitlePart(title: string, label: string, value: string): string {
|
||||||
id: firstWithId.id.trim(),
|
if (!value) return title
|
||||||
title: typeof firstWithId.title === 'string' ? firstWithId.title.trim() : '',
|
if (!title) return `${label}: ${value}`
|
||||||
|
return `${title} | ${label}: ${value}`
|
||||||
|
}
|
||||||
|
|
||||||
|
function buildOAuthIdentityTitle(
|
||||||
|
accountName: string,
|
||||||
|
accountId: string,
|
||||||
|
accountUserId: string,
|
||||||
|
organization: SelectedOAuthOrganization | null,
|
||||||
|
): string {
|
||||||
|
let title = ''
|
||||||
|
|
||||||
|
title = appendTitlePart(title, 'name', accountName)
|
||||||
|
title = appendTitlePart(title, 'account_id', accountId)
|
||||||
|
title = appendTitlePart(title, 'account_user_id', accountUserId)
|
||||||
|
title = appendTitlePart(title, 'org_id', organization?.id || '')
|
||||||
|
title = appendTitlePart(title, 'org_title', organization?.title || '')
|
||||||
|
|
||||||
|
return title
|
||||||
|
}
|
||||||
|
|
||||||
|
function getCachedOAuthOrgBadge(
|
||||||
|
value: OAuthIdentityDisplayValue,
|
||||||
|
organization: SelectedOAuthOrganization | null,
|
||||||
|
): OAuthOrgBadge | null | undefined {
|
||||||
|
if (!value || typeof value !== 'object') return undefined
|
||||||
|
|
||||||
|
const cached = oauthOrgBadgeCache.get(value)
|
||||||
|
if (!cached) return undefined
|
||||||
|
|
||||||
|
const organizations = Array.isArray(value.oauth_organizations)
|
||||||
|
? value.oauth_organizations
|
||||||
|
: null
|
||||||
|
|
||||||
|
if (
|
||||||
|
cached.rawAccountId === value.oauth_account_id
|
||||||
|
&& cached.rawAccountName === value.oauth_account_name
|
||||||
|
&& cached.rawAccountUserId === value.oauth_account_user_id
|
||||||
|
&& cached.rawOrganizations === organizations
|
||||||
|
&& cached.organizationCount === (organizations?.length || 0)
|
||||||
|
&& cached.selectedOrganizationIndex === (organization?.index ?? -1)
|
||||||
|
&& cached.selectedOrganizationRawId === organization?.rawId
|
||||||
|
&& cached.selectedOrganizationRawTitle === organization?.rawTitle
|
||||||
|
&& cached.selectedOrganizationRawIsDefault === organization?.rawIsDefault
|
||||||
|
) {
|
||||||
|
return cached.result
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return undefined
|
||||||
|
}
|
||||||
|
|
||||||
|
function setCachedOAuthOrgBadge(
|
||||||
|
value: OAuthIdentityDisplayValue,
|
||||||
|
organization: SelectedOAuthOrganization | null,
|
||||||
|
result: OAuthOrgBadge | null,
|
||||||
|
): void {
|
||||||
|
if (!value || typeof value !== 'object') return
|
||||||
|
|
||||||
|
const organizations = Array.isArray(value.oauth_organizations)
|
||||||
|
? value.oauth_organizations
|
||||||
|
: null
|
||||||
|
|
||||||
|
oauthOrgBadgeCache.set(value, {
|
||||||
|
rawAccountId: value.oauth_account_id,
|
||||||
|
rawAccountName: value.oauth_account_name,
|
||||||
|
rawAccountUserId: value.oauth_account_user_id,
|
||||||
|
rawOrganizations: organizations,
|
||||||
|
organizationCount: organizations?.length || 0,
|
||||||
|
selectedOrganizationIndex: organization?.index ?? -1,
|
||||||
|
selectedOrganizationRawId: organization?.rawId,
|
||||||
|
selectedOrganizationRawTitle: organization?.rawTitle,
|
||||||
|
selectedOrganizationRawIsDefault: organization?.rawIsDefault,
|
||||||
|
result,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
export function getOAuthOrgBadge(
|
export function getOAuthOrgBadge(
|
||||||
value: OAuthIdentityDisplayValue,
|
value: OAuthIdentityDisplayValue,
|
||||||
): { id: string; label: string; title: string } | null {
|
): OAuthOrgBadge | null {
|
||||||
const org = getPrimaryOAuthOrganization(value)
|
const org = getPrimaryOAuthOrganization(value)
|
||||||
|
const cached = getCachedOAuthOrgBadge(value, org)
|
||||||
|
if (cached !== undefined) return cached
|
||||||
|
|
||||||
const accountId = readStr(value?.oauth_account_id)
|
const accountId = readStr(value?.oauth_account_id)
|
||||||
const accountName = readStr(value?.oauth_account_name)
|
const accountName = readStr(value?.oauth_account_name)
|
||||||
const accountUserId = readStr(value?.oauth_account_user_id)
|
const accountUserId = readStr(value?.oauth_account_user_id)
|
||||||
|
|
||||||
const badgeId = org?.id || accountId || accountUserId || ''
|
const badgeId = org?.id || accountId || accountUserId || ''
|
||||||
const label = org?.id
|
if (!badgeId) {
|
||||||
|
setCachedOAuthOrgBadge(value, org, null)
|
||||||
|
return null
|
||||||
|
}
|
||||||
|
|
||||||
|
const label = org
|
||||||
? formatOAuthOrganizationBadge(org.id)
|
? formatOAuthOrganizationBadge(org.id)
|
||||||
: accountId
|
: accountId
|
||||||
? formatOAuthAccountBadge(accountId)
|
? formatOAuthAccountBadge(accountId)
|
||||||
: accountUserId
|
: accountUserId
|
||||||
? formatOAuthAccountUserBadge(accountUserId)
|
? formatOAuthAccountUserBadge(accountUserId)
|
||||||
: ''
|
: ''
|
||||||
if (!badgeId || !label) return null
|
if (!label) {
|
||||||
|
setCachedOAuthOrgBadge(value, org, null)
|
||||||
|
return null
|
||||||
|
}
|
||||||
|
|
||||||
const titleParts = [
|
const result: OAuthOrgBadge = {
|
||||||
accountName ? `name: ${accountName}` : '',
|
|
||||||
accountId ? `account_id: ${accountId}` : '',
|
|
||||||
accountUserId ? `account_user_id: ${accountUserId}` : '',
|
|
||||||
org?.id ? `org_id: ${org.id}` : '',
|
|
||||||
org?.title ? `org_title: ${org.title}` : '',
|
|
||||||
].filter(Boolean)
|
|
||||||
|
|
||||||
return {
|
|
||||||
id: badgeId,
|
id: badgeId,
|
||||||
label,
|
label,
|
||||||
title: titleParts.join(' | '),
|
title: buildOAuthIdentityTitle(accountName, accountId, accountUserId, org),
|
||||||
}
|
}
|
||||||
|
|
||||||
|
setCachedOAuthOrgBadge(value, org, result)
|
||||||
|
return result
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -406,7 +406,7 @@
|
|||||||
v-for="key in keyPage.keys"
|
v-for="key in keyPage.keys"
|
||||||
:key="key.key_id"
|
:key="key.key_id"
|
||||||
class="border-b border-border/40 last:border-b-0 hover:bg-muted/30 transition-colors"
|
class="border-b border-border/40 last:border-b-0 hover:bg-muted/30 transition-colors"
|
||||||
:class="getRowClass(key)"
|
:class="keyUiStateMap[key.key_id]?.rowClass || ''"
|
||||||
>
|
>
|
||||||
<TableCell
|
<TableCell
|
||||||
class="py-3"
|
class="py-3"
|
||||||
@@ -469,7 +469,7 @@
|
|||||||
size="icon"
|
size="icon"
|
||||||
class="h-4 w-4 shrink-0"
|
class="h-4 w-4 shrink-0"
|
||||||
:disabled="refreshingOAuthKeyId === key.key_id"
|
:disabled="refreshingOAuthKeyId === key.key_id"
|
||||||
:title="getOAuthRefreshButtonTitle(key)"
|
:title="keyUiStateMap[key.key_id]?.oauthRefreshButtonTitle || ''"
|
||||||
@click.stop="handleRefreshOAuth(key)"
|
@click.stop="handleRefreshOAuth(key)"
|
||||||
>
|
>
|
||||||
<RefreshCw
|
<RefreshCw
|
||||||
@@ -478,33 +478,33 @@
|
|||||||
/>
|
/>
|
||||||
</Button>
|
</Button>
|
||||||
<span
|
<span
|
||||||
v-if="getVisibleOAuthState(key)"
|
v-if="keyUiStateMap[key.key_id]?.visibleOAuthState"
|
||||||
class="text-[10px]"
|
class="text-[10px]"
|
||||||
:class="{
|
:class="{
|
||||||
'text-destructive': getVisibleOAuthState(key)?.isInvalid || getVisibleOAuthState(key)?.isExpired,
|
'text-destructive': keyUiStateMap[key.key_id]?.visibleOAuthState?.isInvalid || keyUiStateMap[key.key_id]?.visibleOAuthState?.isExpired,
|
||||||
'text-warning': getVisibleOAuthState(key)?.isExpiringSoon && !getVisibleOAuthState(key)?.isExpired && !getVisibleOAuthState(key)?.isInvalid,
|
'text-warning': keyUiStateMap[key.key_id]?.visibleOAuthState?.isExpiringSoon && !keyUiStateMap[key.key_id]?.visibleOAuthState?.isExpired && !keyUiStateMap[key.key_id]?.visibleOAuthState?.isInvalid,
|
||||||
'text-muted-foreground': !getVisibleOAuthState(key)?.isExpired && !getVisibleOAuthState(key)?.isExpiringSoon && !getVisibleOAuthState(key)?.isInvalid
|
'text-muted-foreground': !keyUiStateMap[key.key_id]?.visibleOAuthState?.isExpired && !keyUiStateMap[key.key_id]?.visibleOAuthState?.isExpiringSoon && !keyUiStateMap[key.key_id]?.visibleOAuthState?.isInvalid
|
||||||
}"
|
}"
|
||||||
:title="getOAuthStatusTitle(key)"
|
:title="keyUiStateMap[key.key_id]?.oauthStatusTitle || ''"
|
||||||
>
|
>
|
||||||
{{ getVisibleOAuthState(key)?.text }}
|
{{ keyUiStateMap[key.key_id]?.visibleOAuthState?.text }}
|
||||||
</span>
|
</span>
|
||||||
</template>
|
</template>
|
||||||
<Badge
|
<Badge
|
||||||
v-if="key.oauth_plan_type"
|
v-if="key.oauth_plan_type"
|
||||||
variant="outline"
|
variant="outline"
|
||||||
class="text-[9px] px-1 py-0 h-4 shrink-0"
|
class="text-[9px] px-1 py-0 h-4 shrink-0"
|
||||||
:class="getOAuthPlanTypeClass(key.oauth_plan_type)"
|
:class="keyUiStateMap[key.key_id]?.planClass || ''"
|
||||||
>
|
>
|
||||||
{{ formatOAuthPlanType(key.oauth_plan_type) }}
|
{{ keyUiStateMap[key.key_id]?.planLabel }}
|
||||||
</Badge>
|
</Badge>
|
||||||
<Badge
|
<Badge
|
||||||
v-if="getOAuthOrgBadge(key)"
|
v-if="keyUiStateMap[key.key_id]?.oauthOrgBadge"
|
||||||
variant="secondary"
|
variant="secondary"
|
||||||
class="text-[9px] px-1 py-0 h-4 shrink-0"
|
class="text-[9px] px-1 py-0 h-4 shrink-0"
|
||||||
:title="getOAuthOrgBadge(key)?.title"
|
:title="keyUiStateMap[key.key_id]?.oauthOrgBadge?.title"
|
||||||
>
|
>
|
||||||
{{ getOAuthOrgBadge(key)?.label }}
|
{{ keyUiStateMap[key.key_id]?.oauthOrgBadge?.label }}
|
||||||
</Badge>
|
</Badge>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
@@ -546,10 +546,10 @@
|
|||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<span
|
<span
|
||||||
v-else-if="getQuotaFallbackText(key)"
|
v-else-if="keyUiStateMap[key.key_id]?.quotaFallbackText"
|
||||||
:class="getQuotaTextClass(getQuotaFallbackText(key) || '')"
|
:class="keyUiStateMap[key.key_id]?.quotaTextClass || ''"
|
||||||
>
|
>
|
||||||
{{ getQuotaFallbackText(key) }}
|
{{ keyUiStateMap[key.key_id]?.quotaFallbackText }}
|
||||||
</span>
|
</span>
|
||||||
<span
|
<span
|
||||||
v-else
|
v-else
|
||||||
@@ -580,16 +580,16 @@
|
|||||||
</TableCell>
|
</TableCell>
|
||||||
<TableCell class="py-3 text-center">
|
<TableCell class="py-3 text-center">
|
||||||
<span class="text-[10px] text-muted-foreground whitespace-nowrap">
|
<span class="text-[10px] text-muted-foreground whitespace-nowrap">
|
||||||
{{ key.last_used_at ? formatRelativeTime(key.last_used_at) : '-' }}
|
{{ keyUiStateMap[key.key_id]?.lastUsedRelative || '-' }}
|
||||||
</span>
|
</span>
|
||||||
</TableCell>
|
</TableCell>
|
||||||
<TableCell class="py-3 text-center">
|
<TableCell class="py-3 text-center">
|
||||||
<Badge
|
<Badge
|
||||||
:variant="getSchedulingBadgeVariant(key)"
|
:variant="keyUiStateMap[key.key_id]?.schedulingBadgeVariant || 'default'"
|
||||||
class="text-[10px]"
|
class="text-[10px]"
|
||||||
:title="getSchedulingTitle(key)"
|
:title="keyUiStateMap[key.key_id]?.schedulingTitle || ''"
|
||||||
>
|
>
|
||||||
{{ getSchedulingBadgeLabel(key) }}
|
{{ keyUiStateMap[key.key_id]?.schedulingBadgeLabel }}
|
||||||
</Badge>
|
</Badge>
|
||||||
</TableCell>
|
</TableCell>
|
||||||
<TableCell class="py-3 px-2 align-middle">
|
<TableCell class="py-3 px-2 align-middle">
|
||||||
@@ -705,7 +705,7 @@
|
|||||||
v-for="key in keyPage.keys"
|
v-for="key in keyPage.keys"
|
||||||
:key="key.key_id"
|
:key="key.key_id"
|
||||||
class="p-4 sm:p-5 hover:bg-muted/30 transition-colors"
|
class="p-4 sm:p-5 hover:bg-muted/30 transition-colors"
|
||||||
:class="getRowClass(key)"
|
:class="keyUiStateMap[key.key_id]?.rowClass || ''"
|
||||||
>
|
>
|
||||||
<div class="space-y-3">
|
<div class="space-y-3">
|
||||||
<div class="text-sm font-medium truncate">
|
<div class="text-sm font-medium truncate">
|
||||||
@@ -714,11 +714,11 @@
|
|||||||
|
|
||||||
<div class="flex flex-wrap items-center gap-1.5">
|
<div class="flex flex-wrap items-center gap-1.5">
|
||||||
<Badge
|
<Badge
|
||||||
:variant="getSchedulingBadgeVariant(key)"
|
:variant="keyUiStateMap[key.key_id]?.schedulingBadgeVariant || 'default'"
|
||||||
class="text-[10px] shrink-0"
|
class="text-[10px] shrink-0"
|
||||||
:title="getSchedulingTitle(key)"
|
:title="keyUiStateMap[key.key_id]?.schedulingTitle || ''"
|
||||||
>
|
>
|
||||||
{{ getSchedulingBadgeLabel(key) }}
|
{{ keyUiStateMap[key.key_id]?.schedulingBadgeLabel }}
|
||||||
</Badge>
|
</Badge>
|
||||||
<span
|
<span
|
||||||
v-if="key.cooldown_ttl_seconds"
|
v-if="key.cooldown_ttl_seconds"
|
||||||
@@ -727,7 +727,7 @@
|
|||||||
冷却 {{ formatTTL(key.cooldown_ttl_seconds) }}
|
冷却 {{ formatTTL(key.cooldown_ttl_seconds) }}
|
||||||
</span>
|
</span>
|
||||||
<template
|
<template
|
||||||
v-for="item in getMobileTagItems(key)"
|
v-for="item in keyUiStateMap[key.key_id]?.mobileTagItems || []"
|
||||||
:key="`${key.key_id}-${item.key}`"
|
:key="`${key.key_id}-${item.key}`"
|
||||||
>
|
>
|
||||||
<button
|
<button
|
||||||
@@ -744,7 +744,7 @@
|
|||||||
v-else-if="item.key === 'plan'"
|
v-else-if="item.key === 'plan'"
|
||||||
variant="outline"
|
variant="outline"
|
||||||
class="text-[9px] px-1 py-0 h-4 shrink-0"
|
class="text-[9px] px-1 py-0 h-4 shrink-0"
|
||||||
:class="key.oauth_plan_type ? getOAuthPlanTypeClass(key.oauth_plan_type) : ''"
|
:class="keyUiStateMap[key.key_id]?.planClass || ''"
|
||||||
>
|
>
|
||||||
{{ item.label }}
|
{{ item.label }}
|
||||||
</Badge>
|
</Badge>
|
||||||
@@ -752,7 +752,7 @@
|
|||||||
v-else-if="item.key === 'org'"
|
v-else-if="item.key === 'org'"
|
||||||
variant="secondary"
|
variant="secondary"
|
||||||
class="text-[9px] px-1 py-0 h-4 shrink-0"
|
class="text-[9px] px-1 py-0 h-4 shrink-0"
|
||||||
:title="getOAuthOrgBadge(key)?.title"
|
:title="keyUiStateMap[key.key_id]?.oauthOrgBadge?.title"
|
||||||
>
|
>
|
||||||
{{ item.label }}
|
{{ item.label }}
|
||||||
</Badge>
|
</Badge>
|
||||||
@@ -775,7 +775,7 @@
|
|||||||
<span class="mx-1.5 text-muted-foreground/40">|</span>
|
<span class="mx-1.5 text-muted-foreground/40">|</span>
|
||||||
<span class="font-medium text-foreground/90">费用:{{ formatStatUsd(key.total_cost_usd) }}</span>
|
<span class="font-medium text-foreground/90">费用:{{ formatStatUsd(key.total_cost_usd) }}</span>
|
||||||
<span class="mx-1.5 text-muted-foreground/40">|</span>
|
<span class="mx-1.5 text-muted-foreground/40">|</span>
|
||||||
<span class="font-medium text-foreground/90">最后使用:{{ key.last_used_at ? formatRelativeTime(key.last_used_at) : '-' }}</span>
|
<span class="font-medium text-foreground/90">最后使用:{{ keyUiStateMap[key.key_id]?.lastUsedRelative || '-' }}</span>
|
||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
@@ -819,10 +819,10 @@
|
|||||||
</div>
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<div
|
<div
|
||||||
v-else-if="getQuotaFallbackText(key)"
|
v-else-if="keyUiStateMap[key.key_id]?.quotaFallbackText"
|
||||||
:class="getQuotaTextClass(getQuotaFallbackText(key) || '')"
|
:class="keyUiStateMap[key.key_id]?.quotaTextClass || ''"
|
||||||
>
|
>
|
||||||
{{ getQuotaFallbackText(key) }}
|
{{ keyUiStateMap[key.key_id]?.quotaFallbackText }}
|
||||||
</div>
|
</div>
|
||||||
<div
|
<div
|
||||||
v-else
|
v-else
|
||||||
@@ -834,7 +834,7 @@
|
|||||||
|
|
||||||
<div class="flex items-center gap-0.5">
|
<div class="flex items-center gap-0.5">
|
||||||
<div
|
<div
|
||||||
v-for="actionId in getMobileActionIds(key)"
|
v-for="actionId in keyUiStateMap[key.key_id]?.mobileActionIds || []"
|
||||||
:key="`${key.key_id}-${actionId}`"
|
:key="`${key.key_id}-${actionId}`"
|
||||||
class="min-w-0 flex-1 flex justify-center"
|
class="min-w-0 flex-1 flex justify-center"
|
||||||
>
|
>
|
||||||
@@ -864,7 +864,7 @@
|
|||||||
size="icon"
|
size="icon"
|
||||||
class="h-7 w-7 shrink-0"
|
class="h-7 w-7 shrink-0"
|
||||||
:disabled="refreshingOAuthKeyId === key.key_id"
|
:disabled="refreshingOAuthKeyId === key.key_id"
|
||||||
:title="getOAuthRefreshButtonTitle(key)"
|
:title="keyUiStateMap[key.key_id]?.oauthRefreshButtonTitle || ''"
|
||||||
@click.stop="handleRefreshOAuth(key)"
|
@click.stop="handleRefreshOAuth(key)"
|
||||||
>
|
>
|
||||||
<RefreshCw
|
<RefreshCw
|
||||||
@@ -1626,6 +1626,24 @@ interface QuotaProgressItem {
|
|||||||
updatedAtSeconds?: number | null
|
updatedAtSeconds?: number | null
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type PoolKeyUiState = {
|
||||||
|
rowClass: string
|
||||||
|
schedulingBadgeLabel: string
|
||||||
|
schedulingBadgeVariant: PoolStatusVariant
|
||||||
|
schedulingTitle: string
|
||||||
|
oauthOrgBadge: ReturnType<typeof getOAuthOrgBadge>
|
||||||
|
visibleOAuthState: ReturnType<typeof getOAuthStatusDisplay>
|
||||||
|
oauthStatusTitle: string
|
||||||
|
oauthRefreshButtonTitle: string
|
||||||
|
planLabel: string
|
||||||
|
planClass: string
|
||||||
|
quotaFallbackText: string | null
|
||||||
|
quotaTextClass: string
|
||||||
|
lastUsedRelative: string
|
||||||
|
mobileTagItems: PoolMobileTagItem[]
|
||||||
|
mobileActionIds: PoolMobileActionId[]
|
||||||
|
}
|
||||||
|
|
||||||
const quotaProgressMap = computed<Record<string, QuotaProgressItem[]>>(() => {
|
const quotaProgressMap = computed<Record<string, QuotaProgressItem[]>>(() => {
|
||||||
const map: Record<string, QuotaProgressItem[]> = {}
|
const map: Record<string, QuotaProgressItem[]> = {}
|
||||||
for (const key of keyPage.value.keys) {
|
for (const key of keyPage.value.keys) {
|
||||||
@@ -1634,6 +1652,42 @@ const quotaProgressMap = computed<Record<string, QuotaProgressItem[]>>(() => {
|
|||||||
return map
|
return map
|
||||||
})
|
})
|
||||||
|
|
||||||
|
const keyUiStateMap = computed<Record<string, PoolKeyUiState>>(() => {
|
||||||
|
const map: Record<string, PoolKeyUiState> = {}
|
||||||
|
|
||||||
|
for (const key of keyPage.value.keys) {
|
||||||
|
const visibleOAuthState = getVisibleOAuthState(key)
|
||||||
|
const oauthOrgBadge = getOAuthOrgBadge(key)
|
||||||
|
const quotaFallbackText = getQuotaFallbackText(key)
|
||||||
|
const canRefreshToken = canRefreshOAuthCredential(key)
|
||||||
|
|
||||||
|
map[key.key_id] = {
|
||||||
|
rowClass: getRowClass(key),
|
||||||
|
schedulingBadgeLabel: getSchedulingBadgeLabel(key),
|
||||||
|
schedulingBadgeVariant: getSchedulingBadgeVariant(key),
|
||||||
|
schedulingTitle: getSchedulingTitle(key),
|
||||||
|
oauthOrgBadge,
|
||||||
|
visibleOAuthState,
|
||||||
|
oauthStatusTitle: visibleOAuthState ? getOAuthStatusTitle(key) : '',
|
||||||
|
oauthRefreshButtonTitle: canRefreshToken ? getOAuthRefreshButtonTitle(key) : '',
|
||||||
|
planLabel: key.oauth_plan_type ? formatOAuthPlanType(key.oauth_plan_type) : '',
|
||||||
|
planClass: key.oauth_plan_type ? getOAuthPlanTypeClass(key.oauth_plan_type) : '',
|
||||||
|
quotaFallbackText,
|
||||||
|
quotaTextClass: quotaFallbackText ? getQuotaTextClass(quotaFallbackText) : '',
|
||||||
|
lastUsedRelative: key.last_used_at ? formatRelativeTime(key.last_used_at) : '-',
|
||||||
|
mobileTagItems: getMobileTagItems(key),
|
||||||
|
mobileActionIds: splitPoolMobileActions({
|
||||||
|
canDownloadOrCopy: true,
|
||||||
|
canRefreshToken,
|
||||||
|
canClearCooldown: Boolean(key.cooldown_reason),
|
||||||
|
hasProxy: true,
|
||||||
|
}).primary,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return map
|
||||||
|
})
|
||||||
|
|
||||||
const quotaRefreshSupported = computed(() => {
|
const quotaRefreshSupported = computed(() => {
|
||||||
return selectedProviderType.value === 'codex'
|
return selectedProviderType.value === 'codex'
|
||||||
|| selectedProviderType.value === 'kiro'
|
|| selectedProviderType.value === 'kiro'
|
||||||
@@ -2506,15 +2560,6 @@ function getMobileTagItems(key: PoolKeyDetail): PoolMobileTagItem[] {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
function getMobileActionIds(key: PoolKeyDetail): PoolMobileActionId[] {
|
|
||||||
return splitPoolMobileActions({
|
|
||||||
canDownloadOrCopy: true,
|
|
||||||
canRefreshToken: canRefreshOAuthCredential(key),
|
|
||||||
canClearCooldown: Boolean(key.cooldown_reason),
|
|
||||||
hasProxy: true,
|
|
||||||
}).primary
|
|
||||||
}
|
|
||||||
|
|
||||||
function getMobileTagClass(item: PoolMobileTagItem): string {
|
function getMobileTagClass(item: PoolMobileTagItem): string {
|
||||||
if (item.tone === 'danger') {
|
if (item.tone === 'danger') {
|
||||||
return 'border-red-500/30 bg-red-500/10 text-red-700 dark:text-red-300'
|
return 'border-red-500/30 bg-red-500/10 text-red-700 dark:text-red-300'
|
||||||
|
|||||||
Reference in New Issue
Block a user