refactor: 抽离 AI pipeline 与调度共享能力逻辑

This commit is contained in:
fawney19
2026-04-10 01:46:14 +08:00
parent b901a6ffc7
commit 5014e2f5fd
255 changed files with 15057 additions and 3115 deletions
@@ -28,19 +28,27 @@ use std::collections::BTreeMap;
use crate::data::auth::GatewayAuthApiKeySnapshot;
use crate::data::candidate_selection::{
read_global_model_names_for_required_capability, MinimalCandidateSelectionRowSource,
read_global_model_names_for_api_format, read_global_model_names_for_required_capability,
MinimalCandidateSelectionRowSource,
};
use crate::GatewayError;
#[cfg_attr(not(test), allow(dead_code))]
const SCHEDULER_AFFINITY_MAX_ENTRIES: usize = 10_000;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum RequiredCapabilityMatchMode {
Compatible,
Exclusive,
}
pub(crate) async fn list_selectable_candidates(
selection_row_source: &(impl MinimalCandidateSelectionRowSource + Sync),
runtime_state: &impl SchedulerRuntimeState,
api_format: &str,
global_model_name: &str,
require_streaming: bool,
required_capabilities: Option<&serde_json::Value>,
auth_snapshot: Option<&GatewayAuthApiKeySnapshot>,
now_unix_secs: u64,
) -> Result<Vec<SchedulerMinimalCandidateSelectionCandidate>, GatewayError> {
@@ -50,6 +58,7 @@ pub(crate) async fn list_selectable_candidates(
api_format,
global_model_name,
require_streaming,
required_capabilities,
auth_snapshot,
now_unix_secs,
)
@@ -70,37 +79,87 @@ pub(crate) async fn list_selectable_candidates_for_required_capability_without_r
return Ok(Vec::new());
}
let model_names = read_global_model_names_for_required_capability(
selection_row_source,
&normalized_api_format,
required_capability,
require_streaming,
auth_snapshot,
)
.await
let capability_mode = required_capability_match_mode(required_capability);
let model_names = match capability_mode {
RequiredCapabilityMatchMode::Exclusive => {
read_global_model_names_for_required_capability(
selection_row_source,
&normalized_api_format,
required_capability,
require_streaming,
auth_snapshot,
)
.await
}
RequiredCapabilityMatchMode::Compatible => {
read_global_model_names_for_api_format(
selection_row_source,
&normalized_api_format,
require_streaming,
auth_snapshot,
)
.await
}
}
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let required_capabilities = build_required_capabilities_object(required_capability);
for global_model_name in model_names {
let candidates = list_selectable_candidates(
let mut candidates = list_selectable_candidates(
selection_row_source,
runtime_state,
&normalized_api_format,
&global_model_name,
require_streaming,
required_capabilities.as_ref(),
auth_snapshot,
now_unix_secs,
)
.await?;
let filtered = candidates
.into_iter()
.filter(|candidate| {
candidate_supports_required_capability(candidate, required_capability)
})
.collect::<Vec<_>>();
if !filtered.is_empty() {
return Ok(filtered);
match capability_mode {
RequiredCapabilityMatchMode::Exclusive => {
let filtered = candidates
.into_iter()
.filter(|candidate| {
candidate_supports_required_capability(candidate, required_capability)
})
.collect::<Vec<_>>();
if !filtered.is_empty() {
return Ok(filtered);
}
}
RequiredCapabilityMatchMode::Compatible => {
if candidates.is_empty() {
continue;
}
candidates.sort_by_key(|candidate| {
!candidate_supports_required_capability(candidate, required_capability)
});
return Ok(candidates);
}
}
}
Ok(Vec::new())
}
fn required_capability_match_mode(required_capability: &str) -> RequiredCapabilityMatchMode {
match required_capability.trim().to_ascii_lowercase().as_str() {
"cache_1h" | "context_1m" => RequiredCapabilityMatchMode::Compatible,
_ => RequiredCapabilityMatchMode::Exclusive,
}
}
fn build_required_capabilities_object(required_capability: &str) -> Option<serde_json::Value> {
let required_capability = required_capability.trim();
if required_capability.is_empty() {
return None;
}
let mut capabilities = serde_json::Map::new();
capabilities.insert(
required_capability.to_string(),
serde_json::Value::Bool(true),
);
Some(serde_json::Value::Object(capabilities))
}