refactor: lazy pool key scheduling

This commit is contained in:
fawney19
2026-05-11 00:12:05 +08:00
parent 1a0f1a7b72
commit bacb14e5f0
26 changed files with 1856 additions and 223 deletions

View File

@@ -206,6 +206,10 @@ fn schedule_pool_group<Candidate>(
.unwrap_or_default();
let active_presets =
normalize_enabled_pool_presets(&pool_config.scheduling_presets, provider_type.as_str());
let lru_distribution_enabled = pool_config.lru_enabled
&& !active_presets
.iter()
.any(|preset| pool_preset_mutex_group(&preset.preset).is_some());
let mut available = Vec::new();
let mut skipped = Vec::new();
@@ -285,7 +289,7 @@ fn schedule_pool_group<Candidate>(
let sort_vectors = build_pool_sort_vectors(
&available,
&active_presets,
pool_config.lru_enabled,
lru_distribution_enabled,
group_sort_seed(
provider_type.as_str(),
available.first().map(|item| &item.item.facts),
@@ -300,7 +304,7 @@ fn schedule_pool_group<Candidate>(
.cmp(&sort_vectors.get(&right.item.facts.key_id))
.then(left.original_index.cmp(&right.original_index))
});
} else if pool_config.lru_enabled {
} else if lru_distribution_enabled {
let lru_ranks = lru_rank_indices(&available, false);
available.sort_by(|left, right| {
lru_ranks
@@ -359,6 +363,16 @@ fn build_pool_sort_vectors<Candidate>(
let lru_ranks = lru_rank_indices(items, false);
let cache_affinity_ranks = lru_rank_indices(items, true);
if lru_enabled {
for item in items {
let key_id = item.item.facts.key_id.clone();
vectors
.entry(key_id.clone())
.or_default()
.push(*lru_ranks.get(&key_id).unwrap_or(&0));
}
}
for preset in presets {
let ranks = match preset.preset.as_str() {
"cache_affinity" => cache_affinity_ranks.clone(),
@@ -385,16 +399,6 @@ fn build_pool_sort_vectors<Candidate>(
}
}
if lru_enabled {
for item in items {
let key_id = item.item.facts.key_id.clone();
vectors
.entry(key_id.clone())
.or_default()
.push(*lru_ranks.get(&key_id).unwrap_or(&0));
}
}
vectors
}
@@ -408,13 +412,13 @@ fn lru_rank_indices<Candidate>(
fn priority_first_ranks<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
lru_ranks: &BTreeMap<String, usize>,
_lru_ranks: &BTreeMap<String, usize>,
) -> BTreeMap<String, usize> {
let scores = collect_metric_scores(items, |item| {
Some(f64::from(item.item.facts.key_internal_priority))
});
if !score_map_has_variation(&scores) {
return lru_ranks.clone();
return neutral_rank_indices(items);
}
rank_indices_from_score_map(items, &scores, false)
}
@@ -422,27 +426,35 @@ fn priority_first_ranks<Candidate>(
fn single_account_ranks<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
) -> BTreeMap<String, usize> {
let n = items.len().saturating_sub(1).max(1) as f64;
let priority_scores = collect_metric_scores(items, |item| {
Some(f64::from(item.item.facts.key_internal_priority))
});
let priority_ranks = rank_indices_from_score_map(items, &priority_scores, false);
let lru_desc_ranks = lru_rank_indices(items, true);
let combined_scores = items
let mut decorated = items
.iter()
.map(|item| {
let key_id = item.item.facts.key_id.clone();
let priority_rank = *priority_ranks.get(&key_id).unwrap_or(&0) as f64 / n;
let lru_rank = *lru_desc_ranks.get(&key_id).unwrap_or(&0) as f64 / n;
(key_id, Some((priority_rank * 0.75) + (lru_rank * 0.25)))
(
item.item.facts.key_internal_priority,
*lru_desc_ranks.get(&key_id).unwrap_or(&0),
item.original_index,
key_id,
)
})
.collect::<BTreeMap<_, _>>();
rank_indices_from_score_map(items, &combined_scores, false)
.collect::<Vec<_>>();
decorated.sort_by(|left, right| {
left.0
.cmp(&right.0)
.then(left.1.cmp(&right.1))
.then(left.2.cmp(&right.2))
});
decorated
.into_iter()
.enumerate()
.map(|(rank, (_, _, _, key_id))| (key_id, rank))
.collect()
}
fn plan_ranks<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
lru_ranks: &BTreeMap<String, usize>,
_lru_ranks: &BTreeMap<String, usize>,
mode: Option<&str>,
) -> BTreeMap<String, usize> {
let scores = items
@@ -458,14 +470,14 @@ fn plan_ranks<Candidate>(
})
.collect::<BTreeMap<_, _>>();
if !score_map_has_variation(&scores) {
return lru_ranks.clone();
return neutral_rank_indices(items);
}
rank_indices_from_score_map(items, &scores, false)
}
fn health_first_ranks<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
lru_ranks: &BTreeMap<String, usize>,
_lru_ranks: &BTreeMap<String, usize>,
) -> BTreeMap<String, usize> {
let scores = collect_metric_scores(items, |item| {
item.item
@@ -474,39 +486,39 @@ fn health_first_ranks<Candidate>(
.map(|score| 1.0 - score.clamp(0.0, 1.0))
});
if !score_map_has_signal(&scores) {
return lru_ranks.clone();
return neutral_rank_indices(items);
}
rank_indices_from_score_map(items, &scores, false)
}
fn latency_first_ranks<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
lru_ranks: &BTreeMap<String, usize>,
_lru_ranks: &BTreeMap<String, usize>,
) -> BTreeMap<String, usize> {
let scores = collect_metric_scores(items, |item| item.item.key_context.latency_avg_ms);
if !score_map_has_signal(&scores) {
return lru_ranks.clone();
return neutral_rank_indices(items);
}
rank_indices_from_score_map(items, &scores, false)
}
fn cost_first_ranks<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
lru_ranks: &BTreeMap<String, usize>,
_lru_ranks: &BTreeMap<String, usize>,
cost_limit_per_key_tokens: Option<u64>,
) -> BTreeMap<String, usize> {
let scores = collect_metric_scores(items, |item| {
cost_penalty(item, cost_limit_per_key_tokens).or(item.item.key_context.quota_usage_ratio)
});
if !score_map_has_signal(&scores) {
return lru_ranks.clone();
return neutral_rank_indices(items);
}
rank_indices_from_score_map(items, &scores, false)
}
fn quota_balanced_ranks<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
lru_ranks: &BTreeMap<String, usize>,
_lru_ranks: &BTreeMap<String, usize>,
cost_limit_per_key_tokens: Option<u64>,
) -> BTreeMap<String, usize> {
let scores = collect_metric_scores(items, |item| {
@@ -516,18 +528,18 @@ fn quota_balanced_ranks<Candidate>(
.or_else(|| cost_penalty(item, cost_limit_per_key_tokens))
});
if !score_map_has_signal(&scores) {
return lru_ranks.clone();
return neutral_rank_indices(items);
}
rank_indices_from_score_map(items, &scores, false)
}
fn recent_refresh_ranks<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
lru_ranks: &BTreeMap<String, usize>,
_lru_ranks: &BTreeMap<String, usize>,
) -> BTreeMap<String, usize> {
let scores = collect_metric_scores(items, |item| item.item.key_context.quota_reset_seconds);
if !score_map_has_signal(&scores) {
return lru_ranks.clone();
return neutral_rank_indices(items);
}
rank_indices_from_score_map(items, &scores, false)
}
@@ -646,6 +658,15 @@ fn rank_indices_from_score_map<Candidate>(
.collect()
}
fn neutral_rank_indices<Candidate>(
items: &[PoolGroupCandidateOrdering<Candidate>],
) -> BTreeMap<String, usize> {
items
.iter()
.map(|item| (item.item.facts.key_id.clone(), 0))
.collect()
}
fn cost_penalty<Candidate>(
item: &PoolGroupCandidateOrdering<Candidate>,
cost_limit_per_key_tokens: Option<u64>,
@@ -740,47 +761,42 @@ fn normalize_enabled_pool_presets(
entries.push((entries.len(), "recent_refresh".to_string(), true, None));
}
let mut group_anchor_index = BTreeMap::<String, usize>::new();
for (index, preset, _, _) in &entries {
let Some(mutex_group) = pool_preset_mutex_group(preset) else {
continue;
};
group_anchor_index
.entry(mutex_group.to_string())
.or_insert(*index);
}
let mut ordered_enabled = Vec::<(usize, usize, String, Option<String>)>::new();
let mut group_enabled = BTreeMap::<String, (usize, usize, String, Option<String>)>::new();
let mut distribution_mode = None::<(usize, String, Option<String>)>;
let mut strategy_presets = Vec::<(usize, String, Option<String>)>::new();
for (index, preset, enabled, mode) in entries {
if !enabled
|| preset == "lru"
|| !pool_preset_supported_for_provider(&preset, &provider_type)
{
if !enabled || !pool_preset_supported_for_provider(&preset, &provider_type) {
continue;
}
let Some(mutex_group) = pool_preset_mutex_group(&preset) else {
ordered_enabled.push((index, index, preset, mode));
strategy_presets.push((index, preset, mode));
continue;
};
let anchor = group_anchor_index
.get(mutex_group)
.copied()
.unwrap_or(index);
let existing = group_enabled.get(mutex_group);
if existing.is_none_or(|current| index < current.1) {
group_enabled.insert(mutex_group.to_string(), (anchor, index, preset, mode));
if mutex_group == "distribution_mode"
&& distribution_mode
.as_ref()
.is_none_or(|current| index < current.0)
{
distribution_mode = Some((index, preset, mode));
}
}
ordered_enabled.extend(group_enabled.into_values());
ordered_enabled.sort_by(|left, right| left.0.cmp(&right.0).then(left.1.cmp(&right.1)));
ordered_enabled
.into_iter()
.map(|(_, _, preset, mode)| NormalizedPoolPreset { preset, mode })
.collect()
let mut normalized = Vec::new();
if let Some((_, preset, mode)) = distribution_mode.filter(|(_, preset, _)| preset != "lru") {
normalized.push(NormalizedPoolPreset { preset, mode });
}
strategy_presets.sort_by_key(|left| left.0);
normalized.extend(
strategy_presets
.into_iter()
.map(|(_, preset, mode)| NormalizedPoolPreset { preset, mode }),
);
normalized
}
fn pool_preset_supported_for_provider(preset: &str, provider_type: &str) -> bool {
@@ -959,7 +975,193 @@ mod tests {
}
#[test]
fn normalizes_distribution_mutex_group_to_first_enabled_member() {
fn pool_scheduler_applies_distribution_mode_before_strategy_presets() {
let key_cache_hit =
sample_candidate("provider-pool", "endpoint-1", "key-cache-hit", 50, true)
.with_presets(vec![
AiPoolSchedulingPreset {
preset: "cache_affinity".to_string(),
enabled: true,
mode: None,
},
AiPoolSchedulingPreset {
preset: "priority_first".to_string(),
enabled: true,
mode: None,
},
]);
let key_high_priority =
sample_candidate("provider-pool", "endpoint-1", "key-high-priority", 10, true)
.with_presets(vec![
AiPoolSchedulingPreset {
preset: "cache_affinity".to_string(),
enabled: true,
mode: None,
},
AiPoolSchedulingPreset {
preset: "priority_first".to_string(),
enabled: true,
mode: None,
},
]);
let runtime_by_provider = BTreeMap::from([(
"provider-pool".to_string(),
AiPoolRuntimeState {
lru_score_by_key: BTreeMap::from([
("key-cache-hit".to_string(), 200.0),
("key-high-priority".to_string(), 10.0),
]),
..AiPoolRuntimeState::default()
},
)]);
let outcome = run_ai_pool_scheduler(
vec![key_cache_hit, key_high_priority],
&runtime_by_provider,
"seed",
);
assert!(outcome.skipped_candidates.is_empty());
assert_eq!(
outcome
.candidates
.iter()
.map(|item| item.candidate.as_str())
.collect::<Vec<_>>(),
vec!["key-cache-hit", "key-high-priority"]
);
}
#[test]
fn load_balance_distribution_is_not_overridden_by_priority_strategy() {
let key_random_first =
sample_candidate("provider-pool", "endpoint-1", "key-random-first", 50, true)
.with_presets(vec![
AiPoolSchedulingPreset {
preset: "load_balance".to_string(),
enabled: true,
mode: None,
},
AiPoolSchedulingPreset {
preset: "priority_first".to_string(),
enabled: true,
mode: None,
},
]);
let key_high_priority =
sample_candidate("provider-pool", "endpoint-1", "key-high-priority", 10, true)
.with_presets(vec![
AiPoolSchedulingPreset {
preset: "load_balance".to_string(),
enabled: true,
mode: None,
},
AiPoolSchedulingPreset {
preset: "priority_first".to_string(),
enabled: true,
mode: None,
},
]);
let nonce = (0..1000)
.map(|index| format!("seed-{index}"))
.find(|nonce| {
let group_seed = format!("codex:provider-pool:endpoint-1:model-1:gpt-5:{nonce}");
stable_hash_score(format!("{group_seed}:key-random-first").as_str())
< stable_hash_score(format!("{group_seed}:key-high-priority").as_str())
})
.expect("test seed should exist");
let outcome = run_ai_pool_scheduler(
vec![key_random_first, key_high_priority],
&BTreeMap::new(),
nonce.as_str(),
);
assert!(outcome.skipped_candidates.is_empty());
assert_eq!(
outcome
.candidates
.iter()
.map(|item| item.candidate.as_str())
.collect::<Vec<_>>(),
vec!["key-random-first", "key-high-priority"]
);
}
#[test]
fn single_account_distribution_orders_by_priority_then_reverse_lru() {
let key_priority_old =
sample_candidate("provider-pool", "endpoint-1", "key-priority-old", 10, true)
.with_presets(vec![AiPoolSchedulingPreset {
preset: "single_account".to_string(),
enabled: true,
mode: None,
}]);
let key_priority_recent = sample_candidate(
"provider-pool",
"endpoint-1",
"key-priority-recent",
10,
true,
)
.with_presets(vec![AiPoolSchedulingPreset {
preset: "single_account".to_string(),
enabled: true,
mode: None,
}]);
let key_lower_priority_recent = sample_candidate(
"provider-pool",
"endpoint-1",
"key-lower-priority-recent",
50,
true,
)
.with_presets(vec![AiPoolSchedulingPreset {
preset: "single_account".to_string(),
enabled: true,
mode: None,
}]);
let runtime_by_provider = BTreeMap::from([(
"provider-pool".to_string(),
AiPoolRuntimeState {
lru_score_by_key: BTreeMap::from([
("key-priority-old".to_string(), 10.0),
("key-priority-recent".to_string(), 200.0),
("key-lower-priority-recent".to_string(), 500.0),
]),
..AiPoolRuntimeState::default()
},
)]);
let outcome = run_ai_pool_scheduler(
vec![
key_priority_old,
key_lower_priority_recent,
key_priority_recent,
],
&runtime_by_provider,
"seed",
);
assert!(outcome.skipped_candidates.is_empty());
assert_eq!(
outcome
.candidates
.iter()
.map(|item| item.candidate.as_str())
.collect::<Vec<_>>(),
vec![
"key-priority-recent",
"key-priority-old",
"key-lower-priority-recent"
]
);
}
#[test]
fn normalizes_distribution_mode_before_strategy_presets() {
let presets = normalize_enabled_ai_pool_presets(
&[
AiPoolSchedulingPreset {
@@ -989,6 +1191,32 @@ mod tests {
assert_eq!(presets, ["single_account", "priority_first"]);
}
#[test]
fn normalizes_lru_as_mutually_exclusive_distribution_mode() {
let presets = normalize_enabled_ai_pool_presets(
&[
AiPoolSchedulingPreset {
preset: "lru".to_string(),
enabled: true,
mode: None,
},
AiPoolSchedulingPreset {
preset: "cache_affinity".to_string(),
enabled: true,
mode: None,
},
AiPoolSchedulingPreset {
preset: "priority_first".to_string(),
enabled: true,
mode: None,
},
],
"openai",
);
assert_eq!(presets, ["priority_first"]);
}
fn sample_candidate(
provider_id: &str,
endpoint_id: &str,