mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
Merge origin/main into codex/gemini-embedding-batch
# Conflicts: # apps/aether-gateway/src/ai_serving/api.rs # apps/aether-gateway/src/ai_serving/planner/passthrough/provider/family/request.rs # apps/aether-gateway/src/ai_serving/planner/standard/family/request.rs # apps/aether-gateway/src/ai_serving/transport.rs # apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/summary.rs # apps/aether-gateway/src/handlers/admin/provider/query/models/model_test/tests.rs # crates/aether-data/src/repository/candidate_selection/postgres.rs # crates/aether-model-fetch/src/strategy.rs
This commit is contained in:
@@ -23,6 +23,7 @@ aether-oauth.workspace = true
|
||||
aether-pool-core.workspace = true
|
||||
aether-provider-pool.workspace = true
|
||||
aether-provider-transport.workspace = true
|
||||
aether-routing-core.workspace = true
|
||||
aether-scheduler-core.workspace = true
|
||||
aether-runtime.workspace = true
|
||||
aether-runtime-state.workspace = true
|
||||
@@ -64,6 +65,8 @@ tracing.workspace = true
|
||||
url.workspace = true
|
||||
uuid.workspace = true
|
||||
webpki-roots.workspace = true
|
||||
wreq.workspace = true
|
||||
wreq-util.workspace = true
|
||||
|
||||
[target.'cfg(not(target_env = "msvc"))'.dependencies]
|
||||
tikv-jemallocator = "0.6"
|
||||
|
||||
@@ -2,14 +2,20 @@ use std::env;
|
||||
use std::process::Command;
|
||||
|
||||
fn main() {
|
||||
println!("cargo:rerun-if-env-changed=AETHER_BUILD_VERSION");
|
||||
println!("cargo:rerun-if-env-changed=AETHER_VERSION");
|
||||
println!("cargo:rerun-if-env-changed=GITHUB_REF_NAME");
|
||||
println!("cargo:rerun-if-changed=../../.git/HEAD");
|
||||
|
||||
let package_version = env::var("CARGO_PKG_VERSION").unwrap_or_else(|_| "unknown".to_string());
|
||||
let version = env::var("AETHER_VERSION")
|
||||
let version = env::var("AETHER_BUILD_VERSION")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.or_else(|| {
|
||||
env::var("AETHER_VERSION")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
})
|
||||
.or_else(|| {
|
||||
env::var("GITHUB_REF_NAME")
|
||||
.ok()
|
||||
|
||||
@@ -41,26 +41,30 @@ pub(crate) use crate::ai_serving::{
|
||||
AiExecutionDecision, AiExecutionPlanPayload, AiStreamAttempt, AiSyncAttempt,
|
||||
};
|
||||
pub(crate) use aether_ai_formats::api::{
|
||||
build_core_error_body_for_client_format, core_error_background_report_kind,
|
||||
core_error_default_client_api_format, core_success_background_report_kind,
|
||||
encode_kiro_sse_events, implicit_sync_finalize_report_kind, is_core_error_finalize_kind,
|
||||
build_core_error_body_for_client_format, convert_standard_chat_response,
|
||||
core_error_background_report_kind, core_error_default_client_api_format,
|
||||
core_success_background_report_kind, encode_kiro_sse_events,
|
||||
implicit_sync_finalize_report_kind, is_core_error_finalize_kind,
|
||||
normalize_provider_private_report_context, normalize_provider_private_response_value,
|
||||
provider_private_response_allows_sync_finalize, resolve_claude_stream_spec,
|
||||
resolve_claude_sync_spec, resolve_gemini_stream_spec, resolve_gemini_sync_spec,
|
||||
resolve_local_image_stream_spec, resolve_local_image_sync_spec,
|
||||
resolve_local_same_format_stream_spec, resolve_local_same_format_sync_spec,
|
||||
resolve_openai_embedding_sync_spec, AiControlPlanRequest, ExecutionRuntimeAuthContext,
|
||||
LocalCoreSyncErrorKind, LocalOpenAiImageSpec, LocalSameFormatProviderFamily,
|
||||
LocalSameFormatProviderSpec, LocalStandardSourceFamily, LocalStandardSourceMode,
|
||||
LocalStandardSpec, StreamingStandardTerminalObserver, EXECUTION_RUNTIME_STREAM_DECISION_ACTION,
|
||||
EXECUTION_RUNTIME_SYNC_DECISION_ACTION, GEMINI_EMBEDDING_SYNC_PLAN_KIND,
|
||||
GEMINI_FILES_DOWNLOAD_PLAN_KIND, GEMINI_VIDEO_CANCEL_SYNC_PLAN_KIND,
|
||||
OPENAI_EMBEDDING_SYNC_PLAN_KIND, OPENAI_IMAGE_STREAM_PLAN_KIND,
|
||||
OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND, OPENAI_IMAGE_SYNC_PLAN_KIND,
|
||||
OPENAI_RERANK_SYNC_PLAN_KIND, OPENAI_VIDEO_CANCEL_SYNC_PLAN_KIND,
|
||||
resolve_openai_embedding_sync_spec, sanitize_request_path_and_query, AiControlPlanRequest,
|
||||
CanonicalContentPart, CanonicalStreamEvent, CanonicalStreamFrame, ClaudeClientEmitter,
|
||||
ExecutionRuntimeAuthContext, LocalCoreSyncErrorKind, LocalOpenAiImageSpec,
|
||||
LocalSameFormatProviderFamily, LocalSameFormatProviderSpec, LocalStandardSourceFamily,
|
||||
LocalStandardSourceMode, LocalStandardSpec, OpenAIChatClientEmitter,
|
||||
OpenAIResponsesClientEmitter, StreamingStandardTerminalObserver,
|
||||
EXECUTION_RUNTIME_STREAM_DECISION_ACTION, EXECUTION_RUNTIME_SYNC_DECISION_ACTION,
|
||||
GEMINI_EMBEDDING_SYNC_PLAN_KIND, GEMINI_FILES_DOWNLOAD_PLAN_KIND,
|
||||
GEMINI_VIDEO_CANCEL_SYNC_PLAN_KIND, OPENAI_EMBEDDING_SYNC_PLAN_KIND,
|
||||
OPENAI_IMAGE_STREAM_PLAN_KIND, OPENAI_IMAGE_SYNC_FINALIZE_REPORT_KIND,
|
||||
OPENAI_IMAGE_SYNC_PLAN_KIND, OPENAI_RERANK_SYNC_PLAN_KIND, OPENAI_VIDEO_CANCEL_SYNC_PLAN_KIND,
|
||||
OPENAI_VIDEO_CONTENT_PLAN_KIND, OPENAI_VIDEO_DELETE_SYNC_PLAN_KIND,
|
||||
OPENAI_VIDEO_REMIX_SYNC_PLAN_KIND,
|
||||
};
|
||||
pub(crate) use aether_ai_formats::protocol::stream::CanonicalUsage as StreamingCanonicalUsage;
|
||||
|
||||
pub(crate) fn parse_direct_request_body(
|
||||
parts: &http::request::Parts,
|
||||
|
||||
@@ -180,7 +180,9 @@ pub(crate) fn resolve_local_decision_execution_runtime_auth_context(
|
||||
decision: &GatewayControlDecision,
|
||||
) -> Option<ExecutionRuntimeAuthContext> {
|
||||
resolve_decision_execution_runtime_auth_context(decision).filter(|auth_context| {
|
||||
!auth_context.user_id.trim().is_empty() && !auth_context.api_key_id.trim().is_empty()
|
||||
auth_context.access_allowed
|
||||
&& !auth_context.user_id.trim().is_empty()
|
||||
&& !auth_context.api_key_id.trim().is_empty()
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -7,7 +7,13 @@ use aether_ai_serving::{
|
||||
AiCandidatePreselectionOutcome, AiSkippedCandidatePersistencePort,
|
||||
};
|
||||
use aether_dispatch_core::{DispatchSequence, DispatchSequenceItem};
|
||||
use aether_scheduler_core::{ClientSessionAffinity, SchedulerMinimalCandidateSelectionCandidate};
|
||||
use aether_routing_core::{
|
||||
rank_vector_for_candidate, CandidateKind, ResolvedRoutingPolicy, RoutingCandidateFacts,
|
||||
RoutingCandidateTrace, RoutingDecisionTrace,
|
||||
};
|
||||
use aether_scheduler_core::{
|
||||
ClientSessionAffinity, SchedulerMinimalCandidateSelectionCandidate, SchedulerRankingOutcome,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
use std::collections::VecDeque;
|
||||
@@ -19,6 +25,7 @@ use tracing::warn;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::ai_serving::planner::candidate_affinity_cache::remember_scheduler_affinity_for_candidate_at_epoch;
|
||||
use crate::ai_serving::planner::candidate_ranking::scheduler_ordering_config_for_routing_policy;
|
||||
use crate::ai_serving::planner::candidate_resolution::{
|
||||
resolve_and_rank_logical_local_execution_candidates, EligibleLocalExecutionCandidate,
|
||||
LocalExecutionCandidateKind, SkippedLocalExecutionCandidate,
|
||||
@@ -36,7 +43,7 @@ use crate::dispatch::refs::dispatch_ref_for_local_candidate;
|
||||
use crate::handlers::shared::provider_pool::admin_provider_pool_config_from_config_value;
|
||||
use crate::orchestration::{local_attempt_slot_count, ExecutionAttemptIdentity};
|
||||
use crate::scheduler::candidate::API_KEY_CONCURRENCY_LIMIT_SKIP_REASON;
|
||||
use crate::scheduler::config::{read_scheduler_ordering_config, SchedulerSchedulingMode};
|
||||
use crate::scheduler::config::SchedulerSchedulingMode;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
const POOL_KEY_RETRY_INDEX_STRIDE: u32 = 100;
|
||||
@@ -189,6 +196,7 @@ struct GatewayLocalCandidateMaterializationPort<'a, F, G> {
|
||||
auth_snapshot: Option<&'a GatewayAuthApiKeySnapshot>,
|
||||
client_session_affinity: Option<&'a ClientSessionAffinity>,
|
||||
required_capabilities: Option<&'a Value>,
|
||||
routing_policy: Option<&'a ResolvedRoutingPolicy>,
|
||||
sticky_session_token: Option<&'a str>,
|
||||
request_auth_channel: Option<&'a str>,
|
||||
persistence_policy: LocalCandidatePersistencePolicy<'a>,
|
||||
@@ -243,6 +251,7 @@ where
|
||||
self.auth_snapshot,
|
||||
self.client_session_affinity,
|
||||
self.required_capabilities,
|
||||
self.routing_policy,
|
||||
self.sticky_session_token,
|
||||
self.request_auth_channel,
|
||||
self.resolution_mode,
|
||||
@@ -281,6 +290,8 @@ where
|
||||
.skipped
|
||||
.record_runtime_miss_diagnostic,
|
||||
candidates,
|
||||
self.routing_policy,
|
||||
self.client_api_format,
|
||||
self.sticky_session_token,
|
||||
self.requested_model,
|
||||
self.request_auth_channel,
|
||||
@@ -294,6 +305,12 @@ where
|
||||
starting_candidate_index: u32,
|
||||
skipped_candidates: Vec<Self::Skipped>,
|
||||
) -> Result<(), Self::Error> {
|
||||
let skipped_candidates = attach_routing_trace_to_skipped_candidates(
|
||||
self.routing_policy,
|
||||
self.client_api_format,
|
||||
starting_candidate_index,
|
||||
skipped_candidates,
|
||||
);
|
||||
persist_skipped_local_execution_candidates_with_context(
|
||||
self.state.app(),
|
||||
self.trace_id,
|
||||
@@ -432,6 +449,7 @@ pub(crate) async fn materialize_local_execution_candidates_with_serving<F, G>(
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
client_session_affinity: Option<&ClientSessionAffinity>,
|
||||
required_capabilities: Option<&Value>,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
sticky_session_token: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
persistence_policy: LocalCandidatePersistencePolicy<'_>,
|
||||
@@ -445,7 +463,8 @@ where
|
||||
F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync,
|
||||
G: Fn(SkippedLocalExecutionCandidate) -> SkippedLocalExecutionCandidate + Send + Sync,
|
||||
{
|
||||
let scheduler_cache_affinity_enabled = scheduler_cache_affinity_enabled(state).await;
|
||||
let scheduler_cache_affinity_enabled =
|
||||
scheduler_cache_affinity_enabled(state, routing_policy).await;
|
||||
let port = GatewayLocalCandidateMaterializationPort {
|
||||
state,
|
||||
trace_id,
|
||||
@@ -454,6 +473,7 @@ where
|
||||
auth_snapshot,
|
||||
client_session_affinity,
|
||||
required_capabilities,
|
||||
routing_policy,
|
||||
sticky_session_token,
|
||||
request_auth_channel,
|
||||
persistence_policy,
|
||||
@@ -478,6 +498,7 @@ pub(crate) async fn build_local_execution_candidate_attempt_source_with_serving<
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
client_session_affinity: Option<&ClientSessionAffinity>,
|
||||
required_capabilities: Option<&Value>,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
sticky_session_token: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
persistence_policy: LocalCandidatePersistencePolicy<'_>,
|
||||
@@ -491,7 +512,8 @@ where
|
||||
F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync,
|
||||
G: Fn(SkippedLocalExecutionCandidate) -> SkippedLocalExecutionCandidate + Send + Sync,
|
||||
{
|
||||
let scheduler_cache_affinity_enabled = scheduler_cache_affinity_enabled(state).await;
|
||||
let scheduler_cache_affinity_enabled =
|
||||
scheduler_cache_affinity_enabled(state, routing_policy).await;
|
||||
let _ = build_available_extra_data;
|
||||
let (candidates, resolved_skipped) = resolve_and_rank_logical_local_execution_candidates(
|
||||
state,
|
||||
@@ -501,6 +523,7 @@ where
|
||||
auth_snapshot,
|
||||
client_session_affinity,
|
||||
required_capabilities,
|
||||
routing_policy,
|
||||
sticky_session_token,
|
||||
request_auth_channel,
|
||||
resolution_mode,
|
||||
@@ -529,7 +552,12 @@ where
|
||||
trace_id,
|
||||
persistence_policy.skipped,
|
||||
u32::try_from(candidates.len()).unwrap_or(u32::MAX),
|
||||
skipped_candidates,
|
||||
attach_routing_trace_to_skipped_candidates(
|
||||
routing_policy,
|
||||
client_api_format,
|
||||
u32::try_from(candidates.len()).unwrap_or(u32::MAX),
|
||||
skipped_candidates,
|
||||
),
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -542,6 +570,7 @@ where
|
||||
sticky_session_token,
|
||||
requested_model,
|
||||
request_auth_channel,
|
||||
routing_policy,
|
||||
);
|
||||
|
||||
(
|
||||
@@ -559,6 +588,7 @@ fn build_logical_candidate_items<'a>(
|
||||
sticky_session_token: Option<&str>,
|
||||
requested_model: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
) -> (VecDeque<LocalExecutionCandidateAttemptSourceItem<'a>>, u32) {
|
||||
let mut items = VecDeque::new();
|
||||
let mut next_candidate_index = starting_candidate_index;
|
||||
@@ -578,12 +608,13 @@ fn build_logical_candidate_items<'a>(
|
||||
}
|
||||
}
|
||||
LocalExecutionCandidateKind::PoolGroup => {
|
||||
let cursor = PoolKeyCursor::new(
|
||||
let cursor = PoolKeyCursor::new_with_routing_policy(
|
||||
state,
|
||||
candidate,
|
||||
sticky_session_token,
|
||||
requested_model,
|
||||
request_auth_channel,
|
||||
routing_policy,
|
||||
);
|
||||
let cursor = if let Some(trace_id) = trace_id {
|
||||
cursor.with_runtime_miss_diagnostic(trace_id, record_runtime_miss_diagnostic)
|
||||
@@ -615,6 +646,7 @@ pub(crate) async fn build_lazy_requested_model_execution_candidate_attempt_sourc
|
||||
auth_snapshot: &GatewayAuthApiKeySnapshot,
|
||||
client_session_affinity: Option<&ClientSessionAffinity>,
|
||||
required_capabilities: Option<&Value>,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
sticky_session_token: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
persistence_policy: LocalCandidatePersistencePolicy<'_>,
|
||||
@@ -628,7 +660,8 @@ where
|
||||
F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync + 'a,
|
||||
G: Fn(SkippedLocalExecutionCandidate) -> SkippedLocalExecutionCandidate + Send + Sync + 'a,
|
||||
{
|
||||
let scheduler_cache_affinity_enabled = scheduler_cache_affinity_enabled(state).await;
|
||||
let scheduler_cache_affinity_enabled =
|
||||
scheduler_cache_affinity_enabled(state, routing_policy).await;
|
||||
let _ = build_available_extra_data;
|
||||
let decorate_skipped_candidate = Arc::new(decorate_skipped_candidate);
|
||||
let record_runtime_miss_diagnostic = persistence_policy.skipped.record_runtime_miss_diagnostic;
|
||||
@@ -639,6 +672,7 @@ where
|
||||
require_streaming,
|
||||
required_capabilities,
|
||||
auth_snapshot,
|
||||
routing_policy,
|
||||
client_session_affinity,
|
||||
use_api_format_alias_match,
|
||||
key_mode,
|
||||
@@ -652,6 +686,7 @@ where
|
||||
auth_snapshot: auth_snapshot.clone(),
|
||||
client_session_affinity: client_session_affinity.cloned(),
|
||||
required_capabilities: required_capabilities.cloned(),
|
||||
routing_policy: routing_policy.cloned(),
|
||||
sticky_session_token: sticky_session_token.map(str::to_string),
|
||||
request_auth_channel: request_auth_channel.map(str::to_string),
|
||||
skipped_user_id: persistence_policy.skipped.user_id.to_string(),
|
||||
@@ -693,6 +728,7 @@ struct RequestedModelAttemptPageCursor<'a> {
|
||||
auth_snapshot: GatewayAuthApiKeySnapshot,
|
||||
client_session_affinity: Option<ClientSessionAffinity>,
|
||||
required_capabilities: Option<Value>,
|
||||
routing_policy: Option<ResolvedRoutingPolicy>,
|
||||
sticky_session_token: Option<String>,
|
||||
request_auth_channel: Option<String>,
|
||||
skipped_user_id: String,
|
||||
@@ -756,6 +792,7 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
|
||||
Some(&self.auth_snapshot),
|
||||
self.client_session_affinity.as_ref(),
|
||||
self.required_capabilities.as_ref(),
|
||||
self.routing_policy.as_ref(),
|
||||
self.sticky_session_token.as_deref(),
|
||||
self.request_auth_channel.as_deref(),
|
||||
self.resolution_mode,
|
||||
@@ -794,6 +831,7 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
|
||||
self.sticky_session_token.as_deref(),
|
||||
Some(&self.requested_model),
|
||||
self.request_auth_channel.as_deref(),
|
||||
self.routing_policy.as_ref(),
|
||||
);
|
||||
self.next_candidate_index = next_candidate_index
|
||||
.saturating_add(u32::try_from(skipped_candidate_count).unwrap_or(u32::MAX));
|
||||
@@ -814,7 +852,12 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
|
||||
&self.trace_id,
|
||||
skipped_persistence,
|
||||
skipped_starting_candidate_index,
|
||||
skipped_candidates,
|
||||
attach_routing_trace_to_skipped_candidates(
|
||||
self.routing_policy.as_ref(),
|
||||
&self.client_api_format,
|
||||
skipped_starting_candidate_index,
|
||||
skipped_candidates,
|
||||
),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
@@ -858,7 +901,12 @@ impl<'a> RequestedModelAttemptPageCursor<'a> {
|
||||
&self.trace_id,
|
||||
skipped_persistence,
|
||||
self.next_candidate_index,
|
||||
skipped_candidates,
|
||||
attach_routing_trace_to_skipped_candidates(
|
||||
self.routing_policy.as_ref(),
|
||||
&self.client_api_format,
|
||||
self.next_candidate_index,
|
||||
skipped_candidates,
|
||||
),
|
||||
)
|
||||
.await;
|
||||
self.next_candidate_index = self
|
||||
@@ -925,19 +973,14 @@ async fn pop_attempt_from_items(
|
||||
}
|
||||
}
|
||||
|
||||
async fn scheduler_cache_affinity_enabled(state: PlannerAppState<'_>) -> bool {
|
||||
match read_scheduler_ordering_config(state.app()).await {
|
||||
Ok(config) => config.scheduling_mode == SchedulerSchedulingMode::CacheAffinity,
|
||||
Err(error) => {
|
||||
warn!(
|
||||
event_name = "planner_scheduler_affinity_config_load_failed",
|
||||
log_type = "event",
|
||||
error = ?error,
|
||||
"failed to load scheduler config while checking cache affinity mode"
|
||||
);
|
||||
SchedulerSchedulingMode::default() == SchedulerSchedulingMode::CacheAffinity
|
||||
}
|
||||
}
|
||||
async fn scheduler_cache_affinity_enabled(
|
||||
state: PlannerAppState<'_>,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
) -> bool {
|
||||
scheduler_ordering_config_for_routing_policy(state, routing_policy)
|
||||
.await
|
||||
.scheduling_mode
|
||||
== SchedulerSchedulingMode::CacheAffinity
|
||||
}
|
||||
|
||||
pub(crate) fn remember_first_local_candidate_affinity(
|
||||
@@ -1038,6 +1081,8 @@ async fn materialize_logical_local_execution_candidate_attempts<F>(
|
||||
context: LocalAvailableCandidatePersistenceContext<'_>,
|
||||
record_runtime_miss_diagnostic: bool,
|
||||
candidates: Vec<EligibleLocalExecutionCandidate>,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
client_api_format: &str,
|
||||
sticky_session_token: Option<&str>,
|
||||
requested_model: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
@@ -1059,18 +1104,21 @@ where
|
||||
context,
|
||||
candidate,
|
||||
candidate_index,
|
||||
routing_policy,
|
||||
client_api_format,
|
||||
build_extra_data,
|
||||
)
|
||||
.await,
|
||||
);
|
||||
}
|
||||
LocalExecutionCandidateKind::PoolGroup => {
|
||||
let mut cursor = PoolKeyCursor::new(
|
||||
let mut cursor = PoolKeyCursor::new_with_routing_policy(
|
||||
state,
|
||||
candidate,
|
||||
sticky_session_token,
|
||||
requested_model,
|
||||
request_auth_channel,
|
||||
routing_policy,
|
||||
)
|
||||
.with_runtime_miss_diagnostic(trace_id, record_runtime_miss_diagnostic);
|
||||
let attempt_count_before_pool = attempts.len();
|
||||
@@ -1097,6 +1145,8 @@ async fn persist_available_local_execution_candidate_at_index<F>(
|
||||
context: LocalAvailableCandidatePersistenceContext<'_>,
|
||||
candidate: EligibleLocalExecutionCandidate,
|
||||
candidate_index: u32,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
client_api_format: &str,
|
||||
build_extra_data: &F,
|
||||
) -> Vec<LocalExecutionCandidateAttempt>
|
||||
where
|
||||
@@ -1107,6 +1157,16 @@ where
|
||||
available_candidate_base_extra_data_with_dispatch_ref(&candidate, build_extra_data),
|
||||
candidate.ranking.as_ref(),
|
||||
);
|
||||
let extra_data = attach_routing_trace_to_extra_data(
|
||||
routing_policy,
|
||||
client_api_format,
|
||||
&candidate.candidate,
|
||||
candidate.kind,
|
||||
candidate.ranking.as_ref(),
|
||||
None,
|
||||
Some(candidate_index),
|
||||
extra_data,
|
||||
);
|
||||
let should_persist = should_persist_available_local_candidate(&candidate);
|
||||
let mut attempts = Vec::with_capacity(attempt_slots as usize);
|
||||
let mut owned_candidate = Some(candidate);
|
||||
@@ -1190,6 +1250,160 @@ where
|
||||
Some(Value::Object(object))
|
||||
}
|
||||
|
||||
fn attach_routing_trace_to_skipped_candidates(
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
client_api_format: &str,
|
||||
starting_candidate_index: u32,
|
||||
skipped_candidates: Vec<SkippedLocalExecutionCandidate>,
|
||||
) -> Vec<SkippedLocalExecutionCandidate> {
|
||||
skipped_candidates
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(offset, skipped)| {
|
||||
let selected_order =
|
||||
starting_candidate_index.saturating_add(u32::try_from(offset).unwrap_or(u32::MAX));
|
||||
attach_routing_trace_to_skipped_candidate(
|
||||
routing_policy,
|
||||
client_api_format,
|
||||
selected_order,
|
||||
skipped,
|
||||
)
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn attach_routing_trace_to_skipped_candidate(
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
client_api_format: &str,
|
||||
selected_order: u32,
|
||||
mut skipped_candidate: SkippedLocalExecutionCandidate,
|
||||
) -> SkippedLocalExecutionCandidate {
|
||||
let kind = if skipped_candidate
|
||||
.transport
|
||||
.as_ref()
|
||||
.is_some_and(|transport| {
|
||||
admin_provider_pool_config_from_config_value(transport.provider.config.as_ref())
|
||||
.is_some()
|
||||
}) {
|
||||
LocalExecutionCandidateKind::PoolGroup
|
||||
} else {
|
||||
LocalExecutionCandidateKind::SingleKey
|
||||
};
|
||||
skipped_candidate.extra_data = attach_routing_trace_to_extra_data(
|
||||
routing_policy,
|
||||
client_api_format,
|
||||
&skipped_candidate.candidate,
|
||||
kind,
|
||||
skipped_candidate.ranking.as_ref(),
|
||||
Some(skipped_candidate.skip_reason),
|
||||
Some(selected_order),
|
||||
skipped_candidate.extra_data,
|
||||
);
|
||||
skipped_candidate
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn attach_routing_trace_to_extra_data(
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
client_api_format: &str,
|
||||
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
||||
kind: LocalExecutionCandidateKind,
|
||||
ranking: Option<&SchedulerRankingOutcome>,
|
||||
skip_reason: Option<&'static str>,
|
||||
selected_order: Option<u32>,
|
||||
extra_data: Option<Value>,
|
||||
) -> Option<Value> {
|
||||
let Some(policy) = routing_policy else {
|
||||
return extra_data;
|
||||
};
|
||||
let routing_trace = routing_trace_for_candidate(
|
||||
policy,
|
||||
client_api_format,
|
||||
candidate,
|
||||
kind,
|
||||
ranking,
|
||||
skip_reason,
|
||||
selected_order,
|
||||
);
|
||||
Some(merge_routing_trace_into_extra_data(
|
||||
extra_data,
|
||||
routing_trace,
|
||||
))
|
||||
}
|
||||
|
||||
fn merge_routing_trace_into_extra_data(
|
||||
extra_data: Option<Value>,
|
||||
routing_trace: RoutingDecisionTrace,
|
||||
) -> Value {
|
||||
let mut object = match extra_data {
|
||||
Some(Value::Object(object)) => object,
|
||||
Some(value) => {
|
||||
let mut object = serde_json::Map::new();
|
||||
object.insert("extra".to_string(), value);
|
||||
object
|
||||
}
|
||||
None => serde_json::Map::new(),
|
||||
};
|
||||
object.insert(
|
||||
"routing_trace".to_string(),
|
||||
serde_json::json!(routing_trace),
|
||||
);
|
||||
Value::Object(object)
|
||||
}
|
||||
|
||||
fn routing_trace_for_candidate(
|
||||
policy: &ResolvedRoutingPolicy,
|
||||
client_api_format: &str,
|
||||
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
||||
kind: LocalExecutionCandidateKind,
|
||||
ranking: Option<&SchedulerRankingOutcome>,
|
||||
skip_reason: Option<&'static str>,
|
||||
selected_order: Option<u32>,
|
||||
) -> RoutingDecisionTrace {
|
||||
let candidate_kind = routing_candidate_kind(kind);
|
||||
let mut trace = crate::routing::build_routing_trace_seed(policy, client_api_format);
|
||||
trace.global_candidates.push(RoutingCandidateTrace {
|
||||
candidate_kind,
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
model_id: candidate.model_id.clone(),
|
||||
key_id: match candidate_kind {
|
||||
CandidateKind::Provider => Some(candidate.key_id.clone()),
|
||||
CandidateKind::PoolGroup => None,
|
||||
},
|
||||
ranking_vector: rank_vector_for_candidate(
|
||||
&policy.ranking_overlay,
|
||||
&RoutingCandidateFacts {
|
||||
candidate_kind,
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
model_id: candidate.model_id.clone(),
|
||||
key_id: match candidate_kind {
|
||||
CandidateKind::Provider => Some(candidate.key_id.clone()),
|
||||
CandidateKind::PoolGroup => None,
|
||||
},
|
||||
provider_priority: candidate.provider_priority,
|
||||
key_priority: candidate
|
||||
.key_global_priority_for_format
|
||||
.unwrap_or(candidate.key_internal_priority),
|
||||
},
|
||||
),
|
||||
skip_reason: skip_reason.map(str::to_string),
|
||||
selected_order,
|
||||
});
|
||||
if let Some(ranking) = ranking {
|
||||
trace.runtime_facts.cache_affinity_hit = ranking.promoted_by == Some("cached_affinity");
|
||||
}
|
||||
trace
|
||||
}
|
||||
|
||||
fn routing_candidate_kind(kind: LocalExecutionCandidateKind) -> CandidateKind {
|
||||
match kind {
|
||||
LocalExecutionCandidateKind::SingleKey => CandidateKind::Provider,
|
||||
LocalExecutionCandidateKind::PoolGroup => CandidateKind::PoolGroup,
|
||||
}
|
||||
}
|
||||
|
||||
fn dispatch_sequence_from_attempts(
|
||||
attempts: Vec<LocalExecutionCandidateAttempt>,
|
||||
) -> DispatchSequence<LocalExecutionCandidateAttempt> {
|
||||
@@ -1644,6 +1858,7 @@ mod tests {
|
||||
auth_snapshot: Some(&auth_snapshot),
|
||||
client_session_affinity: None,
|
||||
required_capabilities: None,
|
||||
routing_policy: None,
|
||||
sticky_session_token: None,
|
||||
request_auth_channel: None,
|
||||
persistence_policy: LocalCandidatePersistencePolicy {
|
||||
@@ -1718,6 +1933,8 @@ mod tests {
|
||||
false,
|
||||
vec![pool_group, sample_eligible("normal-key", None)],
|
||||
None,
|
||||
"openai:chat",
|
||||
None,
|
||||
Some("gpt-5"),
|
||||
None,
|
||||
&|_| None,
|
||||
|
||||
@@ -5,6 +5,7 @@ use aether_ai_serving::{
|
||||
AiCandidateRankingPort, AiRankableCandidateParts, AiRankingContextConfig,
|
||||
AiRankingSchedulingMode,
|
||||
};
|
||||
use aether_routing_core::{ResolvedRoutingPolicy, RoutingSchedulingMode, RoutingSetPriorityMode};
|
||||
use async_trait::async_trait;
|
||||
use tracing::warn;
|
||||
|
||||
@@ -16,12 +17,12 @@ use crate::scheduler::config::{
|
||||
};
|
||||
use aether_scheduler_core::{
|
||||
matches_affinity_target, ClientSessionAffinity, SchedulerAffinityTarget,
|
||||
SchedulerMinimalCandidateSelectionCandidate, SchedulerRankableCandidate,
|
||||
SchedulerMinimalCandidateSelectionCandidate, SchedulerPriorityMode, SchedulerRankableCandidate,
|
||||
SchedulerRankingContext, SchedulerRankingOutcome,
|
||||
};
|
||||
|
||||
use super::candidate_affinity_cache::read_cached_scheduler_affinity_target;
|
||||
use super::candidate_resolution::EligibleLocalExecutionCandidate;
|
||||
use super::candidate_resolution::{EligibleLocalExecutionCandidate, LocalExecutionCandidateKind};
|
||||
use super::candidate_transport_ranking_facts::{
|
||||
resolve_cached_transport_ranking_facts, CandidateTransportRankingFacts,
|
||||
};
|
||||
@@ -33,6 +34,7 @@ struct GatewayLocalCandidateRankingPort<'a> {
|
||||
client_session_affinity: Option<&'a ClientSessionAffinity>,
|
||||
required_capabilities: Option<&'a serde_json::Value>,
|
||||
ordering_config: SchedulerOrderingConfig,
|
||||
routing_policy: Option<&'a ResolvedRoutingPolicy>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -89,8 +91,10 @@ impl AiCandidateRankingPort for GatewayLocalCandidateRankingPort<'_> {
|
||||
self.ordering_config,
|
||||
)
|
||||
.await;
|
||||
let routing_overlaid_candidate =
|
||||
routing_overlaid_candidate(self.routing_policy, candidate.kind, &candidate.candidate);
|
||||
Ok(build_ai_rankable_candidate(AiRankableCandidateParts {
|
||||
candidate: &candidate.candidate,
|
||||
candidate: &routing_overlaid_candidate,
|
||||
original_index,
|
||||
normalized_client_api_format,
|
||||
provider_api_format: candidate.provider_api_format.as_str(),
|
||||
@@ -122,8 +126,9 @@ pub(crate) async fn rank_eligible_local_execution_candidates(
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
client_session_affinity: Option<&ClientSessionAffinity>,
|
||||
required_capabilities: Option<&serde_json::Value>,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
) -> Vec<EligibleLocalExecutionCandidate> {
|
||||
let ordering_config = read_scheduler_ordering_config_or_default(state).await;
|
||||
let ordering_config = scheduler_ordering_config_for_routing_policy(state, routing_policy).await;
|
||||
let port = GatewayLocalCandidateRankingPort {
|
||||
state,
|
||||
requested_model,
|
||||
@@ -131,6 +136,7 @@ pub(crate) async fn rank_eligible_local_execution_candidates(
|
||||
client_session_affinity,
|
||||
required_capabilities,
|
||||
ordering_config,
|
||||
routing_policy,
|
||||
};
|
||||
|
||||
match run_ai_candidate_ranking(&port, candidates, normalized_client_api_format).await {
|
||||
@@ -189,6 +195,58 @@ fn ai_ranking_scheduling_mode(mode: SchedulerSchedulingMode) -> AiRankingSchedul
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn scheduler_ordering_config_for_routing_policy(
|
||||
state: PlannerAppState<'_>,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
) -> SchedulerOrderingConfig {
|
||||
match routing_policy {
|
||||
Some(policy) => scheduler_ordering_config_from_routing_policy(policy),
|
||||
None => read_scheduler_ordering_config_or_default(state).await,
|
||||
}
|
||||
}
|
||||
|
||||
fn scheduler_ordering_config_from_routing_policy(
|
||||
policy: &ResolvedRoutingPolicy,
|
||||
) -> SchedulerOrderingConfig {
|
||||
SchedulerOrderingConfig {
|
||||
priority_mode: match policy.priority_mode {
|
||||
RoutingSetPriorityMode::Provider => SchedulerPriorityMode::Provider,
|
||||
RoutingSetPriorityMode::GlobalKey => SchedulerPriorityMode::GlobalKey,
|
||||
},
|
||||
scheduling_mode: match policy.scheduling_mode {
|
||||
RoutingSchedulingMode::FixedOrder => SchedulerSchedulingMode::FixedOrder,
|
||||
RoutingSchedulingMode::CacheAffinity => SchedulerSchedulingMode::CacheAffinity,
|
||||
RoutingSchedulingMode::LoadBalance => SchedulerSchedulingMode::LoadBalance,
|
||||
},
|
||||
keep_priority_on_conversion: policy.keep_priority_on_conversion,
|
||||
}
|
||||
}
|
||||
|
||||
fn routing_overlaid_candidate(
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
kind: LocalExecutionCandidateKind,
|
||||
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
||||
) -> SchedulerMinimalCandidateSelectionCandidate {
|
||||
let Some(policy) = routing_policy else {
|
||||
return candidate.clone();
|
||||
};
|
||||
let mut overlaid = candidate.clone();
|
||||
overlaid.provider_priority = policy
|
||||
.ranking_overlay
|
||||
.provider_priority_or_unspecified(candidate.provider_id.as_str());
|
||||
let overlaid_key_priority = match kind {
|
||||
LocalExecutionCandidateKind::SingleKey => policy
|
||||
.ranking_overlay
|
||||
.key_priority_or_unspecified(candidate.key_id.as_str()),
|
||||
LocalExecutionCandidateKind::PoolGroup => policy
|
||||
.ranking_overlay
|
||||
.pool_priority_or_unspecified(candidate.provider_id.as_str()),
|
||||
};
|
||||
overlaid.key_internal_priority = overlaid_key_priority;
|
||||
overlaid.key_global_priority_for_format = Some(overlaid_key_priority);
|
||||
overlaid
|
||||
}
|
||||
|
||||
async fn read_scheduler_ordering_config_or_default(
|
||||
state: PlannerAppState<'_>,
|
||||
) -> SchedulerOrderingConfig {
|
||||
@@ -304,6 +362,82 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn routing_policy_priorities_do_not_fall_back_to_candidate_priorities() {
|
||||
let mut candidate = sample_candidate("endpoint-1", "key-1");
|
||||
candidate.provider_priority = 7;
|
||||
candidate.key_internal_priority = 3;
|
||||
candidate.key_global_priority_for_format = Some(2);
|
||||
let policy = aether_routing_core::ResolvedRoutingPolicy {
|
||||
group_id: Some("group-1".to_string()),
|
||||
group_version: Some(1),
|
||||
selection_source: "system_default".to_string(),
|
||||
requested_model: "gpt-5".to_string(),
|
||||
resolved_model: "gpt-5".to_string(),
|
||||
priority_mode: aether_routing_core::RoutingSetPriorityMode::Provider,
|
||||
scheduling_mode: aether_routing_core::RoutingSchedulingMode::CacheAffinity,
|
||||
keep_priority_on_conversion: false,
|
||||
ranking_overlay: aether_routing_core::RankingOverlay::default(),
|
||||
mutation_plan: Default::default(),
|
||||
pool_policy_overrides: BTreeMap::new(),
|
||||
matched_rules: Vec::new(),
|
||||
};
|
||||
|
||||
let overlaid = super::routing_overlaid_candidate(
|
||||
Some(&policy),
|
||||
LocalExecutionCandidateKind::SingleKey,
|
||||
&candidate,
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
overlaid.provider_priority,
|
||||
aether_routing_core::ROUTING_PRIORITY_UNSPECIFIED
|
||||
);
|
||||
assert_eq!(
|
||||
overlaid.key_internal_priority,
|
||||
aether_routing_core::ROUTING_PRIORITY_UNSPECIFIED
|
||||
);
|
||||
assert_eq!(
|
||||
overlaid.key_global_priority_for_format,
|
||||
Some(aether_routing_core::ROUTING_PRIORITY_UNSPECIFIED)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn routing_policy_uses_pool_priority_for_pool_group_global_key_slot() {
|
||||
let mut candidate = sample_candidate("endpoint-1", "representative-key");
|
||||
candidate.provider_priority = 7;
|
||||
candidate.key_internal_priority = 3;
|
||||
candidate.key_global_priority_for_format = Some(2);
|
||||
let policy = aether_routing_core::ResolvedRoutingPolicy {
|
||||
group_id: Some("group-1".to_string()),
|
||||
group_version: Some(1),
|
||||
selection_source: "system_default".to_string(),
|
||||
requested_model: "gpt-5".to_string(),
|
||||
resolved_model: "gpt-5".to_string(),
|
||||
priority_mode: aether_routing_core::RoutingSetPriorityMode::GlobalKey,
|
||||
scheduling_mode: aether_routing_core::RoutingSchedulingMode::CacheAffinity,
|
||||
keep_priority_on_conversion: false,
|
||||
ranking_overlay: aether_routing_core::RankingOverlay {
|
||||
pool_priority_overrides: BTreeMap::from([("provider-1".to_string(), 4)]),
|
||||
key_priority_overrides: BTreeMap::from([("representative-key".to_string(), 1)]),
|
||||
..Default::default()
|
||||
},
|
||||
mutation_plan: Default::default(),
|
||||
pool_policy_overrides: BTreeMap::new(),
|
||||
matched_rules: Vec::new(),
|
||||
};
|
||||
|
||||
let overlaid = super::routing_overlaid_candidate(
|
||||
Some(&policy),
|
||||
LocalExecutionCandidateKind::PoolGroup,
|
||||
&candidate,
|
||||
);
|
||||
|
||||
assert_eq!(overlaid.key_internal_priority, 4);
|
||||
assert_eq!(overlaid.key_global_priority_for_format, Some(4));
|
||||
}
|
||||
|
||||
fn sample_provider() -> StoredProviderCatalogProvider {
|
||||
sample_provider_with_options("provider-1", false, 0)
|
||||
}
|
||||
@@ -1062,6 +1196,7 @@ mod tests {
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -1141,6 +1276,7 @@ mod tests {
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -1216,6 +1352,7 @@ mod tests {
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -1282,6 +1419,7 @@ mod tests {
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -1364,6 +1502,7 @@ mod tests {
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -1438,6 +1577,7 @@ mod tests {
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -1515,6 +1655,7 @@ mod tests {
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -1610,6 +1751,7 @@ mod tests {
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -1713,6 +1855,7 @@ mod tests {
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -1809,6 +1952,7 @@ mod tests {
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
aether_ai_serving::AiCandidateResolutionMode::Standard,
|
||||
)
|
||||
.await;
|
||||
@@ -1901,6 +2045,7 @@ mod tests {
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
aether_ai_serving::AiCandidateResolutionMode::Standard,
|
||||
)
|
||||
.await;
|
||||
|
||||
@@ -4,6 +4,7 @@ use aether_ai_serving::{
|
||||
run_ai_candidate_resolution, AiCandidateResolutionMode, AiCandidateResolutionPort,
|
||||
AiCandidateResolutionRequest,
|
||||
};
|
||||
use aether_routing_core::ResolvedRoutingPolicy;
|
||||
use async_trait::async_trait;
|
||||
use std::convert::Infallible;
|
||||
use tracing::warn;
|
||||
@@ -60,6 +61,7 @@ struct GatewayLocalCandidateResolutionPort<'a> {
|
||||
auth_snapshot: Option<&'a GatewayAuthApiKeySnapshot>,
|
||||
client_session_affinity: Option<&'a ClientSessionAffinity>,
|
||||
required_capabilities: Option<&'a serde_json::Value>,
|
||||
routing_policy: Option<&'a ResolvedRoutingPolicy>,
|
||||
request_auth_channel: Option<&'a str>,
|
||||
}
|
||||
|
||||
@@ -97,6 +99,11 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
|
||||
transport: &Self::Transport,
|
||||
requested_model: Option<&str>,
|
||||
) -> Option<&'static str> {
|
||||
if let Some(skip_reason) =
|
||||
routing_policy_candidate_skip_reason(self.routing_policy, candidate, transport)
|
||||
{
|
||||
return Some(skip_reason);
|
||||
}
|
||||
if provider_transport_uses_pool(transport) {
|
||||
return pool_group_common_transport_skip_reason(candidate, transport);
|
||||
}
|
||||
@@ -172,6 +179,7 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
|
||||
self.auth_snapshot,
|
||||
self.client_session_affinity,
|
||||
self.required_capabilities,
|
||||
self.routing_policy,
|
||||
)
|
||||
.await)
|
||||
}
|
||||
@@ -192,6 +200,7 @@ pub(crate) async fn resolve_and_rank_local_execution_candidates(
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
client_session_affinity: Option<&ClientSessionAffinity>,
|
||||
required_capabilities: Option<&serde_json::Value>,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
_sticky_session_token: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
) -> (
|
||||
@@ -207,6 +216,7 @@ pub(crate) async fn resolve_and_rank_local_execution_candidates(
|
||||
auth_snapshot,
|
||||
client_session_affinity,
|
||||
required_capabilities,
|
||||
routing_policy,
|
||||
None,
|
||||
request_auth_channel,
|
||||
AiCandidateResolutionMode::Standard,
|
||||
@@ -222,6 +232,7 @@ pub(crate) async fn resolve_and_rank_local_execution_candidates_without_transpor
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
client_session_affinity: Option<&ClientSessionAffinity>,
|
||||
required_capabilities: Option<&serde_json::Value>,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
_sticky_session_token: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
) -> (
|
||||
@@ -237,6 +248,7 @@ pub(crate) async fn resolve_and_rank_local_execution_candidates_without_transpor
|
||||
auth_snapshot,
|
||||
client_session_affinity,
|
||||
required_capabilities,
|
||||
routing_policy,
|
||||
None,
|
||||
request_auth_channel,
|
||||
AiCandidateResolutionMode::WithoutTransportPairGate,
|
||||
@@ -252,6 +264,7 @@ pub(crate) async fn resolve_and_rank_logical_local_execution_candidates(
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
client_session_affinity: Option<&ClientSessionAffinity>,
|
||||
required_capabilities: Option<&serde_json::Value>,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
_sticky_session_token: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
mode: AiCandidateResolutionMode,
|
||||
@@ -267,6 +280,7 @@ pub(crate) async fn resolve_and_rank_logical_local_execution_candidates(
|
||||
auth_snapshot,
|
||||
client_session_affinity,
|
||||
required_capabilities,
|
||||
routing_policy,
|
||||
None,
|
||||
request_auth_channel,
|
||||
mode,
|
||||
@@ -283,6 +297,7 @@ async fn resolve_and_rank_local_execution_candidates_with_mode(
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
client_session_affinity: Option<&ClientSessionAffinity>,
|
||||
required_capabilities: Option<&serde_json::Value>,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
_sticky_session_token: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
mode: AiCandidateResolutionMode,
|
||||
@@ -298,6 +313,7 @@ async fn resolve_and_rank_local_execution_candidates_with_mode(
|
||||
auth_snapshot,
|
||||
client_session_affinity,
|
||||
required_capabilities,
|
||||
routing_policy,
|
||||
None,
|
||||
request_auth_channel,
|
||||
mode,
|
||||
@@ -315,6 +331,7 @@ async fn resolve_and_rank_local_execution_candidates_with_pool_expansion(
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
client_session_affinity: Option<&ClientSessionAffinity>,
|
||||
required_capabilities: Option<&serde_json::Value>,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
_sticky_session_token: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
mode: AiCandidateResolutionMode,
|
||||
@@ -330,6 +347,7 @@ async fn resolve_and_rank_local_execution_candidates_with_pool_expansion(
|
||||
auth_snapshot,
|
||||
client_session_affinity,
|
||||
required_capabilities,
|
||||
routing_policy,
|
||||
request_auth_channel,
|
||||
};
|
||||
|
||||
@@ -369,6 +387,28 @@ fn provider_transport_uses_pool(transport: &GatewayProviderTransportSnapshot) ->
|
||||
.is_some()
|
||||
}
|
||||
|
||||
fn routing_policy_candidate_skip_reason(
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<&'static str> {
|
||||
let policy = routing_policy?;
|
||||
if !policy
|
||||
.ranking_overlay
|
||||
.provider_allowed(candidate.provider_id.as_str())
|
||||
{
|
||||
return Some("routing_profile_disallowed_provider");
|
||||
}
|
||||
if !provider_transport_uses_pool(transport)
|
||||
&& !policy
|
||||
.ranking_overlay
|
||||
.key_allowed(candidate.key_id.as_str())
|
||||
{
|
||||
return Some("routing_profile_disallowed_key");
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
fn pool_group_common_transport_skip_reason(
|
||||
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
|
||||
@@ -2,6 +2,7 @@ use aether_ai_serving::{
|
||||
run_ai_candidate_preselection, AiCandidatePreselectionOutcome, AiCandidatePreselectionPort,
|
||||
};
|
||||
use aether_data_contracts::repository::candidate_selection::StoredMinimalCandidateSelectionRow;
|
||||
use aether_routing_core::ResolvedRoutingPolicy;
|
||||
use aether_scheduler_core::{
|
||||
enumerate_minimal_candidate_selection_with_model_directives, normalize_api_format,
|
||||
resolve_requested_global_model_name_with_model_directives,
|
||||
@@ -35,6 +36,7 @@ struct GatewayLocalCandidatePreselectionPort<'a> {
|
||||
require_streaming: bool,
|
||||
required_capabilities: Option<&'a serde_json::Value>,
|
||||
auth_snapshot: &'a GatewayAuthApiKeySnapshot,
|
||||
routing_policy: Option<&'a ResolvedRoutingPolicy>,
|
||||
client_session_affinity: Option<&'a ClientSessionAffinity>,
|
||||
use_api_format_alias_match: bool,
|
||||
key_mode: LocalCandidatePreselectionKeyMode,
|
||||
@@ -100,13 +102,14 @@ impl AiCandidatePreselectionPort for GatewayLocalCandidatePreselectionPort<'_> {
|
||||
let enable_model_directives = self.model_directive_enabled_api_formats.contains(
|
||||
&crate::ai_serving::normalize_api_format_alias(candidate_api_format),
|
||||
);
|
||||
matches_client_format
|
||||
|| auth_snapshot_allows_cross_format_candidate(
|
||||
self.auth_snapshot,
|
||||
self.requested_model,
|
||||
candidate,
|
||||
enable_model_directives,
|
||||
)
|
||||
routing_policy_allows_provider(self.routing_policy, candidate)
|
||||
&& (matches_client_format
|
||||
|| auth_snapshot_allows_cross_format_candidate(
|
||||
self.auth_snapshot,
|
||||
self.requested_model,
|
||||
candidate,
|
||||
enable_model_directives,
|
||||
))
|
||||
}
|
||||
|
||||
fn skipped_candidate_allowed(
|
||||
@@ -118,13 +121,14 @@ impl AiCandidatePreselectionPort for GatewayLocalCandidatePreselectionPort<'_> {
|
||||
let enable_model_directives = self.model_directive_enabled_api_formats.contains(
|
||||
&crate::ai_serving::normalize_api_format_alias(candidate_api_format),
|
||||
);
|
||||
matches_client_format
|
||||
|| auth_snapshot_allows_cross_format_candidate(
|
||||
self.auth_snapshot,
|
||||
self.requested_model,
|
||||
&skipped_candidate.candidate,
|
||||
enable_model_directives,
|
||||
)
|
||||
routing_policy_allows_provider(self.routing_policy, &skipped_candidate.candidate)
|
||||
&& (matches_client_format
|
||||
|| auth_snapshot_allows_cross_format_candidate(
|
||||
self.auth_snapshot,
|
||||
self.requested_model,
|
||||
&skipped_candidate.candidate,
|
||||
enable_model_directives,
|
||||
))
|
||||
}
|
||||
|
||||
fn candidate_key(&self, candidate: &Self::Candidate) -> String {
|
||||
@@ -144,6 +148,7 @@ pub(crate) async fn preselect_local_execution_candidates_with_serving(
|
||||
require_streaming: bool,
|
||||
required_capabilities: Option<&serde_json::Value>,
|
||||
auth_snapshot: &GatewayAuthApiKeySnapshot,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
client_session_affinity: Option<&ClientSessionAffinity>,
|
||||
use_api_format_alias_match: bool,
|
||||
key_mode: LocalCandidatePreselectionKeyMode,
|
||||
@@ -166,6 +171,7 @@ pub(crate) async fn preselect_local_execution_candidates_with_serving(
|
||||
require_streaming,
|
||||
required_capabilities,
|
||||
auth_snapshot,
|
||||
routing_policy,
|
||||
client_session_affinity,
|
||||
use_api_format_alias_match,
|
||||
key_mode,
|
||||
@@ -182,6 +188,7 @@ pub(crate) async fn preselect_local_execution_candidates_for_api_formats_with_se
|
||||
require_streaming: bool,
|
||||
required_capabilities: Option<&serde_json::Value>,
|
||||
auth_snapshot: &GatewayAuthApiKeySnapshot,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
client_session_affinity: Option<&ClientSessionAffinity>,
|
||||
use_api_format_alias_match: bool,
|
||||
key_mode: LocalCandidatePreselectionKeyMode,
|
||||
@@ -213,6 +220,7 @@ pub(crate) async fn preselect_local_execution_candidates_for_api_formats_with_se
|
||||
require_streaming,
|
||||
required_capabilities,
|
||||
auth_snapshot,
|
||||
routing_policy,
|
||||
client_session_affinity,
|
||||
use_api_format_alias_match,
|
||||
key_mode,
|
||||
@@ -230,6 +238,7 @@ pub(crate) struct LocalCandidatePreselectionPageCursor<'a> {
|
||||
require_streaming: bool,
|
||||
required_capabilities: Option<serde_json::Value>,
|
||||
auth_snapshot: GatewayAuthApiKeySnapshot,
|
||||
routing_policy: Option<ResolvedRoutingPolicy>,
|
||||
client_session_affinity: Option<ClientSessionAffinity>,
|
||||
use_api_format_alias_match: bool,
|
||||
key_mode: LocalCandidatePreselectionKeyMode,
|
||||
@@ -253,6 +262,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
require_streaming: bool,
|
||||
required_capabilities: Option<&serde_json::Value>,
|
||||
auth_snapshot: &GatewayAuthApiKeySnapshot,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
client_session_affinity: Option<&ClientSessionAffinity>,
|
||||
use_api_format_alias_match: bool,
|
||||
key_mode: LocalCandidatePreselectionKeyMode,
|
||||
@@ -283,6 +293,7 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
require_streaming,
|
||||
required_capabilities: required_capabilities.cloned(),
|
||||
auth_snapshot: auth_snapshot.clone(),
|
||||
routing_policy: routing_policy.cloned(),
|
||||
client_session_affinity: client_session_affinity.cloned(),
|
||||
use_api_format_alias_match,
|
||||
key_mode,
|
||||
@@ -620,16 +631,17 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
candidate_api_format: &str,
|
||||
enable_model_directives: bool,
|
||||
) -> bool {
|
||||
matches_client_api_format(
|
||||
self.use_api_format_alias_match,
|
||||
candidate_api_format,
|
||||
&self.client_api_format,
|
||||
) || auth_snapshot_allows_cross_format_candidate(
|
||||
&self.auth_snapshot,
|
||||
&self.requested_model,
|
||||
candidate,
|
||||
enable_model_directives,
|
||||
)
|
||||
routing_policy_allows_provider(self.routing_policy.as_ref(), candidate)
|
||||
&& (matches_client_api_format(
|
||||
self.use_api_format_alias_match,
|
||||
candidate_api_format,
|
||||
&self.client_api_format,
|
||||
) || auth_snapshot_allows_cross_format_candidate(
|
||||
&self.auth_snapshot,
|
||||
&self.requested_model,
|
||||
candidate,
|
||||
enable_model_directives,
|
||||
))
|
||||
}
|
||||
|
||||
fn skipped_candidate_allowed_for_page(
|
||||
@@ -638,16 +650,17 @@ impl<'a> LocalCandidatePreselectionPageCursor<'a> {
|
||||
candidate_api_format: &str,
|
||||
enable_model_directives: bool,
|
||||
) -> bool {
|
||||
matches_client_api_format(
|
||||
self.use_api_format_alias_match,
|
||||
candidate_api_format,
|
||||
&self.client_api_format,
|
||||
) || auth_snapshot_allows_cross_format_candidate(
|
||||
&self.auth_snapshot,
|
||||
&self.requested_model,
|
||||
&skipped_candidate.candidate,
|
||||
enable_model_directives,
|
||||
)
|
||||
routing_policy_allows_provider(self.routing_policy.as_ref(), &skipped_candidate.candidate)
|
||||
&& (matches_client_api_format(
|
||||
self.use_api_format_alias_match,
|
||||
candidate_api_format,
|
||||
&self.client_api_format,
|
||||
) || auth_snapshot_allows_cross_format_candidate(
|
||||
&self.auth_snapshot,
|
||||
&self.requested_model,
|
||||
&skipped_candidate.candidate,
|
||||
enable_model_directives,
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -739,6 +752,18 @@ pub(crate) fn auth_snapshot_allows_cross_format_candidate(
|
||||
true
|
||||
}
|
||||
|
||||
fn routing_policy_allows_provider(
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
||||
) -> bool {
|
||||
match routing_policy {
|
||||
Some(policy) => policy
|
||||
.ranking_overlay
|
||||
.provider_allowed(candidate.provider_id.as_str()),
|
||||
None => true,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
@@ -887,6 +912,7 @@ mod tests {
|
||||
None,
|
||||
&auth_snapshot,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
|
||||
)
|
||||
@@ -943,6 +969,7 @@ mod tests {
|
||||
None,
|
||||
&auth_snapshot,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
|
||||
)
|
||||
|
||||
@@ -1,10 +1,27 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use aether_ai_serving::{run_ai_authenticated_decision_input, AiAuthenticatedDecisionInputPort};
|
||||
use aether_routing_core::{
|
||||
rank_vector_for_candidate, CandidateKind, ResolvedRoutingPolicy, RoutingCandidateFacts,
|
||||
RoutingCandidateTrace, RoutingDecisionTrace, RoutingPoolExpansionTrace, RoutingRulePhase,
|
||||
};
|
||||
use aether_scheduler_core::ClientSessionAffinity;
|
||||
use async_trait::async_trait;
|
||||
use http::StatusCode;
|
||||
use http::{HeaderMap, HeaderName, HeaderValue};
|
||||
use serde_json::{json, Value};
|
||||
use tracing::warn;
|
||||
|
||||
use crate::ai_serving::planner::common::extract_standard_requested_model;
|
||||
use crate::ai_serving::{ExecutionRuntimeAuthContext, GatewayAuthApiKeySnapshot, PlannerAppState};
|
||||
use crate::client_session_affinity::client_session_affinity_from_request;
|
||||
use crate::clock::current_unix_secs;
|
||||
use crate::{AppState, GatewayError};
|
||||
use crate::routing::{
|
||||
apply_routing_mutation_plan, build_routing_trace_seed, resolve_gateway_routing_policy,
|
||||
select_gateway_routing_group, GatewayRoutingPolicyInput, GatewayRoutingSelectionError,
|
||||
GatewayRoutingSelectionInput, ROUTING_GROUP_HEADER,
|
||||
};
|
||||
use crate::{AiExecutionDecision, AppState, GatewayError};
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct ResolvedLocalDecisionAuthInput {
|
||||
@@ -21,6 +38,9 @@ pub(crate) struct LocalRequestedModelDecisionInput {
|
||||
pub(crate) required_capabilities: Option<serde_json::Value>,
|
||||
pub(crate) request_auth_channel: Option<String>,
|
||||
pub(crate) client_session_affinity: Option<ClientSessionAffinity>,
|
||||
pub(crate) routing_policy: Option<ResolvedRoutingPolicy>,
|
||||
pub(crate) routing_trace_seed: Option<RoutingDecisionTrace>,
|
||||
pub(crate) routing_context: Option<LocalRoutingRequestContext>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
@@ -31,6 +51,92 @@ pub(crate) struct LocalAuthenticatedDecisionInput {
|
||||
pub(crate) client_session_affinity: Option<ClientSessionAffinity>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct LocalRoutingRequestContext {
|
||||
pub(crate) group_id: Option<String>,
|
||||
pub(crate) group_version: Option<i64>,
|
||||
pub(crate) group_config_json: Value,
|
||||
pub(crate) selection_source: String,
|
||||
pub(crate) client_api_format: String,
|
||||
pub(crate) effective_body_json: Value,
|
||||
pub(crate) effective_headers: HeaderMap,
|
||||
}
|
||||
|
||||
impl LocalRequestedModelDecisionInput {
|
||||
pub(crate) fn effective_body_json<'a>(&'a self, fallback: &'a Value) -> &'a Value {
|
||||
self.routing_context
|
||||
.as_ref()
|
||||
.map(|context| &context.effective_body_json)
|
||||
.unwrap_or(fallback)
|
||||
}
|
||||
|
||||
pub(crate) fn effective_headers<'a>(&'a self, fallback: &'a HeaderMap) -> &'a HeaderMap {
|
||||
self.routing_context
|
||||
.as_ref()
|
||||
.map(|context| &context.effective_headers)
|
||||
.unwrap_or(fallback)
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn apply_provider_request_routing_policy_to_decision(
|
||||
input: &LocalRequestedModelDecisionInput,
|
||||
decision: &mut AiExecutionDecision,
|
||||
) -> Result<(), GatewayError> {
|
||||
let Some(context) = input.routing_context.as_ref() else {
|
||||
return Ok(());
|
||||
};
|
||||
let provider_api_format = decision
|
||||
.provider_api_format
|
||||
.as_deref()
|
||||
.unwrap_or(context.client_api_format.as_str());
|
||||
let resolved_model = decision
|
||||
.mapped_model
|
||||
.as_deref()
|
||||
.or(decision.model_name.as_deref())
|
||||
.unwrap_or(input.requested_model.as_str());
|
||||
let original_provider_request_body = decision.provider_request_body.clone();
|
||||
let mut provider_request_body = original_provider_request_body
|
||||
.clone()
|
||||
.unwrap_or(serde_json::Value::Null);
|
||||
let mut provider_headers = btree_headers_to_header_map(&decision.provider_request_headers)?;
|
||||
let provider_headers_json = headers_to_routing_value(&provider_headers);
|
||||
let policy = resolve_gateway_routing_policy(GatewayRoutingPolicyInput {
|
||||
group_id: context.group_id.as_deref(),
|
||||
group_version: context.group_version,
|
||||
group_config_json: &context.group_config_json,
|
||||
selection_source: context.selection_source.as_str(),
|
||||
requested_model: input.requested_model.as_str(),
|
||||
resolved_model,
|
||||
api_format: provider_api_format,
|
||||
user_id: Some(input.auth_context.user_id.as_str()),
|
||||
api_key_id: Some(input.auth_context.api_key_id.as_str()),
|
||||
headers: &provider_headers_json,
|
||||
body: &provider_request_body,
|
||||
phase: RoutingRulePhase::ProviderRequest,
|
||||
})?;
|
||||
ensure_report_context_routing_trace(input, decision, &policy);
|
||||
if policy.mutation_plan.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
if original_provider_request_body.is_none() && !policy.mutation_plan.body_patch.is_empty() {
|
||||
return Err(GatewayError::Client {
|
||||
status: StatusCode::BAD_REQUEST,
|
||||
message: "routing provider_request body patch cannot be applied to a binary or empty upstream body".to_string(),
|
||||
});
|
||||
}
|
||||
apply_routing_mutation_plan(
|
||||
&mut provider_request_body,
|
||||
&mut provider_headers,
|
||||
&policy.mutation_plan,
|
||||
)?;
|
||||
decision.provider_request_headers = header_map_to_btree_headers(&provider_headers);
|
||||
if original_provider_request_body.is_some() {
|
||||
decision.provider_request_body = Some(provider_request_body);
|
||||
}
|
||||
update_report_context_provider_request_mutation(decision, &policy);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
struct GatewayAuthenticatedDecisionInputPort<'a> {
|
||||
state: PlannerAppState<'a>,
|
||||
now_unix_secs: u64,
|
||||
@@ -99,9 +205,154 @@ pub(crate) fn build_local_requested_model_decision_input(
|
||||
required_capabilities: resolved_input.required_capabilities,
|
||||
request_auth_channel: None,
|
||||
client_session_affinity: None,
|
||||
routing_policy: None,
|
||||
routing_trace_seed: None,
|
||||
routing_context: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn attach_routing_policy_to_local_requested_model_input(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
input: &mut LocalRequestedModelDecisionInput,
|
||||
body_json: &Value,
|
||||
client_api_format: &str,
|
||||
) -> Result<(), GatewayError> {
|
||||
let explicit_group = routing_header_value_str(&parts.headers, ROUTING_GROUP_HEADER);
|
||||
let selected_group = match state.routing_group_read_repository() {
|
||||
Some(repository) => {
|
||||
let user_group_ids = match state
|
||||
.list_user_groups_for_user(&input.auth_context.user_id)
|
||||
.await
|
||||
{
|
||||
Ok(groups) => groups.into_iter().map(|group| group.id).collect::<Vec<_>>(),
|
||||
Err(error) => {
|
||||
warn!(
|
||||
user_id = %input.auth_context.user_id,
|
||||
error = ?error,
|
||||
"gateway routing profile user group lookup failed"
|
||||
);
|
||||
Vec::new()
|
||||
}
|
||||
};
|
||||
let selection = select_gateway_routing_group(
|
||||
repository.as_ref(),
|
||||
GatewayRoutingSelectionInput {
|
||||
explicit_group: explicit_group.as_deref(),
|
||||
user_id: Some(input.auth_context.user_id.as_str()),
|
||||
api_key_id: Some(input.auth_context.api_key_id.as_str()),
|
||||
user_group_ids: &user_group_ids,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.map_err(routing_selection_error)?;
|
||||
selection.group.map(|group| {
|
||||
(
|
||||
Some(group.id),
|
||||
Some(group.version),
|
||||
group.config_json,
|
||||
selection.source,
|
||||
)
|
||||
})
|
||||
}
|
||||
None => {
|
||||
if explicit_group
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.is_some_and(|value| !value.is_empty())
|
||||
{
|
||||
return Err(routing_selection_error(
|
||||
GatewayRoutingSelectionError::NotFound(explicit_group.unwrap_or_default()),
|
||||
));
|
||||
}
|
||||
None
|
||||
}
|
||||
};
|
||||
|
||||
let Some((group_id, group_version, group_config_json, selection_source)) = selected_group
|
||||
else {
|
||||
input.client_session_affinity =
|
||||
client_session_affinity_from_request(&parts.headers, Some(body_json));
|
||||
input.routing_policy = None;
|
||||
input.routing_trace_seed = None;
|
||||
input.routing_context = None;
|
||||
return Ok(());
|
||||
};
|
||||
|
||||
let headers_json = headers_to_routing_value(&parts.headers);
|
||||
let policy = resolve_gateway_routing_policy(GatewayRoutingPolicyInput {
|
||||
group_id: group_id.as_deref(),
|
||||
group_version,
|
||||
group_config_json: &group_config_json,
|
||||
selection_source: selection_source.as_str(),
|
||||
requested_model: input.requested_model.as_str(),
|
||||
resolved_model: input.requested_model.as_str(),
|
||||
api_format: client_api_format,
|
||||
user_id: Some(input.auth_context.user_id.as_str()),
|
||||
api_key_id: Some(input.auth_context.api_key_id.as_str()),
|
||||
headers: &headers_json,
|
||||
body: body_json,
|
||||
phase: RoutingRulePhase::ClientRequest,
|
||||
})?;
|
||||
let mut effective_body_json = body_json.clone();
|
||||
let mut effective_headers = parts.headers.clone();
|
||||
apply_routing_mutation_plan(
|
||||
&mut effective_body_json,
|
||||
&mut effective_headers,
|
||||
&policy.mutation_plan,
|
||||
)?;
|
||||
|
||||
let mut requested_model_changed = false;
|
||||
if let Some(mut mutated_model) = extract_standard_requested_model(&effective_body_json) {
|
||||
mutated_model = mutated_model.trim().to_string();
|
||||
if !mutated_model.is_empty() && mutated_model != input.requested_model {
|
||||
input.requested_model = mutated_model;
|
||||
requested_model_changed = true;
|
||||
}
|
||||
}
|
||||
if requested_model_changed {
|
||||
input.required_capabilities = PlannerAppState::new(state)
|
||||
.resolve_request_candidate_required_capabilities(
|
||||
&input.auth_context.user_id,
|
||||
&input.auth_context.api_key_id,
|
||||
Some(input.requested_model.as_str()),
|
||||
input.required_capabilities.as_ref(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
let effective_headers_json = headers_to_routing_value(&effective_headers);
|
||||
input.client_session_affinity =
|
||||
client_session_affinity_from_request(&effective_headers, Some(&effective_body_json));
|
||||
let mut final_policy = resolve_gateway_routing_policy(GatewayRoutingPolicyInput {
|
||||
group_id: group_id.as_deref(),
|
||||
group_version,
|
||||
group_config_json: &group_config_json,
|
||||
selection_source: selection_source.as_str(),
|
||||
requested_model: input.requested_model.as_str(),
|
||||
resolved_model: input.requested_model.as_str(),
|
||||
api_format: client_api_format,
|
||||
user_id: Some(input.auth_context.user_id.as_str()),
|
||||
api_key_id: Some(input.auth_context.api_key_id.as_str()),
|
||||
headers: &effective_headers_json,
|
||||
body: &effective_body_json,
|
||||
phase: RoutingRulePhase::ClientRequest,
|
||||
})?;
|
||||
final_policy.mutation_plan = policy.mutation_plan.clone();
|
||||
input.routing_trace_seed = Some(build_routing_trace_seed(&final_policy, client_api_format));
|
||||
input.routing_policy = Some(final_policy);
|
||||
input.routing_context = Some(LocalRoutingRequestContext {
|
||||
group_id,
|
||||
group_version,
|
||||
group_config_json,
|
||||
selection_source,
|
||||
client_api_format: client_api_format.to_string(),
|
||||
effective_body_json,
|
||||
effective_headers,
|
||||
});
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) fn build_local_authenticated_decision_input(
|
||||
resolved_input: ResolvedLocalDecisionAuthInput,
|
||||
) -> LocalAuthenticatedDecisionInput {
|
||||
@@ -132,3 +383,506 @@ pub(crate) async fn resolve_local_authenticated_decision_input(
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn routing_selection_error(error: GatewayRoutingSelectionError) -> GatewayError {
|
||||
GatewayError::Client {
|
||||
status: StatusCode::FORBIDDEN,
|
||||
message: error.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
fn headers_to_routing_value(headers: &http::HeaderMap) -> Value {
|
||||
let mut object = serde_json::Map::new();
|
||||
for (name, value) in headers {
|
||||
if let Ok(value) = value.to_str() {
|
||||
object.insert(name.as_str().to_ascii_lowercase(), json!(value));
|
||||
}
|
||||
}
|
||||
Value::Object(object)
|
||||
}
|
||||
|
||||
fn routing_header_value_str(headers: &http::HeaderMap, key: &str) -> Option<String> {
|
||||
headers
|
||||
.get(key)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn btree_headers_to_header_map(
|
||||
headers: &BTreeMap<String, String>,
|
||||
) -> Result<HeaderMap, GatewayError> {
|
||||
let mut output = HeaderMap::new();
|
||||
for (name, value) in headers {
|
||||
let name = HeaderName::from_bytes(name.as_bytes()).map_err(|err| GatewayError::Client {
|
||||
status: StatusCode::BAD_REQUEST,
|
||||
message: format!("invalid provider request header name in routing mutation: {err}"),
|
||||
})?;
|
||||
let value = HeaderValue::from_str(value).map_err(|err| GatewayError::Client {
|
||||
status: StatusCode::BAD_REQUEST,
|
||||
message: format!("invalid provider request header value in routing mutation: {err}"),
|
||||
})?;
|
||||
output.insert(name, value);
|
||||
}
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
fn header_map_to_btree_headers(headers: &HeaderMap) -> BTreeMap<String, String> {
|
||||
headers
|
||||
.iter()
|
||||
.filter_map(|(name, value)| {
|
||||
value
|
||||
.to_str()
|
||||
.ok()
|
||||
.map(|value| (name.as_str().to_string(), value.to_string()))
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn update_report_context_provider_request_mutation(
|
||||
decision: &mut AiExecutionDecision,
|
||||
policy: &ResolvedRoutingPolicy,
|
||||
) {
|
||||
let Some(serde_json::Value::Object(object)) = decision.report_context.as_mut() else {
|
||||
return;
|
||||
};
|
||||
let body_paths = policy
|
||||
.mutation_plan
|
||||
.body_patch
|
||||
.iter()
|
||||
.map(|operation| operation.path().to_string())
|
||||
.collect::<Vec<_>>();
|
||||
let header_names = policy
|
||||
.mutation_plan
|
||||
.header_patch
|
||||
.iter()
|
||||
.map(|operation| operation.name().to_string())
|
||||
.collect::<Vec<_>>();
|
||||
let trace_patch_summary = serde_json::json!({
|
||||
"body_paths": body_paths,
|
||||
"header_names": header_names,
|
||||
});
|
||||
if let Some(serde_json::Value::Object(routing_trace)) = object.get_mut("routing_trace") {
|
||||
routing_trace.insert(
|
||||
"provider_request_patch_summary".to_string(),
|
||||
trace_patch_summary.clone(),
|
||||
);
|
||||
}
|
||||
object.insert(
|
||||
"provider_request_headers".to_string(),
|
||||
serde_json::json!(decision.provider_request_headers),
|
||||
);
|
||||
object.insert(
|
||||
"routing_provider_request_patch_summary".to_string(),
|
||||
serde_json::json!({
|
||||
"body_paths": trace_patch_summary["body_paths"].clone(),
|
||||
"header_names": trace_patch_summary["header_names"].clone(),
|
||||
"matched_rules": policy
|
||||
.matched_rules
|
||||
.iter()
|
||||
.map(|rule| rule.id.clone())
|
||||
.collect::<Vec<_>>()
|
||||
}),
|
||||
);
|
||||
}
|
||||
|
||||
fn ensure_report_context_routing_trace(
|
||||
input: &LocalRequestedModelDecisionInput,
|
||||
decision: &mut AiExecutionDecision,
|
||||
policy: &ResolvedRoutingPolicy,
|
||||
) {
|
||||
let Some(serde_json::Value::Object(object)) = decision.report_context.as_mut() else {
|
||||
return;
|
||||
};
|
||||
if object.get("routing_trace").is_some() {
|
||||
return;
|
||||
}
|
||||
|
||||
let client_api_format = decision
|
||||
.client_api_format
|
||||
.as_deref()
|
||||
.or_else(|| {
|
||||
input
|
||||
.routing_context
|
||||
.as_ref()
|
||||
.map(|context| context.client_api_format.as_str())
|
||||
})
|
||||
.unwrap_or_default();
|
||||
let mut trace = input
|
||||
.routing_trace_seed
|
||||
.clone()
|
||||
.unwrap_or_else(|| build_routing_trace_seed(policy, client_api_format));
|
||||
|
||||
let candidate_group_id = object
|
||||
.get("candidate_group_id")
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned);
|
||||
let pool_key_index = object
|
||||
.get("pool_key_index")
|
||||
.and_then(Value::as_u64)
|
||||
.and_then(|value| u32::try_from(value).ok());
|
||||
let is_pool_expansion = candidate_group_id.is_some() && pool_key_index.is_some();
|
||||
let candidate_kind = if is_pool_expansion {
|
||||
CandidateKind::PoolGroup
|
||||
} else {
|
||||
CandidateKind::Provider
|
||||
};
|
||||
let provider_id = candidate_group_id
|
||||
.clone()
|
||||
.or_else(|| decision.provider_id.clone())
|
||||
.unwrap_or_default();
|
||||
let endpoint_id = decision.endpoint_id.clone().unwrap_or_default();
|
||||
let model_id = object
|
||||
.get("model_id")
|
||||
.and_then(Value::as_str)
|
||||
.map(ToOwned::to_owned)
|
||||
.or_else(|| decision.mapped_model.clone())
|
||||
.or_else(|| decision.model_name.clone())
|
||||
.unwrap_or_else(|| input.requested_model.clone());
|
||||
let key_id = decision.key_id.clone().filter(|_| !is_pool_expansion);
|
||||
let provider_priority = object
|
||||
.get("provider_priority")
|
||||
.and_then(Value::as_i64)
|
||||
.and_then(|value| i32::try_from(value).ok())
|
||||
.unwrap_or_default();
|
||||
let key_priority = object
|
||||
.get("priority_slot")
|
||||
.and_then(Value::as_i64)
|
||||
.and_then(|value| i32::try_from(value).ok())
|
||||
.unwrap_or_default();
|
||||
trace.global_candidates.push(RoutingCandidateTrace {
|
||||
candidate_kind,
|
||||
provider_id: provider_id.clone(),
|
||||
endpoint_id,
|
||||
model_id: model_id.clone(),
|
||||
key_id: key_id.clone(),
|
||||
ranking_vector: rank_vector_for_candidate(
|
||||
&policy.ranking_overlay,
|
||||
&RoutingCandidateFacts {
|
||||
candidate_kind,
|
||||
provider_id: provider_id.clone(),
|
||||
endpoint_id: decision.endpoint_id.clone().unwrap_or_default(),
|
||||
model_id,
|
||||
key_id,
|
||||
provider_priority,
|
||||
key_priority,
|
||||
},
|
||||
),
|
||||
skip_reason: None,
|
||||
selected_order: object
|
||||
.get("candidate_index")
|
||||
.and_then(Value::as_u64)
|
||||
.and_then(|value| u32::try_from(value).ok()),
|
||||
});
|
||||
|
||||
if is_pool_expansion {
|
||||
if let (Some(pool_group_id), Some(key_id)) = (candidate_group_id, decision.key_id.clone()) {
|
||||
trace.pool_expansion.push(RoutingPoolExpansionTrace {
|
||||
pool_group_id,
|
||||
key_id,
|
||||
pool_ranking_vector: Vec::new(),
|
||||
pool_skip_reason: None,
|
||||
selected_order: pool_key_index,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
object.insert("routing_trace".to_string(), serde_json::json!(trace));
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn sample_auth_context() -> ExecutionRuntimeAuthContext {
|
||||
ExecutionRuntimeAuthContext {
|
||||
user_id: "user-1".to_string(),
|
||||
api_key_id: "api-key-1".to_string(),
|
||||
username: None,
|
||||
api_key_name: None,
|
||||
balance_remaining: None,
|
||||
access_allowed: true,
|
||||
api_key_is_standalone: false,
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_auth_snapshot() -> GatewayAuthApiKeySnapshot {
|
||||
GatewayAuthApiKeySnapshot {
|
||||
user_id: "user-1".to_string(),
|
||||
username: "alice".to_string(),
|
||||
email: None,
|
||||
user_role: "user".to_string(),
|
||||
user_auth_source: "local".to_string(),
|
||||
user_is_active: true,
|
||||
user_is_deleted: false,
|
||||
user_rate_limit: None,
|
||||
user_allowed_providers: None,
|
||||
user_allowed_api_formats: None,
|
||||
user_allowed_models: None,
|
||||
api_key_id: "api-key-1".to_string(),
|
||||
api_key_name: Some("default".to_string()),
|
||||
api_key_is_active: true,
|
||||
api_key_is_locked: false,
|
||||
api_key_is_standalone: false,
|
||||
api_key_rate_limit: None,
|
||||
api_key_concurrent_limit: None,
|
||||
api_key_expires_at_unix_secs: None,
|
||||
api_key_allowed_providers: None,
|
||||
api_key_allowed_api_formats: None,
|
||||
api_key_allowed_models: None,
|
||||
currently_usable: true,
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_decision_input() -> LocalRequestedModelDecisionInput {
|
||||
LocalRequestedModelDecisionInput {
|
||||
auth_context: sample_auth_context(),
|
||||
requested_model: "gpt-5".to_string(),
|
||||
auth_snapshot: sample_auth_snapshot(),
|
||||
required_capabilities: None,
|
||||
request_auth_channel: None,
|
||||
client_session_affinity: None,
|
||||
routing_policy: None,
|
||||
routing_trace_seed: None,
|
||||
routing_context: Some(LocalRoutingRequestContext {
|
||||
group_id: Some("group-1".to_string()),
|
||||
group_version: Some(3),
|
||||
selection_source: "explicit_header".to_string(),
|
||||
client_api_format: "openai:chat".to_string(),
|
||||
effective_body_json: json!({"model":"gpt-5"}),
|
||||
effective_headers: HeaderMap::new(),
|
||||
group_config_json: json!({
|
||||
"allowed_models": ["gpt-5"],
|
||||
"rules": [{
|
||||
"id": "provider-patch",
|
||||
"priority": 1,
|
||||
"enabled": true,
|
||||
"phase": "provider_request",
|
||||
"conditions": {},
|
||||
"actions": [
|
||||
{
|
||||
"type": "json_patch_body",
|
||||
"patch": [{
|
||||
"op": "add",
|
||||
"path": "/metadata/routing",
|
||||
"value": "provider"
|
||||
}]
|
||||
},
|
||||
{
|
||||
"type": "patch_headers",
|
||||
"patch": [{
|
||||
"op": "set",
|
||||
"name": "x-provider-route",
|
||||
"value": "provider"
|
||||
}]
|
||||
}
|
||||
]
|
||||
}]
|
||||
}),
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_decision() -> AiExecutionDecision {
|
||||
AiExecutionDecision {
|
||||
action: "execution_runtime_sync_decision".to_string(),
|
||||
decision_kind: Some("openai_chat_sync".to_string()),
|
||||
execution_strategy: None,
|
||||
conversion_mode: None,
|
||||
request_id: Some("trace-1".to_string()),
|
||||
candidate_id: Some("candidate-1".to_string()),
|
||||
provider_name: Some("provider".to_string()),
|
||||
provider_id: Some("provider-1".to_string()),
|
||||
endpoint_id: Some("endpoint-1".to_string()),
|
||||
key_id: Some("key-1".to_string()),
|
||||
upstream_base_url: None,
|
||||
upstream_url: None,
|
||||
provider_request_method: None,
|
||||
auth_header: None,
|
||||
auth_value: None,
|
||||
provider_api_format: Some("openai:chat".to_string()),
|
||||
client_api_format: Some("openai:chat".to_string()),
|
||||
provider_contract: None,
|
||||
client_contract: None,
|
||||
model_name: Some("gpt-5".to_string()),
|
||||
mapped_model: Some("gpt-5".to_string()),
|
||||
prompt_cache_key: None,
|
||||
extra_headers: BTreeMap::new(),
|
||||
provider_request_headers: BTreeMap::from([(
|
||||
"content-type".to_string(),
|
||||
"application/json".to_string(),
|
||||
)]),
|
||||
provider_request_body: Some(json!({"model":"gpt-5","metadata":{}})),
|
||||
provider_request_body_base64: None,
|
||||
content_type: Some("application/json".to_string()),
|
||||
proxy: None,
|
||||
transport_profile: None,
|
||||
timeouts: None,
|
||||
upstream_is_stream: false,
|
||||
report_kind: Some("local_sync_success".to_string()),
|
||||
report_context: Some(json!({
|
||||
"candidate_index": 0,
|
||||
"retry_index": 0,
|
||||
"model_id": "model-1"
|
||||
})),
|
||||
auth_context: Some(sample_auth_context()),
|
||||
}
|
||||
}
|
||||
|
||||
fn set_provider_request_rules(input: &mut LocalRequestedModelDecisionInput, actions: Value) {
|
||||
let config = json!({
|
||||
"allowed_models": ["gpt-5"],
|
||||
"rules": [{
|
||||
"id": "provider-patch",
|
||||
"priority": 1,
|
||||
"enabled": true,
|
||||
"phase": "provider_request",
|
||||
"conditions": {},
|
||||
"actions": actions
|
||||
}]
|
||||
});
|
||||
input
|
||||
.routing_context
|
||||
.as_mut()
|
||||
.expect("sample input should include routing context")
|
||||
.group_config_json = config;
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_request_routing_policy_mutates_decision_body_headers_and_report_context() {
|
||||
let input = sample_decision_input();
|
||||
let mut decision = sample_decision();
|
||||
|
||||
apply_provider_request_routing_policy_to_decision(&input, &mut decision)
|
||||
.expect("provider routing mutation should apply");
|
||||
|
||||
assert_eq!(
|
||||
decision.provider_request_body.as_ref().unwrap()["metadata"]["routing"],
|
||||
json!("provider")
|
||||
);
|
||||
assert_eq!(
|
||||
decision
|
||||
.provider_request_headers
|
||||
.get("x-provider-route")
|
||||
.map(String::as_str),
|
||||
Some("provider")
|
||||
);
|
||||
let report_context = decision.report_context.as_ref().unwrap();
|
||||
assert_eq!(
|
||||
report_context["routing_provider_request_patch_summary"]["matched_rules"],
|
||||
json!(["provider-patch"])
|
||||
);
|
||||
assert_eq!(
|
||||
report_context["routing_trace"]["provider_request_patch_summary"]["body_paths"],
|
||||
json!(["/metadata/routing"])
|
||||
);
|
||||
assert_eq!(
|
||||
report_context["routing_trace"]["global_candidates"][0]["provider_id"],
|
||||
json!("provider-1")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_request_routing_policy_rejects_body_patch_without_json_body() {
|
||||
let input = sample_decision_input();
|
||||
let mut decision = sample_decision();
|
||||
decision.provider_request_body = None;
|
||||
decision.provider_request_body_base64 = Some("AA==".to_string());
|
||||
|
||||
let error = apply_provider_request_routing_policy_to_decision(&input, &mut decision)
|
||||
.expect_err("provider body patch should reject binary upstream bodies");
|
||||
|
||||
match error {
|
||||
GatewayError::Client { status, message } => {
|
||||
assert_eq!(status, StatusCode::BAD_REQUEST);
|
||||
assert!(message.contains("binary or empty upstream body"));
|
||||
}
|
||||
other => panic!("unexpected error: {other:?}"),
|
||||
}
|
||||
assert!(
|
||||
decision
|
||||
.report_context
|
||||
.as_ref()
|
||||
.and_then(|context| context.get("routing_trace"))
|
||||
.is_some(),
|
||||
"failed provider_request mutation should still seed routing trace"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_request_routing_policy_allows_header_patch_without_json_body() {
|
||||
let mut input = sample_decision_input();
|
||||
set_provider_request_rules(
|
||||
&mut input,
|
||||
json!([{
|
||||
"type": "patch_headers",
|
||||
"patch": [{
|
||||
"op": "set",
|
||||
"name": "x-provider-route",
|
||||
"value": "header-only"
|
||||
}]
|
||||
}]),
|
||||
);
|
||||
let mut decision = sample_decision();
|
||||
decision.provider_request_body = None;
|
||||
decision.provider_request_body_base64 = Some("AA==".to_string());
|
||||
|
||||
apply_provider_request_routing_policy_to_decision(&input, &mut decision)
|
||||
.expect("header-only provider routing mutation should apply without JSON body");
|
||||
|
||||
assert_eq!(decision.provider_request_body, None);
|
||||
assert_eq!(
|
||||
decision
|
||||
.provider_request_headers
|
||||
.get("x-provider-route")
|
||||
.map(String::as_str),
|
||||
Some("header-only")
|
||||
);
|
||||
assert_eq!(
|
||||
decision.report_context.as_ref().unwrap()["routing_trace"]
|
||||
["provider_request_patch_summary"]["header_names"],
|
||||
json!(["x-provider-route"])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_request_routing_trace_records_pool_expansion_candidate() {
|
||||
let input = sample_decision_input();
|
||||
let mut decision = sample_decision();
|
||||
decision.report_context = Some(json!({
|
||||
"candidate_index": 2,
|
||||
"retry_index": 2,
|
||||
"model_id": "model-1",
|
||||
"candidate_group_id": "pool-group-1",
|
||||
"pool_key_index": 1,
|
||||
"provider_priority": 7,
|
||||
"priority_slot": 3
|
||||
}));
|
||||
|
||||
apply_provider_request_routing_policy_to_decision(&input, &mut decision)
|
||||
.expect("provider routing mutation should seed pool trace");
|
||||
|
||||
let routing_trace = &decision.report_context.as_ref().unwrap()["routing_trace"];
|
||||
assert_eq!(
|
||||
routing_trace["global_candidates"][0]["candidate_kind"],
|
||||
json!("pool_group")
|
||||
);
|
||||
assert_eq!(
|
||||
routing_trace["global_candidates"][0]["provider_id"],
|
||||
json!("pool-group-1")
|
||||
);
|
||||
assert_eq!(routing_trace["global_candidates"][0]["key_id"], Value::Null);
|
||||
assert_eq!(
|
||||
routing_trace["pool_expansion"][0]["pool_group_id"],
|
||||
json!("pool-group-1")
|
||||
);
|
||||
assert_eq!(routing_trace["pool_expansion"][0]["key_id"], json!("key-1"));
|
||||
assert_eq!(
|
||||
routing_trace["pool_expansion"][0]["selected_order"],
|
||||
json!(1)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -33,7 +33,7 @@ pub(crate) async fn maybe_build_sync_local_same_format_provider_decision_payload
|
||||
let Some(input) = resolve_local_same_format_provider_decision_input(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
@@ -55,6 +55,7 @@ pub(crate) async fn maybe_build_sync_local_same_format_provider_decision_payload
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, spec,
|
||||
)
|
||||
@@ -70,7 +71,7 @@ pub(crate) async fn maybe_build_sync_local_same_format_provider_decision_payload
|
||||
maybe_build_local_same_format_provider_decision_payload_for_candidate(
|
||||
state, parts, trace_id, body_json, &input, attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
{
|
||||
return Ok(Some(payload));
|
||||
}
|
||||
@@ -100,7 +101,7 @@ pub(crate) async fn maybe_build_stream_local_same_format_provider_decision_paylo
|
||||
let Some(input) = resolve_local_same_format_provider_decision_input(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
@@ -122,6 +123,7 @@ pub(crate) async fn maybe_build_stream_local_same_format_provider_decision_paylo
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, spec,
|
||||
)
|
||||
@@ -137,7 +139,7 @@ pub(crate) async fn maybe_build_stream_local_same_format_provider_decision_paylo
|
||||
maybe_build_local_same_format_provider_decision_payload_for_candidate(
|
||||
state, parts, trace_id, body_json, &input, attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
{
|
||||
return Ok(Some(payload));
|
||||
}
|
||||
|
||||
@@ -13,6 +13,7 @@ use crate::ai_serving::planner::candidate_metadata::{
|
||||
use crate::ai_serving::planner::candidate_resolution::SkippedLocalExecutionCandidate;
|
||||
use crate::ai_serving::planner::common::extract_requested_model_from_request;
|
||||
use crate::ai_serving::planner::decision_input::{
|
||||
attach_routing_policy_to_local_requested_model_input,
|
||||
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
|
||||
};
|
||||
use crate::ai_serving::planner::materialization_policy::{
|
||||
@@ -39,19 +40,21 @@ pub(crate) async fn resolve_local_same_format_provider_decision_input(
|
||||
decision: &GatewayControlDecision,
|
||||
body_json: &serde_json::Value,
|
||||
spec: LocalSameFormatProviderSpec,
|
||||
) -> Option<LocalSameFormatProviderDecisionInput> {
|
||||
) -> Result<Option<LocalSameFormatProviderDecisionInput>, GatewayError> {
|
||||
let spec_metadata = local_same_format_provider_spec_metadata(spec);
|
||||
let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else {
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let requested_model = extract_requested_model_from_request(
|
||||
let Some(requested_model) = extract_requested_model_from_request(
|
||||
parts,
|
||||
body_json,
|
||||
spec_metadata
|
||||
.requested_model_family
|
||||
.expect("same-format provider specs should declare requested-model family"),
|
||||
)?;
|
||||
) else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let resolved_input = match resolve_local_authenticated_decision_input(
|
||||
state,
|
||||
@@ -62,7 +65,7 @@ pub(crate) async fn resolve_local_same_format_provider_decision_input(
|
||||
.await
|
||||
{
|
||||
Ok(Some(resolved_input)) => resolved_input,
|
||||
Ok(None) => return None,
|
||||
Ok(None) => return Ok(None),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
@@ -70,14 +73,31 @@ pub(crate) async fn resolve_local_same_format_provider_decision_input(
|
||||
error = ?err,
|
||||
"gateway local same-format decision auth snapshot read failed"
|
||||
);
|
||||
return None;
|
||||
return Err(err);
|
||||
}
|
||||
};
|
||||
|
||||
let mut input = build_local_requested_model_decision_input(resolved_input, requested_model);
|
||||
input.request_auth_channel = decision.request_auth_channel.clone();
|
||||
input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json));
|
||||
Some(input)
|
||||
if let Err(err) = attach_routing_policy_to_local_requested_model_input(
|
||||
state,
|
||||
parts,
|
||||
&mut input,
|
||||
body_json,
|
||||
spec_metadata.api_format,
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
api_format = spec_metadata.api_format,
|
||||
error = ?err,
|
||||
"gateway local same-format decision routing profile resolution failed"
|
||||
);
|
||||
return Err(err);
|
||||
}
|
||||
Ok(Some(input))
|
||||
}
|
||||
|
||||
pub(crate) async fn materialize_local_same_format_provider_candidate_attempts(
|
||||
@@ -114,6 +134,7 @@ pub(crate) async fn materialize_local_same_format_provider_candidate_attempts(
|
||||
Some(&input.auth_snapshot),
|
||||
input.client_session_affinity.as_ref(),
|
||||
input.required_capabilities.as_ref(),
|
||||
input.routing_policy.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
@@ -211,6 +232,7 @@ pub(crate) async fn build_local_same_format_provider_candidate_attempt_source<'a
|
||||
Some(&input.auth_snapshot),
|
||||
input.client_session_affinity.as_ref(),
|
||||
input.required_capabilities.as_ref(),
|
||||
input.routing_policy.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
|
||||
@@ -6,6 +6,7 @@ use crate::ai_serving::planner::candidate_materialization::{
|
||||
mark_skipped_local_execution_candidate, mark_skipped_local_execution_candidate_with_extra_data,
|
||||
mark_skipped_local_execution_candidate_with_failure_diagnostic,
|
||||
};
|
||||
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
|
||||
use crate::ai_serving::planner::materialization_policy::{
|
||||
build_local_candidate_persistence_policy, LocalCandidatePersistencePolicyKind,
|
||||
};
|
||||
@@ -22,7 +23,7 @@ use crate::ai_serving::transport::{
|
||||
};
|
||||
use crate::{
|
||||
append_execution_contract_fields_to_value, append_local_failover_policy_to_value,
|
||||
AiExecutionDecision, AppState,
|
||||
AiExecutionDecision, AppState, GatewayError,
|
||||
};
|
||||
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
|
||||
|
||||
@@ -40,7 +41,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
input: &LocalSameFormatProviderDecisionInput,
|
||||
attempt: LocalSameFormatProviderCandidateAttempt,
|
||||
spec: LocalSameFormatProviderSpec,
|
||||
) -> Option<AiExecutionDecision> {
|
||||
) -> Result<Option<AiExecutionDecision>, GatewayError> {
|
||||
let spec_metadata = local_same_format_provider_spec_metadata(spec);
|
||||
let LocalSameFormatProviderCandidateAttempt {
|
||||
eligible,
|
||||
@@ -51,10 +52,13 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
let candidate = &eligible.candidate;
|
||||
let (execution_strategy, conversion_mode) =
|
||||
ai_local_execution_contract_for_formats(spec_metadata.api_format, spec_metadata.api_format);
|
||||
let resolved = resolve_local_same_format_provider_candidate_payload_parts(
|
||||
let Some(resolved) = resolve_local_same_format_provider_candidate_payload_parts(
|
||||
state, parts, trace_id, body_json, input, &attempt, spec,
|
||||
)
|
||||
.await?;
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let prompt_cache_key = resolved
|
||||
.provider_request_body
|
||||
@@ -66,7 +70,10 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
let proxy = state
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport)
|
||||
.await;
|
||||
let transport_profile = resolve_transport_profile(&resolved.transport);
|
||||
let transport_profile = resolved
|
||||
.transport_profile
|
||||
.clone()
|
||||
.or_else(|| resolve_transport_profile(&resolved.transport));
|
||||
let mut extra_fields = serde_json::Map::new();
|
||||
if let Some(proxy_value) =
|
||||
build_request_trace_proxy_value(Some(&resolved.transport), proxy.as_ref())
|
||||
@@ -85,6 +92,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
);
|
||||
}
|
||||
let provider_api_format = resolved.provider_api_format.clone();
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let report_context = append_local_failover_policy_to_value(
|
||||
append_execution_contract_fields_to_value(
|
||||
build_local_execution_report_context(LocalExecutionReportContextParts {
|
||||
@@ -112,7 +120,7 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
body_rules: resolved.transport.endpoint.body_rules.as_ref(),
|
||||
provider_request_method: Some(serde_json::Value::Null),
|
||||
provider_request_headers: Some(&resolved.provider_request_headers),
|
||||
original_headers: &parts.headers,
|
||||
original_headers: effective_headers,
|
||||
request_path: Some(parts.uri.path()),
|
||||
request_query_string: parts.uri.query(),
|
||||
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
|
||||
@@ -149,43 +157,44 @@ pub(crate) async fn maybe_build_local_same_format_provider_decision_payload_for_
|
||||
upstream_url,
|
||||
provider_request_headers,
|
||||
provider_request_body,
|
||||
transport_profile: _,
|
||||
} = resolved;
|
||||
|
||||
Some(build_ai_execution_decision_response(
|
||||
AiExecutionDecisionResponseParts {
|
||||
decision_is_stream: spec_metadata.require_streaming,
|
||||
decision_kind: spec_metadata.decision_kind.to_string(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: candidate_id.to_string(),
|
||||
provider_name: transport.provider.name.clone(),
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
key_id: candidate.key_id.clone(),
|
||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||
upstream_url,
|
||||
provider_request_method: None,
|
||||
auth_header,
|
||||
auth_value,
|
||||
provider_api_format,
|
||||
client_api_format: spec_metadata.api_format.to_string(),
|
||||
model_name: input.requested_model.clone(),
|
||||
mapped_model,
|
||||
prompt_cache_key,
|
||||
provider_request_headers,
|
||||
provider_request_body: Some(provider_request_body),
|
||||
provider_request_body_base64: None,
|
||||
content_type: Some("application/json".to_string()),
|
||||
proxy,
|
||||
transport_profile,
|
||||
timeouts: resolve_transport_execution_timeouts(&transport),
|
||||
upstream_is_stream,
|
||||
report_kind: Some(report_kind.to_string()),
|
||||
report_context: Some(report_context),
|
||||
auth_context: input.auth_context.clone(),
|
||||
},
|
||||
))
|
||||
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
|
||||
decision_is_stream: spec_metadata.require_streaming,
|
||||
decision_kind: spec_metadata.decision_kind.to_string(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: candidate_id.to_string(),
|
||||
provider_name: transport.provider.name.clone(),
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
key_id: candidate.key_id.clone(),
|
||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||
upstream_url,
|
||||
provider_request_method: None,
|
||||
auth_header,
|
||||
auth_value,
|
||||
provider_api_format,
|
||||
client_api_format: spec_metadata.api_format.to_string(),
|
||||
model_name: input.requested_model.clone(),
|
||||
mapped_model,
|
||||
prompt_cache_key,
|
||||
provider_request_headers,
|
||||
provider_request_body: Some(provider_request_body),
|
||||
provider_request_body_base64: None,
|
||||
content_type: Some("application/json".to_string()),
|
||||
proxy,
|
||||
transport_profile,
|
||||
timeouts: resolve_transport_execution_timeouts(&transport),
|
||||
upstream_is_stream,
|
||||
report_kind: Some(report_kind.to_string()),
|
||||
report_context: Some(report_context),
|
||||
auth_context: input.auth_context.clone(),
|
||||
});
|
||||
apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
|
||||
Ok(Some(decision))
|
||||
}
|
||||
|
||||
pub(super) async fn mark_skipped_local_same_format_provider_candidate(
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use aether_contracts::ResolvedTransportProfile;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::ai_serving::planner::common::{
|
||||
@@ -12,7 +13,8 @@ use crate::ai_serving::transport::antigravity::{
|
||||
AntigravityRequestEnvelopeSupport, AntigravityRequestSideSupport,
|
||||
};
|
||||
use crate::ai_serving::transport::{
|
||||
build_same_format_provider_headers, SameFormatProviderHeadersInput,
|
||||
build_grok_browser_headers, build_grok_upstream_url, build_same_format_provider_headers,
|
||||
GrokHeaderInput, SameFormatProviderHeadersInput, GROK_CHAT_PATH,
|
||||
};
|
||||
use crate::ai_serving::{CandidateFailureDiagnostic, GatewayProviderTransportSnapshot};
|
||||
use crate::AppState;
|
||||
@@ -96,6 +98,7 @@ pub(crate) struct LocalSameFormatProviderCandidatePayloadParts {
|
||||
pub(super) upstream_url: String,
|
||||
pub(super) provider_request_headers: BTreeMap<String, String>,
|
||||
pub(super) provider_request_body: Value,
|
||||
pub(super) transport_profile: Option<ResolvedTransportProfile>,
|
||||
}
|
||||
|
||||
pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
@@ -125,6 +128,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
Some(&input.requested_model),
|
||||
)
|
||||
.await;
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
|
||||
let Some(mut base_provider_request_body) =
|
||||
super::super::request::build_same_format_provider_request_body(
|
||||
@@ -133,7 +137,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
&prepared.mapped_model,
|
||||
spec,
|
||||
prepared.transport.endpoint.body_rules.as_ref(),
|
||||
Some(&parts.headers),
|
||||
Some(effective_headers),
|
||||
prepared.upstream_is_stream,
|
||||
prepared.force_body_stream_field,
|
||||
prepared.kiro_auth.as_ref(),
|
||||
@@ -246,16 +250,28 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
base_provider_request_body
|
||||
};
|
||||
|
||||
let Some(upstream_url) = super::super::request::build_same_format_upstream_url(
|
||||
parts,
|
||||
&prepared.transport,
|
||||
&prepared.mapped_model,
|
||||
prepared.provider_api_format.as_str(),
|
||||
spec,
|
||||
prepared.upstream_is_stream,
|
||||
prepared.kiro_auth.as_ref(),
|
||||
Some(&provider_request_body),
|
||||
) else {
|
||||
let is_grok = prepared
|
||||
.transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("grok");
|
||||
let transport_profile =
|
||||
crate::ai_serving::transport::resolve_transport_profile(&prepared.transport);
|
||||
let Some(upstream_url) = (if is_grok {
|
||||
Some(build_grok_upstream_url(&prepared.transport, GROK_CHAT_PATH))
|
||||
} else {
|
||||
super::super::request::build_same_format_upstream_url(
|
||||
parts,
|
||||
&prepared.transport,
|
||||
&prepared.mapped_model,
|
||||
prepared.provider_api_format.as_str(),
|
||||
spec,
|
||||
prepared.upstream_is_stream,
|
||||
prepared.kiro_auth.as_ref(),
|
||||
Some(&provider_request_body),
|
||||
)
|
||||
}) else {
|
||||
mark_skipped_local_same_format_provider_candidate_with_failure_diagnostic(
|
||||
state,
|
||||
input,
|
||||
@@ -278,9 +294,20 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
.as_ref()
|
||||
.map(build_antigravity_static_identity_headers)
|
||||
.unwrap_or_default();
|
||||
let Some(provider_request_headers) =
|
||||
let Some(provider_request_headers) = (if is_grok {
|
||||
build_grok_browser_headers(GrokHeaderInput {
|
||||
transport: &prepared.transport,
|
||||
transport_profile: transport_profile.as_ref(),
|
||||
request_headers: Some(effective_headers),
|
||||
content_type: "application/json",
|
||||
accept: "text/event-stream",
|
||||
header_rules: prepared.transport.endpoint.header_rules.as_ref(),
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: body_json,
|
||||
})
|
||||
} else {
|
||||
build_same_format_provider_headers(SameFormatProviderHeadersInput {
|
||||
headers: &parts.headers,
|
||||
headers: effective_headers,
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: body_json,
|
||||
header_rules: prepared.transport.endpoint.header_rules.as_ref(),
|
||||
@@ -295,7 +322,7 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
.as_ref()
|
||||
.map(|auth| auth.machine_id.as_str()),
|
||||
})
|
||||
else {
|
||||
}) else {
|
||||
mark_skipped_local_same_format_provider_candidate_with_failure_diagnostic(
|
||||
state,
|
||||
input,
|
||||
@@ -327,5 +354,6 @@ pub(crate) async fn resolve_local_same_format_provider_candidate_payload_parts(
|
||||
upstream_url,
|
||||
provider_request_headers,
|
||||
provider_request_body,
|
||||
transport_profile,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -31,7 +31,7 @@ pub(crate) struct LocalSameFormatProviderSyncAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_json: serde_json::Value,
|
||||
input: LocalSameFormatProviderDecisionInput,
|
||||
spec: LocalSameFormatProviderSpec,
|
||||
requested_model_family: RequestedModelFamily,
|
||||
@@ -42,7 +42,7 @@ pub(crate) struct LocalSameFormatProviderStreamAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_json: serde_json::Value,
|
||||
input: LocalSameFormatProviderDecisionInput,
|
||||
spec: LocalSameFormatProviderSpec,
|
||||
requested_model_family: RequestedModelFamily,
|
||||
@@ -64,7 +64,7 @@ pub(crate) async fn build_local_sync_attempt_source<'a>(
|
||||
let Some(input) = resolve_local_same_format_provider_decision_input(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
@@ -85,8 +85,13 @@ pub(crate) async fn build_local_sync_attempt_source<'a>(
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let effective_body_json = input.effective_body_json(body_json).clone();
|
||||
let (candidates, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, spec,
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
&effective_body_json,
|
||||
spec,
|
||||
)
|
||||
.await?;
|
||||
apply_local_runtime_candidate_evaluation_progress_preserving_candidate_signal(
|
||||
@@ -103,7 +108,7 @@ pub(crate) async fn build_local_sync_attempt_source<'a>(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
body_json: effective_body_json,
|
||||
input,
|
||||
spec,
|
||||
requested_model_family,
|
||||
@@ -128,7 +133,7 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
|
||||
let Some(input) = resolve_local_same_format_provider_decision_input(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
@@ -149,8 +154,13 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let effective_body_json = input.effective_body_json(body_json).clone();
|
||||
let (candidates, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, spec,
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
&effective_body_json,
|
||||
spec,
|
||||
)
|
||||
.await?;
|
||||
apply_local_runtime_candidate_evaluation_progress_preserving_candidate_signal(
|
||||
@@ -167,7 +177,7 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
body_json: effective_body_json,
|
||||
input,
|
||||
spec,
|
||||
requested_model_family,
|
||||
@@ -244,12 +254,12 @@ impl LocalSameFormatProviderSyncAttemptSource<'_> {
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
@@ -257,7 +267,7 @@ impl LocalSameFormatProviderSyncAttemptSource<'_> {
|
||||
match build_sync_plan_from_requested_model_family(
|
||||
self.requested_model_family,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
payload,
|
||||
) {
|
||||
Ok(value) => Ok(value),
|
||||
@@ -282,12 +292,12 @@ impl LocalSameFormatProviderStreamAttemptSource<'_> {
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
@@ -295,7 +305,7 @@ impl LocalSameFormatProviderStreamAttemptSource<'_> {
|
||||
match build_stream_plan_from_requested_model_family(
|
||||
self.requested_model_family,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
payload,
|
||||
) {
|
||||
Ok(value) => Ok(value),
|
||||
@@ -326,7 +336,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
|
||||
let Some(input) = resolve_local_same_format_provider_decision_input(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
@@ -347,6 +357,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, spec,
|
||||
)
|
||||
@@ -365,7 +376,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
|
||||
let Some(payload) = maybe_build_local_same_format_provider_decision_payload_for_candidate(
|
||||
state, parts, trace_id, body_json, &input, attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
@@ -411,7 +422,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
|
||||
let Some(input) = resolve_local_same_format_provider_decision_input(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
@@ -432,6 +443,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
let (mut source, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, spec,
|
||||
)
|
||||
@@ -450,7 +462,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
|
||||
let Some(payload) = maybe_build_local_same_format_provider_decision_payload_for_candidate(
|
||||
state, parts, trace_id, body_json, &input, attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
@@ -82,6 +82,18 @@ pub(crate) fn build_local_execution_report_context(
|
||||
.client_session_affinity
|
||||
.and_then(client_session_affinity_report_context_value)
|
||||
{
|
||||
if let Some(client_family) = value
|
||||
.as_object()
|
||||
.and_then(|object| object.get("client_family"))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|client_family| !client_family.is_empty())
|
||||
{
|
||||
extra_fields.insert(
|
||||
"client_family".to_string(),
|
||||
Value::String(client_family.to_ascii_lowercase()),
|
||||
);
|
||||
}
|
||||
extra_fields.insert(
|
||||
CLIENT_SESSION_AFFINITY_REPORT_CONTEXT_FIELD.to_string(),
|
||||
value,
|
||||
|
||||
@@ -29,7 +29,7 @@ use self::support::{
|
||||
pub(crate) struct LocalGeminiFilesSyncAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_json: serde_json::Value,
|
||||
body_base64: Option<&'a str>,
|
||||
body_is_empty: bool,
|
||||
trace_id: &'a str,
|
||||
@@ -110,10 +110,11 @@ pub(crate) async fn build_local_gemini_files_sync_attempt_source_for_kind<'a>(
|
||||
trace_id,
|
||||
decision,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let effective_body_json = input.effective_body_json(body_json).clone();
|
||||
let (candidates, candidate_count) =
|
||||
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
|
||||
if candidate_count == 0 {
|
||||
@@ -124,7 +125,7 @@ pub(crate) async fn build_local_gemini_files_sync_attempt_source_for_kind<'a>(
|
||||
LocalGeminiFilesSyncAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
body_json: effective_body_json,
|
||||
body_base64,
|
||||
body_is_empty,
|
||||
trace_id,
|
||||
@@ -148,7 +149,7 @@ pub(crate) async fn build_local_gemini_files_stream_attempt_source_for_kind<'a>(
|
||||
};
|
||||
|
||||
let Some(input) =
|
||||
resolve_local_gemini_files_decision_input(state, parts, None, trace_id, decision).await
|
||||
resolve_local_gemini_files_decision_input(state, parts, None, trace_id, decision).await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
@@ -226,7 +227,7 @@ impl LocalGeminiFilesSyncAttemptSource<'_> {
|
||||
let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
self.body_base64,
|
||||
self.body_is_empty,
|
||||
self.trace_id,
|
||||
@@ -234,7 +235,7 @@ impl LocalGeminiFilesSyncAttemptSource<'_> {
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
@@ -272,7 +273,7 @@ impl LocalGeminiFilesStreamAttemptSource<'_> {
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
@@ -313,10 +314,11 @@ pub(crate) async fn maybe_build_sync_local_gemini_files_decision_payload(
|
||||
trace_id,
|
||||
decision,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
|
||||
let (mut source, _) =
|
||||
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
|
||||
@@ -333,7 +335,7 @@ pub(crate) async fn maybe_build_sync_local_gemini_files_decision_payload(
|
||||
attempt,
|
||||
spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
{
|
||||
return Ok(Some(payload));
|
||||
}
|
||||
@@ -354,7 +356,7 @@ pub(crate) async fn maybe_build_stream_local_gemini_files_decision_payload(
|
||||
};
|
||||
|
||||
let Some(input) =
|
||||
resolve_local_gemini_files_decision_input(state, parts, None, trace_id, decision).await
|
||||
resolve_local_gemini_files_decision_input(state, parts, None, trace_id, decision).await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
@@ -375,7 +377,7 @@ pub(crate) async fn maybe_build_stream_local_gemini_files_decision_payload(
|
||||
attempt,
|
||||
spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
{
|
||||
return Ok(Some(payload));
|
||||
}
|
||||
@@ -402,10 +404,11 @@ async fn build_local_sync_plan_and_reports(
|
||||
trace_id,
|
||||
decision,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(Vec::new());
|
||||
};
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
|
||||
let (mut source, _) =
|
||||
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
|
||||
@@ -423,7 +426,7 @@ async fn build_local_sync_plan_and_reports(
|
||||
attempt,
|
||||
spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
@@ -454,7 +457,7 @@ async fn build_local_stream_plan_and_reports(
|
||||
) -> Result<Vec<AiStreamAttempt>, GatewayError> {
|
||||
let spec_metadata = local_gemini_files_spec_metadata(spec);
|
||||
let Some(input) =
|
||||
resolve_local_gemini_files_decision_input(state, parts, None, trace_id, decision).await
|
||||
resolve_local_gemini_files_decision_input(state, parts, None, trace_id, decision).await?
|
||||
else {
|
||||
return Ok(Vec::new());
|
||||
};
|
||||
@@ -476,7 +479,7 @@ async fn build_local_stream_plan_and_reports(
|
||||
attempt,
|
||||
spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use serde_json::json;
|
||||
|
||||
use crate::ai_serving::build_request_trace_proxy_value;
|
||||
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
|
||||
use crate::ai_serving::planner::report_context::{
|
||||
build_local_execution_report_context, LocalExecutionReportContextParts,
|
||||
};
|
||||
@@ -12,7 +13,7 @@ use crate::ai_serving::transport::{
|
||||
resolve_transport_execution_timeouts, resolve_transport_profile,
|
||||
};
|
||||
use crate::ai_serving::{ai_local_execution_contract_for_formats, PlannerAppState};
|
||||
use crate::{AiExecutionDecision, AppState};
|
||||
use crate::{AiExecutionDecision, AppState, GatewayError};
|
||||
|
||||
use super::request::resolve_local_gemini_files_candidate_payload_parts;
|
||||
use super::support::{
|
||||
@@ -31,7 +32,7 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
|
||||
input: &LocalGeminiFilesDecisionInput,
|
||||
attempt: LocalGeminiFilesCandidateAttempt,
|
||||
spec: LocalGeminiFilesSpec,
|
||||
) -> Option<AiExecutionDecision> {
|
||||
) -> Result<Option<AiExecutionDecision>, GatewayError> {
|
||||
let spec_metadata = local_gemini_files_spec_metadata(spec);
|
||||
let planner_state = PlannerAppState::new(state);
|
||||
let attempt_identity = attempt.attempt_identity();
|
||||
@@ -46,7 +47,10 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
|
||||
&attempt,
|
||||
spec,
|
||||
)
|
||||
.await?;
|
||||
.await;
|
||||
let Some(resolved) = resolved else {
|
||||
return Ok(None);
|
||||
};
|
||||
let LocalGeminiFilesCandidateAttempt {
|
||||
eligible,
|
||||
candidate_id,
|
||||
@@ -69,6 +73,7 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
|
||||
}
|
||||
extra_fields.insert("file_key_id".to_string(), json!(candidate.key_id));
|
||||
extra_fields.insert("file_name".to_string(), json!(resolved.file_name));
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let report_context = build_local_execution_report_context(LocalExecutionReportContextParts {
|
||||
auth_context: &input.auth_context,
|
||||
request_id: trace_id,
|
||||
@@ -94,7 +99,7 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
|
||||
body_rules: transport.endpoint.body_rules.as_ref(),
|
||||
provider_request_method: None,
|
||||
provider_request_headers: None,
|
||||
original_headers: &parts.headers,
|
||||
original_headers: effective_headers,
|
||||
request_path: Some(parts.uri.path()),
|
||||
request_query_string: parts.uri.query(),
|
||||
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
|
||||
@@ -119,45 +124,44 @@ pub(super) async fn maybe_build_local_gemini_files_decision_payload_for_candidat
|
||||
file_name: _,
|
||||
} = resolved;
|
||||
|
||||
Some(build_ai_execution_decision_response(
|
||||
AiExecutionDecisionResponseParts {
|
||||
decision_is_stream: spec_metadata.require_streaming,
|
||||
decision_kind: spec_metadata.decision_kind.to_string(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: candidate_id.clone(),
|
||||
provider_name: transport.provider.name.clone(),
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
key_id: candidate.key_id.clone(),
|
||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||
upstream_url,
|
||||
provider_request_method: Some(parts.method.to_string()),
|
||||
auth_header: Some(auth_header),
|
||||
auth_value: Some(auth_value),
|
||||
provider_api_format: GEMINI_FILES_CLIENT_API_FORMAT.to_string(),
|
||||
client_api_format: GEMINI_FILES_CLIENT_API_FORMAT.to_string(),
|
||||
model_name: "gemini-files".to_string(),
|
||||
mapped_model: candidate.selected_provider_model_name.clone(),
|
||||
prompt_cache_key: None,
|
||||
provider_request_headers,
|
||||
provider_request_body,
|
||||
provider_request_body_base64,
|
||||
content_type: parts
|
||||
.headers
|
||||
.get(http::header::CONTENT_TYPE)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned),
|
||||
proxy,
|
||||
transport_profile,
|
||||
timeouts: resolve_transport_execution_timeouts(&transport),
|
||||
upstream_is_stream: spec_metadata.require_streaming,
|
||||
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
||||
report_context: Some(report_context),
|
||||
auth_context: input.auth_context.clone(),
|
||||
},
|
||||
))
|
||||
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
|
||||
decision_is_stream: spec_metadata.require_streaming,
|
||||
decision_kind: spec_metadata.decision_kind.to_string(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: candidate_id.clone(),
|
||||
provider_name: transport.provider.name.clone(),
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
key_id: candidate.key_id.clone(),
|
||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||
upstream_url,
|
||||
provider_request_method: Some(parts.method.to_string()),
|
||||
auth_header: Some(auth_header),
|
||||
auth_value: Some(auth_value),
|
||||
provider_api_format: GEMINI_FILES_CLIENT_API_FORMAT.to_string(),
|
||||
client_api_format: GEMINI_FILES_CLIENT_API_FORMAT.to_string(),
|
||||
model_name: "gemini-files".to_string(),
|
||||
mapped_model: candidate.selected_provider_model_name.clone(),
|
||||
prompt_cache_key: None,
|
||||
provider_request_headers,
|
||||
provider_request_body,
|
||||
provider_request_body_base64,
|
||||
content_type: effective_headers
|
||||
.get(http::header::CONTENT_TYPE)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned),
|
||||
proxy,
|
||||
transport_profile,
|
||||
timeouts: resolve_transport_execution_timeouts(&transport),
|
||||
upstream_is_stream: spec_metadata.require_streaming,
|
||||
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
||||
report_context: Some(report_context),
|
||||
auth_context: input.auth_context.clone(),
|
||||
});
|
||||
apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
|
||||
Ok(Some(decision))
|
||||
}
|
||||
|
||||
@@ -45,6 +45,7 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts(
|
||||
let spec_metadata = local_gemini_files_spec_metadata(spec);
|
||||
let candidate = &attempt.eligible.candidate;
|
||||
let transport = &attempt.eligible.transport;
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
|
||||
if let Some(skip_reason) =
|
||||
gemini_files_transport_unsupported_reason(transport, GEMINI_FILES_CANDIDATE_API_FORMAT)
|
||||
@@ -103,7 +104,7 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts(
|
||||
body_is_empty,
|
||||
spec_metadata.decision_kind == GEMINI_FILES_UPLOAD_PLAN_KIND,
|
||||
transport.endpoint.body_rules.as_ref(),
|
||||
Some(&parts.headers),
|
||||
Some(effective_headers),
|
||||
) {
|
||||
Ok(parts) => parts,
|
||||
Err(GeminiFilesRequestBodyError::BodyRulesUnsupportedForBinaryUpload) => {
|
||||
@@ -145,7 +146,7 @@ pub(super) async fn resolve_local_gemini_files_candidate_payload_parts(
|
||||
};
|
||||
|
||||
let Some(provider_request_headers) = build_gemini_files_headers(GeminiFilesHeadersInput {
|
||||
headers: &parts.headers,
|
||||
headers: effective_headers,
|
||||
auth_header: &auth_header,
|
||||
auth_value: &auth_value,
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
|
||||
@@ -14,7 +14,8 @@ use crate::ai_serving::planner::candidate_metadata::{
|
||||
build_local_execution_candidate_metadata_for_candidate, LocalExecutionCandidateMetadataParts,
|
||||
};
|
||||
use crate::ai_serving::planner::decision_input::{
|
||||
build_local_authenticated_decision_input, resolve_local_authenticated_decision_input,
|
||||
attach_routing_policy_to_local_requested_model_input,
|
||||
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
|
||||
};
|
||||
use crate::ai_serving::planner::materialization_policy::{
|
||||
build_local_candidate_persistence_policy, LocalCandidatePersistencePolicyKind,
|
||||
@@ -29,11 +30,12 @@ use crate::{AppState, GatewayError};
|
||||
|
||||
pub(super) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttempt as LocalGeminiFilesCandidateAttempt;
|
||||
pub(super) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttemptSource as LocalGeminiFilesCandidateAttemptSource;
|
||||
pub(super) use crate::ai_serving::planner::decision_input::LocalAuthenticatedDecisionInput as LocalGeminiFilesDecisionInput;
|
||||
pub(super) use crate::ai_serving::planner::decision_input::LocalRequestedModelDecisionInput as LocalGeminiFilesDecisionInput;
|
||||
|
||||
pub(super) const GEMINI_FILES_CANDIDATE_API_FORMAT: &str = "gemini:files";
|
||||
pub(super) const GEMINI_FILES_CLIENT_API_FORMAT: &str = "gemini:files";
|
||||
pub(super) const GEMINI_FILES_REQUIRED_CAPABILITY: &str = "gemini_files";
|
||||
pub(super) const GEMINI_FILES_ROUTING_MODEL: &str = "gemini-files";
|
||||
|
||||
pub(super) async fn resolve_local_gemini_files_decision_input(
|
||||
state: &AppState,
|
||||
@@ -41,9 +43,9 @@ pub(super) async fn resolve_local_gemini_files_decision_input(
|
||||
body_json: Option<&serde_json::Value>,
|
||||
trace_id: &str,
|
||||
decision: &GatewayControlDecision,
|
||||
) -> Option<LocalGeminiFilesDecisionInput> {
|
||||
) -> Result<Option<LocalGeminiFilesDecisionInput>, GatewayError> {
|
||||
let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else {
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let explicit_required_capabilities = json!({ "gemini_files": true });
|
||||
@@ -56,20 +58,33 @@ pub(super) async fn resolve_local_gemini_files_decision_input(
|
||||
.await
|
||||
{
|
||||
Ok(Some(resolved_input)) => resolved_input,
|
||||
Ok(None) => return None,
|
||||
Ok(None) => return Ok(None),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
error = ?err,
|
||||
"gateway local gemini files decision auth snapshot read failed"
|
||||
);
|
||||
return None;
|
||||
return Err(err);
|
||||
}
|
||||
};
|
||||
|
||||
let mut input = build_local_authenticated_decision_input(resolved_input);
|
||||
let routing_body_json = body_json.cloned().unwrap_or(serde_json::Value::Null);
|
||||
let mut input = build_local_requested_model_decision_input(
|
||||
resolved_input,
|
||||
GEMINI_FILES_ROUTING_MODEL.to_string(),
|
||||
);
|
||||
input.request_auth_channel = decision.request_auth_channel.clone();
|
||||
input.client_session_affinity = client_session_affinity_from_parts(parts, body_json);
|
||||
Some(input)
|
||||
attach_routing_policy_to_local_requested_model_input(
|
||||
state,
|
||||
parts,
|
||||
&mut input,
|
||||
&routing_body_json,
|
||||
GEMINI_FILES_CLIENT_API_FORMAT,
|
||||
)
|
||||
.await?;
|
||||
Ok(Some(input))
|
||||
}
|
||||
|
||||
pub(super) async fn materialize_local_gemini_files_candidate_attempts(
|
||||
@@ -101,8 +116,9 @@ pub(super) async fn materialize_local_gemini_files_candidate_attempts(
|
||||
Some(&input.auth_snapshot),
|
||||
input.client_session_affinity.as_ref(),
|
||||
input.required_capabilities.as_ref(),
|
||||
input.routing_policy.as_ref(),
|
||||
None,
|
||||
None,
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
candidates,
|
||||
Vec::new(),
|
||||
@@ -173,8 +189,9 @@ pub(super) async fn build_local_gemini_files_candidate_attempt_source<'a>(
|
||||
Some(&input.auth_snapshot),
|
||||
input.client_session_affinity.as_ref(),
|
||||
input.required_capabilities.as_ref(),
|
||||
input.routing_policy.as_ref(),
|
||||
None,
|
||||
None,
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
candidates,
|
||||
Vec::new(),
|
||||
|
||||
@@ -32,7 +32,7 @@ pub(super) use crate::ai_serving::LocalOpenAiImageSpec;
|
||||
pub(crate) struct LocalOpenAiImageSyncAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_json: serde_json::Value,
|
||||
body_base64: Option<&'a str>,
|
||||
trace_id: &'a str,
|
||||
input: LocalOpenAiImageDecisionInput,
|
||||
@@ -43,7 +43,7 @@ pub(crate) struct LocalOpenAiImageSyncAttemptSource<'a> {
|
||||
pub(crate) struct LocalOpenAiImageStreamAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_json: serde_json::Value,
|
||||
body_base64: Option<&'a str>,
|
||||
trace_id: &'a str,
|
||||
input: LocalOpenAiImageDecisionInput,
|
||||
@@ -152,16 +152,17 @@ pub(crate) async fn build_local_image_sync_attempt_source_for_kind<'a>(
|
||||
trace_id,
|
||||
decision,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let effective_body_json = input.effective_body_json(body_json).clone();
|
||||
let Some((candidates, candidate_count)) = build_local_openai_image_candidate_attempt_source(
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
body_json,
|
||||
&effective_body_json,
|
||||
spec_metadata.api_format,
|
||||
spec_metadata.decision_kind,
|
||||
)
|
||||
@@ -178,7 +179,7 @@ pub(crate) async fn build_local_image_sync_attempt_source_for_kind<'a>(
|
||||
LocalOpenAiImageSyncAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
body_json: effective_body_json,
|
||||
body_base64,
|
||||
trace_id,
|
||||
input,
|
||||
@@ -211,16 +212,17 @@ pub(crate) async fn build_local_image_stream_attempt_source_for_kind<'a>(
|
||||
trace_id,
|
||||
decision,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let effective_body_json = input.effective_body_json(body_json).clone();
|
||||
let Some((candidates, candidate_count)) = build_local_openai_image_candidate_attempt_source(
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
body_json,
|
||||
&effective_body_json,
|
||||
spec_metadata.api_format,
|
||||
spec_metadata.decision_kind,
|
||||
)
|
||||
@@ -237,7 +239,7 @@ pub(crate) async fn build_local_image_stream_attempt_source_for_kind<'a>(
|
||||
LocalOpenAiImageStreamAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
body_json: effective_body_json,
|
||||
body_base64,
|
||||
trace_id,
|
||||
input,
|
||||
@@ -303,21 +305,21 @@ impl LocalOpenAiImageSyncAttemptSource<'_> {
|
||||
let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
self.body_base64,
|
||||
self.trace_id,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let provider_api_format = payload.provider_api_format.as_deref().unwrap_or_default();
|
||||
let built = if provider_api_format == "gemini:generate_content" {
|
||||
build_gemini_sync_plan_from_decision(self.parts, self.body_json, payload)
|
||||
build_gemini_sync_plan_from_decision(self.parts, &self.body_json, payload)
|
||||
} else {
|
||||
build_passthrough_sync_plan_from_decision(self.parts, payload)
|
||||
};
|
||||
@@ -345,23 +347,23 @@ impl LocalOpenAiImageStreamAttemptSource<'_> {
|
||||
let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
self.body_base64,
|
||||
self.trace_id,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let provider_api_format = payload.provider_api_format.as_deref().unwrap_or_default();
|
||||
let built = if provider_api_format == "gemini:generate_content" {
|
||||
build_gemini_stream_plan_from_decision(self.parts, self.body_json, payload)
|
||||
build_gemini_stream_plan_from_decision(self.parts, &self.body_json, payload)
|
||||
} else {
|
||||
build_standard_stream_plan_from_decision(self.parts, self.body_json, payload, false)
|
||||
build_standard_stream_plan_from_decision(self.parts, &self.body_json, payload, false)
|
||||
};
|
||||
match built {
|
||||
Ok(value) => Ok(value),
|
||||
@@ -400,10 +402,11 @@ pub(crate) async fn maybe_build_sync_local_image_decision_payload(
|
||||
trace_id,
|
||||
decision,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
|
||||
let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source(
|
||||
state,
|
||||
@@ -429,7 +432,7 @@ pub(crate) async fn maybe_build_sync_local_image_decision_payload(
|
||||
attempt,
|
||||
spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
{
|
||||
return Ok(Some(payload));
|
||||
}
|
||||
@@ -460,10 +463,11 @@ pub(crate) async fn maybe_build_stream_local_image_decision_payload(
|
||||
trace_id,
|
||||
decision,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
|
||||
let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source(
|
||||
state,
|
||||
@@ -489,7 +493,7 @@ pub(crate) async fn maybe_build_stream_local_image_decision_payload(
|
||||
attempt,
|
||||
spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
{
|
||||
return Ok(Some(payload));
|
||||
}
|
||||
@@ -516,10 +520,11 @@ async fn build_local_sync_plan_and_reports(
|
||||
trace_id,
|
||||
decision,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(Vec::new());
|
||||
};
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
|
||||
let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source(
|
||||
state,
|
||||
@@ -546,7 +551,7 @@ async fn build_local_sync_plan_and_reports(
|
||||
attempt,
|
||||
spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
@@ -592,10 +597,11 @@ async fn build_local_stream_plan_and_reports(
|
||||
trace_id,
|
||||
decision,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(Vec::new());
|
||||
};
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
|
||||
let Some((mut source, _)) = build_local_openai_image_candidate_attempt_source(
|
||||
state,
|
||||
@@ -622,7 +628,7 @@ async fn build_local_stream_plan_and_reports(
|
||||
attempt,
|
||||
spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
use crate::ai_serving::build_request_trace_proxy_value;
|
||||
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
|
||||
use crate::ai_serving::planner::report_context::{
|
||||
build_local_execution_report_context, LocalExecutionReportContextParts,
|
||||
};
|
||||
@@ -10,7 +11,9 @@ use crate::ai_serving::transport::{
|
||||
resolve_transport_execution_timeouts, resolve_transport_profile,
|
||||
};
|
||||
use crate::ai_serving::{ai_local_execution_contract_for_formats, PlannerAppState};
|
||||
use crate::{append_execution_contract_fields_to_value, AiExecutionDecision, AppState};
|
||||
use crate::{
|
||||
append_execution_contract_fields_to_value, AiExecutionDecision, AppState, GatewayError,
|
||||
};
|
||||
|
||||
use super::request::resolve_local_openai_image_candidate_payload_parts;
|
||||
use super::support::{LocalOpenAiImageCandidateAttempt, LocalOpenAiImageDecisionInput};
|
||||
@@ -25,11 +28,11 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
|
||||
input: &LocalOpenAiImageDecisionInput,
|
||||
attempt: LocalOpenAiImageCandidateAttempt,
|
||||
spec: LocalOpenAiImageSpec,
|
||||
) -> Option<AiExecutionDecision> {
|
||||
) -> Result<Option<AiExecutionDecision>, GatewayError> {
|
||||
let spec_metadata = local_openai_image_spec_metadata(spec);
|
||||
let planner_state = PlannerAppState::new(state);
|
||||
let attempt_identity = attempt.attempt_identity();
|
||||
let resolved = resolve_local_openai_image_candidate_payload_parts(
|
||||
let Some(resolved) = resolve_local_openai_image_candidate_payload_parts(
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
@@ -39,7 +42,10 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
|
||||
&attempt,
|
||||
spec,
|
||||
)
|
||||
.await?;
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let LocalOpenAiImageCandidateAttempt {
|
||||
eligible,
|
||||
candidate_id,
|
||||
@@ -57,7 +63,10 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
|
||||
.app()
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&transport)
|
||||
.await;
|
||||
let transport_profile = resolve_transport_profile(&transport);
|
||||
let transport_profile = resolved
|
||||
.transport_profile
|
||||
.clone()
|
||||
.or_else(|| resolve_transport_profile(&transport));
|
||||
let mut extra_fields = serde_json::Map::new();
|
||||
if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) {
|
||||
extra_fields.insert("proxy".to_string(), proxy_value);
|
||||
@@ -88,6 +97,7 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
|
||||
.get("stream")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(spec_metadata.require_streaming);
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let report_context = append_execution_contract_fields_to_value(
|
||||
build_local_execution_report_context(LocalExecutionReportContextParts {
|
||||
auth_context: &input.auth_context,
|
||||
@@ -114,7 +124,7 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
|
||||
body_rules: transport.endpoint.body_rules.as_ref(),
|
||||
provider_request_method: Some(serde_json::Value::String(parts.method.to_string())),
|
||||
provider_request_headers: Some(&resolved.provider_request_headers),
|
||||
original_headers: &parts.headers,
|
||||
original_headers: effective_headers,
|
||||
request_path: Some(parts.uri.path()),
|
||||
request_query_string: parts.uri.query(),
|
||||
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
|
||||
@@ -134,39 +144,39 @@ pub(super) async fn maybe_build_local_openai_image_decision_payload_for_candidat
|
||||
provider_api_format.as_str(),
|
||||
);
|
||||
|
||||
Some(build_ai_execution_decision_response(
|
||||
AiExecutionDecisionResponseParts {
|
||||
decision_is_stream: spec_metadata.require_streaming,
|
||||
decision_kind: spec_metadata.decision_kind.to_string(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: candidate_id.clone(),
|
||||
provider_name: transport.provider.name.clone(),
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
key_id: candidate.key_id.clone(),
|
||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||
upstream_url: resolved.upstream_url,
|
||||
provider_request_method: Some(parts.method.to_string()),
|
||||
auth_header: Some(resolved.auth_header),
|
||||
auth_value: Some(resolved.auth_value),
|
||||
provider_api_format,
|
||||
client_api_format: spec_metadata.api_format.to_string(),
|
||||
model_name: resolved.requested_model,
|
||||
mapped_model: resolved.mapped_model,
|
||||
prompt_cache_key: None,
|
||||
provider_request_headers: resolved.provider_request_headers,
|
||||
provider_request_body: Some(resolved.provider_request_body),
|
||||
provider_request_body_base64: None,
|
||||
content_type: Some("application/json".to_string()),
|
||||
proxy,
|
||||
transport_profile,
|
||||
timeouts: resolve_transport_execution_timeouts(&transport),
|
||||
upstream_is_stream,
|
||||
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
||||
report_context: Some(report_context),
|
||||
auth_context: input.auth_context.clone(),
|
||||
},
|
||||
))
|
||||
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
|
||||
decision_is_stream: spec_metadata.require_streaming,
|
||||
decision_kind: spec_metadata.decision_kind.to_string(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: candidate_id.clone(),
|
||||
provider_name: transport.provider.name.clone(),
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
key_id: candidate.key_id.clone(),
|
||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||
upstream_url: resolved.upstream_url,
|
||||
provider_request_method: Some(parts.method.to_string()),
|
||||
auth_header: Some(resolved.auth_header),
|
||||
auth_value: Some(resolved.auth_value),
|
||||
provider_api_format,
|
||||
client_api_format: spec_metadata.api_format.to_string(),
|
||||
model_name: resolved.requested_model,
|
||||
mapped_model: resolved.mapped_model,
|
||||
prompt_cache_key: None,
|
||||
provider_request_headers: resolved.provider_request_headers,
|
||||
provider_request_body: Some(resolved.provider_request_body),
|
||||
provider_request_body_base64: None,
|
||||
content_type: Some("application/json".to_string()),
|
||||
proxy,
|
||||
transport_profile,
|
||||
timeouts: resolve_transport_execution_timeouts(&transport),
|
||||
upstream_is_stream,
|
||||
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
||||
report_context: Some(report_context),
|
||||
auth_context: input.auth_context.clone(),
|
||||
});
|
||||
apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
|
||||
Ok(Some(decision))
|
||||
}
|
||||
|
||||
@@ -1,17 +1,19 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use aether_contracts::ResolvedTransportProfile;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::ai_serving::planner::candidate_preparation::{
|
||||
prepare_header_authenticated_candidate, OauthPreparationContext,
|
||||
};
|
||||
use crate::ai_serving::planner::spec_metadata::local_openai_image_spec_metadata;
|
||||
use crate::ai_serving::pure::normalize_openai_image_request_with_options;
|
||||
use crate::ai_serving::transport::{
|
||||
build_openai_image_headers, build_openai_image_upstream_url,
|
||||
build_standard_provider_request_headers, openai_image_transport_unsupported_reason,
|
||||
resolve_openai_image_auth, ProviderOpenAiImageHeadersInput,
|
||||
StandardProviderRequestHeadersInput,
|
||||
build_grok_browser_headers, build_grok_upstream_url, build_openai_image_headers,
|
||||
build_openai_image_upstream_url, build_standard_provider_request_headers,
|
||||
openai_image_transport_unsupported_reason, resolve_openai_image_auth, GrokHeaderInput,
|
||||
ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput, GROK_CHAT_PATH,
|
||||
};
|
||||
use crate::ai_serving::{
|
||||
apply_codex_openai_responses_special_body_edits, apply_codex_openai_responses_special_headers,
|
||||
@@ -21,6 +23,7 @@ use crate::ai_serving::{
|
||||
normalize_openai_image_request, request_conversion_direct_auth, CandidateFailureDiagnostic,
|
||||
GatewayProviderTransportSnapshot, PlannerAppState, RequestConversionKind,
|
||||
};
|
||||
use crate::image_capabilities::openai_image_normalize_options_for_provider;
|
||||
use crate::AppState;
|
||||
|
||||
use super::support::{
|
||||
@@ -43,6 +46,7 @@ pub(super) struct LocalOpenAiImageCandidatePayloadParts {
|
||||
pub(super) provider_request_body: Value,
|
||||
pub(super) upstream_url: String,
|
||||
pub(super) input_summary: Value,
|
||||
pub(super) transport_profile: Option<ResolvedTransportProfile>,
|
||||
}
|
||||
|
||||
pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
|
||||
@@ -59,6 +63,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
|
||||
let candidate = &attempt.eligible.candidate;
|
||||
let transport = &attempt.eligible.transport;
|
||||
let provider_api_format = attempt.eligible.provider_api_format.as_str();
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
|
||||
if provider_api_format == "gemini:generate_content" {
|
||||
return resolve_local_openai_image_to_gemini_candidate_payload_parts(
|
||||
@@ -120,8 +125,13 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
|
||||
let auth_header = prepared_candidate.auth_header;
|
||||
let auth_value = prepared_candidate.auth_value;
|
||||
|
||||
let Some(normalized_request) = normalize_openai_image_request(parts, body_json, body_base64)
|
||||
else {
|
||||
let normalized_request = normalize_openai_image_request_with_options(
|
||||
parts,
|
||||
body_json,
|
||||
body_base64,
|
||||
openai_image_normalize_options_for_provider(&transport.provider.provider_type),
|
||||
);
|
||||
let Some(normalized_request) = normalized_request else {
|
||||
mark_skipped_local_openai_image_candidate_with_failure_diagnostic(
|
||||
state,
|
||||
input,
|
||||
@@ -145,8 +155,16 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("chatgpt_web");
|
||||
let is_grok = transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("grok");
|
||||
let transport_profile = crate::ai_serving::transport::resolve_transport_profile(transport);
|
||||
let upstream_url = if is_chatgpt_web {
|
||||
chatgpt_web_image_internal_url(&transport.endpoint.base_url)
|
||||
} else if is_grok {
|
||||
build_grok_upstream_url(transport, GROK_CHAT_PATH)
|
||||
} else {
|
||||
build_openai_image_upstream_url(transport, parts.uri.query())
|
||||
};
|
||||
@@ -168,16 +186,27 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
|
||||
);
|
||||
}
|
||||
|
||||
let Some(mut provider_request_headers) =
|
||||
let Some(mut provider_request_headers) = (if is_grok {
|
||||
build_grok_browser_headers(GrokHeaderInput {
|
||||
transport,
|
||||
transport_profile: transport_profile.as_ref(),
|
||||
request_headers: Some(effective_headers),
|
||||
content_type: "application/json",
|
||||
accept: "*/*",
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: body_json,
|
||||
})
|
||||
} else {
|
||||
build_openai_image_headers(ProviderOpenAiImageHeadersInput {
|
||||
headers: &parts.headers,
|
||||
headers: effective_headers,
|
||||
auth_header: &auth_header,
|
||||
auth_value: &auth_value,
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: body_json,
|
||||
})
|
||||
else {
|
||||
}) else {
|
||||
mark_skipped_local_openai_image_candidate_with_failure_diagnostic(
|
||||
state,
|
||||
input,
|
||||
@@ -197,11 +226,12 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
|
||||
};
|
||||
if is_chatgpt_web {
|
||||
provider_request_headers.insert("x-aether-chatgpt-web-image".to_string(), "1".to_string());
|
||||
} else if is_grok {
|
||||
} else {
|
||||
apply_codex_openai_responses_special_headers(
|
||||
&mut provider_request_headers,
|
||||
&provider_request_body,
|
||||
&parts.headers,
|
||||
effective_headers,
|
||||
transport.provider.provider_type.as_str(),
|
||||
spec_metadata.api_format,
|
||||
Some(trace_id),
|
||||
@@ -222,7 +252,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
|
||||
.unwrap_or_default()
|
||||
.to_string();
|
||||
|
||||
let input_summary = if is_chatgpt_web {
|
||||
let input_summary = if is_chatgpt_web || is_grok {
|
||||
provider_request_body.clone()
|
||||
} else {
|
||||
normalized_request.summary_json
|
||||
@@ -239,6 +269,7 @@ pub(super) async fn resolve_local_openai_image_candidate_payload_parts(
|
||||
provider_request_body,
|
||||
upstream_url,
|
||||
input_summary,
|
||||
transport_profile,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -256,6 +287,7 @@ async fn resolve_local_openai_image_to_gemini_candidate_payload_parts(
|
||||
let candidate = &attempt.eligible.candidate;
|
||||
let transport = &attempt.eligible.transport;
|
||||
let provider_api_format = "gemini:generate_content";
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
|
||||
let prepared_candidate = match prepare_header_authenticated_candidate(
|
||||
PlannerAppState::new(state),
|
||||
@@ -332,7 +364,7 @@ async fn resolve_local_openai_image_to_gemini_candidate_payload_parts(
|
||||
converted.body_json,
|
||||
transport.endpoint.body_rules.as_ref(),
|
||||
body_json,
|
||||
&parts.headers,
|
||||
effective_headers,
|
||||
) {
|
||||
Some(body) => body,
|
||||
None => {
|
||||
@@ -385,7 +417,7 @@ async fn resolve_local_openai_image_to_gemini_candidate_payload_parts(
|
||||
transport,
|
||||
provider_api_format,
|
||||
same_format: false,
|
||||
headers: &parts.headers,
|
||||
headers: effective_headers,
|
||||
auth_header: &prepared_candidate.auth_header,
|
||||
auth_value: &prepared_candidate.auth_value,
|
||||
extra_headers: &BTreeMap::new(),
|
||||
@@ -424,6 +456,7 @@ async fn resolve_local_openai_image_to_gemini_candidate_payload_parts(
|
||||
provider_request_body: converted.body_json,
|
||||
upstream_url,
|
||||
input_summary: converted.summary_json,
|
||||
transport_profile: None,
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -13,6 +13,7 @@ use crate::ai_serving::planner::candidate_metadata::{
|
||||
use crate::ai_serving::planner::candidate_resolution::SkippedLocalExecutionCandidate;
|
||||
use crate::ai_serving::planner::candidate_source::auth_snapshot_allows_cross_format_candidate;
|
||||
use crate::ai_serving::planner::decision_input::{
|
||||
attach_routing_policy_to_local_requested_model_input,
|
||||
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
|
||||
};
|
||||
use crate::ai_serving::planner::materialization_policy::{
|
||||
@@ -42,12 +43,16 @@ pub(super) async fn resolve_local_openai_image_decision_input(
|
||||
body_base64: Option<&str>,
|
||||
trace_id: &str,
|
||||
decision: &GatewayControlDecision,
|
||||
) -> Option<LocalOpenAiImageDecisionInput> {
|
||||
) -> Result<Option<LocalOpenAiImageDecisionInput>, GatewayError> {
|
||||
let Some(auth_context) = resolve_local_openai_image_auth_context(decision) else {
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let requested_model = resolve_requested_image_model_for_request(parts, body_json, body_base64)?;
|
||||
let Some(requested_model) =
|
||||
resolve_requested_image_model_for_request(parts, body_json, body_base64)
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let resolved_input = match resolve_local_authenticated_decision_input(
|
||||
state,
|
||||
@@ -58,21 +63,37 @@ pub(super) async fn resolve_local_openai_image_decision_input(
|
||||
.await
|
||||
{
|
||||
Ok(Some(resolved_input)) => resolved_input,
|
||||
Ok(None) => return None,
|
||||
Ok(None) => return Ok(None),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
error = ?err,
|
||||
"gateway local openai image decision auth snapshot read failed"
|
||||
);
|
||||
return None;
|
||||
return Err(err);
|
||||
}
|
||||
};
|
||||
|
||||
let mut input = build_local_requested_model_decision_input(resolved_input, requested_model);
|
||||
input.request_auth_channel = decision.request_auth_channel.clone();
|
||||
input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json));
|
||||
Some(input)
|
||||
if let Err(err) = attach_routing_policy_to_local_requested_model_input(
|
||||
state,
|
||||
parts,
|
||||
&mut input,
|
||||
body_json,
|
||||
"openai:image",
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
error = ?err,
|
||||
"gateway local openai image decision routing profile resolution failed"
|
||||
);
|
||||
return Err(err);
|
||||
}
|
||||
Ok(Some(input))
|
||||
}
|
||||
|
||||
fn resolve_local_openai_image_auth_context(
|
||||
@@ -229,6 +250,7 @@ pub(super) async fn build_local_openai_image_candidate_attempt_source<'a>(
|
||||
Some(&input.auth_snapshot),
|
||||
input.client_session_affinity.as_ref(),
|
||||
input.required_capabilities.as_ref(),
|
||||
input.routing_policy.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
@@ -305,6 +327,7 @@ async fn materialize_local_openai_image_candidate_attempts(
|
||||
Some(&input.auth_snapshot),
|
||||
input.client_session_affinity.as_ref(),
|
||||
input.required_capabilities.as_ref(),
|
||||
input.routing_policy.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
|
||||
@@ -27,7 +27,7 @@ use self::support::{
|
||||
pub(crate) struct LocalVideoCreateSyncAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_json: serde_json::Value,
|
||||
trace_id: &'a str,
|
||||
input: LocalVideoCreateDecisionInput,
|
||||
spec: LocalVideoCreateSpec,
|
||||
@@ -65,16 +65,17 @@ pub(crate) async fn build_local_video_sync_attempt_source_for_kind<'a>(
|
||||
let Some(input) = resolve_local_video_create_decision_input(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let effective_body_json = input.effective_body_json(body_json).clone();
|
||||
let Some((candidates, candidate_count)) = build_local_video_create_candidate_attempt_source(
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
body_json,
|
||||
&effective_body_json,
|
||||
spec_metadata.api_format,
|
||||
spec_metadata.decision_kind,
|
||||
)
|
||||
@@ -91,7 +92,7 @@ pub(crate) async fn build_local_video_sync_attempt_source_for_kind<'a>(
|
||||
LocalVideoCreateSyncAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
body_json: effective_body_json,
|
||||
trace_id,
|
||||
input,
|
||||
spec,
|
||||
@@ -133,13 +134,13 @@ impl LocalVideoCreateSyncAttemptSource<'_> {
|
||||
let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
self.trace_id,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
@@ -175,10 +176,11 @@ pub(crate) async fn maybe_build_sync_local_video_decision_payload(
|
||||
let Some(input) = resolve_local_video_create_decision_input(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
|
||||
let Some((mut source, _)) = build_local_video_create_candidate_attempt_source(
|
||||
state,
|
||||
@@ -197,7 +199,7 @@ pub(crate) async fn maybe_build_sync_local_video_decision_payload(
|
||||
if let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate(
|
||||
state, parts, body_json, trace_id, &input, attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
{
|
||||
return Ok(Some(payload));
|
||||
}
|
||||
@@ -218,10 +220,11 @@ async fn build_local_sync_plan_and_reports(
|
||||
let Some(input) = resolve_local_video_create_decision_input(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(Vec::new());
|
||||
};
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
|
||||
let Some((mut source, _)) = build_local_video_create_candidate_attempt_source(
|
||||
state,
|
||||
@@ -241,7 +244,7 @@ async fn build_local_sync_plan_and_reports(
|
||||
let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate(
|
||||
state, parts, body_json, trace_id, &input, attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
use crate::ai_serving::build_request_trace_proxy_value;
|
||||
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
|
||||
use crate::ai_serving::planner::report_context::{
|
||||
build_local_execution_report_context, LocalExecutionReportContextParts,
|
||||
};
|
||||
@@ -10,7 +11,7 @@ use crate::ai_serving::transport::{
|
||||
resolve_transport_execution_timeouts, resolve_transport_profile,
|
||||
};
|
||||
use crate::ai_serving::{ai_local_execution_contract_for_formats, PlannerAppState};
|
||||
use crate::{AiExecutionDecision, AppState};
|
||||
use crate::{AiExecutionDecision, AppState, GatewayError};
|
||||
|
||||
use super::request::resolve_local_video_create_candidate_payload_parts;
|
||||
use super::support::{LocalVideoCreateCandidateAttempt, LocalVideoCreateDecisionInput};
|
||||
@@ -24,14 +25,17 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
|
||||
input: &LocalVideoCreateDecisionInput,
|
||||
attempt: LocalVideoCreateCandidateAttempt,
|
||||
spec: LocalVideoCreateSpec,
|
||||
) -> Option<AiExecutionDecision> {
|
||||
) -> Result<Option<AiExecutionDecision>, GatewayError> {
|
||||
let spec_metadata = local_video_create_spec_metadata(spec);
|
||||
let planner_state = PlannerAppState::new(state);
|
||||
let attempt_identity = attempt.attempt_identity();
|
||||
let resolved = resolve_local_video_create_candidate_payload_parts(
|
||||
let Some(resolved) = resolve_local_video_create_candidate_payload_parts(
|
||||
state, parts, body_json, trace_id, input, &attempt, spec,
|
||||
)
|
||||
.await?;
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let LocalVideoCreateCandidateAttempt {
|
||||
eligible,
|
||||
candidate_id,
|
||||
@@ -50,6 +54,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
|
||||
if let Some(proxy_value) = build_request_trace_proxy_value(Some(&transport), proxy.as_ref()) {
|
||||
extra_fields.insert("proxy".to_string(), proxy_value);
|
||||
}
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let report_context = build_local_execution_report_context(LocalExecutionReportContextParts {
|
||||
auth_context: &input.auth_context,
|
||||
request_id: trace_id,
|
||||
@@ -75,7 +80,7 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
|
||||
body_rules: transport.endpoint.body_rules.as_ref(),
|
||||
provider_request_method: None,
|
||||
provider_request_headers: None,
|
||||
original_headers: &parts.headers,
|
||||
original_headers: effective_headers,
|
||||
request_path: Some(parts.uri.path()),
|
||||
request_query_string: parts.uri.query(),
|
||||
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
|
||||
@@ -99,45 +104,45 @@ pub(super) async fn maybe_build_local_video_create_decision_payload_for_candidat
|
||||
upstream_url,
|
||||
} = resolved;
|
||||
|
||||
Some(build_ai_execution_decision_response(
|
||||
AiExecutionDecisionResponseParts {
|
||||
decision_is_stream: false,
|
||||
decision_kind: spec_metadata.decision_kind.to_string(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: candidate_id.clone(),
|
||||
provider_name: transport.provider.name.clone(),
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
key_id: candidate.key_id.clone(),
|
||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||
upstream_url,
|
||||
provider_request_method: Some(parts.method.to_string()),
|
||||
auth_header: Some(auth_header),
|
||||
auth_value: Some(auth_value),
|
||||
provider_api_format: spec_metadata.api_format.to_string(),
|
||||
client_api_format: spec_metadata.api_format.to_string(),
|
||||
model_name: input.requested_model.clone(),
|
||||
mapped_model,
|
||||
prompt_cache_key: None,
|
||||
provider_request_headers,
|
||||
provider_request_body: Some(provider_request_body),
|
||||
provider_request_body_base64: None,
|
||||
content_type: parts
|
||||
.headers
|
||||
.get(http::header::CONTENT_TYPE)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned),
|
||||
proxy,
|
||||
transport_profile,
|
||||
timeouts: resolve_transport_execution_timeouts(&transport),
|
||||
upstream_is_stream: false,
|
||||
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
||||
report_context: Some(report_context),
|
||||
auth_context: input.auth_context.clone(),
|
||||
},
|
||||
))
|
||||
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
|
||||
decision_is_stream: false,
|
||||
decision_kind: spec_metadata.decision_kind.to_string(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: candidate_id.clone(),
|
||||
provider_name: transport.provider.name.clone(),
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
key_id: candidate.key_id.clone(),
|
||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||
upstream_url,
|
||||
provider_request_method: Some(parts.method.to_string()),
|
||||
auth_header: Some(auth_header),
|
||||
auth_value: Some(auth_value),
|
||||
provider_api_format: spec_metadata.api_format.to_string(),
|
||||
client_api_format: spec_metadata.api_format.to_string(),
|
||||
model_name: input.requested_model.clone(),
|
||||
mapped_model,
|
||||
prompt_cache_key: None,
|
||||
provider_request_headers,
|
||||
provider_request_body: Some(provider_request_body),
|
||||
provider_request_body_base64: None,
|
||||
content_type: parts
|
||||
.headers
|
||||
.get(http::header::CONTENT_TYPE)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned),
|
||||
proxy,
|
||||
transport_profile,
|
||||
timeouts: resolve_transport_execution_timeouts(&transport),
|
||||
upstream_is_stream: false,
|
||||
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
||||
report_context: Some(report_context),
|
||||
auth_context: input.auth_context.clone(),
|
||||
});
|
||||
apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
|
||||
Ok(Some(decision))
|
||||
}
|
||||
|
||||
@@ -41,6 +41,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
||||
let spec_metadata = local_video_create_spec_metadata(spec);
|
||||
let candidate = &attempt.eligible.candidate;
|
||||
let transport = &attempt.eligible.transport;
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
|
||||
let provider_family = provider_video_create_family(spec.family);
|
||||
let transport_unsupported_reason = video_create_transport_unsupported_reason(
|
||||
@@ -124,7 +125,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
||||
provider_family,
|
||||
&mapped_model,
|
||||
transport.endpoint.body_rules.as_ref(),
|
||||
Some(&parts.headers),
|
||||
Some(effective_headers),
|
||||
) else {
|
||||
mark_skipped_local_video_candidate_with_failure_diagnostic(
|
||||
state,
|
||||
@@ -146,7 +147,7 @@ pub(super) async fn resolve_local_video_create_candidate_payload_parts(
|
||||
|
||||
let Some(provider_request_headers) =
|
||||
build_video_create_headers(ProviderVideoCreateHeadersInput {
|
||||
headers: &parts.headers,
|
||||
headers: effective_headers,
|
||||
auth_header: &auth_header,
|
||||
auth_value: &auth_value,
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
|
||||
@@ -15,6 +15,7 @@ use crate::ai_serving::planner::candidate_metadata::{
|
||||
use crate::ai_serving::planner::candidate_resolution::SkippedLocalExecutionCandidate;
|
||||
use crate::ai_serving::planner::common::extract_requested_model_from_request;
|
||||
use crate::ai_serving::planner::decision_input::{
|
||||
attach_routing_policy_to_local_requested_model_input,
|
||||
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
|
||||
};
|
||||
use crate::ai_serving::planner::materialization_policy::{
|
||||
@@ -41,19 +42,21 @@ pub(super) async fn resolve_local_video_create_decision_input(
|
||||
decision: &GatewayControlDecision,
|
||||
body_json: &serde_json::Value,
|
||||
spec: LocalVideoCreateSpec,
|
||||
) -> Option<LocalVideoCreateDecisionInput> {
|
||||
) -> Result<Option<LocalVideoCreateDecisionInput>, GatewayError> {
|
||||
let spec_metadata = local_video_create_spec_metadata(spec);
|
||||
let Some(auth_context) = resolve_local_video_create_auth_context(decision, spec.family) else {
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let requested_model = extract_requested_model_from_request(
|
||||
let Some(requested_model) = extract_requested_model_from_request(
|
||||
parts,
|
||||
body_json,
|
||||
spec_metadata
|
||||
.requested_model_family
|
||||
.expect("video specs should declare requested-model family"),
|
||||
)?;
|
||||
) else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let resolved_input = match resolve_local_authenticated_decision_input(
|
||||
state,
|
||||
@@ -64,7 +67,7 @@ pub(super) async fn resolve_local_video_create_decision_input(
|
||||
.await
|
||||
{
|
||||
Ok(Some(resolved_input)) => resolved_input,
|
||||
Ok(None) => return None,
|
||||
Ok(None) => return Ok(None),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
@@ -72,14 +75,31 @@ pub(super) async fn resolve_local_video_create_decision_input(
|
||||
error = ?err,
|
||||
"gateway local video decision auth snapshot read failed"
|
||||
);
|
||||
return None;
|
||||
return Err(err);
|
||||
}
|
||||
};
|
||||
|
||||
let mut input = build_local_requested_model_decision_input(resolved_input, requested_model);
|
||||
input.request_auth_channel = decision.request_auth_channel.clone();
|
||||
input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json));
|
||||
Some(input)
|
||||
if let Err(err) = attach_routing_policy_to_local_requested_model_input(
|
||||
state,
|
||||
parts,
|
||||
&mut input,
|
||||
body_json,
|
||||
spec_metadata.api_format,
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
decision_kind = spec_metadata.decision_kind,
|
||||
error = ?err,
|
||||
"gateway local video decision routing profile resolution failed"
|
||||
);
|
||||
return Err(err);
|
||||
}
|
||||
Ok(Some(input))
|
||||
}
|
||||
|
||||
fn resolve_local_video_create_auth_context(
|
||||
@@ -196,6 +216,7 @@ pub(super) async fn build_local_video_create_candidate_attempt_source<'a>(
|
||||
Some(&input.auth_snapshot),
|
||||
input.client_session_affinity.as_ref(),
|
||||
input.required_capabilities.as_ref(),
|
||||
input.routing_policy.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
@@ -261,6 +282,7 @@ async fn materialize_local_video_create_candidate_attempts(
|
||||
Some(&input.auth_snapshot),
|
||||
input.client_session_affinity.as_ref(),
|
||||
input.required_capabilities.as_ref(),
|
||||
input.routing_policy.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
|
||||
@@ -29,7 +29,7 @@ pub(crate) struct LocalStandardSyncAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_json: serde_json::Value,
|
||||
input: LocalStandardDecisionInput,
|
||||
spec: LocalStandardSpec,
|
||||
requested_model_family: RequestedModelFamily,
|
||||
@@ -40,7 +40,7 @@ pub(crate) struct LocalStandardStreamAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_json: serde_json::Value,
|
||||
input: LocalStandardDecisionInput,
|
||||
spec: LocalStandardSpec,
|
||||
requested_model_family: RequestedModelFamily,
|
||||
@@ -61,7 +61,7 @@ pub(crate) async fn build_local_sync_attempt_source<'a>(
|
||||
.expect("standard spec metadata should include requested-model family");
|
||||
let Some(input) =
|
||||
resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
@@ -82,9 +82,15 @@ pub(crate) async fn build_local_sync_attempt_source<'a>(
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let (candidates, candidate_count) =
|
||||
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec)
|
||||
.await?;
|
||||
let effective_body_json = input.effective_body_json(body_json).clone();
|
||||
let (candidates, candidate_count) = build_local_standard_candidate_attempt_source(
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
&effective_body_json,
|
||||
spec,
|
||||
)
|
||||
.await?;
|
||||
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
|
||||
if candidate_count == 0 {
|
||||
return Ok(None);
|
||||
@@ -95,7 +101,7 @@ pub(crate) async fn build_local_sync_attempt_source<'a>(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
body_json: effective_body_json,
|
||||
input,
|
||||
spec,
|
||||
requested_model_family,
|
||||
@@ -119,7 +125,7 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
|
||||
.expect("standard spec metadata should include requested-model family");
|
||||
let Some(input) =
|
||||
resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
@@ -140,9 +146,15 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let (candidates, candidate_count) =
|
||||
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec)
|
||||
.await?;
|
||||
let effective_body_json = input.effective_body_json(body_json).clone();
|
||||
let (candidates, candidate_count) = build_local_standard_candidate_attempt_source(
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
&effective_body_json,
|
||||
spec,
|
||||
)
|
||||
.await?;
|
||||
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
|
||||
if candidate_count == 0 {
|
||||
return Ok(None);
|
||||
@@ -153,7 +165,7 @@ pub(crate) async fn build_local_stream_attempt_source<'a>(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
body_json: effective_body_json,
|
||||
input,
|
||||
spec,
|
||||
requested_model_family,
|
||||
@@ -228,19 +240,19 @@ impl LocalStandardSyncAttemptSource<'_> {
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
match build_sync_plan_from_requested_model_family(
|
||||
self.requested_model_family,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
payload,
|
||||
) {
|
||||
Ok(value) => Ok(value),
|
||||
@@ -265,19 +277,19 @@ impl LocalStandardStreamAttemptSource<'_> {
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
match build_stream_plan_from_requested_model_family(
|
||||
self.requested_model_family,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
payload,
|
||||
) {
|
||||
Ok(value) => Ok(value),
|
||||
@@ -309,7 +321,7 @@ pub(crate) async fn maybe_build_sync_via_standard_family_payload(
|
||||
|
||||
let Some(input) =
|
||||
resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
@@ -322,6 +334,7 @@ pub(crate) async fn maybe_build_sync_via_standard_family_payload(
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
let (mut source, candidate_count) =
|
||||
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec)
|
||||
.await?;
|
||||
@@ -331,7 +344,7 @@ pub(crate) async fn maybe_build_sync_via_standard_family_payload(
|
||||
if let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
|
||||
state, parts, trace_id, body_json, &input, attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
{
|
||||
return Ok(Some(payload));
|
||||
}
|
||||
@@ -358,7 +371,7 @@ pub(crate) async fn maybe_build_stream_via_standard_family_payload(
|
||||
|
||||
let Some(input) =
|
||||
resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
@@ -371,6 +384,7 @@ pub(crate) async fn maybe_build_stream_via_standard_family_payload(
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
let (mut source, candidate_count) =
|
||||
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec)
|
||||
.await?;
|
||||
@@ -380,7 +394,7 @@ pub(crate) async fn maybe_build_stream_via_standard_family_payload(
|
||||
if let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
|
||||
state, parts, trace_id, body_json, &input, attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
{
|
||||
return Ok(Some(payload));
|
||||
}
|
||||
@@ -405,7 +419,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
|
||||
.expect("standard spec metadata should include requested-model family");
|
||||
let Some(input) =
|
||||
resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
@@ -426,6 +440,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
let (mut source, candidate_count) =
|
||||
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec)
|
||||
.await?;
|
||||
@@ -438,7 +453,7 @@ pub(crate) async fn build_local_sync_plan_and_reports(
|
||||
let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
|
||||
state, parts, trace_id, body_json, &input, attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
@@ -479,7 +494,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
|
||||
.expect("standard spec metadata should include requested-model family");
|
||||
let Some(input) =
|
||||
resolve_local_standard_decision_input(state, parts, trace_id, decision, body_json, spec)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
@@ -500,6 +515,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
let (mut source, candidate_count) =
|
||||
build_local_standard_candidate_attempt_source(state, trace_id, &input, body_json, spec)
|
||||
.await?;
|
||||
@@ -512,7 +528,7 @@ pub(crate) async fn build_local_stream_plan_and_reports(
|
||||
let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
|
||||
state, parts, trace_id, body_json, &input, attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
@@ -16,6 +16,7 @@ use crate::ai_serving::planner::candidate_source::{
|
||||
};
|
||||
use crate::ai_serving::planner::common::extract_requested_model_from_request;
|
||||
use crate::ai_serving::planner::decision_input::{
|
||||
attach_routing_policy_to_local_requested_model_input,
|
||||
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
|
||||
};
|
||||
use crate::ai_serving::planner::materialization_policy::{
|
||||
@@ -39,19 +40,21 @@ pub(super) async fn resolve_local_standard_decision_input(
|
||||
decision: &GatewayControlDecision,
|
||||
body_json: &serde_json::Value,
|
||||
spec: LocalStandardSpec,
|
||||
) -> Option<LocalStandardDecisionInput> {
|
||||
) -> Result<Option<LocalStandardDecisionInput>, GatewayError> {
|
||||
let spec_metadata = local_standard_spec_metadata(spec);
|
||||
let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else {
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let requested_model = extract_requested_model_from_request(
|
||||
let Some(requested_model) = extract_requested_model_from_request(
|
||||
parts,
|
||||
body_json,
|
||||
spec_metadata
|
||||
.requested_model_family
|
||||
.expect("standard specs should declare requested-model family"),
|
||||
)?;
|
||||
) else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let resolved_input = match resolve_local_authenticated_decision_input(
|
||||
state,
|
||||
@@ -62,7 +65,7 @@ pub(super) async fn resolve_local_standard_decision_input(
|
||||
.await
|
||||
{
|
||||
Ok(Some(resolved_input)) => resolved_input,
|
||||
Ok(None) => return None,
|
||||
Ok(None) => return Ok(None),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
@@ -70,14 +73,31 @@ pub(super) async fn resolve_local_standard_decision_input(
|
||||
error = ?err,
|
||||
"gateway local standard decision auth snapshot read failed"
|
||||
);
|
||||
return None;
|
||||
return Err(err);
|
||||
}
|
||||
};
|
||||
|
||||
let mut input = build_local_requested_model_decision_input(resolved_input, requested_model);
|
||||
input.request_auth_channel = decision.request_auth_channel.clone();
|
||||
input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json));
|
||||
Some(input)
|
||||
if let Err(err) = attach_routing_policy_to_local_requested_model_input(
|
||||
state,
|
||||
parts,
|
||||
&mut input,
|
||||
body_json,
|
||||
spec_metadata.api_format,
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
api_format = spec_metadata.api_format,
|
||||
error = ?err,
|
||||
"gateway local standard decision routing profile resolution failed"
|
||||
);
|
||||
return Err(err);
|
||||
}
|
||||
Ok(Some(input))
|
||||
}
|
||||
|
||||
pub(super) async fn materialize_local_standard_candidate_attempts(
|
||||
@@ -104,6 +124,7 @@ pub(super) async fn materialize_local_standard_candidate_attempts(
|
||||
spec_metadata.require_streaming,
|
||||
input.required_capabilities.as_ref(),
|
||||
&input.auth_snapshot,
|
||||
input.routing_policy.as_ref(),
|
||||
input.client_session_affinity.as_ref(),
|
||||
false,
|
||||
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
|
||||
@@ -128,6 +149,7 @@ pub(super) async fn materialize_local_standard_candidate_attempts(
|
||||
Some(&input.auth_snapshot),
|
||||
input.client_session_affinity.as_ref(),
|
||||
input.required_capabilities.as_ref(),
|
||||
input.routing_policy.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
@@ -228,6 +250,7 @@ pub(super) async fn build_local_standard_candidate_attempt_source<'a>(
|
||||
&input.auth_snapshot,
|
||||
input.client_session_affinity.as_ref(),
|
||||
input.required_capabilities.as_ref(),
|
||||
input.routing_policy.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
@@ -320,6 +343,7 @@ async fn maybe_append_gemini_image_openai_image_preselection(
|
||||
spec_metadata.require_streaming,
|
||||
input.required_capabilities.as_ref(),
|
||||
&input.auth_snapshot,
|
||||
input.routing_policy.as_ref(),
|
||||
input.client_session_affinity.as_ref(),
|
||||
false,
|
||||
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
|
||||
|
||||
@@ -3,6 +3,7 @@ use crate::ai_serving::planner::candidate_materialization::{
|
||||
mark_skipped_local_execution_candidate, mark_skipped_local_execution_candidate_with_extra_data,
|
||||
mark_skipped_local_execution_candidate_with_failure_diagnostic,
|
||||
};
|
||||
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
|
||||
use crate::ai_serving::planner::materialization_policy::{
|
||||
build_local_candidate_persistence_policy, LocalCandidatePersistencePolicyKind,
|
||||
};
|
||||
@@ -24,7 +25,7 @@ use crate::ai_serving::{
|
||||
};
|
||||
use crate::{
|
||||
append_execution_contract_fields_to_value, append_local_failover_policy_to_value,
|
||||
AiExecutionDecision, AppState,
|
||||
AiExecutionDecision, AppState, GatewayError,
|
||||
};
|
||||
|
||||
use super::request::resolve_local_standard_candidate_payload_parts;
|
||||
@@ -38,7 +39,7 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
|
||||
input: &LocalStandardDecisionInput,
|
||||
attempt: LocalStandardCandidateAttempt,
|
||||
spec: LocalStandardSpec,
|
||||
) -> Option<AiExecutionDecision> {
|
||||
) -> Result<Option<AiExecutionDecision>, GatewayError> {
|
||||
let spec_metadata = local_standard_spec_metadata(spec);
|
||||
if api_format_alias_matches(
|
||||
&attempt.eligible.provider_api_format,
|
||||
@@ -70,10 +71,13 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
|
||||
..
|
||||
} = &attempt;
|
||||
let candidate = &eligible.candidate;
|
||||
let resolved = resolve_local_standard_candidate_payload_parts(
|
||||
let Some(resolved) = resolve_local_standard_candidate_payload_parts(
|
||||
state, parts, trace_id, body_json, input, &attempt, spec,
|
||||
)
|
||||
.await?;
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let proxy = state
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport)
|
||||
.await;
|
||||
@@ -93,6 +97,7 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
|
||||
spec_metadata.api_format,
|
||||
resolved.provider_api_format.as_str(),
|
||||
);
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let report_context = append_local_failover_policy_to_value(
|
||||
append_execution_contract_fields_to_value(
|
||||
build_local_execution_report_context(LocalExecutionReportContextParts {
|
||||
@@ -120,7 +125,7 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
|
||||
body_rules: resolved.transport.endpoint.body_rules.as_ref(),
|
||||
provider_request_method: Some(serde_json::Value::Null),
|
||||
provider_request_headers: Some(&resolved.provider_request_headers),
|
||||
original_headers: &parts.headers,
|
||||
original_headers: effective_headers,
|
||||
request_path: Some(parts.uri.path()),
|
||||
request_query_string: parts.uri.query(),
|
||||
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
|
||||
@@ -144,7 +149,10 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
|
||||
),
|
||||
&resolved.transport,
|
||||
);
|
||||
let transport_profile = resolve_transport_profile(&resolved.transport);
|
||||
let transport_profile = resolved
|
||||
.transport_profile
|
||||
.clone()
|
||||
.or_else(|| resolve_transport_profile(&resolved.transport));
|
||||
let timeouts = resolve_transport_execution_timeouts(&resolved.transport);
|
||||
let super::request::LocalStandardCandidatePayloadParts {
|
||||
auth_header,
|
||||
@@ -157,43 +165,44 @@ pub(super) async fn maybe_build_local_standard_decision_payload_for_candidate(
|
||||
upstream_is_stream,
|
||||
envelope_name: _,
|
||||
transport,
|
||||
transport_profile: _,
|
||||
} = resolved;
|
||||
|
||||
Some(build_ai_execution_decision_response(
|
||||
AiExecutionDecisionResponseParts {
|
||||
decision_is_stream: spec_metadata.require_streaming,
|
||||
decision_kind: spec_metadata.decision_kind.to_string(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: candidate_id.to_string(),
|
||||
provider_name: candidate.provider_name.clone(),
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
key_id: candidate.key_id.clone(),
|
||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||
upstream_url,
|
||||
provider_request_method: None,
|
||||
auth_header: Some(auth_header),
|
||||
auth_value: Some(auth_value),
|
||||
provider_api_format,
|
||||
client_api_format: spec_metadata.api_format.to_string(),
|
||||
model_name: input.requested_model.clone(),
|
||||
mapped_model,
|
||||
prompt_cache_key: None,
|
||||
provider_request_headers,
|
||||
provider_request_body: Some(provider_request_body),
|
||||
provider_request_body_base64: None,
|
||||
content_type: Some("application/json".to_string()),
|
||||
proxy,
|
||||
transport_profile,
|
||||
timeouts,
|
||||
upstream_is_stream,
|
||||
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
||||
report_context: Some(report_context),
|
||||
auth_context: input.auth_context.clone(),
|
||||
},
|
||||
))
|
||||
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
|
||||
decision_is_stream: spec_metadata.require_streaming,
|
||||
decision_kind: spec_metadata.decision_kind.to_string(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: candidate_id.to_string(),
|
||||
provider_name: candidate.provider_name.clone(),
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
key_id: candidate.key_id.clone(),
|
||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||
upstream_url,
|
||||
provider_request_method: None,
|
||||
auth_header: Some(auth_header),
|
||||
auth_value: Some(auth_value),
|
||||
provider_api_format,
|
||||
client_api_format: spec_metadata.api_format.to_string(),
|
||||
model_name: input.requested_model.clone(),
|
||||
mapped_model,
|
||||
prompt_cache_key: None,
|
||||
provider_request_headers,
|
||||
provider_request_body: Some(provider_request_body),
|
||||
provider_request_body_base64: None,
|
||||
content_type: Some("application/json".to_string()),
|
||||
proxy,
|
||||
transport_profile,
|
||||
timeouts,
|
||||
upstream_is_stream,
|
||||
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
||||
report_context: Some(report_context),
|
||||
auth_context: input.auth_context.clone(),
|
||||
});
|
||||
apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
|
||||
Ok(Some(decision))
|
||||
}
|
||||
|
||||
pub(super) async fn mark_skipped_local_standard_candidate(
|
||||
@@ -347,6 +356,9 @@ mod tests {
|
||||
required_capabilities: None,
|
||||
request_auth_channel: None,
|
||||
client_session_affinity: None,
|
||||
routing_policy: None,
|
||||
routing_trace_seed: None,
|
||||
routing_context: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -512,6 +524,7 @@ mod tests {
|
||||
claude_stream_spec(),
|
||||
)
|
||||
.await
|
||||
.expect("same-format candidate should not fail routing mutation")
|
||||
.expect("same-format candidate should build a standard-family payload");
|
||||
|
||||
assert_eq!(payload.endpoint_id.as_deref(), Some("endpoint-claude"));
|
||||
@@ -547,6 +560,7 @@ mod tests {
|
||||
claude_stream_spec(),
|
||||
)
|
||||
.await
|
||||
.expect("cross-format candidate should not fail routing mutation")
|
||||
.expect("cross-format candidate should still build after the same-format candidate");
|
||||
|
||||
assert_eq!(
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use aether_contracts::ResolvedTransportProfile;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::ai_serving::planner::candidate_preparation::{
|
||||
@@ -21,10 +22,11 @@ use crate::ai_serving::transport::kiro::{
|
||||
KIRO_ENVELOPE_NAME,
|
||||
};
|
||||
use crate::ai_serving::transport::{
|
||||
build_kiro_cross_format_upstream_url, build_openai_image_headers,
|
||||
build_openai_image_upstream_url, build_standard_provider_request_headers,
|
||||
openai_image_transport_unsupported_reason, resolve_openai_image_auth,
|
||||
ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput,
|
||||
build_grok_browser_headers, build_grok_upstream_url, build_kiro_cross_format_upstream_url,
|
||||
build_openai_image_headers, build_openai_image_upstream_url,
|
||||
build_standard_provider_request_headers, openai_image_transport_unsupported_reason,
|
||||
resolve_grok_session_auth, resolve_openai_image_auth, GrokHeaderInput,
|
||||
ProviderOpenAiImageHeadersInput, StandardProviderRequestHeadersInput, GROK_CHAT_PATH,
|
||||
};
|
||||
use crate::ai_serving::{
|
||||
build_openai_image_request_body_from_gemini_image_request, gemini_request_is_image_generation,
|
||||
@@ -49,6 +51,14 @@ pub(crate) struct LocalStandardCandidatePayloadParts {
|
||||
pub(super) upstream_is_stream: bool,
|
||||
pub(super) envelope_name: Option<&'static str>,
|
||||
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
|
||||
pub(super) transport_profile: Option<ResolvedTransportProfile>,
|
||||
}
|
||||
|
||||
fn is_grok_text_provider_api_format(provider_api_format: &str) -> bool {
|
||||
matches!(
|
||||
crate::ai_serving::normalize_api_format_alias(provider_api_format).as_str(),
|
||||
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages"
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
@@ -64,7 +74,14 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
let planner_state = crate::ai_serving::PlannerAppState::new(state);
|
||||
let candidate = &attempt.eligible.candidate;
|
||||
let transport = &attempt.eligible.transport;
|
||||
let transport_profile = crate::ai_serving::transport::resolve_transport_profile(transport);
|
||||
let provider_api_format = attempt.eligible.provider_api_format.as_str();
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let is_grok = transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("grok");
|
||||
if spec_metadata.api_format == "gemini:generate_content"
|
||||
&& provider_api_format == "openai:image"
|
||||
&& gemini_request_is_image_generation(body_json)
|
||||
@@ -75,11 +92,116 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
.await;
|
||||
}
|
||||
let is_kiro_claude_cli = is_kiro_claude_messages_transport(transport, provider_api_format);
|
||||
if !crate::ai_serving::request_pair_allowed_for_transport(
|
||||
if is_grok && is_grok_text_provider_api_format(provider_api_format) {
|
||||
let prepared_candidate = match prepare_header_authenticated_candidate(
|
||||
planner_state,
|
||||
transport,
|
||||
candidate,
|
||||
resolve_grok_session_auth(transport),
|
||||
OauthPreparationContext {
|
||||
trace_id,
|
||||
api_format: provider_api_format,
|
||||
operation: "standard_family_grok_text_request",
|
||||
},
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(prepared) => prepared,
|
||||
Err(skip_reason) => {
|
||||
mark_skipped_local_standard_candidate(
|
||||
state,
|
||||
input,
|
||||
trace_id,
|
||||
candidate,
|
||||
attempt.candidate_index,
|
||||
&attempt.candidate_id,
|
||||
skip_reason,
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
}
|
||||
};
|
||||
|
||||
let mut provider_request_body = body_json.clone();
|
||||
if let Some(object) = provider_request_body.as_object_mut() {
|
||||
object.insert(
|
||||
"model".to_string(),
|
||||
serde_json::Value::String(prepared_candidate.mapped_model.clone()),
|
||||
);
|
||||
}
|
||||
|
||||
let upstream_is_stream = resolve_upstream_is_stream_for_provider(
|
||||
transport.endpoint.config.as_ref(),
|
||||
transport.provider.provider_type.as_str(),
|
||||
provider_api_format,
|
||||
spec_metadata.require_streaming,
|
||||
false,
|
||||
);
|
||||
let force_body_stream_field =
|
||||
endpoint_config_forces_body_stream_field(transport.endpoint.config.as_ref());
|
||||
enforce_provider_body_stream_policy(
|
||||
&mut provider_request_body,
|
||||
provider_api_format,
|
||||
upstream_is_stream,
|
||||
request_requires_body_stream_field(body_json, force_body_stream_field),
|
||||
);
|
||||
|
||||
let upstream_url = build_grok_upstream_url(transport, GROK_CHAT_PATH);
|
||||
let Some(provider_request_headers) = build_grok_browser_headers(GrokHeaderInput {
|
||||
transport,
|
||||
transport_profile: transport_profile.as_ref(),
|
||||
request_headers: Some(effective_headers),
|
||||
content_type: "application/json",
|
||||
accept: "text/event-stream",
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: body_json,
|
||||
}) else {
|
||||
mark_skipped_local_standard_candidate_with_failure_diagnostic(
|
||||
state,
|
||||
input,
|
||||
trace_id,
|
||||
candidate,
|
||||
attempt.candidate_index,
|
||||
&attempt.candidate_id,
|
||||
"transport_header_rules_apply_failed",
|
||||
CandidateFailureDiagnostic::header_rules_apply_failed(
|
||||
spec_metadata.api_format,
|
||||
provider_api_format,
|
||||
"grok_standard_family_headers",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
};
|
||||
|
||||
return Some(LocalStandardCandidatePayloadParts {
|
||||
auth_header: prepared_candidate.auth_header,
|
||||
auth_value: prepared_candidate.auth_value,
|
||||
mapped_model: prepared_candidate.mapped_model,
|
||||
provider_api_format: provider_api_format.to_string(),
|
||||
provider_request_body,
|
||||
provider_request_headers,
|
||||
upstream_url,
|
||||
upstream_is_stream,
|
||||
envelope_name: None,
|
||||
transport: Arc::clone(transport),
|
||||
transport_profile,
|
||||
});
|
||||
}
|
||||
|
||||
let Some(conversion_kind) =
|
||||
crate::ai_serving::request_conversion_kind(spec_metadata.api_format, provider_api_format)
|
||||
else {
|
||||
return None;
|
||||
};
|
||||
|
||||
if crate::ai_serving::request_conversion_transport_unsupported_reason(
|
||||
transport,
|
||||
spec_metadata.api_format,
|
||||
provider_api_format,
|
||||
) {
|
||||
conversion_kind,
|
||||
)
|
||||
.is_some()
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
@@ -212,7 +334,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
transport.endpoint.body_rules.as_ref()
|
||||
},
|
||||
Some(input.auth_context.api_key_id.as_str()),
|
||||
Some(&parts.headers),
|
||||
Some(effective_headers),
|
||||
enable_model_directives,
|
||||
) {
|
||||
Some(body) => body,
|
||||
@@ -362,7 +484,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
transport,
|
||||
provider_api_format,
|
||||
same_format: false,
|
||||
headers: &parts.headers,
|
||||
headers: effective_headers,
|
||||
auth_header: &prepared_candidate.auth_header,
|
||||
auth_value: &prepared_candidate.auth_value,
|
||||
extra_headers: &BTreeMap::new(),
|
||||
@@ -393,7 +515,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
apply_codex_openai_responses_special_headers(
|
||||
&mut provider_request_headers,
|
||||
&provider_request_body,
|
||||
&parts.headers,
|
||||
effective_headers,
|
||||
transport.provider.provider_type.as_str(),
|
||||
provider_api_format,
|
||||
Some(trace_id),
|
||||
@@ -411,6 +533,7 @@ pub(crate) async fn resolve_local_standard_candidate_payload_parts(
|
||||
upstream_is_stream,
|
||||
envelope_name: None,
|
||||
transport: Arc::clone(transport),
|
||||
transport_profile: None,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -510,9 +633,10 @@ async fn resolve_local_gemini_image_to_openai_image_candidate_payload_parts(
|
||||
|
||||
let upstream_is_stream = true;
|
||||
let upstream_url = build_openai_image_upstream_url(transport, None);
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let Some(mut provider_request_headers) =
|
||||
build_openai_image_headers(ProviderOpenAiImageHeadersInput {
|
||||
headers: &parts.headers,
|
||||
headers: effective_headers,
|
||||
auth_header: &prepared_candidate.auth_header,
|
||||
auth_value: &prepared_candidate.auth_value,
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
@@ -540,7 +664,7 @@ async fn resolve_local_gemini_image_to_openai_image_candidate_payload_parts(
|
||||
apply_codex_openai_responses_special_headers(
|
||||
&mut provider_request_headers,
|
||||
&converted.body_json,
|
||||
&parts.headers,
|
||||
effective_headers,
|
||||
transport.provider.provider_type.as_str(),
|
||||
provider_api_format,
|
||||
Some(trace_id),
|
||||
@@ -558,6 +682,7 @@ async fn resolve_local_gemini_image_to_openai_image_candidate_payload_parts(
|
||||
upstream_is_stream,
|
||||
envelope_name: None,
|
||||
transport: Arc::clone(transport),
|
||||
transport_profile: None,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -579,12 +704,13 @@ async fn build_kiro_cross_format_payload_parts(
|
||||
kiro_auth: &KiroRequestAuth,
|
||||
) -> Option<LocalStandardCandidatePayloadParts> {
|
||||
let candidate = &attempt.eligible.candidate;
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let provider_request_body = match build_kiro_provider_request_body(
|
||||
&claude_request_body,
|
||||
&mapped_model,
|
||||
&kiro_auth.auth_config,
|
||||
transport.endpoint.body_rules.as_ref(),
|
||||
Some(&parts.headers),
|
||||
Some(effective_headers),
|
||||
) {
|
||||
Some(body) => body,
|
||||
None => {
|
||||
@@ -635,7 +761,7 @@ async fn build_kiro_cross_format_payload_parts(
|
||||
}
|
||||
};
|
||||
let provider_request_headers = match build_kiro_provider_headers(KiroProviderHeadersInput {
|
||||
headers: &parts.headers,
|
||||
headers: effective_headers,
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: original_body_json,
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
@@ -676,5 +802,6 @@ async fn build_kiro_cross_format_payload_parts(
|
||||
upstream_is_stream,
|
||||
envelope_name: Some(KIRO_ENVELOPE_NAME),
|
||||
transport: Arc::clone(transport),
|
||||
transport_profile: None,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
use crate::ai_serving::build_request_trace_proxy_value;
|
||||
use crate::ai_serving::planner::common::OPENAI_CHAT_STREAM_PLAN_KIND;
|
||||
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
|
||||
use crate::ai_serving::planner::report_context::{
|
||||
build_local_execution_report_context, insert_provider_stream_event_api_format,
|
||||
LocalExecutionReportContextParts,
|
||||
@@ -67,7 +68,10 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
|
||||
let proxy = state
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport)
|
||||
.await;
|
||||
let transport_profile = resolve_transport_profile(&resolved.transport);
|
||||
let transport_profile = resolved
|
||||
.transport_profile
|
||||
.clone()
|
||||
.or_else(|| resolve_transport_profile(&resolved.transport));
|
||||
let timeouts = resolve_transport_execution_timeouts(&resolved.transport);
|
||||
let mut extra_fields = serde_json::Map::new();
|
||||
if let Some(proxy_value) =
|
||||
@@ -99,12 +103,14 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
|
||||
envelope_name,
|
||||
transport,
|
||||
request_redacted,
|
||||
transport_profile: _,
|
||||
} = resolved;
|
||||
let original_request_body_json = if request_redacted {
|
||||
Some(&provider_request_body)
|
||||
} else {
|
||||
Some(body_json)
|
||||
};
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let report_context = append_local_failover_policy_to_value(
|
||||
append_execution_contract_fields_to_value(
|
||||
build_local_execution_report_context(LocalExecutionReportContextParts {
|
||||
@@ -132,7 +138,7 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
|
||||
body_rules: transport.endpoint.body_rules.as_ref(),
|
||||
provider_request_method: Some(serde_json::Value::Null),
|
||||
provider_request_headers: Some(&provider_request_headers),
|
||||
original_headers: &parts.headers,
|
||||
original_headers: effective_headers,
|
||||
request_path: Some(parts.uri.path()),
|
||||
request_query_string: parts.uri.query(),
|
||||
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
|
||||
@@ -160,39 +166,39 @@ pub(crate) async fn maybe_build_local_openai_chat_decision_payload_for_candidate
|
||||
&transport,
|
||||
);
|
||||
|
||||
Ok(Some(build_ai_execution_decision_response(
|
||||
AiExecutionDecisionResponseParts {
|
||||
decision_is_stream,
|
||||
decision_kind: decision_kind.to_string(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: candidate_id.clone(),
|
||||
provider_name: transport.provider.name.clone(),
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
key_id: candidate.key_id.clone(),
|
||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||
upstream_url,
|
||||
provider_request_method: None,
|
||||
auth_header: Some(auth_header),
|
||||
auth_value: Some(auth_value),
|
||||
provider_api_format,
|
||||
client_api_format: "openai:chat".to_string(),
|
||||
model_name: input.requested_model.clone(),
|
||||
mapped_model,
|
||||
prompt_cache_key,
|
||||
provider_request_headers,
|
||||
provider_request_body: Some(provider_request_body),
|
||||
provider_request_body_base64: None,
|
||||
content_type: Some("application/json".to_string()),
|
||||
proxy,
|
||||
transport_profile,
|
||||
timeouts,
|
||||
upstream_is_stream,
|
||||
report_kind: Some(report_kind),
|
||||
report_context: Some(report_context),
|
||||
auth_context: input.auth_context.clone(),
|
||||
},
|
||||
)))
|
||||
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
|
||||
decision_is_stream,
|
||||
decision_kind: decision_kind.to_string(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: candidate_id.clone(),
|
||||
provider_name: transport.provider.name.clone(),
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
key_id: candidate.key_id.clone(),
|
||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||
upstream_url,
|
||||
provider_request_method: None,
|
||||
auth_header: Some(auth_header),
|
||||
auth_value: Some(auth_value),
|
||||
provider_api_format,
|
||||
client_api_format: "openai:chat".to_string(),
|
||||
model_name: input.requested_model.clone(),
|
||||
mapped_model,
|
||||
prompt_cache_key,
|
||||
provider_request_headers,
|
||||
provider_request_body: Some(provider_request_body),
|
||||
provider_request_body_base64: None,
|
||||
content_type: Some("application/json".to_string()),
|
||||
proxy,
|
||||
transport_profile,
|
||||
timeouts,
|
||||
upstream_is_stream,
|
||||
report_kind: Some(report_kind),
|
||||
report_context: Some(report_context),
|
||||
auth_context: input.auth_context.clone(),
|
||||
});
|
||||
apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
|
||||
Ok(Some(decision))
|
||||
}
|
||||
|
||||
@@ -3,6 +3,7 @@ use std::collections::BTreeMap;
|
||||
use std::sync::Arc;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use aether_contracts::ResolvedTransportProfile;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::ai_serving::planner::candidate_preparation::{
|
||||
@@ -27,8 +28,9 @@ use crate::ai_serving::transport::kiro::{
|
||||
};
|
||||
use crate::ai_serving::transport::local_openai_chat_transport_unsupported_reason;
|
||||
use crate::ai_serving::transport::{
|
||||
build_kiro_cross_format_upstream_url, build_standard_provider_request_headers,
|
||||
StandardProviderRequestHeadersInput,
|
||||
build_grok_browser_headers, build_grok_upstream_url, build_kiro_cross_format_upstream_url,
|
||||
build_standard_provider_request_headers, GrokHeaderInput, StandardProviderRequestHeadersInput,
|
||||
GROK_CHAT_PATH,
|
||||
};
|
||||
use crate::ai_serving::{
|
||||
ai_local_execution_contract_for_formats, request_conversion_direct_auth,
|
||||
@@ -64,6 +66,14 @@ pub(crate) struct LocalOpenAiChatCandidatePayloadParts {
|
||||
pub(super) envelope_name: Option<&'static str>,
|
||||
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
|
||||
pub(super) request_redacted: bool,
|
||||
pub(super) transport_profile: Option<ResolvedTransportProfile>,
|
||||
}
|
||||
|
||||
fn is_grok_text_provider_api_format(provider_api_format: &str) -> bool {
|
||||
matches!(
|
||||
crate::ai_serving::normalize_api_format_alias(provider_api_format).as_str(),
|
||||
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages"
|
||||
)
|
||||
}
|
||||
|
||||
fn request_identity_response_encoding_when_redacted(
|
||||
@@ -177,6 +187,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
||||
let candidate = &eligible.candidate;
|
||||
let provider_api_format = eligible.provider_api_format.as_str();
|
||||
let transport = &eligible.transport;
|
||||
let transport_profile = crate::ai_serving::transport::resolve_transport_profile(transport);
|
||||
let force_body_stream_field =
|
||||
endpoint_config_forces_body_stream_field(transport.endpoint.config.as_ref());
|
||||
let enable_model_directives =
|
||||
@@ -190,6 +201,130 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
||||
resolve_provider_chat_request_redaction(state, parts, body_json, input, candidate_id)
|
||||
.await?;
|
||||
let body_json = redaction.body_json.as_ref();
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let is_grok = transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("grok");
|
||||
|
||||
if is_grok && is_grok_text_provider_api_format(provider_api_format) {
|
||||
let prepared_candidate = match prepare_header_authenticated_candidate(
|
||||
planner_state,
|
||||
transport,
|
||||
candidate,
|
||||
crate::ai_serving::transport::resolve_grok_session_auth(transport),
|
||||
OauthPreparationContext {
|
||||
trace_id,
|
||||
api_format: provider_api_format,
|
||||
operation: "openai_chat_same_format",
|
||||
},
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(prepared) => prepared,
|
||||
Err(skip_reason) => {
|
||||
mark_skipped_local_openai_chat_candidate(
|
||||
state,
|
||||
input,
|
||||
trace_id,
|
||||
candidate,
|
||||
candidate_index,
|
||||
candidate_id,
|
||||
skip_reason,
|
||||
)
|
||||
.await;
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
|
||||
let Some(provider_request_body) = build_local_openai_chat_request_body(
|
||||
body_json,
|
||||
&prepared_candidate.mapped_model,
|
||||
upstream_is_stream,
|
||||
force_body_stream_field,
|
||||
transport.endpoint.body_rules.as_ref(),
|
||||
effective_headers,
|
||||
enable_model_directives,
|
||||
) else {
|
||||
mark_skipped_local_openai_chat_candidate_with_extra_data(
|
||||
state,
|
||||
input,
|
||||
trace_id,
|
||||
candidate,
|
||||
candidate_index,
|
||||
candidate_id,
|
||||
"provider_request_body_build_failed",
|
||||
request_body_build_failure_extra_data(
|
||||
body_json,
|
||||
"openai:chat",
|
||||
provider_api_format,
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let upstream_url = build_grok_upstream_url(transport, GROK_CHAT_PATH);
|
||||
let Some(mut provider_request_headers) = build_grok_browser_headers(GrokHeaderInput {
|
||||
transport,
|
||||
transport_profile: transport_profile.as_ref(),
|
||||
request_headers: Some(effective_headers),
|
||||
content_type: "application/json",
|
||||
accept: "text/event-stream",
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: body_json,
|
||||
}) else {
|
||||
mark_skipped_local_openai_chat_candidate_with_failure_diagnostic(
|
||||
state,
|
||||
input,
|
||||
trace_id,
|
||||
candidate,
|
||||
candidate_index,
|
||||
candidate_id,
|
||||
"transport_header_rules_apply_failed",
|
||||
CandidateFailureDiagnostic::header_rules_apply_failed(
|
||||
"openai:chat",
|
||||
provider_api_format,
|
||||
"grok_openai_chat_headers",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let (execution_strategy, conversion_mode) =
|
||||
ai_local_execution_contract_for_formats("openai:chat", provider_api_format);
|
||||
let resolved_report_kind =
|
||||
if decision_kind == OPENAI_CHAT_STREAM_PLAN_KIND || !upstream_is_stream {
|
||||
report_kind.to_string()
|
||||
} else {
|
||||
"openai_chat_sync_finalize".to_string()
|
||||
};
|
||||
|
||||
request_identity_response_encoding_when_redacted(
|
||||
&mut provider_request_headers,
|
||||
redaction.redacted,
|
||||
);
|
||||
|
||||
return Ok(Some(LocalOpenAiChatCandidatePayloadParts {
|
||||
auth_header: prepared_candidate.auth_header,
|
||||
auth_value: prepared_candidate.auth_value,
|
||||
mapped_model: prepared_candidate.mapped_model,
|
||||
provider_api_format: provider_api_format.to_string(),
|
||||
provider_request_body,
|
||||
provider_request_headers,
|
||||
upstream_url,
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
report_kind: resolved_report_kind,
|
||||
envelope_name: None,
|
||||
transport: Arc::clone(transport),
|
||||
request_redacted: redaction.redacted,
|
||||
transport_profile,
|
||||
}));
|
||||
}
|
||||
|
||||
if provider_api_format == "openai:chat" {
|
||||
if let Some(skip_reason) = local_openai_chat_transport_unsupported_reason(transport) {
|
||||
@@ -241,7 +376,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
||||
upstream_is_stream,
|
||||
force_body_stream_field,
|
||||
transport.endpoint.body_rules.as_ref(),
|
||||
&parts.headers,
|
||||
effective_headers,
|
||||
enable_model_directives,
|
||||
) else {
|
||||
mark_skipped_local_openai_chat_candidate_with_extra_data(
|
||||
@@ -286,7 +421,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
||||
transport,
|
||||
provider_api_format,
|
||||
same_format: true,
|
||||
headers: &parts.headers,
|
||||
headers: effective_headers,
|
||||
auth_header: &prepared_candidate.auth_header,
|
||||
auth_value: &prepared_candidate.auth_value,
|
||||
extra_headers: &BTreeMap::new(),
|
||||
@@ -317,7 +452,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
||||
apply_codex_openai_responses_special_headers(
|
||||
&mut provider_request_headers,
|
||||
&provider_request_body,
|
||||
&parts.headers,
|
||||
effective_headers,
|
||||
transport.provider.provider_type.as_str(),
|
||||
transport.endpoint.api_format.as_str(),
|
||||
Some(trace_id),
|
||||
@@ -351,6 +486,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
||||
envelope_name: None,
|
||||
transport: Arc::clone(transport),
|
||||
request_redacted: redaction.redacted,
|
||||
transport_profile,
|
||||
}));
|
||||
};
|
||||
|
||||
@@ -480,7 +616,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
||||
transport.endpoint.body_rules.as_ref()
|
||||
},
|
||||
Some(input.auth_context.api_key_id.as_str()),
|
||||
&parts.headers,
|
||||
effective_headers,
|
||||
enable_model_directives,
|
||||
) else {
|
||||
mark_skipped_local_openai_chat_candidate_with_extra_data(
|
||||
@@ -575,7 +711,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
||||
transport,
|
||||
provider_api_format: provider_api_format.as_str(),
|
||||
same_format: false,
|
||||
headers: &parts.headers,
|
||||
headers: effective_headers,
|
||||
auth_header: &prepared_candidate.auth_header,
|
||||
auth_value: &prepared_candidate.auth_value,
|
||||
extra_headers: &BTreeMap::new(),
|
||||
@@ -606,7 +742,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
||||
apply_codex_openai_responses_special_headers(
|
||||
&mut provider_request_headers,
|
||||
&provider_request_body,
|
||||
&parts.headers,
|
||||
effective_headers,
|
||||
transport.provider.provider_type.as_str(),
|
||||
provider_api_format.as_str(),
|
||||
Some(trace_id),
|
||||
@@ -639,6 +775,7 @@ pub(crate) async fn resolve_local_openai_chat_candidate_payload_parts(
|
||||
envelope_name: None,
|
||||
transport: Arc::clone(transport),
|
||||
request_redacted: redaction.redacted,
|
||||
transport_profile: None,
|
||||
}))
|
||||
}
|
||||
|
||||
@@ -664,12 +801,13 @@ async fn build_kiro_openai_chat_cross_format_payload_parts(
|
||||
request_redacted: bool,
|
||||
) -> Option<LocalOpenAiChatCandidatePayloadParts> {
|
||||
let candidate = &eligible.candidate;
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let provider_request_body = match build_kiro_provider_request_body(
|
||||
&claude_request_body,
|
||||
&mapped_model,
|
||||
&kiro_auth.auth_config,
|
||||
transport.endpoint.body_rules.as_ref(),
|
||||
Some(&parts.headers),
|
||||
Some(effective_headers),
|
||||
) {
|
||||
Some(body) => body,
|
||||
None => {
|
||||
@@ -720,7 +858,7 @@ async fn build_kiro_openai_chat_cross_format_payload_parts(
|
||||
}
|
||||
};
|
||||
let mut provider_request_headers = match build_kiro_provider_headers(KiroProviderHeadersInput {
|
||||
headers: &parts.headers,
|
||||
headers: effective_headers,
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: original_body_json,
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
@@ -775,6 +913,7 @@ async fn build_kiro_openai_chat_cross_format_payload_parts(
|
||||
envelope_name: Some(KIRO_ENVELOPE_NAME),
|
||||
transport: Arc::clone(transport),
|
||||
request_redacted,
|
||||
transport_profile: None,
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -140,6 +140,7 @@ pub(crate) async fn materialize_local_openai_chat_candidate_attempts(
|
||||
Some(&input.auth_snapshot),
|
||||
input.client_session_affinity.as_ref(),
|
||||
input.required_capabilities.as_ref(),
|
||||
input.routing_policy.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
@@ -220,6 +221,7 @@ pub(crate) async fn build_local_openai_chat_candidate_attempt_source<'a>(
|
||||
Some(&input.auth_snapshot),
|
||||
input.client_session_affinity.as_ref(),
|
||||
input.required_capabilities.as_ref(),
|
||||
input.routing_policy.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
@@ -298,6 +300,7 @@ pub(crate) async fn build_lazy_local_openai_chat_candidate_attempt_source<'a>(
|
||||
&input.auth_snapshot,
|
||||
input.client_session_affinity.as_ref(),
|
||||
input.required_capabilities.as_ref(),
|
||||
input.routing_policy.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
|
||||
@@ -136,10 +136,11 @@ pub(crate) async fn maybe_build_sync_local_decision_payload(
|
||||
let Some(input) = resolve_local_openai_chat_decision_input(
|
||||
state, parts, trace_id, decision, body_json, plan_kind, false,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
|
||||
let (mut source, _) = build_lazy_local_openai_chat_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, false,
|
||||
@@ -187,10 +188,11 @@ pub(crate) async fn maybe_build_stream_local_decision_payload(
|
||||
let Some(input) = resolve_local_openai_chat_decision_input(
|
||||
state, parts, trace_id, decision, body_json, plan_kind, false,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
|
||||
let (mut source, _) = build_lazy_local_openai_chat_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, true,
|
||||
|
||||
@@ -26,6 +26,7 @@ pub(crate) async fn list_local_openai_chat_candidates(
|
||||
require_streaming,
|
||||
input.required_capabilities.as_ref(),
|
||||
&input.auth_snapshot,
|
||||
input.routing_policy.as_ref(),
|
||||
input.client_session_affinity.as_ref(),
|
||||
false,
|
||||
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModel,
|
||||
|
||||
@@ -4,11 +4,12 @@ use super::super::{GatewayControlDecision, LocalOpenAiChatDecisionInput};
|
||||
use super::diagnostic::set_local_openai_chat_miss_diagnostic;
|
||||
use crate::ai_serving::planner::common::extract_standard_requested_model;
|
||||
use crate::ai_serving::planner::decision_input::{
|
||||
attach_routing_policy_to_local_requested_model_input,
|
||||
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
|
||||
};
|
||||
use crate::ai_serving::resolve_local_decision_execution_runtime_auth_context;
|
||||
use crate::client_session_affinity::client_session_affinity_from_parts;
|
||||
use crate::AppState;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
pub(crate) async fn resolve_local_openai_chat_decision_input(
|
||||
state: &AppState,
|
||||
@@ -18,7 +19,7 @@ pub(crate) async fn resolve_local_openai_chat_decision_input(
|
||||
body_json: &serde_json::Value,
|
||||
plan_kind: &str,
|
||||
record_miss_diagnostic: bool,
|
||||
) -> Option<LocalOpenAiChatDecisionInput> {
|
||||
) -> Result<Option<LocalOpenAiChatDecisionInput>, GatewayError> {
|
||||
let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else {
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
@@ -37,7 +38,7 @@ pub(crate) async fn resolve_local_openai_chat_decision_input(
|
||||
"missing_auth_context",
|
||||
);
|
||||
}
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let Some(requested_model) = extract_standard_requested_model(body_json) else {
|
||||
@@ -55,7 +56,7 @@ pub(crate) async fn resolve_local_openai_chat_decision_input(
|
||||
"missing_requested_model",
|
||||
);
|
||||
}
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let resolved_input = match resolve_local_authenticated_decision_input(
|
||||
@@ -84,7 +85,7 @@ pub(crate) async fn resolve_local_openai_chat_decision_input(
|
||||
"auth_snapshot_missing",
|
||||
);
|
||||
}
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
Err(err) => {
|
||||
warn!(
|
||||
@@ -102,12 +103,28 @@ pub(crate) async fn resolve_local_openai_chat_decision_input(
|
||||
"auth_snapshot_read_failed",
|
||||
);
|
||||
}
|
||||
return None;
|
||||
return Err(err);
|
||||
}
|
||||
};
|
||||
|
||||
let mut input = build_local_requested_model_decision_input(resolved_input, requested_model);
|
||||
input.request_auth_channel = decision.request_auth_channel.clone();
|
||||
input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json));
|
||||
Some(input)
|
||||
if let Err(err) = attach_routing_policy_to_local_requested_model_input(
|
||||
state,
|
||||
parts,
|
||||
&mut input,
|
||||
body_json,
|
||||
"openai:chat",
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
error = ?err,
|
||||
"gateway local openai chat decision routing profile resolution failed"
|
||||
);
|
||||
return Err(err);
|
||||
}
|
||||
Ok(Some(input))
|
||||
}
|
||||
|
||||
@@ -23,7 +23,7 @@ pub(crate) struct LocalOpenAiChatStreamAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_json: serde_json::Value,
|
||||
input: LocalOpenAiChatDecisionInput,
|
||||
candidates: LocalOpenAiChatCandidateAttemptSource<'a>,
|
||||
}
|
||||
@@ -43,13 +43,18 @@ pub(crate) async fn build_local_openai_chat_stream_attempt_source<'a>(
|
||||
let Some(input) = resolve_local_openai_chat_decision_input(
|
||||
state, parts, trace_id, decision, body_json, plan_kind, true,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let effective_body_json = input.effective_body_json(body_json).clone();
|
||||
|
||||
let (candidates, candidate_count) = build_lazy_local_openai_chat_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, true,
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
&effective_body_json,
|
||||
true,
|
||||
)
|
||||
.await;
|
||||
if candidate_count == 0 {
|
||||
@@ -77,7 +82,7 @@ pub(crate) async fn build_local_openai_chat_stream_attempt_source<'a>(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
body_json: effective_body_json,
|
||||
input,
|
||||
candidates,
|
||||
},
|
||||
@@ -127,7 +132,7 @@ impl LocalOpenAiChatStreamAttemptSource<'_> {
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
OPENAI_CHAT_STREAM_PLAN_KIND,
|
||||
@@ -139,7 +144,7 @@ impl LocalOpenAiChatStreamAttemptSource<'_> {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match build_openai_chat_stream_plan_from_decision(self.parts, self.body_json, payload) {
|
||||
match build_openai_chat_stream_plan_from_decision(self.parts, &self.body_json, payload) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
@@ -168,7 +173,7 @@ pub(crate) async fn build_local_openai_chat_stream_plan_and_reports(
|
||||
let Some(input) = resolve_local_openai_chat_decision_input(
|
||||
state, parts, trace_id, decision, body_json, plan_kind, true,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(Vec::new());
|
||||
};
|
||||
|
||||
@@ -23,7 +23,7 @@ pub(crate) struct LocalOpenAiChatSyncAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_json: serde_json::Value,
|
||||
input: LocalOpenAiChatDecisionInput,
|
||||
candidates: LocalOpenAiChatCandidateAttemptSource<'a>,
|
||||
}
|
||||
@@ -43,13 +43,18 @@ pub(crate) async fn build_local_openai_chat_sync_attempt_source<'a>(
|
||||
let Some(input) = resolve_local_openai_chat_decision_input(
|
||||
state, parts, trace_id, decision, body_json, plan_kind, true,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let effective_body_json = input.effective_body_json(body_json).clone();
|
||||
|
||||
let (candidates, candidate_count) = build_lazy_local_openai_chat_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, false,
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
&effective_body_json,
|
||||
false,
|
||||
)
|
||||
.await;
|
||||
if candidate_count == 0 {
|
||||
@@ -77,7 +82,7 @@ pub(crate) async fn build_local_openai_chat_sync_attempt_source<'a>(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
body_json: effective_body_json,
|
||||
input,
|
||||
candidates,
|
||||
},
|
||||
@@ -127,7 +132,7 @@ impl LocalOpenAiChatSyncAttemptSource<'_> {
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
OPENAI_CHAT_SYNC_PLAN_KIND,
|
||||
@@ -139,7 +144,7 @@ impl LocalOpenAiChatSyncAttemptSource<'_> {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match build_openai_chat_sync_plan_from_decision(self.parts, self.body_json, payload) {
|
||||
match build_openai_chat_sync_plan_from_decision(self.parts, &self.body_json, payload) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
@@ -168,7 +173,7 @@ pub(crate) async fn build_local_openai_chat_sync_plan_and_reports(
|
||||
let Some(input) = resolve_local_openai_chat_decision_input(
|
||||
state, parts, trace_id, decision, body_json, plan_kind, true,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(Vec::new());
|
||||
};
|
||||
|
||||
@@ -2,6 +2,7 @@ use serde_json::json;
|
||||
use tracing::debug;
|
||||
|
||||
use crate::ai_serving::build_request_trace_proxy_value;
|
||||
use crate::ai_serving::planner::decision_input::apply_provider_request_routing_policy_to_decision;
|
||||
use crate::ai_serving::planner::report_context::{
|
||||
build_local_execution_report_context, insert_provider_stream_event_api_format,
|
||||
LocalExecutionReportContextParts,
|
||||
@@ -15,7 +16,7 @@ use crate::ai_serving::transport::{
|
||||
};
|
||||
use crate::{
|
||||
append_execution_contract_fields_to_value, append_local_failover_policy_to_value,
|
||||
AiExecutionDecision, AppState,
|
||||
AiExecutionDecision, AppState, GatewayError,
|
||||
};
|
||||
|
||||
use super::request::resolve_local_openai_responses_candidate_payload_parts;
|
||||
@@ -30,7 +31,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
|
||||
input: &LocalOpenAiResponsesDecisionInput,
|
||||
attempt: LocalOpenAiResponsesCandidateAttempt,
|
||||
spec: LocalOpenAiResponsesSpec,
|
||||
) -> Option<AiExecutionDecision> {
|
||||
) -> Result<Option<AiExecutionDecision>, GatewayError> {
|
||||
let spec_metadata = local_openai_responses_spec_metadata(spec);
|
||||
let attempt_identity = attempt.attempt_identity();
|
||||
let LocalOpenAiResponsesCandidateAttempt {
|
||||
@@ -39,7 +40,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
|
||||
candidate_id,
|
||||
..
|
||||
} = attempt;
|
||||
let resolved = resolve_local_openai_responses_candidate_payload_parts(
|
||||
let Some(resolved) = resolve_local_openai_responses_candidate_payload_parts(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
@@ -50,7 +51,10 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
|
||||
&candidate_id,
|
||||
spec,
|
||||
)
|
||||
.await?;
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let candidate = &eligible.candidate;
|
||||
|
||||
let prompt_cache_key = resolved
|
||||
@@ -63,7 +67,10 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
|
||||
let proxy = state
|
||||
.resolve_transport_proxy_snapshot_with_tunnel_affinity(&resolved.transport)
|
||||
.await;
|
||||
let transport_profile = resolve_transport_profile(&resolved.transport);
|
||||
let transport_profile = resolved
|
||||
.transport_profile
|
||||
.clone()
|
||||
.or_else(|| resolve_transport_profile(&resolved.transport));
|
||||
let timeouts = resolve_transport_execution_timeouts(&resolved.transport);
|
||||
let mut extra_fields = serde_json::Map::new();
|
||||
if let Some(proxy_value) =
|
||||
@@ -78,6 +85,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
|
||||
&mut extra_fields,
|
||||
resolved.transport.provider.provider_type.as_str(),
|
||||
);
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let report_context = append_local_failover_policy_to_value(
|
||||
append_execution_contract_fields_to_value(
|
||||
build_local_execution_report_context(LocalExecutionReportContextParts {
|
||||
@@ -105,7 +113,7 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
|
||||
body_rules: resolved.transport.endpoint.body_rules.as_ref(),
|
||||
provider_request_method: Some(serde_json::Value::Null),
|
||||
provider_request_headers: Some(&resolved.provider_request_headers),
|
||||
original_headers: &parts.headers,
|
||||
original_headers: effective_headers,
|
||||
request_path: Some(parts.uri.path()),
|
||||
request_query_string: parts.uri.query(),
|
||||
request_origin: Some(crate::ai_serving::request_origin_from_parts(parts)),
|
||||
@@ -170,41 +178,42 @@ pub(crate) async fn maybe_build_local_openai_responses_decision_payload_for_cand
|
||||
envelope_name: _,
|
||||
upstream_is_stream,
|
||||
transport,
|
||||
transport_profile: _,
|
||||
} = resolved;
|
||||
|
||||
Some(build_ai_execution_decision_response(
|
||||
AiExecutionDecisionResponseParts {
|
||||
decision_is_stream: spec_metadata.require_streaming,
|
||||
decision_kind: spec_metadata.decision_kind.to_string(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: candidate_id.clone(),
|
||||
provider_name: transport.provider.name.clone(),
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
key_id: candidate.key_id.clone(),
|
||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||
upstream_url,
|
||||
provider_request_method: None,
|
||||
auth_header: Some(auth_header),
|
||||
auth_value: Some(auth_value),
|
||||
provider_api_format,
|
||||
client_api_format: spec_metadata.api_format.to_string(),
|
||||
model_name: input.requested_model.clone(),
|
||||
mapped_model,
|
||||
prompt_cache_key,
|
||||
provider_request_headers,
|
||||
provider_request_body: Some(provider_request_body),
|
||||
provider_request_body_base64: None,
|
||||
content_type: Some("application/json".to_string()),
|
||||
proxy,
|
||||
transport_profile,
|
||||
timeouts,
|
||||
upstream_is_stream,
|
||||
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
||||
report_context: Some(report_context),
|
||||
auth_context: input.auth_context.clone(),
|
||||
},
|
||||
))
|
||||
let mut decision = build_ai_execution_decision_response(AiExecutionDecisionResponseParts {
|
||||
decision_is_stream: spec_metadata.require_streaming,
|
||||
decision_kind: spec_metadata.decision_kind.to_string(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
request_id: trace_id.to_string(),
|
||||
candidate_id: candidate_id.clone(),
|
||||
provider_name: transport.provider.name.clone(),
|
||||
provider_id: candidate.provider_id.clone(),
|
||||
endpoint_id: candidate.endpoint_id.clone(),
|
||||
key_id: candidate.key_id.clone(),
|
||||
upstream_base_url: transport.endpoint.base_url.clone(),
|
||||
upstream_url,
|
||||
provider_request_method: None,
|
||||
auth_header: Some(auth_header),
|
||||
auth_value: Some(auth_value),
|
||||
provider_api_format,
|
||||
client_api_format: spec_metadata.api_format.to_string(),
|
||||
model_name: input.requested_model.clone(),
|
||||
mapped_model,
|
||||
prompt_cache_key,
|
||||
provider_request_headers,
|
||||
provider_request_body: Some(provider_request_body),
|
||||
provider_request_body_base64: None,
|
||||
content_type: Some("application/json".to_string()),
|
||||
proxy,
|
||||
transport_profile,
|
||||
timeouts,
|
||||
upstream_is_stream,
|
||||
report_kind: spec_metadata.report_kind.map(ToOwned::to_owned),
|
||||
report_context: Some(report_context),
|
||||
auth_context: input.auth_context.clone(),
|
||||
});
|
||||
apply_provider_request_routing_policy_to_decision(input, &mut decision)?;
|
||||
Ok(Some(decision))
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
use aether_contracts::ResolvedTransportProfile;
|
||||
use serde_json::Value;
|
||||
use tracing::debug;
|
||||
|
||||
@@ -35,8 +36,10 @@ use crate::ai_serving::transport::kiro::{
|
||||
KiroRequestAuth, KIRO_ENVELOPE_NAME,
|
||||
};
|
||||
use crate::ai_serving::transport::{
|
||||
build_kiro_cross_format_upstream_url, build_standard_provider_request_headers,
|
||||
local_standard_transport_unsupported_reason_with_network, StandardProviderRequestHeadersInput,
|
||||
build_grok_browser_headers, build_grok_upstream_url, build_kiro_cross_format_upstream_url,
|
||||
build_standard_provider_request_headers,
|
||||
local_standard_transport_unsupported_reason_with_network, GrokHeaderInput,
|
||||
StandardProviderRequestHeadersInput, GROK_CHAT_PATH,
|
||||
};
|
||||
use crate::ai_serving::{
|
||||
ai_local_execution_contract_for_formats, request_conversion_direct_auth,
|
||||
@@ -56,6 +59,13 @@ use super::LocalOpenAiResponsesSpec;
|
||||
|
||||
const ANTIGRAVITY_ENVELOPE_NAME: &str = "antigravity:v1internal";
|
||||
|
||||
fn is_grok_text_provider_api_format(provider_api_format: &str) -> bool {
|
||||
matches!(
|
||||
crate::ai_serving::normalize_api_format_alias(provider_api_format).as_str(),
|
||||
"openai:chat" | "openai:responses" | "openai:responses:compact" | "claude:messages"
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) struct LocalOpenAiResponsesCandidatePayloadParts {
|
||||
pub(super) auth_header: String,
|
||||
pub(super) auth_value: String,
|
||||
@@ -70,6 +80,7 @@ pub(crate) struct LocalOpenAiResponsesCandidatePayloadParts {
|
||||
pub(super) envelope_name: Option<&'static str>,
|
||||
pub(super) upstream_is_stream: bool,
|
||||
pub(super) transport: Arc<GatewayProviderTransportSnapshot>,
|
||||
pub(super) transport_profile: Option<ResolvedTransportProfile>,
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
@@ -90,12 +101,22 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
let candidate = &eligible.candidate;
|
||||
let provider_api_format = eligible.provider_api_format.as_str();
|
||||
let transport = &eligible.transport;
|
||||
let transport_profile = crate::ai_serving::transport::resolve_transport_profile(transport);
|
||||
let is_antigravity = is_antigravity_provider_transport(transport);
|
||||
let is_kiro_claude_cli = is_kiro_claude_messages_transport(transport, provider_api_format);
|
||||
let is_grok = transport
|
||||
.provider
|
||||
.provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("grok");
|
||||
|
||||
let same_format = api_format_alias_matches(provider_api_format, &client_api_format);
|
||||
let conversion_kind = request_conversion_kind(spec_metadata.api_format, provider_api_format);
|
||||
let transport_unsupported_reason = if same_format && is_kiro_claude_cli {
|
||||
let transport_unsupported_reason = if is_grok
|
||||
&& is_grok_text_provider_api_format(provider_api_format)
|
||||
{
|
||||
None
|
||||
} else if same_format && is_kiro_claude_cli {
|
||||
local_kiro_request_transport_unsupported_reason_with_network(transport)
|
||||
} else if same_format {
|
||||
local_standard_transport_unsupported_reason_with_network(transport, provider_api_format)
|
||||
@@ -154,7 +175,9 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
None
|
||||
};
|
||||
|
||||
let direct_auth = if kiro_auth.is_some() {
|
||||
let direct_auth = if is_grok && is_grok_text_provider_api_format(provider_api_format) {
|
||||
crate::ai_serving::transport::resolve_grok_session_auth(transport)
|
||||
} else if kiro_auth.is_some() {
|
||||
None
|
||||
} else if same_format {
|
||||
match crate::ai_serving::normalize_api_format_alias(provider_api_format).as_str() {
|
||||
@@ -236,42 +259,58 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
);
|
||||
let force_body_stream_field =
|
||||
endpoint_config_forces_body_stream_field(transport.endpoint.config.as_ref());
|
||||
let Some(mut base_provider_request_body) = (if needs_bidirectional_conversion {
|
||||
build_cross_format_openai_responses_request_body(
|
||||
body_json,
|
||||
&mapped_model,
|
||||
spec_metadata.api_format,
|
||||
provider_api_format,
|
||||
upstream_is_stream,
|
||||
force_body_stream_field,
|
||||
transport.provider.provider_type.as_str(),
|
||||
if is_kiro_claude_cli {
|
||||
None
|
||||
} else {
|
||||
transport.endpoint.body_rules.as_ref()
|
||||
},
|
||||
Some(input.auth_context.api_key_id.as_str()),
|
||||
&parts.headers,
|
||||
enable_model_directives,
|
||||
)
|
||||
} else {
|
||||
build_local_openai_responses_request_body(
|
||||
body_json,
|
||||
&mapped_model,
|
||||
upstream_is_stream,
|
||||
force_body_stream_field,
|
||||
transport.provider.provider_type.as_str(),
|
||||
provider_api_format,
|
||||
if is_kiro_claude_cli {
|
||||
None
|
||||
} else {
|
||||
transport.endpoint.body_rules.as_ref()
|
||||
},
|
||||
Some(input.auth_context.api_key_id.as_str()),
|
||||
&parts.headers,
|
||||
enable_model_directives,
|
||||
)
|
||||
}) else {
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let Some(mut base_provider_request_body) =
|
||||
(if is_grok && is_grok_text_provider_api_format(provider_api_format) {
|
||||
build_local_openai_responses_request_body(
|
||||
body_json,
|
||||
&mapped_model,
|
||||
upstream_is_stream,
|
||||
force_body_stream_field,
|
||||
transport.provider.provider_type.as_str(),
|
||||
spec_metadata.api_format,
|
||||
transport.endpoint.body_rules.as_ref(),
|
||||
Some(input.auth_context.api_key_id.as_str()),
|
||||
effective_headers,
|
||||
enable_model_directives,
|
||||
)
|
||||
} else if needs_bidirectional_conversion {
|
||||
build_cross_format_openai_responses_request_body(
|
||||
body_json,
|
||||
&mapped_model,
|
||||
spec_metadata.api_format,
|
||||
provider_api_format,
|
||||
upstream_is_stream,
|
||||
force_body_stream_field,
|
||||
transport.provider.provider_type.as_str(),
|
||||
if is_kiro_claude_cli {
|
||||
None
|
||||
} else {
|
||||
transport.endpoint.body_rules.as_ref()
|
||||
},
|
||||
Some(input.auth_context.api_key_id.as_str()),
|
||||
effective_headers,
|
||||
enable_model_directives,
|
||||
)
|
||||
} else {
|
||||
build_local_openai_responses_request_body(
|
||||
body_json,
|
||||
&mapped_model,
|
||||
upstream_is_stream,
|
||||
force_body_stream_field,
|
||||
transport.provider.provider_type.as_str(),
|
||||
provider_api_format,
|
||||
if is_kiro_claude_cli {
|
||||
None
|
||||
} else {
|
||||
transport.endpoint.body_rules.as_ref()
|
||||
},
|
||||
Some(input.auth_context.api_key_id.as_str()),
|
||||
effective_headers,
|
||||
enable_model_directives,
|
||||
)
|
||||
})
|
||||
else {
|
||||
mark_skipped_local_openai_responses_candidate_with_extra_data(
|
||||
state,
|
||||
input,
|
||||
@@ -390,7 +429,9 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
.await;
|
||||
}
|
||||
|
||||
let Some(upstream_url) = (if needs_bidirectional_conversion {
|
||||
let Some(upstream_url) = (if is_grok && is_grok_text_provider_api_format(provider_api_format) {
|
||||
Some(build_grok_upstream_url(transport, GROK_CHAT_PATH))
|
||||
} else if needs_bidirectional_conversion {
|
||||
build_cross_format_openai_responses_upstream_url(
|
||||
parts,
|
||||
transport,
|
||||
@@ -427,48 +468,86 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
.as_ref()
|
||||
.map(build_antigravity_static_identity_headers)
|
||||
.unwrap_or_default();
|
||||
let Some(resolved_headers) =
|
||||
build_standard_provider_request_headers(StandardProviderRequestHeadersInput {
|
||||
let resolved_headers = if is_grok && is_grok_text_provider_api_format(provider_api_format) {
|
||||
let Some(headers) = build_grok_browser_headers(GrokHeaderInput {
|
||||
transport,
|
||||
provider_api_format,
|
||||
same_format,
|
||||
headers: &parts.headers,
|
||||
auth_header: &auth_header,
|
||||
auth_value: &auth_value,
|
||||
extra_headers: &extra_headers,
|
||||
transport_profile: transport_profile.as_ref(),
|
||||
request_headers: Some(effective_headers),
|
||||
content_type: "application/json",
|
||||
accept: "text/event-stream",
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: body_json,
|
||||
upstream_is_stream,
|
||||
})
|
||||
else {
|
||||
mark_skipped_local_openai_responses_candidate_with_failure_diagnostic(
|
||||
state,
|
||||
input,
|
||||
trace_id,
|
||||
candidate,
|
||||
candidate_index,
|
||||
candidate_id,
|
||||
"transport_header_rules_apply_failed",
|
||||
CandidateFailureDiagnostic::header_rules_apply_failed(
|
||||
spec_metadata.api_format,
|
||||
}) else {
|
||||
mark_skipped_local_openai_responses_candidate_with_failure_diagnostic(
|
||||
state,
|
||||
input,
|
||||
trace_id,
|
||||
candidate,
|
||||
candidate_index,
|
||||
candidate_id,
|
||||
"transport_header_rules_apply_failed",
|
||||
CandidateFailureDiagnostic::header_rules_apply_failed(
|
||||
spec_metadata.api_format,
|
||||
provider_api_format,
|
||||
"grok_openai_responses_headers",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
};
|
||||
crate::ai_serving::transport::StandardProviderRequestHeaders {
|
||||
headers,
|
||||
auth_header: auth_header.clone(),
|
||||
auth_value: auth_value.clone(),
|
||||
}
|
||||
} else {
|
||||
let Some(resolved_headers) =
|
||||
build_standard_provider_request_headers(StandardProviderRequestHeadersInput {
|
||||
transport,
|
||||
provider_api_format,
|
||||
"openai_responses_headers",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
same_format,
|
||||
headers: effective_headers,
|
||||
auth_header: &auth_header,
|
||||
auth_value: &auth_value,
|
||||
extra_headers: &extra_headers,
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: body_json,
|
||||
upstream_is_stream,
|
||||
})
|
||||
else {
|
||||
mark_skipped_local_openai_responses_candidate_with_failure_diagnostic(
|
||||
state,
|
||||
input,
|
||||
trace_id,
|
||||
candidate,
|
||||
candidate_index,
|
||||
candidate_id,
|
||||
"transport_header_rules_apply_failed",
|
||||
CandidateFailureDiagnostic::header_rules_apply_failed(
|
||||
spec_metadata.api_format,
|
||||
provider_api_format,
|
||||
"openai_responses_headers",
|
||||
),
|
||||
)
|
||||
.await;
|
||||
return None;
|
||||
};
|
||||
resolved_headers
|
||||
};
|
||||
let mut provider_request_headers = resolved_headers.headers;
|
||||
apply_codex_openai_responses_special_headers(
|
||||
&mut provider_request_headers,
|
||||
&provider_request_body,
|
||||
&parts.headers,
|
||||
transport.provider.provider_type.as_str(),
|
||||
provider_api_format,
|
||||
Some(trace_id),
|
||||
transport.key.decrypted_auth_config.as_deref(),
|
||||
);
|
||||
if !is_grok {
|
||||
apply_codex_openai_responses_special_headers(
|
||||
&mut provider_request_headers,
|
||||
&provider_request_body,
|
||||
effective_headers,
|
||||
transport.provider.provider_type.as_str(),
|
||||
provider_api_format,
|
||||
Some(trace_id),
|
||||
transport.key.decrypted_auth_config.as_deref(),
|
||||
);
|
||||
}
|
||||
|
||||
let (execution_strategy, conversion_mode) =
|
||||
ai_local_execution_contract_for_formats(spec_metadata.api_format, provider_api_format);
|
||||
@@ -516,6 +595,7 @@ pub(crate) async fn resolve_local_openai_responses_candidate_payload_parts(
|
||||
},
|
||||
upstream_is_stream,
|
||||
transport: Arc::clone(transport),
|
||||
transport_profile,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -545,12 +625,13 @@ async fn build_kiro_openai_responses_payload_parts(
|
||||
kiro_auth: &KiroRequestAuth,
|
||||
) -> Option<LocalOpenAiResponsesCandidatePayloadParts> {
|
||||
let candidate = &eligible.candidate;
|
||||
let effective_headers = input.effective_headers(&parts.headers);
|
||||
let provider_request_body = match build_kiro_provider_request_body(
|
||||
&claude_request_body,
|
||||
&mapped_model,
|
||||
&kiro_auth.auth_config,
|
||||
transport.endpoint.body_rules.as_ref(),
|
||||
Some(&parts.headers),
|
||||
Some(effective_headers),
|
||||
) {
|
||||
Some(body) => body,
|
||||
None => {
|
||||
@@ -601,7 +682,7 @@ async fn build_kiro_openai_responses_payload_parts(
|
||||
}
|
||||
};
|
||||
let provider_request_headers = match build_kiro_provider_headers(KiroProviderHeadersInput {
|
||||
headers: &parts.headers,
|
||||
headers: effective_headers,
|
||||
provider_request_body: &provider_request_body,
|
||||
original_request_body: original_body_json,
|
||||
header_rules: transport.endpoint.header_rules.as_ref(),
|
||||
@@ -666,5 +747,6 @@ async fn build_kiro_openai_responses_payload_parts(
|
||||
envelope_name: Some(KIRO_ENVELOPE_NAME),
|
||||
upstream_is_stream,
|
||||
transport: Arc::clone(transport),
|
||||
transport_profile: None,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -19,6 +19,7 @@ use crate::ai_serving::planner::candidate_source::{
|
||||
};
|
||||
use crate::ai_serving::planner::common::extract_standard_requested_model;
|
||||
use crate::ai_serving::planner::decision_input::{
|
||||
attach_routing_policy_to_local_requested_model_input,
|
||||
build_local_requested_model_decision_input, resolve_local_authenticated_decision_input,
|
||||
};
|
||||
use crate::ai_serving::planner::materialization_policy::{
|
||||
@@ -48,7 +49,7 @@ pub(crate) async fn resolve_local_openai_responses_decision_input(
|
||||
decision: &GatewayControlDecision,
|
||||
body_json: &serde_json::Value,
|
||||
plan_kind: &str,
|
||||
) -> Option<LocalOpenAiResponsesDecisionInput> {
|
||||
) -> Result<Option<LocalOpenAiResponsesDecisionInput>, GatewayError> {
|
||||
let Some(auth_context) = resolve_local_decision_execution_runtime_auth_context(decision) else {
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
@@ -65,7 +66,7 @@ pub(crate) async fn resolve_local_openai_responses_decision_input(
|
||||
extract_standard_requested_model(body_json).as_deref(),
|
||||
"missing_auth_context",
|
||||
);
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let Some(requested_model) = extract_standard_requested_model(body_json) else {
|
||||
@@ -81,7 +82,7 @@ pub(crate) async fn resolve_local_openai_responses_decision_input(
|
||||
None,
|
||||
"missing_requested_model",
|
||||
);
|
||||
return None;
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let resolved_input = match resolve_local_authenticated_decision_input(
|
||||
@@ -108,7 +109,7 @@ pub(crate) async fn resolve_local_openai_responses_decision_input(
|
||||
Some(requested_model.as_str()),
|
||||
"auth_snapshot_missing",
|
||||
);
|
||||
return None;
|
||||
return Ok(None);
|
||||
}
|
||||
Err(err) => {
|
||||
warn!(
|
||||
@@ -124,14 +125,30 @@ pub(crate) async fn resolve_local_openai_responses_decision_input(
|
||||
Some(requested_model.as_str()),
|
||||
"auth_snapshot_read_failed",
|
||||
);
|
||||
return None;
|
||||
return Err(err);
|
||||
}
|
||||
};
|
||||
|
||||
let mut input = build_local_requested_model_decision_input(resolved_input, requested_model);
|
||||
input.request_auth_channel = decision.request_auth_channel.clone();
|
||||
input.client_session_affinity = client_session_affinity_from_parts(parts, Some(body_json));
|
||||
Some(input)
|
||||
if let Err(err) = attach_routing_policy_to_local_requested_model_input(
|
||||
state,
|
||||
parts,
|
||||
&mut input,
|
||||
body_json,
|
||||
"openai:responses",
|
||||
)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
error = ?err,
|
||||
"gateway local openai responses decision routing profile resolution failed"
|
||||
);
|
||||
return Err(err);
|
||||
}
|
||||
Ok(Some(input))
|
||||
}
|
||||
|
||||
pub(crate) async fn materialize_local_openai_responses_candidate_attempts(
|
||||
@@ -157,6 +174,7 @@ pub(crate) async fn materialize_local_openai_responses_candidate_attempts(
|
||||
spec_metadata.require_streaming,
|
||||
input.required_capabilities.as_ref(),
|
||||
&input.auth_snapshot,
|
||||
input.routing_policy.as_ref(),
|
||||
input.client_session_affinity.as_ref(),
|
||||
true,
|
||||
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
|
||||
@@ -170,6 +188,7 @@ pub(crate) async fn materialize_local_openai_responses_candidate_attempts(
|
||||
Some(&input.auth_snapshot),
|
||||
input.client_session_affinity.as_ref(),
|
||||
input.required_capabilities.as_ref(),
|
||||
input.routing_policy.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
@@ -256,6 +275,7 @@ pub(crate) async fn build_local_openai_responses_candidate_attempt_source<'a>(
|
||||
&input.auth_snapshot,
|
||||
input.client_session_affinity.as_ref(),
|
||||
input.required_capabilities.as_ref(),
|
||||
input.routing_policy.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
|
||||
@@ -103,10 +103,11 @@ pub(crate) async fn maybe_build_sync_local_openai_responses_decision_payload(
|
||||
let Some(input) = resolve_local_openai_responses_decision_input(
|
||||
state, parts, trace_id, decision, body_json, plan_kind,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
|
||||
let (mut source, _) = build_local_openai_responses_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, spec,
|
||||
@@ -117,7 +118,7 @@ pub(crate) async fn maybe_build_sync_local_openai_responses_decision_payload(
|
||||
if let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
|
||||
state, parts, trace_id, body_json, &input, attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
{
|
||||
return Ok(Some(payload));
|
||||
}
|
||||
@@ -141,10 +142,11 @@ pub(crate) async fn maybe_build_stream_local_openai_responses_decision_payload(
|
||||
let Some(input) = resolve_local_openai_responses_decision_input(
|
||||
state, parts, trace_id, decision, body_json, plan_kind,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let body_json = input.effective_body_json(body_json);
|
||||
|
||||
let (mut source, _) = build_local_openai_responses_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, spec,
|
||||
@@ -155,7 +157,7 @@ pub(crate) async fn maybe_build_stream_local_openai_responses_decision_payload(
|
||||
if let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
|
||||
state, parts, trace_id, body_json, &input, attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
{
|
||||
return Ok(Some(payload));
|
||||
}
|
||||
|
||||
@@ -29,7 +29,7 @@ pub(crate) struct LocalOpenAiResponsesSyncAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_json: serde_json::Value,
|
||||
input: LocalOpenAiResponsesDecisionInput,
|
||||
spec: LocalOpenAiResponsesSpec,
|
||||
candidates: LocalOpenAiResponsesCandidateAttemptSource<'a>,
|
||||
@@ -39,7 +39,7 @@ pub(crate) struct LocalOpenAiResponsesStreamAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_json: serde_json::Value,
|
||||
input: LocalOpenAiResponsesDecisionInput,
|
||||
spec: LocalOpenAiResponsesSpec,
|
||||
candidates: LocalOpenAiResponsesCandidateAttemptSource<'a>,
|
||||
@@ -62,7 +62,7 @@ pub(super) async fn build_local_sync_attempt_source<'a>(
|
||||
body_json,
|
||||
spec_metadata.decision_kind,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
@@ -74,8 +74,13 @@ pub(super) async fn build_local_sync_attempt_source<'a>(
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let effective_body_json = input.effective_body_json(body_json).clone();
|
||||
let (candidates, candidate_count) = build_local_openai_responses_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, spec,
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
&effective_body_json,
|
||||
spec,
|
||||
)
|
||||
.await?;
|
||||
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
|
||||
@@ -88,7 +93,7 @@ pub(super) async fn build_local_sync_attempt_source<'a>(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
body_json: effective_body_json,
|
||||
input,
|
||||
spec,
|
||||
candidates,
|
||||
@@ -114,7 +119,7 @@ pub(super) async fn build_local_stream_attempt_source<'a>(
|
||||
body_json,
|
||||
spec_metadata.decision_kind,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
@@ -126,8 +131,13 @@ pub(super) async fn build_local_stream_attempt_source<'a>(
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let effective_body_json = input.effective_body_json(body_json).clone();
|
||||
let (candidates, candidate_count) = build_local_openai_responses_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, spec,
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
&effective_body_json,
|
||||
spec,
|
||||
)
|
||||
.await?;
|
||||
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
|
||||
@@ -140,7 +150,7 @@ pub(super) async fn build_local_stream_attempt_source<'a>(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
body_json: effective_body_json,
|
||||
input,
|
||||
spec,
|
||||
candidates,
|
||||
@@ -214,19 +224,19 @@ impl LocalOpenAiResponsesSyncAttemptSource<'_> {
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match build_openai_responses_sync_plan_from_decision(
|
||||
self.parts,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
payload,
|
||||
self.spec.compact,
|
||||
) {
|
||||
@@ -252,19 +262,19 @@ impl LocalOpenAiResponsesStreamAttemptSource<'_> {
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match build_openai_responses_stream_plan_from_decision(
|
||||
self.parts,
|
||||
self.body_json,
|
||||
&self.body_json,
|
||||
payload,
|
||||
self.spec.compact,
|
||||
) {
|
||||
@@ -298,7 +308,7 @@ pub(super) async fn build_local_sync_plan_and_reports(
|
||||
body_json,
|
||||
spec_metadata.decision_kind,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(Vec::new());
|
||||
};
|
||||
@@ -325,7 +335,7 @@ pub(super) async fn build_local_sync_plan_and_reports(
|
||||
let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
|
||||
state, parts, trace_id, body_json, &input, attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
@@ -370,7 +380,7 @@ pub(super) async fn build_local_stream_plan_and_reports(
|
||||
body_json,
|
||||
spec_metadata.decision_kind,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
return Ok(Vec::new());
|
||||
};
|
||||
@@ -397,7 +407,7 @@ pub(super) async fn build_local_stream_plan_and_reports(
|
||||
let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
|
||||
state, parts, trace_id, body_json, &input, attempt, spec,
|
||||
)
|
||||
.await
|
||||
.await?
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
@@ -61,6 +61,7 @@ pub(crate) use aether_ai_formats::api::{
|
||||
maybe_build_standard_sync_finalize_product_from_normalized_payload, model_directive_base_model,
|
||||
normalize_api_format_alias, normalize_claude_request_to_openai_chat_request,
|
||||
normalize_gemini_request_to_openai_chat_request, normalize_openai_image_request,
|
||||
normalize_openai_image_request_with_options,
|
||||
normalize_openai_responses_request_to_openai_chat_request,
|
||||
normalize_provider_private_report_context, normalize_provider_private_response_value,
|
||||
normalize_standard_request_to_openai_chat_request, openai_image_operation_from_path,
|
||||
@@ -97,10 +98,10 @@ pub(crate) use aether_ai_formats::api::{
|
||||
LocalStandardSourceMode, LocalStandardSpec, LocalSyncReportParts, LocalVideoCreateFamily,
|
||||
LocalVideoCreateSpec, NormalizedOpenAiImageRequest, OpenAIChatClientEmitter,
|
||||
OpenAIChatProviderState, OpenAIResponsesClientEmitter, OpenAIResponsesProviderState,
|
||||
OpenAiImageOperation, OpenAiImageRequestForGemini, OpenAiImageResponseFormat,
|
||||
OpenAiImageStreamState, OpenAiImageSyncFinalizeProduct, ProviderAdaptationDescriptor,
|
||||
ProviderAdaptationSurface, ProviderPrivateStreamNormalizer, RequestConversionKind,
|
||||
StandardCrossFormatSyncProduct, StandardSyncFinalizeNormalizedProduct,
|
||||
OpenAiImageNormalizeOptions, OpenAiImageOperation, OpenAiImageRequestForGemini,
|
||||
OpenAiImageResponseFormat, OpenAiImageStreamState, OpenAiImageSyncFinalizeProduct,
|
||||
ProviderAdaptationDescriptor, ProviderAdaptationSurface, ProviderPrivateStreamNormalizer,
|
||||
RequestConversionKind, StandardCrossFormatSyncProduct, StandardSyncFinalizeNormalizedProduct,
|
||||
StreamingStandardFormatMatrix, SyncChatResponseConversionKind, SyncCliResponseConversionKind,
|
||||
SyncToStreamBridgeOutcome, ANTIGRAVITY_V1INTERNAL_ENVELOPE_NAME, CLAUDE_CHAT_STREAM_PLAN_KIND,
|
||||
CLAUDE_CHAT_STREAM_SUCCESS_REPORT_KIND, CLAUDE_CHAT_SYNC_ERROR_REPORT_KIND,
|
||||
|
||||
@@ -14,6 +14,10 @@ pub(crate) mod kiro {
|
||||
pub(crate) use aether_provider_transport::kiro::*;
|
||||
}
|
||||
|
||||
pub(crate) mod grok {
|
||||
pub(crate) use aether_provider_transport::grok::*;
|
||||
}
|
||||
|
||||
pub(crate) mod oauth_refresh {
|
||||
pub(crate) use aether_provider_transport::oauth_refresh::*;
|
||||
}
|
||||
@@ -55,6 +59,7 @@ pub(crate) use aether_provider_transport::{
|
||||
body_rules_handle_path, body_rules_have_enabled_rules,
|
||||
build_cross_format_openai_chat_upstream_url, build_cross_format_openai_responses_upstream_url,
|
||||
build_gemini_files_headers, build_gemini_files_request_body, build_gemini_files_upstream_url,
|
||||
build_grok_app_chat_body, build_grok_browser_headers, build_grok_upstream_url,
|
||||
build_kiro_cross_format_upstream_url, build_local_openai_chat_upstream_url,
|
||||
build_local_openai_responses_upstream_url, build_openai_image_headers,
|
||||
build_openai_image_upstream_url, build_passthrough_headers, build_request_trace_proxy_value,
|
||||
@@ -74,7 +79,7 @@ pub(crate) use aether_provider_transport::{
|
||||
request_conversion_enabled_for_transport, request_conversion_transport_supported,
|
||||
request_conversion_transport_unsupported_reason, request_pair_allowed_for_transport,
|
||||
request_pair_direct_auth, request_pair_transport_unsupported_reason, resolve_gemini_files_auth,
|
||||
resolve_openai_image_auth, resolve_same_format_provider_direct_auth,
|
||||
resolve_grok_session_auth, resolve_openai_image_auth, resolve_same_format_provider_direct_auth,
|
||||
resolve_transport_execution_timeouts, resolve_transport_profile,
|
||||
resolve_transport_proxy_snapshot, resolve_transport_proxy_snapshot_with_tunnel_affinity,
|
||||
resolve_video_create_auth, same_format_provider_transport_supported,
|
||||
@@ -84,12 +89,12 @@ pub(crate) use aether_provider_transport::{
|
||||
supports_local_oauth_request_auth_resolution, transport_proxy_is_locally_supported,
|
||||
video_create_transport_unsupported_reason, CandidateTransportPolicyFacts,
|
||||
GatewayProviderTransportSnapshot, GeminiFilesHeadersInput, GeminiFilesRequestBodyError,
|
||||
GeminiFilesRequestBodyParts, LocalResolvedOAuthRequestAuth, ProviderOpenAiImageHeadersInput,
|
||||
ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput, SameFormatProviderFamily,
|
||||
SameFormatProviderHeadersInput, SameFormatProviderRequestBehavior,
|
||||
GeminiFilesRequestBodyParts, GrokHeaderInput, LocalResolvedOAuthRequestAuth,
|
||||
ProviderOpenAiImageHeadersInput, ProviderVideoCreateFamily, ProviderVideoCreateHeadersInput,
|
||||
SameFormatProviderFamily, SameFormatProviderHeadersInput, SameFormatProviderRequestBehavior,
|
||||
SameFormatProviderRequestBehaviorParams, SameFormatProviderRequestBodyInput,
|
||||
SameFormatProviderUpstreamUrlParams, StandardPlanFallbackAcceptPolicy,
|
||||
StandardPlanFallbackHeadersInput, StandardProviderRequestHeaders,
|
||||
StandardProviderRequestHeadersInput, TransportRequestBodySemanticsError,
|
||||
TransportRequestUrlParams,
|
||||
TransportRequestUrlParams, GROK_CHAT_PATH, GROK_INTERNAL_HEADER, GROK_RATE_LIMITS_PATH,
|
||||
};
|
||||
|
||||
@@ -17,7 +17,6 @@ const AI_POST_ROUTE_PATTERNS: &[&str] = &[
|
||||
"/v1/responses/compact",
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits",
|
||||
"/v1/images/variations",
|
||||
];
|
||||
|
||||
const AI_ANY_ROUTE_PATTERNS: &[&str] = &[
|
||||
|
||||
@@ -115,7 +115,6 @@ pub(crate) const RUST_FRONTDOOR_OWNED_ROUTE_PATTERNS: &[&str] = &[
|
||||
"/v1/rerank",
|
||||
"/v1/images/generations",
|
||||
"/v1/images/edits",
|
||||
"/v1/images/variations",
|
||||
"/v1/messages",
|
||||
"/v1/messages/count_tokens",
|
||||
"/v1/responses",
|
||||
|
||||
@@ -106,7 +106,7 @@ async fn balance_capacity_rejection(
|
||||
requested_model: Option<&str>,
|
||||
body: &Bytes,
|
||||
) -> Result<Option<GatewayLocalAuthRejection>, GatewayError> {
|
||||
if auth_context.api_key_is_standalone || auth_context.admin_bypass_limits {
|
||||
if auth_context.api_key_is_standalone {
|
||||
return Ok(None);
|
||||
}
|
||||
if auth_context.local_rejection.is_some() {
|
||||
@@ -816,6 +816,43 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn admin_bypass_limits_does_not_skip_exhausted_daily_quota_capacity() {
|
||||
let context = billing_context_with_pricing(
|
||||
Some(json!({
|
||||
"tiers": [{
|
||||
"up_to": null,
|
||||
"input_price_per_1m": 1.0,
|
||||
"output_price_per_1m": 2.0
|
||||
}]
|
||||
})),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
);
|
||||
let state = state_with_quota_and_wallet(quota_availability(0.0, false), context);
|
||||
let mut decision = decision_with_allowed_models(vec!["gpt-5".to_string()]);
|
||||
if let Some(auth_context) = decision.auth_context.as_mut() {
|
||||
auth_context.admin_bypass_limits = true;
|
||||
}
|
||||
let uri: Uri = "/v1/chat/completions".parse().expect("uri should parse");
|
||||
let body = Bytes::from_static(
|
||||
br#"{"model":"gpt-5","messages":[{"role":"user","content":"hi"}],"stream":true}"#,
|
||||
);
|
||||
|
||||
let rejection =
|
||||
request_model_local_rejection(&state, Some(&decision), &uri, &json_headers(), &body)
|
||||
.await
|
||||
.expect("quota rejection should resolve");
|
||||
|
||||
assert_eq!(
|
||||
rejection,
|
||||
Some(GatewayLocalAuthRejection::BalanceDenied {
|
||||
remaining: Some(0.0),
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn positive_balance_still_denies_known_cost_above_available_capacity() {
|
||||
let context = billing_context_with_pricing(
|
||||
|
||||
@@ -9,7 +9,8 @@ pub(crate) use gate::{
|
||||
request_model_local_rejection, should_buffer_request_for_local_auth,
|
||||
trusted_auth_local_rejection, GatewayLocalAuthRejection,
|
||||
};
|
||||
pub(super) use resolution::{resolve_control_decision_auth, ControlDecisionAuthResolution};
|
||||
pub(crate) use resolution::{
|
||||
resolve_execution_runtime_auth_context, GatewayAdminPrincipalContext, GatewayControlAuthContext,
|
||||
refresh_execution_runtime_auth_context, resolve_execution_runtime_auth_context,
|
||||
GatewayAdminPrincipalContext, GatewayControlAuthContext,
|
||||
};
|
||||
pub(super) use resolution::{resolve_control_decision_auth, ControlDecisionAuthResolution};
|
||||
|
||||
@@ -433,7 +433,14 @@ pub(crate) async fn resolve_execution_runtime_auth_context(
|
||||
let _ = trace_id;
|
||||
|
||||
if let Some(auth_context) = decision.auth_context.clone() {
|
||||
return Ok(Some(auth_context));
|
||||
return Ok(Some(
|
||||
refresh_execution_runtime_auth_context(
|
||||
state,
|
||||
auth_context,
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
)
|
||||
.await?,
|
||||
));
|
||||
}
|
||||
|
||||
let Some(auth_endpoint_signature) = decision.auth_endpoint_signature.as_deref() else {
|
||||
@@ -445,7 +452,14 @@ pub(crate) async fn resolve_execution_runtime_auth_context(
|
||||
};
|
||||
|
||||
if let Some(auth_context) = get_cached_auth_context(state, &cache_key) {
|
||||
return Ok(Some(auth_context));
|
||||
let refreshed = refresh_execution_runtime_auth_context(
|
||||
state,
|
||||
auth_context,
|
||||
Some(auth_endpoint_signature),
|
||||
)
|
||||
.await?;
|
||||
put_cached_auth_context(state, cache_key, refreshed.clone());
|
||||
return Ok(Some(refreshed));
|
||||
}
|
||||
|
||||
if let Some(auth_context) =
|
||||
@@ -461,6 +475,56 @@ pub(crate) async fn resolve_execution_runtime_auth_context(
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
pub(crate) async fn refresh_execution_runtime_auth_context(
|
||||
state: &AppState,
|
||||
auth_context: GatewayControlAuthContext,
|
||||
auth_endpoint_signature: Option<&str>,
|
||||
) -> Result<GatewayControlAuthContext, GatewayError> {
|
||||
if auth_context.local_rejection.is_some() || !auth_context.access_allowed {
|
||||
return Ok(auth_context);
|
||||
}
|
||||
let Some(auth_endpoint_signature) = auth_endpoint_signature
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return Ok(auth_context);
|
||||
};
|
||||
if !state.has_auth_api_key_reader()
|
||||
|| auth_context.user_id.trim().is_empty()
|
||||
|| auth_context.api_key_id.trim().is_empty()
|
||||
{
|
||||
return Ok(auth_context);
|
||||
}
|
||||
|
||||
let snapshot = state
|
||||
.data
|
||||
.read_auth_api_key_snapshot(
|
||||
&auth_context.user_id,
|
||||
&auth_context.api_key_id,
|
||||
current_unix_secs(),
|
||||
)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
let Some(snapshot) = snapshot else {
|
||||
let mut denied = auth_context;
|
||||
denied.access_allowed = false;
|
||||
denied.local_rejection = Some(GatewayLocalAuthRejection::InvalidApiKey);
|
||||
denied.balance_remaining = None;
|
||||
return Ok(denied);
|
||||
};
|
||||
|
||||
let wallet_access = resolve_wallet_auth_gate(state, &snapshot).await?;
|
||||
Ok(build_data_backed_auth_context(
|
||||
state,
|
||||
snapshot,
|
||||
auth_endpoint_signature,
|
||||
Some(true),
|
||||
auth_context.balance_remaining,
|
||||
wallet_access,
|
||||
)
|
||||
.await)
|
||||
}
|
||||
|
||||
fn put_cached_auth_context(
|
||||
state: &AppState,
|
||||
cache_key: String,
|
||||
@@ -609,7 +673,7 @@ async fn build_data_backed_auth_context(
|
||||
.api_key_expires_at_unix_secs
|
||||
.is_some_and(|expires_at| expires_at < current_unix_secs());
|
||||
let locked_api_key = snapshot.api_key_is_locked && !snapshot.api_key_is_standalone;
|
||||
let access_allowed = header_access_allowed
|
||||
let key_access_allowed = header_access_allowed
|
||||
.map(|value| value && snapshot.currently_usable)
|
||||
.unwrap_or(snapshot.currently_usable);
|
||||
let wallet_remaining = wallet_access
|
||||
@@ -656,7 +720,7 @@ async fn build_data_backed_auth_context(
|
||||
user_id: snapshot.user_id,
|
||||
api_key_id: snapshot.api_key_id,
|
||||
balance_remaining: wallet_remaining.or(balance_remaining),
|
||||
access_allowed,
|
||||
access_allowed: key_access_allowed && local_rejection.is_none(),
|
||||
user_rate_limit: snapshot.user_rate_limit,
|
||||
api_key_rate_limit: snapshot.api_key_rate_limit,
|
||||
api_key_is_standalone: snapshot.api_key_is_standalone,
|
||||
@@ -835,13 +899,20 @@ mod tests {
|
||||
InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot,
|
||||
};
|
||||
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
|
||||
use aether_data::repository::wallet::{
|
||||
InMemoryWalletRepository, StoredWalletSnapshot, WalletReadRepository,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
|
||||
};
|
||||
use axum::http::{HeaderMap, Uri};
|
||||
|
||||
use super::{resolve_data_backed_auth_context, GatewayLocalAuthRejection};
|
||||
use super::{
|
||||
resolve_data_backed_auth_context, resolve_execution_runtime_auth_context,
|
||||
GatewayLocalAuthRejection,
|
||||
};
|
||||
use crate::control::auth::credentials::hash_api_key;
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::data::GatewayDataState;
|
||||
use crate::AppState;
|
||||
|
||||
@@ -946,6 +1017,154 @@ mod tests {
|
||||
assert_eq!(repository.touch_count("key-1"), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn data_backed_auth_context_marks_wallet_denial_as_not_allowed() {
|
||||
let api_key = "sk-test-empty-wallet";
|
||||
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
|
||||
Some(hash_api_key(api_key)),
|
||||
sample_snapshot("key-empty-wallet", "user-empty-wallet"),
|
||||
)]));
|
||||
let wallet_repository = Arc::new(InMemoryWalletRepository::seed(vec![
|
||||
StoredWalletSnapshot::new(
|
||||
"wallet-empty".to_string(),
|
||||
Some("user-empty-wallet".to_string()),
|
||||
None,
|
||||
0.0,
|
||||
0.0,
|
||||
"finite".to_string(),
|
||||
"USD".to_string(),
|
||||
"active".to_string(),
|
||||
0.0,
|
||||
0.0,
|
||||
0.0,
|
||||
0.0,
|
||||
100,
|
||||
)
|
||||
.expect("wallet should build"),
|
||||
]));
|
||||
let data =
|
||||
GatewayDataState::with_auth_and_wallet_for_tests(auth_repository, wallet_repository);
|
||||
let state = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(data);
|
||||
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert(
|
||||
http::header::AUTHORIZATION,
|
||||
format!("Bearer {api_key}").parse().unwrap(),
|
||||
);
|
||||
|
||||
let auth_context = resolve_data_backed_auth_context(
|
||||
&state,
|
||||
&headers,
|
||||
&uri("/v1/chat/completions"),
|
||||
Some("openai:chat"),
|
||||
)
|
||||
.await
|
||||
.expect("resolution should succeed")
|
||||
.expect("auth context should exist");
|
||||
|
||||
assert_eq!(
|
||||
auth_context.local_rejection,
|
||||
Some(GatewayLocalAuthRejection::BalanceDenied {
|
||||
remaining: Some(0.0),
|
||||
})
|
||||
);
|
||||
assert!(!auth_context.access_allowed);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn execution_runtime_auth_context_revalidates_cached_wallet_state() {
|
||||
let api_key = "sk-test-runtime-wallet-cache";
|
||||
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::seed(vec![(
|
||||
Some(hash_api_key(api_key)),
|
||||
sample_snapshot("key-runtime-wallet-cache", "user-runtime-wallet-cache"),
|
||||
)]));
|
||||
let wallet_repository = Arc::new(InMemoryWalletRepository::seed(vec![
|
||||
StoredWalletSnapshot::new(
|
||||
"wallet-runtime-cache".to_string(),
|
||||
Some("user-runtime-wallet-cache".to_string()),
|
||||
None,
|
||||
10.0,
|
||||
0.0,
|
||||
"finite".to_string(),
|
||||
"USD".to_string(),
|
||||
"active".to_string(),
|
||||
10.0,
|
||||
0.0,
|
||||
0.0,
|
||||
0.0,
|
||||
100,
|
||||
)
|
||||
.expect("wallet should build"),
|
||||
]));
|
||||
let data = GatewayDataState::with_auth_and_wallet_for_tests(
|
||||
auth_repository,
|
||||
Arc::clone(&wallet_repository),
|
||||
);
|
||||
let state = AppState::new()
|
||||
.expect("state should build")
|
||||
.with_data_state_for_tests(data);
|
||||
let decision = GatewayControlDecision::synthetic(
|
||||
"/v1/chat/completions",
|
||||
Some("ai_public".to_string()),
|
||||
Some("openai".to_string()),
|
||||
Some("chat".to_string()),
|
||||
Some("openai:chat".to_string()),
|
||||
);
|
||||
let mut headers = HeaderMap::new();
|
||||
headers.insert("x-api-key", api_key.parse().unwrap());
|
||||
|
||||
let first = resolve_execution_runtime_auth_context(
|
||||
&state,
|
||||
&decision,
|
||||
&headers,
|
||||
&uri("/v1/chat/completions"),
|
||||
"trace-runtime-wallet-cache",
|
||||
)
|
||||
.await
|
||||
.expect("resolution should succeed")
|
||||
.expect("auth context should exist");
|
||||
assert!(first.access_allowed);
|
||||
|
||||
wallet_repository
|
||||
.update_auth_user_wallet_snapshot(
|
||||
"user-runtime-wallet-cache",
|
||||
0.0,
|
||||
0.0,
|
||||
"finite",
|
||||
"USD",
|
||||
"active",
|
||||
10.0,
|
||||
10.0,
|
||||
0.0,
|
||||
0.0,
|
||||
Some(101),
|
||||
)
|
||||
.await
|
||||
.expect("wallet update should succeed")
|
||||
.expect("wallet should exist");
|
||||
|
||||
let second = resolve_execution_runtime_auth_context(
|
||||
&state,
|
||||
&decision,
|
||||
&headers,
|
||||
&uri("/v1/chat/completions"),
|
||||
"trace-runtime-wallet-cache",
|
||||
)
|
||||
.await
|
||||
.expect("resolution should succeed")
|
||||
.expect("auth context should exist");
|
||||
|
||||
assert_eq!(
|
||||
second.local_rejection,
|
||||
Some(GatewayLocalAuthRejection::BalanceDenied {
|
||||
remaining: Some(0.0),
|
||||
})
|
||||
);
|
||||
assert!(!second.access_allowed);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn data_backed_auth_context_allows_provider_id_for_matching_provider_type() {
|
||||
let api_key = "sk-test-provider-id";
|
||||
|
||||
@@ -137,6 +137,11 @@ const PERMISSION_GROUPS: &[PermissionGroup] = &[
|
||||
label: "代理节点",
|
||||
assignable: true,
|
||||
},
|
||||
PermissionGroup {
|
||||
scope: "routing_profiles",
|
||||
label: "调度分组",
|
||||
assignable: true,
|
||||
},
|
||||
PermissionGroup {
|
||||
scope: "security",
|
||||
label: "安全",
|
||||
@@ -446,6 +451,9 @@ fn permission_key(scope: &str, access: &str) -> &'static str {
|
||||
("proxy_nodes", "read") => "admin:proxy_nodes:read",
|
||||
("proxy_nodes", "write") => "admin:proxy_nodes:write",
|
||||
("proxy_nodes", "admin") => "admin:proxy_nodes:admin",
|
||||
("routing_profiles", "read") => "admin:routing_profiles:read",
|
||||
("routing_profiles", "write") => "admin:routing_profiles:write",
|
||||
("routing_profiles", "admin") => "admin:routing_profiles:admin",
|
||||
("security", "read") => "admin:security:read",
|
||||
("security", "write") => "admin:security:write",
|
||||
("security", "admin") => "admin:security:admin",
|
||||
|
||||
@@ -8,9 +8,10 @@ mod public;
|
||||
mod route;
|
||||
|
||||
pub(crate) use auth::{
|
||||
extract_requested_model, request_model_local_rejection, resolve_execution_runtime_auth_context,
|
||||
should_buffer_request_for_local_auth, trusted_auth_local_rejection,
|
||||
GatewayAdminPrincipalContext, GatewayControlAuthContext, GatewayLocalAuthRejection,
|
||||
extract_requested_model, refresh_execution_runtime_auth_context, request_model_local_rejection,
|
||||
resolve_execution_runtime_auth_context, should_buffer_request_for_local_auth,
|
||||
trusted_auth_local_rejection, GatewayAdminPrincipalContext, GatewayControlAuthContext,
|
||||
GatewayLocalAuthRejection,
|
||||
};
|
||||
pub(crate) use execute::{allows_control_execute_emergency, maybe_execute_via_control};
|
||||
pub(crate) use management_token_permissions::{
|
||||
|
||||
@@ -14,6 +14,8 @@ mod observability_families;
|
||||
mod operations_families;
|
||||
#[path = "admin/provider_ops_routes.rs"]
|
||||
mod provider_ops_routes;
|
||||
#[path = "admin/routing_families.rs"]
|
||||
mod routing_families;
|
||||
#[path = "admin/system_families.rs"]
|
||||
mod system_families;
|
||||
|
||||
@@ -23,6 +25,7 @@ use model_provider_families::classify_admin_model_provider_family_route;
|
||||
use observability_families::classify_admin_observability_family_route;
|
||||
use operations_families::classify_admin_operations_family_route;
|
||||
use provider_ops_routes::classify_admin_provider_ops_routes;
|
||||
use routing_families::classify_admin_routing_family_route;
|
||||
use system_families::classify_admin_system_family_route;
|
||||
|
||||
pub(super) fn classify_admin_route(
|
||||
@@ -67,6 +70,10 @@ pub(super) fn classify_admin_route(
|
||||
classify_admin_system_family_route(method, normalized_path, normalized_path_no_trailing)
|
||||
{
|
||||
Some(route)
|
||||
} else if let Some(route) =
|
||||
classify_admin_routing_family_route(method, normalized_path_no_trailing)
|
||||
{
|
||||
Some(route)
|
||||
} else if let Some(route) = classify_admin_provider_ops_routes(method, normalized_path) {
|
||||
Some(route)
|
||||
} else if let Some(route) = classify_admin_model_provider_family_route(method, normalized_path)
|
||||
|
||||
@@ -8,6 +8,56 @@ pub(super) fn classify_admin_operations_family_route(
|
||||
normalized_path_no_trailing: &str,
|
||||
) -> Option<ClassifiedRoute> {
|
||||
if method == http::Method::GET
|
||||
&& matches!(
|
||||
normalized_path,
|
||||
"/api/admin/referrals" | "/api/admin/referrals/"
|
||||
)
|
||||
{
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"referrals_manage",
|
||||
"list_referrals",
|
||||
"admin:billing",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::GET
|
||||
&& matches!(
|
||||
normalized_path,
|
||||
"/api/admin/referral-rewards" | "/api/admin/referral-rewards/"
|
||||
)
|
||||
{
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"referrals_manage",
|
||||
"list_referral_rewards",
|
||||
"admin:billing",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
&& normalized_path.starts_with("/api/admin/referral-rewards/")
|
||||
&& normalized_path.ends_with("/retry")
|
||||
&& normalized_path.matches('/').count() == 5
|
||||
{
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"referrals_manage",
|
||||
"retry_referral_reward",
|
||||
"admin:billing",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
&& normalized_path.starts_with("/api/admin/referral-rewards/")
|
||||
&& normalized_path.ends_with("/void")
|
||||
&& normalized_path.matches('/').count() == 5
|
||||
{
|
||||
Some(classified(
|
||||
"admin_proxy",
|
||||
"referrals_manage",
|
||||
"void_referral_reward",
|
||||
"admin:billing",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::GET
|
||||
&& matches!(
|
||||
normalized_path,
|
||||
"/api/admin/provider-ops/architectures" | "/api/admin/provider-ops/architectures/"
|
||||
|
||||
@@ -0,0 +1,74 @@
|
||||
use axum::http;
|
||||
|
||||
use super::{classified, ClassifiedRoute};
|
||||
|
||||
pub(super) fn classify_admin_routing_family_route(
|
||||
method: &http::Method,
|
||||
normalized_path_no_trailing: &str,
|
||||
) -> Option<ClassifiedRoute> {
|
||||
let path = normalized_path_no_trailing;
|
||||
if method == http::Method::GET && path == "/api/admin/routing/groups" {
|
||||
Some(routing_route("list_groups"))
|
||||
} else if method == http::Method::POST && path == "/api/admin/routing/groups" {
|
||||
Some(routing_route("create_group"))
|
||||
} else if method == http::Method::GET
|
||||
&& path.starts_with("/api/admin/routing/groups/")
|
||||
&& path.ends_with("/versions")
|
||||
&& path.matches('/').count() == 6
|
||||
{
|
||||
Some(routing_route("list_group_versions"))
|
||||
} else if method == http::Method::POST
|
||||
&& path.starts_with("/api/admin/routing/groups/")
|
||||
&& path.ends_with("/publish")
|
||||
&& path.matches('/').count() == 6
|
||||
{
|
||||
Some(routing_route("publish_group"))
|
||||
} else if method == http::Method::POST
|
||||
&& path.starts_with("/api/admin/routing/groups/")
|
||||
&& path.ends_with("/dry-run")
|
||||
&& path.matches('/').count() == 6
|
||||
{
|
||||
Some(routing_route("dry_run_group"))
|
||||
} else if method == http::Method::GET
|
||||
&& path.starts_with("/api/admin/routing/groups/")
|
||||
&& path.matches('/').count() == 5
|
||||
{
|
||||
Some(routing_route("get_group"))
|
||||
} else if method == http::Method::PATCH
|
||||
&& path.starts_with("/api/admin/routing/groups/")
|
||||
&& path.matches('/').count() == 5
|
||||
{
|
||||
Some(routing_route("update_group"))
|
||||
} else if method == http::Method::DELETE
|
||||
&& path.starts_with("/api/admin/routing/groups/")
|
||||
&& path.matches('/').count() == 5
|
||||
{
|
||||
Some(routing_route("delete_group"))
|
||||
} else if method == http::Method::GET && path == "/api/admin/routing/bindings" {
|
||||
Some(routing_route("list_bindings"))
|
||||
} else if method == http::Method::POST && path == "/api/admin/routing/bindings" {
|
||||
Some(routing_route("create_binding"))
|
||||
} else if method == http::Method::PATCH
|
||||
&& path.starts_with("/api/admin/routing/bindings/")
|
||||
&& path.matches('/').count() == 5
|
||||
{
|
||||
Some(routing_route("update_binding"))
|
||||
} else if method == http::Method::DELETE
|
||||
&& path.starts_with("/api/admin/routing/bindings/")
|
||||
&& path.matches('/').count() == 5
|
||||
{
|
||||
Some(routing_route("delete_binding"))
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
fn routing_route(route_kind: &'static str) -> ClassifiedRoute {
|
||||
classified(
|
||||
"admin_proxy",
|
||||
"routing_profiles_manage",
|
||||
route_kind,
|
||||
"admin:routing_profiles",
|
||||
false,
|
||||
)
|
||||
}
|
||||
@@ -55,7 +55,7 @@ pub(super) fn classify_ai_public_route(
|
||||
} else if method == http::Method::POST
|
||||
&& matches!(
|
||||
normalized_path,
|
||||
"/v1/images/generations" | "/v1/images/edits" | "/v1/images/variations"
|
||||
"/v1/images/generations" | "/v1/images/edits"
|
||||
)
|
||||
{
|
||||
Some(classified(
|
||||
|
||||
@@ -279,6 +279,20 @@ pub(super) fn classify_public_support_route(
|
||||
"user:announcements",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::GET
|
||||
&& matches!(
|
||||
normalized_path,
|
||||
"/api/announcements/users/me/required-unread"
|
||||
| "/api/announcements/users/me/required-unread/"
|
||||
)
|
||||
{
|
||||
Some(classified(
|
||||
"public_support",
|
||||
"announcement_user",
|
||||
"required_unread",
|
||||
"user:announcements",
|
||||
false,
|
||||
))
|
||||
} else if method == http::Method::POST
|
||||
&& matches!(
|
||||
normalized_path,
|
||||
@@ -462,6 +476,7 @@ pub(super) fn classify_public_support_route(
|
||||
| "/api/users/me/available-models"
|
||||
| "/api/users/me/endpoint-status"
|
||||
| "/api/users/me/preferences"
|
||||
| "/api/users/me/referral"
|
||||
| "/api/users/me/model-capabilities"
|
||||
)
|
||||
{
|
||||
@@ -477,6 +492,7 @@ pub(super) fn classify_public_support_route(
|
||||
"/api/users/me/available-models" => "available_models",
|
||||
"/api/users/me/endpoint-status" => "endpoint_status",
|
||||
"/api/users/me/preferences" => "preferences",
|
||||
"/api/users/me/referral" => "referral",
|
||||
"/api/users/me/model-capabilities" => "model_capabilities",
|
||||
_ => "detail",
|
||||
};
|
||||
|
||||
92
apps/aether-gateway/src/control/tests/admin_routing.rs
Normal file
92
apps/aether-gateway/src/control/tests/admin_routing.rs
Normal file
@@ -0,0 +1,92 @@
|
||||
use http::Uri;
|
||||
|
||||
use crate::handlers::shared::local_proxy_route_requires_buffered_body;
|
||||
|
||||
use super::{classify_control_route, headers, GatewayPublicRequestContext};
|
||||
|
||||
#[test]
|
||||
fn classifies_admin_routing_group_routes_as_admin_proxy_route() {
|
||||
let headers = headers(&[]);
|
||||
|
||||
let list_uri: Uri = "/api/admin/routing/groups"
|
||||
.parse()
|
||||
.expect("uri should parse");
|
||||
let list = classify_control_route(&http::Method::GET, &list_uri, &headers)
|
||||
.expect("route should classify");
|
||||
assert_eq!(list.route_class.as_deref(), Some("admin_proxy"));
|
||||
assert_eq!(
|
||||
list.route_family.as_deref(),
|
||||
Some("routing_profiles_manage")
|
||||
);
|
||||
assert_eq!(list.route_kind.as_deref(), Some("list_groups"));
|
||||
assert_eq!(
|
||||
list.auth_endpoint_signature.as_deref(),
|
||||
Some("admin:routing_profiles")
|
||||
);
|
||||
|
||||
let create_uri: Uri = "/api/admin/routing/groups"
|
||||
.parse()
|
||||
.expect("uri should parse");
|
||||
let create = classify_control_route(&http::Method::POST, &create_uri, &headers)
|
||||
.expect("route should classify");
|
||||
assert_eq!(
|
||||
create.route_family.as_deref(),
|
||||
Some("routing_profiles_manage")
|
||||
);
|
||||
assert_eq!(create.route_kind.as_deref(), Some("create_group"));
|
||||
|
||||
let update_uri: Uri = "/api/admin/routing/groups/group-1"
|
||||
.parse()
|
||||
.expect("uri should parse");
|
||||
let update = classify_control_route(&http::Method::PATCH, &update_uri, &headers)
|
||||
.expect("route should classify");
|
||||
assert_eq!(
|
||||
update.route_family.as_deref(),
|
||||
Some("routing_profiles_manage")
|
||||
);
|
||||
assert_eq!(update.route_kind.as_deref(), Some("update_group"));
|
||||
|
||||
let dry_run_uri: Uri = "/api/admin/routing/groups/group-1/dry-run"
|
||||
.parse()
|
||||
.expect("uri should parse");
|
||||
let dry_run = classify_control_route(&http::Method::POST, &dry_run_uri, &headers)
|
||||
.expect("route should classify");
|
||||
assert_eq!(
|
||||
dry_run.route_family.as_deref(),
|
||||
Some("routing_profiles_manage")
|
||||
);
|
||||
assert_eq!(dry_run.route_kind.as_deref(), Some("dry_run_group"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn admin_routing_write_routes_buffer_request_body() {
|
||||
let headers = headers(&[]);
|
||||
let routes = [
|
||||
(http::Method::POST, "/api/admin/routing/groups"),
|
||||
(http::Method::PATCH, "/api/admin/routing/groups/group-1"),
|
||||
(
|
||||
http::Method::POST,
|
||||
"/api/admin/routing/groups/group-1/dry-run",
|
||||
),
|
||||
(http::Method::POST, "/api/admin/routing/bindings"),
|
||||
(http::Method::PATCH, "/api/admin/routing/bindings/binding-1"),
|
||||
];
|
||||
|
||||
for (method, path) in routes {
|
||||
let uri: Uri = path.parse().expect("uri should parse");
|
||||
let decision =
|
||||
classify_control_route(&method, &uri, &headers).expect("route should classify");
|
||||
let context = GatewayPublicRequestContext::from_request_parts(
|
||||
"trace-routing-write",
|
||||
&method,
|
||||
&uri,
|
||||
&headers,
|
||||
Some(decision),
|
||||
);
|
||||
|
||||
assert!(
|
||||
local_proxy_route_requires_buffered_body(&context),
|
||||
"{method} {path} should buffer request body"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -80,6 +80,28 @@ fn classifies_openai_chat_and_responses_separately_from_embedding() {
|
||||
assert_ne!(responses.route_kind.as_deref(), Some("embedding"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_openai_image_generation_and_edit_but_not_variation() {
|
||||
let headers = headers(&[("authorization", "Bearer sk-test")]);
|
||||
|
||||
for path in ["/v1/images/generations", "/v1/images/edits"] {
|
||||
let uri: Uri = path.parse().expect("uri should parse");
|
||||
let decision = classify_control_route(&http::Method::POST, &uri, &headers)
|
||||
.expect("image route should classify");
|
||||
|
||||
assert_eq!(decision.route_family.as_deref(), Some("openai"));
|
||||
assert_eq!(decision.route_kind.as_deref(), Some("image"));
|
||||
assert_eq!(
|
||||
decision.auth_endpoint_signature.as_deref(),
|
||||
Some("openai:image")
|
||||
);
|
||||
assert!(decision.is_execution_runtime_candidate());
|
||||
}
|
||||
|
||||
let variation_uri: Uri = "/v1/images/variations".parse().expect("uri should parse");
|
||||
assert!(classify_control_route(&http::Method::POST, &variation_uri, &headers).is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn classifies_models_list_as_claude_when_headers_match() {
|
||||
let headers = headers(&[
|
||||
|
||||
@@ -86,6 +86,7 @@ mod admin_provider_query;
|
||||
mod admin_provider_strategy;
|
||||
mod admin_providers_models;
|
||||
mod admin_proxy_nodes;
|
||||
mod admin_routing;
|
||||
mod admin_security;
|
||||
mod admin_stats;
|
||||
mod admin_usage;
|
||||
|
||||
@@ -1,15 +1,15 @@
|
||||
use super::{
|
||||
AuthApiKeyLookupKey, CreateManagementTokenRecord, DataLayerError, GatewayAuthApiKeySnapshot,
|
||||
GatewayDataState, ManagementTokenListQuery, ProxyNodeHeartbeatMutation,
|
||||
ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, ProxyNodeRegistrationMutation,
|
||||
ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation, ProxyNodeTunnelStatusMutation,
|
||||
RegenerateManagementTokenSecret, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot,
|
||||
StoredLdapModuleConfig, StoredManagementToken, StoredManagementTokenListPage,
|
||||
StoredManagementTokenWithUser, StoredOAuthProviderConfig, StoredOAuthProviderModuleConfig,
|
||||
StoredProxyFleetMetricsBucket, StoredProxyNode, StoredProxyNodeEvent,
|
||||
StoredProxyNodeMetricsBucket, StoredUserAuthRecord, StoredUserOAuthLinkSummary,
|
||||
StoredUserPreferenceRecord, StoredUserSessionRecord, StoredWalletSnapshot,
|
||||
UpdateManagementTokenRecord, UpsertOAuthProviderConfigRecord,
|
||||
GatewayDataState, ManagementTokenCounterDelta, ManagementTokenListQuery, ProxyNodeCounterDelta,
|
||||
ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation,
|
||||
ProxyNodeRegistrationMutation, ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation,
|
||||
ProxyNodeTunnelStatusMutation, RegenerateManagementTokenSecret, StoredAuthApiKeyExportRecord,
|
||||
StoredAuthApiKeySnapshot, StoredLdapModuleConfig, StoredManagementToken,
|
||||
StoredManagementTokenListPage, StoredManagementTokenWithUser, StoredOAuthProviderConfig,
|
||||
StoredOAuthProviderModuleConfig, StoredProxyFleetMetricsBucket, StoredProxyNode,
|
||||
StoredProxyNodeEvent, StoredProxyNodeMetricsBucket, StoredUserAuthRecord,
|
||||
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
|
||||
StoredWalletSnapshot, UpdateManagementTokenRecord, UpsertOAuthProviderConfigRecord,
|
||||
};
|
||||
use crate::LocalMutationOutcome;
|
||||
use aether_data::repository::auth::{
|
||||
@@ -1117,6 +1117,20 @@ impl GatewayDataState {
|
||||
token_id: &str,
|
||||
last_used_ip: Option<&str>,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
if let Some(repository) = &self.usage_writer {
|
||||
let enqueued = repository
|
||||
.enqueue_management_token_counter_delta(ManagementTokenCounterDelta {
|
||||
token_id: token_id.to_string(),
|
||||
usage_count_delta: 1,
|
||||
last_used_at_unix_secs: Some(chrono::Utc::now().timestamp().max(0) as u64),
|
||||
last_used_ip: last_used_ip.map(ToOwned::to_owned),
|
||||
})
|
||||
.await?;
|
||||
if enqueued {
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
|
||||
match &self.management_token_writer {
|
||||
Some(repository) => {
|
||||
repository
|
||||
@@ -1278,6 +1292,21 @@ impl GatewayDataState {
|
||||
&self,
|
||||
mutation: &ProxyNodeTrafficMutation,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
if let Some(repository) = &self.usage_writer {
|
||||
let enqueued = repository
|
||||
.enqueue_proxy_node_counter_delta(ProxyNodeCounterDelta {
|
||||
node_id: mutation.node_id.clone(),
|
||||
total_requests_delta: mutation.total_requests_delta,
|
||||
failed_requests_delta: mutation.failed_requests_delta,
|
||||
dns_failures_delta: mutation.dns_failures_delta,
|
||||
stream_errors_delta: mutation.stream_errors_delta,
|
||||
})
|
||||
.await?;
|
||||
if enqueued {
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
|
||||
match &self.proxy_node_writer {
|
||||
Some(repository) => repository.record_traffic(mutation).await,
|
||||
None => Ok(false),
|
||||
@@ -1925,14 +1954,38 @@ fn resolve_effective_list_policy(
|
||||
&aether_data::repository::users::StoredUserGroup,
|
||||
) -> (&str, Option<Vec<String>>),
|
||||
) -> Option<Vec<String>> {
|
||||
let group_policy = groups.iter().fold(None, |effective, group| {
|
||||
let (mode, values) = group_field(group);
|
||||
intersect_list_policies(effective, list_restriction_from_mode(mode, values))
|
||||
});
|
||||
let group_policy = union_group_list_policies(groups, group_field);
|
||||
let user_policy = list_restriction_from_mode(user_mode, user_values);
|
||||
intersect_list_policies(group_policy, user_policy)
|
||||
}
|
||||
|
||||
fn union_group_list_policies(
|
||||
groups: &[aether_data::repository::users::StoredUserGroup],
|
||||
group_field: impl Fn(
|
||||
&aether_data::repository::users::StoredUserGroup,
|
||||
) -> (&str, Option<Vec<String>>),
|
||||
) -> Option<Vec<String>> {
|
||||
let mut saw_restrictive_group = false;
|
||||
let mut values = std::collections::BTreeSet::new();
|
||||
|
||||
for group in groups {
|
||||
let (mode, group_values) = group_field(group);
|
||||
match mode {
|
||||
"unrestricted" => return None,
|
||||
"specific" => {
|
||||
saw_restrictive_group = true;
|
||||
values.extend(group_values.unwrap_or_default());
|
||||
}
|
||||
"deny_all" => {
|
||||
saw_restrictive_group = true;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
saw_restrictive_group.then(|| values.into_iter().collect())
|
||||
}
|
||||
|
||||
fn list_restriction_from_mode(mode: &str, values: Option<Vec<String>>) -> Option<Vec<String>> {
|
||||
match mode {
|
||||
"specific" => Some(values.unwrap_or_default()),
|
||||
@@ -2148,7 +2201,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn list_policy_intersects_group_and_user_restrictions() {
|
||||
fn list_policy_intersects_unrestricted_group_union_with_user_restriction() {
|
||||
let groups = vec![
|
||||
sample_group("default", 0, None, "unrestricted", None, "system"),
|
||||
sample_group(
|
||||
@@ -2168,11 +2221,14 @@ mod tests {
|
||||
|group| (&group.allowed_models_mode, group.allowed_models.clone()),
|
||||
);
|
||||
|
||||
assert_eq!(policy, Some(vec!["gpt-4.1".to_string()]));
|
||||
assert_eq!(
|
||||
policy,
|
||||
Some(vec!["gpt-4.1".to_string(), "gemini-2.5-pro".to_string()])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn list_policy_intersects_multiple_group_restrictions() {
|
||||
fn list_policy_unions_multiple_group_restrictions_legacy_case() {
|
||||
let groups = vec![
|
||||
sample_group(
|
||||
"team-a",
|
||||
@@ -2196,7 +2252,91 @@ mod tests {
|
||||
(&group.allowed_models_mode, group.allowed_models.clone())
|
||||
});
|
||||
|
||||
assert_eq!(policy, Some(vec!["gpt-4.1".to_string()]));
|
||||
assert_eq!(
|
||||
policy,
|
||||
Some(vec![
|
||||
"gemini-2.5-pro".to_string(),
|
||||
"gpt-4.1".to_string(),
|
||||
"gpt-5".to_string()
|
||||
])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn list_policy_unions_multiple_group_restrictions() {
|
||||
let groups = vec![
|
||||
sample_group(
|
||||
"team-a",
|
||||
10,
|
||||
Some(vec!["gpt-5", "gpt-4.1"]),
|
||||
"specific",
|
||||
None,
|
||||
"system",
|
||||
),
|
||||
sample_group(
|
||||
"team-b",
|
||||
20,
|
||||
Some(vec!["gpt-4.1", "gemini-2.5-pro"]),
|
||||
"specific",
|
||||
None,
|
||||
"system",
|
||||
),
|
||||
];
|
||||
|
||||
let policy = resolve_effective_list_policy(None, "unrestricted", &groups, |group| {
|
||||
(&group.allowed_models_mode, group.allowed_models.clone())
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
policy,
|
||||
Some(vec![
|
||||
"gemini-2.5-pro".to_string(),
|
||||
"gpt-4.1".to_string(),
|
||||
"gpt-5".to_string()
|
||||
])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unrestricted_group_makes_group_policy_unrestricted() {
|
||||
let groups = vec![
|
||||
sample_group(
|
||||
"restricted",
|
||||
10,
|
||||
Some(vec!["gpt-5"]),
|
||||
"specific",
|
||||
None,
|
||||
"system",
|
||||
),
|
||||
sample_group("unrestricted", 20, None, "unrestricted", None, "system"),
|
||||
];
|
||||
|
||||
let policy = resolve_effective_list_policy(None, "unrestricted", &groups, |group| {
|
||||
(&group.allowed_models_mode, group.allowed_models.clone())
|
||||
});
|
||||
|
||||
assert_eq!(policy, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn deny_all_group_does_not_remove_other_group_grants() {
|
||||
let groups = vec![
|
||||
sample_group("deny", 10, None, "deny_all", None, "system"),
|
||||
sample_group(
|
||||
"restricted",
|
||||
20,
|
||||
Some(vec!["gpt-5"]),
|
||||
"specific",
|
||||
None,
|
||||
"system",
|
||||
),
|
||||
];
|
||||
|
||||
let policy = resolve_effective_list_policy(None, "unrestricted", &groups, |group| {
|
||||
(&group.allowed_models_mode, group.allowed_models.clone())
|
||||
});
|
||||
|
||||
assert_eq!(policy, Some(vec!["gpt-5".to_string()]));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
401
apps/aether-gateway/src/data/state/candidate_cache.rs
Normal file
401
apps/aether-gateway/src/data/state/candidate_cache.rs
Normal file
@@ -0,0 +1,401 @@
|
||||
use std::collections::HashSet;
|
||||
use std::future::Future;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::sync::{Arc, Mutex};
|
||||
use std::time::Duration;
|
||||
|
||||
use aether_cache::ExpiringMap;
|
||||
use aether_data::DataLayerError;
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
|
||||
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
|
||||
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use tokio::sync::Notify;
|
||||
|
||||
const CANDIDATE_SELECTION_CACHE_TTL: Duration = Duration::from_secs(5);
|
||||
const CANDIDATE_SELECTION_CACHE_MAX_ENTRIES: usize = 4096;
|
||||
|
||||
pub(super) struct CachedMinimalCandidateSelectionReadRepository {
|
||||
inner: Arc<dyn MinimalCandidateSelectionReadRepository>,
|
||||
entries: ExpiringMap<CandidateSelectionCacheKey, Vec<StoredMinimalCandidateSelectionRow>>,
|
||||
inflight: Mutex<HashSet<CandidateSelectionCacheKey>>,
|
||||
inflight_notify: Notify,
|
||||
epoch: AtomicU64,
|
||||
}
|
||||
|
||||
impl CachedMinimalCandidateSelectionReadRepository {
|
||||
pub(super) fn new(inner: Arc<dyn MinimalCandidateSelectionReadRepository>) -> Self {
|
||||
Self {
|
||||
inner,
|
||||
entries: ExpiringMap::new(),
|
||||
inflight: Mutex::new(HashSet::new()),
|
||||
inflight_notify: Notify::new(),
|
||||
epoch: AtomicU64::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
async fn get_or_load<F, Fut>(
|
||||
&self,
|
||||
key: CandidateSelectionCacheKey,
|
||||
load: F,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>
|
||||
where
|
||||
F: Fn() -> Fut,
|
||||
Fut: Future<Output = Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>>,
|
||||
{
|
||||
if let Some(rows) = self.entries.get_fresh(&key, CANDIDATE_SELECTION_CACHE_TTL) {
|
||||
return Ok(rows);
|
||||
}
|
||||
|
||||
loop {
|
||||
let notified = self.inflight_notify.notified();
|
||||
match self.register_inflight(&key) {
|
||||
InflightRegistration::Bypass => return load().await,
|
||||
InflightRegistration::Follower => {
|
||||
notified.await;
|
||||
if let Some(rows) = self.entries.get_fresh(&key, CANDIDATE_SELECTION_CACHE_TTL)
|
||||
{
|
||||
return Ok(rows);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
InflightRegistration::Leader => {}
|
||||
}
|
||||
|
||||
let load_epoch = self.epoch.load(Ordering::Acquire);
|
||||
let result = load().await;
|
||||
if let Ok(rows) = &result {
|
||||
if load_epoch == self.epoch.load(Ordering::Acquire) {
|
||||
self.entries.insert(
|
||||
key.clone(),
|
||||
rows.clone(),
|
||||
CANDIDATE_SELECTION_CACHE_TTL,
|
||||
CANDIDATE_SELECTION_CACHE_MAX_ENTRIES,
|
||||
);
|
||||
}
|
||||
}
|
||||
self.finish_inflight(&key);
|
||||
return result;
|
||||
}
|
||||
}
|
||||
|
||||
fn register_inflight(&self, key: &CandidateSelectionCacheKey) -> InflightRegistration {
|
||||
match self.inflight.lock() {
|
||||
Ok(mut inflight) => {
|
||||
if inflight.insert(key.clone()) {
|
||||
InflightRegistration::Leader
|
||||
} else {
|
||||
InflightRegistration::Follower
|
||||
}
|
||||
}
|
||||
Err(_) => InflightRegistration::Bypass,
|
||||
}
|
||||
}
|
||||
|
||||
fn finish_inflight(&self, key: &CandidateSelectionCacheKey) {
|
||||
if let Ok(mut inflight) = self.inflight.lock() {
|
||||
inflight.remove(key);
|
||||
}
|
||||
self.inflight_notify.notify_waiters();
|
||||
}
|
||||
|
||||
fn clear(&self) {
|
||||
self.epoch.fetch_add(1, Ordering::AcqRel);
|
||||
self.entries.clear();
|
||||
}
|
||||
}
|
||||
|
||||
enum InflightRegistration {
|
||||
Leader,
|
||||
Follower,
|
||||
Bypass,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl MinimalCandidateSelectionReadRepository for CachedMinimalCandidateSelectionReadRepository {
|
||||
fn clear_local_cache(&self) {
|
||||
self.clear();
|
||||
self.inner.clear_local_cache();
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let key = CandidateSelectionCacheKey::ApiFormat {
|
||||
api_format: normalize_api_format_key(api_format),
|
||||
};
|
||||
self.get_or_load(key, || self.inner.list_for_exact_api_format(api_format))
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_global_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
global_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let key = CandidateSelectionCacheKey::ApiFormatAndGlobalModel {
|
||||
api_format: normalize_api_format_key(api_format),
|
||||
global_model_name: global_model_name.to_string(),
|
||||
};
|
||||
self.get_or_load(key, || {
|
||||
self.inner
|
||||
.list_for_exact_api_format_and_global_model(api_format, global_model_name)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let key = CandidateSelectionCacheKey::ApiFormatAndRequestedModel {
|
||||
api_format: normalize_api_format_key(api_format),
|
||||
requested_model_name: requested_model_name.to_string(),
|
||||
};
|
||||
self.get_or_load(key, || {
|
||||
self.inner
|
||||
.list_for_exact_api_format_and_requested_model(api_format, requested_model_name)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model_page(
|
||||
&self,
|
||||
query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let key = CandidateSelectionCacheKey::RequestedModelPage {
|
||||
api_format: normalize_api_format_key(&query.api_format),
|
||||
requested_model_name: query.requested_model_name.clone(),
|
||||
offset: query.offset,
|
||||
limit: query.limit,
|
||||
};
|
||||
self.get_or_load(key, || {
|
||||
self.inner
|
||||
.list_for_exact_api_format_and_requested_model_page(query)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group(
|
||||
&self,
|
||||
query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let key = CandidateSelectionCacheKey::PoolKeyRowsForGroup {
|
||||
api_format: normalize_api_format_key(&query.api_format),
|
||||
provider_id: query.provider_id.clone(),
|
||||
endpoint_id: query.endpoint_id.clone(),
|
||||
model_id: query.model_id.clone(),
|
||||
selected_provider_model_name: query.selected_provider_model_name.clone(),
|
||||
order: CandidateSelectionPoolOrderKey::from(&query.order),
|
||||
offset: query.offset,
|
||||
limit: query.limit,
|
||||
};
|
||||
self.get_or_load(key, || self.inner.list_pool_key_rows_for_group(query))
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group_key_ids(
|
||||
&self,
|
||||
query: &StoredPoolKeyCandidateRowsByKeyIdsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let key = CandidateSelectionCacheKey::PoolKeyRowsForGroupKeyIds {
|
||||
api_format: normalize_api_format_key(&query.api_format),
|
||||
provider_id: query.provider_id.clone(),
|
||||
endpoint_id: query.endpoint_id.clone(),
|
||||
model_id: query.model_id.clone(),
|
||||
selected_provider_model_name: query.selected_provider_model_name.clone(),
|
||||
key_ids: query.key_ids.clone(),
|
||||
};
|
||||
self.get_or_load(key, || {
|
||||
self.inner.list_pool_key_rows_for_group_key_ids(query)
|
||||
})
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
|
||||
enum CandidateSelectionCacheKey {
|
||||
ApiFormat {
|
||||
api_format: String,
|
||||
},
|
||||
ApiFormatAndGlobalModel {
|
||||
api_format: String,
|
||||
global_model_name: String,
|
||||
},
|
||||
ApiFormatAndRequestedModel {
|
||||
api_format: String,
|
||||
requested_model_name: String,
|
||||
},
|
||||
RequestedModelPage {
|
||||
api_format: String,
|
||||
requested_model_name: String,
|
||||
offset: u32,
|
||||
limit: u32,
|
||||
},
|
||||
PoolKeyRowsForGroup {
|
||||
api_format: String,
|
||||
provider_id: String,
|
||||
endpoint_id: String,
|
||||
model_id: String,
|
||||
selected_provider_model_name: String,
|
||||
order: CandidateSelectionPoolOrderKey,
|
||||
offset: u32,
|
||||
limit: u32,
|
||||
},
|
||||
PoolKeyRowsForGroupKeyIds {
|
||||
api_format: String,
|
||||
provider_id: String,
|
||||
endpoint_id: String,
|
||||
model_id: String,
|
||||
selected_provider_model_name: String,
|
||||
key_ids: Vec<String>,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
|
||||
enum CandidateSelectionPoolOrderKey {
|
||||
InternalPriority,
|
||||
Lru,
|
||||
CacheAffinity,
|
||||
SingleAccount,
|
||||
LoadBalance { seed: String },
|
||||
}
|
||||
|
||||
impl From<&StoredPoolKeyCandidateOrder> for CandidateSelectionPoolOrderKey {
|
||||
fn from(order: &StoredPoolKeyCandidateOrder) -> Self {
|
||||
match order {
|
||||
StoredPoolKeyCandidateOrder::InternalPriority => Self::InternalPriority,
|
||||
StoredPoolKeyCandidateOrder::Lru => Self::Lru,
|
||||
StoredPoolKeyCandidateOrder::CacheAffinity => Self::CacheAffinity,
|
||||
StoredPoolKeyCandidateOrder::SingleAccount => Self::SingleAccount,
|
||||
StoredPoolKeyCandidateOrder::LoadBalance { seed } => {
|
||||
Self::LoadBalance { seed: seed.clone() }
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_api_format_key(api_format: &str) -> String {
|
||||
crate::ai_serving::normalize_api_format_alias(api_format.trim())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::sync::atomic::AtomicUsize;
|
||||
|
||||
struct StubCandidateSelectionRepository {
|
||||
calls: AtomicUsize,
|
||||
delay: Duration,
|
||||
}
|
||||
|
||||
impl StubCandidateSelectionRepository {
|
||||
fn new(delay: Duration) -> Self {
|
||||
Self {
|
||||
calls: AtomicUsize::new(0),
|
||||
delay,
|
||||
}
|
||||
}
|
||||
|
||||
fn calls(&self) -> usize {
|
||||
self.calls.load(Ordering::SeqCst)
|
||||
}
|
||||
|
||||
async fn load(&self) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.calls.fetch_add(1, Ordering::SeqCst);
|
||||
if !self.delay.is_zero() {
|
||||
tokio::time::sleep(self.delay).await;
|
||||
}
|
||||
Ok(Vec::new())
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl MinimalCandidateSelectionReadRepository for StubCandidateSelectionRepository {
|
||||
async fn list_for_exact_api_format(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.load().await
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_global_model(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
_global_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.load().await
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
_requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.load().await
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model_page(
|
||||
&self,
|
||||
_query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.load().await
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group(
|
||||
&self,
|
||||
_query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.load().await
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group_key_ids(
|
||||
&self,
|
||||
_query: &StoredPoolKeyCandidateRowsByKeyIdsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.load().await
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn candidate_selection_cache_coalesces_concurrent_loads() {
|
||||
let inner = Arc::new(StubCandidateSelectionRepository::new(
|
||||
Duration::from_millis(25),
|
||||
));
|
||||
let cache = Arc::new(CachedMinimalCandidateSelectionReadRepository::new(
|
||||
inner.clone(),
|
||||
));
|
||||
let mut tasks = Vec::new();
|
||||
|
||||
for _ in 0..16 {
|
||||
let cache = cache.clone();
|
||||
tasks.push(tokio::spawn(async move {
|
||||
cache.list_for_exact_api_format("openai").await.unwrap();
|
||||
}));
|
||||
}
|
||||
|
||||
for task in tasks {
|
||||
task.await.unwrap();
|
||||
}
|
||||
|
||||
assert_eq!(inner.calls(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn candidate_selection_cache_clear_invalidates_entries() {
|
||||
let inner = Arc::new(StubCandidateSelectionRepository::new(Duration::ZERO));
|
||||
let cache = CachedMinimalCandidateSelectionReadRepository::new(inner.clone());
|
||||
|
||||
cache.list_for_exact_api_format("openai").await.unwrap();
|
||||
cache.list_for_exact_api_format("openai").await.unwrap();
|
||||
assert_eq!(inner.calls(), 1);
|
||||
|
||||
cache.clear_local_cache();
|
||||
cache.list_for_exact_api_format("openai").await.unwrap();
|
||||
assert_eq!(inner.calls(), 2);
|
||||
}
|
||||
}
|
||||
@@ -1,10 +1,10 @@
|
||||
use super::{
|
||||
DataLayerError, GatewayDataState, GeminiFileMappingListQuery, GeminiFileMappingStats,
|
||||
ProviderCatalogKeyListQuery, PublicHealthStatusCount, PublicHealthTimelineBucket,
|
||||
StoredGeminiFileMapping, StoredGeminiFileMappingListPage, StoredProviderCatalogEndpoint,
|
||||
StoredProviderCatalogKey, StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats,
|
||||
StoredProviderCatalogProvider, StoredRequestCandidate, UpsertGeminiFileMappingRecord,
|
||||
UpsertRequestCandidateRecord,
|
||||
ApiKeyLastUsedDelta, DataLayerError, GatewayDataState, GeminiFileMappingListQuery,
|
||||
GeminiFileMappingStats, ProviderCatalogKeyListQuery, PublicHealthStatusCount,
|
||||
PublicHealthTimelineBucket, StoredGeminiFileMapping, StoredGeminiFileMappingListPage,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogKeyPage,
|
||||
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider, StoredRequestCandidate,
|
||||
UpsertGeminiFileMappingRecord, UpsertRequestCandidateRecord,
|
||||
};
|
||||
|
||||
impl GatewayDataState {
|
||||
@@ -121,6 +121,18 @@ impl GatewayDataState {
|
||||
&self,
|
||||
api_key_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
if let Some(repository) = &self.usage_writer {
|
||||
let enqueued = repository
|
||||
.enqueue_api_key_last_used_delta(ApiKeyLastUsedDelta {
|
||||
api_key_id: api_key_id.to_string(),
|
||||
last_used_at_unix_secs: chrono::Utc::now().timestamp().max(0) as u64,
|
||||
})
|
||||
.await?;
|
||||
if enqueued {
|
||||
return Ok(true);
|
||||
}
|
||||
}
|
||||
|
||||
match &self.auth_api_key_writer {
|
||||
Some(repository) => repository.touch_last_used_at(api_key_id).await,
|
||||
None => Ok(false),
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
use aether_data::{DataBackends, DataLayerError, DatabaseDriver};
|
||||
use aether_data_contracts::repository::candidate_selection::MinimalCandidateSelectionReadRepository;
|
||||
use aether_runtime_state::RuntimeQueueStore;
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -49,6 +50,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -82,7 +85,17 @@ impl GatewayDataState {
|
||||
let gemini_file_mapping_reader = backends.read().gemini_file_mappings();
|
||||
let global_model_reader = backends.read().global_models();
|
||||
let global_model_writer = backends.write().global_models();
|
||||
let minimal_candidate_selection_reader = backends.read().minimal_candidate_selection();
|
||||
let minimal_candidate_selection_reader =
|
||||
backends
|
||||
.read()
|
||||
.minimal_candidate_selection()
|
||||
.map(|repository| {
|
||||
Arc::new(
|
||||
super::candidate_cache::CachedMinimalCandidateSelectionReadRepository::new(
|
||||
repository,
|
||||
),
|
||||
) as Arc<dyn MinimalCandidateSelectionReadRepository>
|
||||
});
|
||||
let request_candidate_reader = backends.read().request_candidates();
|
||||
let request_candidate_writer = backends.write().request_candidates();
|
||||
let gemini_file_mapping_writer = backends.write().gemini_file_mappings();
|
||||
@@ -92,6 +105,8 @@ impl GatewayDataState {
|
||||
let pool_score_writer = backends.write().pool_scores();
|
||||
let provider_quota_reader = backends.read().provider_quotas();
|
||||
let provider_quota_writer = backends.write().provider_quotas();
|
||||
let routing_group_reader = backends.read().routing_groups();
|
||||
let routing_group_writer = backends.write().routing_groups();
|
||||
let usage_reader = backends.read().usage();
|
||||
let usage_writer = backends.write().usage();
|
||||
let user_reader = backends.read().users();
|
||||
@@ -133,6 +148,8 @@ impl GatewayDataState {
|
||||
pool_score_writer,
|
||||
provider_quota_reader,
|
||||
provider_quota_writer,
|
||||
routing_group_reader,
|
||||
routing_group_writer,
|
||||
usage_reader,
|
||||
usage_writer,
|
||||
user_reader,
|
||||
@@ -253,6 +270,12 @@ impl GatewayDataState {
|
||||
self.minimal_candidate_selection_reader.is_some()
|
||||
}
|
||||
|
||||
pub(crate) fn clear_minimal_candidate_selection_cache(&self) {
|
||||
if let Some(repository) = &self.minimal_candidate_selection_reader {
|
||||
repository.clear_local_cache();
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn has_request_candidate_reader(&self) -> bool {
|
||||
self.request_candidate_reader.is_some()
|
||||
}
|
||||
@@ -261,6 +284,14 @@ impl GatewayDataState {
|
||||
self.request_candidate_writer.is_some()
|
||||
}
|
||||
|
||||
pub(crate) fn has_routing_group_reader(&self) -> bool {
|
||||
self.routing_group_reader.is_some()
|
||||
}
|
||||
|
||||
pub(crate) fn has_routing_group_writer(&self) -> bool {
|
||||
self.routing_group_writer.is_some()
|
||||
}
|
||||
|
||||
pub(crate) fn has_provider_catalog_reader(&self) -> bool {
|
||||
self.provider_catalog_reader.is_some()
|
||||
}
|
||||
@@ -315,6 +346,10 @@ impl GatewayDataState {
|
||||
self.usage_writer.is_some()
|
||||
}
|
||||
|
||||
pub(crate) fn has_usage_counter_flush_backend(&self) -> bool {
|
||||
self.has_usage_writer() && self.database_driver() == Some(DatabaseDriver::Postgres)
|
||||
}
|
||||
|
||||
pub(crate) fn has_usage_worker_queue(&self) -> bool {
|
||||
self.usage_worker_queue.is_some()
|
||||
}
|
||||
|
||||
@@ -12,7 +12,9 @@ use aether_data_contracts::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use aether_data_contracts::repository::settlement::{StoredUsageSettlement, UsageSettlementInput};
|
||||
use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UpsertUsageRecord};
|
||||
use aether_data_contracts::repository::usage::{
|
||||
ProxyNodeCounterDelta, StoredRequestUsageAudit, UpsertUsageRecord,
|
||||
};
|
||||
use aether_data_contracts::repository::video_tasks::{StoredVideoTask, VideoTaskLookupKey};
|
||||
use aether_runtime_state::RuntimeQueueStore;
|
||||
use aether_usage_runtime::{
|
||||
@@ -284,6 +286,21 @@ impl aether_usage_runtime::ManualProxyNodeCounter for GatewayDataState {
|
||||
failed_delta: i64,
|
||||
latency_ms: Option<i64>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
if let Some(repository) = &self.usage_writer {
|
||||
let enqueued = repository
|
||||
.enqueue_proxy_node_counter_delta(ProxyNodeCounterDelta {
|
||||
node_id: node_id.to_string(),
|
||||
total_requests_delta: total_delta,
|
||||
failed_requests_delta: failed_delta,
|
||||
dns_failures_delta: 0,
|
||||
stream_errors_delta: 0,
|
||||
})
|
||||
.await?;
|
||||
if enqueued {
|
||||
return Ok(());
|
||||
}
|
||||
}
|
||||
|
||||
match &self.proxy_node_writer {
|
||||
Some(repository) => {
|
||||
repository
|
||||
|
||||
@@ -125,12 +125,16 @@ use aether_data_contracts::repository::provider_catalog::{
|
||||
use aether_data_contracts::repository::quota::{
|
||||
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
|
||||
};
|
||||
use aether_data_contracts::repository::routing_profiles::{
|
||||
RoutingGroupReadRepository, RoutingGroupWriteRepository,
|
||||
};
|
||||
use aether_data_contracts::repository::settlement::{
|
||||
SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput,
|
||||
};
|
||||
use aether_data_contracts::repository::usage::{
|
||||
PendingUsageCleanupSummary, StoredProviderUsageSummary, StoredRequestUsageAudit,
|
||||
UpsertUsageRecord, UsageReadRepository, UsageWriteRepository,
|
||||
ApiKeyLastUsedDelta, ManagementTokenCounterDelta, PendingUsageCleanupSummary,
|
||||
ProxyNodeCounterDelta, StoredProviderUsageSummary, StoredRequestUsageAudit, UpsertUsageRecord,
|
||||
UsageReadRepository, UsageWriteRepository,
|
||||
};
|
||||
use aether_data_contracts::repository::video_tasks::{
|
||||
StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount,
|
||||
@@ -138,6 +142,12 @@ use aether_data_contracts::repository::video_tasks::{
|
||||
};
|
||||
use aether_runtime_state::RuntimeQueueStore;
|
||||
|
||||
pub(crate) use self::referrals::{
|
||||
ReferralAdminStats, ReferralMutationStatus, ReferralRelationshipListQuery,
|
||||
ReferralRelationshipRecord, ReferralRewardConfig, ReferralRewardListQuery,
|
||||
ReferralRewardRecord, ReferralUserDashboard,
|
||||
};
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub(crate) struct GatewayDataState {
|
||||
config: GatewayDataConfig,
|
||||
@@ -170,6 +180,8 @@ pub(crate) struct GatewayDataState {
|
||||
pool_score_writer: Option<Arc<dyn PoolMemberScoreWriteRepository>>,
|
||||
provider_quota_reader: Option<Arc<dyn ProviderQuotaReadRepository>>,
|
||||
provider_quota_writer: Option<Arc<dyn ProviderQuotaWriteRepository>>,
|
||||
routing_group_reader: Option<Arc<dyn RoutingGroupReadRepository>>,
|
||||
routing_group_writer: Option<Arc<dyn RoutingGroupWriteRepository>>,
|
||||
usage_reader: Option<Arc<dyn UsageReadRepository>>,
|
||||
usage_writer: Option<Arc<dyn UsageWriteRepository>>,
|
||||
user_reader: Option<Arc<dyn UserReadRepository>>,
|
||||
@@ -279,6 +291,14 @@ impl fmt::Debug for GatewayDataState {
|
||||
"has_provider_quota_writer",
|
||||
&self.provider_quota_writer.is_some(),
|
||||
)
|
||||
.field(
|
||||
"has_routing_group_reader",
|
||||
&self.routing_group_reader.is_some(),
|
||||
)
|
||||
.field(
|
||||
"has_routing_group_writer",
|
||||
&self.routing_group_writer.is_some(),
|
||||
)
|
||||
.field("has_usage_reader", &self.usage_reader.is_some())
|
||||
.field("has_usage_writer", &self.usage_writer.is_some())
|
||||
.field("has_user_preferences", &self.user_preferences.is_some())
|
||||
@@ -297,11 +317,14 @@ impl fmt::Debug for GatewayDataState {
|
||||
}
|
||||
|
||||
mod auth;
|
||||
mod candidate_cache;
|
||||
mod catalog;
|
||||
mod core;
|
||||
mod integrations;
|
||||
mod models;
|
||||
mod pool_scores;
|
||||
mod referrals;
|
||||
mod routing_profiles;
|
||||
mod runtime;
|
||||
#[cfg(test)]
|
||||
mod testing;
|
||||
|
||||
2652
apps/aether-gateway/src/data/state/referrals.rs
Normal file
2652
apps/aether-gateway/src/data/state/referrals.rs
Normal file
File diff suppressed because it is too large
Load Diff
131
apps/aether-gateway/src/data/state/routing_profiles.rs
Normal file
131
apps/aether-gateway/src/data/state/routing_profiles.rs
Normal file
@@ -0,0 +1,131 @@
|
||||
use aether_data_contracts::repository::routing_profiles::{
|
||||
CreateRoutingGroupBindingRecord, CreateRoutingGroupRecord, CreateRoutingGroupVersionRecord,
|
||||
RoutingGroupBindingQuery, RoutingGroupLookupKey, RoutingGroupReadRepository,
|
||||
StoredRoutingGroup, StoredRoutingGroupBinding, StoredRoutingGroupVersion,
|
||||
UpdateRoutingGroupBindingRecord, UpdateRoutingGroupRecord,
|
||||
};
|
||||
use std::sync::Arc;
|
||||
|
||||
use super::{DataLayerError, GatewayDataState};
|
||||
|
||||
impl GatewayDataState {
|
||||
pub(crate) fn routing_group_read_repository(
|
||||
&self,
|
||||
) -> Option<Arc<dyn RoutingGroupReadRepository>> {
|
||||
self.routing_group_reader.clone()
|
||||
}
|
||||
|
||||
pub(crate) async fn list_routing_groups(
|
||||
&self,
|
||||
) -> Result<Vec<StoredRoutingGroup>, DataLayerError> {
|
||||
match &self.routing_group_reader {
|
||||
Some(repository) => repository.list_routing_groups().await,
|
||||
None => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn find_routing_group(
|
||||
&self,
|
||||
lookup: RoutingGroupLookupKey<'_>,
|
||||
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
|
||||
match &self.routing_group_reader {
|
||||
Some(repository) => repository.find_routing_group(lookup).await,
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn list_routing_group_bindings(
|
||||
&self,
|
||||
query: &RoutingGroupBindingQuery,
|
||||
) -> Result<Vec<StoredRoutingGroupBinding>, DataLayerError> {
|
||||
match &self.routing_group_reader {
|
||||
Some(repository) => repository.list_routing_group_bindings(query).await,
|
||||
None => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn list_routing_group_versions(
|
||||
&self,
|
||||
group_id: &str,
|
||||
) -> Result<Vec<StoredRoutingGroupVersion>, DataLayerError> {
|
||||
match &self.routing_group_reader {
|
||||
Some(repository) => repository.list_routing_group_versions(group_id).await,
|
||||
None => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn create_routing_group(
|
||||
&self,
|
||||
record: CreateRoutingGroupRecord,
|
||||
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
|
||||
match &self.routing_group_writer {
|
||||
Some(repository) => repository.create_routing_group(record).await.map(Some),
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn update_routing_group(
|
||||
&self,
|
||||
id: &str,
|
||||
patch: UpdateRoutingGroupRecord,
|
||||
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
|
||||
match &self.routing_group_writer {
|
||||
Some(repository) => repository.update_routing_group(id, patch).await,
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn delete_routing_group(&self, id: &str) -> Result<bool, DataLayerError> {
|
||||
match &self.routing_group_writer {
|
||||
Some(repository) => repository.delete_routing_group(id).await,
|
||||
None => Ok(false),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn create_routing_group_binding(
|
||||
&self,
|
||||
record: CreateRoutingGroupBindingRecord,
|
||||
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
|
||||
match &self.routing_group_writer {
|
||||
Some(repository) => repository
|
||||
.create_routing_group_binding(record)
|
||||
.await
|
||||
.map(Some),
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn update_routing_group_binding(
|
||||
&self,
|
||||
id: &str,
|
||||
patch: UpdateRoutingGroupBindingRecord,
|
||||
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
|
||||
match &self.routing_group_writer {
|
||||
Some(repository) => repository.update_routing_group_binding(id, patch).await,
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn delete_routing_group_binding(
|
||||
&self,
|
||||
id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
match &self.routing_group_writer {
|
||||
Some(repository) => repository.delete_routing_group_binding(id).await,
|
||||
None => Ok(false),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn create_routing_group_version(
|
||||
&self,
|
||||
record: CreateRoutingGroupVersionRecord,
|
||||
) -> Result<Option<StoredRoutingGroupVersion>, DataLayerError> {
|
||||
match &self.routing_group_writer {
|
||||
Some(repository) => repository
|
||||
.create_routing_group_version(record)
|
||||
.await
|
||||
.map(Some),
|
||||
None => Ok(None),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -38,7 +38,7 @@ use aether_data_contracts::repository::usage::{
|
||||
PendingUsageCleanupSummary, ProviderApiKeyWindowUsageRequest,
|
||||
StoredProviderApiKeyWindowUsageSummary, StoredUsageDailySummary, UsageAuditListQuery,
|
||||
UsageCleanupExecutionMode, UsageCleanupSummary, UsageCleanupTargets, UsageCleanupWindow,
|
||||
UsageDailyHeatmapQuery,
|
||||
UsageCounterFlushSummary, UsageCounterHealthSnapshot, UsageDailyHeatmapQuery,
|
||||
};
|
||||
use aether_runtime_state::RuntimeQueueStore;
|
||||
use aether_video_tasks_core::read_data_backed_video_task_response;
|
||||
@@ -262,6 +262,22 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn list_required_unread_active_announcements(
|
||||
&self,
|
||||
user_id: &str,
|
||||
now_unix_secs: u64,
|
||||
limit: usize,
|
||||
) -> Result<Vec<StoredAnnouncement>, DataLayerError> {
|
||||
match &self.announcement_reader {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.list_required_unread_active_announcements(user_id, now_unix_secs, limit)
|
||||
.await
|
||||
}
|
||||
None => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn create_announcement(
|
||||
&self,
|
||||
record: CreateAnnouncementRecord,
|
||||
@@ -954,6 +970,31 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn flush_usage_counter_deltas(
|
||||
&self,
|
||||
batch_size: usize,
|
||||
) -> Result<UsageCounterFlushSummary, DataLayerError> {
|
||||
match &self.usage_writer {
|
||||
Some(repository) => repository.flush_usage_counter_deltas(batch_size).await,
|
||||
None => Ok(UsageCounterFlushSummary::default()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn cleanup_processed_usage_counter_deltas(
|
||||
&self,
|
||||
cutoff_unix_secs: u64,
|
||||
batch_size: usize,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
match &self.usage_writer {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.cleanup_processed_usage_counter_deltas(cutoff_unix_secs, batch_size)
|
||||
.await
|
||||
}
|
||||
None => Ok(0),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn cleanup_stale_pending_requests(
|
||||
&self,
|
||||
cutoff_unix_secs: u64,
|
||||
@@ -1119,6 +1160,15 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn read_usage_counter_health(
|
||||
&self,
|
||||
) -> Result<UsageCounterHealthSnapshot, DataLayerError> {
|
||||
match &self.usage_reader {
|
||||
Some(repository) => repository.read_usage_counter_health().await,
|
||||
None => Ok(UsageCounterHealthSnapshot::default()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn summarize_usage_totals_by_user_ids(
|
||||
&self,
|
||||
user_ids: &[String],
|
||||
|
||||
@@ -38,6 +38,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -92,6 +94,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
|
||||
@@ -73,6 +73,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -126,6 +128,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -175,6 +179,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -312,6 +318,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -389,6 +397,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -447,6 +457,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: Some(provider_quota_reader),
|
||||
provider_quota_writer: Some(provider_quota_writer),
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -514,6 +526,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: Some(provider_quota_reader),
|
||||
provider_quota_writer: Some(provider_quota_writer),
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -563,6 +577,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -613,6 +629,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -674,6 +692,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: Some(usage_reader),
|
||||
usage_writer: Some(usage_writer),
|
||||
user_reader: None,
|
||||
@@ -737,6 +757,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -784,6 +806,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: Some(repository),
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -846,6 +870,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: Some(repository),
|
||||
@@ -901,6 +927,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: Some(user_repository),
|
||||
@@ -961,6 +989,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: Some(usage_reader),
|
||||
usage_writer: Some(usage_writer),
|
||||
user_reader: Some(user_repository),
|
||||
@@ -1022,6 +1052,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: Some(user_repository),
|
||||
@@ -1082,6 +1114,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: Some(usage_reader),
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1131,6 +1165,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1180,6 +1216,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1241,6 +1279,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1307,6 +1347,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1356,6 +1398,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1410,6 +1454,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1481,6 +1527,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1547,6 +1595,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1597,6 +1647,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1647,6 +1699,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1699,6 +1753,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: Some(usage_repository),
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1749,6 +1805,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1799,6 +1857,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1849,6 +1909,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1907,6 +1969,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: Some(provider_quota_reader),
|
||||
provider_quota_writer: Some(provider_quota_writer),
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -1966,6 +2030,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: Some(provider_quota_reader),
|
||||
provider_quota_writer: Some(provider_quota_writer),
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -2028,6 +2094,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: Some(provider_quota_reader),
|
||||
provider_quota_writer: Some(provider_quota_writer),
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -2096,6 +2164,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -2165,6 +2235,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -2238,6 +2310,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: Some(usage_reader),
|
||||
usage_writer: Some(usage_writer),
|
||||
user_reader: None,
|
||||
@@ -2318,6 +2392,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: Some(usage_reader),
|
||||
usage_writer: Some(usage_writer),
|
||||
user_reader: None,
|
||||
@@ -2380,6 +2456,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: Some(provider_quota_reader),
|
||||
provider_quota_writer: Some(provider_quota_writer),
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -2433,6 +2511,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: Some(usage_reader),
|
||||
usage_writer: Some(usage_writer),
|
||||
user_reader: None,
|
||||
@@ -2482,6 +2562,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -2537,6 +2619,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -2596,6 +2680,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: Some(usage_reader),
|
||||
usage_writer: Some(usage_writer),
|
||||
user_reader: None,
|
||||
@@ -2656,6 +2742,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: Some(usage_reader),
|
||||
usage_writer: Some(usage_writer),
|
||||
user_reader: None,
|
||||
@@ -2709,6 +2797,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: Some(provider_quota_reader),
|
||||
provider_quota_writer: Some(provider_quota_writer),
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
|
||||
@@ -42,6 +42,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -98,6 +100,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -151,6 +155,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -208,6 +214,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -269,6 +277,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
@@ -339,6 +349,8 @@ impl GatewayDataState {
|
||||
pool_score_writer: None,
|
||||
provider_quota_reader: None,
|
||||
provider_quota_writer: None,
|
||||
routing_group_reader: None,
|
||||
routing_group_writer: None,
|
||||
usage_reader: None,
|
||||
usage_writer: None,
|
||||
user_reader: None,
|
||||
|
||||
@@ -16,6 +16,7 @@ use aether_pool_core::{
|
||||
PoolMemberSignals, PoolRuntimeState, PoolSchedulingConfig, PoolSchedulingPreset,
|
||||
};
|
||||
use aether_provider_pool::ProviderPoolService;
|
||||
use aether_routing_core::{RankingOverlay, ResolvedRoutingPolicy};
|
||||
use tracing::warn;
|
||||
|
||||
use crate::ai_serving::{
|
||||
@@ -40,6 +41,7 @@ use crate::orchestration::LocalExecutionCandidateMetadata;
|
||||
|
||||
static LOAD_BALANCE_SEQUENCE: AtomicU64 = AtomicU64::new(0);
|
||||
const POOL_ACTIVE_PROBE_SEALED_SKIP_REASON: &str = "pool_active_probe_sealed";
|
||||
const ROUTING_PROFILE_DISALLOWED_KEY_SKIP_REASON: &str = "routing_profile_disallowed_key";
|
||||
|
||||
type PoolCatalogKeyContext = PoolMemberSignals;
|
||||
|
||||
@@ -187,6 +189,7 @@ pub(crate) struct PoolKeyCursor<'a> {
|
||||
sticky_session_token: Option<String>,
|
||||
requested_model: Option<String>,
|
||||
request_auth_channel: Option<String>,
|
||||
routing_overlay: Option<RankingOverlay>,
|
||||
runtime_miss_trace_id: Option<String>,
|
||||
record_runtime_miss_diagnostic: bool,
|
||||
pool_key_order: StoredPoolKeyCandidateOrder,
|
||||
@@ -216,7 +219,26 @@ impl<'a> PoolKeyCursor<'a> {
|
||||
requested_model: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
) -> Self {
|
||||
let pool_key_order = pool_key_candidate_order_for_group(&group);
|
||||
Self::new_with_routing_policy(
|
||||
state,
|
||||
group,
|
||||
sticky_session_token,
|
||||
requested_model,
|
||||
request_auth_channel,
|
||||
None,
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) fn new_with_routing_policy(
|
||||
state: PlannerAppState<'a>,
|
||||
group: EligibleLocalExecutionCandidate,
|
||||
sticky_session_token: Option<&str>,
|
||||
requested_model: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
) -> Self {
|
||||
let pool_key_order = pool_key_candidate_order_for_group(&group, routing_policy);
|
||||
let routing_overlay = routing_policy.map(|policy| policy.ranking_overlay.clone());
|
||||
let pool_config = pool_config_for_candidate(&group);
|
||||
let score_top_n = pool_config
|
||||
.as_ref()
|
||||
@@ -236,6 +258,7 @@ impl<'a> PoolKeyCursor<'a> {
|
||||
sticky_session_token: sticky_session_token.map(str::to_string),
|
||||
requested_model: requested_model.map(str::to_string),
|
||||
request_auth_channel: request_auth_channel.map(str::to_string),
|
||||
routing_overlay,
|
||||
runtime_miss_trace_id: None,
|
||||
record_runtime_miss_diagnostic: false,
|
||||
pool_key_order,
|
||||
@@ -551,6 +574,9 @@ impl<'a> PoolKeyCursor<'a> {
|
||||
async fn next_queued_candidate(&mut self) -> Option<EligibleLocalExecutionCandidate> {
|
||||
while let Some(candidate) = self.queued_candidates.pop_front() {
|
||||
let mut candidate = candidate;
|
||||
if self.skip_candidate_if_routing_profile_disallowed(&candidate) {
|
||||
continue;
|
||||
}
|
||||
if self.skip_candidate_if_runtime_cooldown(&candidate).await {
|
||||
continue;
|
||||
}
|
||||
@@ -562,6 +588,28 @@ impl<'a> PoolKeyCursor<'a> {
|
||||
None
|
||||
}
|
||||
|
||||
fn skip_candidate_if_routing_profile_disallowed(
|
||||
&mut self,
|
||||
candidate: &EligibleLocalExecutionCandidate,
|
||||
) -> bool {
|
||||
let Some(overlay) = self.routing_overlay.as_ref() else {
|
||||
return false;
|
||||
};
|
||||
if overlay.key_allowed(candidate.candidate.key_id.as_str()) {
|
||||
return false;
|
||||
}
|
||||
self.record_skip_reason(ROUTING_PROFILE_DISALLOWED_KEY_SKIP_REASON);
|
||||
self.skipped_candidates
|
||||
.push(SkippedLocalExecutionCandidate {
|
||||
candidate: candidate.candidate.clone(),
|
||||
skip_reason: ROUTING_PROFILE_DISALLOWED_KEY_SKIP_REASON,
|
||||
transport: Some(candidate.transport.clone()),
|
||||
ranking: candidate.ranking.clone(),
|
||||
extra_data: None,
|
||||
});
|
||||
true
|
||||
}
|
||||
|
||||
async fn skip_candidate_if_runtime_cooldown(
|
||||
&mut self,
|
||||
candidate: &EligibleLocalExecutionCandidate,
|
||||
@@ -958,19 +1006,38 @@ fn should_trigger_active_probe_burst_for_request(
|
||||
|
||||
fn pool_key_candidate_order_for_group(
|
||||
group: &EligibleLocalExecutionCandidate,
|
||||
routing_policy: Option<&ResolvedRoutingPolicy>,
|
||||
) -> StoredPoolKeyCandidateOrder {
|
||||
let Some(pool_config) = pool_config_for_candidate(group) else {
|
||||
return StoredPoolKeyCandidateOrder::InternalPriority;
|
||||
};
|
||||
let presets = pool_config
|
||||
.scheduling_presets
|
||||
.iter()
|
||||
.map(|preset| PoolSchedulingPreset {
|
||||
preset: preset.preset.clone(),
|
||||
enabled: preset.enabled,
|
||||
mode: preset.mode.clone(),
|
||||
let override_presets = routing_policy
|
||||
.and_then(|policy| {
|
||||
policy
|
||||
.pool_policy_overrides
|
||||
.get(group.candidate.provider_id.as_str())
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
.filter(|override_policy| !override_policy.scheduling_presets.is_empty());
|
||||
let presets = match override_presets {
|
||||
Some(override_policy) => override_policy
|
||||
.scheduling_presets
|
||||
.iter()
|
||||
.map(|preset| PoolSchedulingPreset {
|
||||
preset: preset.preset.clone(),
|
||||
enabled: preset.enabled,
|
||||
mode: preset.mode.clone(),
|
||||
})
|
||||
.collect::<Vec<_>>(),
|
||||
None => pool_config
|
||||
.scheduling_presets
|
||||
.iter()
|
||||
.map(|preset| PoolSchedulingPreset {
|
||||
preset: preset.preset.clone(),
|
||||
enabled: preset.enabled,
|
||||
mode: preset.mode.clone(),
|
||||
})
|
||||
.collect::<Vec<_>>(),
|
||||
};
|
||||
let active_presets = ProviderPoolService::with_builtin_adapters()
|
||||
.normalize_scheduling_presets(group.transport.provider.provider_type.as_str(), &presets)
|
||||
.into_iter()
|
||||
@@ -1071,6 +1138,7 @@ mod tests {
|
||||
apply_local_execution_pool_scheduler_with_runtime_map, build_pool_catalog_key_context,
|
||||
pool_config_for_candidate, should_trigger_active_probe_burst_for_request,
|
||||
PoolCatalogKeyContext, PoolKeyCursor, POOL_ACTIVE_PROBE_SEALED_SKIP_REASON,
|
||||
ROUTING_PROFILE_DISALLOWED_KEY_SKIP_REASON,
|
||||
};
|
||||
use crate::ai_serving::{
|
||||
apply_local_runtime_candidate_terminal_reason, EligibleLocalExecutionCandidate,
|
||||
@@ -1096,6 +1164,9 @@ mod tests {
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider,
|
||||
};
|
||||
use aether_routing_core::{
|
||||
RankingOverlay, ResolvedRoutingPolicy, RoutingSchedulingMode, RoutingSetPriorityMode,
|
||||
};
|
||||
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
|
||||
use serde_json::json;
|
||||
use std::collections::{BTreeMap, BTreeSet, VecDeque};
|
||||
@@ -2103,6 +2174,59 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pool_key_cursor_filters_expanded_keys_by_routing_profile_allowed_keys() {
|
||||
let app = AppState::new().expect("state should build");
|
||||
let provider_config = Some(json!({ "pool_advanced": { "lru_enabled": true } }));
|
||||
let group = sample_eligible_candidate(
|
||||
"provider-pool",
|
||||
"endpoint-1",
|
||||
"pool-group",
|
||||
10,
|
||||
provider_config.clone(),
|
||||
);
|
||||
let routing_policy = routing_policy_with_allowed_keys(["key-b"]);
|
||||
let mut cursor = PoolKeyCursor::new_with_routing_policy(
|
||||
PlannerAppState::new(&app),
|
||||
group,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
Some(&routing_policy),
|
||||
);
|
||||
cursor.queued_candidates = VecDeque::from([
|
||||
sample_eligible_candidate(
|
||||
"provider-pool",
|
||||
"endpoint-1",
|
||||
"key-a",
|
||||
10,
|
||||
provider_config.clone(),
|
||||
),
|
||||
sample_eligible_candidate("provider-pool", "endpoint-1", "key-b", 10, provider_config),
|
||||
]);
|
||||
|
||||
let candidate = cursor
|
||||
.next_key()
|
||||
.await
|
||||
.expect("cursor should skip disallowed pool key and return allowed key");
|
||||
assert_eq!(candidate.candidate.key_id, "key-b");
|
||||
assert_eq!(candidate.orchestration.pool_key_index, Some(0));
|
||||
assert_eq!(
|
||||
cursor
|
||||
.skip_reason_counts
|
||||
.get(ROUTING_PROFILE_DISALLOWED_KEY_SKIP_REASON),
|
||||
Some(&1)
|
||||
);
|
||||
let skipped = cursor.take_skipped_candidates();
|
||||
assert_eq!(
|
||||
skipped
|
||||
.iter()
|
||||
.map(|item| (item.candidate.key_id.as_str(), item.skip_reason))
|
||||
.collect::<Vec<_>>(),
|
||||
vec![("key-a", ROUTING_PROFILE_DISALLOWED_KEY_SKIP_REASON)]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pool_key_cursor_allows_parallel_requests_to_use_same_healthy_key() {
|
||||
let app = AppState::new().expect("state should build");
|
||||
@@ -2668,6 +2792,28 @@ mod tests {
|
||||
(provider, endpoint, keys, rows)
|
||||
}
|
||||
|
||||
fn routing_policy_with_allowed_keys<const N: usize>(
|
||||
key_ids: [&str; N],
|
||||
) -> ResolvedRoutingPolicy {
|
||||
ResolvedRoutingPolicy {
|
||||
group_id: Some("routing-group-1".to_string()),
|
||||
group_version: Some(1),
|
||||
selection_source: "test".to_string(),
|
||||
requested_model: "gpt-5".to_string(),
|
||||
resolved_model: "gpt-5".to_string(),
|
||||
priority_mode: RoutingSetPriorityMode::Provider,
|
||||
scheduling_mode: RoutingSchedulingMode::CacheAffinity,
|
||||
keep_priority_on_conversion: false,
|
||||
ranking_overlay: RankingOverlay {
|
||||
allowed_keys: key_ids.into_iter().map(str::to_string).collect(),
|
||||
..RankingOverlay::default()
|
||||
},
|
||||
mutation_plan: Default::default(),
|
||||
pool_policy_overrides: BTreeMap::new(),
|
||||
matched_rules: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_eligible_candidate(
|
||||
provider_id: &str,
|
||||
endpoint_id: &str,
|
||||
|
||||
@@ -3,9 +3,10 @@ use std::io::Error as IoError;
|
||||
use std::time::Instant;
|
||||
|
||||
use aether_contracts::{
|
||||
ExecutionPlan, ExecutionResult, ExecutionTelemetry, RequestBody, ResponseBody, StreamFrame,
|
||||
StreamFramePayload, StreamFrameType, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER,
|
||||
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
|
||||
ExecutionPlan, ExecutionResult, ExecutionTelemetry, RequestBody, ResolvedTransportProfile,
|
||||
ResponseBody, StreamFrame, StreamFramePayload, StreamFrameType,
|
||||
EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER, EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER,
|
||||
TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_HTTP_MODE_AUTO, TRANSPORT_POOL_SCOPE_KEY,
|
||||
};
|
||||
use axum::body::Bytes;
|
||||
use base64::Engine as _;
|
||||
@@ -30,6 +31,7 @@ const CHATGPT_WEB_CLIENT_VERSION: &str = "prod-be885abbfcfe7b1f511e88b3003d9ee44
|
||||
const CHATGPT_WEB_BUILD_NUMBER: &str = "5955942";
|
||||
const CHATGPT_WEB_SEC_CH_UA: &str =
|
||||
r#""Microsoft Edge";v="143", "Chromium";v="143", "Not A(Brand";v="24""#;
|
||||
const CHATGPT_WEB_BROWSER_PROFILE: &str = "chrome143";
|
||||
|
||||
pub(crate) struct ChatGptWebImageStream {
|
||||
pub(crate) frame_stream: BoxStream<'static, Result<Bytes, IoError>>,
|
||||
@@ -921,7 +923,7 @@ async fn execute_subrequest(
|
||||
provider_api_format: plan.provider_api_format.clone(),
|
||||
model_name: plan.model_name.clone(),
|
||||
proxy: plan.proxy.clone(),
|
||||
transport_profile: plan.transport_profile.clone(),
|
||||
transport_profile: chatgpt_web_image_transport_profile(plan),
|
||||
timeouts: plan.timeouts.clone(),
|
||||
};
|
||||
DirectSyncExecutionRuntime::new()
|
||||
@@ -929,6 +931,34 @@ async fn execute_subrequest(
|
||||
.await
|
||||
}
|
||||
|
||||
fn chatgpt_web_image_transport_profile(plan: &ExecutionPlan) -> Option<ResolvedTransportProfile> {
|
||||
match plan.transport_profile.as_ref() {
|
||||
Some(profile)
|
||||
if profile
|
||||
.backend
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(TRANSPORT_BACKEND_BROWSER_WREQ) =>
|
||||
{
|
||||
Some(profile.clone())
|
||||
}
|
||||
_ => Some(default_chatgpt_web_image_transport_profile()),
|
||||
}
|
||||
}
|
||||
|
||||
fn default_chatgpt_web_image_transport_profile() -> ResolvedTransportProfile {
|
||||
ResolvedTransportProfile {
|
||||
profile_id: CHATGPT_WEB_BROWSER_PROFILE.to_string(),
|
||||
backend: TRANSPORT_BACKEND_BROWSER_WREQ.to_string(),
|
||||
http_mode: TRANSPORT_HTTP_MODE_AUTO.to_string(),
|
||||
pool_scope: TRANSPORT_POOL_SCOPE_KEY.to_string(),
|
||||
header_fingerprint: None,
|
||||
extra: Some(json!({
|
||||
"browser_profile": CHATGPT_WEB_BROWSER_PROFILE,
|
||||
"source": "chatgpt_web_image_default",
|
||||
})),
|
||||
}
|
||||
}
|
||||
|
||||
fn web_base_headers(fp: &WebFingerprint, token: &str, path: &str) -> BTreeMap<String, String> {
|
||||
let mut headers = BTreeMap::from([
|
||||
("user-agent".to_string(), fp.user_agent.to_string()),
|
||||
@@ -1962,6 +1992,30 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn chatgpt_web_image_subrequests_default_to_browser_wreq_transport() {
|
||||
let plan = sample_plan(
|
||||
CHATGPT_WEB_DEFAULT_BASE_URL,
|
||||
json!({"prompt": "draw a small test image"}),
|
||||
false,
|
||||
);
|
||||
|
||||
let profile = chatgpt_web_image_transport_profile(&plan).expect("transport profile");
|
||||
|
||||
assert_eq!(profile.backend, TRANSPORT_BACKEND_BROWSER_WREQ);
|
||||
assert_eq!(profile.profile_id, CHATGPT_WEB_BROWSER_PROFILE);
|
||||
assert_eq!(profile.http_mode, TRANSPORT_HTTP_MODE_AUTO);
|
||||
assert_eq!(profile.pool_scope, TRANSPORT_POOL_SCOPE_KEY);
|
||||
assert_eq!(
|
||||
profile
|
||||
.extra
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("source"))
|
||||
.and_then(Value::as_str),
|
||||
Some("chatgpt_web_image_default")
|
||||
);
|
||||
}
|
||||
|
||||
async fn start_mock_chatgpt_web() -> (String, tokio::task::JoinHandle<()>) {
|
||||
let app = Router::new().fallback(any(|request: Request| async move {
|
||||
let path = request.uri().path().to_string();
|
||||
|
||||
4061
apps/aether-gateway/src/execution_runtime/grok.rs
Normal file
4061
apps/aether-gateway/src/execution_runtime/grok.rs
Normal file
File diff suppressed because it is too large
Load Diff
@@ -6,6 +6,7 @@ use serde_json::{Map, Value};
|
||||
mod chatgpt_web_image;
|
||||
mod constants;
|
||||
mod fallback;
|
||||
mod grok;
|
||||
mod kiro_web_search;
|
||||
pub(crate) mod ndjson;
|
||||
mod oauth_retry;
|
||||
@@ -54,6 +55,7 @@ pub(crate) use sync::{
|
||||
resolve_local_sync_success_background_report_kind, LocalVideoSyncSuccessBuild,
|
||||
LocalVideoSyncSuccessOutcome,
|
||||
};
|
||||
pub(crate) use transport::execute_sync_plan_with_report_context as execute_execution_runtime_sync_plan_with_report_context;
|
||||
pub(crate) use transport::{
|
||||
execute_sync_plan as execute_execution_runtime_sync_plan, DirectSyncExecutionRuntime,
|
||||
DirectUpstreamStreamExecution, ExecutionRuntimeTransportError,
|
||||
|
||||
@@ -370,6 +370,8 @@ impl IntoResponse for ExecutionRuntimeAppError {
|
||||
) => StatusCode::BAD_REQUEST,
|
||||
ExecutionRuntimeServerError::Transport(
|
||||
ExecutionRuntimeTransportError::ClientBuild(_)
|
||||
| ExecutionRuntimeTransportError::BrowserClientBuild(_)
|
||||
| ExecutionRuntimeTransportError::BrowserBody(_)
|
||||
| ExecutionRuntimeTransportError::UpstreamRequest(_)
|
||||
| ExecutionRuntimeTransportError::RelayError(_)
|
||||
| ExecutionRuntimeTransportError::InvalidJson(_),
|
||||
|
||||
@@ -60,6 +60,7 @@ use crate::constants::{CONTROL_CANDIDATE_ID_HEADER, CONTROL_REQUEST_ID_HEADER};
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::execution_runtime::build_direct_execution_frame_stream;
|
||||
use crate::execution_runtime::chatgpt_web_image::maybe_execute_chatgpt_web_image_stream;
|
||||
use crate::execution_runtime::grok::maybe_execute_grok_stream;
|
||||
use crate::execution_runtime::kiro_web_search::maybe_execute_kiro_web_search_stream;
|
||||
use crate::execution_runtime::oauth_retry::refresh_oauth_plan_auth_for_retry;
|
||||
#[cfg(test)]
|
||||
@@ -525,6 +526,58 @@ pub(crate) async fn execute_execution_runtime_stream(
|
||||
key_id.as_str(),
|
||||
)
|
||||
.await;
|
||||
match maybe_execute_grok_stream(&plan, report_context.as_ref()).await {
|
||||
Ok(Some(grok_stream)) => {
|
||||
return execute_stream_from_frame_stream(
|
||||
state,
|
||||
plan,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
report_kind,
|
||||
grok_stream.report_context.or(report_context),
|
||||
candidate_started_unix_secs,
|
||||
stream_started_at,
|
||||
grok_stream.frame_stream,
|
||||
provider_pool_in_flight_guard.take(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
Ok(None) => {}
|
||||
Err(err) => {
|
||||
info!(
|
||||
event_name = "grok_execution_unavailable",
|
||||
log_type = "ops",
|
||||
trace_id = %trace_id,
|
||||
request_id = %plan_request_id_for_log,
|
||||
candidate_id = ?plan.candidate_id,
|
||||
provider_name = provider_name.as_str(),
|
||||
endpoint_id = %endpoint_id,
|
||||
key_id = %key_id,
|
||||
model_name = model_name.as_str(),
|
||||
candidate_index = candidate_index.as_str(),
|
||||
error = %err,
|
||||
"gateway Grok stream execution unavailable"
|
||||
);
|
||||
let terminal_unix_secs = current_request_candidate_unix_ms();
|
||||
record_local_request_candidate_status(
|
||||
state,
|
||||
&plan,
|
||||
report_context.as_ref(),
|
||||
SchedulerRequestCandidateStatusUpdate {
|
||||
status: RequestCandidateStatus::Failed,
|
||||
status_code: None,
|
||||
error_type: Some("grok_execution_unavailable".to_string()),
|
||||
error_message: Some(format!("{err:?}")),
|
||||
latency_ms: None,
|
||||
started_at_unix_ms: Some(candidate_started_unix_secs),
|
||||
finished_at_unix_ms: Some(terminal_unix_secs),
|
||||
},
|
||||
)
|
||||
.await;
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
match maybe_execute_kiro_web_search_stream(state, &plan, report_context.as_ref()).await {
|
||||
Ok(Some(kiro_web_search)) => {
|
||||
return execute_stream_from_frame_stream(
|
||||
|
||||
@@ -18,7 +18,9 @@ use crate::ai_serving::api::{
|
||||
normalize_provider_private_report_context, StreamingStandardTerminalObserver,
|
||||
};
|
||||
use crate::execution_runtime::ndjson::encode_stream_frame_ndjson;
|
||||
use crate::execution_runtime::transport::DirectUpstreamResponse;
|
||||
use crate::execution_runtime::transport::{
|
||||
format_wreq_upstream_request_error, DirectUpstreamResponse,
|
||||
};
|
||||
use crate::execution_runtime::DirectUpstreamStreamExecution;
|
||||
use crate::GatewayError;
|
||||
|
||||
@@ -235,6 +237,62 @@ pub(crate) fn build_direct_execution_frame_stream(
|
||||
}
|
||||
}
|
||||
}
|
||||
DirectUpstreamResponse::BrowserWreq(response) => {
|
||||
let mut bytes_stream = response.bytes_stream();
|
||||
while let Some(item) = bytes_stream.next().await {
|
||||
match item {
|
||||
Ok(chunk) => {
|
||||
if ttfb_ms.is_none() {
|
||||
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
|
||||
}
|
||||
if !first_chunk_telemetry_emitted {
|
||||
match encode_telemetry_frame(ttfb_ms, ttfb_ms, upstream_bytes) {
|
||||
Ok(frame) => yield Ok(frame),
|
||||
Err(err) => {
|
||||
yield Err(err);
|
||||
return;
|
||||
}
|
||||
}
|
||||
first_chunk_telemetry_emitted = true;
|
||||
}
|
||||
upstream_bytes += chunk.len() as u64;
|
||||
observe_stream_chunk(
|
||||
&mut stream_terminal_observer,
|
||||
&normalized_observer_context,
|
||||
private_stream_normalizer.as_mut(),
|
||||
&mut observer_buffered,
|
||||
chunk.as_ref(),
|
||||
);
|
||||
match encode_data_frame(&chunk) {
|
||||
Ok(frame) => yield Ok(frame),
|
||||
Err(err) => {
|
||||
yield Err(err);
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(err) => {
|
||||
let message = format_wreq_upstream_request_error(&err);
|
||||
warn!(
|
||||
event_name = "stream_pump_body_read_error",
|
||||
log_type = "ops",
|
||||
status_code,
|
||||
upstream_bytes,
|
||||
error = %message,
|
||||
"upstream body stream read error"
|
||||
);
|
||||
match encode_error_frame(status_code, message) {
|
||||
Ok(frame) => yield Ok(frame),
|
||||
Err(encode_err) => {
|
||||
yield Err(encode_err);
|
||||
return;
|
||||
}
|
||||
}
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
DirectUpstreamResponse::LocalTunnel(mut response) => loop {
|
||||
match response.next_chunk().await {
|
||||
Ok(Some(chunk)) => {
|
||||
@@ -454,6 +512,35 @@ async fn buffer_non_sse_upstream_body(
|
||||
}
|
||||
}
|
||||
}
|
||||
DirectUpstreamResponse::BrowserWreq(response) => {
|
||||
let mut bytes_stream = response.bytes_stream();
|
||||
while let Some(item) = bytes_stream.next().await {
|
||||
match item {
|
||||
Ok(chunk) => {
|
||||
if ttfb_ms.is_none() {
|
||||
ttfb_ms = Some(started_at.elapsed().as_millis() as u64);
|
||||
}
|
||||
upstream_bytes += chunk.len() as u64;
|
||||
body_bytes.extend_from_slice(&chunk);
|
||||
}
|
||||
Err(err) => {
|
||||
let message = format_wreq_upstream_request_error(&err);
|
||||
warn!(
|
||||
event_name = "stream_pump_body_read_error",
|
||||
log_type = "ops",
|
||||
upstream_bytes,
|
||||
error = %message,
|
||||
"upstream body stream read error"
|
||||
);
|
||||
return Err(BufferedUpstreamBodyError {
|
||||
message,
|
||||
ttfb_ms,
|
||||
upstream_bytes,
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
DirectUpstreamResponse::LocalTunnel(mut response) => loop {
|
||||
match response.next_chunk().await {
|
||||
Ok(Some(chunk)) => {
|
||||
|
||||
@@ -39,14 +39,15 @@ use crate::api::response::{
|
||||
use crate::clock::current_unix_ms as current_request_candidate_unix_ms;
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::execution_runtime::chatgpt_web_image::maybe_execute_chatgpt_web_image_sync;
|
||||
use crate::execution_runtime::grok::maybe_execute_grok_sync;
|
||||
use crate::execution_runtime::oauth_retry::refresh_oauth_plan_auth_for_retry;
|
||||
#[cfg(test)]
|
||||
use crate::execution_runtime::remote_compat::post_sync_plan_to_remote_execution_runtime;
|
||||
use crate::execution_runtime::submission::submit_local_core_error_or_sync_finalize;
|
||||
use crate::execution_runtime::transport::{
|
||||
build_request_body, collect_response_headers, decode_response_body_bytes,
|
||||
response_body_is_json, send_request, DirectSyncExecutionRuntime,
|
||||
ExecutionRuntimeTransportError,
|
||||
format_upstream_request_error, format_wreq_upstream_request_error, response_body_is_json,
|
||||
send_request, DirectHttpResponse, DirectSyncExecutionRuntime, ExecutionRuntimeTransportError,
|
||||
};
|
||||
use crate::execution_runtime::{
|
||||
analyze_local_candidate_failover_sync, apply_endpoint_response_header_rules,
|
||||
@@ -734,23 +735,46 @@ async fn execute_openai_image_sync_upstream_sse_candidate(
|
||||
.await
|
||||
.map_err(SyncExecutionFailure::from_transport)?;
|
||||
let ttfb_ms = started_at.elapsed().as_millis() as u64;
|
||||
let status_code = response.status().as_u16();
|
||||
let headers = collect_response_headers(response.headers());
|
||||
let status_code = response.status_code();
|
||||
let headers = response.headers();
|
||||
progress.record_response_started(status_code, ttfb_ms).await;
|
||||
|
||||
let mut upstream_stream = response.bytes_stream();
|
||||
let mut body_bytes = Vec::new();
|
||||
while let Some(chunk) = upstream_stream.next().await {
|
||||
let chunk = chunk.map_err(|err| {
|
||||
SyncExecutionFailure::from_transport(ExecutionRuntimeTransportError::UpstreamRequest(
|
||||
crate::execution_runtime::transport::format_upstream_request_error(&err),
|
||||
))
|
||||
})?;
|
||||
let elapsed_ms = started_at.elapsed().as_millis() as u64;
|
||||
progress
|
||||
.observe_chunk(&chunk, status_code, elapsed_ms)
|
||||
.await;
|
||||
body_bytes.extend_from_slice(&chunk);
|
||||
match response {
|
||||
DirectHttpResponse::Reqwest(response) => {
|
||||
let mut upstream_stream = response.bytes_stream();
|
||||
while let Some(chunk) = upstream_stream.next().await {
|
||||
let chunk = chunk.map_err(|err| {
|
||||
SyncExecutionFailure::from_transport(
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(
|
||||
format_upstream_request_error(&err),
|
||||
),
|
||||
)
|
||||
})?;
|
||||
let elapsed_ms = started_at.elapsed().as_millis() as u64;
|
||||
progress
|
||||
.observe_chunk(&chunk, status_code, elapsed_ms)
|
||||
.await;
|
||||
body_bytes.extend_from_slice(&chunk);
|
||||
}
|
||||
}
|
||||
DirectHttpResponse::BrowserWreq(response) => {
|
||||
let mut upstream_stream = response.bytes_stream();
|
||||
while let Some(chunk) = upstream_stream.next().await {
|
||||
let chunk = chunk.map_err(|err| {
|
||||
SyncExecutionFailure::from_transport(
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(
|
||||
format_wreq_upstream_request_error(&err),
|
||||
),
|
||||
)
|
||||
})?;
|
||||
let elapsed_ms = started_at.elapsed().as_millis() as u64;
|
||||
progress
|
||||
.observe_chunk(&chunk, status_code, elapsed_ms)
|
||||
.await;
|
||||
body_bytes.extend_from_slice(&chunk);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let decoded_body_bytes =
|
||||
@@ -1129,64 +1153,106 @@ async fn execute_execution_runtime_sync_impl(
|
||||
.await;
|
||||
#[cfg(not(test))]
|
||||
let mut result = {
|
||||
match maybe_execute_chatgpt_web_image_sync(state, &plan, report_context.as_ref()).await {
|
||||
match maybe_execute_grok_sync(&plan, report_context.as_ref()).await {
|
||||
Ok(Some(result)) => result,
|
||||
Ok(None) => match execute_direct_sync_runtime_candidate(
|
||||
state,
|
||||
&plan,
|
||||
report_context.as_ref(),
|
||||
trace_id,
|
||||
plan_kind,
|
||||
plan_request_id_for_log.as_str(),
|
||||
plan_candidate_id.as_deref(),
|
||||
provider_name.as_str(),
|
||||
endpoint_id.as_str(),
|
||||
key_id.as_str(),
|
||||
model_name.as_str(),
|
||||
candidate_index.as_str(),
|
||||
progress_snapshot.clone(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(err) => {
|
||||
warn!(
|
||||
event_name = "sync_execution_runtime_unavailable",
|
||||
log_type = "ops",
|
||||
trace_id = %trace_id,
|
||||
request_id = %plan_request_id_for_log,
|
||||
candidate_id = ?plan_candidate_id,
|
||||
provider_name,
|
||||
endpoint_id,
|
||||
key_id,
|
||||
model_name,
|
||||
candidate_index = candidate_index.as_str(),
|
||||
error_type = err.error_type,
|
||||
error = %err.message,
|
||||
"gateway in-process sync execution unavailable"
|
||||
);
|
||||
let terminal_unix_secs = current_request_candidate_unix_ms();
|
||||
record_local_request_candidate_status(
|
||||
Ok(None) => {
|
||||
match maybe_execute_chatgpt_web_image_sync(state, &plan, report_context.as_ref())
|
||||
.await
|
||||
{
|
||||
Ok(Some(result)) => result,
|
||||
Ok(None) => match execute_direct_sync_runtime_candidate(
|
||||
state,
|
||||
&plan,
|
||||
report_context.as_ref(),
|
||||
SchedulerRequestCandidateStatusUpdate {
|
||||
status: RequestCandidateStatus::Failed,
|
||||
status_code: err.status_code,
|
||||
error_type: Some(err.error_type.to_string()),
|
||||
error_message: Some(err.message),
|
||||
latency_ms: err.latency_ms,
|
||||
started_at_unix_ms: Some(candidate_started_unix_secs),
|
||||
finished_at_unix_ms: Some(terminal_unix_secs),
|
||||
},
|
||||
trace_id,
|
||||
plan_kind,
|
||||
plan_request_id_for_log.as_str(),
|
||||
plan_candidate_id.as_deref(),
|
||||
provider_name.as_str(),
|
||||
endpoint_id.as_str(),
|
||||
key_id.as_str(),
|
||||
model_name.as_str(),
|
||||
candidate_index.as_str(),
|
||||
progress_snapshot.clone(),
|
||||
)
|
||||
.await;
|
||||
return Ok(None);
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(err) => {
|
||||
warn!(
|
||||
event_name = "sync_execution_runtime_unavailable",
|
||||
log_type = "ops",
|
||||
trace_id = %trace_id,
|
||||
request_id = %plan_request_id_for_log,
|
||||
candidate_id = ?plan_candidate_id,
|
||||
provider_name,
|
||||
endpoint_id,
|
||||
key_id,
|
||||
model_name,
|
||||
candidate_index = candidate_index.as_str(),
|
||||
error_type = err.error_type,
|
||||
error = %err.message,
|
||||
"gateway in-process sync execution unavailable"
|
||||
);
|
||||
let terminal_unix_secs = current_request_candidate_unix_ms();
|
||||
record_local_request_candidate_status(
|
||||
state,
|
||||
&plan,
|
||||
report_context.as_ref(),
|
||||
SchedulerRequestCandidateStatusUpdate {
|
||||
status: RequestCandidateStatus::Failed,
|
||||
status_code: err.status_code,
|
||||
error_type: Some(err.error_type.to_string()),
|
||||
error_message: Some(err.message),
|
||||
latency_ms: err.latency_ms,
|
||||
started_at_unix_ms: Some(candidate_started_unix_secs),
|
||||
finished_at_unix_ms: Some(terminal_unix_secs),
|
||||
},
|
||||
)
|
||||
.await;
|
||||
return Ok(None);
|
||||
}
|
||||
},
|
||||
Err(err) => {
|
||||
warn!(
|
||||
event_name = "chatgpt_web_image_execution_unavailable",
|
||||
log_type = "ops",
|
||||
trace_id = %trace_id,
|
||||
request_id = %plan_request_id_for_log,
|
||||
candidate_id = ?plan_candidate_id,
|
||||
provider_name,
|
||||
endpoint_id,
|
||||
key_id,
|
||||
model_name,
|
||||
candidate_index = candidate_index.as_str(),
|
||||
error = %err,
|
||||
"gateway ChatGPT-Web image execution unavailable"
|
||||
);
|
||||
let terminal_unix_secs = current_request_candidate_unix_ms();
|
||||
record_local_request_candidate_status(
|
||||
state,
|
||||
&plan,
|
||||
report_context.as_ref(),
|
||||
SchedulerRequestCandidateStatusUpdate {
|
||||
status: RequestCandidateStatus::Failed,
|
||||
status_code: None,
|
||||
error_type: Some(
|
||||
"chatgpt_web_image_execution_unavailable".to_string(),
|
||||
),
|
||||
error_message: Some(err.to_string()),
|
||||
latency_ms: None,
|
||||
started_at_unix_ms: Some(candidate_started_unix_secs),
|
||||
finished_at_unix_ms: Some(terminal_unix_secs),
|
||||
},
|
||||
)
|
||||
.await;
|
||||
return Ok(None);
|
||||
}
|
||||
}
|
||||
},
|
||||
}
|
||||
Err(err) => {
|
||||
warn!(
|
||||
event_name = "chatgpt_web_image_execution_unavailable",
|
||||
event_name = "grok_execution_unavailable",
|
||||
log_type = "ops",
|
||||
trace_id = %trace_id,
|
||||
request_id = %plan_request_id_for_log,
|
||||
@@ -1197,7 +1263,7 @@ async fn execute_execution_runtime_sync_impl(
|
||||
model_name,
|
||||
candidate_index = candidate_index.as_str(),
|
||||
error = %err,
|
||||
"gateway ChatGPT-Web image execution unavailable"
|
||||
"gateway Grok execution unavailable"
|
||||
);
|
||||
let terminal_unix_secs = current_request_candidate_unix_ms();
|
||||
record_local_request_candidate_status(
|
||||
@@ -1207,7 +1273,7 @@ async fn execute_execution_runtime_sync_impl(
|
||||
SchedulerRequestCandidateStatusUpdate {
|
||||
status: RequestCandidateStatus::Failed,
|
||||
status_code: None,
|
||||
error_type: Some("chatgpt_web_image_execution_unavailable".to_string()),
|
||||
error_type: Some("grok_execution_unavailable".to_string()),
|
||||
error_message: Some(err.to_string()),
|
||||
latency_ms: None,
|
||||
started_at_unix_ms: Some(candidate_started_unix_secs),
|
||||
@@ -1264,30 +1330,72 @@ async fn execute_execution_runtime_sync_impl(
|
||||
.trim()
|
||||
.is_empty()
|
||||
{
|
||||
match maybe_execute_chatgpt_web_image_sync(state, &plan, report_context.as_ref()).await
|
||||
{
|
||||
match maybe_execute_grok_sync(&plan, report_context.as_ref()).await {
|
||||
Ok(Some(result)) => result,
|
||||
Ok(None) => match execute_direct_sync_runtime_candidate(
|
||||
Ok(None) => match maybe_execute_chatgpt_web_image_sync(
|
||||
state,
|
||||
&plan,
|
||||
report_context.as_ref(),
|
||||
trace_id,
|
||||
plan_kind,
|
||||
plan_request_id_for_log.as_str(),
|
||||
plan_candidate_id.as_deref(),
|
||||
provider_name.as_str(),
|
||||
endpoint_id.as_str(),
|
||||
key_id.as_str(),
|
||||
model_name.as_str(),
|
||||
candidate_index.as_str(),
|
||||
progress_snapshot.clone(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Ok(Some(result)) => result,
|
||||
Ok(None) => match execute_direct_sync_runtime_candidate(
|
||||
state,
|
||||
&plan,
|
||||
report_context.as_ref(),
|
||||
trace_id,
|
||||
plan_kind,
|
||||
plan_request_id_for_log.as_str(),
|
||||
plan_candidate_id.as_deref(),
|
||||
provider_name.as_str(),
|
||||
endpoint_id.as_str(),
|
||||
key_id.as_str(),
|
||||
model_name.as_str(),
|
||||
candidate_index.as_str(),
|
||||
progress_snapshot.clone(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(result) => result,
|
||||
Err(err) => {
|
||||
warn!(
|
||||
event_name = "sync_execution_runtime_unavailable",
|
||||
log_type = "ops",
|
||||
trace_id = %trace_id,
|
||||
request_id = %plan_request_id_for_log,
|
||||
candidate_id = ?plan_candidate_id,
|
||||
provider_name,
|
||||
endpoint_id,
|
||||
key_id,
|
||||
model_name,
|
||||
candidate_index = candidate_index.as_str(),
|
||||
error_type = err.error_type,
|
||||
error = %err.message,
|
||||
"gateway in-process sync execution unavailable"
|
||||
);
|
||||
let terminal_unix_secs = current_request_candidate_unix_ms();
|
||||
record_local_request_candidate_status(
|
||||
state,
|
||||
&plan,
|
||||
report_context.as_ref(),
|
||||
SchedulerRequestCandidateStatusUpdate {
|
||||
status: RequestCandidateStatus::Failed,
|
||||
status_code: err.status_code,
|
||||
error_type: Some(err.error_type.to_string()),
|
||||
error_message: Some(err.message),
|
||||
latency_ms: err.latency_ms,
|
||||
started_at_unix_ms: Some(candidate_started_unix_secs),
|
||||
finished_at_unix_ms: Some(terminal_unix_secs),
|
||||
},
|
||||
)
|
||||
.await;
|
||||
return Ok(None);
|
||||
}
|
||||
},
|
||||
Err(err) => {
|
||||
warn!(
|
||||
event_name = "sync_execution_runtime_unavailable",
|
||||
event_name = "chatgpt_web_image_execution_unavailable",
|
||||
log_type = "ops",
|
||||
trace_id = %trace_id,
|
||||
request_id = %plan_request_id_for_log,
|
||||
@@ -1297,9 +1405,8 @@ async fn execute_execution_runtime_sync_impl(
|
||||
key_id,
|
||||
model_name,
|
||||
candidate_index = candidate_index.as_str(),
|
||||
error_type = err.error_type,
|
||||
error = %err.message,
|
||||
"gateway in-process sync execution unavailable"
|
||||
error = %err,
|
||||
"gateway ChatGPT-Web image execution unavailable"
|
||||
);
|
||||
let terminal_unix_secs = current_request_candidate_unix_ms();
|
||||
record_local_request_candidate_status(
|
||||
@@ -1308,10 +1415,12 @@ async fn execute_execution_runtime_sync_impl(
|
||||
report_context.as_ref(),
|
||||
SchedulerRequestCandidateStatusUpdate {
|
||||
status: RequestCandidateStatus::Failed,
|
||||
status_code: err.status_code,
|
||||
error_type: Some(err.error_type.to_string()),
|
||||
error_message: Some(err.message),
|
||||
latency_ms: err.latency_ms,
|
||||
status_code: None,
|
||||
error_type: Some(
|
||||
"chatgpt_web_image_execution_unavailable".to_string(),
|
||||
),
|
||||
error_message: Some(err.to_string()),
|
||||
latency_ms: None,
|
||||
started_at_unix_ms: Some(candidate_started_unix_secs),
|
||||
finished_at_unix_ms: Some(terminal_unix_secs),
|
||||
},
|
||||
@@ -1322,7 +1431,7 @@ async fn execute_execution_runtime_sync_impl(
|
||||
},
|
||||
Err(err) => {
|
||||
warn!(
|
||||
event_name = "chatgpt_web_image_execution_unavailable",
|
||||
event_name = "grok_execution_unavailable",
|
||||
log_type = "ops",
|
||||
trace_id = %trace_id,
|
||||
request_id = %plan_request_id_for_log,
|
||||
@@ -1333,7 +1442,7 @@ async fn execute_execution_runtime_sync_impl(
|
||||
model_name,
|
||||
candidate_index = candidate_index.as_str(),
|
||||
error = %err,
|
||||
"gateway ChatGPT-Web image execution unavailable"
|
||||
"gateway Grok execution unavailable"
|
||||
);
|
||||
let terminal_unix_secs = current_request_candidate_unix_ms();
|
||||
record_local_request_candidate_status(
|
||||
@@ -1343,7 +1452,7 @@ async fn execute_execution_runtime_sync_impl(
|
||||
SchedulerRequestCandidateStatusUpdate {
|
||||
status: RequestCandidateStatus::Failed,
|
||||
status_code: None,
|
||||
error_type: Some("chatgpt_web_image_execution_unavailable".to_string()),
|
||||
error_type: Some("grok_execution_unavailable".to_string()),
|
||||
error_message: Some(err.to_string()),
|
||||
latency_ms: None,
|
||||
started_at_unix_ms: Some(candidate_started_unix_secs),
|
||||
|
||||
@@ -8,7 +8,8 @@ use aether_contracts::{
|
||||
ExecutionPlan, ExecutionResult, ExecutionTelemetry, ProxySnapshot, ResolvedTransportProfile,
|
||||
ResponseBody, EXECUTION_REQUEST_ACCEPT_INVALID_CERTS_HEADER,
|
||||
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
|
||||
TRANSPORT_BACKEND_REQWEST_RUSTLS, TRANSPORT_HTTP_MODE_HTTP1_ONLY,
|
||||
TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_BACKEND_REQWEST_RUSTLS,
|
||||
TRANSPORT_HTTP_MODE_HTTP1_ONLY,
|
||||
};
|
||||
use aether_data::repository::proxy_nodes::ProxyNodeTrafficMutation;
|
||||
use aether_http::{apply_http_client_config, HttpClientConfig};
|
||||
@@ -81,6 +82,52 @@ pub(crate) fn format_upstream_request_error(err: &reqwest::Error) -> String {
|
||||
detail
|
||||
}
|
||||
|
||||
pub(crate) fn format_wreq_upstream_request_error(err: &wreq::Error) -> String {
|
||||
let mut kinds = Vec::new();
|
||||
if err.is_connect() {
|
||||
kinds.push("connect");
|
||||
}
|
||||
if err.is_timeout() {
|
||||
kinds.push("timeout");
|
||||
}
|
||||
if err.is_redirect() {
|
||||
kinds.push("redirect");
|
||||
}
|
||||
if err.is_body() {
|
||||
kinds.push("body");
|
||||
}
|
||||
if err.is_decode() {
|
||||
kinds.push("decode");
|
||||
}
|
||||
if err.is_request() {
|
||||
kinds.push("request");
|
||||
}
|
||||
|
||||
let mut detail = err.to_string();
|
||||
let mut source = err.source();
|
||||
while let Some(cause) = source {
|
||||
let cause_text = cause.to_string();
|
||||
if !cause_text.is_empty() && !detail.contains(&cause_text) {
|
||||
detail.push_str(": ");
|
||||
detail.push_str(&cause_text);
|
||||
}
|
||||
source = cause.source();
|
||||
}
|
||||
|
||||
if let Some(uri) = err.uri() {
|
||||
detail.push_str(" [uri=");
|
||||
detail.push_str(&uri.to_string());
|
||||
detail.push(']');
|
||||
}
|
||||
if !kinds.is_empty() {
|
||||
detail.push_str(" [kind=");
|
||||
detail.push_str(&kinds.join(","));
|
||||
detail.push(']');
|
||||
}
|
||||
|
||||
detail
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub(crate) enum ExecutionRuntimeTransportError {
|
||||
#[error("stream execution is not supported for this plan")]
|
||||
@@ -107,6 +154,10 @@ pub(crate) enum ExecutionRuntimeTransportError {
|
||||
BodyEncode(serde_json::Error),
|
||||
#[error("failed to build HTTP client: {0}")]
|
||||
ClientBuild(reqwest::Error),
|
||||
#[error("failed to build browser impersonation HTTP client: {0}")]
|
||||
BrowserClientBuild(wreq::Error),
|
||||
#[error("browser impersonation response body failed: {0}")]
|
||||
BrowserBody(String),
|
||||
#[error("failed to execute upstream request: {0}")]
|
||||
UpstreamRequest(String),
|
||||
#[error("hub relay request failed: {0}")]
|
||||
@@ -136,7 +187,7 @@ struct RelayRequestMeta {
|
||||
pub(crate) struct DirectSyncExecutionRuntime;
|
||||
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
struct ExecutionTransportControls {
|
||||
pub(crate) struct ExecutionTransportControls {
|
||||
follow_redirects: Option<bool>,
|
||||
http1_only: bool,
|
||||
accept_invalid_certs: bool,
|
||||
@@ -144,6 +195,7 @@ struct ExecutionTransportControls {
|
||||
|
||||
pub(crate) enum DirectUpstreamResponse {
|
||||
Reqwest(reqwest::Response),
|
||||
BrowserWreq(wreq::Response),
|
||||
LocalTunnel(tunnel::DirectRelayResponse),
|
||||
}
|
||||
|
||||
@@ -172,11 +224,9 @@ impl DirectSyncExecutionRuntime {
|
||||
let started_at = Instant::now();
|
||||
let response = send_request(plan, body_bytes).await?;
|
||||
let ttfb_ms = started_at.elapsed().as_millis() as u64;
|
||||
let status_code = response.status().as_u16();
|
||||
let headers = collect_response_headers(response.headers());
|
||||
let body_bytes = response.bytes().await.map_err(|err| {
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err))
|
||||
})?;
|
||||
let status_code = response.status_code();
|
||||
let headers = response.headers();
|
||||
let body_bytes = response.bytes().await?;
|
||||
let decoded_body_bytes = decode_response_body_bytes(&headers, &body_bytes)
|
||||
.unwrap_or_else(|| body_bytes.to_vec());
|
||||
let elapsed_ms = started_at.elapsed().as_millis() as u64;
|
||||
@@ -230,8 +280,8 @@ impl DirectSyncExecutionRuntime {
|
||||
|
||||
let started_at = Instant::now();
|
||||
let response = send_request(plan, body_bytes).await?;
|
||||
let status_code = response.status().as_u16();
|
||||
let headers = collect_response_headers(response.headers());
|
||||
let status_code = response.status_code();
|
||||
let headers = response.headers();
|
||||
|
||||
let stream_summary_report_context = build_stream_summary_report_context(plan);
|
||||
|
||||
@@ -242,7 +292,7 @@ impl DirectSyncExecutionRuntime {
|
||||
headers,
|
||||
provider_api_format: plan.provider_api_format.clone(),
|
||||
stream_summary_report_context,
|
||||
response: DirectUpstreamResponse::Reqwest(response),
|
||||
response: response.into_direct_upstream_response(),
|
||||
started_at,
|
||||
})
|
||||
}
|
||||
@@ -252,6 +302,15 @@ pub(crate) async fn execute_sync_plan(
|
||||
state: &AppState,
|
||||
trace_id: Option<&str>,
|
||||
plan: &ExecutionPlan,
|
||||
) -> Result<ExecutionResult, GatewayError> {
|
||||
execute_sync_plan_with_report_context(state, trace_id, plan, None).await
|
||||
}
|
||||
|
||||
pub(crate) async fn execute_sync_plan_with_report_context(
|
||||
state: &AppState,
|
||||
trace_id: Option<&str>,
|
||||
plan: &ExecutionPlan,
|
||||
report_context: Option<&serde_json::Value>,
|
||||
) -> Result<ExecutionResult, GatewayError> {
|
||||
#[cfg(test)]
|
||||
{
|
||||
@@ -275,6 +334,18 @@ pub(crate) async fn execute_sync_plan(
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()));
|
||||
}
|
||||
|
||||
match super::grok::maybe_execute_grok_sync(plan, report_context).await {
|
||||
Ok(Some(result)) => {
|
||||
record_manual_proxy_request_outcome(state, plan, result.status_code).await;
|
||||
return Ok(result);
|
||||
}
|
||||
Ok(None) => {}
|
||||
Err(err) => {
|
||||
record_manual_proxy_request_failure(state, plan).await;
|
||||
return Err(GatewayError::Internal(err.to_string()));
|
||||
}
|
||||
}
|
||||
|
||||
let _ = trace_id;
|
||||
match DirectSyncExecutionRuntime::new().execute_sync(plan).await {
|
||||
Ok(result) => {
|
||||
@@ -554,7 +625,7 @@ fn build_direct_tunnel_request_meta(
|
||||
pub(crate) async fn send_request(
|
||||
plan: &ExecutionPlan,
|
||||
body_bytes: Vec<u8>,
|
||||
) -> Result<reqwest::Response, ExecutionRuntimeTransportError> {
|
||||
) -> Result<DirectHttpResponse, ExecutionRuntimeTransportError> {
|
||||
if let Some(detail) = gateway_frontdoor_self_loop_guard_error(plan.url.as_str()) {
|
||||
return Err(ExecutionRuntimeTransportError::UpstreamRequest(detail));
|
||||
}
|
||||
@@ -572,6 +643,18 @@ pub(crate) async fn send_request(
|
||||
.and_then(|timeouts| timeouts.total_ms)
|
||||
.map(Duration::from_millis);
|
||||
|
||||
if transport_profile_uses_browser_wreq(plan.transport_profile.as_ref()) {
|
||||
return send_via_browser_wreq_transport(
|
||||
plan,
|
||||
method,
|
||||
headers,
|
||||
body_bytes,
|
||||
total_timeout,
|
||||
transport_controls,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
if let Some(node_id) = resolve_tunnel_node_id(plan.proxy.as_ref()) {
|
||||
return send_via_tunnel_relay(
|
||||
plan,
|
||||
@@ -582,7 +665,8 @@ pub(crate) async fn send_request(
|
||||
total_timeout,
|
||||
transport_controls,
|
||||
)
|
||||
.await;
|
||||
.await
|
||||
.map(DirectHttpResponse::Reqwest);
|
||||
}
|
||||
|
||||
let client = build_client(
|
||||
@@ -596,9 +680,95 @@ pub(crate) async fn send_request(
|
||||
if let Some(timeout) = total_timeout {
|
||||
request = request.timeout(timeout);
|
||||
}
|
||||
request.send().await.map_err(|err| {
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err))
|
||||
})
|
||||
request
|
||||
.send()
|
||||
.await
|
||||
.map(DirectHttpResponse::Reqwest)
|
||||
.map_err(|err| {
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err))
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) enum DirectHttpResponse {
|
||||
Reqwest(reqwest::Response),
|
||||
BrowserWreq(wreq::Response),
|
||||
}
|
||||
|
||||
impl DirectHttpResponse {
|
||||
pub(crate) fn status_code(&self) -> u16 {
|
||||
match self {
|
||||
DirectHttpResponse::Reqwest(response) => response.status().as_u16(),
|
||||
DirectHttpResponse::BrowserWreq(response) => response.status().as_u16(),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn headers(&self) -> BTreeMap<String, String> {
|
||||
match self {
|
||||
DirectHttpResponse::Reqwest(response) => collect_response_headers(response.headers()),
|
||||
DirectHttpResponse::BrowserWreq(response) => {
|
||||
collect_response_headers(response.headers())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn bytes(self) -> Result<Bytes, ExecutionRuntimeTransportError> {
|
||||
match self {
|
||||
DirectHttpResponse::Reqwest(response) => response.bytes().await.map_err(|err| {
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(format_upstream_request_error(&err))
|
||||
}),
|
||||
DirectHttpResponse::BrowserWreq(response) => response.bytes().await.map_err(|err| {
|
||||
ExecutionRuntimeTransportError::BrowserBody(format_wreq_upstream_request_error(
|
||||
&err,
|
||||
))
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn into_direct_upstream_response(self) -> DirectUpstreamResponse {
|
||||
match self {
|
||||
DirectHttpResponse::Reqwest(response) => DirectUpstreamResponse::Reqwest(response),
|
||||
DirectHttpResponse::BrowserWreq(response) => {
|
||||
DirectUpstreamResponse::BrowserWreq(response)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn send_via_browser_wreq_transport(
|
||||
plan: &ExecutionPlan,
|
||||
method: reqwest::Method,
|
||||
headers: HeaderMap,
|
||||
body_bytes: Vec<u8>,
|
||||
total_timeout: Option<Duration>,
|
||||
transport_controls: ExecutionTransportControls,
|
||||
) -> Result<DirectHttpResponse, ExecutionRuntimeTransportError> {
|
||||
let profile = plan.transport_profile.as_ref().ok_or_else(|| {
|
||||
ExecutionRuntimeTransportError::UnsupportedTransportProfile(String::new())
|
||||
})?;
|
||||
let client = build_browser_wreq_client(
|
||||
plan.timeouts.as_ref(),
|
||||
plan.proxy.as_ref(),
|
||||
profile,
|
||||
transport_controls,
|
||||
)?;
|
||||
let method = wreq::Method::from_bytes(method.as_str().as_bytes())
|
||||
.map_err(ExecutionRuntimeTransportError::InvalidMethod)?;
|
||||
let mut request = client
|
||||
.request(method, plan.url.as_str())
|
||||
.headers(headers)
|
||||
.body(body_bytes);
|
||||
if let Some(timeout) = total_timeout {
|
||||
request = request.timeout(timeout);
|
||||
}
|
||||
request
|
||||
.send()
|
||||
.await
|
||||
.map(DirectHttpResponse::BrowserWreq)
|
||||
.map_err(|err| {
|
||||
ExecutionRuntimeTransportError::UpstreamRequest(format_wreq_upstream_request_error(
|
||||
&err,
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
async fn send_via_tunnel_relay(
|
||||
@@ -905,6 +1075,96 @@ fn build_client(
|
||||
.map_err(ExecutionRuntimeTransportError::ClientBuild)
|
||||
}
|
||||
|
||||
pub(crate) fn build_browser_wreq_client(
|
||||
timeouts: Option<&aether_contracts::ExecutionTimeouts>,
|
||||
proxy: Option<&ProxySnapshot>,
|
||||
transport_profile: &ResolvedTransportProfile,
|
||||
transport_controls: ExecutionTransportControls,
|
||||
) -> Result<wreq::Client, ExecutionRuntimeTransportError> {
|
||||
let emulation = browser_wreq_emulation_from_profile(transport_profile)?;
|
||||
let mut builder = wreq::Client::builder().emulation(emulation);
|
||||
if transport_controls.follow_redirects == Some(true) {
|
||||
builder = builder.redirect(wreq::redirect::Policy::limited(10));
|
||||
}
|
||||
if transport_controls.http1_only || transport_profile_http1_only(Some(transport_profile)) {
|
||||
builder = builder.http1_only();
|
||||
}
|
||||
if transport_controls.accept_invalid_certs {
|
||||
builder = builder.cert_verification(false).verify_hostname(false);
|
||||
}
|
||||
if let Some(connect_ms) = timeouts.and_then(|timeouts| timeouts.connect_ms) {
|
||||
builder = builder.connect_timeout(Duration::from_millis(connect_ms));
|
||||
}
|
||||
if let Some(total_ms) = timeouts.and_then(|timeouts| timeouts.total_ms) {
|
||||
builder = builder.timeout(Duration::from_millis(total_ms));
|
||||
}
|
||||
if let Some(read_ms) = timeouts.and_then(|timeouts| timeouts.read_ms) {
|
||||
builder = builder.read_timeout(Duration::from_millis(read_ms));
|
||||
}
|
||||
if let Some(proxy_url) = resolve_proxy_url(proxy)? {
|
||||
let proxy = wreq::Proxy::all(proxy_url.as_str())
|
||||
.map_err(ExecutionRuntimeTransportError::BrowserClientBuild)?;
|
||||
builder = builder.proxy(proxy);
|
||||
}
|
||||
builder
|
||||
.build()
|
||||
.map_err(ExecutionRuntimeTransportError::BrowserClientBuild)
|
||||
}
|
||||
|
||||
fn browser_wreq_emulation_from_profile(
|
||||
profile: &ResolvedTransportProfile,
|
||||
) -> Result<wreq_util::Emulation, ExecutionRuntimeTransportError> {
|
||||
match normalize_browser_profile_name(browser_transport_profile_name(profile)).as_str() {
|
||||
"chrome100" => Ok(wreq_util::Emulation::Chrome100),
|
||||
"chrome101" => Ok(wreq_util::Emulation::Chrome101),
|
||||
"chrome104" => Ok(wreq_util::Emulation::Chrome104),
|
||||
"chrome105" => Ok(wreq_util::Emulation::Chrome105),
|
||||
"chrome106" => Ok(wreq_util::Emulation::Chrome106),
|
||||
"chrome107" => Ok(wreq_util::Emulation::Chrome107),
|
||||
"chrome108" => Ok(wreq_util::Emulation::Chrome108),
|
||||
"chrome109" => Ok(wreq_util::Emulation::Chrome109),
|
||||
"chrome110" => Ok(wreq_util::Emulation::Chrome110),
|
||||
"chrome114" => Ok(wreq_util::Emulation::Chrome114),
|
||||
"chrome116" => Ok(wreq_util::Emulation::Chrome116),
|
||||
"chrome117" => Ok(wreq_util::Emulation::Chrome117),
|
||||
"chrome118" => Ok(wreq_util::Emulation::Chrome118),
|
||||
"chrome119" => Ok(wreq_util::Emulation::Chrome119),
|
||||
"chrome120" => Ok(wreq_util::Emulation::Chrome120),
|
||||
"chrome123" => Ok(wreq_util::Emulation::Chrome123),
|
||||
"chrome124" => Ok(wreq_util::Emulation::Chrome124),
|
||||
"chrome126" => Ok(wreq_util::Emulation::Chrome126),
|
||||
"chrome127" => Ok(wreq_util::Emulation::Chrome127),
|
||||
"chrome128" => Ok(wreq_util::Emulation::Chrome128),
|
||||
"chrome129" => Ok(wreq_util::Emulation::Chrome129),
|
||||
"chrome130" => Ok(wreq_util::Emulation::Chrome130),
|
||||
"chrome131" => Ok(wreq_util::Emulation::Chrome131),
|
||||
"chrome132" => Ok(wreq_util::Emulation::Chrome132),
|
||||
"chrome133" => Ok(wreq_util::Emulation::Chrome133),
|
||||
"chrome134" => Ok(wreq_util::Emulation::Chrome134),
|
||||
"chrome135" => Ok(wreq_util::Emulation::Chrome135),
|
||||
"chrome136" => Ok(wreq_util::Emulation::Chrome136),
|
||||
"chrome137" => Ok(wreq_util::Emulation::Chrome137),
|
||||
"chrome138" => Ok(wreq_util::Emulation::Chrome138),
|
||||
"chrome139" => Ok(wreq_util::Emulation::Chrome139),
|
||||
"chrome140" => Ok(wreq_util::Emulation::Chrome140),
|
||||
"chrome141" => Ok(wreq_util::Emulation::Chrome141),
|
||||
"chrome142" => Ok(wreq_util::Emulation::Chrome142),
|
||||
"chrome143" => Ok(wreq_util::Emulation::Chrome143),
|
||||
"chrome144" => Ok(wreq_util::Emulation::Chrome144),
|
||||
"chrome145" => Ok(wreq_util::Emulation::Chrome145),
|
||||
other => Err(ExecutionRuntimeTransportError::UnsupportedTransportProfile(
|
||||
format!("browser_wreq:{other}"),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_browser_profile_name(value: String) -> String {
|
||||
value
|
||||
.trim()
|
||||
.to_ascii_lowercase()
|
||||
.replace(['_', '-', ' '], "")
|
||||
}
|
||||
|
||||
fn validate_reqwest_transport_profile(
|
||||
transport_profile: Option<&ResolvedTransportProfile>,
|
||||
) -> Result<(), ExecutionRuntimeTransportError> {
|
||||
@@ -923,6 +1183,56 @@ fn validate_reqwest_transport_profile(
|
||||
))
|
||||
}
|
||||
|
||||
fn transport_profile_uses_browser_wreq(
|
||||
transport_profile: Option<&ResolvedTransportProfile>,
|
||||
) -> bool {
|
||||
transport_profile
|
||||
.map(|profile| {
|
||||
profile
|
||||
.backend
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(TRANSPORT_BACKEND_BROWSER_WREQ)
|
||||
})
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn browser_transport_profile_name(profile: &ResolvedTransportProfile) -> String {
|
||||
profile
|
||||
.extra
|
||||
.as_ref()
|
||||
.and_then(|value| {
|
||||
value
|
||||
.get("browser_profile")
|
||||
.or_else(|| value.get("impersonate"))
|
||||
.and_then(Value::as_str)
|
||||
})
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.or_else(|| {
|
||||
profile
|
||||
.profile_id
|
||||
.trim()
|
||||
.is_empty()
|
||||
.then_some("chrome136".to_string())
|
||||
.or_else(|| Some(profile.profile_id.trim().to_string()))
|
||||
})
|
||||
.unwrap_or_else(|| "chrome136".to_string())
|
||||
}
|
||||
|
||||
fn insert_browser_control_header(
|
||||
headers: &mut HeaderMap,
|
||||
name: &'static str,
|
||||
value: &str,
|
||||
) -> Result<(), ExecutionRuntimeTransportError> {
|
||||
headers.insert(
|
||||
HeaderName::from_static(name),
|
||||
HeaderValue::from_str(value)
|
||||
.map_err(|_| ExecutionRuntimeTransportError::InvalidHeaderValue(name.to_string()))?,
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn transport_profile_http1_only(transport_profile: Option<&ResolvedTransportProfile>) -> bool {
|
||||
transport_profile
|
||||
.map(|profile| {
|
||||
@@ -991,7 +1301,7 @@ fn resolve_proxy_url(
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
fn build_request_headers(
|
||||
pub(crate) fn build_request_headers(
|
||||
headers: &BTreeMap<String, String>,
|
||||
content_encoding: Option<&str>,
|
||||
allow_passthrough_content_encoding: bool,
|
||||
@@ -1180,22 +1490,22 @@ mod tests {
|
||||
use aether_contracts::{
|
||||
ExecutionPlan, ExecutionTimeouts, ProxySnapshot, RequestBody, ResolvedTransportProfile,
|
||||
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER, EXECUTION_REQUEST_HTTP1_ONLY_HEADER,
|
||||
TRANSPORT_BACKEND_REQWEST_RUSTLS,
|
||||
TRANSPORT_BACKEND_BROWSER_WREQ, TRANSPORT_BACKEND_REQWEST_RUSTLS,
|
||||
};
|
||||
use aether_data::repository::proxy_nodes::{
|
||||
InMemoryProxyNodeRepository, ProxyNodeReadRepository, StoredProxyNode,
|
||||
};
|
||||
use axum::body::Bytes;
|
||||
use axum::body::{Body, Bytes};
|
||||
use axum::extract::ws::Message;
|
||||
use axum::extract::Path;
|
||||
use axum::http::HeaderMap as AxumHeaderMap;
|
||||
use axum::routing::post;
|
||||
use axum::routing::{any, post};
|
||||
use axum::{Json, Router};
|
||||
use serde_json::json;
|
||||
use tokio::sync::watch;
|
||||
|
||||
use super::{
|
||||
build_client, build_request_headers, execute_sync_plan,
|
||||
build_browser_wreq_client, build_client, build_request_headers, execute_sync_plan,
|
||||
record_manual_proxy_request_failure, record_manual_proxy_request_outcome,
|
||||
record_manual_proxy_request_success, record_manual_proxy_stream_error,
|
||||
resolve_execution_transport_controls, DirectSyncExecutionRuntime,
|
||||
@@ -1431,6 +1741,228 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn direct_sync_execution_runtime_routes_browser_wreq_transport_in_process() {
|
||||
async fn browser_upstream(headers: AxumHeaderMap, body: Bytes) -> axum::response::Response {
|
||||
assert_eq!(
|
||||
headers
|
||||
.get("content-type")
|
||||
.and_then(|value| value.to_str().ok()),
|
||||
Some("application/json")
|
||||
);
|
||||
assert!(
|
||||
headers
|
||||
.get(EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER)
|
||||
.is_none(),
|
||||
"internal execution control headers must not leak upstream"
|
||||
);
|
||||
assert_eq!(body.as_ref(), br#"{"modelName":"auto"}"#);
|
||||
axum::response::Response::builder()
|
||||
.status(http::StatusCode::ACCEPTED)
|
||||
.header("content-type", "application/json")
|
||||
.body(Body::from(
|
||||
json!({
|
||||
"ok": true,
|
||||
"via": "browser_wreq"
|
||||
})
|
||||
.to_string(),
|
||||
))
|
||||
.expect("response should build")
|
||||
}
|
||||
|
||||
let listener = crate::test_support::bind_loopback_listener()
|
||||
.await
|
||||
.expect("listener should bind");
|
||||
let addr = listener.local_addr().expect("local addr should resolve");
|
||||
let app = Router::new().route("/request", any(browser_upstream));
|
||||
let server = tokio::spawn(async move {
|
||||
axum::serve(listener, app)
|
||||
.await
|
||||
.expect("test server should run");
|
||||
});
|
||||
|
||||
let plan = ExecutionPlan {
|
||||
request_id: "req-browser-wreq".into(),
|
||||
candidate_id: None,
|
||||
provider_name: Some("grok".into()),
|
||||
provider_id: "provider-1".into(),
|
||||
endpoint_id: "endpoint-1".into(),
|
||||
key_id: "key-1".into(),
|
||||
method: "POST".into(),
|
||||
url: format!("http://{addr}/request"),
|
||||
headers: BTreeMap::from([
|
||||
("content-type".into(), "application/json".into()),
|
||||
(
|
||||
EXECUTION_REQUEST_FOLLOW_REDIRECTS_HEADER.into(),
|
||||
"true".into(),
|
||||
),
|
||||
]),
|
||||
content_type: Some("application/json".into()),
|
||||
content_encoding: None,
|
||||
body: RequestBody::from_json(json!({"modelName":"auto"})),
|
||||
stream: false,
|
||||
client_api_format: "openai:responses".into(),
|
||||
provider_api_format: "grok:rate_limits".into(),
|
||||
model_name: Some("grok-quota".into()),
|
||||
proxy: None,
|
||||
transport_profile: Some(ResolvedTransportProfile {
|
||||
profile_id: "chrome136".into(),
|
||||
backend: TRANSPORT_BACKEND_BROWSER_WREQ.into(),
|
||||
http_mode: "auto".into(),
|
||||
pool_scope: "key".into(),
|
||||
header_fingerprint: None,
|
||||
extra: Some(json!({
|
||||
"browser_profile": "chrome136"
|
||||
})),
|
||||
}),
|
||||
timeouts: Some(ExecutionTimeouts {
|
||||
total_ms: Some(5_000),
|
||||
..ExecutionTimeouts::default()
|
||||
}),
|
||||
};
|
||||
|
||||
let result = DirectSyncExecutionRuntime::new()
|
||||
.execute_sync(&plan)
|
||||
.await
|
||||
.expect("browser wreq transport plan should execute in-process");
|
||||
|
||||
server.abort();
|
||||
|
||||
assert_eq!(result.status_code, http::StatusCode::ACCEPTED.as_u16());
|
||||
assert_eq!(
|
||||
result
|
||||
.body
|
||||
.and_then(|body| body.json_body)
|
||||
.and_then(|body| body.get("via").cloned()),
|
||||
Some(json!("browser_wreq"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn browser_wreq_transport_rejects_unknown_profile() {
|
||||
let profile = ResolvedTransportProfile {
|
||||
profile_id: "firefox999".into(),
|
||||
backend: TRANSPORT_BACKEND_BROWSER_WREQ.into(),
|
||||
http_mode: "auto".into(),
|
||||
pool_scope: "key".into(),
|
||||
header_fingerprint: None,
|
||||
extra: None,
|
||||
};
|
||||
|
||||
let error = match build_browser_wreq_client(
|
||||
None,
|
||||
None,
|
||||
&profile,
|
||||
ExecutionTransportControls::default(),
|
||||
) {
|
||||
Ok(_) => panic!("unknown browser profile should fail loudly"),
|
||||
Err(error) => error,
|
||||
};
|
||||
|
||||
assert!(matches!(
|
||||
error,
|
||||
ExecutionRuntimeTransportError::UnsupportedTransportProfile(backend)
|
||||
if backend == "browser_wreq:firefox999"
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn execute_sync_plan_routes_grok_marker_through_grok_runtime() {
|
||||
let listener = crate::test_support::bind_loopback_listener()
|
||||
.await
|
||||
.expect("listener should bind");
|
||||
let addr = listener.local_addr().expect("local addr should resolve");
|
||||
let app = Router::new().route(
|
||||
"/rest/app-chat/conversations/new",
|
||||
post(|body: Bytes| async move {
|
||||
let body_json: serde_json::Value =
|
||||
serde_json::from_slice(&body).expect("request body should be json");
|
||||
if body_json.get("message").and_then(serde_json::Value::as_str)
|
||||
!= Some("[user]: hello")
|
||||
{
|
||||
return (
|
||||
axum::http::StatusCode::BAD_REQUEST,
|
||||
Json(json!({
|
||||
"error": {
|
||||
"message": "expected grok app-chat message",
|
||||
"body": body_json,
|
||||
}
|
||||
})),
|
||||
);
|
||||
}
|
||||
(
|
||||
axum::http::StatusCode::OK,
|
||||
Json(json!({
|
||||
"result": {
|
||||
"response": {
|
||||
"token": "pong",
|
||||
"messageTag": "final"
|
||||
}
|
||||
}
|
||||
})),
|
||||
)
|
||||
}),
|
||||
);
|
||||
let server = tokio::spawn(async move {
|
||||
axum::serve(listener, app)
|
||||
.await
|
||||
.expect("test server should run");
|
||||
});
|
||||
let plan = ExecutionPlan {
|
||||
request_id: "req-grok-runtime".into(),
|
||||
candidate_id: Some("cand-grok".into()),
|
||||
provider_name: Some("grok".into()),
|
||||
provider_id: "provider-grok".into(),
|
||||
endpoint_id: "endpoint-grok".into(),
|
||||
key_id: "key-grok".into(),
|
||||
method: "POST".into(),
|
||||
url: format!("http://{addr}/rest/app-chat/conversations/new"),
|
||||
headers: BTreeMap::from([
|
||||
("content-type".into(), "application/json".into()),
|
||||
(
|
||||
aether_provider_transport::GROK_INTERNAL_HEADER.into(),
|
||||
"1".into(),
|
||||
),
|
||||
]),
|
||||
content_type: Some("application/json".into()),
|
||||
content_encoding: None,
|
||||
body: RequestBody::from_json(json!({
|
||||
"model": "grok-4.20-0309-non-reasoning",
|
||||
"messages": [{"role": "user", "content": "hello"}],
|
||||
})),
|
||||
stream: true,
|
||||
client_api_format: "openai:chat".into(),
|
||||
provider_api_format: "openai:chat".into(),
|
||||
model_name: Some("grok-4.20-0309-non-reasoning".into()),
|
||||
proxy: None,
|
||||
transport_profile: None,
|
||||
timeouts: Some(ExecutionTimeouts {
|
||||
connect_ms: Some(5_000),
|
||||
total_ms: Some(5_000),
|
||||
..ExecutionTimeouts::default()
|
||||
}),
|
||||
};
|
||||
let report_context = json!({"mapped_model": "grok-4.20-fast"});
|
||||
|
||||
let result = super::super::grok::maybe_execute_grok_sync(&plan, Some(&report_context))
|
||||
.await
|
||||
.expect("grok runtime plan should execute")
|
||||
.expect("grok runtime should handle marked plan");
|
||||
|
||||
server.abort();
|
||||
|
||||
assert_eq!(result.status_code, http::StatusCode::OK.as_u16());
|
||||
assert_eq!(
|
||||
result
|
||||
.body
|
||||
.and_then(|body| body.json_body)
|
||||
.and_then(|body| body["choices"][0]["message"]["content"]
|
||||
.as_str()
|
||||
.map(str::to_string)),
|
||||
Some("pong".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn execute_sync_plan_records_manual_proxy_success() {
|
||||
let repository = Arc::new(InMemoryProxyNodeRepository::seed(vec![
|
||||
|
||||
@@ -38,7 +38,7 @@ pub(super) async fn build_admin_create_api_key_install_session_response(
|
||||
Err(_) => {
|
||||
return Ok(build_admin_api_keys_bad_request_response(
|
||||
"请求数据验证失败",
|
||||
))
|
||||
));
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -16,6 +16,7 @@ use axum::{
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
use tracing::warn;
|
||||
|
||||
pub(super) async fn maybe_build_local_admin_payment_orders_response(
|
||||
state: &AdminAppState<'_>,
|
||||
@@ -211,6 +212,19 @@ async fn build_admin_payment_credit_order_response(
|
||||
.await?
|
||||
{
|
||||
crate::AdminWalletMutationOutcome::Applied((order, credited)) => {
|
||||
if credited {
|
||||
if let Err(err) = state
|
||||
.app()
|
||||
.apply_referral_rewards_for_payment_order_id(&order.id)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
error = ?err,
|
||||
order_id = %order.id,
|
||||
"failed to apply referral rewards for admin-credited payment order"
|
||||
);
|
||||
}
|
||||
}
|
||||
Ok(attach_admin_audit_response(
|
||||
Json(json!({
|
||||
"order": build_admin_payment_order_payload(&order),
|
||||
|
||||
@@ -14,6 +14,7 @@ use axum::{
|
||||
Json,
|
||||
};
|
||||
use serde_json::json;
|
||||
use tracing::warn;
|
||||
|
||||
pub(in super::super) async fn build_admin_wallet_complete_refund_response(
|
||||
state: &AdminAppState<'_>,
|
||||
@@ -86,6 +87,20 @@ pub(in super::super) async fn build_admin_wallet_complete_refund_response(
|
||||
.await?
|
||||
{
|
||||
crate::AdminWalletMutationOutcome::Applied(refund) => {
|
||||
if let Some(order_id) = refund.payment_order_id.as_deref() {
|
||||
if let Err(err) = state
|
||||
.app()
|
||||
.reverse_referral_rewards_for_order(order_id, refund.amount_usd)
|
||||
.await
|
||||
{
|
||||
warn!(
|
||||
error = ?err,
|
||||
order_id = %order_id,
|
||||
refund_id = %refund.id,
|
||||
"failed to reverse referral rewards for completed refund"
|
||||
);
|
||||
}
|
||||
}
|
||||
let response = Json(json!({
|
||||
"refund": build_admin_wallet_refund_payload(&wallet, &owner, &refund),
|
||||
}))
|
||||
|
||||
@@ -6,6 +6,8 @@ pub(super) mod features;
|
||||
mod model;
|
||||
pub(super) mod observability;
|
||||
pub(super) mod provider;
|
||||
mod referrals;
|
||||
mod routing;
|
||||
mod system;
|
||||
mod users;
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ use super::route_filters::{
|
||||
};
|
||||
use crate::constants::INTERNAL_GATEWAY_PATH_PREFIXES;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::build_admin_usage_counter_health_payload;
|
||||
use crate::GatewayError;
|
||||
use aether_admin::observability::monitoring::{
|
||||
admin_monitoring_bad_request_response, admin_monitoring_user_behavior_user_id_from_path,
|
||||
@@ -189,6 +190,13 @@ pub(super) async fn build_admin_monitoring_system_status_response(
|
||||
)
|
||||
.unwrap_or(usize::MAX);
|
||||
let tunnel = state.tunnel.stats();
|
||||
let usage_counter_snapshot = state
|
||||
.data
|
||||
.read_usage_counter_health()
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
let usage_counter =
|
||||
build_admin_usage_counter_health_payload(&usage_counter_snapshot, now_unix_secs);
|
||||
|
||||
Ok(build_admin_monitoring_system_status_payload_response(
|
||||
now,
|
||||
@@ -206,5 +214,6 @@ pub(super) async fn build_admin_monitoring_system_status_response(
|
||||
tunnel.active_streams,
|
||||
INTERNAL_GATEWAY_PATH_PREFIXES,
|
||||
recent_errors,
|
||||
usage_counter,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -249,9 +249,10 @@ async fn admin_monitoring_resilience_status_returns_local_payload() {
|
||||
let recommendations = payload["recommendations"]
|
||||
.as_array()
|
||||
.expect("recommendations should be array");
|
||||
assert!(recommendations.iter().any(|item| item
|
||||
.as_str()
|
||||
.is_some_and(|value| value.contains("prod-key"))));
|
||||
assert!(recommendations.iter().any(|item| {
|
||||
item.as_str()
|
||||
.is_some_and(|value| value.contains("prod-key"))
|
||||
}));
|
||||
assert!(payload["timestamp"].as_str().is_some());
|
||||
}
|
||||
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
use super::range::{build_comparison_range, parse_bounded_u32};
|
||||
use super::resolve_admin_usage_time_range;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::{query_param_optional_bool, query_param_value};
|
||||
use crate::handlers::admin::shared::{
|
||||
build_admin_usage_counter_health_payload, query_param_optional_bool, query_param_value,
|
||||
};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::observability::stats::{
|
||||
admin_stats_bad_request_response, admin_stats_comparison_empty_response,
|
||||
@@ -21,6 +23,22 @@ use aether_data_contracts::repository::usage::{
|
||||
};
|
||||
use axum::{body::Body, http, response::Response};
|
||||
|
||||
async fn build_usage_counter_health_payload(
|
||||
state: &AdminAppState<'_>,
|
||||
) -> Result<serde_json::Value, GatewayError> {
|
||||
let now_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
|
||||
let snapshot = state
|
||||
.as_ref()
|
||||
.data
|
||||
.read_usage_counter_health()
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))?;
|
||||
Ok(build_admin_usage_counter_health_payload(
|
||||
&snapshot,
|
||||
now_unix_secs,
|
||||
))
|
||||
}
|
||||
|
||||
fn usage_summary_to_admin_stats_aggregate(
|
||||
summary: &aether_data_contracts::repository::usage::StoredUsageAuditSummary,
|
||||
) -> AdminStatsAggregate {
|
||||
@@ -211,13 +229,18 @@ pub(super) async fn maybe_build_local_admin_stats_analytics_response(
|
||||
Ok(value) => u64::from(value.unwrap_or(10_000)),
|
||||
Err(detail) => return Ok(Some(admin_stats_bad_request_response(detail))),
|
||||
};
|
||||
let usage_counter = build_usage_counter_health_payload(state).await?;
|
||||
if !state.has_usage_data_reader() {
|
||||
return Ok(Some(admin_stats_provider_performance_empty_response()));
|
||||
return Ok(Some(admin_stats_provider_performance_empty_response(
|
||||
usage_counter,
|
||||
)));
|
||||
}
|
||||
|
||||
let Some((created_from_unix_secs, created_until_unix_secs)) = time_range.to_unix_bounds()
|
||||
else {
|
||||
return Ok(Some(admin_stats_provider_performance_empty_response()));
|
||||
return Ok(Some(admin_stats_provider_performance_empty_response(
|
||||
usage_counter,
|
||||
)));
|
||||
};
|
||||
let performance = state
|
||||
.summarize_usage_provider_performance(&UsageProviderPerformanceQuery {
|
||||
@@ -240,6 +263,7 @@ pub(super) async fn maybe_build_local_admin_stats_analytics_response(
|
||||
.await?;
|
||||
return Ok(Some(build_admin_stats_provider_performance_response(
|
||||
&performance,
|
||||
usage_counter,
|
||||
)));
|
||||
}
|
||||
|
||||
|
||||
@@ -5,12 +5,13 @@ use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
use crate::handlers::admin::shared::query_param_value;
|
||||
use crate::GatewayError;
|
||||
use aether_admin::observability::usage::{
|
||||
admin_usage_bad_request_response, admin_usage_data_unavailable_response,
|
||||
admin_usage_has_fallback, admin_usage_is_failed, admin_usage_matches_search,
|
||||
admin_usage_matches_username, admin_usage_parse_ids, admin_usage_parse_limit,
|
||||
admin_usage_parse_offset, admin_usage_provider_key_name, admin_usage_record_json,
|
||||
build_admin_usage_active_requests_response, build_admin_usage_records_response,
|
||||
build_admin_usage_summary_stats_response_from_summary, ADMIN_USAGE_DATA_UNAVAILABLE_DETAIL,
|
||||
admin_usage_bad_request_response, admin_usage_client_family,
|
||||
admin_usage_data_unavailable_response, admin_usage_has_fallback, admin_usage_is_failed,
|
||||
admin_usage_matches_search, admin_usage_matches_username, admin_usage_parse_ids,
|
||||
admin_usage_parse_limit, admin_usage_parse_offset, admin_usage_provider_key_name,
|
||||
admin_usage_record_json, build_admin_usage_active_requests_response,
|
||||
build_admin_usage_records_response, build_admin_usage_summary_stats_response_from_summary,
|
||||
ADMIN_USAGE_DATA_UNAVAILABLE_DETAIL,
|
||||
};
|
||||
use aether_data::repository::users::StoredUserSummary;
|
||||
use aether_data_contracts::repository::{
|
||||
@@ -263,6 +264,19 @@ fn admin_usage_matches_attempt_status(
|
||||
}
|
||||
}
|
||||
|
||||
fn admin_usage_matches_client_family(
|
||||
item: &StoredRequestUsageAudit,
|
||||
client_family: Option<&str>,
|
||||
) -> bool {
|
||||
let Some(client_family) = client_family
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return true;
|
||||
};
|
||||
admin_usage_client_family(item).is_some_and(|value| value.eq_ignore_ascii_case(client_family))
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn build_admin_usage_records_response_with_attempt_flags(
|
||||
items: &[StoredRequestUsageAudit],
|
||||
@@ -502,7 +516,9 @@ pub(super) async fn maybe_build_local_admin_usage_summary_response(
|
||||
.summarize_usage_audits(&UsageAuditSummaryQuery {
|
||||
created_from_unix_secs,
|
||||
created_until_unix_secs,
|
||||
..Default::default()
|
||||
user_id: query_param_value(query, "user_id"),
|
||||
provider_name: query_param_value(query, "provider"),
|
||||
model: query_param_value(query, "model"),
|
||||
})
|
||||
.await?;
|
||||
return Ok(Some(build_admin_usage_summary_stats_response_from_summary(
|
||||
@@ -598,6 +614,7 @@ pub(super) async fn maybe_build_local_admin_usage_summary_response(
|
||||
admin_usage_attempt_status_filter(query_param_value(query, "status").as_deref());
|
||||
let search = query_param_value(query, "search");
|
||||
let username_filter = query_param_value(query, "username");
|
||||
let client_family_filter = query_param_value(query, "client_family");
|
||||
let limit = match admin_usage_parse_limit(query) {
|
||||
Ok(value) => value,
|
||||
Err(detail) => return Ok(Some(admin_usage_bad_request_response(detail))),
|
||||
@@ -632,7 +649,12 @@ pub(super) async fn maybe_build_local_admin_usage_summary_response(
|
||||
let active_username_filter = username_filter
|
||||
.as_deref()
|
||||
.filter(|value| !value.trim().is_empty());
|
||||
let (usage, total) = if let Some(attempt_status) = attempt_status_filter {
|
||||
let active_client_family_filter = client_family_filter
|
||||
.as_deref()
|
||||
.filter(|value| !value.trim().is_empty());
|
||||
let (usage, total) = if attempt_status_filter.is_some()
|
||||
|| active_client_family_filter.is_some()
|
||||
{
|
||||
let mut usage = state.list_usage_audits(&base_query).await?;
|
||||
let user_ids: Vec<String> = usage
|
||||
.iter()
|
||||
@@ -662,12 +684,14 @@ pub(super) async fn maybe_build_local_admin_usage_summary_response(
|
||||
active_username_filter,
|
||||
&users_by_id,
|
||||
state.has_auth_user_data_reader(),
|
||||
) && admin_usage_matches_attempt_status(
|
||||
item,
|
||||
attempt_status,
|
||||
&attempt_flags_by_usage_id,
|
||||
request_candidate_reader_available,
|
||||
)
|
||||
) && attempt_status_filter.is_none_or(|attempt_status| {
|
||||
admin_usage_matches_attempt_status(
|
||||
item,
|
||||
attempt_status,
|
||||
&attempt_flags_by_usage_id,
|
||||
request_candidate_reader_available,
|
||||
)
|
||||
}) && admin_usage_matches_client_family(item, active_client_family_filter)
|
||||
});
|
||||
sort_usage_newest_first(&mut usage);
|
||||
let total = usage.len();
|
||||
|
||||
@@ -13,12 +13,18 @@ pub(super) fn key_api_formats_without_entry(
|
||||
}
|
||||
|
||||
pub(super) fn endpoint_key_counts_by_format(
|
||||
provider_type: &str,
|
||||
endpoints: &[StoredProviderCatalogEndpoint],
|
||||
keys: &[StoredProviderCatalogKey],
|
||||
) -> (
|
||||
std::collections::BTreeMap<String, usize>,
|
||||
std::collections::BTreeMap<String, usize>,
|
||||
) {
|
||||
admin_provider_endpoints_pure::endpoint_key_counts_by_format(keys)
|
||||
admin_provider_endpoints_pure::endpoint_key_counts_by_format(provider_type, endpoints, keys)
|
||||
}
|
||||
|
||||
pub(super) fn normalize_endpoint_api_format(api_format: &str) -> String {
|
||||
admin_provider_endpoints_pure::normalize_endpoint_api_format(api_format)
|
||||
}
|
||||
|
||||
pub(super) fn build_admin_provider_endpoint_response(
|
||||
|
||||
@@ -4,7 +4,10 @@ use aether_data_contracts::repository::provider_catalog::{
|
||||
};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use super::payloads::{build_admin_provider_endpoint_response, endpoint_key_counts_by_format};
|
||||
use super::payloads::{
|
||||
build_admin_provider_endpoint_response, endpoint_key_counts_by_format,
|
||||
normalize_endpoint_api_format,
|
||||
};
|
||||
|
||||
pub(crate) async fn build_admin_provider_endpoints_payload(
|
||||
state: &AdminAppState<'_>,
|
||||
@@ -38,7 +41,8 @@ pub(crate) async fn build_admin_provider_endpoints_payload(
|
||||
.await
|
||||
.ok()
|
||||
.unwrap_or_default();
|
||||
let (total_keys_by_format, active_keys_by_format) = endpoint_key_counts_by_format(&keys);
|
||||
let (total_keys_by_format, active_keys_by_format) =
|
||||
endpoint_key_counts_by_format(&provider.provider_type, &endpoints, &keys);
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
@@ -51,15 +55,16 @@ pub(crate) async fn build_admin_provider_endpoints_payload(
|
||||
.skip(skip)
|
||||
.take(limit)
|
||||
.map(|endpoint| {
|
||||
let endpoint_api_format = normalize_endpoint_api_format(&endpoint.api_format);
|
||||
build_admin_provider_endpoint_response(
|
||||
&endpoint,
|
||||
&provider.name,
|
||||
total_keys_by_format
|
||||
.get(endpoint.api_format.as_str())
|
||||
.get(endpoint_api_format.as_str())
|
||||
.copied()
|
||||
.unwrap_or(0),
|
||||
active_keys_by_format
|
||||
.get(endpoint.api_format.as_str())
|
||||
.get(endpoint_api_format.as_str())
|
||||
.copied()
|
||||
.unwrap_or(0),
|
||||
now_unix_secs,
|
||||
@@ -92,22 +97,27 @@ pub(crate) async fn build_admin_endpoint_payload(
|
||||
.await
|
||||
.ok()
|
||||
.unwrap_or_default();
|
||||
let (total_keys_by_format, active_keys_by_format) = endpoint_key_counts_by_format(&keys);
|
||||
let (total_keys_by_format, active_keys_by_format) = endpoint_key_counts_by_format(
|
||||
&provider.provider_type,
|
||||
std::slice::from_ref(&endpoint),
|
||||
&keys,
|
||||
);
|
||||
let now_unix_secs = SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
let endpoint_api_format = normalize_endpoint_api_format(&endpoint.api_format);
|
||||
|
||||
Some(build_admin_provider_endpoint_response(
|
||||
&endpoint,
|
||||
&provider.name,
|
||||
total_keys_by_format
|
||||
.get(endpoint.api_format.as_str())
|
||||
.get(endpoint_api_format.as_str())
|
||||
.copied()
|
||||
.unwrap_or(0),
|
||||
active_keys_by_format
|
||||
.get(endpoint.api_format.as_str())
|
||||
.get(endpoint_api_format.as_str())
|
||||
.copied()
|
||||
.unwrap_or(0),
|
||||
now_unix_secs,
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use super::extractors::admin_endpoint_id;
|
||||
use super::payloads::{
|
||||
build_admin_provider_endpoint_response, endpoint_key_counts_by_format,
|
||||
AdminProviderEndpointUpdatePatch,
|
||||
normalize_endpoint_api_format, AdminProviderEndpointUpdatePatch,
|
||||
};
|
||||
use super::support::build_admin_endpoints_data_unavailable_response;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
|
||||
@@ -147,18 +147,23 @@ pub(super) async fn maybe_handle(
|
||||
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
|
||||
.await
|
||||
.unwrap_or_default();
|
||||
let (total_keys_by_format, active_keys_by_format) = endpoint_key_counts_by_format(&keys);
|
||||
let (total_keys_by_format, active_keys_by_format) = endpoint_key_counts_by_format(
|
||||
&provider.provider_type,
|
||||
std::slice::from_ref(&updated),
|
||||
&keys,
|
||||
);
|
||||
let updated_api_format = normalize_endpoint_api_format(&updated.api_format);
|
||||
|
||||
Ok(Some(
|
||||
Json(build_admin_provider_endpoint_response(
|
||||
&updated,
|
||||
&provider.name,
|
||||
total_keys_by_format
|
||||
.get(updated.api_format.as_str())
|
||||
.get(updated_api_format.as_str())
|
||||
.copied()
|
||||
.unwrap_or(0),
|
||||
active_keys_by_format
|
||||
.get(updated.api_format.as_str())
|
||||
.get(updated_api_format.as_str())
|
||||
.copied()
|
||||
.unwrap_or(0),
|
||||
now_unix_secs,
|
||||
|
||||
@@ -32,7 +32,7 @@ pub(super) async fn maybe_handle(
|
||||
.into_response(),
|
||||
));
|
||||
};
|
||||
let Some(_provider) = state
|
||||
let Some(provider) = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
@@ -129,7 +129,9 @@ pub(super) async fn maybe_handle(
|
||||
Json(serde_json::Value::Array(
|
||||
created
|
||||
.iter()
|
||||
.map(|model| build_admin_provider_model_response(model, now_unix_secs))
|
||||
.map(|model| {
|
||||
build_admin_provider_model_response(&provider, model, now_unix_secs)
|
||||
})
|
||||
.collect(),
|
||||
))
|
||||
.into_response(),
|
||||
|
||||
@@ -31,7 +31,7 @@ pub(super) async fn maybe_handle(
|
||||
.into_response(),
|
||||
));
|
||||
};
|
||||
let Some(_provider) = state
|
||||
let Some(provider) = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
@@ -90,8 +90,12 @@ pub(super) async fn maybe_handle(
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
Json(build_admin_provider_model_response(&created, now_unix_secs))
|
||||
.into_response()
|
||||
Json(build_admin_provider_model_response(
|
||||
&provider,
|
||||
&created,
|
||||
now_unix_secs,
|
||||
))
|
||||
.into_response()
|
||||
}
|
||||
None => (
|
||||
http::StatusCode::INTERNAL_SERVER_ERROR,
|
||||
|
||||
@@ -1,9 +1,13 @@
|
||||
use crate::handlers::admin::provider::shared::model_test_capabilities::{
|
||||
admin_provider_model_supports_image_generation, admin_provider_model_test_capabilities_payload,
|
||||
};
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::models as admin_provider_models_pure;
|
||||
use aether_data_contracts::repository::global_models::{
|
||||
AdminProviderModelListQuery, StoredAdminProviderModel,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogProvider;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
pub(super) fn admin_provider_model_effective_input_price(
|
||||
@@ -26,10 +30,32 @@ pub(super) fn admin_provider_model_effective_capability(
|
||||
}
|
||||
|
||||
pub(super) fn build_admin_provider_model_response(
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
model: &StoredAdminProviderModel,
|
||||
now_unix_secs: u64,
|
||||
) -> serde_json::Value {
|
||||
admin_provider_models_pure::build_admin_provider_model_response(model, now_unix_secs)
|
||||
let mut payload =
|
||||
admin_provider_models_pure::build_admin_provider_model_response(model, now_unix_secs);
|
||||
let fallback_supports_image_generation = payload
|
||||
.get("effective_supports_image_generation")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
let supports_image_generation = admin_provider_model_supports_image_generation(
|
||||
&provider.provider_type,
|
||||
&model.provider_model_name,
|
||||
fallback_supports_image_generation,
|
||||
);
|
||||
if let Some(object) = payload.as_object_mut() {
|
||||
object.insert(
|
||||
"model_test_capabilities".to_string(),
|
||||
admin_provider_model_test_capabilities_payload(
|
||||
&provider.provider_type,
|
||||
&model.provider_model_name,
|
||||
supports_image_generation,
|
||||
),
|
||||
);
|
||||
}
|
||||
payload
|
||||
}
|
||||
|
||||
pub(super) async fn build_admin_provider_models_payload(
|
||||
@@ -48,9 +74,10 @@ pub(super) async fn build_admin_provider_models_payload(
|
||||
.ok()?
|
||||
.into_iter()
|
||||
.next()?;
|
||||
let provider_id = provider.id.clone();
|
||||
let mut models = state
|
||||
.list_admin_provider_models(&AdminProviderModelListQuery {
|
||||
provider_id: provider.id,
|
||||
provider_id,
|
||||
is_active,
|
||||
offset: skip,
|
||||
limit,
|
||||
@@ -70,7 +97,7 @@ pub(super) async fn build_admin_provider_models_payload(
|
||||
Some(serde_json::Value::Array(
|
||||
models
|
||||
.iter()
|
||||
.map(|model| build_admin_provider_model_response(model, now_unix_secs))
|
||||
.map(|model| build_admin_provider_model_response(&provider, model, now_unix_secs))
|
||||
.collect(),
|
||||
))
|
||||
}
|
||||
@@ -80,9 +107,15 @@ pub(super) async fn build_admin_provider_model_payload(
|
||||
provider_id: &str,
|
||||
model_id: &str,
|
||||
) -> Option<serde_json::Value> {
|
||||
if !state.has_global_model_data_reader() {
|
||||
if !state.has_provider_catalog_data_reader() || !state.has_global_model_data_reader() {
|
||||
return None;
|
||||
}
|
||||
let provider = state
|
||||
.read_provider_catalog_providers_by_ids(&[provider_id.to_string()])
|
||||
.await
|
||||
.ok()?
|
||||
.into_iter()
|
||||
.next()?;
|
||||
let model = state
|
||||
.get_admin_provider_model(provider_id, model_id)
|
||||
.await
|
||||
@@ -92,7 +125,11 @@ pub(super) async fn build_admin_provider_model_payload(
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
Some(build_admin_provider_model_response(&model, now_unix_secs))
|
||||
Some(build_admin_provider_model_response(
|
||||
&provider,
|
||||
&model,
|
||||
now_unix_secs,
|
||||
))
|
||||
}
|
||||
|
||||
pub(super) async fn admin_provider_model_name_exists(
|
||||
|
||||
@@ -33,6 +33,20 @@ pub(super) async fn maybe_handle(
|
||||
.into_response(),
|
||||
));
|
||||
};
|
||||
let Some(provider) = state
|
||||
.read_provider_catalog_providers_by_ids(std::slice::from_ref(&provider_id))
|
||||
.await?
|
||||
.into_iter()
|
||||
.next()
|
||||
else {
|
||||
return Ok(Some(
|
||||
(
|
||||
http::StatusCode::NOT_FOUND,
|
||||
Json(json!({ "detail": format!("Provider {provider_id} 不存在") })),
|
||||
)
|
||||
.into_response(),
|
||||
));
|
||||
};
|
||||
let Some(existing) = state
|
||||
.get_admin_provider_model(&provider_id, &model_id)
|
||||
.await?
|
||||
@@ -110,8 +124,12 @@ pub(super) async fn maybe_handle(
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0);
|
||||
Json(build_admin_provider_model_response(&updated, now_unix_secs))
|
||||
.into_response()
|
||||
Json(build_admin_provider_model_response(
|
||||
&provider,
|
||||
&updated,
|
||||
now_unix_secs,
|
||||
))
|
||||
.into_response()
|
||||
}
|
||||
None => (
|
||||
http::StatusCode::NOT_FOUND,
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
use super::super::helpers::admin_provider_oauth_key_name_from_auth_config;
|
||||
use super::super::token_import::{
|
||||
build_provider_access_token_import_auth_config, provider_type_supports_access_token_import,
|
||||
};
|
||||
@@ -24,13 +25,11 @@ use crate::handlers::admin::provider::oauth::runtime::{
|
||||
use crate::handlers::admin::provider::oauth::state::{
|
||||
admin_provider_oauth_template, exchange_admin_provider_oauth_refresh_token,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::support::ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL;
|
||||
use crate::handlers::admin::request::{AdminAppState, AdminProviderOAuthTemplate};
|
||||
use crate::GatewayError;
|
||||
use aether_admin::provider::oauth::parse_admin_provider_oauth_kiro_batch_import_entries;
|
||||
use aether_contracts::ProxySnapshot;
|
||||
use serde_json::{json, Map, Value};
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
struct AdminProviderOAuthResolvedBatchImport {
|
||||
access_token: String,
|
||||
@@ -45,7 +44,7 @@ pub(super) fn estimate_admin_provider_oauth_batch_import_total(
|
||||
if provider_type.eq_ignore_ascii_case("kiro") {
|
||||
parse_admin_provider_oauth_kiro_batch_import_entries(raw_credentials).len()
|
||||
} else {
|
||||
parse_admin_provider_oauth_batch_import_entries(raw_credentials).len()
|
||||
parse_admin_provider_oauth_batch_import_entries(provider_type, raw_credentials).len()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -67,7 +66,8 @@ pub(super) async fn execute_admin_provider_oauth_batch_import_for_provider_type(
|
||||
)
|
||||
.await
|
||||
} else {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(raw_credentials);
|
||||
let entries =
|
||||
parse_admin_provider_oauth_batch_import_entries(provider_type, raw_credentials);
|
||||
execute_admin_provider_oauth_batch_import(
|
||||
state,
|
||||
provider_id,
|
||||
@@ -82,7 +82,7 @@ pub(super) async fn execute_admin_provider_oauth_batch_import_for_provider_type(
|
||||
|
||||
async fn resolve_admin_provider_oauth_batch_import_tokens(
|
||||
state: &AdminAppState<'_>,
|
||||
template: AdminProviderOAuthTemplate,
|
||||
template: Option<AdminProviderOAuthTemplate>,
|
||||
provider_type: &str,
|
||||
entry: &AdminProviderOAuthBatchImportEntry,
|
||||
request_proxy: Option<ProxySnapshot>,
|
||||
@@ -99,6 +99,29 @@ async fn resolve_admin_provider_oauth_batch_import_tokens(
|
||||
.filter(|value| !value.is_empty());
|
||||
|
||||
if let Some(refresh_token) = refresh_token {
|
||||
let Some(template) = template else {
|
||||
if provider_type_supports_access_token_import(provider_type) {
|
||||
if let Some(access_token) = access_token {
|
||||
let (auth_config, expires_at) = build_provider_access_token_import_auth_config(
|
||||
provider_type,
|
||||
access_token,
|
||||
Some(refresh_token),
|
||||
entry.expires_at,
|
||||
Some("Provider 不支持 Refresh Token 交换,已回退为 Session Token 导入"),
|
||||
);
|
||||
return Ok(AdminProviderOAuthResolvedBatchImport {
|
||||
access_token: access_token.to_string(),
|
||||
auth_config,
|
||||
expires_at,
|
||||
});
|
||||
}
|
||||
}
|
||||
return Err(
|
||||
"该 Provider 不支持 Refresh Token 导入,请提供 sso_token 或 access_token"
|
||||
.to_string(),
|
||||
);
|
||||
};
|
||||
|
||||
let token_payload = match exchange_admin_provider_oauth_refresh_token(
|
||||
state,
|
||||
template,
|
||||
@@ -152,7 +175,7 @@ async fn resolve_admin_provider_oauth_batch_import_tokens(
|
||||
|
||||
if let Some(access_token) = access_token {
|
||||
if !provider_type_supports_access_token_import(provider_type) {
|
||||
return Err("Access Token 导入仅支持 Codex / ChatGPT Web Provider".to_string());
|
||||
return Err("Access Token 导入仅支持 Codex / ChatGPT Web / Grok Provider".to_string());
|
||||
}
|
||||
let (auth_config, expires_at) = build_provider_access_token_import_auth_config(
|
||||
provider_type,
|
||||
@@ -204,25 +227,7 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
});
|
||||
};
|
||||
|
||||
let Some(template) = admin_provider_oauth_template(provider_type) else {
|
||||
return Ok(AdminProviderOAuthBatchImportOutcome {
|
||||
total: entries.len(),
|
||||
success: 0,
|
||||
failed: entries.len(),
|
||||
results: entries
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(index, _)| {
|
||||
json!({
|
||||
"index": index,
|
||||
"status": "error",
|
||||
"error": ADMIN_PROVIDER_OAUTH_DATA_UNAVAILABLE_DETAIL,
|
||||
"replaced": false,
|
||||
})
|
||||
})
|
||||
.collect(),
|
||||
});
|
||||
};
|
||||
let template = admin_provider_oauth_template(provider_type);
|
||||
|
||||
let endpoint_resolution =
|
||||
resolve_provider_oauth_runtime_endpoints(state, &provider, provider_type).await?;
|
||||
@@ -340,24 +345,11 @@ pub(super) async fn execute_admin_provider_oauth_batch_import(
|
||||
}
|
||||
}
|
||||
} else {
|
||||
let key_name = auth_config
|
||||
.get("email")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|email| format!("{provider_type}_{email}"))
|
||||
.unwrap_or_else(|| {
|
||||
format!(
|
||||
"{}_{}_{}",
|
||||
provider_type,
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
.unwrap_or(0),
|
||||
index
|
||||
)
|
||||
});
|
||||
let key_name = admin_provider_oauth_key_name_from_auth_config(
|
||||
provider_type,
|
||||
&auth_config,
|
||||
Some(index),
|
||||
);
|
||||
match create_provider_oauth_catalog_key(
|
||||
state,
|
||||
provider_id,
|
||||
|
||||
@@ -8,7 +8,7 @@ use super::parse::{
|
||||
};
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::state::{
|
||||
admin_provider_oauth_template, build_admin_provider_oauth_backend_unavailable_response,
|
||||
build_admin_provider_oauth_backend_unavailable_response,
|
||||
is_fixed_provider_type_for_provider_oauth,
|
||||
};
|
||||
use crate::handlers::admin::provider::shared::paths::admin_provider_oauth_batch_import_provider_id;
|
||||
@@ -60,10 +60,6 @@ pub(in super::super) async fn handle_admin_provider_oauth_batch_import(
|
||||
"该 Provider 不是固定类型,无法使用 provider-oauth",
|
||||
));
|
||||
}
|
||||
if provider_type != "kiro" && admin_provider_oauth_template(&provider_type).is_none() {
|
||||
return Ok(build_admin_provider_oauth_backend_unavailable_response());
|
||||
}
|
||||
|
||||
let total = estimate_admin_provider_oauth_batch_import_total(
|
||||
&provider_type,
|
||||
payload.credentials.as_str(),
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
use super::super::token_import::{import_tokens_from_raw_token, normalize_single_import_tokens};
|
||||
use super::super::token_import::{import_tokens_from_raw_token, normalize_provider_import_tokens};
|
||||
use crate::handlers::admin::provider::oauth::errors::build_internal_control_error_response;
|
||||
use crate::handlers::admin::provider::oauth::state::{current_unix_secs, json_u64_value};
|
||||
use axum::{
|
||||
@@ -25,8 +25,15 @@ pub(super) struct AdminProviderOAuthBatchImportEntry {
|
||||
pub account_id: Option<String>,
|
||||
pub account_user_id: Option<String>,
|
||||
pub plan_type: Option<String>,
|
||||
pub pool_tier: Option<String>,
|
||||
pub user_id: Option<String>,
|
||||
pub email: Option<String>,
|
||||
pub account_name: Option<String>,
|
||||
pub sso_rw_token: Option<String>,
|
||||
pub cf_cookies: Option<String>,
|
||||
pub cf_clearance: Option<String>,
|
||||
pub user_agent: Option<String>,
|
||||
pub browser_profile: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
@@ -67,16 +74,72 @@ fn coerce_admin_provider_oauth_import_str(value: Option<&serde_json::Value>) ->
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn grok_cookie_value(raw: &str, name: &str) -> Option<String> {
|
||||
raw.trim()
|
||||
.strip_prefix("Cookie:")
|
||||
.unwrap_or_else(|| raw.trim())
|
||||
.split(';')
|
||||
.filter_map(|segment| segment.trim().split_once('='))
|
||||
.find_map(|(cookie_name, cookie_value)| {
|
||||
cookie_name
|
||||
.trim()
|
||||
.eq_ignore_ascii_case(name)
|
||||
.then(|| cookie_value.trim())
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
})
|
||||
}
|
||||
|
||||
fn grok_cookie_profile(raw: &str) -> Option<String> {
|
||||
let raw = raw
|
||||
.trim()
|
||||
.strip_prefix("Cookie:")
|
||||
.unwrap_or_else(|| raw.trim());
|
||||
let parts = raw
|
||||
.split(';')
|
||||
.filter_map(|segment| {
|
||||
let (cookie_name, cookie_value) = segment.trim().split_once('=')?;
|
||||
let cookie_name = cookie_name.trim();
|
||||
let cookie_value = cookie_value.trim();
|
||||
if cookie_name.is_empty()
|
||||
|| cookie_value.is_empty()
|
||||
|| cookie_name.eq_ignore_ascii_case("sso")
|
||||
|| cookie_name.eq_ignore_ascii_case("sso-rw")
|
||||
{
|
||||
return None;
|
||||
}
|
||||
Some(format!("{cookie_name}={cookie_value}"))
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
(!parts.is_empty()).then(|| parts.join("; "))
|
||||
}
|
||||
|
||||
fn grok_cookie_session_token(provider_type: &str, raw: &str) -> Option<String> {
|
||||
provider_type
|
||||
.trim()
|
||||
.eq_ignore_ascii_case("grok")
|
||||
.then(|| grok_cookie_value(raw, "sso"))
|
||||
.flatten()
|
||||
}
|
||||
|
||||
fn extract_admin_provider_oauth_batch_import_entry(
|
||||
provider_type: &str,
|
||||
item: &serde_json::Value,
|
||||
) -> Option<AdminProviderOAuthBatchImportEntry> {
|
||||
match item {
|
||||
serde_json::Value::String(value) => {
|
||||
let refresh_token = value.trim();
|
||||
if refresh_token.is_empty() {
|
||||
let raw_token = value.trim();
|
||||
if raw_token.is_empty() {
|
||||
None
|
||||
} else {
|
||||
let (refresh_token, access_token) = import_tokens_from_raw_token(refresh_token);
|
||||
let sso_from_cookie = grok_cookie_session_token(provider_type, raw_token);
|
||||
let token_input = sso_from_cookie.as_deref().unwrap_or(raw_token);
|
||||
let (refresh_token, access_token) = import_tokens_from_raw_token(token_input);
|
||||
let (refresh_token, access_token) = normalize_provider_import_tokens(
|
||||
provider_type,
|
||||
refresh_token.as_deref(),
|
||||
access_token.as_deref(),
|
||||
);
|
||||
Some(AdminProviderOAuthBatchImportEntry {
|
||||
refresh_token,
|
||||
access_token,
|
||||
@@ -84,8 +147,15 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
account_id: None,
|
||||
account_user_id: None,
|
||||
plan_type: None,
|
||||
user_id: None,
|
||||
pool_tier: None,
|
||||
user_id: grok_cookie_value(raw_token, "x-userid"),
|
||||
email: None,
|
||||
account_name: None,
|
||||
sso_rw_token: grok_cookie_value(raw_token, "sso-rw"),
|
||||
cf_cookies: grok_cookie_profile(raw_token),
|
||||
cf_clearance: grok_cookie_value(raw_token, "cf_clearance"),
|
||||
user_agent: None,
|
||||
browser_profile: None,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -100,8 +170,34 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
.get("access_token")
|
||||
.or_else(|| object.get("accessToken")),
|
||||
);
|
||||
let (refresh_token, access_token) =
|
||||
normalize_single_import_tokens(refresh_token.as_deref(), access_token.as_deref());
|
||||
let grok_token_alias = if provider_type.trim().eq_ignore_ascii_case("grok") {
|
||||
object.get("token")
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let grok_cookie = if provider_type.trim().eq_ignore_ascii_case("grok") {
|
||||
coerce_admin_provider_oauth_import_str(
|
||||
object.get("cookie").or_else(|| object.get("cookieHeader")),
|
||||
)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let session_token = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("sso_token")
|
||||
.or_else(|| object.get("ssoToken"))
|
||||
.or(grok_token_alias),
|
||||
)
|
||||
.or_else(|| {
|
||||
grok_cookie
|
||||
.as_deref()
|
||||
.and_then(|cookie| grok_cookie_value(cookie, "sso"))
|
||||
});
|
||||
let (refresh_token, access_token) = normalize_provider_import_tokens(
|
||||
provider_type,
|
||||
refresh_token.as_deref(),
|
||||
access_token.as_deref().or(session_token.as_deref()),
|
||||
);
|
||||
if refresh_token.is_none() && access_token.is_none() {
|
||||
return None;
|
||||
}
|
||||
@@ -129,14 +225,65 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
.or_else(|| object.get("chatgptPlanType")),
|
||||
)
|
||||
.map(|value| value.to_ascii_lowercase());
|
||||
let pool_tier = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("pool_tier")
|
||||
.or_else(|| object.get("poolTier"))
|
||||
.or_else(|| object.get("tier")),
|
||||
)
|
||||
.map(|value| value.to_ascii_lowercase());
|
||||
let user_id = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("user_id")
|
||||
.or_else(|| object.get("userId"))
|
||||
.or_else(|| object.get("chatgpt_user_id"))
|
||||
.or_else(|| object.get("chatgptUserId")),
|
||||
);
|
||||
)
|
||||
.or_else(|| {
|
||||
grok_cookie
|
||||
.as_deref()
|
||||
.and_then(|cookie| grok_cookie_value(cookie, "x-userid"))
|
||||
});
|
||||
let email = coerce_admin_provider_oauth_import_str(object.get("email"));
|
||||
let account_name = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("account_name")
|
||||
.or_else(|| object.get("accountName")),
|
||||
);
|
||||
let sso_rw_token = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("sso_rw_token")
|
||||
.or_else(|| object.get("ssoRwToken")),
|
||||
)
|
||||
.or_else(|| {
|
||||
grok_cookie
|
||||
.as_deref()
|
||||
.and_then(|cookie| grok_cookie_value(cookie, "sso-rw"))
|
||||
});
|
||||
let cf_clearance = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("cf_clearance")
|
||||
.or_else(|| object.get("cfClearance")),
|
||||
)
|
||||
.or_else(|| {
|
||||
grok_cookie
|
||||
.as_deref()
|
||||
.and_then(|cookie| grok_cookie_value(cookie, "cf_clearance"))
|
||||
});
|
||||
let cf_cookies = coerce_admin_provider_oauth_import_str(
|
||||
object.get("cf_cookies").or_else(|| object.get("cfCookies")),
|
||||
)
|
||||
.or_else(|| grok_cookie.as_deref().and_then(grok_cookie_profile));
|
||||
let user_agent = coerce_admin_provider_oauth_import_str(
|
||||
object.get("user_agent").or_else(|| object.get("userAgent")),
|
||||
);
|
||||
let browser_profile = coerce_admin_provider_oauth_import_str(
|
||||
object
|
||||
.get("browser_profile")
|
||||
.or_else(|| object.get("browserProfile"))
|
||||
.or_else(|| object.get("browser"))
|
||||
.or_else(|| object.get("impersonate")),
|
||||
);
|
||||
Some(AdminProviderOAuthBatchImportEntry {
|
||||
refresh_token,
|
||||
access_token,
|
||||
@@ -144,8 +291,15 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
account_id,
|
||||
account_user_id,
|
||||
plan_type,
|
||||
pool_tier,
|
||||
user_id,
|
||||
email,
|
||||
account_name,
|
||||
sso_rw_token,
|
||||
cf_cookies,
|
||||
cf_clearance,
|
||||
user_agent,
|
||||
browser_profile,
|
||||
})
|
||||
}
|
||||
_ => None,
|
||||
@@ -153,6 +307,7 @@ fn extract_admin_provider_oauth_batch_import_entry(
|
||||
}
|
||||
|
||||
pub(super) fn parse_admin_provider_oauth_batch_import_entries(
|
||||
provider_type: &str,
|
||||
raw_credentials: &str,
|
||||
) -> Vec<AdminProviderOAuthBatchImportEntry> {
|
||||
let raw = raw_credentials.trim();
|
||||
@@ -165,7 +320,9 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries(
|
||||
{
|
||||
return items
|
||||
.iter()
|
||||
.filter_map(extract_admin_provider_oauth_batch_import_entry)
|
||||
.filter_map(|item| {
|
||||
extract_admin_provider_oauth_batch_import_entry(provider_type, item)
|
||||
})
|
||||
.collect();
|
||||
}
|
||||
}
|
||||
@@ -174,7 +331,7 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries(
|
||||
if let Ok(value @ serde_json::Value::Object(_)) =
|
||||
serde_json::from_str::<serde_json::Value>(raw)
|
||||
{
|
||||
return extract_admin_provider_oauth_batch_import_entry(&value)
|
||||
return extract_admin_provider_oauth_batch_import_entry(provider_type, &value)
|
||||
.into_iter()
|
||||
.collect();
|
||||
}
|
||||
@@ -183,18 +340,19 @@ pub(super) fn parse_admin_provider_oauth_batch_import_entries(
|
||||
raw.lines()
|
||||
.map(str::trim)
|
||||
.filter(|line| !line.is_empty() && !line.starts_with('#'))
|
||||
.map(|token| {
|
||||
let (refresh_token, access_token) = import_tokens_from_raw_token(token);
|
||||
AdminProviderOAuthBatchImportEntry {
|
||||
refresh_token,
|
||||
access_token,
|
||||
expires_at: None,
|
||||
account_id: None,
|
||||
account_user_id: None,
|
||||
plan_type: None,
|
||||
user_id: None,
|
||||
email: None,
|
||||
.filter_map(|line| {
|
||||
if line.starts_with('{') {
|
||||
return serde_json::from_str::<serde_json::Value>(line)
|
||||
.ok()
|
||||
.and_then(|value| {
|
||||
extract_admin_provider_oauth_batch_import_entry(provider_type, &value)
|
||||
});
|
||||
}
|
||||
|
||||
extract_admin_provider_oauth_batch_import_entry(
|
||||
provider_type,
|
||||
&serde_json::Value::String(line.to_string()),
|
||||
)
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
@@ -204,10 +362,8 @@ pub(super) fn apply_admin_provider_oauth_batch_import_hints(
|
||||
entry: &AdminProviderOAuthBatchImportEntry,
|
||||
auth_config: &mut serde_json::Map<String, serde_json::Value>,
|
||||
) {
|
||||
if !matches!(
|
||||
provider_type.trim().to_ascii_lowercase().as_str(),
|
||||
"codex" | "chatgpt_web"
|
||||
) {
|
||||
let provider_type = provider_type.trim().to_ascii_lowercase();
|
||||
if !matches!(provider_type.as_str(), "codex" | "chatgpt_web" | "grok") {
|
||||
return;
|
||||
}
|
||||
if let Some(account_id) = entry.account_id.as_ref() {
|
||||
@@ -225,6 +381,11 @@ pub(super) fn apply_admin_provider_oauth_batch_import_hints(
|
||||
.entry("plan_type".to_string())
|
||||
.or_insert_with(|| json!(plan_type));
|
||||
}
|
||||
if let Some(pool_tier) = entry.pool_tier.as_ref() {
|
||||
auth_config
|
||||
.entry("pool_tier".to_string())
|
||||
.or_insert_with(|| json!(pool_tier));
|
||||
}
|
||||
if let Some(user_id) = entry.user_id.as_ref() {
|
||||
auth_config
|
||||
.entry("user_id".to_string())
|
||||
@@ -235,6 +396,36 @@ pub(super) fn apply_admin_provider_oauth_batch_import_hints(
|
||||
.entry("email".to_string())
|
||||
.or_insert_with(|| json!(email));
|
||||
}
|
||||
if let Some(account_name) = entry.account_name.as_ref() {
|
||||
auth_config
|
||||
.entry("account_name".to_string())
|
||||
.or_insert_with(|| json!(account_name));
|
||||
}
|
||||
if let Some(sso_rw_token) = entry.sso_rw_token.as_ref() {
|
||||
auth_config
|
||||
.entry("sso_rw_token".to_string())
|
||||
.or_insert_with(|| json!(sso_rw_token));
|
||||
}
|
||||
if let Some(cf_cookies) = entry.cf_cookies.as_ref() {
|
||||
auth_config
|
||||
.entry("cf_cookies".to_string())
|
||||
.or_insert_with(|| json!(cf_cookies));
|
||||
}
|
||||
if let Some(cf_clearance) = entry.cf_clearance.as_ref() {
|
||||
auth_config
|
||||
.entry("cf_clearance".to_string())
|
||||
.or_insert_with(|| json!(cf_clearance));
|
||||
}
|
||||
if let Some(user_agent) = entry.user_agent.as_ref() {
|
||||
auth_config
|
||||
.entry("user_agent".to_string())
|
||||
.or_insert_with(|| json!(user_agent));
|
||||
}
|
||||
if let Some(browser_profile) = entry.browser_profile.as_ref() {
|
||||
auth_config
|
||||
.entry("browser_profile".to_string())
|
||||
.or_insert_with(|| json!(browser_profile));
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn extract_admin_provider_oauth_batch_error_detail(
|
||||
@@ -337,6 +528,7 @@ mod tests {
|
||||
#[test]
|
||||
fn parses_access_token_only_entry() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"codex",
|
||||
r#"[{"accessToken":"at_1","expiresAt":2100000000,"accountId":"acc-1","email":"u@example.com"}]"#,
|
||||
);
|
||||
|
||||
@@ -356,10 +548,89 @@ mod tests {
|
||||
"exp": 2_000_000_000u64,
|
||||
}));
|
||||
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(&token);
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries("codex", &token);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some(token.as_str()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_grok_jsonl_session_entries() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"grok",
|
||||
r#"{"sso_token":"sso-1","cf_clearance":"cf-1","pool_tier":"heavy","email":"grok@example.com","browser_profile":"chrome136"}"#,
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some("sso-1"));
|
||||
assert_eq!(entries[0].cf_clearance.as_deref(), Some("cf-1"));
|
||||
assert_eq!(entries[0].pool_tier.as_deref(), Some("heavy"));
|
||||
assert_eq!(entries[0].email.as_deref(), Some("grok@example.com"));
|
||||
assert_eq!(entries[0].browser_profile.as_deref(), Some("chrome136"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_grok_token_alias_with_account_traits() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"grok",
|
||||
r#"[{"token":"sso-1","planType":"super","tier":"heavy","accountName":"Grok Heavy"}]"#,
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some("sso-1"));
|
||||
assert_eq!(entries[0].plan_type.as_deref(), Some("super"));
|
||||
assert_eq!(entries[0].pool_tier.as_deref(), Some("heavy"));
|
||||
assert_eq!(entries[0].account_name.as_deref(), Some("Grok Heavy"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_grok_plain_line_as_session_token() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries("grok", "opaque-sso-token");
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some("opaque-sso-token"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_grok_cookie_line_as_session_metadata() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"grok",
|
||||
"i18nextLng=zh; cf_clearance=cf-1; sso-rw=rw-1; sso=sso-1; x-userid=user-1",
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some("sso-1"));
|
||||
assert_eq!(entries[0].sso_rw_token.as_deref(), Some("rw-1"));
|
||||
assert_eq!(
|
||||
entries[0].cf_cookies.as_deref(),
|
||||
Some("i18nextLng=zh; cf_clearance=cf-1; x-userid=user-1")
|
||||
);
|
||||
assert_eq!(entries[0].cf_clearance.as_deref(), Some("cf-1"));
|
||||
assert_eq!(entries[0].user_id.as_deref(), Some("user-1"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_grok_cookie_object_as_session_metadata() {
|
||||
let entries = parse_admin_provider_oauth_batch_import_entries(
|
||||
"grok",
|
||||
r#"[{"cookie":"cf_clearance=cf-1; sso-rw=rw-1; sso=sso-1; x-userid=user-1","tier":"heavy"}]"#,
|
||||
);
|
||||
|
||||
assert_eq!(entries.len(), 1);
|
||||
assert_eq!(entries[0].refresh_token, None);
|
||||
assert_eq!(entries[0].access_token.as_deref(), Some("sso-1"));
|
||||
assert_eq!(entries[0].sso_rw_token.as_deref(), Some("rw-1"));
|
||||
assert_eq!(
|
||||
entries[0].cf_cookies.as_deref(),
|
||||
Some("cf_clearance=cf-1; x-userid=user-1")
|
||||
);
|
||||
assert_eq!(entries[0].cf_clearance.as_deref(), Some("cf-1"));
|
||||
assert_eq!(entries[0].user_id.as_deref(), Some("user-1"));
|
||||
assert_eq!(entries[0].pool_tier.as_deref(), Some("heavy"));
|
||||
}
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user