mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
Refactor pool candidate scheduling
This commit is contained in:
@@ -2,19 +2,30 @@ use crate::ai_serving::{is_json_request, GatewayControlDecision};
|
||||
|
||||
pub(crate) use crate::ai_serving::{
|
||||
build_gemini_stream_plan_from_decision, build_gemini_sync_plan_from_decision,
|
||||
build_local_gemini_files_stream_attempt_source_for_kind,
|
||||
build_local_gemini_files_stream_plan_and_reports_for_kind,
|
||||
build_local_gemini_files_sync_attempt_source_for_kind,
|
||||
build_local_gemini_files_sync_plan_and_reports_for_kind,
|
||||
build_local_image_stream_attempt_source_for_kind,
|
||||
build_local_image_stream_plan_and_reports_for_kind,
|
||||
build_local_image_sync_attempt_source_for_kind,
|
||||
build_local_image_sync_plan_and_reports_for_kind,
|
||||
build_local_openai_chat_stream_attempt_source_for_kind,
|
||||
build_local_openai_chat_stream_plan_and_reports_for_kind,
|
||||
build_local_openai_chat_sync_attempt_source_for_kind,
|
||||
build_local_openai_chat_sync_plan_and_reports_for_kind,
|
||||
build_local_openai_responses_stream_attempt_source_for_kind,
|
||||
build_local_openai_responses_stream_plan_and_reports_for_kind,
|
||||
build_local_openai_responses_sync_attempt_source_for_kind,
|
||||
build_local_openai_responses_sync_plan_and_reports_for_kind,
|
||||
build_local_same_format_stream_plan_and_reports, build_local_same_format_sync_plan_and_reports,
|
||||
build_local_same_format_stream_attempt_source, build_local_same_format_stream_plan_and_reports,
|
||||
build_local_same_format_sync_attempt_source, build_local_same_format_sync_plan_and_reports,
|
||||
build_local_video_sync_attempt_source_for_kind,
|
||||
build_local_video_sync_plan_and_reports_for_kind,
|
||||
build_openai_responses_stream_plan_from_decision,
|
||||
build_openai_responses_sync_plan_from_decision, build_passthrough_sync_plan_from_decision,
|
||||
build_standard_family_stream_plan_and_reports, build_standard_family_sync_plan_and_reports,
|
||||
build_standard_family_stream_attempt_source, build_standard_family_stream_plan_and_reports,
|
||||
build_standard_family_sync_attempt_source, build_standard_family_sync_plan_and_reports,
|
||||
build_standard_stream_plan_from_decision, build_standard_sync_plan_from_decision,
|
||||
maybe_build_stream_decision_payload, maybe_build_stream_plan_payload,
|
||||
maybe_build_sync_decision_payload, maybe_build_sync_plan_payload,
|
||||
|
||||
@@ -21,26 +21,37 @@ pub(crate) use self::finalize::internal::{
|
||||
};
|
||||
pub(crate) use self::planner::{
|
||||
build_gemini_stream_plan_from_decision, build_gemini_sync_plan_from_decision,
|
||||
build_local_gemini_files_stream_attempt_source_for_kind,
|
||||
build_local_gemini_files_stream_plan_and_reports_for_kind,
|
||||
build_local_gemini_files_sync_attempt_source_for_kind,
|
||||
build_local_gemini_files_sync_plan_and_reports_for_kind,
|
||||
build_local_image_stream_attempt_source_for_kind,
|
||||
build_local_image_stream_plan_and_reports_for_kind,
|
||||
build_local_image_sync_attempt_source_for_kind,
|
||||
build_local_image_sync_plan_and_reports_for_kind,
|
||||
build_local_openai_chat_stream_attempt_source_for_kind,
|
||||
build_local_openai_chat_stream_plan_and_reports_for_kind,
|
||||
build_local_openai_chat_sync_attempt_source_for_kind,
|
||||
build_local_openai_chat_sync_plan_and_reports_for_kind,
|
||||
build_local_openai_responses_stream_attempt_source_for_kind,
|
||||
build_local_openai_responses_stream_plan_and_reports_for_kind,
|
||||
build_local_openai_responses_sync_attempt_source_for_kind,
|
||||
build_local_openai_responses_sync_plan_and_reports_for_kind,
|
||||
build_local_same_format_stream_plan_and_reports, build_local_same_format_sync_plan_and_reports,
|
||||
build_local_same_format_stream_attempt_source, build_local_same_format_stream_plan_and_reports,
|
||||
build_local_same_format_sync_attempt_source, build_local_same_format_sync_plan_and_reports,
|
||||
build_local_video_sync_attempt_source_for_kind,
|
||||
build_local_video_sync_plan_and_reports_for_kind,
|
||||
build_openai_responses_stream_plan_from_decision,
|
||||
build_openai_responses_sync_plan_from_decision, build_passthrough_sync_plan_from_decision,
|
||||
build_standard_family_stream_plan_and_reports, build_standard_family_sync_plan_and_reports,
|
||||
build_standard_family_stream_attempt_source, build_standard_family_stream_plan_and_reports,
|
||||
build_standard_family_sync_attempt_source, build_standard_family_sync_plan_and_reports,
|
||||
build_standard_stream_plan_from_decision, build_standard_sync_plan_from_decision,
|
||||
extract_pool_sticky_session_token, maybe_build_stream_decision_payload,
|
||||
maybe_build_stream_plan_payload, maybe_build_sync_decision_payload,
|
||||
maybe_build_sync_plan_payload, planner_is_matching_stream_request,
|
||||
set_local_openai_chat_execution_exhausted_diagnostic, CandidateFailureDiagnostic,
|
||||
CandidateFailureDiagnosticKind, GatewayAuthApiKeySnapshot, GatewayProviderTransportSnapshot,
|
||||
LocalResolvedOAuthRequestAuth, PlannerAppState,
|
||||
LocalExecutionAttemptSource, LocalResolvedOAuthRequestAuth, PlannerAppState,
|
||||
};
|
||||
pub(crate) use self::pure::*;
|
||||
pub(crate) use self::transport::{
|
||||
|
||||
@@ -9,6 +9,7 @@ use aether_ai_serving::{
|
||||
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
use std::collections::VecDeque;
|
||||
use std::convert::Infallible;
|
||||
use uuid::Uuid;
|
||||
|
||||
@@ -16,16 +17,18 @@ use crate::ai_serving::planner::candidate_affinity_cache::remember_scheduler_aff
|
||||
use crate::ai_serving::planner::candidate_resolution::{
|
||||
resolve_and_rank_local_execution_candidates,
|
||||
resolve_and_rank_local_execution_candidates_without_transport_pair_gate,
|
||||
EligibleLocalExecutionCandidate, SkippedLocalExecutionCandidate,
|
||||
resolve_and_rank_logical_local_execution_candidates, EligibleLocalExecutionCandidate,
|
||||
LocalExecutionCandidateKind, SkippedLocalExecutionCandidate,
|
||||
};
|
||||
use crate::ai_serving::planner::materialization_policy::LocalCandidatePersistencePolicy;
|
||||
use crate::ai_serving::planner::pool_scheduler::PoolKeyCursor;
|
||||
use crate::ai_serving::planner::runtime_miss::record_local_runtime_candidate_skip_reason;
|
||||
use crate::ai_serving::planner::CandidateFailureDiagnostic;
|
||||
use crate::ai_serving::{GatewayAuthApiKeySnapshot, PlannerAppState};
|
||||
use crate::clock::current_unix_ms;
|
||||
use crate::handlers::shared::provider_pool::admin_provider_pool_config_from_config_value;
|
||||
use crate::orchestration::{local_attempt_slot_count, ExecutionAttemptIdentity};
|
||||
use crate::AppState;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub(crate) struct LocalExecutionCandidateAttempt {
|
||||
@@ -35,6 +38,71 @@ pub(crate) struct LocalExecutionCandidateAttempt {
|
||||
pub(crate) candidate_id: String,
|
||||
}
|
||||
|
||||
pub(crate) struct LocalExecutionCandidateAttemptSource<'a> {
|
||||
items: VecDeque<LocalExecutionCandidateAttemptSourceItem<'a>>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub(crate) trait LocalExecutionAttemptSource<T>: Send {
|
||||
async fn next_execution_attempt(&mut self) -> Result<Option<T>, GatewayError>;
|
||||
|
||||
async fn drain_execution_attempts(&mut self) -> Result<Vec<T>, GatewayError>;
|
||||
}
|
||||
|
||||
enum LocalExecutionCandidateAttemptSourceItem<'a> {
|
||||
Static {
|
||||
attempts: VecDeque<LocalExecutionCandidateAttempt>,
|
||||
},
|
||||
Pool {
|
||||
cursor: PoolKeyCursor<'a>,
|
||||
candidate_index: u32,
|
||||
pending_attempts: VecDeque<LocalExecutionCandidateAttempt>,
|
||||
},
|
||||
}
|
||||
|
||||
impl<'a> LocalExecutionCandidateAttemptSource<'a> {
|
||||
pub(crate) async fn next_attempt(&mut self) -> Option<LocalExecutionCandidateAttempt> {
|
||||
loop {
|
||||
let front = self.items.front_mut()?;
|
||||
match front {
|
||||
LocalExecutionCandidateAttemptSourceItem::Static { attempts } => {
|
||||
if let Some(attempt) = attempts.pop_front() {
|
||||
if attempts.is_empty() {
|
||||
self.items.pop_front();
|
||||
}
|
||||
return Some(attempt);
|
||||
}
|
||||
self.items.pop_front();
|
||||
}
|
||||
LocalExecutionCandidateAttemptSourceItem::Pool {
|
||||
cursor,
|
||||
candidate_index,
|
||||
pending_attempts,
|
||||
} => {
|
||||
if let Some(attempt) = pending_attempts.pop_front() {
|
||||
return Some(attempt);
|
||||
}
|
||||
let Some(candidate) = cursor.next_key().await else {
|
||||
cursor.log_exhausted();
|
||||
let _ = cursor.take_skipped_candidates();
|
||||
self.items.pop_front();
|
||||
continue;
|
||||
};
|
||||
*pending_attempts = build_unpersisted_local_execution_candidate_attempts(
|
||||
candidate,
|
||||
*candidate_index,
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn drain_static_attempts(&mut self) -> Vec<LocalExecutionCandidateAttempt> {
|
||||
self.items.clear();
|
||||
Vec::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalExecutionCandidateAttempt {
|
||||
pub(crate) fn attempt_identity(&self) -> ExecutionAttemptIdentity {
|
||||
ExecutionAttemptIdentity::new(self.candidate_index, self.retry_index)
|
||||
@@ -350,6 +418,99 @@ where
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub(crate) async fn build_local_execution_candidate_attempt_source_with_serving<'a, F, G>(
|
||||
state: PlannerAppState<'a>,
|
||||
trace_id: &str,
|
||||
client_api_format: &str,
|
||||
requested_model: Option<&str>,
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
required_capabilities: Option<&Value>,
|
||||
sticky_session_token: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
persistence_policy: LocalCandidatePersistencePolicy<'_>,
|
||||
candidates: Vec<SchedulerMinimalCandidateSelectionCandidate>,
|
||||
preselection_skipped: Vec<SkippedLocalExecutionCandidate>,
|
||||
resolution_mode: LocalCandidateResolutionMode,
|
||||
build_available_extra_data: F,
|
||||
decorate_skipped_candidate: G,
|
||||
) -> (LocalExecutionCandidateAttemptSource<'a>, usize)
|
||||
where
|
||||
F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync,
|
||||
G: Fn(SkippedLocalExecutionCandidate) -> SkippedLocalExecutionCandidate + Send + Sync,
|
||||
{
|
||||
let (candidates, resolved_skipped) = resolve_and_rank_logical_local_execution_candidates(
|
||||
state,
|
||||
candidates,
|
||||
client_api_format,
|
||||
requested_model,
|
||||
auth_snapshot,
|
||||
required_capabilities,
|
||||
sticky_session_token,
|
||||
request_auth_channel,
|
||||
resolution_mode,
|
||||
)
|
||||
.await;
|
||||
let skipped_candidate_count = preselection_skipped.len() + resolved_skipped.len();
|
||||
let skipped_candidates = preselection_skipped
|
||||
.into_iter()
|
||||
.chain(resolved_skipped)
|
||||
.map(decorate_skipped_candidate)
|
||||
.collect::<Vec<_>>();
|
||||
let candidate_count = candidates.len() + skipped_candidate_count;
|
||||
if persistence_policy.skipped.record_runtime_miss_diagnostic {
|
||||
for skipped_candidate in &skipped_candidates {
|
||||
record_local_runtime_candidate_skip_reason(
|
||||
state.app(),
|
||||
trace_id,
|
||||
skipped_candidate.skip_reason,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
remember_first_local_candidate_affinity(
|
||||
state,
|
||||
auth_snapshot,
|
||||
client_api_format,
|
||||
requested_model,
|
||||
&candidates,
|
||||
);
|
||||
|
||||
let mut items = VecDeque::new();
|
||||
for (candidate_index, candidate) in candidates.into_iter().enumerate() {
|
||||
let candidate_index = candidate_index as u32;
|
||||
match candidate.kind {
|
||||
LocalExecutionCandidateKind::SingleKey => {
|
||||
let attempts = build_unpersisted_local_execution_candidate_attempts(
|
||||
candidate,
|
||||
candidate_index,
|
||||
);
|
||||
if !attempts.is_empty() {
|
||||
items.push_back(LocalExecutionCandidateAttemptSourceItem::Static { attempts });
|
||||
}
|
||||
}
|
||||
LocalExecutionCandidateKind::PoolGroup => {
|
||||
items.push_back(LocalExecutionCandidateAttemptSourceItem::Pool {
|
||||
cursor: PoolKeyCursor::new(
|
||||
state,
|
||||
candidate,
|
||||
sticky_session_token,
|
||||
requested_model,
|
||||
request_auth_channel,
|
||||
),
|
||||
candidate_index,
|
||||
pending_attempts: VecDeque::new(),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
(
|
||||
LocalExecutionCandidateAttemptSource { items },
|
||||
candidate_count,
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) fn remember_first_local_candidate_affinity(
|
||||
state: PlannerAppState<'_>,
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
@@ -438,6 +599,99 @@ where
|
||||
.await
|
||||
}
|
||||
|
||||
async fn persist_available_local_execution_candidate_at_index<F>(
|
||||
state: PlannerAppState<'_>,
|
||||
trace_id: &str,
|
||||
context: LocalAvailableCandidatePersistenceContext<'_>,
|
||||
candidate: EligibleLocalExecutionCandidate,
|
||||
candidate_index: u32,
|
||||
build_extra_data: &F,
|
||||
) -> Vec<LocalExecutionCandidateAttempt>
|
||||
where
|
||||
F: Fn(&EligibleLocalExecutionCandidate) -> Option<Value> + Send + Sync,
|
||||
{
|
||||
let attempt_slots = local_attempt_slot_count(&candidate.transport).max(1);
|
||||
let extra_data = ai_candidate_extra_data_with_ranking(
|
||||
build_extra_data(&candidate),
|
||||
candidate.ranking.as_ref(),
|
||||
);
|
||||
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);
|
||||
|
||||
for retry_index in 0..attempt_slots {
|
||||
let candidate_ref = owned_candidate
|
||||
.as_ref()
|
||||
.expect("candidate should remain available until final retry");
|
||||
let generated_candidate_id = Uuid::new_v4().to_string();
|
||||
let candidate_id = if should_persist {
|
||||
state
|
||||
.persist_available_local_candidate(
|
||||
trace_id,
|
||||
context.user_id,
|
||||
context.api_key_id,
|
||||
&candidate_ref.candidate,
|
||||
candidate_index,
|
||||
retry_index,
|
||||
generated_candidate_id.as_str(),
|
||||
context.required_capabilities,
|
||||
extra_data.clone(),
|
||||
current_unix_ms(),
|
||||
context.error_context,
|
||||
)
|
||||
.await
|
||||
} else {
|
||||
generated_candidate_id
|
||||
};
|
||||
|
||||
let candidate = if retry_index + 1 == attempt_slots {
|
||||
owned_candidate
|
||||
.take()
|
||||
.expect("final retry should consume owned candidate")
|
||||
} else {
|
||||
candidate_ref.clone()
|
||||
};
|
||||
attempts.push(LocalExecutionCandidateAttempt {
|
||||
eligible: candidate,
|
||||
candidate_index,
|
||||
retry_index,
|
||||
candidate_id,
|
||||
});
|
||||
}
|
||||
|
||||
attempts
|
||||
}
|
||||
|
||||
fn build_unpersisted_local_execution_candidate_attempts(
|
||||
candidate: EligibleLocalExecutionCandidate,
|
||||
candidate_index: u32,
|
||||
) -> VecDeque<LocalExecutionCandidateAttempt> {
|
||||
let attempt_slots = local_attempt_slot_count(&candidate.transport).max(1);
|
||||
let mut attempts = VecDeque::with_capacity(attempt_slots as usize);
|
||||
let mut owned_candidate = Some(candidate);
|
||||
|
||||
for retry_index in 0..attempt_slots {
|
||||
let candidate = if retry_index + 1 == attempt_slots {
|
||||
owned_candidate
|
||||
.take()
|
||||
.expect("final retry should consume owned candidate")
|
||||
} else {
|
||||
owned_candidate
|
||||
.as_ref()
|
||||
.expect("candidate should remain available until final retry")
|
||||
.clone()
|
||||
};
|
||||
attempts.push_back(LocalExecutionCandidateAttempt {
|
||||
eligible: candidate,
|
||||
candidate_index,
|
||||
retry_index,
|
||||
candidate_id: Uuid::new_v4().to_string(),
|
||||
});
|
||||
}
|
||||
|
||||
attempts
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub(crate) async fn persist_skipped_local_execution_candidate(
|
||||
state: &AppState,
|
||||
@@ -604,6 +858,7 @@ pub(crate) async fn persist_skipped_local_execution_candidates_with_context(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::collections::VecDeque;
|
||||
use std::sync::Arc;
|
||||
|
||||
use aether_data::repository::candidates::InMemoryRequestCandidateRepository;
|
||||
@@ -707,6 +962,7 @@ mod tests {
|
||||
pool_key_index: Option<u32>,
|
||||
) -> EligibleLocalExecutionCandidate {
|
||||
EligibleLocalExecutionCandidate {
|
||||
kind: LocalExecutionCandidateKind::SingleKey,
|
||||
candidate: sample_candidate(key_id),
|
||||
transport: sample_transport(
|
||||
key_id,
|
||||
@@ -722,7 +978,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pool_group_representatives_are_persisted_as_available_before_attempt() {
|
||||
async fn pool_group_keys_are_not_persisted_as_available_before_attempt() {
|
||||
let repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
||||
let app = AppState::new()
|
||||
.expect("state should build")
|
||||
@@ -753,11 +1009,9 @@ mod tests {
|
||||
.read_request_candidates_by_request_id("trace-pool-lazy")
|
||||
.await
|
||||
.expect("request candidates should read");
|
||||
assert_eq!(stored.len(), 2);
|
||||
assert_eq!(stored[0].key_id.as_deref(), Some("pool-key"));
|
||||
assert_eq!(stored[0].candidate_index, 0);
|
||||
assert_eq!(stored[1].key_id.as_deref(), Some("normal-key"));
|
||||
assert_eq!(stored[1].candidate_index, 2);
|
||||
assert_eq!(stored.len(), 1);
|
||||
assert_eq!(stored[0].key_id.as_deref(), Some("normal-key"));
|
||||
assert_eq!(stored[0].candidate_index, 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -818,6 +1072,28 @@ mod tests {
|
||||
assert_eq!(extra_data.get("demoted_by"), Some(&json!("cross_format")));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn dynamic_attempt_source_does_not_drain_unexecuted_single_keys() {
|
||||
let mut source = LocalExecutionCandidateAttemptSource {
|
||||
items: VecDeque::from([LocalExecutionCandidateAttemptSourceItem::Static {
|
||||
attempts: build_unpersisted_local_execution_candidate_attempts(
|
||||
sample_eligible("normal-key", None),
|
||||
0,
|
||||
),
|
||||
}]),
|
||||
};
|
||||
|
||||
let first = source
|
||||
.next_attempt()
|
||||
.await
|
||||
.expect("first attempt should be available");
|
||||
assert_eq!(first.eligible.candidate.key_id, "normal-key");
|
||||
|
||||
let remaining = source.drain_static_attempts();
|
||||
assert!(remaining.is_empty());
|
||||
assert!(source.next_attempt().await.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pool_internal_skipped_candidates_are_not_persisted() {
|
||||
let repository = Arc::new(InMemoryRequestCandidateRepository::default());
|
||||
|
||||
@@ -219,7 +219,10 @@ mod tests {
|
||||
use super::super::candidate_affinity_cache::remember_scheduler_affinity_for_candidate;
|
||||
use super::super::candidate_transport_ranking_facts::resolve_cached_candidate_transport_ranking_facts;
|
||||
use super::{PlannerAppState, SchedulerMinimalCandidateSelectionCandidate};
|
||||
use crate::ai_serving::planner::candidate_resolution::resolve_and_rank_local_execution_candidates;
|
||||
use crate::ai_serving::planner::candidate_resolution::{
|
||||
resolve_and_rank_local_execution_candidates,
|
||||
resolve_and_rank_logical_local_execution_candidates, LocalExecutionCandidateKind,
|
||||
};
|
||||
use crate::data::auth::GatewayAuthApiKeySnapshot;
|
||||
use crate::data::GatewayDataState;
|
||||
use crate::tunnel::TunnelAttachmentRecord;
|
||||
@@ -1708,7 +1711,7 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pool_key_affinity_promotes_pool_group_when_cached_key_is_inactive() {
|
||||
async fn pool_key_affinity_promotes_logical_pool_group_when_cached_key_is_inactive() {
|
||||
let provider_catalog = InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![
|
||||
sample_provider_with_options("provider-priority", false, 0),
|
||||
@@ -1764,18 +1767,10 @@ mod tests {
|
||||
&cached_candidate,
|
||||
);
|
||||
|
||||
let (ranked, skipped) = resolve_and_rank_local_execution_candidates(
|
||||
let (ranked, skipped) = resolve_and_rank_logical_local_execution_candidates(
|
||||
PlannerAppState::new(&state),
|
||||
vec![
|
||||
cached_candidate,
|
||||
sample_priority_candidate(
|
||||
"provider-pool",
|
||||
"endpoint-pool",
|
||||
"key-fallback",
|
||||
"openai:chat",
|
||||
Some(10),
|
||||
10,
|
||||
),
|
||||
sample_priority_candidate(
|
||||
"provider-priority",
|
||||
"endpoint-priority",
|
||||
@@ -1786,16 +1781,18 @@ mod tests {
|
||||
),
|
||||
],
|
||||
"openai:chat",
|
||||
"gpt-4.1",
|
||||
Some("gpt-4.1"),
|
||||
Some(&auth_snapshot),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
aether_ai_serving::AiCandidateResolutionMode::Standard,
|
||||
)
|
||||
.await;
|
||||
|
||||
assert_eq!(ranked[0].candidate.key_id, "key-fallback");
|
||||
assert_eq!(ranked[0].orchestration.pool_key_index, Some(0));
|
||||
assert_eq!(ranked[0].candidate.key_id, "key-cached");
|
||||
assert_eq!(ranked[0].kind, LocalExecutionCandidateKind::PoolGroup);
|
||||
assert_eq!(ranked[0].orchestration.pool_key_index, None);
|
||||
assert_eq!(
|
||||
ranked[0]
|
||||
.ranking
|
||||
@@ -1803,17 +1800,11 @@ mod tests {
|
||||
.and_then(|ranking| ranking.promoted_by),
|
||||
Some(RANKING_REASON_CACHED_AFFINITY)
|
||||
);
|
||||
assert_eq!(
|
||||
skipped
|
||||
.iter()
|
||||
.map(|item| (item.candidate.key_id.as_str(), item.skip_reason))
|
||||
.collect::<Vec<_>>(),
|
||||
vec![("key-cached", "key_inactive")]
|
||||
);
|
||||
assert!(skipped.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pool_key_affinity_promotes_pool_group_when_cached_key_is_blocked() {
|
||||
async fn pool_key_affinity_promotes_logical_pool_group_when_cached_key_is_blocked() {
|
||||
let mut cached_key = sample_key_for_provider("provider-pool", "key-cached", "");
|
||||
cached_key.oauth_invalid_reason =
|
||||
Some("[ACCOUNT_BLOCK] account has been deactivated".to_string());
|
||||
@@ -1866,18 +1857,10 @@ mod tests {
|
||||
&cached_candidate,
|
||||
);
|
||||
|
||||
let (ranked, skipped) = resolve_and_rank_local_execution_candidates(
|
||||
let (ranked, skipped) = resolve_and_rank_logical_local_execution_candidates(
|
||||
PlannerAppState::new(&state),
|
||||
vec![
|
||||
cached_candidate,
|
||||
sample_priority_candidate(
|
||||
"provider-pool",
|
||||
"endpoint-pool",
|
||||
"key-fallback",
|
||||
"openai:chat",
|
||||
Some(10),
|
||||
10,
|
||||
),
|
||||
sample_priority_candidate(
|
||||
"provider-priority",
|
||||
"endpoint-priority",
|
||||
@@ -1888,15 +1871,18 @@ mod tests {
|
||||
),
|
||||
],
|
||||
"openai:chat",
|
||||
"gpt-4.1",
|
||||
Some("gpt-4.1"),
|
||||
Some(&auth_snapshot),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
aether_ai_serving::AiCandidateResolutionMode::Standard,
|
||||
)
|
||||
.await;
|
||||
|
||||
assert_eq!(ranked[0].candidate.key_id, "key-fallback");
|
||||
assert_eq!(ranked[0].candidate.key_id, "key-cached");
|
||||
assert_eq!(ranked[0].kind, LocalExecutionCandidateKind::PoolGroup);
|
||||
assert_eq!(ranked[0].orchestration.pool_key_index, None);
|
||||
assert_eq!(
|
||||
ranked[0]
|
||||
.ranking
|
||||
@@ -1904,13 +1890,7 @@ mod tests {
|
||||
.and_then(|ranking| ranking.promoted_by),
|
||||
Some(RANKING_REASON_CACHED_AFFINITY)
|
||||
);
|
||||
assert_eq!(
|
||||
skipped
|
||||
.iter()
|
||||
.map(|item| (item.candidate.key_id.as_str(), item.skip_reason))
|
||||
.collect::<Vec<_>>(),
|
||||
vec![("key-cached", "pool_account_blocked")]
|
||||
);
|
||||
assert!(skipped.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -22,6 +22,7 @@ use super::pool_scheduler::apply_local_execution_pool_scheduler;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub(crate) struct EligibleLocalExecutionCandidate {
|
||||
pub(crate) kind: LocalExecutionCandidateKind,
|
||||
pub(crate) candidate: SchedulerMinimalCandidateSelectionCandidate,
|
||||
pub(crate) transport: Arc<GatewayProviderTransportSnapshot>,
|
||||
pub(crate) provider_api_format: String,
|
||||
@@ -29,6 +30,13 @@ pub(crate) struct EligibleLocalExecutionCandidate {
|
||||
pub(crate) ranking: Option<SchedulerRankingOutcome>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
|
||||
pub(crate) enum LocalExecutionCandidateKind {
|
||||
#[default]
|
||||
SingleKey,
|
||||
PoolGroup,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub(crate) struct SkippedLocalExecutionCandidate {
|
||||
pub(crate) candidate: SchedulerMinimalCandidateSelectionCandidate,
|
||||
@@ -87,6 +95,9 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
|
||||
transport: &Self::Transport,
|
||||
requested_model: Option<&str>,
|
||||
) -> Option<&'static str> {
|
||||
if provider_transport_uses_pool(transport) {
|
||||
return pool_group_common_transport_skip_reason(candidate, transport);
|
||||
}
|
||||
if let Some(skip_reason) =
|
||||
candidate_auth_channel_skip_reason(transport, self.request_auth_channel)
|
||||
{
|
||||
@@ -131,7 +142,13 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
|
||||
transport: Self::Transport,
|
||||
) -> Self::Eligible {
|
||||
let provider_api_format = transport.endpoint.api_format.trim().to_ascii_lowercase();
|
||||
let kind = if provider_transport_uses_pool(&transport) {
|
||||
LocalExecutionCandidateKind::PoolGroup
|
||||
} else {
|
||||
LocalExecutionCandidateKind::SingleKey
|
||||
};
|
||||
EligibleLocalExecutionCandidate {
|
||||
kind,
|
||||
candidate,
|
||||
transport: Arc::new(transport),
|
||||
provider_api_format,
|
||||
@@ -160,10 +177,14 @@ impl AiCandidateResolutionPort for GatewayLocalCandidateResolutionPort<'_> {
|
||||
&self,
|
||||
candidates: Vec<Self::Eligible>,
|
||||
) -> Result<(Vec<Self::Eligible>, Vec<Self::Skipped>), Self::Error> {
|
||||
Ok(
|
||||
apply_local_execution_pool_scheduler(self.state, candidates, self.sticky_session_token)
|
||||
.await,
|
||||
Ok(apply_local_execution_pool_scheduler(
|
||||
self.state,
|
||||
candidates,
|
||||
self.sticky_session_token,
|
||||
self.requested_model,
|
||||
self.request_auth_channel,
|
||||
)
|
||||
.await)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -223,6 +244,35 @@ pub(crate) async fn resolve_and_rank_local_execution_candidates_without_transpor
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn resolve_and_rank_logical_local_execution_candidates(
|
||||
state: PlannerAppState<'_>,
|
||||
candidates: Vec<SchedulerMinimalCandidateSelectionCandidate>,
|
||||
client_api_format: &str,
|
||||
requested_model: Option<&str>,
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
required_capabilities: Option<&serde_json::Value>,
|
||||
sticky_session_token: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
mode: AiCandidateResolutionMode,
|
||||
) -> (
|
||||
Vec<EligibleLocalExecutionCandidate>,
|
||||
Vec<SkippedLocalExecutionCandidate>,
|
||||
) {
|
||||
resolve_and_rank_local_execution_candidates_with_pool_expansion(
|
||||
state,
|
||||
candidates,
|
||||
client_api_format,
|
||||
requested_model,
|
||||
auth_snapshot,
|
||||
required_capabilities,
|
||||
sticky_session_token,
|
||||
request_auth_channel,
|
||||
mode,
|
||||
false,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn resolve_and_rank_local_execution_candidates_with_mode(
|
||||
state: PlannerAppState<'_>,
|
||||
candidates: Vec<SchedulerMinimalCandidateSelectionCandidate>,
|
||||
@@ -236,6 +286,37 @@ async fn resolve_and_rank_local_execution_candidates_with_mode(
|
||||
) -> (
|
||||
Vec<EligibleLocalExecutionCandidate>,
|
||||
Vec<SkippedLocalExecutionCandidate>,
|
||||
) {
|
||||
resolve_and_rank_local_execution_candidates_with_pool_expansion(
|
||||
state,
|
||||
candidates,
|
||||
client_api_format,
|
||||
requested_model,
|
||||
auth_snapshot,
|
||||
required_capabilities,
|
||||
sticky_session_token,
|
||||
request_auth_channel,
|
||||
mode,
|
||||
true,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn resolve_and_rank_local_execution_candidates_with_pool_expansion(
|
||||
state: PlannerAppState<'_>,
|
||||
candidates: Vec<SchedulerMinimalCandidateSelectionCandidate>,
|
||||
client_api_format: &str,
|
||||
requested_model: Option<&str>,
|
||||
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
|
||||
required_capabilities: Option<&serde_json::Value>,
|
||||
sticky_session_token: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
mode: AiCandidateResolutionMode,
|
||||
expand_pool_groups: bool,
|
||||
) -> (
|
||||
Vec<EligibleLocalExecutionCandidate>,
|
||||
Vec<SkippedLocalExecutionCandidate>,
|
||||
) {
|
||||
let port = GatewayLocalCandidateResolutionPort {
|
||||
state,
|
||||
@@ -250,6 +331,7 @@ async fn resolve_and_rank_local_execution_candidates_with_mode(
|
||||
client_api_format,
|
||||
requested_model,
|
||||
mode,
|
||||
expand_pool_groups,
|
||||
};
|
||||
|
||||
match run_ai_candidate_resolution(&port, candidates, request).await {
|
||||
@@ -269,7 +351,33 @@ fn candidate_transport_policy_facts(
|
||||
}
|
||||
}
|
||||
|
||||
fn candidate_auth_channel_skip_reason(
|
||||
fn provider_transport_uses_pool(transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
crate::handlers::shared::provider_pool::admin_provider_pool_config_from_config_value(
|
||||
transport.provider.config.as_ref(),
|
||||
)
|
||||
.is_some()
|
||||
}
|
||||
|
||||
fn pool_group_common_transport_skip_reason(
|
||||
candidate: &SchedulerMinimalCandidateSelectionCandidate,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<&'static str> {
|
||||
if !transport.provider.is_active {
|
||||
return Some("provider_inactive");
|
||||
}
|
||||
if !transport.endpoint.is_active {
|
||||
return Some("endpoint_inactive");
|
||||
}
|
||||
if !crate::ai_serving::api_format_alias_matches(
|
||||
candidate.endpoint_api_format.as_str(),
|
||||
transport.endpoint.api_format.trim(),
|
||||
) {
|
||||
return Some("endpoint_api_format_changed");
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
pub(crate) fn candidate_auth_channel_skip_reason(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
request_auth_channel: Option<&str>,
|
||||
) -> Option<&'static str> {
|
||||
@@ -381,12 +489,13 @@ pub(crate) async fn read_candidate_transport_snapshot(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::candidate_auth_channel_skip_reason;
|
||||
use super::{candidate_auth_channel_skip_reason, pool_group_common_transport_skip_reason};
|
||||
use crate::ai_serving::GatewayProviderTransportSnapshot;
|
||||
use aether_provider_transport::snapshot::{
|
||||
GatewayProviderTransportEndpoint, GatewayProviderTransportKey,
|
||||
GatewayProviderTransportProvider,
|
||||
};
|
||||
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
|
||||
use serde_json::json;
|
||||
|
||||
fn sample_transport(auth_type: &str) -> GatewayProviderTransportSnapshot {
|
||||
@@ -444,6 +553,28 @@ mod tests {
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_candidate() -> SchedulerMinimalCandidateSelectionCandidate {
|
||||
SchedulerMinimalCandidateSelectionCandidate {
|
||||
provider_id: "provider-1".to_string(),
|
||||
provider_name: "provider".to_string(),
|
||||
provider_type: "custom".to_string(),
|
||||
provider_priority: 10,
|
||||
endpoint_id: "endpoint-1".to_string(),
|
||||
endpoint_api_format: "claude:messages".to_string(),
|
||||
key_id: "key-1".to_string(),
|
||||
key_name: "key".to_string(),
|
||||
key_auth_type: "bearer".to_string(),
|
||||
key_internal_priority: 10,
|
||||
key_global_priority_for_format: None,
|
||||
key_capabilities: None,
|
||||
model_id: "model-1".to_string(),
|
||||
global_model_id: "global-model-1".to_string(),
|
||||
global_model_name: "claude-sonnet".to_string(),
|
||||
selected_provider_model_name: "claude-sonnet".to_string(),
|
||||
mapping_matched_model: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auth_channel_gate_skips_mismatched_raw_secret_auth() {
|
||||
let transport = sample_transport("bearer");
|
||||
@@ -476,4 +607,17 @@ mod tests {
|
||||
Some("auth_channel_mismatch")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn pool_group_common_gate_ignores_representative_key_model_policy() {
|
||||
let candidate = sample_candidate();
|
||||
let mut transport = sample_transport("bearer");
|
||||
transport.key.allowed_models = Some(vec!["different-model".to_string()]);
|
||||
transport.key.api_formats = Some(vec!["different:format".to_string()]);
|
||||
|
||||
assert_eq!(
|
||||
pool_group_common_transport_skip_reason(&candidate, &transport),
|
||||
None
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -24,8 +24,10 @@ mod specialized;
|
||||
mod standard;
|
||||
mod state;
|
||||
|
||||
pub(crate) use self::candidate_materialization::LocalExecutionAttemptSource;
|
||||
pub(crate) use self::passthrough::{
|
||||
build_local_same_format_stream_plan_and_reports, build_local_same_format_sync_plan_and_reports,
|
||||
build_local_same_format_stream_attempt_source, build_local_same_format_stream_plan_and_reports,
|
||||
build_local_same_format_sync_attempt_source, build_local_same_format_sync_plan_and_reports,
|
||||
};
|
||||
pub(crate) use self::plan_builders::{
|
||||
build_gemini_stream_plan_from_decision, build_gemini_sync_plan_from_decision,
|
||||
@@ -36,18 +38,29 @@ pub(crate) use self::plan_builders::{
|
||||
};
|
||||
pub(crate) use self::route::is_matching_stream_request as planner_is_matching_stream_request;
|
||||
pub(crate) use self::specialized::{
|
||||
build_local_gemini_files_stream_attempt_source_for_kind,
|
||||
build_local_gemini_files_stream_plan_and_reports_for_kind,
|
||||
build_local_gemini_files_sync_attempt_source_for_kind,
|
||||
build_local_gemini_files_sync_plan_and_reports_for_kind,
|
||||
build_local_image_stream_attempt_source_for_kind,
|
||||
build_local_image_stream_plan_and_reports_for_kind,
|
||||
build_local_image_sync_attempt_source_for_kind,
|
||||
build_local_image_sync_plan_and_reports_for_kind,
|
||||
build_local_video_sync_attempt_source_for_kind,
|
||||
build_local_video_sync_plan_and_reports_for_kind,
|
||||
};
|
||||
pub(crate) use self::standard::{
|
||||
build_local_openai_chat_stream_attempt_source_for_kind,
|
||||
build_local_openai_chat_stream_plan_and_reports_for_kind,
|
||||
build_local_openai_chat_sync_attempt_source_for_kind,
|
||||
build_local_openai_chat_sync_plan_and_reports_for_kind,
|
||||
build_local_openai_responses_stream_attempt_source_for_kind,
|
||||
build_local_openai_responses_stream_plan_and_reports_for_kind,
|
||||
build_local_openai_responses_sync_attempt_source_for_kind,
|
||||
build_local_openai_responses_sync_plan_and_reports_for_kind,
|
||||
build_local_stream_attempt_source as build_standard_family_stream_attempt_source,
|
||||
build_local_stream_plan_and_reports as build_standard_family_stream_plan_and_reports,
|
||||
build_local_sync_attempt_source as build_standard_family_sync_attempt_source,
|
||||
build_local_sync_plan_and_reports as build_standard_family_sync_plan_and_reports,
|
||||
set_local_openai_chat_execution_exhausted_diagnostic,
|
||||
};
|
||||
|
||||
@@ -3,7 +3,9 @@
|
||||
mod provider;
|
||||
|
||||
pub(crate) use self::provider::{
|
||||
build_local_stream_attempt_source as build_local_same_format_stream_attempt_source,
|
||||
build_local_stream_plan_and_reports as build_local_same_format_stream_plan_and_reports,
|
||||
build_local_sync_attempt_source as build_local_same_format_sync_attempt_source,
|
||||
build_local_sync_plan_and_reports as build_local_same_format_sync_plan_and_reports,
|
||||
maybe_build_local_same_format_provider_decision_payload_for_candidate,
|
||||
maybe_build_stream_local_same_format_provider_decision_payload,
|
||||
|
||||
@@ -62,17 +62,20 @@ mod plans;
|
||||
mod request;
|
||||
|
||||
pub(crate) use self::family::{
|
||||
build_local_same_format_provider_candidate_attempt_source,
|
||||
materialize_local_same_format_provider_candidate_attempts,
|
||||
maybe_build_local_same_format_provider_decision_payload_for_candidate,
|
||||
resolve_local_same_format_provider_decision_input, LocalSameFormatProviderFamily,
|
||||
LocalSameFormatProviderSpec,
|
||||
resolve_local_same_format_provider_decision_input, LocalSameFormatProviderCandidateAttempt,
|
||||
LocalSameFormatProviderCandidateAttemptSource, LocalSameFormatProviderDecisionInput,
|
||||
LocalSameFormatProviderFamily, LocalSameFormatProviderSpec,
|
||||
};
|
||||
pub(crate) use self::family::{
|
||||
maybe_build_stream_local_same_format_provider_decision_payload,
|
||||
maybe_build_sync_local_same_format_provider_decision_payload,
|
||||
};
|
||||
pub(crate) use self::plans::{
|
||||
build_local_stream_plan_and_reports, build_local_sync_plan_and_reports,
|
||||
build_local_stream_attempt_source, build_local_stream_plan_and_reports,
|
||||
build_local_sync_attempt_source, build_local_sync_plan_and_reports,
|
||||
};
|
||||
|
||||
const ANTIGRAVITY_ENVELOPE_NAME: &str = "antigravity:v1internal";
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
use tracing::warn;
|
||||
|
||||
use crate::ai_serving::planner::candidate_materialization::{
|
||||
build_local_execution_candidate_attempt_source_with_serving,
|
||||
materialize_local_execution_candidates_with_serving, LocalCandidateResolutionMode,
|
||||
LocalExecutionCandidateAttemptSource,
|
||||
};
|
||||
use crate::ai_serving::planner::candidate_metadata::{
|
||||
build_local_execution_candidate_contract_metadata,
|
||||
@@ -169,3 +171,96 @@ pub(crate) async fn materialize_local_same_format_provider_candidate_attempts(
|
||||
|
||||
Ok((outcome.attempts, outcome.candidate_count))
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_same_format_provider_candidate_attempt_source<'a>(
|
||||
state: &'a AppState,
|
||||
trace_id: &str,
|
||||
input: &LocalSameFormatProviderDecisionInput,
|
||||
body_json: &serde_json::Value,
|
||||
spec: LocalSameFormatProviderSpec,
|
||||
) -> Result<(LocalExecutionCandidateAttemptSource<'a>, usize), GatewayError> {
|
||||
let spec_metadata = local_same_format_provider_spec_metadata(spec);
|
||||
let planner_state = PlannerAppState::new(state);
|
||||
let sticky_session_token = extract_pool_sticky_session_token(body_json);
|
||||
let persistence_policy = build_local_candidate_persistence_policy(
|
||||
&input.auth_context,
|
||||
input.required_capabilities.as_ref(),
|
||||
LocalCandidatePersistencePolicyKind::SameFormatProviderDecision,
|
||||
);
|
||||
let (candidates, preselection_skipped) = planner_state
|
||||
.list_selectable_candidates_with_skip_reasons(
|
||||
spec_metadata.api_format,
|
||||
&input.requested_model,
|
||||
spec_metadata.require_streaming,
|
||||
input.required_capabilities.as_ref(),
|
||||
Some(&input.auth_snapshot),
|
||||
current_unix_secs(),
|
||||
)
|
||||
.await?;
|
||||
|
||||
Ok(build_local_execution_candidate_attempt_source_with_serving(
|
||||
planner_state,
|
||||
trace_id,
|
||||
spec_metadata.api_format,
|
||||
Some(&input.requested_model),
|
||||
Some(&input.auth_snapshot),
|
||||
input.required_capabilities.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
candidates,
|
||||
preselection_skipped
|
||||
.into_iter()
|
||||
.map(|item| SkippedLocalExecutionCandidate {
|
||||
candidate: item.candidate,
|
||||
skip_reason: item.skip_reason,
|
||||
transport: None,
|
||||
ranking: None,
|
||||
extra_data: None,
|
||||
})
|
||||
.collect(),
|
||||
LocalCandidateResolutionMode::Standard,
|
||||
|eligible| {
|
||||
let (execution_strategy, conversion_mode) = ai_local_execution_contract_for_formats(
|
||||
spec_metadata.api_format,
|
||||
spec_metadata.api_format,
|
||||
);
|
||||
Some(build_local_execution_candidate_contract_metadata(
|
||||
LocalExecutionCandidateMetadataParts {
|
||||
eligible,
|
||||
provider_api_format: spec_metadata.api_format,
|
||||
client_api_format: spec_metadata.api_format,
|
||||
extra_fields: serde_json::Map::new(),
|
||||
},
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
spec_metadata.api_format,
|
||||
))
|
||||
},
|
||||
|mut skipped_candidate| {
|
||||
let provider_api_format = skipped_candidate
|
||||
.transport
|
||||
.as_ref()
|
||||
.map(|transport| transport.endpoint.api_format.trim().to_ascii_lowercase())
|
||||
.unwrap_or_else(|| spec_metadata.api_format.to_string());
|
||||
let (execution_strategy, conversion_mode) = ai_local_execution_contract_for_formats(
|
||||
spec_metadata.api_format,
|
||||
provider_api_format.as_str(),
|
||||
);
|
||||
skipped_candidate.extra_data = Some(
|
||||
build_local_execution_candidate_contract_metadata_for_candidate(
|
||||
&skipped_candidate.candidate,
|
||||
skipped_candidate.transport_ref(),
|
||||
provider_api_format.as_str(),
|
||||
spec_metadata.api_format,
|
||||
serde_json::Map::new(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
provider_api_format.as_str(),
|
||||
),
|
||||
);
|
||||
skipped_candidate
|
||||
},
|
||||
)
|
||||
.await)
|
||||
}
|
||||
|
||||
@@ -8,10 +8,12 @@ pub(crate) use self::build::{
|
||||
maybe_build_sync_local_same_format_provider_decision_payload,
|
||||
};
|
||||
pub(crate) use self::candidates::{
|
||||
build_local_same_format_provider_candidate_attempt_source,
|
||||
materialize_local_same_format_provider_candidate_attempts,
|
||||
resolve_local_same_format_provider_decision_input,
|
||||
};
|
||||
pub(crate) use self::payload::maybe_build_local_same_format_provider_decision_payload_for_candidate;
|
||||
pub(crate) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttempt as LocalSameFormatProviderCandidateAttempt;
|
||||
pub(crate) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttemptSource as LocalSameFormatProviderCandidateAttemptSource;
|
||||
pub(crate) use crate::ai_serving::planner::decision_input::LocalRequestedModelDecisionInput as LocalSameFormatProviderDecisionInput;
|
||||
pub(crate) use crate::ai_serving::{LocalSameFormatProviderFamily, LocalSameFormatProviderSpec};
|
||||
|
||||
@@ -1,6 +1,10 @@
|
||||
use async_trait::async_trait;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::ai_serving::planner::common::extract_requested_model_from_request;
|
||||
use crate::ai_serving::planner::candidate_materialization::LocalExecutionAttemptSource;
|
||||
use crate::ai_serving::planner::common::{
|
||||
extract_requested_model_from_request, RequestedModelFamily,
|
||||
};
|
||||
use crate::ai_serving::planner::runtime_miss::{
|
||||
apply_local_runtime_candidate_evaluation_progress_preserving_candidate_signal,
|
||||
apply_local_runtime_candidate_terminal_reason, set_local_runtime_miss_diagnostic_reason,
|
||||
@@ -15,12 +19,299 @@ pub(crate) use crate::ai_serving::{
|
||||
};
|
||||
|
||||
use super::{
|
||||
build_local_same_format_provider_candidate_attempt_source,
|
||||
materialize_local_same_format_provider_candidate_attempts,
|
||||
maybe_build_local_same_format_provider_decision_payload_for_candidate,
|
||||
resolve_local_same_format_provider_decision_input, AiStreamAttempt, AiSyncAttempt, AppState,
|
||||
GatewayControlDecision, GatewayError, LocalSameFormatProviderSpec,
|
||||
GatewayControlDecision, GatewayError, LocalSameFormatProviderCandidateAttempt,
|
||||
LocalSameFormatProviderCandidateAttemptSource, LocalSameFormatProviderDecisionInput,
|
||||
LocalSameFormatProviderSpec,
|
||||
};
|
||||
|
||||
pub(crate) struct LocalSameFormatProviderSyncAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
input: LocalSameFormatProviderDecisionInput,
|
||||
spec: LocalSameFormatProviderSpec,
|
||||
requested_model_family: RequestedModelFamily,
|
||||
candidates: LocalSameFormatProviderCandidateAttemptSource<'a>,
|
||||
}
|
||||
|
||||
pub(crate) struct LocalSameFormatProviderStreamAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
input: LocalSameFormatProviderDecisionInput,
|
||||
spec: LocalSameFormatProviderSpec,
|
||||
requested_model_family: RequestedModelFamily,
|
||||
candidates: LocalSameFormatProviderCandidateAttemptSource<'a>,
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_sync_attempt_source<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
body_json: &'a serde_json::Value,
|
||||
spec: LocalSameFormatProviderSpec,
|
||||
) -> Result<Option<(LocalSameFormatProviderSyncAttemptSource<'a>, usize)>, GatewayError> {
|
||||
let spec_metadata = local_same_format_provider_spec_metadata(spec);
|
||||
let requested_model_family = spec_metadata
|
||||
.requested_model_family
|
||||
.expect("same-format provider spec metadata should include requested-model family");
|
||||
let Some(input) = resolve_local_same_format_provider_decision_input(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
spec_metadata.decision_kind,
|
||||
extract_requested_model_from_request(parts, body_json, requested_model_family)
|
||||
.as_deref(),
|
||||
"decision_input_unavailable",
|
||||
);
|
||||
return Ok(None);
|
||||
};
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
spec_metadata.decision_kind,
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let (candidates, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, spec,
|
||||
)
|
||||
.await?;
|
||||
apply_local_runtime_candidate_evaluation_progress_preserving_candidate_signal(
|
||||
state,
|
||||
trace_id,
|
||||
candidate_count,
|
||||
);
|
||||
if candidate_count == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
Ok(Some((
|
||||
LocalSameFormatProviderSyncAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
input,
|
||||
spec,
|
||||
requested_model_family,
|
||||
candidates,
|
||||
},
|
||||
candidate_count,
|
||||
)))
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_stream_attempt_source<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
body_json: &'a serde_json::Value,
|
||||
spec: LocalSameFormatProviderSpec,
|
||||
) -> Result<Option<(LocalSameFormatProviderStreamAttemptSource<'a>, usize)>, GatewayError> {
|
||||
let spec_metadata = local_same_format_provider_spec_metadata(spec);
|
||||
let requested_model_family = spec_metadata
|
||||
.requested_model_family
|
||||
.expect("same-format provider spec metadata should include requested-model family");
|
||||
let Some(input) = resolve_local_same_format_provider_decision_input(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
spec_metadata.decision_kind,
|
||||
extract_requested_model_from_request(parts, body_json, requested_model_family)
|
||||
.as_deref(),
|
||||
"decision_input_unavailable",
|
||||
);
|
||||
return Ok(None);
|
||||
};
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
spec_metadata.decision_kind,
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let (candidates, candidate_count) = build_local_same_format_provider_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, spec,
|
||||
)
|
||||
.await?;
|
||||
apply_local_runtime_candidate_evaluation_progress_preserving_candidate_signal(
|
||||
state,
|
||||
trace_id,
|
||||
candidate_count,
|
||||
);
|
||||
if candidate_count == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
Ok(Some((
|
||||
LocalSameFormatProviderStreamAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
input,
|
||||
spec,
|
||||
requested_model_family,
|
||||
candidates,
|
||||
},
|
||||
candidate_count,
|
||||
)))
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalSameFormatProviderSyncAttemptSource<'_> {
|
||||
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
|
||||
while let Some(attempt) = self.candidates.next_attempt().await {
|
||||
match self.build_sync_attempt(attempt).await? {
|
||||
Some(attempt) => return Ok(Some(attempt)),
|
||||
None => continue,
|
||||
}
|
||||
}
|
||||
apply_local_runtime_candidate_terminal_reason(
|
||||
self.state,
|
||||
self.trace_id,
|
||||
"no_local_sync_plans",
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn drain_execution_attempts(&mut self) -> Result<Vec<AiSyncAttempt>, GatewayError> {
|
||||
let mut drained = Vec::new();
|
||||
for attempt in self.candidates.drain_static_attempts() {
|
||||
if let Some(attempt) = self.build_sync_attempt(attempt).await? {
|
||||
drained.push(attempt);
|
||||
}
|
||||
}
|
||||
Ok(drained)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalExecutionAttemptSource<AiStreamAttempt>
|
||||
for LocalSameFormatProviderStreamAttemptSource<'_>
|
||||
{
|
||||
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
|
||||
while let Some(attempt) = self.candidates.next_attempt().await {
|
||||
match self.build_stream_attempt(attempt).await? {
|
||||
Some(attempt) => return Ok(Some(attempt)),
|
||||
None => continue,
|
||||
}
|
||||
}
|
||||
apply_local_runtime_candidate_terminal_reason(
|
||||
self.state,
|
||||
self.trace_id,
|
||||
"no_local_stream_plans",
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn drain_execution_attempts(&mut self) -> Result<Vec<AiStreamAttempt>, GatewayError> {
|
||||
let mut drained = Vec::new();
|
||||
for attempt in self.candidates.drain_static_attempts() {
|
||||
if let Some(attempt) = self.build_stream_attempt(attempt).await? {
|
||||
drained.push(attempt);
|
||||
}
|
||||
}
|
||||
Ok(drained)
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalSameFormatProviderSyncAttemptSource<'_> {
|
||||
async fn build_sync_attempt(
|
||||
&self,
|
||||
attempt: LocalSameFormatProviderCandidateAttempt,
|
||||
) -> Result<Option<AiSyncAttempt>, GatewayError> {
|
||||
let Some(payload) = maybe_build_local_same_format_provider_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match build_sync_plan_from_requested_model_family(
|
||||
self.requested_model_family,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
payload,
|
||||
) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %self.trace_id,
|
||||
error = ?err,
|
||||
"gateway local same-format sync decision plan build failed"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalSameFormatProviderStreamAttemptSource<'_> {
|
||||
async fn build_stream_attempt(
|
||||
&self,
|
||||
attempt: LocalSameFormatProviderCandidateAttempt,
|
||||
) -> Result<Option<AiStreamAttempt>, GatewayError> {
|
||||
let Some(payload) = maybe_build_local_same_format_provider_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match build_stream_plan_from_requested_model_family(
|
||||
self.requested_model_family,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
payload,
|
||||
) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %self.trace_id,
|
||||
error = ?err,
|
||||
"gateway local same-format stream decision plan build failed"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_sync_plan_and_reports(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
use std::collections::{btree_map::Entry, BTreeMap, BTreeSet};
|
||||
use std::collections::{btree_map::Entry, BTreeMap, BTreeSet, VecDeque};
|
||||
use std::sync::atomic::{AtomicU64, Ordering as AtomicOrdering};
|
||||
|
||||
use aether_ai_serving::{
|
||||
@@ -6,14 +6,20 @@ use aether_ai_serving::{
|
||||
AiPoolCandidateOrchestration, AiPoolCatalogKeyContext, AiPoolRuntimeState,
|
||||
AiPoolSchedulingConfig, AiPoolSchedulingPreset,
|
||||
};
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery,
|
||||
};
|
||||
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
|
||||
use serde_json::{Map, Value};
|
||||
use tracing::warn;
|
||||
|
||||
use crate::ai_serving::planner::candidate_resolution::{
|
||||
EligibleLocalExecutionCandidate, SkippedLocalExecutionCandidate,
|
||||
candidate_auth_channel_skip_reason, read_candidate_transport_snapshot,
|
||||
EligibleLocalExecutionCandidate, LocalExecutionCandidateKind, SkippedLocalExecutionCandidate,
|
||||
};
|
||||
use crate::ai_serving::{
|
||||
candidate_common_transport_skip_reason, CandidateTransportPolicyFacts, PlannerAppState,
|
||||
};
|
||||
use crate::ai_serving::PlannerAppState;
|
||||
use crate::clock::current_unix_ms;
|
||||
use crate::handlers::shared::provider_pool::admin_provider_pool_config_from_config_value;
|
||||
use crate::handlers::shared::provider_pool::read_admin_provider_pool_runtime_state;
|
||||
@@ -28,6 +34,8 @@ use crate::orchestration::LocalExecutionCandidateMetadata;
|
||||
use crate::provider_key_auth::provider_key_auth_semantics;
|
||||
|
||||
static LOAD_BALANCE_SEQUENCE: AtomicU64 = AtomicU64::new(0);
|
||||
const DEFAULT_POOL_KEY_PAGE_SIZE: u32 = 128;
|
||||
const DEFAULT_POOL_MAX_SCANNED_KEYS: u32 = 1024;
|
||||
|
||||
type PoolCatalogKeyContext = AiPoolCatalogKeyContext;
|
||||
|
||||
@@ -35,6 +43,8 @@ pub(crate) async fn apply_local_execution_pool_scheduler(
|
||||
state: PlannerAppState<'_>,
|
||||
candidates: Vec<EligibleLocalExecutionCandidate>,
|
||||
sticky_session_token: Option<&str>,
|
||||
requested_model: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
) -> (
|
||||
Vec<EligibleLocalExecutionCandidate>,
|
||||
Vec<SkippedLocalExecutionCandidate>,
|
||||
@@ -47,6 +57,40 @@ pub(crate) async fn apply_local_execution_pool_scheduler(
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty());
|
||||
|
||||
let mut scheduled = Vec::new();
|
||||
let mut skipped = Vec::new();
|
||||
for candidate in candidates {
|
||||
if candidate.kind == LocalExecutionCandidateKind::PoolGroup {
|
||||
let mut expanded = expand_pool_group_candidate(
|
||||
state,
|
||||
candidate,
|
||||
sticky_session_token,
|
||||
requested_model,
|
||||
request_auth_channel,
|
||||
)
|
||||
.await;
|
||||
scheduled.append(&mut expanded.0);
|
||||
skipped.append(&mut expanded.1);
|
||||
} else {
|
||||
scheduled.push(candidate);
|
||||
}
|
||||
}
|
||||
|
||||
(scheduled, skipped)
|
||||
}
|
||||
|
||||
async fn schedule_pool_page_candidates(
|
||||
state: PlannerAppState<'_>,
|
||||
candidates: Vec<EligibleLocalExecutionCandidate>,
|
||||
sticky_session_token: Option<&str>,
|
||||
) -> (
|
||||
Vec<EligibleLocalExecutionCandidate>,
|
||||
Vec<SkippedLocalExecutionCandidate>,
|
||||
) {
|
||||
if candidates.is_empty() {
|
||||
return (Vec::new(), Vec::new());
|
||||
}
|
||||
|
||||
let mut provider_runtime_requirements =
|
||||
BTreeMap::<String, (AdminProviderPoolConfig, BTreeSet<String>)>::new();
|
||||
for candidate in &candidates {
|
||||
@@ -88,6 +132,249 @@ pub(crate) async fn apply_local_execution_pool_scheduler(
|
||||
)
|
||||
}
|
||||
|
||||
async fn expand_pool_group_candidate(
|
||||
state: PlannerAppState<'_>,
|
||||
group: EligibleLocalExecutionCandidate,
|
||||
sticky_session_token: Option<&str>,
|
||||
requested_model: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
) -> (
|
||||
Vec<EligibleLocalExecutionCandidate>,
|
||||
Vec<SkippedLocalExecutionCandidate>,
|
||||
) {
|
||||
let mut cursor = PoolKeyCursor::new(
|
||||
state,
|
||||
group,
|
||||
sticky_session_token,
|
||||
requested_model,
|
||||
request_auth_channel,
|
||||
);
|
||||
let mut scheduled = Vec::new();
|
||||
let mut skipped = Vec::new();
|
||||
|
||||
while let Some(candidate) = cursor.next_key().await {
|
||||
scheduled.push(candidate);
|
||||
}
|
||||
skipped.append(&mut cursor.take_skipped_candidates());
|
||||
|
||||
if scheduled.is_empty() {
|
||||
cursor.log_exhausted();
|
||||
}
|
||||
(scheduled, skipped)
|
||||
}
|
||||
|
||||
pub(crate) struct PoolKeyCursor<'a> {
|
||||
state: PlannerAppState<'a>,
|
||||
group: EligibleLocalExecutionCandidate,
|
||||
sticky_session_token: Option<String>,
|
||||
requested_model: Option<String>,
|
||||
request_auth_channel: Option<String>,
|
||||
next_offset: u32,
|
||||
scanned_keys: u32,
|
||||
page_size: u32,
|
||||
max_scanned_keys: u32,
|
||||
skip_reason_counts: BTreeMap<&'static str, u32>,
|
||||
next_pool_key_index: u32,
|
||||
queued_candidates: VecDeque<EligibleLocalExecutionCandidate>,
|
||||
skipped_candidates: Vec<SkippedLocalExecutionCandidate>,
|
||||
exhausted_logged: bool,
|
||||
}
|
||||
|
||||
impl<'a> PoolKeyCursor<'a> {
|
||||
pub(crate) fn new(
|
||||
state: PlannerAppState<'a>,
|
||||
group: EligibleLocalExecutionCandidate,
|
||||
sticky_session_token: Option<&str>,
|
||||
requested_model: Option<&str>,
|
||||
request_auth_channel: Option<&str>,
|
||||
) -> Self {
|
||||
Self {
|
||||
state,
|
||||
group,
|
||||
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),
|
||||
next_offset: 0,
|
||||
scanned_keys: 0,
|
||||
page_size: DEFAULT_POOL_KEY_PAGE_SIZE,
|
||||
max_scanned_keys: DEFAULT_POOL_MAX_SCANNED_KEYS,
|
||||
skip_reason_counts: BTreeMap::new(),
|
||||
next_pool_key_index: 0,
|
||||
queued_candidates: VecDeque::new(),
|
||||
skipped_candidates: Vec::new(),
|
||||
exhausted_logged: false,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn next_key(&mut self) -> Option<EligibleLocalExecutionCandidate> {
|
||||
loop {
|
||||
if let Some(candidate) = self.queued_candidates.pop_front() {
|
||||
return Some(candidate);
|
||||
}
|
||||
|
||||
let page_candidates = self.next_page_candidates().await?;
|
||||
let (mut page_scheduled, mut page_skipped) = schedule_pool_page_candidates(
|
||||
self.state,
|
||||
page_candidates,
|
||||
self.sticky_session_token.as_deref(),
|
||||
)
|
||||
.await;
|
||||
for skipped_candidate in &page_skipped {
|
||||
self.record_skip_reason(skipped_candidate.skip_reason);
|
||||
}
|
||||
for candidate in &mut page_scheduled {
|
||||
candidate.orchestration.pool_key_index = Some(self.next_pool_key_index);
|
||||
self.next_pool_key_index = self.next_pool_key_index.saturating_add(1);
|
||||
}
|
||||
self.queued_candidates.extend(page_scheduled);
|
||||
self.skipped_candidates.append(&mut page_skipped);
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn take_skipped_candidates(&mut self) -> Vec<SkippedLocalExecutionCandidate> {
|
||||
std::mem::take(&mut self.skipped_candidates)
|
||||
}
|
||||
|
||||
pub(crate) fn log_exhausted(&mut self) {
|
||||
if self.exhausted_logged {
|
||||
return;
|
||||
}
|
||||
self.exhausted_logged = true;
|
||||
warn!(
|
||||
event_name = "pool_group_exhausted",
|
||||
log_type = "event",
|
||||
provider_id = %self.group.candidate.provider_id,
|
||||
endpoint_id = %self.group.candidate.endpoint_id,
|
||||
model_id = %self.group.candidate.model_id,
|
||||
scanned_keys = self.scanned_keys,
|
||||
skip_reason_counts = ?self.skip_reason_counts,
|
||||
"gateway pool scheduler exhausted pool group without a schedulable key"
|
||||
);
|
||||
}
|
||||
|
||||
async fn next_page_candidates(&mut self) -> Option<Vec<EligibleLocalExecutionCandidate>> {
|
||||
if self.scanned_keys >= self.max_scanned_keys {
|
||||
return None;
|
||||
}
|
||||
|
||||
let limit = self
|
||||
.page_size
|
||||
.min(self.max_scanned_keys - self.scanned_keys);
|
||||
let query = StoredPoolKeyCandidateRowsQuery {
|
||||
api_format: self.group.candidate.endpoint_api_format.clone(),
|
||||
provider_id: self.group.candidate.provider_id.clone(),
|
||||
endpoint_id: self.group.candidate.endpoint_id.clone(),
|
||||
model_id: self.group.candidate.model_id.clone(),
|
||||
selected_provider_model_name: self.group.candidate.selected_provider_model_name.clone(),
|
||||
offset: self.next_offset,
|
||||
limit,
|
||||
};
|
||||
let rows = match self
|
||||
.state
|
||||
.app()
|
||||
.list_pool_key_candidate_rows_for_group(&query)
|
||||
.await
|
||||
{
|
||||
Ok(rows) => rows,
|
||||
Err(err) => {
|
||||
warn!(
|
||||
event_name = "pool_group_key_page_load_failed",
|
||||
log_type = "event",
|
||||
provider_id = %self.group.candidate.provider_id,
|
||||
endpoint_id = %self.group.candidate.endpoint_id,
|
||||
model_id = %self.group.candidate.model_id,
|
||||
selected_provider_model_name = %self.group.candidate.selected_provider_model_name,
|
||||
offset = self.next_offset,
|
||||
limit,
|
||||
error = ?err,
|
||||
"gateway pool scheduler failed to read pool key page"
|
||||
);
|
||||
return None;
|
||||
}
|
||||
};
|
||||
if rows.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
self.scanned_keys += rows.len() as u32;
|
||||
self.next_offset = self.next_offset.saturating_add(rows.len() as u32);
|
||||
Some(self.build_page_eligible_candidates(rows).await)
|
||||
}
|
||||
|
||||
async fn build_page_eligible_candidates(
|
||||
&mut self,
|
||||
rows: Vec<StoredMinimalCandidateSelectionRow>,
|
||||
) -> Vec<EligibleLocalExecutionCandidate> {
|
||||
let mut candidates = Vec::with_capacity(rows.len());
|
||||
for row in rows {
|
||||
let candidate = pool_candidate_from_row(&self.group, row);
|
||||
let Some(transport) = read_candidate_transport_snapshot(self.state, &candidate).await
|
||||
else {
|
||||
self.record_skip_reason("transport_snapshot_missing");
|
||||
continue;
|
||||
};
|
||||
if let Some(skip_reason) =
|
||||
candidate_auth_channel_skip_reason(&transport, self.request_auth_channel.as_deref())
|
||||
{
|
||||
self.record_skip_reason(skip_reason);
|
||||
continue;
|
||||
}
|
||||
if let Some(skip_reason) = candidate_common_transport_skip_reason(
|
||||
&transport,
|
||||
pool_candidate_transport_policy_facts(&candidate),
|
||||
self.requested_model.as_deref(),
|
||||
) {
|
||||
self.record_skip_reason(skip_reason);
|
||||
continue;
|
||||
}
|
||||
candidates.push(EligibleLocalExecutionCandidate {
|
||||
kind: LocalExecutionCandidateKind::SingleKey,
|
||||
candidate,
|
||||
provider_api_format: transport.endpoint.api_format.trim().to_ascii_lowercase(),
|
||||
transport: std::sync::Arc::new(transport),
|
||||
orchestration: LocalExecutionCandidateMetadata::default(),
|
||||
ranking: self.group.ranking.clone(),
|
||||
});
|
||||
}
|
||||
candidates
|
||||
}
|
||||
|
||||
fn record_skip_reason(&mut self, reason: &'static str) {
|
||||
*self.skip_reason_counts.entry(reason).or_insert(0) += 1;
|
||||
}
|
||||
}
|
||||
|
||||
fn pool_candidate_transport_policy_facts(
|
||||
candidate: &aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate,
|
||||
) -> CandidateTransportPolicyFacts<'_> {
|
||||
CandidateTransportPolicyFacts {
|
||||
endpoint_api_format: candidate.endpoint_api_format.as_str(),
|
||||
global_model_name: candidate.global_model_name.as_str(),
|
||||
selected_provider_model_name: candidate.selected_provider_model_name.as_str(),
|
||||
mapping_matched_model: candidate.mapping_matched_model.as_deref(),
|
||||
}
|
||||
}
|
||||
|
||||
fn pool_candidate_from_row(
|
||||
group: &EligibleLocalExecutionCandidate,
|
||||
row: StoredMinimalCandidateSelectionRow,
|
||||
) -> aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate {
|
||||
let mut candidate = group.candidate.clone();
|
||||
candidate.key_id = row.key_id;
|
||||
candidate.key_name = row.key_name;
|
||||
candidate.key_auth_type = row.key_auth_type;
|
||||
candidate.key_internal_priority = row.key_internal_priority;
|
||||
candidate.key_global_priority_for_format =
|
||||
aether_scheduler_core::extract_global_priority_for_format(
|
||||
row.key_global_priority_by_format.as_ref(),
|
||||
group.candidate.endpoint_api_format.as_str(),
|
||||
)
|
||||
.ok()
|
||||
.flatten();
|
||||
candidate.key_capabilities = row.key_capabilities;
|
||||
candidate
|
||||
}
|
||||
|
||||
async fn read_pool_catalog_key_contexts_by_id(
|
||||
state: PlannerAppState<'_>,
|
||||
candidates: &[EligibleLocalExecutionCandidate],
|
||||
@@ -402,7 +689,9 @@ mod tests {
|
||||
apply_local_execution_pool_scheduler_with_runtime_map, build_pool_catalog_key_context,
|
||||
PoolCatalogKeyContext,
|
||||
};
|
||||
use crate::ai_serving::planner::candidate_resolution::EligibleLocalExecutionCandidate;
|
||||
use crate::ai_serving::planner::candidate_resolution::{
|
||||
EligibleLocalExecutionCandidate, LocalExecutionCandidateKind,
|
||||
};
|
||||
use crate::ai_serving::PlannerAppState;
|
||||
use crate::data::GatewayDataState;
|
||||
use crate::handlers::shared::provider_pool::AdminProviderPoolRuntimeState;
|
||||
@@ -1170,6 +1459,11 @@ mod tests {
|
||||
provider_config: Option<serde_json::Value>,
|
||||
) -> EligibleLocalExecutionCandidate {
|
||||
EligibleLocalExecutionCandidate {
|
||||
kind: if provider_config.is_some() {
|
||||
LocalExecutionCandidateKind::PoolGroup
|
||||
} else {
|
||||
LocalExecutionCandidateKind::SingleKey
|
||||
},
|
||||
candidate: SchedulerMinimalCandidateSelectionCandidate {
|
||||
provider_id: provider_id.to_string(),
|
||||
provider_name: provider_id.to_string(),
|
||||
|
||||
@@ -2,8 +2,10 @@ mod decision;
|
||||
mod request;
|
||||
mod support;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::ai_serving::planner::candidate_materialization::LocalExecutionAttemptSource;
|
||||
use crate::ai_serving::planner::plan_builders::{
|
||||
build_passthrough_stream_plan_from_decision, build_passthrough_sync_plan_from_decision,
|
||||
AiStreamAttempt, AiSyncAttempt,
|
||||
@@ -18,9 +20,33 @@ use crate::{AiExecutionDecision, AppState, GatewayError};
|
||||
|
||||
use self::decision::maybe_build_local_gemini_files_decision_payload_for_candidate;
|
||||
use self::support::{
|
||||
build_local_gemini_files_candidate_attempt_source,
|
||||
materialize_local_gemini_files_candidate_attempts, resolve_local_gemini_files_decision_input,
|
||||
LocalGeminiFilesCandidateAttempt, LocalGeminiFilesCandidateAttemptSource,
|
||||
LocalGeminiFilesDecisionInput,
|
||||
};
|
||||
|
||||
pub(crate) struct LocalGeminiFilesSyncAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_base64: Option<&'a str>,
|
||||
body_is_empty: bool,
|
||||
trace_id: &'a str,
|
||||
input: LocalGeminiFilesDecisionInput,
|
||||
spec: LocalGeminiFilesSpec,
|
||||
candidates: LocalGeminiFilesCandidateAttemptSource<'a>,
|
||||
}
|
||||
|
||||
pub(crate) struct LocalGeminiFilesStreamAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
input: LocalGeminiFilesDecisionInput,
|
||||
spec: LocalGeminiFilesSpec,
|
||||
candidates: LocalGeminiFilesCandidateAttemptSource<'a>,
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_gemini_files_sync_plan_and_reports_for_kind(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
@@ -62,6 +88,202 @@ pub(crate) async fn build_local_gemini_files_stream_plan_and_reports_for_kind(
|
||||
build_local_stream_plan_and_reports(state, parts, trace_id, decision, spec).await
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub(crate) async fn build_local_gemini_files_sync_attempt_source_for_kind<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_base64: Option<&'a str>,
|
||||
body_is_empty: bool,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
plan_kind: &str,
|
||||
) -> Result<Option<(LocalGeminiFilesSyncAttemptSource<'a>, usize)>, GatewayError> {
|
||||
let Some(spec) = resolve_sync_spec(plan_kind) else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let Some(input) = resolve_local_gemini_files_decision_input(state, trace_id, decision).await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let (candidates, candidate_count) =
|
||||
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
|
||||
if candidate_count == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
Ok(Some((
|
||||
LocalGeminiFilesSyncAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
body_base64,
|
||||
body_is_empty,
|
||||
trace_id,
|
||||
input,
|
||||
spec,
|
||||
candidates,
|
||||
},
|
||||
candidate_count,
|
||||
)))
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_gemini_files_stream_attempt_source_for_kind<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
plan_kind: &str,
|
||||
) -> Result<Option<(LocalGeminiFilesStreamAttemptSource<'a>, usize)>, GatewayError> {
|
||||
let Some(spec) = resolve_stream_spec(plan_kind) else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let Some(input) = resolve_local_gemini_files_decision_input(state, trace_id, decision).await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let (candidates, candidate_count) =
|
||||
build_local_gemini_files_candidate_attempt_source(state, trace_id, &input).await?;
|
||||
if candidate_count == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
Ok(Some((
|
||||
LocalGeminiFilesStreamAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
input,
|
||||
spec,
|
||||
candidates,
|
||||
},
|
||||
candidate_count,
|
||||
)))
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalGeminiFilesSyncAttemptSource<'_> {
|
||||
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
|
||||
while let Some(attempt) = self.candidates.next_attempt().await {
|
||||
match self.build_sync_attempt(attempt).await? {
|
||||
Some(attempt) => return Ok(Some(attempt)),
|
||||
None => continue,
|
||||
}
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn drain_execution_attempts(&mut self) -> Result<Vec<AiSyncAttempt>, GatewayError> {
|
||||
let mut drained = Vec::new();
|
||||
for attempt in self.candidates.drain_static_attempts() {
|
||||
if let Some(attempt) = self.build_sync_attempt(attempt).await? {
|
||||
drained.push(attempt);
|
||||
}
|
||||
}
|
||||
Ok(drained)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalGeminiFilesStreamAttemptSource<'_> {
|
||||
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
|
||||
while let Some(attempt) = self.candidates.next_attempt().await {
|
||||
match self.build_stream_attempt(attempt).await? {
|
||||
Some(attempt) => return Ok(Some(attempt)),
|
||||
None => continue,
|
||||
}
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn drain_execution_attempts(&mut self) -> Result<Vec<AiStreamAttempt>, GatewayError> {
|
||||
let mut drained = Vec::new();
|
||||
for attempt in self.candidates.drain_static_attempts() {
|
||||
if let Some(attempt) = self.build_stream_attempt(attempt).await? {
|
||||
drained.push(attempt);
|
||||
}
|
||||
}
|
||||
Ok(drained)
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalGeminiFilesSyncAttemptSource<'_> {
|
||||
async fn build_sync_attempt(
|
||||
&self,
|
||||
attempt: LocalGeminiFilesCandidateAttempt,
|
||||
) -> Result<Option<AiSyncAttempt>, GatewayError> {
|
||||
let spec_metadata = local_gemini_files_spec_metadata(self.spec);
|
||||
let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
self.body_base64,
|
||||
self.body_is_empty,
|
||||
self.trace_id,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match build_passthrough_sync_plan_from_decision(self.parts, payload) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %self.trace_id,
|
||||
decision_kind = spec_metadata.decision_kind,
|
||||
error = ?err,
|
||||
"gateway local gemini files sync decision plan build failed"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalGeminiFilesStreamAttemptSource<'_> {
|
||||
async fn build_stream_attempt(
|
||||
&self,
|
||||
attempt: LocalGeminiFilesCandidateAttempt,
|
||||
) -> Result<Option<AiStreamAttempt>, GatewayError> {
|
||||
let spec_metadata = local_gemini_files_spec_metadata(self.spec);
|
||||
let empty_body_json = serde_json::Value::Null;
|
||||
let Some(payload) = maybe_build_local_gemini_files_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
&empty_body_json,
|
||||
None,
|
||||
true,
|
||||
self.trace_id,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match build_passthrough_stream_plan_from_decision(self.parts, payload) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %self.trace_id,
|
||||
decision_kind = spec_metadata.decision_kind,
|
||||
error = ?err,
|
||||
"gateway local gemini files stream decision plan build failed"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn maybe_build_sync_local_gemini_files_decision_payload(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
|
||||
@@ -3,9 +3,11 @@ use serde_json::json;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::ai_serving::planner::candidate_materialization::{
|
||||
build_local_execution_candidate_attempt_source_with_serving,
|
||||
mark_skipped_local_execution_candidate,
|
||||
mark_skipped_local_execution_candidate_with_failure_diagnostic,
|
||||
materialize_local_execution_candidates_with_serving, LocalCandidateResolutionMode,
|
||||
LocalExecutionCandidateAttemptSource,
|
||||
};
|
||||
use crate::ai_serving::planner::candidate_metadata::{
|
||||
build_local_execution_candidate_metadata,
|
||||
@@ -25,6 +27,7 @@ use crate::clock::current_unix_secs;
|
||||
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) const GEMINI_FILES_CANDIDATE_API_FORMAT: &str = "gemini:files";
|
||||
@@ -134,6 +137,74 @@ pub(super) async fn materialize_local_gemini_files_candidate_attempts(
|
||||
Ok(outcome.attempts)
|
||||
}
|
||||
|
||||
pub(super) async fn build_local_gemini_files_candidate_attempt_source<'a>(
|
||||
state: &'a AppState,
|
||||
trace_id: &str,
|
||||
input: &LocalGeminiFilesDecisionInput,
|
||||
) -> Result<(LocalGeminiFilesCandidateAttemptSource<'a>, usize), GatewayError> {
|
||||
let planner_state = PlannerAppState::new(state);
|
||||
let persistence_policy = build_local_candidate_persistence_policy(
|
||||
&input.auth_context,
|
||||
input.required_capabilities.as_ref(),
|
||||
LocalCandidatePersistencePolicyKind::GeminiFilesDecision,
|
||||
);
|
||||
let candidates = planner_state
|
||||
.list_selectable_candidates_for_required_capability_without_requested_model(
|
||||
GEMINI_FILES_CANDIDATE_API_FORMAT,
|
||||
GEMINI_FILES_REQUIRED_CAPABILITY,
|
||||
false,
|
||||
Some(&input.auth_snapshot),
|
||||
current_unix_secs(),
|
||||
)
|
||||
.await?;
|
||||
Ok(build_local_execution_candidate_attempt_source_with_serving(
|
||||
planner_state,
|
||||
trace_id,
|
||||
GEMINI_FILES_CLIENT_API_FORMAT,
|
||||
None,
|
||||
Some(&input.auth_snapshot),
|
||||
input.required_capabilities.as_ref(),
|
||||
None,
|
||||
None,
|
||||
persistence_policy,
|
||||
candidates,
|
||||
Vec::new(),
|
||||
LocalCandidateResolutionMode::WithoutTransportPairGate,
|
||||
|eligible| {
|
||||
let mut extra_fields = serde_json::Map::new();
|
||||
extra_fields.insert(
|
||||
"candidate_api_format".to_string(),
|
||||
json!(GEMINI_FILES_CANDIDATE_API_FORMAT),
|
||||
);
|
||||
Some(build_local_execution_candidate_metadata(
|
||||
LocalExecutionCandidateMetadataParts {
|
||||
eligible,
|
||||
provider_api_format: GEMINI_FILES_CLIENT_API_FORMAT,
|
||||
client_api_format: GEMINI_FILES_CLIENT_API_FORMAT,
|
||||
extra_fields,
|
||||
},
|
||||
))
|
||||
},
|
||||
|mut skipped_candidate| {
|
||||
let mut extra_fields = serde_json::Map::new();
|
||||
extra_fields.insert(
|
||||
"candidate_api_format".to_string(),
|
||||
json!(GEMINI_FILES_CANDIDATE_API_FORMAT),
|
||||
);
|
||||
skipped_candidate.extra_data =
|
||||
Some(build_local_execution_candidate_metadata_for_candidate(
|
||||
&skipped_candidate.candidate,
|
||||
skipped_candidate.transport_ref(),
|
||||
GEMINI_FILES_CLIENT_API_FORMAT,
|
||||
GEMINI_FILES_CLIENT_API_FORMAT,
|
||||
extra_fields,
|
||||
));
|
||||
skipped_candidate
|
||||
},
|
||||
)
|
||||
.await)
|
||||
}
|
||||
|
||||
pub(super) async fn mark_skipped_local_gemini_files_candidate(
|
||||
state: &AppState,
|
||||
input: &LocalGeminiFilesDecisionInput,
|
||||
|
||||
@@ -2,8 +2,10 @@ mod decision;
|
||||
mod request;
|
||||
mod support;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::ai_serving::planner::candidate_materialization::LocalExecutionAttemptSource;
|
||||
use crate::ai_serving::planner::plan_builders::{
|
||||
build_passthrough_sync_plan_from_decision, build_standard_stream_plan_from_decision,
|
||||
AiStreamAttempt, AiSyncAttempt,
|
||||
@@ -18,11 +20,35 @@ use crate::{AiExecutionDecision, AppState, GatewayError};
|
||||
|
||||
use self::decision::maybe_build_local_openai_image_decision_payload_for_candidate;
|
||||
use self::support::{
|
||||
list_local_openai_image_candidate_attempts, resolve_local_openai_image_decision_input,
|
||||
build_local_openai_image_candidate_attempt_source, list_local_openai_image_candidate_attempts,
|
||||
resolve_local_openai_image_decision_input, LocalOpenAiImageCandidateAttempt,
|
||||
LocalOpenAiImageCandidateAttemptSource, LocalOpenAiImageDecisionInput,
|
||||
};
|
||||
|
||||
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_base64: Option<&'a str>,
|
||||
trace_id: &'a str,
|
||||
input: LocalOpenAiImageDecisionInput,
|
||||
spec: LocalOpenAiImageSpec,
|
||||
candidates: LocalOpenAiImageCandidateAttemptSource<'a>,
|
||||
}
|
||||
|
||||
pub(crate) struct LocalOpenAiImageStreamAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_base64: Option<&'a str>,
|
||||
trace_id: &'a str,
|
||||
input: LocalOpenAiImageDecisionInput,
|
||||
spec: LocalOpenAiImageSpec,
|
||||
candidates: LocalOpenAiImageCandidateAttemptSource<'a>,
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_image_sync_plan_and_reports_for_kind(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
@@ -73,6 +99,242 @@ pub(crate) async fn build_local_image_stream_plan_and_reports_for_kind(
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_image_sync_attempt_source_for_kind<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_base64: Option<&'a str>,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
plan_kind: &str,
|
||||
) -> Result<Option<(LocalOpenAiImageSyncAttemptSource<'a>, usize)>, GatewayError> {
|
||||
let Some(spec) = resolve_sync_spec(plan_kind) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let spec_metadata = local_openai_image_spec_metadata(spec);
|
||||
|
||||
let Some(input) = resolve_local_openai_image_decision_input(
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
body_base64,
|
||||
trace_id,
|
||||
decision,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let Some((candidates, candidate_count)) = build_local_openai_image_candidate_attempt_source(
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
body_json,
|
||||
spec_metadata.api_format,
|
||||
spec_metadata.decision_kind,
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
if candidate_count == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
Ok(Some((
|
||||
LocalOpenAiImageSyncAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
body_base64,
|
||||
trace_id,
|
||||
input,
|
||||
spec,
|
||||
candidates,
|
||||
},
|
||||
candidate_count,
|
||||
)))
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_image_stream_attempt_source_for_kind<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
body_json: &'a serde_json::Value,
|
||||
body_base64: Option<&'a str>,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
plan_kind: &str,
|
||||
) -> Result<Option<(LocalOpenAiImageStreamAttemptSource<'a>, usize)>, GatewayError> {
|
||||
let Some(spec) = resolve_stream_spec(plan_kind) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let spec_metadata = local_openai_image_spec_metadata(spec);
|
||||
|
||||
let Some(input) = resolve_local_openai_image_decision_input(
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
body_base64,
|
||||
trace_id,
|
||||
decision,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let Some((candidates, candidate_count)) = build_local_openai_image_candidate_attempt_source(
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
body_json,
|
||||
spec_metadata.api_format,
|
||||
spec_metadata.decision_kind,
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
if candidate_count == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
Ok(Some((
|
||||
LocalOpenAiImageStreamAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
body_base64,
|
||||
trace_id,
|
||||
input,
|
||||
spec,
|
||||
candidates,
|
||||
},
|
||||
candidate_count,
|
||||
)))
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiImageSyncAttemptSource<'_> {
|
||||
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
|
||||
while let Some(attempt) = self.candidates.next_attempt().await {
|
||||
match self.build_sync_attempt(attempt).await? {
|
||||
Some(attempt) => return Ok(Some(attempt)),
|
||||
None => continue,
|
||||
}
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn drain_execution_attempts(&mut self) -> Result<Vec<AiSyncAttempt>, GatewayError> {
|
||||
let mut drained = Vec::new();
|
||||
for attempt in self.candidates.drain_static_attempts() {
|
||||
if let Some(attempt) = self.build_sync_attempt(attempt).await? {
|
||||
drained.push(attempt);
|
||||
}
|
||||
}
|
||||
Ok(drained)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiImageStreamAttemptSource<'_> {
|
||||
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
|
||||
while let Some(attempt) = self.candidates.next_attempt().await {
|
||||
match self.build_stream_attempt(attempt).await? {
|
||||
Some(attempt) => return Ok(Some(attempt)),
|
||||
None => continue,
|
||||
}
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn drain_execution_attempts(&mut self) -> Result<Vec<AiStreamAttempt>, GatewayError> {
|
||||
let mut drained = Vec::new();
|
||||
for attempt in self.candidates.drain_static_attempts() {
|
||||
if let Some(attempt) = self.build_stream_attempt(attempt).await? {
|
||||
drained.push(attempt);
|
||||
}
|
||||
}
|
||||
Ok(drained)
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalOpenAiImageSyncAttemptSource<'_> {
|
||||
async fn build_sync_attempt(
|
||||
&self,
|
||||
attempt: LocalOpenAiImageCandidateAttempt,
|
||||
) -> Result<Option<AiSyncAttempt>, GatewayError> {
|
||||
let spec_metadata = local_openai_image_spec_metadata(self.spec);
|
||||
let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
self.body_base64,
|
||||
self.trace_id,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match build_passthrough_sync_plan_from_decision(self.parts, payload) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %self.trace_id,
|
||||
decision_kind = spec_metadata.decision_kind,
|
||||
error = ?err,
|
||||
"gateway local openai image sync decision plan build failed"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalOpenAiImageStreamAttemptSource<'_> {
|
||||
async fn build_stream_attempt(
|
||||
&self,
|
||||
attempt: LocalOpenAiImageCandidateAttempt,
|
||||
) -> Result<Option<AiStreamAttempt>, GatewayError> {
|
||||
let spec_metadata = local_openai_image_spec_metadata(self.spec);
|
||||
let Some(payload) = maybe_build_local_openai_image_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
self.body_base64,
|
||||
self.trace_id,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match build_standard_stream_plan_from_decision(self.parts, self.body_json, payload, false) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %self.trace_id,
|
||||
decision_kind = spec_metadata.decision_kind,
|
||||
error = ?err,
|
||||
"gateway local openai image stream decision plan build failed"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn maybe_build_sync_local_image_decision_payload(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use tracing::warn;
|
||||
|
||||
use crate::ai_serving::planner::candidate_materialization::{
|
||||
build_local_execution_candidate_attempt_source_with_serving,
|
||||
mark_skipped_local_execution_candidate,
|
||||
mark_skipped_local_execution_candidate_with_failure_diagnostic,
|
||||
materialize_local_execution_candidates_with_serving, LocalCandidateResolutionMode,
|
||||
@@ -23,10 +24,11 @@ use crate::ai_serving::{
|
||||
PlannerAppState,
|
||||
};
|
||||
use crate::clock::current_unix_secs;
|
||||
use crate::AppState;
|
||||
use crate::{AppState, GatewayError};
|
||||
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
|
||||
|
||||
pub(super) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttempt as LocalOpenAiImageCandidateAttempt;
|
||||
pub(super) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttemptSource as LocalOpenAiImageCandidateAttemptSource;
|
||||
pub(super) use crate::ai_serving::planner::decision_input::LocalRequestedModelDecisionInput as LocalOpenAiImageDecisionInput;
|
||||
|
||||
use super::request::resolve_requested_image_model_for_request;
|
||||
@@ -131,6 +133,94 @@ pub(super) async fn list_local_openai_image_candidate_attempts(
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) async fn build_local_openai_image_candidate_attempt_source<'a>(
|
||||
state: &'a AppState,
|
||||
trace_id: &str,
|
||||
input: &LocalOpenAiImageDecisionInput,
|
||||
body_json: &serde_json::Value,
|
||||
api_format: &str,
|
||||
decision_kind: &str,
|
||||
) -> Result<Option<(LocalOpenAiImageCandidateAttemptSource<'a>, usize)>, GatewayError> {
|
||||
let planner_state = PlannerAppState::new(state);
|
||||
let (candidates, preselection_skipped) = match planner_state
|
||||
.list_selectable_candidates_with_skip_reasons(
|
||||
api_format,
|
||||
&input.requested_model,
|
||||
false,
|
||||
input.required_capabilities.as_ref(),
|
||||
Some(&input.auth_snapshot),
|
||||
current_unix_secs(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(candidates) => candidates,
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
decision_kind,
|
||||
error = ?err,
|
||||
"gateway local openai image decision scheduler selection failed"
|
||||
);
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
|
||||
let sticky_session_token = extract_pool_sticky_session_token(body_json);
|
||||
let persistence_policy = build_local_candidate_persistence_policy(
|
||||
&input.auth_context,
|
||||
input.required_capabilities.as_ref(),
|
||||
LocalCandidatePersistencePolicyKind::ImageDecision,
|
||||
);
|
||||
|
||||
let (source, candidate_count) = build_local_execution_candidate_attempt_source_with_serving(
|
||||
planner_state,
|
||||
trace_id,
|
||||
api_format,
|
||||
Some(&input.requested_model),
|
||||
Some(&input.auth_snapshot),
|
||||
input.required_capabilities.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
candidates,
|
||||
preselection_skipped
|
||||
.into_iter()
|
||||
.map(|item| SkippedLocalExecutionCandidate {
|
||||
candidate: item.candidate,
|
||||
skip_reason: item.skip_reason,
|
||||
transport: None,
|
||||
ranking: None,
|
||||
extra_data: None,
|
||||
})
|
||||
.collect(),
|
||||
LocalCandidateResolutionMode::Standard,
|
||||
|eligible| {
|
||||
Some(build_local_execution_candidate_metadata(
|
||||
LocalExecutionCandidateMetadataParts {
|
||||
eligible,
|
||||
provider_api_format: api_format,
|
||||
client_api_format: api_format,
|
||||
extra_fields: serde_json::Map::new(),
|
||||
},
|
||||
))
|
||||
},
|
||||
|mut skipped_candidate| {
|
||||
skipped_candidate.extra_data =
|
||||
Some(build_local_execution_candidate_metadata_for_candidate(
|
||||
&skipped_candidate.candidate,
|
||||
skipped_candidate.transport_ref(),
|
||||
api_format,
|
||||
api_format,
|
||||
serde_json::Map::new(),
|
||||
));
|
||||
skipped_candidate
|
||||
},
|
||||
)
|
||||
.await;
|
||||
|
||||
Ok(Some((source, candidate_count)))
|
||||
}
|
||||
|
||||
async fn materialize_local_openai_image_candidate_attempts(
|
||||
state: PlannerAppState<'_>,
|
||||
trace_id: &str,
|
||||
|
||||
@@ -5,16 +5,22 @@ mod image;
|
||||
mod video;
|
||||
|
||||
pub(crate) use self::files::{
|
||||
build_local_gemini_files_stream_attempt_source_for_kind,
|
||||
build_local_gemini_files_stream_plan_and_reports_for_kind,
|
||||
build_local_gemini_files_sync_attempt_source_for_kind,
|
||||
build_local_gemini_files_sync_plan_and_reports_for_kind,
|
||||
maybe_build_stream_local_gemini_files_decision_payload,
|
||||
maybe_build_sync_local_gemini_files_decision_payload,
|
||||
};
|
||||
pub(crate) use self::image::{
|
||||
build_local_image_stream_attempt_source_for_kind,
|
||||
build_local_image_stream_plan_and_reports_for_kind,
|
||||
build_local_image_sync_attempt_source_for_kind,
|
||||
build_local_image_sync_plan_and_reports_for_kind,
|
||||
maybe_build_stream_local_image_decision_payload, maybe_build_sync_local_image_decision_payload,
|
||||
};
|
||||
pub(crate) use self::video::{
|
||||
build_local_video_sync_plan_and_reports_for_kind, maybe_build_sync_local_video_decision_payload,
|
||||
build_local_video_sync_attempt_source_for_kind,
|
||||
build_local_video_sync_plan_and_reports_for_kind,
|
||||
maybe_build_sync_local_video_decision_payload,
|
||||
};
|
||||
|
||||
@@ -2,8 +2,10 @@ mod decision;
|
||||
mod request;
|
||||
mod support;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::ai_serving::planner::candidate_materialization::LocalExecutionAttemptSource;
|
||||
use crate::ai_serving::planner::plan_builders::{
|
||||
build_passthrough_sync_plan_from_decision, AiSyncAttempt,
|
||||
};
|
||||
@@ -17,9 +19,21 @@ use crate::{AiExecutionDecision, AppState, GatewayError};
|
||||
|
||||
use self::decision::maybe_build_local_video_create_decision_payload_for_candidate;
|
||||
use self::support::{
|
||||
list_local_video_create_candidate_attempts, resolve_local_video_create_decision_input,
|
||||
build_local_video_create_candidate_attempt_source, list_local_video_create_candidate_attempts,
|
||||
resolve_local_video_create_decision_input, LocalVideoCreateCandidateAttempt,
|
||||
LocalVideoCreateCandidateAttemptSource, LocalVideoCreateDecisionInput,
|
||||
};
|
||||
|
||||
pub(crate) struct LocalVideoCreateSyncAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
body_json: &'a serde_json::Value,
|
||||
trace_id: &'a str,
|
||||
input: LocalVideoCreateDecisionInput,
|
||||
spec: LocalVideoCreateSpec,
|
||||
candidates: LocalVideoCreateCandidateAttemptSource<'a>,
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_video_sync_plan_and_reports_for_kind(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
@@ -35,6 +49,116 @@ pub(crate) async fn build_local_video_sync_plan_and_reports_for_kind(
|
||||
build_local_sync_plan_and_reports(state, parts, body_json, trace_id, decision, spec).await
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_video_sync_attempt_source_for_kind<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
body_json: &'a serde_json::Value,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
plan_kind: &str,
|
||||
) -> Result<Option<(LocalVideoCreateSyncAttemptSource<'a>, usize)>, GatewayError> {
|
||||
let Some(spec) = resolve_sync_spec(plan_kind) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let spec_metadata = local_video_create_spec_metadata(spec);
|
||||
|
||||
let Some(input) = resolve_local_video_create_decision_input(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let Some((candidates, candidate_count)) = build_local_video_create_candidate_attempt_source(
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
body_json,
|
||||
spec_metadata.api_format,
|
||||
spec_metadata.decision_kind,
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
if candidate_count == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
Ok(Some((
|
||||
LocalVideoCreateSyncAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
trace_id,
|
||||
input,
|
||||
spec,
|
||||
candidates,
|
||||
},
|
||||
candidate_count,
|
||||
)))
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalVideoCreateSyncAttemptSource<'_> {
|
||||
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
|
||||
while let Some(attempt) = self.candidates.next_attempt().await {
|
||||
match self.build_sync_attempt(attempt).await? {
|
||||
Some(attempt) => return Ok(Some(attempt)),
|
||||
None => continue,
|
||||
}
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn drain_execution_attempts(&mut self) -> Result<Vec<AiSyncAttempt>, GatewayError> {
|
||||
let mut drained = Vec::new();
|
||||
for attempt in self.candidates.drain_static_attempts() {
|
||||
if let Some(attempt) = self.build_sync_attempt(attempt).await? {
|
||||
drained.push(attempt);
|
||||
}
|
||||
}
|
||||
Ok(drained)
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalVideoCreateSyncAttemptSource<'_> {
|
||||
async fn build_sync_attempt(
|
||||
&self,
|
||||
attempt: LocalVideoCreateCandidateAttempt,
|
||||
) -> Result<Option<AiSyncAttempt>, GatewayError> {
|
||||
let spec_metadata = local_video_create_spec_metadata(self.spec);
|
||||
let Some(payload) = maybe_build_local_video_create_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
self.trace_id,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match build_passthrough_sync_plan_from_decision(self.parts, payload) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %self.trace_id,
|
||||
decision_kind = spec_metadata.decision_kind,
|
||||
error = ?err,
|
||||
"gateway local video sync decision plan build failed"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn maybe_build_sync_local_video_decision_payload(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
|
||||
@@ -3,6 +3,7 @@ use tracing::warn;
|
||||
|
||||
use super::{LocalVideoCreateFamily, LocalVideoCreateSpec};
|
||||
use crate::ai_serving::planner::candidate_materialization::{
|
||||
build_local_execution_candidate_attempt_source_with_serving,
|
||||
mark_skipped_local_execution_candidate,
|
||||
mark_skipped_local_execution_candidate_with_failure_diagnostic,
|
||||
materialize_local_execution_candidates_with_serving, LocalCandidateResolutionMode,
|
||||
@@ -26,9 +27,10 @@ use crate::ai_serving::{
|
||||
PlannerAppState,
|
||||
};
|
||||
use crate::clock::current_unix_secs;
|
||||
use crate::AppState;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
pub(super) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttempt as LocalVideoCreateCandidateAttempt;
|
||||
pub(super) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttemptSource as LocalVideoCreateCandidateAttemptSource;
|
||||
pub(super) use crate::ai_serving::planner::decision_input::LocalRequestedModelDecisionInput as LocalVideoCreateDecisionInput;
|
||||
|
||||
pub(super) async fn resolve_local_video_create_decision_input(
|
||||
@@ -143,6 +145,94 @@ pub(super) async fn list_local_video_create_candidate_attempts(
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) async fn build_local_video_create_candidate_attempt_source<'a>(
|
||||
state: &'a AppState,
|
||||
trace_id: &str,
|
||||
input: &LocalVideoCreateDecisionInput,
|
||||
body_json: &serde_json::Value,
|
||||
api_format: &str,
|
||||
decision_kind: &str,
|
||||
) -> Result<Option<(LocalVideoCreateCandidateAttemptSource<'a>, usize)>, GatewayError> {
|
||||
let planner_state = PlannerAppState::new(state);
|
||||
let (candidates, preselection_skipped) = match planner_state
|
||||
.list_selectable_candidates_with_skip_reasons(
|
||||
api_format,
|
||||
&input.requested_model,
|
||||
false,
|
||||
input.required_capabilities.as_ref(),
|
||||
Some(&input.auth_snapshot),
|
||||
current_unix_secs(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(candidates) => candidates,
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
decision_kind = decision_kind,
|
||||
error = ?err,
|
||||
"gateway local video decision scheduler selection failed"
|
||||
);
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
|
||||
let sticky_session_token = extract_pool_sticky_session_token(body_json);
|
||||
let persistence_policy = build_local_candidate_persistence_policy(
|
||||
&input.auth_context,
|
||||
input.required_capabilities.as_ref(),
|
||||
LocalCandidatePersistencePolicyKind::VideoDecision,
|
||||
);
|
||||
|
||||
let (source, candidate_count) = build_local_execution_candidate_attempt_source_with_serving(
|
||||
planner_state,
|
||||
trace_id,
|
||||
api_format,
|
||||
Some(&input.requested_model),
|
||||
Some(&input.auth_snapshot),
|
||||
input.required_capabilities.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
candidates,
|
||||
preselection_skipped
|
||||
.into_iter()
|
||||
.map(|item| SkippedLocalExecutionCandidate {
|
||||
candidate: item.candidate,
|
||||
skip_reason: item.skip_reason,
|
||||
transport: None,
|
||||
ranking: None,
|
||||
extra_data: None,
|
||||
})
|
||||
.collect(),
|
||||
LocalCandidateResolutionMode::Standard,
|
||||
|eligible| {
|
||||
Some(build_local_execution_candidate_metadata(
|
||||
LocalExecutionCandidateMetadataParts {
|
||||
eligible,
|
||||
provider_api_format: api_format,
|
||||
client_api_format: api_format,
|
||||
extra_fields: serde_json::Map::new(),
|
||||
},
|
||||
))
|
||||
},
|
||||
|mut skipped_candidate| {
|
||||
skipped_candidate.extra_data =
|
||||
Some(build_local_execution_candidate_metadata_for_candidate(
|
||||
&skipped_candidate.candidate,
|
||||
skipped_candidate.transport_ref(),
|
||||
api_format,
|
||||
api_format,
|
||||
serde_json::Map::new(),
|
||||
));
|
||||
skipped_candidate
|
||||
},
|
||||
)
|
||||
.await;
|
||||
|
||||
Ok(Some((source, candidate_count)))
|
||||
}
|
||||
|
||||
async fn materialize_local_video_create_candidate_attempts(
|
||||
state: PlannerAppState<'_>,
|
||||
trace_id: &str,
|
||||
|
||||
@@ -1,6 +1,12 @@
|
||||
use async_trait::async_trait;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::ai_serving::planner::common::extract_requested_model_from_request;
|
||||
use crate::ai_serving::planner::candidate_materialization::{
|
||||
LocalExecutionAttemptSource, LocalExecutionCandidateAttemptSource,
|
||||
};
|
||||
use crate::ai_serving::planner::common::{
|
||||
extract_requested_model_from_request, RequestedModelFamily,
|
||||
};
|
||||
use crate::ai_serving::planner::plan_builders::{AiStreamAttempt, AiSyncAttempt};
|
||||
use crate::ai_serving::planner::runtime_miss::{
|
||||
apply_local_runtime_candidate_evaluation_progress,
|
||||
@@ -14,10 +20,279 @@ use crate::ai_serving::GatewayControlDecision;
|
||||
use crate::{AiExecutionDecision, AppState, GatewayError};
|
||||
|
||||
use super::candidates::{
|
||||
materialize_local_standard_candidate_attempts, resolve_local_standard_decision_input,
|
||||
build_local_standard_candidate_attempt_source, materialize_local_standard_candidate_attempts,
|
||||
resolve_local_standard_decision_input,
|
||||
};
|
||||
use super::payload::maybe_build_local_standard_decision_payload_for_candidate;
|
||||
use super::LocalStandardSpec;
|
||||
use super::{LocalStandardDecisionInput, LocalStandardSpec};
|
||||
|
||||
pub(crate) struct LocalStandardSyncAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
input: LocalStandardDecisionInput,
|
||||
spec: LocalStandardSpec,
|
||||
requested_model_family: RequestedModelFamily,
|
||||
candidates: LocalExecutionCandidateAttemptSource<'a>,
|
||||
}
|
||||
|
||||
pub(crate) struct LocalStandardStreamAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
input: LocalStandardDecisionInput,
|
||||
spec: LocalStandardSpec,
|
||||
requested_model_family: RequestedModelFamily,
|
||||
candidates: LocalExecutionCandidateAttemptSource<'a>,
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_sync_attempt_source<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
body_json: &'a serde_json::Value,
|
||||
spec: LocalStandardSpec,
|
||||
) -> Result<Option<(LocalStandardSyncAttemptSource<'a>, usize)>, GatewayError> {
|
||||
let spec_metadata = local_standard_spec_metadata(spec);
|
||||
let requested_model_family = spec_metadata
|
||||
.requested_model_family
|
||||
.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
|
||||
else {
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
spec_metadata.decision_kind,
|
||||
extract_requested_model_from_request(parts, body_json, requested_model_family)
|
||||
.as_deref(),
|
||||
"decision_input_unavailable",
|
||||
);
|
||||
return Ok(None);
|
||||
};
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
spec_metadata.decision_kind,
|
||||
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?;
|
||||
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
|
||||
if candidate_count == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
Ok(Some((
|
||||
LocalStandardSyncAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
input,
|
||||
spec,
|
||||
requested_model_family,
|
||||
candidates,
|
||||
},
|
||||
candidate_count,
|
||||
)))
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_stream_attempt_source<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
body_json: &'a serde_json::Value,
|
||||
spec: LocalStandardSpec,
|
||||
) -> Result<Option<(LocalStandardStreamAttemptSource<'a>, usize)>, GatewayError> {
|
||||
let spec_metadata = local_standard_spec_metadata(spec);
|
||||
let requested_model_family = spec_metadata
|
||||
.requested_model_family
|
||||
.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
|
||||
else {
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
spec_metadata.decision_kind,
|
||||
extract_requested_model_from_request(parts, body_json, requested_model_family)
|
||||
.as_deref(),
|
||||
"decision_input_unavailable",
|
||||
);
|
||||
return Ok(None);
|
||||
};
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
spec_metadata.decision_kind,
|
||||
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?;
|
||||
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
|
||||
if candidate_count == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
Ok(Some((
|
||||
LocalStandardStreamAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
input,
|
||||
spec,
|
||||
requested_model_family,
|
||||
candidates,
|
||||
},
|
||||
candidate_count,
|
||||
)))
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalStandardSyncAttemptSource<'_> {
|
||||
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
|
||||
while let Some(attempt) = self.candidates.next_attempt().await {
|
||||
match self.build_sync_attempt(attempt).await? {
|
||||
Some(attempt) => return Ok(Some(attempt)),
|
||||
None => continue,
|
||||
}
|
||||
}
|
||||
apply_local_runtime_candidate_terminal_reason(
|
||||
self.state,
|
||||
self.trace_id,
|
||||
"no_local_sync_plans",
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn drain_execution_attempts(&mut self) -> Result<Vec<AiSyncAttempt>, GatewayError> {
|
||||
let mut drained = Vec::new();
|
||||
for attempt in self.candidates.drain_static_attempts() {
|
||||
if let Some(attempt) = self.build_sync_attempt(attempt).await? {
|
||||
drained.push(attempt);
|
||||
}
|
||||
}
|
||||
Ok(drained)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalStandardStreamAttemptSource<'_> {
|
||||
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
|
||||
while let Some(attempt) = self.candidates.next_attempt().await {
|
||||
match self.build_stream_attempt(attempt).await? {
|
||||
Some(attempt) => return Ok(Some(attempt)),
|
||||
None => continue,
|
||||
}
|
||||
}
|
||||
apply_local_runtime_candidate_terminal_reason(
|
||||
self.state,
|
||||
self.trace_id,
|
||||
"no_local_stream_plans",
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn drain_execution_attempts(&mut self) -> Result<Vec<AiStreamAttempt>, GatewayError> {
|
||||
let mut drained = Vec::new();
|
||||
for attempt in self.candidates.drain_static_attempts() {
|
||||
if let Some(attempt) = self.build_stream_attempt(attempt).await? {
|
||||
drained.push(attempt);
|
||||
}
|
||||
}
|
||||
Ok(drained)
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalStandardSyncAttemptSource<'_> {
|
||||
async fn build_sync_attempt(
|
||||
&self,
|
||||
attempt: super::LocalStandardCandidateAttempt,
|
||||
) -> Result<Option<AiSyncAttempt>, GatewayError> {
|
||||
let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
match build_sync_plan_from_requested_model_family(
|
||||
self.requested_model_family,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
payload,
|
||||
) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %self.trace_id,
|
||||
error = ?err,
|
||||
"gateway local standard sync plan build failed"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalStandardStreamAttemptSource<'_> {
|
||||
async fn build_stream_attempt(
|
||||
&self,
|
||||
attempt: super::LocalStandardCandidateAttempt,
|
||||
) -> Result<Option<AiStreamAttempt>, GatewayError> {
|
||||
let Some(payload) = maybe_build_local_standard_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
match build_stream_plan_from_requested_model_family(
|
||||
self.requested_model_family,
|
||||
self.parts,
|
||||
self.body_json,
|
||||
payload,
|
||||
) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %self.trace_id,
|
||||
error = ?err,
|
||||
"gateway local standard stream plan build failed"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn maybe_build_sync_via_standard_family_payload(
|
||||
state: &AppState,
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
use tracing::warn;
|
||||
|
||||
use crate::ai_serving::planner::candidate_materialization::{
|
||||
build_local_execution_candidate_attempt_source_with_serving,
|
||||
materialize_local_execution_candidates_with_serving, LocalCandidateResolutionMode,
|
||||
LocalExecutionCandidateAttemptSource,
|
||||
};
|
||||
use crate::ai_serving::planner::candidate_metadata::{
|
||||
build_local_execution_candidate_contract_metadata,
|
||||
@@ -166,3 +168,95 @@ pub(super) async fn materialize_local_standard_candidate_attempts(
|
||||
|
||||
Ok((outcome.attempts, outcome.candidate_count))
|
||||
}
|
||||
|
||||
pub(super) async fn build_local_standard_candidate_attempt_source<'a>(
|
||||
state: &'a AppState,
|
||||
trace_id: &str,
|
||||
input: &LocalStandardDecisionInput,
|
||||
body_json: &serde_json::Value,
|
||||
spec: LocalStandardSpec,
|
||||
) -> Result<(LocalExecutionCandidateAttemptSource<'a>, usize), GatewayError> {
|
||||
let spec_metadata = local_standard_spec_metadata(spec);
|
||||
let planner_state = PlannerAppState::new(state);
|
||||
let sticky_session_token = extract_pool_sticky_session_token(body_json);
|
||||
let persistence_policy = build_local_candidate_persistence_policy(
|
||||
&input.auth_context,
|
||||
input.required_capabilities.as_ref(),
|
||||
LocalCandidatePersistencePolicyKind::StandardDecision,
|
||||
);
|
||||
let preselection = preselect_local_execution_candidates_with_serving(
|
||||
planner_state,
|
||||
spec_metadata.api_format,
|
||||
&input.requested_model,
|
||||
spec_metadata.require_streaming,
|
||||
input.required_capabilities.as_ref(),
|
||||
&input.auth_snapshot,
|
||||
false,
|
||||
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
|
||||
)
|
||||
.await?;
|
||||
|
||||
Ok(build_local_execution_candidate_attempt_source_with_serving(
|
||||
planner_state,
|
||||
trace_id,
|
||||
spec_metadata.api_format,
|
||||
Some(&input.requested_model),
|
||||
Some(&input.auth_snapshot),
|
||||
input.required_capabilities.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
preselection.candidates,
|
||||
preselection.skipped_candidates,
|
||||
LocalCandidateResolutionMode::Standard,
|
||||
|eligible| {
|
||||
let provider_api_format = eligible.provider_api_format.clone();
|
||||
let (execution_strategy, conversion_mode) = ai_local_execution_contract_for_formats(
|
||||
spec_metadata.api_format,
|
||||
&provider_api_format,
|
||||
);
|
||||
Some(build_local_execution_candidate_contract_metadata(
|
||||
LocalExecutionCandidateMetadataParts {
|
||||
eligible,
|
||||
provider_api_format: provider_api_format.as_str(),
|
||||
client_api_format: spec_metadata.api_format,
|
||||
extra_fields: serde_json::Map::new(),
|
||||
},
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
eligible.candidate.endpoint_api_format.as_str(),
|
||||
))
|
||||
},
|
||||
|mut skipped_candidate| {
|
||||
let provider_api_format = skipped_candidate
|
||||
.transport
|
||||
.as_ref()
|
||||
.map(|transport| transport.endpoint.api_format.trim().to_ascii_lowercase())
|
||||
.unwrap_or_else(|| {
|
||||
skipped_candidate
|
||||
.candidate
|
||||
.endpoint_api_format
|
||||
.trim()
|
||||
.to_ascii_lowercase()
|
||||
});
|
||||
let (execution_strategy, conversion_mode) = ai_local_execution_contract_for_formats(
|
||||
spec_metadata.api_format,
|
||||
&provider_api_format,
|
||||
);
|
||||
skipped_candidate.extra_data = Some(
|
||||
build_local_execution_candidate_contract_metadata_for_candidate(
|
||||
&skipped_candidate.candidate,
|
||||
skipped_candidate.transport_ref(),
|
||||
provider_api_format.as_str(),
|
||||
spec_metadata.api_format,
|
||||
serde_json::Map::new(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
provider_api_format.as_str(),
|
||||
),
|
||||
);
|
||||
skipped_candidate
|
||||
},
|
||||
)
|
||||
.await)
|
||||
}
|
||||
|
||||
@@ -4,7 +4,8 @@ mod payload;
|
||||
mod request;
|
||||
|
||||
pub(crate) use self::build::{
|
||||
build_local_stream_plan_and_reports, build_local_sync_plan_and_reports,
|
||||
build_local_stream_attempt_source, build_local_stream_plan_and_reports,
|
||||
build_local_sync_attempt_source, build_local_sync_plan_and_reports,
|
||||
maybe_build_stream_via_standard_family_payload, maybe_build_sync_via_standard_family_payload,
|
||||
};
|
||||
pub(super) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttempt as LocalStandardCandidateAttempt;
|
||||
|
||||
@@ -288,7 +288,9 @@ mod tests {
|
||||
|
||||
use super::maybe_build_local_standard_decision_payload_for_candidate;
|
||||
use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttempt;
|
||||
use crate::ai_serving::planner::candidate_resolution::EligibleLocalExecutionCandidate;
|
||||
use crate::ai_serving::planner::candidate_resolution::{
|
||||
EligibleLocalExecutionCandidate, LocalExecutionCandidateKind,
|
||||
};
|
||||
use crate::ai_serving::planner::decision_input::LocalRequestedModelDecisionInput;
|
||||
use crate::ai_serving::{
|
||||
ExecutionRuntimeAuthContext, GatewayAuthApiKeySnapshot, LocalStandardSourceFamily,
|
||||
@@ -452,6 +454,7 @@ mod tests {
|
||||
) -> LocalExecutionCandidateAttempt {
|
||||
LocalExecutionCandidateAttempt {
|
||||
eligible: EligibleLocalExecutionCandidate {
|
||||
kind: LocalExecutionCandidateKind::SingleKey,
|
||||
candidate: sample_candidate(api_format, endpoint_id),
|
||||
transport: Arc::new(sample_transport(api_format, endpoint_id)),
|
||||
provider_api_format: api_format.to_string(),
|
||||
|
||||
@@ -15,7 +15,8 @@ mod openai;
|
||||
|
||||
pub(crate) use self::codex::apply_codex_openai_responses_special_headers;
|
||||
pub(crate) use self::family::{
|
||||
build_local_stream_plan_and_reports, build_local_sync_plan_and_reports,
|
||||
build_local_stream_attempt_source, build_local_stream_plan_and_reports,
|
||||
build_local_sync_attempt_source, build_local_sync_plan_and_reports,
|
||||
};
|
||||
pub(crate) use self::normalize::{
|
||||
build_cross_format_openai_chat_request_body, build_cross_format_openai_chat_upstream_url,
|
||||
@@ -25,9 +26,13 @@ pub(crate) use self::normalize::{
|
||||
build_local_openai_responses_upstream_url,
|
||||
};
|
||||
pub(crate) use self::openai::{
|
||||
build_local_openai_chat_stream_attempt_source_for_kind,
|
||||
build_local_openai_chat_stream_plan_and_reports_for_kind,
|
||||
build_local_openai_chat_sync_attempt_source_for_kind,
|
||||
build_local_openai_chat_sync_plan_and_reports_for_kind,
|
||||
build_local_openai_responses_stream_attempt_source_for_kind,
|
||||
build_local_openai_responses_stream_plan_and_reports_for_kind,
|
||||
build_local_openai_responses_sync_attempt_source_for_kind,
|
||||
build_local_openai_responses_sync_plan_and_reports_for_kind, copy_request_number_field,
|
||||
copy_request_number_field_as, map_openai_reasoning_effort_to_claude_output,
|
||||
map_openai_reasoning_effort_to_gemini_budget, maybe_build_stream_local_decision_payload,
|
||||
|
||||
@@ -7,6 +7,7 @@ mod support;
|
||||
|
||||
pub(super) use self::payload::maybe_build_local_openai_chat_decision_payload_for_candidate;
|
||||
pub(super) use self::support::{
|
||||
build_local_openai_chat_candidate_attempt_source,
|
||||
materialize_local_openai_chat_candidate_attempts, LocalOpenAiChatCandidateAttempt,
|
||||
LocalOpenAiChatDecisionInput,
|
||||
LocalOpenAiChatCandidateAttemptSource, LocalOpenAiChatDecisionInput,
|
||||
};
|
||||
|
||||
@@ -1,9 +1,11 @@
|
||||
use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
|
||||
|
||||
use crate::ai_serving::planner::candidate_materialization::{
|
||||
build_local_execution_candidate_attempt_source_with_serving,
|
||||
mark_skipped_local_execution_candidate, mark_skipped_local_execution_candidate_with_extra_data,
|
||||
mark_skipped_local_execution_candidate_with_failure_diagnostic,
|
||||
materialize_local_execution_candidates_with_serving, LocalCandidateResolutionMode,
|
||||
LocalExecutionCandidateAttemptSource,
|
||||
};
|
||||
use crate::ai_serving::planner::candidate_metadata::{
|
||||
build_local_execution_candidate_contract_metadata,
|
||||
@@ -22,6 +24,7 @@ use crate::ai_serving::{
|
||||
use crate::AppState;
|
||||
|
||||
pub(crate) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttempt as LocalOpenAiChatCandidateAttempt;
|
||||
pub(crate) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttemptSource as LocalOpenAiChatCandidateAttemptSource;
|
||||
pub(crate) use crate::ai_serving::planner::decision_input::LocalRequestedModelDecisionInput as LocalOpenAiChatDecisionInput;
|
||||
|
||||
pub(crate) async fn mark_skipped_local_openai_chat_candidate(
|
||||
@@ -189,3 +192,80 @@ pub(crate) async fn materialize_local_openai_chat_candidate_attempts(
|
||||
|
||||
outcome.attempts
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_openai_chat_candidate_attempt_source<'a>(
|
||||
state: &'a AppState,
|
||||
trace_id: &str,
|
||||
input: &LocalOpenAiChatDecisionInput,
|
||||
body_json: &serde_json::Value,
|
||||
candidates: Vec<SchedulerMinimalCandidateSelectionCandidate>,
|
||||
preselection_skipped: Vec<SkippedLocalExecutionCandidate>,
|
||||
) -> (LocalOpenAiChatCandidateAttemptSource<'a>, usize) {
|
||||
let planner_state = PlannerAppState::new(state);
|
||||
let sticky_session_token = extract_pool_sticky_session_token(body_json);
|
||||
let auth_context: &ExecutionRuntimeAuthContext = &input.auth_context;
|
||||
let persistence_policy = build_local_candidate_persistence_policy(
|
||||
auth_context,
|
||||
input.required_capabilities.as_ref(),
|
||||
LocalCandidatePersistencePolicyKind::OpenAiChatDecision,
|
||||
);
|
||||
build_local_execution_candidate_attempt_source_with_serving(
|
||||
planner_state,
|
||||
trace_id,
|
||||
"openai:chat",
|
||||
Some(&input.requested_model),
|
||||
Some(&input.auth_snapshot),
|
||||
input.required_capabilities.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
candidates,
|
||||
preselection_skipped,
|
||||
LocalCandidateResolutionMode::Standard,
|
||||
|eligible| {
|
||||
let provider_api_format = eligible.provider_api_format.clone();
|
||||
let (execution_strategy, conversion_mode) =
|
||||
ai_local_execution_contract_for_formats("openai:chat", &provider_api_format);
|
||||
Some(build_local_execution_candidate_contract_metadata(
|
||||
LocalExecutionCandidateMetadataParts {
|
||||
eligible,
|
||||
provider_api_format: provider_api_format.as_str(),
|
||||
client_api_format: "openai:chat",
|
||||
extra_fields: serde_json::Map::new(),
|
||||
},
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
eligible.candidate.endpoint_api_format.trim(),
|
||||
))
|
||||
},
|
||||
|mut skipped_candidate| {
|
||||
let provider_api_format = skipped_candidate
|
||||
.transport
|
||||
.as_ref()
|
||||
.map(|transport| transport.endpoint.api_format.trim().to_ascii_lowercase())
|
||||
.unwrap_or_else(|| {
|
||||
skipped_candidate
|
||||
.candidate
|
||||
.endpoint_api_format
|
||||
.trim()
|
||||
.to_ascii_lowercase()
|
||||
});
|
||||
let (execution_strategy, conversion_mode) =
|
||||
ai_local_execution_contract_for_formats("openai:chat", &provider_api_format);
|
||||
skipped_candidate.extra_data = Some(
|
||||
build_local_execution_candidate_contract_metadata_for_candidate(
|
||||
&skipped_candidate.candidate,
|
||||
skipped_candidate.transport_ref(),
|
||||
provider_api_format.as_str(),
|
||||
"openai:chat",
|
||||
serde_json::Map::new(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
provider_api_format.as_str(),
|
||||
),
|
||||
);
|
||||
skipped_candidate
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
@@ -12,11 +12,14 @@ mod decision;
|
||||
mod plans;
|
||||
|
||||
use self::decision::{
|
||||
build_local_openai_chat_candidate_attempt_source,
|
||||
materialize_local_openai_chat_candidate_attempts,
|
||||
maybe_build_local_openai_chat_decision_payload_for_candidate, LocalOpenAiChatDecisionInput,
|
||||
maybe_build_local_openai_chat_decision_payload_for_candidate, LocalOpenAiChatCandidateAttempt,
|
||||
LocalOpenAiChatCandidateAttemptSource, LocalOpenAiChatDecisionInput,
|
||||
};
|
||||
use self::plans::{
|
||||
build_local_openai_chat_stream_plan_and_reports, build_local_openai_chat_sync_plan_and_reports,
|
||||
build_local_openai_chat_stream_attempt_source, build_local_openai_chat_stream_plan_and_reports,
|
||||
build_local_openai_chat_sync_attempt_source, build_local_openai_chat_sync_plan_and_reports,
|
||||
list_local_openai_chat_candidates, resolve_local_openai_chat_decision_input,
|
||||
set_local_openai_chat_miss_diagnostic,
|
||||
};
|
||||
@@ -49,6 +52,50 @@ pub(crate) async fn build_local_openai_chat_stream_plan_and_reports_for_kind(
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_openai_chat_sync_attempt_source_for_kind<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
body_json: &'a serde_json::Value,
|
||||
plan_kind: &str,
|
||||
) -> Result<
|
||||
Option<(
|
||||
impl crate::ai_serving::planner::LocalExecutionAttemptSource<
|
||||
crate::ai_serving::planner::plan_builders::AiSyncAttempt,
|
||||
> + 'a,
|
||||
usize,
|
||||
)>,
|
||||
GatewayError,
|
||||
> {
|
||||
build_local_openai_chat_sync_attempt_source(
|
||||
state, parts, trace_id, decision, body_json, plan_kind,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_openai_chat_stream_attempt_source_for_kind<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
body_json: &'a serde_json::Value,
|
||||
plan_kind: &str,
|
||||
) -> Result<
|
||||
Option<(
|
||||
impl crate::ai_serving::planner::LocalExecutionAttemptSource<
|
||||
crate::ai_serving::planner::plan_builders::AiStreamAttempt,
|
||||
> + 'a,
|
||||
usize,
|
||||
)>,
|
||||
GatewayError,
|
||||
> {
|
||||
build_local_openai_chat_stream_attempt_source(
|
||||
state, parts, trace_id, decision, body_json, plan_kind,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) fn set_local_openai_chat_execution_exhausted_diagnostic(
|
||||
state: &AppState,
|
||||
trace_id: &str,
|
||||
|
||||
@@ -12,5 +12,9 @@ mod sync;
|
||||
pub(super) use self::candidates::list_local_openai_chat_candidates;
|
||||
pub(super) use self::diagnostic::set_local_openai_chat_miss_diagnostic;
|
||||
pub(super) use self::resolve::resolve_local_openai_chat_decision_input;
|
||||
pub(super) use self::stream::build_local_openai_chat_stream_plan_and_reports;
|
||||
pub(super) use self::sync::build_local_openai_chat_sync_plan_and_reports;
|
||||
pub(super) use self::stream::{
|
||||
build_local_openai_chat_stream_attempt_source, build_local_openai_chat_stream_plan_and_reports,
|
||||
};
|
||||
pub(super) use self::sync::{
|
||||
build_local_openai_chat_sync_attempt_source, build_local_openai_chat_sync_plan_and_reports,
|
||||
};
|
||||
|
||||
@@ -1,21 +1,181 @@
|
||||
use async_trait::async_trait;
|
||||
use tracing::warn;
|
||||
|
||||
use super::super::{
|
||||
build_local_openai_chat_candidate_attempt_source,
|
||||
materialize_local_openai_chat_candidate_attempts,
|
||||
maybe_build_local_openai_chat_decision_payload_for_candidate, AppState, GatewayControlDecision,
|
||||
GatewayError,
|
||||
GatewayError, LocalOpenAiChatCandidateAttempt, LocalOpenAiChatCandidateAttemptSource,
|
||||
LocalOpenAiChatDecisionInput,
|
||||
};
|
||||
use super::candidates::list_local_openai_chat_candidates;
|
||||
use super::diagnostic::{
|
||||
set_local_openai_chat_candidate_evaluation_diagnostic, set_local_openai_chat_miss_diagnostic,
|
||||
};
|
||||
use super::resolve::resolve_local_openai_chat_decision_input;
|
||||
use crate::ai_serving::planner::candidate_materialization::LocalExecutionAttemptSource;
|
||||
use crate::ai_serving::planner::common::OPENAI_CHAT_STREAM_PLAN_KIND;
|
||||
use crate::ai_serving::planner::plan_builders::{
|
||||
build_openai_chat_stream_plan_from_decision, AiStreamAttempt,
|
||||
};
|
||||
use crate::ai_serving::planner::runtime_miss::apply_local_runtime_candidate_terminal_reason;
|
||||
|
||||
pub(crate) struct LocalOpenAiChatStreamAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
input: LocalOpenAiChatDecisionInput,
|
||||
candidates: LocalOpenAiChatCandidateAttemptSource<'a>,
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_openai_chat_stream_attempt_source<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
body_json: &'a serde_json::Value,
|
||||
plan_kind: &str,
|
||||
) -> Result<Option<(LocalOpenAiChatStreamAttemptSource<'a>, usize)>, GatewayError> {
|
||||
if plan_kind != OPENAI_CHAT_STREAM_PLAN_KIND {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let Some(input) = resolve_local_openai_chat_decision_input(
|
||||
state, trace_id, decision, body_json, plan_kind, true,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let (candidates, skipped_candidates) =
|
||||
match list_local_openai_chat_candidates(state, &input, true).await {
|
||||
Ok(value) => value,
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
error = ?err,
|
||||
"gateway local openai chat stream decision scheduler selection failed"
|
||||
);
|
||||
set_local_openai_chat_miss_diagnostic(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
Some(input.requested_model.as_str()),
|
||||
"scheduler_selection_failed",
|
||||
);
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
let candidate_count = candidates.len() + skipped_candidates.len();
|
||||
if candidate_count == 0 {
|
||||
set_local_openai_chat_candidate_evaluation_diagnostic(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
Some(input.requested_model.as_str()),
|
||||
0,
|
||||
);
|
||||
return Ok(None);
|
||||
}
|
||||
set_local_openai_chat_candidate_evaluation_diagnostic(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
Some(input.requested_model.as_str()),
|
||||
candidate_count,
|
||||
);
|
||||
|
||||
let (candidates, candidate_count) = build_local_openai_chat_candidate_attempt_source(
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
body_json,
|
||||
candidates,
|
||||
skipped_candidates,
|
||||
)
|
||||
.await;
|
||||
|
||||
Ok(Some((
|
||||
LocalOpenAiChatStreamAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
input,
|
||||
candidates,
|
||||
},
|
||||
candidate_count,
|
||||
)))
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiChatStreamAttemptSource<'_> {
|
||||
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
|
||||
while let Some(attempt) = self.candidates.next_attempt().await {
|
||||
match self.build_stream_attempt(attempt).await? {
|
||||
Some(attempt) => return Ok(Some(attempt)),
|
||||
None => continue,
|
||||
}
|
||||
}
|
||||
apply_local_runtime_candidate_terminal_reason(
|
||||
self.state,
|
||||
self.trace_id,
|
||||
"no_local_stream_plans",
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn drain_execution_attempts(&mut self) -> Result<Vec<AiStreamAttempt>, GatewayError> {
|
||||
let mut drained = Vec::new();
|
||||
for attempt in self.candidates.drain_static_attempts() {
|
||||
if let Some(attempt) = self.build_stream_attempt(attempt).await? {
|
||||
drained.push(attempt);
|
||||
}
|
||||
}
|
||||
Ok(drained)
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalOpenAiChatStreamAttemptSource<'_> {
|
||||
async fn build_stream_attempt(
|
||||
&self,
|
||||
attempt: LocalOpenAiChatCandidateAttempt,
|
||||
) -> Result<Option<AiStreamAttempt>, GatewayError> {
|
||||
let Some(payload) = maybe_build_local_openai_chat_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
OPENAI_CHAT_STREAM_PLAN_KIND,
|
||||
"openai_chat_stream_success",
|
||||
true,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match build_openai_chat_stream_plan_from_decision(self.parts, self.body_json, payload) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %self.trace_id,
|
||||
error = ?err,
|
||||
"gateway local openai chat stream decision plan build failed"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_openai_chat_stream_plan_and_reports(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
|
||||
@@ -1,15 +1,19 @@
|
||||
use async_trait::async_trait;
|
||||
use tracing::warn;
|
||||
|
||||
use super::super::{
|
||||
build_local_openai_chat_candidate_attempt_source,
|
||||
materialize_local_openai_chat_candidate_attempts,
|
||||
maybe_build_local_openai_chat_decision_payload_for_candidate, AppState, GatewayControlDecision,
|
||||
GatewayError,
|
||||
GatewayError, LocalOpenAiChatCandidateAttempt, LocalOpenAiChatCandidateAttemptSource,
|
||||
LocalOpenAiChatDecisionInput,
|
||||
};
|
||||
use super::candidates::list_local_openai_chat_candidates;
|
||||
use super::diagnostic::{
|
||||
set_local_openai_chat_candidate_evaluation_diagnostic, set_local_openai_chat_miss_diagnostic,
|
||||
};
|
||||
use super::resolve::resolve_local_openai_chat_decision_input;
|
||||
use crate::ai_serving::planner::candidate_materialization::LocalExecutionAttemptSource;
|
||||
use crate::ai_serving::planner::common::{
|
||||
force_upstream_streaming_for_provider, OPENAI_CHAT_SYNC_PLAN_KIND,
|
||||
};
|
||||
@@ -18,6 +22,15 @@ use crate::ai_serving::planner::plan_builders::{
|
||||
};
|
||||
use crate::ai_serving::planner::runtime_miss::apply_local_runtime_candidate_terminal_reason;
|
||||
|
||||
pub(crate) struct LocalOpenAiChatSyncAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
input: LocalOpenAiChatDecisionInput,
|
||||
candidates: LocalOpenAiChatCandidateAttemptSource<'a>,
|
||||
}
|
||||
|
||||
fn openai_chat_sync_upstream_is_stream_for_candidate(
|
||||
provider_type: &str,
|
||||
provider_api_format: &str,
|
||||
@@ -25,6 +38,157 @@ fn openai_chat_sync_upstream_is_stream_for_candidate(
|
||||
force_upstream_streaming_for_provider(provider_type, provider_api_format)
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_openai_chat_sync_attempt_source<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
body_json: &'a serde_json::Value,
|
||||
plan_kind: &str,
|
||||
) -> Result<Option<(LocalOpenAiChatSyncAttemptSource<'a>, usize)>, GatewayError> {
|
||||
if plan_kind != OPENAI_CHAT_SYNC_PLAN_KIND {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let Some(input) = resolve_local_openai_chat_decision_input(
|
||||
state, trace_id, decision, body_json, plan_kind, true,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let (candidates, skipped_candidates) =
|
||||
match list_local_openai_chat_candidates(state, &input, false).await {
|
||||
Ok(value) => value,
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %trace_id,
|
||||
error = ?err,
|
||||
"gateway local openai chat sync decision scheduler selection failed"
|
||||
);
|
||||
set_local_openai_chat_miss_diagnostic(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
Some(input.requested_model.as_str()),
|
||||
"scheduler_selection_failed",
|
||||
);
|
||||
return Ok(None);
|
||||
}
|
||||
};
|
||||
let candidate_count = candidates.len() + skipped_candidates.len();
|
||||
if candidate_count == 0 {
|
||||
set_local_openai_chat_candidate_evaluation_diagnostic(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
Some(input.requested_model.as_str()),
|
||||
0,
|
||||
);
|
||||
return Ok(None);
|
||||
}
|
||||
set_local_openai_chat_candidate_evaluation_diagnostic(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
Some(input.requested_model.as_str()),
|
||||
candidate_count,
|
||||
);
|
||||
|
||||
let (candidates, candidate_count) = build_local_openai_chat_candidate_attempt_source(
|
||||
state,
|
||||
trace_id,
|
||||
&input,
|
||||
body_json,
|
||||
candidates,
|
||||
skipped_candidates,
|
||||
)
|
||||
.await;
|
||||
|
||||
Ok(Some((
|
||||
LocalOpenAiChatSyncAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
input,
|
||||
candidates,
|
||||
},
|
||||
candidate_count,
|
||||
)))
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiChatSyncAttemptSource<'_> {
|
||||
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
|
||||
while let Some(attempt) = self.candidates.next_attempt().await {
|
||||
match self.build_sync_attempt(attempt).await? {
|
||||
Some(attempt) => return Ok(Some(attempt)),
|
||||
None => continue,
|
||||
}
|
||||
}
|
||||
apply_local_runtime_candidate_terminal_reason(
|
||||
self.state,
|
||||
self.trace_id,
|
||||
"no_local_sync_plans",
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn drain_execution_attempts(&mut self) -> Result<Vec<AiSyncAttempt>, GatewayError> {
|
||||
let mut drained = Vec::new();
|
||||
for attempt in self.candidates.drain_static_attempts() {
|
||||
if let Some(attempt) = self.build_sync_attempt(attempt).await? {
|
||||
drained.push(attempt);
|
||||
}
|
||||
}
|
||||
Ok(drained)
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalOpenAiChatSyncAttemptSource<'_> {
|
||||
async fn build_sync_attempt(
|
||||
&self,
|
||||
attempt: LocalOpenAiChatCandidateAttempt,
|
||||
) -> Result<Option<AiSyncAttempt>, GatewayError> {
|
||||
let upstream_is_stream = openai_chat_sync_upstream_is_stream_for_candidate(
|
||||
attempt.eligible.transport.provider.provider_type.as_str(),
|
||||
attempt.eligible.provider_api_format.as_str(),
|
||||
);
|
||||
let Some(payload) = maybe_build_local_openai_chat_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
OPENAI_CHAT_SYNC_PLAN_KIND,
|
||||
"openai_chat_sync_success",
|
||||
upstream_is_stream,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match build_openai_chat_sync_plan_from_decision(self.parts, self.body_json, payload) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %self.trace_id,
|
||||
error = ?err,
|
||||
"gateway local openai chat sync decision plan build failed"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_openai_chat_sync_plan_and_reports(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
|
||||
@@ -7,13 +7,17 @@ pub(crate) use crate::ai_serving::{
|
||||
parse_openai_stop_sequences, resolve_openai_chat_max_tokens, value_as_u64,
|
||||
};
|
||||
pub(crate) use chat::{
|
||||
build_local_openai_chat_stream_attempt_source_for_kind,
|
||||
build_local_openai_chat_stream_plan_and_reports_for_kind,
|
||||
build_local_openai_chat_sync_attempt_source_for_kind,
|
||||
build_local_openai_chat_sync_plan_and_reports_for_kind,
|
||||
maybe_build_stream_local_decision_payload, maybe_build_sync_local_decision_payload,
|
||||
set_local_openai_chat_execution_exhausted_diagnostic,
|
||||
};
|
||||
pub(crate) use responses::{
|
||||
build_local_openai_responses_stream_attempt_source_for_kind,
|
||||
build_local_openai_responses_stream_plan_and_reports_for_kind,
|
||||
build_local_openai_responses_sync_attempt_source_for_kind,
|
||||
build_local_openai_responses_sync_plan_and_reports_for_kind,
|
||||
maybe_build_stream_local_openai_responses_decision_payload,
|
||||
maybe_build_sync_local_openai_responses_decision_payload,
|
||||
|
||||
@@ -7,8 +7,9 @@ mod support;
|
||||
|
||||
pub(super) use self::payload::maybe_build_local_openai_responses_decision_payload_for_candidate;
|
||||
pub(super) use self::support::{
|
||||
build_local_openai_responses_candidate_attempt_source,
|
||||
materialize_local_openai_responses_candidate_attempts,
|
||||
resolve_local_openai_responses_decision_input, LocalOpenAiResponsesCandidateAttempt,
|
||||
LocalOpenAiResponsesDecisionInput,
|
||||
LocalOpenAiResponsesCandidateAttemptSource, LocalOpenAiResponsesDecisionInput,
|
||||
};
|
||||
pub(super) use crate::ai_serving::LocalOpenAiResponsesSpec;
|
||||
|
||||
@@ -2,9 +2,11 @@ use aether_scheduler_core::SchedulerMinimalCandidateSelectionCandidate;
|
||||
use tracing::warn;
|
||||
|
||||
use crate::ai_serving::planner::candidate_materialization::{
|
||||
build_local_execution_candidate_attempt_source_with_serving,
|
||||
mark_skipped_local_execution_candidate, mark_skipped_local_execution_candidate_with_extra_data,
|
||||
mark_skipped_local_execution_candidate_with_failure_diagnostic,
|
||||
materialize_local_execution_candidates_with_serving, LocalCandidateResolutionMode,
|
||||
LocalExecutionCandidateAttemptSource,
|
||||
};
|
||||
use crate::ai_serving::planner::candidate_metadata::{
|
||||
build_local_execution_candidate_contract_metadata,
|
||||
@@ -34,6 +36,7 @@ use crate::{AppState, GatewayError};
|
||||
use super::LocalOpenAiResponsesSpec;
|
||||
|
||||
pub(crate) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttempt as LocalOpenAiResponsesCandidateAttempt;
|
||||
pub(crate) use crate::ai_serving::planner::candidate_materialization::LocalExecutionCandidateAttemptSource as LocalOpenAiResponsesCandidateAttemptSource;
|
||||
pub(crate) use crate::ai_serving::planner::decision_input::LocalRequestedModelDecisionInput as LocalOpenAiResponsesDecisionInput;
|
||||
|
||||
pub(crate) async fn resolve_local_openai_responses_decision_input(
|
||||
@@ -221,6 +224,98 @@ pub(crate) async fn materialize_local_openai_responses_candidate_attempts(
|
||||
Ok((outcome.attempts, outcome.candidate_count))
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_openai_responses_candidate_attempt_source<'a>(
|
||||
state: &'a AppState,
|
||||
trace_id: &str,
|
||||
input: &LocalOpenAiResponsesDecisionInput,
|
||||
body_json: &serde_json::Value,
|
||||
spec: LocalOpenAiResponsesSpec,
|
||||
) -> Result<(LocalOpenAiResponsesCandidateAttemptSource<'a>, usize), GatewayError> {
|
||||
let spec_metadata = local_openai_responses_spec_metadata(spec);
|
||||
let planner_state = PlannerAppState::new(state);
|
||||
let sticky_session_token = extract_pool_sticky_session_token(body_json);
|
||||
let auth_context: &ExecutionRuntimeAuthContext = &input.auth_context;
|
||||
let persistence_policy = build_local_candidate_persistence_policy(
|
||||
auth_context,
|
||||
input.required_capabilities.as_ref(),
|
||||
LocalCandidatePersistencePolicyKind::OpenAiResponsesDecision,
|
||||
);
|
||||
let preselection = preselect_local_execution_candidates_with_serving(
|
||||
planner_state,
|
||||
spec_metadata.api_format,
|
||||
&input.requested_model,
|
||||
spec_metadata.require_streaming,
|
||||
input.required_capabilities.as_ref(),
|
||||
&input.auth_snapshot,
|
||||
true,
|
||||
LocalCandidatePreselectionKeyMode::ProviderEndpointKeyModelAndApiFormat,
|
||||
)
|
||||
.await?;
|
||||
Ok(build_local_execution_candidate_attempt_source_with_serving(
|
||||
planner_state,
|
||||
trace_id,
|
||||
spec_metadata.api_format,
|
||||
Some(&input.requested_model),
|
||||
Some(&input.auth_snapshot),
|
||||
input.required_capabilities.as_ref(),
|
||||
sticky_session_token.as_deref(),
|
||||
input.request_auth_channel.as_deref(),
|
||||
persistence_policy,
|
||||
preselection.candidates,
|
||||
preselection.skipped_candidates,
|
||||
LocalCandidateResolutionMode::Standard,
|
||||
|eligible| {
|
||||
let provider_api_format = eligible.provider_api_format.clone();
|
||||
let (execution_strategy, conversion_mode) = ai_local_execution_contract_for_formats(
|
||||
spec_metadata.api_format,
|
||||
&provider_api_format,
|
||||
);
|
||||
Some(build_local_execution_candidate_contract_metadata(
|
||||
LocalExecutionCandidateMetadataParts {
|
||||
eligible,
|
||||
provider_api_format: provider_api_format.as_str(),
|
||||
client_api_format: spec_metadata.api_format,
|
||||
extra_fields: serde_json::Map::new(),
|
||||
},
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
eligible.candidate.endpoint_api_format.as_str(),
|
||||
))
|
||||
},
|
||||
|mut skipped_candidate| {
|
||||
let provider_api_format = skipped_candidate
|
||||
.transport
|
||||
.as_ref()
|
||||
.map(|transport| transport.endpoint.api_format.trim().to_ascii_lowercase())
|
||||
.unwrap_or_else(|| {
|
||||
skipped_candidate
|
||||
.candidate
|
||||
.endpoint_api_format
|
||||
.trim()
|
||||
.to_ascii_lowercase()
|
||||
});
|
||||
let (execution_strategy, conversion_mode) = ai_local_execution_contract_for_formats(
|
||||
spec_metadata.api_format,
|
||||
&provider_api_format,
|
||||
);
|
||||
skipped_candidate.extra_data = Some(
|
||||
build_local_execution_candidate_contract_metadata_for_candidate(
|
||||
&skipped_candidate.candidate,
|
||||
skipped_candidate.transport_ref(),
|
||||
provider_api_format.as_str(),
|
||||
spec_metadata.api_format,
|
||||
serde_json::Map::new(),
|
||||
execution_strategy,
|
||||
conversion_mode,
|
||||
provider_api_format.as_str(),
|
||||
),
|
||||
);
|
||||
skipped_candidate
|
||||
},
|
||||
)
|
||||
.await)
|
||||
}
|
||||
|
||||
pub(crate) async fn mark_skipped_local_openai_responses_candidate(
|
||||
state: &AppState,
|
||||
input: &LocalOpenAiResponsesDecisionInput,
|
||||
|
||||
@@ -11,7 +11,8 @@ use self::decision::{
|
||||
resolve_local_openai_responses_decision_input,
|
||||
};
|
||||
use self::plans::{
|
||||
build_local_stream_plan_and_reports, build_local_sync_plan_and_reports, resolve_stream_spec,
|
||||
build_local_stream_attempt_source, build_local_stream_plan_and_reports,
|
||||
build_local_sync_attempt_source, build_local_sync_plan_and_reports, resolve_stream_spec,
|
||||
resolve_sync_spec,
|
||||
};
|
||||
|
||||
@@ -45,6 +46,48 @@ pub(crate) async fn build_local_openai_responses_stream_plan_and_reports_for_kin
|
||||
build_local_stream_plan_and_reports(state, parts, trace_id, decision, body_json, spec).await
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_openai_responses_sync_attempt_source_for_kind<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
body_json: &'a serde_json::Value,
|
||||
plan_kind: &str,
|
||||
) -> Result<
|
||||
Option<(
|
||||
impl crate::ai_serving::planner::LocalExecutionAttemptSource<AiSyncAttempt> + 'a,
|
||||
usize,
|
||||
)>,
|
||||
GatewayError,
|
||||
> {
|
||||
let Some(spec) = resolve_sync_spec(plan_kind) else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
build_local_sync_attempt_source(state, parts, trace_id, decision, body_json, spec).await
|
||||
}
|
||||
|
||||
pub(crate) async fn build_local_openai_responses_stream_attempt_source_for_kind<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
body_json: &'a serde_json::Value,
|
||||
plan_kind: &str,
|
||||
) -> Result<
|
||||
Option<(
|
||||
impl crate::ai_serving::planner::LocalExecutionAttemptSource<AiStreamAttempt> + 'a,
|
||||
usize,
|
||||
)>,
|
||||
GatewayError,
|
||||
> {
|
||||
let Some(spec) = resolve_stream_spec(plan_kind) else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
build_local_stream_attempt_source(state, parts, trace_id, decision, body_json, spec).await
|
||||
}
|
||||
|
||||
pub(crate) async fn maybe_build_sync_local_openai_responses_decision_payload(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
|
||||
@@ -1,10 +1,15 @@
|
||||
use async_trait::async_trait;
|
||||
use tracing::warn;
|
||||
|
||||
use super::decision::{
|
||||
build_local_openai_responses_candidate_attempt_source,
|
||||
materialize_local_openai_responses_candidate_attempts,
|
||||
maybe_build_local_openai_responses_decision_payload_for_candidate,
|
||||
resolve_local_openai_responses_decision_input, LocalOpenAiResponsesSpec,
|
||||
resolve_local_openai_responses_decision_input, LocalOpenAiResponsesCandidateAttempt,
|
||||
LocalOpenAiResponsesCandidateAttemptSource, LocalOpenAiResponsesDecisionInput,
|
||||
LocalOpenAiResponsesSpec,
|
||||
};
|
||||
use crate::ai_serving::planner::candidate_materialization::LocalExecutionAttemptSource;
|
||||
use crate::ai_serving::planner::plan_builders::{
|
||||
build_openai_responses_stream_plan_from_decision,
|
||||
build_openai_responses_sync_plan_from_decision, AiStreamAttempt, AiSyncAttempt,
|
||||
@@ -21,6 +26,260 @@ pub(crate) use crate::ai_serving::{
|
||||
};
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
pub(crate) struct LocalOpenAiResponsesSyncAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
input: LocalOpenAiResponsesDecisionInput,
|
||||
spec: LocalOpenAiResponsesSpec,
|
||||
candidates: LocalOpenAiResponsesCandidateAttemptSource<'a>,
|
||||
}
|
||||
|
||||
pub(crate) struct LocalOpenAiResponsesStreamAttemptSource<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
body_json: &'a serde_json::Value,
|
||||
input: LocalOpenAiResponsesDecisionInput,
|
||||
spec: LocalOpenAiResponsesSpec,
|
||||
candidates: LocalOpenAiResponsesCandidateAttemptSource<'a>,
|
||||
}
|
||||
|
||||
pub(super) async fn build_local_sync_attempt_source<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
body_json: &'a serde_json::Value,
|
||||
spec: LocalOpenAiResponsesSpec,
|
||||
) -> Result<Option<(LocalOpenAiResponsesSyncAttemptSource<'a>, usize)>, GatewayError> {
|
||||
let spec_metadata = local_openai_responses_spec_metadata(spec);
|
||||
let Some(input) = resolve_local_openai_responses_decision_input(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
body_json,
|
||||
spec_metadata.decision_kind,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
spec_metadata.decision_kind,
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let (candidates, candidate_count) = build_local_openai_responses_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, spec,
|
||||
)
|
||||
.await?;
|
||||
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
|
||||
if candidate_count == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
Ok(Some((
|
||||
LocalOpenAiResponsesSyncAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
input,
|
||||
spec,
|
||||
candidates,
|
||||
},
|
||||
candidate_count,
|
||||
)))
|
||||
}
|
||||
|
||||
pub(super) async fn build_local_stream_attempt_source<'a>(
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
trace_id: &'a str,
|
||||
decision: &'a GatewayControlDecision,
|
||||
body_json: &'a serde_json::Value,
|
||||
spec: LocalOpenAiResponsesSpec,
|
||||
) -> Result<Option<(LocalOpenAiResponsesStreamAttemptSource<'a>, usize)>, GatewayError> {
|
||||
let spec_metadata = local_openai_responses_spec_metadata(spec);
|
||||
let Some(input) = resolve_local_openai_responses_decision_input(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
body_json,
|
||||
spec_metadata.decision_kind,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
set_local_runtime_miss_diagnostic_reason(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
spec_metadata.decision_kind,
|
||||
Some(input.requested_model.as_str()),
|
||||
"candidate_evaluation_incomplete",
|
||||
);
|
||||
let (candidates, candidate_count) = build_local_openai_responses_candidate_attempt_source(
|
||||
state, trace_id, &input, body_json, spec,
|
||||
)
|
||||
.await?;
|
||||
apply_local_runtime_candidate_evaluation_progress(state, trace_id, candidate_count);
|
||||
if candidate_count == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
Ok(Some((
|
||||
LocalOpenAiResponsesStreamAttemptSource {
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
body_json,
|
||||
input,
|
||||
spec,
|
||||
candidates,
|
||||
},
|
||||
candidate_count,
|
||||
)))
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalExecutionAttemptSource<AiSyncAttempt> for LocalOpenAiResponsesSyncAttemptSource<'_> {
|
||||
async fn next_execution_attempt(&mut self) -> Result<Option<AiSyncAttempt>, GatewayError> {
|
||||
while let Some(attempt) = self.candidates.next_attempt().await {
|
||||
match self.build_sync_attempt(attempt).await? {
|
||||
Some(attempt) => return Ok(Some(attempt)),
|
||||
None => continue,
|
||||
}
|
||||
}
|
||||
apply_local_runtime_candidate_terminal_reason(
|
||||
self.state,
|
||||
self.trace_id,
|
||||
"no_local_sync_plans",
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn drain_execution_attempts(&mut self) -> Result<Vec<AiSyncAttempt>, GatewayError> {
|
||||
let mut drained = Vec::new();
|
||||
for attempt in self.candidates.drain_static_attempts() {
|
||||
if let Some(attempt) = self.build_sync_attempt(attempt).await? {
|
||||
drained.push(attempt);
|
||||
}
|
||||
}
|
||||
Ok(drained)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalExecutionAttemptSource<AiStreamAttempt> for LocalOpenAiResponsesStreamAttemptSource<'_> {
|
||||
async fn next_execution_attempt(&mut self) -> Result<Option<AiStreamAttempt>, GatewayError> {
|
||||
while let Some(attempt) = self.candidates.next_attempt().await {
|
||||
match self.build_stream_attempt(attempt).await? {
|
||||
Some(attempt) => return Ok(Some(attempt)),
|
||||
None => continue,
|
||||
}
|
||||
}
|
||||
apply_local_runtime_candidate_terminal_reason(
|
||||
self.state,
|
||||
self.trace_id,
|
||||
"no_local_stream_plans",
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
async fn drain_execution_attempts(&mut self) -> Result<Vec<AiStreamAttempt>, GatewayError> {
|
||||
let mut drained = Vec::new();
|
||||
for attempt in self.candidates.drain_static_attempts() {
|
||||
if let Some(attempt) = self.build_stream_attempt(attempt).await? {
|
||||
drained.push(attempt);
|
||||
}
|
||||
}
|
||||
Ok(drained)
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalOpenAiResponsesSyncAttemptSource<'_> {
|
||||
async fn build_sync_attempt(
|
||||
&self,
|
||||
attempt: LocalOpenAiResponsesCandidateAttempt,
|
||||
) -> Result<Option<AiSyncAttempt>, GatewayError> {
|
||||
let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match build_openai_responses_sync_plan_from_decision(
|
||||
self.parts,
|
||||
self.body_json,
|
||||
payload,
|
||||
self.spec.compact,
|
||||
) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %self.trace_id,
|
||||
error = ?err,
|
||||
"gateway local openai responses sync decision plan build failed"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl LocalOpenAiResponsesStreamAttemptSource<'_> {
|
||||
async fn build_stream_attempt(
|
||||
&self,
|
||||
attempt: LocalOpenAiResponsesCandidateAttempt,
|
||||
) -> Result<Option<AiStreamAttempt>, GatewayError> {
|
||||
let Some(payload) = maybe_build_local_openai_responses_decision_payload_for_candidate(
|
||||
self.state,
|
||||
self.parts,
|
||||
self.trace_id,
|
||||
self.body_json,
|
||||
&self.input,
|
||||
attempt,
|
||||
self.spec,
|
||||
)
|
||||
.await
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
match build_openai_responses_stream_plan_from_decision(
|
||||
self.parts,
|
||||
self.body_json,
|
||||
payload,
|
||||
self.spec.compact,
|
||||
) {
|
||||
Ok(value) => Ok(value),
|
||||
Err(err) => {
|
||||
warn!(
|
||||
trace_id = %self.trace_id,
|
||||
error = ?err,
|
||||
"gateway local openai responses stream decision plan build failed"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) async fn build_local_sync_plan_and_reports(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
|
||||
@@ -1,5 +1,8 @@
|
||||
use aether_data::DataLayerError;
|
||||
use aether_data_contracts::repository::candidate_selection::StoredMinimalCandidateSelectionRow;
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery,
|
||||
StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
use aether_scheduler_core::{
|
||||
auth_constraints_allow_api_format, collect_global_model_names_for_required_capability,
|
||||
enumerate_minimal_candidate_selection_with_model_directives, normalize_api_format,
|
||||
@@ -20,32 +23,61 @@ pub(crate) trait MinimalCandidateSelectionRowSource {
|
||||
global_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>;
|
||||
|
||||
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>;
|
||||
|
||||
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model_page(
|
||||
&self,
|
||||
query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>;
|
||||
|
||||
async fn read_minimal_candidate_selection_rows_for_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>;
|
||||
|
||||
async fn read_pool_key_candidate_rows_for_group(
|
||||
&self,
|
||||
query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError>;
|
||||
}
|
||||
|
||||
const REQUESTED_MODEL_CANDIDATE_PAGE_SIZE: u32 = 256;
|
||||
const REQUESTED_MODEL_MAX_SCANNED_ROWS: u32 = 2048;
|
||||
|
||||
pub(crate) async fn read_requested_model_rows(
|
||||
state: &(impl MinimalCandidateSelectionRowSource + Sync),
|
||||
api_format: &str,
|
||||
requested_model_name: &str,
|
||||
enable_model_directives: bool,
|
||||
) -> Result<Option<(String, Vec<StoredMinimalCandidateSelectionRow>)>, DataLayerError> {
|
||||
let rows = state
|
||||
.read_minimal_candidate_selection_rows_for_api_format(api_format)
|
||||
.await?;
|
||||
let rows = rows
|
||||
.into_iter()
|
||||
.filter(|row| {
|
||||
row_supports_requested_model_with_model_directives(
|
||||
row,
|
||||
requested_model_name,
|
||||
api_format,
|
||||
enable_model_directives,
|
||||
)
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
let fast_rows = read_requested_model_rows_fast_path(
|
||||
state,
|
||||
api_format,
|
||||
requested_model_name,
|
||||
enable_model_directives,
|
||||
)
|
||||
.await?;
|
||||
let mut rows = filter_rows_for_requested_model(
|
||||
fast_rows,
|
||||
requested_model_name,
|
||||
api_format,
|
||||
enable_model_directives,
|
||||
);
|
||||
if rows.is_empty() {
|
||||
let fallback_rows = state
|
||||
.read_minimal_candidate_selection_rows_for_api_format(api_format)
|
||||
.await?;
|
||||
rows = filter_rows_for_requested_model(
|
||||
fallback_rows,
|
||||
requested_model_name,
|
||||
api_format,
|
||||
enable_model_directives,
|
||||
);
|
||||
}
|
||||
if rows.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
@@ -69,6 +101,85 @@ pub(crate) async fn read_requested_model_rows(
|
||||
Ok(Some((resolved_global_model_name, resolved_rows)))
|
||||
}
|
||||
|
||||
fn filter_rows_for_requested_model(
|
||||
rows: Vec<StoredMinimalCandidateSelectionRow>,
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
enable_model_directives: bool,
|
||||
) -> Vec<StoredMinimalCandidateSelectionRow> {
|
||||
rows.into_iter()
|
||||
.filter(|row| {
|
||||
row_supports_requested_model_with_model_directives(
|
||||
row,
|
||||
requested_model_name,
|
||||
api_format,
|
||||
enable_model_directives,
|
||||
)
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
async fn read_requested_model_rows_fast_path(
|
||||
state: &(impl MinimalCandidateSelectionRowSource + Sync),
|
||||
api_format: &str,
|
||||
requested_model_name: &str,
|
||||
enable_model_directives: bool,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let mut requested_names = vec![requested_model_name.trim().to_string()];
|
||||
if enable_model_directives {
|
||||
if let Some(base_model) =
|
||||
crate::ai_serving::model_directive_base_model(requested_model_name)
|
||||
{
|
||||
if !requested_names.iter().any(|value| value == &base_model) {
|
||||
requested_names.push(base_model);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut rows = Vec::new();
|
||||
let mut seen = BTreeSet::new();
|
||||
for requested_name in requested_names {
|
||||
if requested_name.is_empty() {
|
||||
continue;
|
||||
}
|
||||
let mut offset = 0;
|
||||
let mut scanned = 0;
|
||||
while scanned < REQUESTED_MODEL_MAX_SCANNED_ROWS {
|
||||
let limit =
|
||||
REQUESTED_MODEL_CANDIDATE_PAGE_SIZE.min(REQUESTED_MODEL_MAX_SCANNED_ROWS - scanned);
|
||||
let page = state
|
||||
.read_minimal_candidate_selection_rows_for_api_format_and_requested_model_page(
|
||||
&StoredRequestedModelCandidateRowsQuery {
|
||||
api_format: api_format.to_string(),
|
||||
requested_model_name: requested_name.clone(),
|
||||
offset,
|
||||
limit,
|
||||
},
|
||||
)
|
||||
.await?;
|
||||
if page.is_empty() {
|
||||
break;
|
||||
}
|
||||
let page_len = page.len() as u32;
|
||||
for row in page {
|
||||
if seen.insert((
|
||||
row.endpoint_id.clone(),
|
||||
row.key_id.clone(),
|
||||
row.model_id.clone(),
|
||||
)) {
|
||||
rows.push(row);
|
||||
}
|
||||
}
|
||||
scanned = scanned.saturating_add(page_len);
|
||||
if page_len < limit {
|
||||
break;
|
||||
}
|
||||
offset = offset.saturating_add(limit);
|
||||
}
|
||||
}
|
||||
Ok(rows)
|
||||
}
|
||||
|
||||
pub(crate) async fn enumerate_minimal_candidate_selection_with_required_capabilities(
|
||||
state: &(impl MinimalCandidateSelectionRowSource + Sync),
|
||||
api_format: &str,
|
||||
@@ -210,3 +321,180 @@ fn auth_snapshot_constraints(snapshot: &GatewayAuthApiKeySnapshot) -> SchedulerA
|
||||
.map(|items| items.to_vec()),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
read_requested_model_rows, MinimalCandidateSelectionRowSource,
|
||||
StoredMinimalCandidateSelectionRow,
|
||||
};
|
||||
use aether_data::DataLayerError;
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
use async_trait::async_trait;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
struct CountingSelectionSource {
|
||||
fast_rows: Vec<StoredMinimalCandidateSelectionRow>,
|
||||
fallback_rows: Vec<StoredMinimalCandidateSelectionRow>,
|
||||
fast_calls: AtomicUsize,
|
||||
fallback_calls: AtomicUsize,
|
||||
}
|
||||
|
||||
impl CountingSelectionSource {
|
||||
fn new(
|
||||
fast_rows: Vec<StoredMinimalCandidateSelectionRow>,
|
||||
fallback_rows: Vec<StoredMinimalCandidateSelectionRow>,
|
||||
) -> Self {
|
||||
Self {
|
||||
fast_rows,
|
||||
fallback_rows,
|
||||
fast_calls: AtomicUsize::new(0),
|
||||
fallback_calls: AtomicUsize::new(0),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl MinimalCandidateSelectionRowSource for CountingSelectionSource {
|
||||
async fn read_minimal_candidate_selection_rows_for_api_format_and_global_model(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
_global_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
_requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
|
||||
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model_page(
|
||||
&self,
|
||||
query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.fast_calls.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(self
|
||||
.fast_rows
|
||||
.iter()
|
||||
.skip(query.offset as usize)
|
||||
.take(query.limit as usize)
|
||||
.cloned()
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn read_minimal_candidate_selection_rows_for_api_format(
|
||||
&self,
|
||||
_api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.fallback_calls.fetch_add(1, Ordering::SeqCst);
|
||||
Ok(self.fallback_rows.clone())
|
||||
}
|
||||
|
||||
async fn read_pool_key_candidate_rows_for_group(
|
||||
&self,
|
||||
_query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Ok(Vec::new())
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_row(global_model_name: &str) -> StoredMinimalCandidateSelectionRow {
|
||||
StoredMinimalCandidateSelectionRow {
|
||||
provider_id: "provider-1".to_string(),
|
||||
provider_name: "provider".to_string(),
|
||||
provider_type: "custom".to_string(),
|
||||
provider_priority: 10,
|
||||
provider_is_active: true,
|
||||
endpoint_id: "endpoint-1".to_string(),
|
||||
endpoint_api_format: "openai:chat".to_string(),
|
||||
endpoint_api_family: Some("openai".to_string()),
|
||||
endpoint_kind: Some("chat".to_string()),
|
||||
endpoint_is_active: true,
|
||||
key_id: "key-1".to_string(),
|
||||
key_name: "key".to_string(),
|
||||
key_auth_type: "api_key".to_string(),
|
||||
key_is_active: true,
|
||||
key_api_formats: Some(vec!["openai:chat".to_string()]),
|
||||
key_allowed_models: None,
|
||||
key_capabilities: None,
|
||||
key_internal_priority: 10,
|
||||
key_global_priority_by_format: None,
|
||||
model_id: "model-1".to_string(),
|
||||
global_model_id: "global-model-1".to_string(),
|
||||
global_model_name: global_model_name.to_string(),
|
||||
global_model_mappings: None,
|
||||
global_model_supports_streaming: Some(true),
|
||||
model_provider_model_name: global_model_name.to_string(),
|
||||
model_provider_model_mappings: None,
|
||||
model_supports_streaming: Some(true),
|
||||
model_is_active: true,
|
||||
model_is_available: true,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn requested_model_rows_use_fast_path_without_full_format_scan() {
|
||||
let source = CountingSelectionSource::new(vec![sample_row("gpt-5")], Vec::new());
|
||||
|
||||
let result = read_requested_model_rows(&source, "openai:chat", "gpt-5", false)
|
||||
.await
|
||||
.expect("read should succeed")
|
||||
.expect("rows should resolve");
|
||||
|
||||
assert_eq!(result.0, "gpt-5");
|
||||
assert_eq!(result.1.len(), 1);
|
||||
assert_eq!(source.fast_calls.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(source.fallback_calls.load(Ordering::SeqCst), 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn requested_model_rows_fall_back_to_full_format_scan_when_fast_path_misses() {
|
||||
let source = CountingSelectionSource::new(Vec::new(), vec![sample_row("gpt-5")]);
|
||||
|
||||
let result = read_requested_model_rows(&source, "openai:chat", "gpt-5", false)
|
||||
.await
|
||||
.expect("read should succeed")
|
||||
.expect("rows should resolve");
|
||||
|
||||
assert_eq!(result.0, "gpt-5");
|
||||
assert_eq!(result.1.len(), 1);
|
||||
assert_eq!(source.fast_calls.load(Ordering::SeqCst), 1);
|
||||
assert_eq!(source.fallback_calls.load(Ordering::SeqCst), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn requested_model_rows_fast_path_stops_at_scan_limit() {
|
||||
let mut rows = Vec::new();
|
||||
for index in 0..(super::REQUESTED_MODEL_MAX_SCANNED_ROWS + 5) {
|
||||
let mut row = sample_row("gpt-5");
|
||||
row.provider_id = format!("provider-{index}");
|
||||
row.endpoint_id = format!("endpoint-{index}");
|
||||
row.key_id = format!("key-{index}");
|
||||
row.model_id = format!("model-{index}");
|
||||
rows.push(row);
|
||||
}
|
||||
let source = CountingSelectionSource::new(rows, Vec::new());
|
||||
|
||||
let result = read_requested_model_rows(&source, "openai:chat", "gpt-5", false)
|
||||
.await
|
||||
.expect("read should succeed")
|
||||
.expect("rows should resolve");
|
||||
|
||||
assert_eq!(
|
||||
result.1.len(),
|
||||
super::REQUESTED_MODEL_MAX_SCANNED_ROWS as usize
|
||||
);
|
||||
assert_eq!(
|
||||
source.fast_calls.load(Ordering::SeqCst),
|
||||
(super::REQUESTED_MODEL_MAX_SCANNED_ROWS / super::REQUESTED_MODEL_CANDIDATE_PAGE_SIZE)
|
||||
as usize
|
||||
);
|
||||
assert_eq!(source.fallback_calls.load(Ordering::SeqCst), 0);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -146,6 +146,26 @@ impl MinimalCandidateSelectionRowSource for GatewayDataState {
|
||||
.await
|
||||
}
|
||||
|
||||
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.list_minimal_candidate_selection_rows_for_requested_model(
|
||||
api_format,
|
||||
requested_model_name,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn read_minimal_candidate_selection_rows_for_api_format_and_requested_model_page(
|
||||
&self,
|
||||
query: &aether_data_contracts::repository::candidate_selection::StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.list_minimal_candidate_selection_rows_for_requested_model_page(query)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn read_minimal_candidate_selection_rows_for_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
@@ -153,6 +173,13 @@ impl MinimalCandidateSelectionRowSource for GatewayDataState {
|
||||
self.list_minimal_candidate_selection_rows_for_api_format(api_format)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn read_pool_key_candidate_rows_for_group(
|
||||
&self,
|
||||
query: &aether_data_contracts::repository::candidate_selection::StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.list_pool_key_candidate_rows_for_group(query).await
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
|
||||
@@ -77,6 +77,7 @@ use aether_data_contracts::repository::billing::{
|
||||
};
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
|
||||
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
use aether_data_contracts::repository::candidates::{
|
||||
PublicHealthStatusCount, PublicHealthTimelineBucket, RequestCandidateReadRepository,
|
||||
|
||||
@@ -2,9 +2,10 @@ use super::{
|
||||
AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord,
|
||||
DataLayerError, GatewayDataState, PublicCatalogModelListQuery, PublicCatalogModelSearchQuery,
|
||||
PublicGlobalModelQuery, StoredAdminGlobalModel, StoredAdminGlobalModelPage,
|
||||
StoredAdminProviderModel, StoredMinimalCandidateSelectionRow, StoredProviderActiveGlobalModel,
|
||||
StoredProviderModelStats, StoredPublicCatalogModel, StoredPublicGlobalModel,
|
||||
StoredPublicGlobalModelPage, UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
|
||||
StoredAdminProviderModel, StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery,
|
||||
StoredProviderActiveGlobalModel, StoredProviderModelStats, StoredPublicCatalogModel,
|
||||
StoredPublicGlobalModel, StoredPublicGlobalModelPage, StoredRequestedModelCandidateRowsQuery,
|
||||
UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
|
||||
};
|
||||
|
||||
impl GatewayDataState {
|
||||
@@ -23,6 +24,35 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn list_minimal_candidate_selection_rows_for_requested_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
match &self.minimal_candidate_selection_reader {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.list_for_exact_api_format_and_requested_model(api_format, requested_model_name)
|
||||
.await
|
||||
}
|
||||
None => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn list_minimal_candidate_selection_rows_for_requested_model_page(
|
||||
&self,
|
||||
query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
match &self.minimal_candidate_selection_reader {
|
||||
Some(repository) => {
|
||||
repository
|
||||
.list_for_exact_api_format_and_requested_model_page(query)
|
||||
.await
|
||||
}
|
||||
None => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn list_minimal_candidate_selection_rows_for_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
@@ -33,6 +63,16 @@ impl GatewayDataState {
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn list_pool_key_candidate_rows_for_group(
|
||||
&self,
|
||||
query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
match &self.minimal_candidate_selection_reader {
|
||||
Some(repository) => repository.list_pool_key_rows_for_group(query).await,
|
||||
None => Ok(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn list_public_global_models(
|
||||
&self,
|
||||
query: &PublicGlobalModelQuery,
|
||||
|
||||
@@ -11,6 +11,7 @@ use axum::http::Response;
|
||||
use tokio::time::{timeout, Duration};
|
||||
use tracing::{debug, warn, Instrument};
|
||||
|
||||
use crate::ai_serving::LocalExecutionAttemptSource;
|
||||
use crate::clock::current_unix_ms;
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::execution_runtime::{execute_execution_runtime_stream, execute_execution_runtime_sync};
|
||||
@@ -80,6 +81,42 @@ where
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn execute_sync_attempt_source<T, S>(
|
||||
state: &AppState,
|
||||
parts: &http::request::Parts,
|
||||
trace_id: &str,
|
||||
decision: &GatewayControlDecision,
|
||||
plan_kind: &str,
|
||||
mut source: S,
|
||||
) -> Result<LocalExecutionRequestOutcome, GatewayError>
|
||||
where
|
||||
T: AiExecutionAttempt + Send + Sync + 'static,
|
||||
S: LocalExecutionAttemptSource<T>,
|
||||
{
|
||||
let span = tracing::debug_span!("candidates", trace_id = %trace_id, plan_kind);
|
||||
|
||||
async move {
|
||||
tracing::debug!(
|
||||
event_name = "candidate_loop_started",
|
||||
log_type = "event",
|
||||
trace_id = %trace_id,
|
||||
plan_kind,
|
||||
"dynamic candidate loop started"
|
||||
);
|
||||
|
||||
let port = SyncAttemptLoopPort {
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
};
|
||||
run_dynamic_attempt_loop(&port, &mut source).await
|
||||
}
|
||||
.instrument(span)
|
||||
.await
|
||||
}
|
||||
|
||||
struct SyncAttemptLoopPort<'a> {
|
||||
state: &'a AppState,
|
||||
parts: &'a http::request::Parts,
|
||||
@@ -182,6 +219,75 @@ where
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn execute_stream_attempt_source<T, S>(
|
||||
state: &AppState,
|
||||
trace_id: &str,
|
||||
decision: &GatewayControlDecision,
|
||||
plan_kind: &str,
|
||||
mut source: S,
|
||||
) -> Result<LocalExecutionRequestOutcome, GatewayError>
|
||||
where
|
||||
T: AiExecutionAttempt + Send + Sync + 'static,
|
||||
S: LocalExecutionAttemptSource<T>,
|
||||
{
|
||||
let span = tracing::debug_span!("candidates", trace_id = %trace_id, plan_kind);
|
||||
|
||||
async move {
|
||||
tracing::debug!(
|
||||
event_name = "candidate_loop_started",
|
||||
log_type = "event",
|
||||
trace_id = %trace_id,
|
||||
plan_kind,
|
||||
"dynamic candidate loop started"
|
||||
);
|
||||
|
||||
let port = StreamAttemptLoopPort {
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
};
|
||||
run_dynamic_attempt_loop(&port, &mut source).await
|
||||
}
|
||||
.instrument(span)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn run_dynamic_attempt_loop<Port, Source, Attempt>(
|
||||
port: &Port,
|
||||
source: &mut Source,
|
||||
) -> Result<LocalExecutionRequestOutcome, GatewayError>
|
||||
where
|
||||
Port: AiAttemptLoopPort<
|
||||
Attempt,
|
||||
Response = Response<Body>,
|
||||
Exhaustion = crate::executor::LocalExecutionExhaustion,
|
||||
Error = GatewayError,
|
||||
>,
|
||||
Source: LocalExecutionAttemptSource<Attempt>,
|
||||
Attempt: AiExecutionAttempt + Send + Sync + 'static,
|
||||
{
|
||||
let mut last_attempted = None;
|
||||
|
||||
while let Some(attempt) = source.next_execution_attempt().await? {
|
||||
last_attempted = Some((attempt.execution_plan().clone(), attempt.report_context()));
|
||||
if let Some(response) = port.execute_attempt(&attempt).await? {
|
||||
let remaining = source.drain_execution_attempts().await?;
|
||||
port.mark_unused_attempts(remaining).await?;
|
||||
return Ok(LocalExecutionRequestOutcome::responded(response));
|
||||
}
|
||||
}
|
||||
|
||||
let Some((last_plan, last_report_context)) = last_attempted else {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
};
|
||||
|
||||
Ok(LocalExecutionRequestOutcome::Exhausted(
|
||||
port.build_exhaustion(last_plan, last_report_context)
|
||||
.await?,
|
||||
))
|
||||
}
|
||||
|
||||
struct StreamAttemptLoopPort<'a> {
|
||||
state: &'a AppState,
|
||||
trace_id: &'a str,
|
||||
@@ -307,10 +413,7 @@ where
|
||||
|
||||
fn should_skip_unused_persistence(report_context: Option<&serde_json::Value>) -> bool {
|
||||
let metadata = local_execution_candidate_metadata_from_report_context(report_context);
|
||||
metadata.candidate_group_id.is_some()
|
||||
&& metadata
|
||||
.pool_key_index
|
||||
.is_some_and(|pool_key_index| pool_key_index > 0)
|
||||
metadata.candidate_group_id.is_some() && metadata.pool_key_index.is_some()
|
||||
}
|
||||
|
||||
fn resolve_stream_candidate_watchdog_timeout(plan: &aether_contracts::ExecutionPlan) -> Duration {
|
||||
@@ -520,13 +623,13 @@ mod tests {
|
||||
#[test]
|
||||
fn unused_persistence_skips_pool_internal_candidates() {
|
||||
assert!(should_skip_unused_persistence(Some(&json!({
|
||||
"candidate_group_id": "pool-group",
|
||||
"pool_key_index": 1,
|
||||
}))));
|
||||
assert!(!should_skip_unused_persistence(Some(&json!({
|
||||
"candidate_group_id": "pool-group",
|
||||
"pool_key_index": 0,
|
||||
}))));
|
||||
assert!(should_skip_unused_persistence(Some(&json!({
|
||||
"candidate_group_id": "pool-group",
|
||||
"pool_key_index": 1,
|
||||
}))));
|
||||
assert!(!should_skip_unused_persistence(Some(&json!({
|
||||
"candidate_group_id": "pool-group",
|
||||
}))));
|
||||
|
||||
@@ -1,24 +1,29 @@
|
||||
use crate::ai_serving::api::{
|
||||
build_local_gemini_files_stream_plan_and_reports_for_kind,
|
||||
build_local_gemini_files_sync_plan_and_reports_for_kind,
|
||||
build_local_image_stream_plan_and_reports_for_kind,
|
||||
build_local_image_sync_plan_and_reports_for_kind,
|
||||
build_local_gemini_files_stream_attempt_source_for_kind,
|
||||
build_local_gemini_files_sync_attempt_source_for_kind,
|
||||
build_local_image_stream_attempt_source_for_kind,
|
||||
build_local_image_sync_attempt_source_for_kind,
|
||||
build_local_openai_chat_stream_attempt_source_for_kind,
|
||||
build_local_openai_chat_stream_plan_and_reports_for_kind,
|
||||
build_local_openai_chat_sync_attempt_source_for_kind,
|
||||
build_local_openai_chat_sync_plan_and_reports_for_kind,
|
||||
build_local_openai_responses_stream_attempt_source_for_kind,
|
||||
build_local_openai_responses_stream_plan_and_reports_for_kind,
|
||||
build_local_openai_responses_sync_attempt_source_for_kind,
|
||||
build_local_openai_responses_sync_plan_and_reports_for_kind,
|
||||
build_local_same_format_stream_plan_and_reports, build_local_same_format_sync_plan_and_reports,
|
||||
build_local_video_sync_plan_and_reports_for_kind,
|
||||
build_standard_family_stream_plan_and_reports, build_standard_family_sync_plan_and_reports,
|
||||
parse_direct_request_body, resolve_claude_stream_spec, resolve_claude_sync_spec,
|
||||
resolve_gemini_stream_spec, resolve_gemini_sync_spec, resolve_local_same_format_stream_spec,
|
||||
build_local_same_format_stream_attempt_source, build_local_same_format_stream_plan_and_reports,
|
||||
build_local_same_format_sync_attempt_source, build_local_same_format_sync_plan_and_reports,
|
||||
build_local_video_sync_attempt_source_for_kind, build_standard_family_stream_attempt_source,
|
||||
build_standard_family_sync_attempt_source, parse_direct_request_body,
|
||||
resolve_claude_stream_spec, resolve_claude_sync_spec, resolve_gemini_stream_spec,
|
||||
resolve_gemini_sync_spec, resolve_local_same_format_stream_spec,
|
||||
resolve_local_same_format_sync_spec, set_local_openai_chat_execution_exhausted_diagnostic,
|
||||
AiStreamAttempt, AiSyncAttempt, LocalStandardSpec, EXECUTION_RUNTIME_STREAM_DECISION_ACTION,
|
||||
EXECUTION_RUNTIME_SYNC_DECISION_ACTION,
|
||||
};
|
||||
use crate::control::GatewayControlDecision;
|
||||
use crate::executor::candidate_loop::{
|
||||
execute_stream_plan_and_reports, execute_sync_plan_and_reports,
|
||||
execute_stream_attempt_source, execute_sync_attempt_source, execute_sync_plan_and_reports,
|
||||
};
|
||||
use crate::executor::LocalExecutionRequestOutcome;
|
||||
use crate::{AiExecutionDecision, AppState, GatewayError};
|
||||
@@ -52,28 +57,33 @@ pub(crate) async fn maybe_execute_sync_via_local_decision(
|
||||
body_json: &serde_json::Value,
|
||||
plan_kind: &str,
|
||||
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
|
||||
let plan_and_reports = build_local_openai_chat_sync_plan_and_reports_for_kind(
|
||||
state, parts, trace_id, decision, body_json, plan_kind,
|
||||
)
|
||||
.await?;
|
||||
if plan_and_reports.is_empty() {
|
||||
let Some((attempt_source, candidate_count)) =
|
||||
build_local_openai_chat_sync_attempt_source_for_kind(
|
||||
state, parts, trace_id, decision, body_json, plan_kind,
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
};
|
||||
|
||||
let plan_count = plan_and_reports.len();
|
||||
let outcome = execute_sync_plan_and_reports(
|
||||
let outcome = execute_sync_attempt_source::<AiSyncAttempt, _>(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
plan_and_reports,
|
||||
attempt_source,
|
||||
)
|
||||
.await?;
|
||||
|
||||
if let LocalExecutionRequestOutcome::Exhausted(_) = &outcome {
|
||||
set_local_openai_chat_execution_exhausted_diagnostic(
|
||||
state, trace_id, decision, plan_kind, body_json, plan_count,
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
body_json,
|
||||
candidate_count,
|
||||
);
|
||||
}
|
||||
|
||||
@@ -88,22 +98,32 @@ pub(crate) async fn maybe_execute_stream_via_local_decision(
|
||||
body_json: &serde_json::Value,
|
||||
plan_kind: &str,
|
||||
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
|
||||
let plan_and_reports = build_local_openai_chat_stream_plan_and_reports_for_kind(
|
||||
state, parts, trace_id, decision, body_json, plan_kind,
|
||||
let Some((attempt_source, candidate_count)) =
|
||||
build_local_openai_chat_stream_attempt_source_for_kind(
|
||||
state, parts, trace_id, decision, body_json, plan_kind,
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
};
|
||||
|
||||
let outcome = execute_stream_attempt_source::<AiStreamAttempt, _>(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
attempt_source,
|
||||
)
|
||||
.await?;
|
||||
if plan_and_reports.is_empty() {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
|
||||
let plan_count = plan_and_reports.len();
|
||||
let outcome =
|
||||
execute_stream_plan_and_reports(state, trace_id, decision, plan_kind, plan_and_reports)
|
||||
.await?;
|
||||
|
||||
if let LocalExecutionRequestOutcome::Exhausted(_) = &outcome {
|
||||
set_local_openai_chat_execution_exhausted_diagnostic(
|
||||
state, trace_id, decision, plan_kind, body_json, plan_count,
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
body_json,
|
||||
candidate_count,
|
||||
);
|
||||
}
|
||||
|
||||
@@ -118,22 +138,22 @@ pub(crate) async fn maybe_execute_sync_via_local_openai_responses_decision(
|
||||
body_json: &serde_json::Value,
|
||||
plan_kind: &str,
|
||||
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
|
||||
let plan_and_reports: Vec<AiSyncAttempt> =
|
||||
build_local_openai_responses_sync_plan_and_reports_for_kind(
|
||||
let Some((attempt_source, _candidate_count)) =
|
||||
build_local_openai_responses_sync_attempt_source_for_kind(
|
||||
state, parts, trace_id, decision, body_json, plan_kind,
|
||||
)
|
||||
.await?;
|
||||
if plan_and_reports.is_empty() {
|
||||
.await?
|
||||
else {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
};
|
||||
|
||||
execute_sync_plan_and_reports(
|
||||
execute_sync_attempt_source::<AiSyncAttempt, _>(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
plan_and_reports,
|
||||
attempt_source,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -146,16 +166,23 @@ pub(crate) async fn maybe_execute_stream_via_local_openai_responses_decision(
|
||||
body_json: &serde_json::Value,
|
||||
plan_kind: &str,
|
||||
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
|
||||
let plan_and_reports: Vec<AiStreamAttempt> =
|
||||
build_local_openai_responses_stream_plan_and_reports_for_kind(
|
||||
let Some((attempt_source, _candidate_count)) =
|
||||
build_local_openai_responses_stream_attempt_source_for_kind(
|
||||
state, parts, trace_id, decision, body_json, plan_kind,
|
||||
)
|
||||
.await?;
|
||||
if plan_and_reports.is_empty() {
|
||||
.await?
|
||||
else {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
};
|
||||
|
||||
execute_stream_plan_and_reports(state, trace_id, decision, plan_kind, plan_and_reports).await
|
||||
execute_stream_attempt_source::<AiStreamAttempt, _>(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
attempt_source,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn maybe_execute_sync_via_standard_family_decision(
|
||||
@@ -171,21 +198,21 @@ pub(crate) async fn maybe_execute_sync_via_standard_family_decision(
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
};
|
||||
|
||||
let plan_and_reports: Vec<AiSyncAttempt> = build_standard_family_sync_plan_and_reports(
|
||||
let Some((attempt_source, _candidate_count)) = build_standard_family_sync_attempt_source(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await?;
|
||||
if plan_and_reports.is_empty() {
|
||||
.await?
|
||||
else {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
};
|
||||
|
||||
execute_sync_plan_and_reports(
|
||||
execute_sync_attempt_source::<AiSyncAttempt, _>(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
plan_and_reports,
|
||||
attempt_source,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -203,15 +230,22 @@ pub(crate) async fn maybe_execute_stream_via_standard_family_decision(
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
};
|
||||
|
||||
let plan_and_reports: Vec<AiStreamAttempt> = build_standard_family_stream_plan_and_reports(
|
||||
let Some((attempt_source, _candidate_count)) = build_standard_family_stream_attempt_source(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await?;
|
||||
if plan_and_reports.is_empty() {
|
||||
.await?
|
||||
else {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
};
|
||||
|
||||
execute_stream_plan_and_reports(state, trace_id, decision, plan_kind, plan_and_reports).await
|
||||
execute_stream_attempt_source::<AiStreamAttempt, _>(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
attempt_source,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn maybe_execute_sync_via_local_standard_decision(
|
||||
@@ -328,21 +362,21 @@ pub(crate) async fn maybe_execute_sync_via_local_same_format_provider_decision(
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
};
|
||||
|
||||
let plan_and_reports: Vec<AiSyncAttempt> = build_local_same_format_sync_plan_and_reports(
|
||||
let Some((attempt_source, _candidate_count)) = build_local_same_format_sync_attempt_source(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await?;
|
||||
if plan_and_reports.is_empty() {
|
||||
.await?
|
||||
else {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
};
|
||||
|
||||
execute_sync_plan_and_reports(
|
||||
execute_sync_attempt_source::<AiSyncAttempt, _>(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
plan_and_reports,
|
||||
attempt_source,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -359,15 +393,22 @@ pub(crate) async fn maybe_execute_stream_via_local_same_format_provider_decision
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
};
|
||||
|
||||
let plan_and_reports: Vec<AiStreamAttempt> = build_local_same_format_stream_plan_and_reports(
|
||||
let Some((attempt_source, _candidate_count)) = build_local_same_format_stream_attempt_source(
|
||||
state, parts, trace_id, decision, body_json, spec,
|
||||
)
|
||||
.await?;
|
||||
if plan_and_reports.is_empty() {
|
||||
.await?
|
||||
else {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
};
|
||||
|
||||
execute_stream_plan_and_reports(state, trace_id, decision, plan_kind, plan_and_reports).await
|
||||
execute_stream_attempt_source::<AiStreamAttempt, _>(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
attempt_source,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn maybe_execute_sync_via_local_gemini_files_decision(
|
||||
@@ -380,8 +421,8 @@ pub(crate) async fn maybe_execute_sync_via_local_gemini_files_decision(
|
||||
decision: &GatewayControlDecision,
|
||||
plan_kind: &str,
|
||||
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
|
||||
let plan_and_reports: Vec<AiSyncAttempt> =
|
||||
build_local_gemini_files_sync_plan_and_reports_for_kind(
|
||||
let Some((attempt_source, _candidate_count)) =
|
||||
build_local_gemini_files_sync_attempt_source_for_kind(
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
@@ -391,18 +432,18 @@ pub(crate) async fn maybe_execute_sync_via_local_gemini_files_decision(
|
||||
decision,
|
||||
plan_kind,
|
||||
)
|
||||
.await?;
|
||||
if plan_and_reports.is_empty() {
|
||||
.await?
|
||||
else {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
};
|
||||
|
||||
execute_sync_plan_and_reports(
|
||||
execute_sync_attempt_source::<AiSyncAttempt, _>(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
plan_and_reports,
|
||||
attempt_source,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -416,7 +457,7 @@ pub(crate) async fn maybe_execute_sync_via_local_image_decision(
|
||||
decision: &GatewayControlDecision,
|
||||
plan_kind: &str,
|
||||
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
|
||||
let plan_and_reports: Vec<AiSyncAttempt> = build_local_image_sync_plan_and_reports_for_kind(
|
||||
let Some((attempt_source, _candidate_count)) = build_local_image_sync_attempt_source_for_kind(
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
@@ -425,18 +466,18 @@ pub(crate) async fn maybe_execute_sync_via_local_image_decision(
|
||||
decision,
|
||||
plan_kind,
|
||||
)
|
||||
.await?;
|
||||
if plan_and_reports.is_empty() {
|
||||
.await?
|
||||
else {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
};
|
||||
|
||||
execute_sync_plan_and_reports(
|
||||
execute_sync_attempt_source::<AiSyncAttempt, _>(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
plan_and_reports,
|
||||
attempt_source,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -448,16 +489,23 @@ pub(crate) async fn maybe_execute_stream_via_local_gemini_files_decision(
|
||||
decision: &GatewayControlDecision,
|
||||
plan_kind: &str,
|
||||
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
|
||||
let plan_and_reports: Vec<AiStreamAttempt> =
|
||||
build_local_gemini_files_stream_plan_and_reports_for_kind(
|
||||
let Some((attempt_source, _candidate_count)) =
|
||||
build_local_gemini_files_stream_attempt_source_for_kind(
|
||||
state, parts, trace_id, decision, plan_kind,
|
||||
)
|
||||
.await?;
|
||||
if plan_and_reports.is_empty() {
|
||||
.await?
|
||||
else {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
};
|
||||
|
||||
execute_stream_plan_and_reports(state, trace_id, decision, plan_kind, plan_and_reports).await
|
||||
execute_stream_attempt_source::<AiStreamAttempt, _>(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
attempt_source,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn maybe_execute_stream_via_local_image_decision(
|
||||
@@ -469,8 +517,8 @@ pub(crate) async fn maybe_execute_stream_via_local_image_decision(
|
||||
decision: &GatewayControlDecision,
|
||||
plan_kind: &str,
|
||||
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
|
||||
let plan_and_reports: Vec<AiStreamAttempt> =
|
||||
build_local_image_stream_plan_and_reports_for_kind(
|
||||
let Some((attempt_source, _candidate_count)) =
|
||||
build_local_image_stream_attempt_source_for_kind(
|
||||
state,
|
||||
parts,
|
||||
body_json,
|
||||
@@ -479,12 +527,19 @@ pub(crate) async fn maybe_execute_stream_via_local_image_decision(
|
||||
decision,
|
||||
plan_kind,
|
||||
)
|
||||
.await?;
|
||||
if plan_and_reports.is_empty() {
|
||||
.await?
|
||||
else {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
};
|
||||
|
||||
execute_stream_plan_and_reports(state, trace_id, decision, plan_kind, plan_and_reports).await
|
||||
execute_stream_attempt_source::<AiStreamAttempt, _>(
|
||||
state,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
attempt_source,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn maybe_execute_sync_via_local_video_decision(
|
||||
@@ -495,21 +550,21 @@ pub(crate) async fn maybe_execute_sync_via_local_video_decision(
|
||||
decision: &GatewayControlDecision,
|
||||
plan_kind: &str,
|
||||
) -> Result<LocalExecutionRequestOutcome, GatewayError> {
|
||||
let plan_and_reports: Vec<AiSyncAttempt> = build_local_video_sync_plan_and_reports_for_kind(
|
||||
let Some((attempt_source, _candidate_count)) = build_local_video_sync_attempt_source_for_kind(
|
||||
state, parts, body_json, trace_id, decision, plan_kind,
|
||||
)
|
||||
.await?;
|
||||
if plan_and_reports.is_empty() {
|
||||
.await?
|
||||
else {
|
||||
return Ok(LocalExecutionRequestOutcome::NoPath);
|
||||
}
|
||||
};
|
||||
|
||||
execute_sync_plan_and_reports(
|
||||
execute_sync_attempt_source::<AiSyncAttempt, _>(
|
||||
state,
|
||||
parts,
|
||||
trace_id,
|
||||
decision,
|
||||
plan_kind,
|
||||
plan_and_reports,
|
||||
attempt_source,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
@@ -26,9 +26,16 @@ pub(super) fn rank_scheduler_candidates(
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(index, candidate)| {
|
||||
let provider_key = runtime_snapshot
|
||||
.provider_key_rpm_states
|
||||
.get(&candidate.key_id);
|
||||
let pool_group = runtime_snapshot
|
||||
.pool_provider_ids
|
||||
.contains(candidate.provider_id.as_str());
|
||||
let provider_key = (!pool_group)
|
||||
.then(|| {
|
||||
runtime_snapshot
|
||||
.provider_key_rpm_states
|
||||
.get(&candidate.key_id)
|
||||
})
|
||||
.flatten();
|
||||
SchedulerRankableCandidate::from_candidate(candidate, index)
|
||||
.with_capability_priority(requested_capability_priority_for_candidate(
|
||||
required_capabilities,
|
||||
|
||||
@@ -22,6 +22,7 @@ pub(super) struct CandidateRuntimeSelectionSnapshot {
|
||||
pub(super) recent_candidates: Vec<StoredRequestCandidate>,
|
||||
pub(super) provider_concurrent_limits: BTreeMap<String, usize>,
|
||||
pub(super) provider_key_rpm_states: BTreeMap<String, StoredProviderCatalogKey>,
|
||||
pub(super) pool_provider_ids: BTreeSet<String>,
|
||||
provider_quota_blocks_requests: BTreeMap<String, bool>,
|
||||
key_account_quota_exhausted: BTreeMap<String, bool>,
|
||||
key_oauth_invalid: BTreeMap<String, bool>,
|
||||
@@ -35,8 +36,15 @@ pub(super) async fn read_candidate_runtime_selection_snapshot(
|
||||
) -> Result<CandidateRuntimeSelectionSnapshot, GatewayError> {
|
||||
let recent_candidates = state.read_recent_request_candidates(128).await?;
|
||||
let provider_concurrent_limits = read_provider_concurrent_limits(state, candidates).await?;
|
||||
let provider_skip_exhausted_accounts =
|
||||
read_provider_skip_exhausted_account_map(state, candidates).await?;
|
||||
let provider_pool_state = read_provider_pool_state_map(state, candidates).await?;
|
||||
let provider_skip_exhausted_accounts = provider_pool_state
|
||||
.iter()
|
||||
.map(|(provider_id, state)| (provider_id.clone(), state.skip_exhausted_accounts))
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
let pool_provider_ids = provider_pool_state
|
||||
.iter()
|
||||
.filter_map(|(provider_id, state)| state.pool_enabled.then_some(provider_id.clone()))
|
||||
.collect::<BTreeSet<_>>();
|
||||
let provider_key_rpm_states = read_provider_key_rpm_states(state, candidates).await?;
|
||||
let key_account_quota_exhausted = read_key_account_quota_exhaustion_map(
|
||||
candidates,
|
||||
@@ -54,6 +62,7 @@ pub(super) async fn read_candidate_runtime_selection_snapshot(
|
||||
recent_candidates,
|
||||
provider_concurrent_limits,
|
||||
provider_key_rpm_states,
|
||||
pool_provider_ids,
|
||||
provider_quota_blocks_requests,
|
||||
key_account_quota_exhausted,
|
||||
key_oauth_invalid,
|
||||
@@ -93,6 +102,9 @@ pub(super) fn is_candidate_selectable(
|
||||
now_unix_secs: u64,
|
||||
cached_affinity_target: Option<&SchedulerAffinityTarget>,
|
||||
) -> bool {
|
||||
let pool_group = snapshot
|
||||
.pool_provider_ids
|
||||
.contains(candidate.provider_id.as_str());
|
||||
candidate_is_selectable_with_runtime_state(CandidateRuntimeSelectabilityInput {
|
||||
candidate,
|
||||
recent_candidates: &snapshot.recent_candidates,
|
||||
@@ -105,20 +117,26 @@ pub(super) fn is_candidate_selectable(
|
||||
.get(candidate.provider_id.as_str())
|
||||
.copied()
|
||||
.unwrap_or(false),
|
||||
account_quota_exhausted: snapshot
|
||||
.key_account_quota_exhausted
|
||||
.get(candidate.key_id.as_str())
|
||||
.copied()
|
||||
.unwrap_or(false),
|
||||
oauth_invalid: snapshot
|
||||
.key_oauth_invalid
|
||||
.get(candidate.key_id.as_str())
|
||||
.copied()
|
||||
.unwrap_or(false),
|
||||
rpm_reset_at: snapshot
|
||||
.provider_key_rpm_reset_ats
|
||||
.get(candidate.key_id.as_str())
|
||||
.copied()
|
||||
account_quota_exhausted: !pool_group
|
||||
&& snapshot
|
||||
.key_account_quota_exhausted
|
||||
.get(candidate.key_id.as_str())
|
||||
.copied()
|
||||
.unwrap_or(false),
|
||||
oauth_invalid: !pool_group
|
||||
&& snapshot
|
||||
.key_oauth_invalid
|
||||
.get(candidate.key_id.as_str())
|
||||
.copied()
|
||||
.unwrap_or(false),
|
||||
rpm_reset_at: (!pool_group)
|
||||
.then(|| {
|
||||
snapshot
|
||||
.provider_key_rpm_reset_ats
|
||||
.get(candidate.key_id.as_str())
|
||||
.copied()
|
||||
.flatten()
|
||||
})
|
||||
.flatten(),
|
||||
})
|
||||
}
|
||||
@@ -129,15 +147,22 @@ pub(super) fn current_candidate_runtime_skip_reason(
|
||||
now_unix_secs: u64,
|
||||
cached_affinity_target: Option<&SchedulerAffinityTarget>,
|
||||
) -> Option<&'static str> {
|
||||
let pool_group = snapshot
|
||||
.pool_provider_ids
|
||||
.contains(candidate.provider_id.as_str());
|
||||
let provider_quota_blocks_requests = snapshot
|
||||
.provider_quota_blocks_requests
|
||||
.get(candidate.provider_id.as_str())
|
||||
.copied()
|
||||
.unwrap_or(false);
|
||||
let rpm_reset_at = snapshot
|
||||
.provider_key_rpm_reset_ats
|
||||
.get(candidate.key_id.as_str())
|
||||
.copied()
|
||||
let rpm_reset_at = (!pool_group)
|
||||
.then(|| {
|
||||
snapshot
|
||||
.provider_key_rpm_reset_ats
|
||||
.get(candidate.key_id.as_str())
|
||||
.copied()
|
||||
.flatten()
|
||||
})
|
||||
.flatten();
|
||||
|
||||
candidate_runtime_skip_reason_with_state(CandidateRuntimeSelectabilityInput {
|
||||
@@ -148,16 +173,18 @@ pub(super) fn current_candidate_runtime_skip_reason(
|
||||
now_unix_secs,
|
||||
cached_affinity_target,
|
||||
provider_quota_blocks_requests,
|
||||
account_quota_exhausted: snapshot
|
||||
.key_account_quota_exhausted
|
||||
.get(candidate.key_id.as_str())
|
||||
.copied()
|
||||
.unwrap_or(false),
|
||||
oauth_invalid: snapshot
|
||||
.key_oauth_invalid
|
||||
.get(candidate.key_id.as_str())
|
||||
.copied()
|
||||
.unwrap_or(false),
|
||||
account_quota_exhausted: !pool_group
|
||||
&& snapshot
|
||||
.key_account_quota_exhausted
|
||||
.get(candidate.key_id.as_str())
|
||||
.copied()
|
||||
.unwrap_or(false),
|
||||
oauth_invalid: !pool_group
|
||||
&& snapshot
|
||||
.key_oauth_invalid
|
||||
.get(candidate.key_id.as_str())
|
||||
.copied()
|
||||
.unwrap_or(false),
|
||||
rpm_reset_at,
|
||||
})
|
||||
}
|
||||
@@ -228,10 +255,16 @@ async fn read_provider_quota_block_map(
|
||||
Ok(quota_blocks)
|
||||
}
|
||||
|
||||
async fn read_provider_skip_exhausted_account_map(
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
struct ProviderPoolState {
|
||||
pool_enabled: bool,
|
||||
skip_exhausted_accounts: bool,
|
||||
}
|
||||
|
||||
async fn read_provider_pool_state_map(
|
||||
state: &(impl SchedulerRuntimeState + ?Sized),
|
||||
candidates: &[SchedulerMinimalCandidateSelectionCandidate],
|
||||
) -> Result<BTreeMap<String, bool>, GatewayError> {
|
||||
) -> Result<BTreeMap<String, ProviderPoolState>, GatewayError> {
|
||||
let provider_ids = candidates
|
||||
.iter()
|
||||
.map(|candidate| candidate.provider_id.clone())
|
||||
@@ -248,15 +281,22 @@ async fn read_provider_skip_exhausted_account_map(
|
||||
Ok(providers
|
||||
.into_iter()
|
||||
.map(|provider| {
|
||||
let skip_exhausted_accounts = provider
|
||||
let pool_advanced = provider
|
||||
.config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("pool_advanced"))
|
||||
.and_then(|value| value.get("pool_advanced"));
|
||||
let skip_exhausted_accounts = pool_advanced
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|value| value.get("skip_exhausted_accounts"))
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false);
|
||||
(provider.id, skip_exhausted_accounts)
|
||||
(
|
||||
provider.id,
|
||||
ProviderPoolState {
|
||||
pool_enabled: pool_advanced.is_some(),
|
||||
skip_exhausted_accounts,
|
||||
},
|
||||
)
|
||||
})
|
||||
.collect())
|
||||
}
|
||||
|
||||
@@ -1638,11 +1638,14 @@ async fn skips_codex_candidate_when_account_quota_is_exhausted_and_pool_flag_ena
|
||||
.await
|
||||
.expect("selection should succeed");
|
||||
|
||||
assert_eq!(selected.len(), 1);
|
||||
assert_eq!(selected[0].provider_id, "provider-openai");
|
||||
assert_eq!(skipped.len(), 1);
|
||||
assert_eq!(skipped[0].candidate.provider_id, "provider-codex");
|
||||
assert_eq!(skipped[0].skip_reason, "account_quota_exhausted");
|
||||
assert_eq!(selected.len(), 2);
|
||||
assert!(selected
|
||||
.iter()
|
||||
.any(|item| item.provider_id == "provider-codex"));
|
||||
assert!(selected
|
||||
.iter()
|
||||
.any(|item| item.provider_id == "provider-openai"));
|
||||
assert!(skipped.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -2240,11 +2243,14 @@ async fn skips_kiro_candidate_when_account_quota_is_exhausted_and_pool_flag_enab
|
||||
.await
|
||||
.expect("selection should succeed");
|
||||
|
||||
assert_eq!(selected.len(), 1);
|
||||
assert_eq!(selected[0].provider_id, "provider-openai");
|
||||
assert_eq!(skipped.len(), 1);
|
||||
assert_eq!(skipped[0].candidate.provider_id, "provider-kiro");
|
||||
assert_eq!(skipped[0].skip_reason, "account_quota_exhausted");
|
||||
assert_eq!(selected.len(), 2);
|
||||
assert!(selected
|
||||
.iter()
|
||||
.any(|item| item.provider_id == "provider-kiro"));
|
||||
assert!(selected
|
||||
.iter()
|
||||
.any(|item| item.provider_id == "provider-openai"));
|
||||
assert!(skipped.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -23,6 +23,40 @@ impl AppState {
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn list_minimal_candidate_selection_rows_for_api_format_and_requested_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
requested_model_name: &str,
|
||||
) -> Result<Vec<candidate_selection::StoredMinimalCandidateSelectionRow>, GatewayError> {
|
||||
self.data
|
||||
.list_minimal_candidate_selection_rows_for_requested_model(
|
||||
api_format,
|
||||
requested_model_name,
|
||||
)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn list_minimal_candidate_selection_rows_for_api_format_and_requested_model_page(
|
||||
&self,
|
||||
query: &candidate_selection::StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<candidate_selection::StoredMinimalCandidateSelectionRow>, GatewayError> {
|
||||
self.data
|
||||
.list_minimal_candidate_selection_rows_for_requested_model_page(query)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn list_pool_key_candidate_rows_for_group(
|
||||
&self,
|
||||
query: &candidate_selection::StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<candidate_selection::StoredMinimalCandidateSelectionRow>, GatewayError> {
|
||||
self.data
|
||||
.list_pool_key_candidate_rows_for_group(query)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(err.to_string()))
|
||||
}
|
||||
|
||||
pub(crate) async fn read_provider_quota_snapshot(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
|
||||
@@ -410,7 +410,7 @@ async fn gateway_executes_openai_chat_stream_via_local_decision_gate_without_exe
|
||||
.list_by_request_id("trace-openai-chat-local-stream-123")
|
||||
.await
|
||||
.expect("request candidate trace should read");
|
||||
assert_eq!(stored_candidates.len(), 2);
|
||||
assert_eq!(stored_candidates.len(), 1);
|
||||
assert_eq!(
|
||||
stored_candidates
|
||||
.iter()
|
||||
@@ -907,10 +907,8 @@ async fn gateway_executes_openai_chat_stream_via_local_openai_responses_cross_fo
|
||||
.extra_data
|
||||
.as_ref()
|
||||
.expect("request candidate extra_data should exist");
|
||||
assert_eq!(extra_data["execution_strategy"], "local_cross_format");
|
||||
assert_eq!(extra_data["conversion_mode"], "bidirectional");
|
||||
assert_eq!(extra_data["client_contract"], "openai:chat");
|
||||
assert_eq!(extra_data["provider_contract"], "openai:responses");
|
||||
assert_eq!(extra_data["client_api_format"], "openai:chat");
|
||||
assert_eq!(extra_data["provider_api_format"], "openai:responses");
|
||||
|
||||
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
|
||||
assert!(
|
||||
|
||||
@@ -359,7 +359,7 @@ async fn gateway_executes_openai_chat_sync_via_local_decision_gate_without_execu
|
||||
.list_by_request_id("trace-openai-chat-local-123")
|
||||
.await
|
||||
.expect("request candidate trace should read");
|
||||
assert_eq!(stored_candidates.len(), 2);
|
||||
assert_eq!(stored_candidates.len(), 1);
|
||||
assert_eq!(
|
||||
stored_candidates
|
||||
.iter()
|
||||
@@ -933,7 +933,7 @@ async fn gateway_executes_openai_chat_sync_via_local_cross_format_gemini_candida
|
||||
.list_by_request_id("trace-openai-chat-gemini-local-123")
|
||||
.await
|
||||
.expect("request candidate trace should read");
|
||||
assert_eq!(stored_candidates.len(), 2);
|
||||
assert_eq!(stored_candidates.len(), 1);
|
||||
assert_eq!(
|
||||
stored_candidates
|
||||
.iter()
|
||||
@@ -941,24 +941,6 @@ async fn gateway_executes_openai_chat_sync_via_local_cross_format_gemini_candida
|
||||
.count(),
|
||||
1
|
||||
);
|
||||
let skipped_candidate = stored_candidates
|
||||
.iter()
|
||||
.find(|candidate| candidate.status == RequestCandidateStatus::Skipped)
|
||||
.expect("disabled conversion candidate should be persisted as skipped");
|
||||
assert_eq!(
|
||||
skipped_candidate.skip_reason.as_deref(),
|
||||
Some("format_conversion_disabled")
|
||||
);
|
||||
let extra_data = skipped_candidate
|
||||
.extra_data
|
||||
.as_ref()
|
||||
.expect("skipped cross-format candidate extra_data should exist");
|
||||
assert_eq!(extra_data["execution_strategy"], "local_cross_format");
|
||||
assert_eq!(
|
||||
extra_data["transport_diagnostics"]["request_pair"]["conversion_enabled"],
|
||||
false
|
||||
);
|
||||
|
||||
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
|
||||
assert!(
|
||||
!*seen_report.lock().expect("mutex should lock"),
|
||||
|
||||
@@ -141,7 +141,7 @@ where
|
||||
}
|
||||
|
||||
pub fn ai_should_persist_available_candidate_for_pool_key(pool_key_index: Option<u32>) -> bool {
|
||||
pool_key_index.is_none_or(|index| index == 0)
|
||||
pool_key_index.is_none()
|
||||
}
|
||||
|
||||
pub fn ai_should_persist_skipped_candidate_for_pool_membership(is_pool_candidate: bool) -> bool {
|
||||
@@ -387,9 +387,9 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn pool_candidate_persistence_policy_persists_representatives_only() {
|
||||
fn pool_candidate_persistence_policy_skips_pool_keys_until_execution() {
|
||||
assert!(ai_should_persist_available_candidate_for_pool_key(None));
|
||||
assert!(ai_should_persist_available_candidate_for_pool_key(Some(0)));
|
||||
assert!(!ai_should_persist_available_candidate_for_pool_key(Some(0)));
|
||||
assert!(!ai_should_persist_available_candidate_for_pool_key(Some(1)));
|
||||
|
||||
assert!(ai_should_persist_skipped_candidate_for_pool_membership(
|
||||
|
||||
@@ -11,6 +11,7 @@ pub struct AiCandidateResolutionRequest<'a> {
|
||||
pub client_api_format: &'a str,
|
||||
pub requested_model: Option<&'a str>,
|
||||
pub mode: AiCandidateResolutionMode,
|
||||
pub expand_pool_groups: bool,
|
||||
}
|
||||
|
||||
impl<'a> AiCandidateResolutionRequest<'a> {
|
||||
@@ -19,6 +20,7 @@ impl<'a> AiCandidateResolutionRequest<'a> {
|
||||
client_api_format,
|
||||
requested_model,
|
||||
mode: AiCandidateResolutionMode::Standard,
|
||||
expand_pool_groups: true,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -30,8 +32,14 @@ impl<'a> AiCandidateResolutionRequest<'a> {
|
||||
client_api_format,
|
||||
requested_model,
|
||||
mode: AiCandidateResolutionMode::WithoutTransportPairGate,
|
||||
expand_pool_groups: true,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn logical_pool_groups(mut self) -> Self {
|
||||
self.expand_pool_groups = false;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
@@ -140,8 +148,13 @@ where
|
||||
let ranked = port
|
||||
.rank_eligible_candidates(eligible, normalized_client_api_format.as_str())
|
||||
.await?;
|
||||
let (ranked, pool_skipped) = port.apply_pool_scheduler(ranked).await?;
|
||||
skipped.extend(pool_skipped);
|
||||
let ranked = if request.expand_pool_groups {
|
||||
let (ranked, pool_skipped) = port.apply_pool_scheduler(ranked).await?;
|
||||
skipped.extend(pool_skipped);
|
||||
ranked
|
||||
} else {
|
||||
ranked
|
||||
};
|
||||
|
||||
Ok(AiCandidateResolutionOutcome {
|
||||
eligible_candidates: ranked,
|
||||
@@ -385,6 +398,38 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn resolution_can_keep_pool_groups_logical() {
|
||||
let port = TestPort::default();
|
||||
|
||||
let outcome = run_ai_candidate_resolution(
|
||||
&port,
|
||||
vec!["first", "second"],
|
||||
AiCandidateResolutionRequest::standard("openai:chat", Some("gpt-4.1"))
|
||||
.logical_pool_groups(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
outcome.eligible_candidates,
|
||||
["eligible:second", "eligible:first"]
|
||||
);
|
||||
assert!(outcome.skipped_candidates.is_empty());
|
||||
assert_eq!(
|
||||
port.calls.lock().unwrap().as_slice(),
|
||||
[
|
||||
"transport:first",
|
||||
"common:first:gpt-4.1",
|
||||
"pair:first:openai:chat:gpt-4.1",
|
||||
"transport:second",
|
||||
"common:second:gpt-4.1",
|
||||
"pair:second:openai:chat:gpt-4.1",
|
||||
"rank:openai:chat",
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sticky_session_token_is_extracted_from_known_request_fields() {
|
||||
assert_eq!(
|
||||
|
||||
@@ -2,5 +2,6 @@ mod types;
|
||||
|
||||
pub use types::{
|
||||
MinimalCandidateSelectionReadRepository, MinimalCandidateSelectionRepository,
|
||||
StoredMinimalCandidateSelectionRow, StoredProviderModelMapping,
|
||||
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery,
|
||||
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
|
||||
@@ -40,6 +40,25 @@ pub struct StoredMinimalCandidateSelectionRow {
|
||||
pub model_is_available: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredPoolKeyCandidateRowsQuery {
|
||||
pub api_format: String,
|
||||
pub provider_id: String,
|
||||
pub endpoint_id: String,
|
||||
pub model_id: String,
|
||||
pub selected_provider_model_name: String,
|
||||
pub offset: u32,
|
||||
pub limit: u32,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredRequestedModelCandidateRowsQuery {
|
||||
pub api_format: String,
|
||||
pub requested_model_name: String,
|
||||
pub offset: u32,
|
||||
pub limit: u32,
|
||||
}
|
||||
|
||||
impl StoredMinimalCandidateSelectionRow {
|
||||
pub fn supports_streaming(&self) -> bool {
|
||||
self.model_supports_streaming
|
||||
@@ -73,6 +92,22 @@ pub trait MinimalCandidateSelectionReadRepository: Send + Sync {
|
||||
api_format: &str,
|
||||
global_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model_page(
|
||||
&self,
|
||||
query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
|
||||
|
||||
async fn list_pool_key_rows_for_group(
|
||||
&self,
|
||||
query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
|
||||
}
|
||||
|
||||
pub trait MinimalCandidateSelectionRepository:
|
||||
|
||||
@@ -2,7 +2,10 @@ use std::sync::RwLock;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::{MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow};
|
||||
use super::{
|
||||
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
|
||||
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
@@ -68,6 +71,76 @@ impl MinimalCandidateSelectionReadRepository for InMemoryMinimalCandidateSelecti
|
||||
.filter(|row| row.global_model_name == global_model_name)
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.list_for_exact_api_format_and_requested_model_page(
|
||||
&StoredRequestedModelCandidateRowsQuery {
|
||||
api_format: api_format.to_string(),
|
||||
requested_model_name: requested_model_name.to_string(),
|
||||
offset: 0,
|
||||
limit: u32::MAX,
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model_page(
|
||||
&self,
|
||||
query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let rows = self.list_for_exact_api_format(&query.api_format).await?;
|
||||
let mut rows = rows
|
||||
.into_iter()
|
||||
.filter(|row| {
|
||||
row_matches_requested_model(row, &query.requested_model_name, &query.api_format)
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
rows.sort_by(|left, right| {
|
||||
left.global_model_name
|
||||
.cmp(&right.global_model_name)
|
||||
.then(left.provider_priority.cmp(&right.provider_priority))
|
||||
.then(left.key_internal_priority.cmp(&right.key_internal_priority))
|
||||
.then(left.provider_id.cmp(&right.provider_id))
|
||||
.then(left.endpoint_id.cmp(&right.endpoint_id))
|
||||
.then(left.key_id.cmp(&right.key_id))
|
||||
.then(left.model_id.cmp(&right.model_id))
|
||||
});
|
||||
Ok(rows
|
||||
.into_iter()
|
||||
.skip(query.offset as usize)
|
||||
.take(query.limit as usize)
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group(
|
||||
&self,
|
||||
query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let mut rows = self
|
||||
.list_for_exact_api_format(&query.api_format)
|
||||
.await?
|
||||
.into_iter()
|
||||
.filter(|row| {
|
||||
row.provider_id == query.provider_id
|
||||
&& row.endpoint_id == query.endpoint_id
|
||||
&& row.model_id == query.model_id
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
rows.sort_by(|left, right| {
|
||||
left.key_internal_priority
|
||||
.cmp(&right.key_internal_priority)
|
||||
.then(left.key_id.cmp(&right.key_id))
|
||||
});
|
||||
Ok(rows
|
||||
.into_iter()
|
||||
.skip(query.offset as usize)
|
||||
.take(query.limit as usize)
|
||||
.collect())
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_api_format(value: &str) -> String {
|
||||
@@ -78,6 +151,27 @@ fn api_format_matches(left: &str, right: &str) -> bool {
|
||||
aether_ai_formats::api_format_alias_matches(left, right)
|
||||
}
|
||||
|
||||
fn row_matches_requested_model(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
row.global_model_name == requested_model_name
|
||||
|| row.model_provider_model_name == requested_model_name
|
||||
|| row
|
||||
.model_provider_model_mappings
|
||||
.as_ref()
|
||||
.is_some_and(|mappings| {
|
||||
mappings.iter().any(|mapping| {
|
||||
mapping.api_formats.as_ref().is_none_or(|formats| {
|
||||
formats
|
||||
.iter()
|
||||
.any(|value| api_format_matches(value, api_format))
|
||||
}) && mapping.name == requested_model_name
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn key_auth_channel_matches(row: &StoredMinimalCandidateSelectionRow, api_format: &str) -> bool {
|
||||
let provider_type = row.provider_type.trim().to_ascii_lowercase();
|
||||
let auth_type = row.key_auth_type.trim().to_ascii_lowercase();
|
||||
@@ -114,6 +208,7 @@ mod tests {
|
||||
use super::InMemoryMinimalCandidateSelectionReadRepository;
|
||||
use crate::repository::candidate_selection::{
|
||||
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
|
||||
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
|
||||
fn sample_row(
|
||||
@@ -174,6 +269,61 @@ mod tests {
|
||||
assert_eq!(rows[1].provider_id, "provider-2");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn filters_by_exact_api_format_and_requested_model_aliases() {
|
||||
let mut mapped = sample_row("provider-1", "openai:chat", "gpt-4.1", 10);
|
||||
mapped.model_provider_model_name = "provider-gpt-4.1".to_string();
|
||||
mapped.model_provider_model_mappings = Some(vec![
|
||||
crate::repository::candidate_selection::StoredProviderModelMapping {
|
||||
name: "alias-gpt-4.1".to_string(),
|
||||
priority: 0,
|
||||
api_formats: Some(vec!["openai:chat".to_string()]),
|
||||
},
|
||||
]);
|
||||
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||
mapped,
|
||||
sample_row("provider-2", "openai:chat", "gpt-4.1-mini", 20),
|
||||
]);
|
||||
|
||||
let rows = repository
|
||||
.list_for_exact_api_format_and_requested_model("openai:chat", "alias-gpt-4.1")
|
||||
.await
|
||||
.expect("list should succeed");
|
||||
|
||||
assert_eq!(rows.len(), 1);
|
||||
assert_eq!(rows[0].provider_id, "provider-1");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn requested_model_page_returns_requested_slice_only() {
|
||||
let mut rows = Vec::new();
|
||||
for index in 0..5 {
|
||||
let mut row = sample_row(&format!("provider-{index}"), "openai:chat", "gpt-5", index);
|
||||
row.key_internal_priority = index;
|
||||
rows.push(row);
|
||||
}
|
||||
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(rows);
|
||||
|
||||
let page = repository
|
||||
.list_for_exact_api_format_and_requested_model_page(
|
||||
&StoredRequestedModelCandidateRowsQuery {
|
||||
api_format: "openai:chat".to_string(),
|
||||
requested_model_name: "gpt-5".to_string(),
|
||||
offset: 2,
|
||||
limit: 2,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("page should load");
|
||||
|
||||
assert_eq!(
|
||||
page.iter()
|
||||
.map(|row| row.provider_id.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
vec!["provider-2", "provider-3"]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn filters_by_exact_api_format_only() {
|
||||
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||
@@ -191,4 +341,38 @@ mod tests {
|
||||
assert_eq!(rows[0].provider_id, "provider-1");
|
||||
assert_eq!(rows[1].provider_id, "provider-2");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn list_pool_key_rows_for_group_returns_requested_page_only() {
|
||||
let mut rows = Vec::new();
|
||||
for index in 0..5 {
|
||||
let mut row = sample_row("provider-pool", "openai:chat", "gpt-5", 10);
|
||||
row.endpoint_id = "endpoint-pool".to_string();
|
||||
row.model_id = "model-pool".to_string();
|
||||
row.key_id = format!("key-{index}");
|
||||
row.key_internal_priority = index;
|
||||
rows.push(row);
|
||||
}
|
||||
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(rows);
|
||||
|
||||
let page = repository
|
||||
.list_pool_key_rows_for_group(&StoredPoolKeyCandidateRowsQuery {
|
||||
api_format: "openai:chat".to_string(),
|
||||
provider_id: "provider-pool".to_string(),
|
||||
endpoint_id: "endpoint-pool".to_string(),
|
||||
model_id: "model-pool".to_string(),
|
||||
selected_provider_model_name: "gpt-5".to_string(),
|
||||
offset: 2,
|
||||
limit: 2,
|
||||
})
|
||||
.await
|
||||
.expect("pool key page should load");
|
||||
|
||||
assert_eq!(
|
||||
page.iter()
|
||||
.map(|row| row.key_id.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
vec!["key-2", "key-3"]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,7 +4,8 @@ mod sql;
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use aether_data_contracts::repository::candidate_selection::{
|
||||
MinimalCandidateSelectionReadRepository, MinimalCandidateSelectionRepository,
|
||||
StoredMinimalCandidateSelectionRow, StoredProviderModelMapping,
|
||||
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateRowsQuery,
|
||||
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
pub use memory::InMemoryMinimalCandidateSelectionReadRepository;
|
||||
pub use sql::SqlxMinimalCandidateSelectionReadRepository;
|
||||
|
||||
@@ -5,11 +5,13 @@ use std::collections::BTreeSet;
|
||||
|
||||
use super::{
|
||||
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
|
||||
StoredProviderModelMapping,
|
||||
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
|
||||
StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
|
||||
const LIST_FOR_EXACT_API_FORMAT_SQL: &str = r#"
|
||||
WITH candidate_rows AS (
|
||||
SELECT
|
||||
p.id AS provider_id,
|
||||
p.name AS provider_name,
|
||||
@@ -46,7 +48,8 @@ SELECT
|
||||
m.provider_model_mappings AS model_provider_model_mappings,
|
||||
m.supports_streaming AS model_supports_streaming,
|
||||
m.is_active AS model_is_active,
|
||||
m.is_available AS model_is_available
|
||||
m.is_available AS model_is_available,
|
||||
(p.config -> 'pool_advanced') IS NOT NULL AS provider_pool_enabled
|
||||
FROM providers p
|
||||
INNER JOIN provider_endpoints pe
|
||||
ON pe.provider_id = p.id
|
||||
@@ -124,17 +127,67 @@ WHERE p.is_active = TRUE
|
||||
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
||||
)
|
||||
)
|
||||
),
|
||||
pool_rows AS (
|
||||
SELECT DISTINCT ON (provider_id, endpoint_id, model_id)
|
||||
*
|
||||
FROM candidate_rows
|
||||
WHERE provider_pool_enabled
|
||||
ORDER BY
|
||||
provider_id ASC,
|
||||
endpoint_id ASC,
|
||||
model_id ASC,
|
||||
key_internal_priority ASC,
|
||||
key_id ASC
|
||||
),
|
||||
selected_rows AS (
|
||||
SELECT * FROM candidate_rows WHERE NOT provider_pool_enabled
|
||||
UNION ALL
|
||||
SELECT * FROM pool_rows
|
||||
)
|
||||
SELECT
|
||||
provider_id,
|
||||
provider_name,
|
||||
provider_type,
|
||||
provider_priority,
|
||||
provider_is_active,
|
||||
endpoint_id,
|
||||
endpoint_api_format,
|
||||
endpoint_api_family,
|
||||
endpoint_kind,
|
||||
endpoint_is_active,
|
||||
key_id,
|
||||
key_name,
|
||||
key_auth_type,
|
||||
key_is_active,
|
||||
key_api_formats,
|
||||
key_allowed_models,
|
||||
key_capabilities,
|
||||
key_internal_priority,
|
||||
key_global_priority_by_format,
|
||||
model_id,
|
||||
global_model_id,
|
||||
global_model_name,
|
||||
global_model_mappings,
|
||||
global_model_supports_streaming,
|
||||
model_provider_model_name,
|
||||
model_provider_model_mappings,
|
||||
model_supports_streaming,
|
||||
model_is_active,
|
||||
model_is_available
|
||||
FROM selected_rows
|
||||
ORDER BY
|
||||
gm.name ASC,
|
||||
p.provider_priority ASC,
|
||||
pak.internal_priority ASC,
|
||||
p.id ASC,
|
||||
pe.id ASC,
|
||||
pak.id ASC,
|
||||
m.id ASC
|
||||
global_model_name ASC,
|
||||
provider_priority ASC,
|
||||
key_internal_priority ASC,
|
||||
provider_id ASC,
|
||||
endpoint_id ASC,
|
||||
key_id ASC,
|
||||
model_id ASC
|
||||
"#;
|
||||
|
||||
const LIST_FOR_EXACT_API_FORMAT_AND_GLOBAL_MODEL_SQL: &str = r#"
|
||||
WITH candidate_rows AS (
|
||||
SELECT
|
||||
p.id AS provider_id,
|
||||
p.name AS provider_name,
|
||||
@@ -171,7 +224,8 @@ SELECT
|
||||
m.provider_model_mappings AS model_provider_model_mappings,
|
||||
m.supports_streaming AS model_supports_streaming,
|
||||
m.is_active AS model_is_active,
|
||||
m.is_available AS model_is_available
|
||||
m.is_available AS model_is_available,
|
||||
(p.config -> 'pool_advanced') IS NOT NULL AS provider_pool_enabled
|
||||
FROM providers p
|
||||
INNER JOIN provider_endpoints pe
|
||||
ON pe.provider_id = p.id
|
||||
@@ -250,13 +304,187 @@ WHERE p.is_active = TRUE
|
||||
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
||||
)
|
||||
)
|
||||
),
|
||||
pool_rows AS (
|
||||
SELECT DISTINCT ON (provider_id, endpoint_id, model_id)
|
||||
*
|
||||
FROM candidate_rows
|
||||
WHERE provider_pool_enabled
|
||||
ORDER BY
|
||||
provider_id ASC,
|
||||
endpoint_id ASC,
|
||||
model_id ASC,
|
||||
key_internal_priority ASC,
|
||||
key_id ASC
|
||||
),
|
||||
selected_rows AS (
|
||||
SELECT * FROM candidate_rows WHERE NOT provider_pool_enabled
|
||||
UNION ALL
|
||||
SELECT * FROM pool_rows
|
||||
)
|
||||
SELECT
|
||||
provider_id,
|
||||
provider_name,
|
||||
provider_type,
|
||||
provider_priority,
|
||||
provider_is_active,
|
||||
endpoint_id,
|
||||
endpoint_api_format,
|
||||
endpoint_api_family,
|
||||
endpoint_kind,
|
||||
endpoint_is_active,
|
||||
key_id,
|
||||
key_name,
|
||||
key_auth_type,
|
||||
key_is_active,
|
||||
key_api_formats,
|
||||
key_allowed_models,
|
||||
key_capabilities,
|
||||
key_internal_priority,
|
||||
key_global_priority_by_format,
|
||||
model_id,
|
||||
global_model_id,
|
||||
global_model_name,
|
||||
global_model_mappings,
|
||||
global_model_supports_streaming,
|
||||
model_provider_model_name,
|
||||
model_provider_model_mappings,
|
||||
model_supports_streaming,
|
||||
model_is_active,
|
||||
model_is_available
|
||||
FROM selected_rows
|
||||
ORDER BY
|
||||
provider_priority ASC,
|
||||
key_internal_priority ASC,
|
||||
provider_id ASC,
|
||||
endpoint_id ASC,
|
||||
key_id ASC,
|
||||
model_id ASC
|
||||
"#;
|
||||
|
||||
const LIST_POOL_KEYS_FOR_GROUP_SQL: &str = r#"
|
||||
SELECT
|
||||
p.id AS provider_id,
|
||||
p.name AS provider_name,
|
||||
p.provider_type AS provider_type,
|
||||
p.provider_priority AS provider_priority,
|
||||
p.is_active AS provider_is_active,
|
||||
pe.id AS endpoint_id,
|
||||
pe.api_format AS endpoint_api_format,
|
||||
pe.api_family AS endpoint_api_family,
|
||||
pe.endpoint_kind AS endpoint_kind,
|
||||
pe.is_active AS endpoint_is_active,
|
||||
pak.id AS key_id,
|
||||
pak.name AS key_name,
|
||||
pak.auth_type AS key_auth_type,
|
||||
pak.is_active AS key_is_active,
|
||||
pak.api_formats AS key_api_formats,
|
||||
pak.allowed_models AS key_allowed_models,
|
||||
pak.capabilities AS key_capabilities,
|
||||
pak.internal_priority AS key_internal_priority,
|
||||
pak.global_priority_by_format AS key_global_priority_by_format,
|
||||
m.id AS model_id,
|
||||
m.global_model_id AS global_model_id,
|
||||
gm.name AS global_model_name,
|
||||
CASE
|
||||
WHEN gm.config IS NOT NULL THEN gm.config -> 'model_mappings'
|
||||
ELSE NULL
|
||||
END AS global_model_mappings,
|
||||
CASE
|
||||
WHEN gm.config IS NOT NULL AND gm.config ? 'streaming'
|
||||
THEN (gm.config ->> 'streaming')::BOOLEAN
|
||||
ELSE NULL
|
||||
END AS global_model_supports_streaming,
|
||||
m.provider_model_name AS model_provider_model_name,
|
||||
m.provider_model_mappings AS model_provider_model_mappings,
|
||||
m.supports_streaming AS model_supports_streaming,
|
||||
m.is_active AS model_is_active,
|
||||
m.is_available AS model_is_available
|
||||
FROM providers p
|
||||
INNER JOIN provider_endpoints pe
|
||||
ON pe.provider_id = p.id
|
||||
INNER JOIN provider_api_keys pak
|
||||
ON pak.provider_id = p.id
|
||||
INNER JOIN models m
|
||||
ON m.provider_id = p.id
|
||||
INNER JOIN global_models gm
|
||||
ON gm.id = m.global_model_id
|
||||
WHERE p.is_active = TRUE
|
||||
AND pe.is_active = TRUE
|
||||
AND pak.is_active = TRUE
|
||||
AND m.is_active = TRUE
|
||||
AND m.is_available = TRUE
|
||||
AND gm.is_active = TRUE
|
||||
AND LOWER(pe.api_format) = LOWER($1)
|
||||
AND p.id = $2
|
||||
AND pe.id = $3
|
||||
AND m.id = $4
|
||||
AND (
|
||||
pak.api_formats IS NULL
|
||||
OR EXISTS (
|
||||
SELECT 1
|
||||
FROM json_array_elements_text(pak.api_formats) AS fmt(value)
|
||||
WHERE LOWER(BTRIM(fmt.value)) = ANY($5::text[])
|
||||
)
|
||||
)
|
||||
AND (
|
||||
(
|
||||
LOWER(BTRIM(p.provider_type)) = 'codex'
|
||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||
AND LOWER($6) IN ('openai:responses', 'openai:responses:compact', 'openai:image')
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(p.provider_type)) = 'claude_code'
|
||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||
AND LOWER($6) = 'claude:messages'
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(p.provider_type)) = 'kiro'
|
||||
AND LOWER($6) = 'claude:messages'
|
||||
AND (
|
||||
LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||
OR (
|
||||
LOWER(BTRIM(pak.auth_type)) = 'bearer'
|
||||
AND pak.auth_config IS NOT NULL
|
||||
AND BTRIM(pak.auth_config) <> ''
|
||||
)
|
||||
)
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
||||
AND LOWER(BTRIM(pak.auth_type)) = 'oauth'
|
||||
AND LOWER($6) = 'gemini:generate_content'
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(p.provider_type)) = 'vertex_ai'
|
||||
AND (
|
||||
(
|
||||
LOWER(BTRIM(pak.auth_type)) = 'api_key'
|
||||
AND LOWER($6) = 'gemini:generate_content'
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(pak.auth_type)) IN ('service_account', 'vertex_ai')
|
||||
AND LOWER($6) IN ('claude:messages', 'gemini:generate_content')
|
||||
)
|
||||
)
|
||||
)
|
||||
OR (
|
||||
LOWER(BTRIM(p.provider_type)) NOT IN (
|
||||
'claude_code',
|
||||
'codex',
|
||||
'gemini_cli',
|
||||
'vertex_ai',
|
||||
'antigravity',
|
||||
'kiro'
|
||||
)
|
||||
AND LOWER(BTRIM(pak.auth_type)) <> 'oauth'
|
||||
)
|
||||
)
|
||||
ORDER BY
|
||||
p.provider_priority ASC,
|
||||
pak.internal_priority ASC,
|
||||
p.id ASC,
|
||||
pe.id ASC,
|
||||
pak.id ASC,
|
||||
m.id ASC
|
||||
pak.id ASC
|
||||
LIMIT $7
|
||||
OFFSET $8
|
||||
"#;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
@@ -336,6 +564,130 @@ impl SqlxMinimalCandidateSelectionReadRepository {
|
||||
}
|
||||
Ok(dedupe_candidate_selection_rows(rows))
|
||||
}
|
||||
|
||||
pub async fn list_for_exact_api_format_and_requested_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let mut rows = Vec::new();
|
||||
let canonical_api_format = normalize_api_format(api_format);
|
||||
let storage_aliases = api_format_aliases(&canonical_api_format);
|
||||
let sql_match_aliases = sql_match_aliases(&storage_aliases);
|
||||
let sql = requested_model_selection_sql();
|
||||
for api_format in storage_aliases {
|
||||
rows.extend(
|
||||
Self::collect_query_rows(
|
||||
sqlx::query(sql.as_str())
|
||||
.bind(api_format)
|
||||
.bind(requested_model_name)
|
||||
.bind(sql_match_aliases.clone())
|
||||
.bind(canonical_api_format.clone())
|
||||
.fetch(&self.pool),
|
||||
map_candidate_selection_row,
|
||||
)
|
||||
.await?,
|
||||
);
|
||||
}
|
||||
Ok(dedupe_candidate_selection_rows(rows))
|
||||
}
|
||||
|
||||
pub async fn list_for_exact_api_format_and_requested_model_page(
|
||||
&self,
|
||||
query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let mut rows = Vec::new();
|
||||
let canonical_api_format = normalize_api_format(&query.api_format);
|
||||
let storage_aliases = api_format_aliases(&canonical_api_format);
|
||||
let sql_match_aliases = sql_match_aliases(&storage_aliases);
|
||||
let limit = i64::from(query.limit.max(1));
|
||||
let offset = i64::from(query.offset);
|
||||
let sql = requested_model_selection_page_sql();
|
||||
for api_format in storage_aliases {
|
||||
rows.extend(
|
||||
Self::collect_query_rows(
|
||||
sqlx::query(sql.as_str())
|
||||
.bind(api_format)
|
||||
.bind(query.requested_model_name.as_str())
|
||||
.bind(sql_match_aliases.clone())
|
||||
.bind(canonical_api_format.clone())
|
||||
.bind(limit)
|
||||
.bind(offset)
|
||||
.fetch(&self.pool),
|
||||
map_candidate_selection_row,
|
||||
)
|
||||
.await?,
|
||||
);
|
||||
}
|
||||
Ok(dedupe_candidate_selection_rows(rows))
|
||||
}
|
||||
|
||||
pub async fn list_pool_key_rows_for_group(
|
||||
&self,
|
||||
query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let mut rows = Vec::new();
|
||||
let canonical_api_format = normalize_api_format(&query.api_format);
|
||||
let storage_aliases = api_format_aliases(&canonical_api_format);
|
||||
let sql_match_aliases = sql_match_aliases(&storage_aliases);
|
||||
let limit = i64::from(query.limit.max(1));
|
||||
let offset = i64::from(query.offset);
|
||||
for api_format in storage_aliases {
|
||||
rows.extend(
|
||||
Self::collect_query_rows(
|
||||
sqlx::query(LIST_POOL_KEYS_FOR_GROUP_SQL)
|
||||
.bind(api_format)
|
||||
.bind(query.provider_id.as_str())
|
||||
.bind(query.endpoint_id.as_str())
|
||||
.bind(query.model_id.as_str())
|
||||
.bind(sql_match_aliases.clone())
|
||||
.bind(canonical_api_format.clone())
|
||||
.bind(limit)
|
||||
.bind(offset)
|
||||
.fetch(&self.pool),
|
||||
map_candidate_selection_row,
|
||||
)
|
||||
.await?,
|
||||
);
|
||||
}
|
||||
Ok(dedupe_candidate_selection_rows(rows))
|
||||
}
|
||||
}
|
||||
|
||||
fn requested_model_selection_sql() -> String {
|
||||
LIST_FOR_EXACT_API_FORMAT_AND_GLOBAL_MODEL_SQL
|
||||
.replace(
|
||||
"AND gm.name = $2",
|
||||
r#"AND (
|
||||
gm.name = $2
|
||||
OR m.provider_model_name = $2
|
||||
OR (
|
||||
jsonb_typeof(m.provider_model_mappings) = 'array'
|
||||
AND EXISTS (
|
||||
SELECT 1
|
||||
FROM jsonb_array_elements(m.provider_model_mappings) AS mapping(value)
|
||||
WHERE mapping.value ->> 'name' = $2
|
||||
AND (
|
||||
mapping.value -> 'api_formats' IS NULL
|
||||
OR jsonb_typeof(mapping.value -> 'api_formats') <> 'array'
|
||||
OR EXISTS (
|
||||
SELECT 1
|
||||
FROM jsonb_array_elements_text(mapping.value -> 'api_formats') AS fmt(value)
|
||||
WHERE LOWER(BTRIM(fmt.value)) = ANY($3::text[])
|
||||
)
|
||||
)
|
||||
)
|
||||
)
|
||||
)"#,
|
||||
)
|
||||
.replace(
|
||||
"ORDER BY\n provider_priority ASC,",
|
||||
"ORDER BY\n global_model_name ASC,\n provider_priority ASC,",
|
||||
)
|
||||
}
|
||||
|
||||
fn requested_model_selection_page_sql() -> String {
|
||||
format!("{}\nLIMIT $5\nOFFSET $6", requested_model_selection_sql())
|
||||
}
|
||||
|
||||
fn api_format_aliases(api_format: &str) -> Vec<String> {
|
||||
@@ -384,6 +736,29 @@ impl MinimalCandidateSelectionReadRepository for SqlxMinimalCandidateSelectionRe
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Self::list_for_exact_api_format_and_global_model(self, 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> {
|
||||
Self::list_for_exact_api_format_and_requested_model(self, api_format, requested_model_name)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model_page(
|
||||
&self,
|
||||
query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Self::list_for_exact_api_format_and_requested_model_page(self, query).await
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group(
|
||||
&self,
|
||||
query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Self::list_pool_key_rows_for_group(self, query).await
|
||||
}
|
||||
}
|
||||
|
||||
fn map_candidate_selection_row(
|
||||
@@ -630,8 +1005,8 @@ mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
parse_provider_model_mappings, parse_string_list,
|
||||
SqlxMinimalCandidateSelectionReadRepository,
|
||||
parse_provider_model_mappings, parse_string_list, requested_model_selection_page_sql,
|
||||
requested_model_selection_sql, SqlxMinimalCandidateSelectionReadRepository,
|
||||
};
|
||||
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
|
||||
use crate::repository::candidate_selection::StoredProviderModelMapping;
|
||||
@@ -655,6 +1030,25 @@ mod tests {
|
||||
let _ = repository.pool();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn requested_model_selection_sql_filters_before_row_materialization() {
|
||||
let sql = requested_model_selection_sql();
|
||||
|
||||
assert!(sql.contains("m.provider_model_name = $2"));
|
||||
assert!(sql.contains("jsonb_array_elements(m.provider_model_mappings)"));
|
||||
assert!(!sql.contains("json_typeof(m.provider_model_mappings)"));
|
||||
assert!(!sql.contains("json_array_elements_text(gm.config -> 'model_mappings')"));
|
||||
assert!(sql.contains("ORDER BY\n global_model_name ASC,"));
|
||||
assert!(!sql.contains("AND gm.name = $2\n AND"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn requested_model_selection_page_sql_adds_limit_and_offset() {
|
||||
let sql = requested_model_selection_page_sql();
|
||||
|
||||
assert!(sql.ends_with("LIMIT $5\nOFFSET $6"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_string_list_accepts_stringified_array() {
|
||||
let parsed = parse_string_list(
|
||||
|
||||
@@ -309,6 +309,29 @@ pub fn build_execution_request_candidate_seed(
|
||||
Value::String(plan.endpoint_id.clone()),
|
||||
);
|
||||
context.insert("key_id".to_string(), Value::String(plan.key_id.clone()));
|
||||
let mut extra_data = parse_request_candidate_report_context(Some(&Value::Object(
|
||||
context.clone(),
|
||||
)))
|
||||
.and_then(|metadata| {
|
||||
build_report_candidate_extra_data(ReportCandidateExtraDataInput {
|
||||
client_api_format: metadata.client_api_format,
|
||||
provider_api_format: metadata.provider_api_format,
|
||||
upstream_url: metadata.upstream_url,
|
||||
mapped_model: metadata.mapped_model,
|
||||
key_name: metadata.key_name,
|
||||
header_rules: metadata.header_rules,
|
||||
body_rules: metadata.body_rules,
|
||||
proxy: metadata.proxy,
|
||||
error_flow: metadata.error_flow,
|
||||
ranking_mode: metadata.ranking_mode,
|
||||
priority_mode: metadata.priority_mode,
|
||||
ranking_index: metadata.ranking_index,
|
||||
priority_slot: metadata.priority_slot,
|
||||
promoted_by: metadata.promoted_by,
|
||||
demoted_by: metadata.demoted_by,
|
||||
})
|
||||
});
|
||||
append_seed_extra_data_from_report_context(&mut extra_data, &context);
|
||||
|
||||
SchedulerExecutionRequestCandidateSeed {
|
||||
upsert_record: UpsertRequestCandidateRecord {
|
||||
@@ -331,7 +354,7 @@ pub fn build_execution_request_candidate_seed(
|
||||
error_message: None,
|
||||
latency_ms: None,
|
||||
concurrent_requests: None,
|
||||
extra_data: None,
|
||||
extra_data,
|
||||
required_capabilities: None,
|
||||
created_at_unix_ms: Some(started_at_unix_ms),
|
||||
started_at_unix_ms: Some(started_at_unix_ms),
|
||||
@@ -341,6 +364,33 @@ pub fn build_execution_request_candidate_seed(
|
||||
}
|
||||
}
|
||||
|
||||
fn append_seed_extra_data_from_report_context(
|
||||
extra_data: &mut Option<Value>,
|
||||
context: &Map<String, Value>,
|
||||
) {
|
||||
const PASSTHROUGH_FIELDS: &[&str] = &[
|
||||
"execution_strategy",
|
||||
"conversion_mode",
|
||||
"client_contract",
|
||||
"provider_contract",
|
||||
"transport_diagnostics",
|
||||
];
|
||||
|
||||
let mut object = extra_data
|
||||
.take()
|
||||
.and_then(|value| match value {
|
||||
Value::Object(object) => Some(object),
|
||||
_ => None,
|
||||
})
|
||||
.unwrap_or_default();
|
||||
for field in PASSTHROUGH_FIELDS {
|
||||
if let Some(value) = context.get(*field).filter(|value| !value.is_null()) {
|
||||
object.insert((*field).to_string(), value.clone());
|
||||
}
|
||||
}
|
||||
*extra_data = (!object.is_empty()).then_some(Value::Object(object));
|
||||
}
|
||||
|
||||
pub fn build_local_request_candidate_status_record(
|
||||
input: LocalRequestCandidateStatusRecordInput<'_>,
|
||||
) -> Option<UpsertRequestCandidateRecord> {
|
||||
|
||||
Reference in New Issue
Block a user