mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-10 03:09:50 +08:00
refactor: 抽离 AI pipeline 与调度共享能力逻辑
This commit is contained in:
@@ -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))
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user