Fix admin pool sorting and OAuth refresh

This commit is contained in:
fawney19
2026-05-13 09:16:52 +08:00
parent 4387a9cdd5
commit 3c2497f019
15 changed files with 498 additions and 195 deletions

View File

@@ -15,7 +15,7 @@ use aether_data::repository::oauth_providers::{
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository;
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
ProviderCatalogReadRepository, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint,
};
use axum::body::{to_bytes, Body, Bytes};
use axum::response::{IntoResponse, Response};
@@ -5267,6 +5267,39 @@ fn gateway_manual_kiro_oauth_refresh_reconciles_missing_fixed_endpoint() {
}
async fn gateway_manual_kiro_oauth_refresh_reconciles_missing_fixed_endpoint_impl() {
run_gateway_manual_kiro_oauth_refresh_maintenance_endpoint_test(None, None, true).await;
}
#[test]
fn gateway_manual_kiro_oauth_refresh_uses_disabled_fixed_endpoint_for_maintenance() {
run_manual_kiro_oauth_refresh_test(
"gateway_manual_kiro_oauth_refresh_uses_disabled_fixed_endpoint_for_maintenance",
gateway_manual_kiro_oauth_refresh_uses_disabled_fixed_endpoint_for_maintenance_impl,
);
}
async fn gateway_manual_kiro_oauth_refresh_uses_disabled_fixed_endpoint_for_maintenance_impl() {
let mut endpoint = sample_endpoint(
"endpoint-kiro-disabled-maintenance",
"provider-kiro-oauth-refresh",
"claude:messages",
"https://q.{region}.amazonaws.com",
);
endpoint.is_active = false;
run_gateway_manual_kiro_oauth_refresh_maintenance_endpoint_test(
Some(endpoint),
Some("endpoint-kiro-disabled-maintenance"),
false,
)
.await;
}
async fn run_gateway_manual_kiro_oauth_refresh_maintenance_endpoint_test(
initial_endpoint: Option<StoredProviderCatalogEndpoint>,
expected_endpoint_id: Option<&str>,
expected_endpoint_active: bool,
) {
let refreshed_access_token = sample_kiro_device_access_token("kiro-refresh@example.com");
let expected_access_token = refreshed_access_token.clone();
let refreshed_refresh_token = "s".repeat(120);
@@ -5362,7 +5395,7 @@ async fn gateway_manual_kiro_oauth_refresh_reconciles_missing_fixed_endpoint_imp
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
Vec::new(),
initial_endpoint.into_iter().collect(),
vec![key],
));
@@ -5421,6 +5454,10 @@ async fn gateway_manual_kiro_oauth_refresh_reconciles_missing_fixed_endpoint_imp
assert_eq!(endpoints.len(), 1);
assert_eq!(endpoints[0].api_format, "claude:messages");
assert_eq!(endpoints[0].base_url, "https://q.{region}.amazonaws.com");
assert_eq!(endpoints[0].is_active, expected_endpoint_active);
if let Some(expected_endpoint_id) = expected_endpoint_id {
assert_eq!(endpoints[0].id, expected_endpoint_id);
}
assert_eq!(
*seen_endpoint_id.lock().expect("mutex should lock"),
Some(endpoints[0].id.clone())

View File

@@ -77,6 +77,36 @@ async fn local_admin_pool_response(
.expect("pool route should resolve locally")
}
fn sample_pool_member_score(provider_id: &str, key_id: &str, score: f64) -> StoredPoolMemberScore {
let score_scope = provider_key_pool_score_scope();
let score_identity = PoolMemberIdentity::provider_api_key(provider_id, key_id);
StoredPoolMemberScore {
id: provider_key_pool_score_id(&score_identity, &score_scope),
pool_kind: score_identity.pool_kind.clone(),
pool_id: score_identity.pool_id.clone(),
member_kind: score_identity.member_kind.clone(),
member_id: score_identity.member_id.clone(),
capability: score_scope.capability.clone(),
scope_kind: score_scope.scope_kind.clone(),
scope_id: score_scope.scope_id.clone(),
score,
hard_state: PoolMemberHardState::Available,
score_version: 1,
score_reason: json!({ "weights": { "manual_priority": score } }),
last_ranked_at: Some(1_700_000_000),
last_scheduled_at: None,
last_success_at: None,
last_failure_at: None,
failure_count: 0,
last_probe_attempt_at: None,
last_probe_success_at: None,
last_probe_failure_at: None,
probe_failure_count: 0,
probe_status: PoolMemberProbeStatus::Ok,
updated_at: 1_700_000_050,
}
}
fn sample_provider_key_usage_row(
id: &str,
request_id: &str,
@@ -928,11 +958,16 @@ async fn gateway_sorts_admin_pool_keys_by_imported_and_last_used_time() {
Vec::new(),
vec![old_key, fresh_key, active_key],
));
let pool_score_repository = Arc::new(InMemoryPoolMemberScoreRepository::seed(vec![
sample_pool_member_score("provider-openai", "key-openai-fresh", 0.35),
sample_pool_member_score("provider-openai", "key-openai-active", 0.92),
]));
let state = AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(GatewayDataState::with_provider_catalog_reader_for_tests(
provider_catalog_repository,
));
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_reader_for_tests(provider_catalog_repository)
.with_pool_score_repository_for_tests(pool_score_repository),
);
let default_response = local_admin_pool_response(
&state,
@@ -999,6 +1034,30 @@ async fn gateway_sorts_admin_pool_keys_by_imported_and_last_used_time() {
.map(|item| item["key_name"].as_str().unwrap_or_default())
.collect::<Vec<_>>();
assert_eq!(last_used_names, vec!["active", "old", "fresh"]);
let score_response = local_admin_pool_response(
&state,
http::Method::GET,
"/api/admin/pool/provider-openai/keys?page=1&page_size=50&status=all&sort_by=score&sort_order=desc",
None,
)
.await;
assert_eq!(score_response.status(), StatusCode::OK);
let score_payload: serde_json::Value = serde_json::from_slice(
&to_bytes(score_response.into_body(), usize::MAX)
.await
.expect("body should read"),
)
.expect("json body should parse");
let score_names = score_payload["keys"]
.as_array()
.expect("keys should be array")
.iter()
.map(|item| item["key_name"].as_str().unwrap_or_default())
.collect::<Vec<_>>();
assert_eq!(score_names, vec!["active", "fresh", "old"]);
assert_eq!(score_payload["keys"][0]["pool_score"]["score"], json!(0.92));
assert!(score_payload["keys"][2]["pool_score"].is_null());
}
#[tokio::test]