feat: unify provider scheduling workspace

This commit is contained in:
elky
2026-10-07 00:34:18 +08:00
parent e7de935e61
commit 466c7918a1
69 changed files with 5810 additions and 2666 deletions
@@ -35,6 +35,7 @@ mod overview_fact_metadata;
mod overview_migration_safety;
mod policy_nulls;
mod provider_expenses;
mod scoped_provider_creation;
/// A clean PostgreSQL database is bootstrapped from the schema snapshot first;
/// migrations after the privacy/security frontier are intentionally left
@@ -0,0 +1,131 @@
use super::*;
use aether_data_contracts::repository::{
provider_catalog::{ProviderCatalogWriteRepository, StoredProviderCatalogProvider},
routing_profiles::{
CreateRoutingGroupRecord, RoutingGroupReadRepository, RoutingGroupWriteRepository,
UpdateRoutingGroupRecord,
},
};
use serde_json::json;
#[tokio::test]
async fn postgres_scoped_provider_creation_rolls_back_and_serializes_group_saves() {
let Some(server) = ManagedPostgresServer::try_start()
.await
.expect("local postgres should start or skip")
else {
return;
};
let pool = PgPool::connect(server.database_url()).await.unwrap();
prepare_and_apply_clean_postgres_database(&pool).await;
let groups =
crate::repository::routing_profiles::PostgresRoutingGroupRepository::new(pool.clone());
let providers =
crate::repository::provider_catalog::SqlxProviderCatalogReadRepository::new(pool.clone());
for id in ["selected", "other"] {
groups
.create_routing_group(CreateRoutingGroupRecord {
id: id.into(),
name: id.into(),
description: None,
enabled: true,
is_system_default: id == "selected",
sort_order: 0,
config_json: json!({"disabled_providers": ["existing-disabled"]}),
version: 1,
created_at: 1,
updated_at: 1,
published_at: None,
})
.await
.unwrap();
}
let provider =
StoredProviderCatalogProvider::new("new".into(), "new".into(), None, "custom".into())
.unwrap();
assert!(providers
.create_provider_in_routing_group(&provider, None, "missing")
.await
.is_err());
assert_eq!(
sqlx::query_scalar::<_, i64>("SELECT COUNT(*) FROM providers")
.fetch_one(&pool)
.await
.unwrap(),
0
);
assert!(groups
.list_routing_groups()
.await
.unwrap()
.iter()
.all(|group| group.version == 1));
providers
.create_provider_in_routing_group(&provider, None, "selected")
.await
.unwrap();
let before = groups.list_routing_groups().await.unwrap();
for group in &before {
assert_eq!(group.version, if group.id == "selected" { 1 } else { 2 });
assert_eq!(
group.config_json["disabled_providers"],
if group.id == "selected" {
json!(["existing-disabled"])
} else {
json!(["existing-disabled", "new"])
}
);
}
// The INSERT fails after group updates execute. Its transaction must undo
// every exclusion and version change along with any priority shifts.
assert!(providers
.create_provider_in_routing_group(&provider, Some(0), "other")
.await
.is_err());
assert_eq!(groups.list_routing_groups().await.unwrap(), before);
assert!(groups
.update_routing_group(
"other",
UpdateRoutingGroupRecord {
expected_version: Some(1),
config_json: Some(json!({"disabled_providers": []})),
..Default::default()
}
)
.await
.is_err());
assert_eq!(groups.list_routing_groups().await.unwrap(), before);
let concurrent_provider = StoredProviderCatalogProvider::new(
"concurrent".into(),
"concurrent".into(),
None,
"custom".into(),
)
.unwrap();
let (created, edited) = tokio::join!(
providers.create_provider_in_routing_group(&concurrent_provider, None, "selected"),
groups.update_routing_group(
"other",
UpdateRoutingGroupRecord {
expected_version: Some(2),
config_json: Some(json!({"disabled_providers": ["new"]})),
..Default::default()
}
)
);
created.unwrap();
let other = groups
.list_routing_groups()
.await
.unwrap()
.into_iter()
.find(|group| group.id == "other")
.unwrap();
assert!(other.config_json["disabled_providers"]
.as_array()
.unwrap()
.contains(&json!("concurrent")));
assert_eq!(other.version, if edited.is_ok() { 4 } else { 3 });
}
@@ -1,5 +1,5 @@
use std::collections::BTreeMap;
use std::sync::RwLock;
use std::sync::{Arc, RwLock};
use std::time::{SystemTime, UNIX_EPOCH};
use async_trait::async_trait;
@@ -31,6 +31,8 @@ struct MemoryProviderCatalogIndex {
#[derive(Debug, Default)]
pub struct InMemoryProviderCatalogReadRepository {
index: RwLock<MemoryProviderCatalogIndex>,
routing_groups:
Option<Arc<crate::repository::routing_profiles::InMemoryRoutingGroupRepository>>,
}
impl InMemoryProviderCatalogReadRepository {
@@ -40,6 +42,7 @@ impl InMemoryProviderCatalogReadRepository {
keys: Vec<StoredProviderCatalogKey>,
) -> Self {
Self {
routing_groups: None,
index: RwLock::new(MemoryProviderCatalogIndex {
providers: providers
.into_iter()
@@ -54,6 +57,14 @@ impl InMemoryProviderCatalogReadRepository {
}
}
pub fn with_routing_groups(
mut self,
repository: Arc<crate::repository::routing_profiles::InMemoryRoutingGroupRepository>,
) -> Self {
self.routing_groups = Some(repository);
self
}
fn snapshot(&self) -> ProviderCatalogSnapshot {
let index = self.index.read().expect("provider catalog repository lock");
ProviderCatalogSnapshot::new(
@@ -416,6 +427,46 @@ impl ProviderCatalogReadRepository for InMemoryProviderCatalogReadRepository {
#[async_trait]
impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository {
async fn create_provider_in_routing_group(
&self,
provider: &StoredProviderCatalogProvider,
shift_existing_priorities_from: Option<i32>,
routing_group_id: &str,
) -> Result<StoredProviderCatalogProvider, DataLayerError> {
let groups = self.routing_groups.as_ref().ok_or_else(|| {
DataLayerError::InvalidConfiguration(
"atomic provider creation requires a shared routing group repository".to_string(),
)
})?;
groups.create_scoped_provider(routing_group_id, &provider.id, || {
let mut index = self
.index
.write()
.expect("provider catalog repository lock");
if index.providers.contains_key(&provider.id)
|| index
.providers
.values()
.any(|existing| existing.name == provider.name)
{
return Err(DataLayerError::InvalidInput(
"provider already exists".to_string(),
));
}
if let Some(target_priority) = shift_existing_priorities_from {
for existing in index.providers.values_mut() {
if existing.provider_priority >= target_priority {
existing.provider_priority += 1;
}
}
}
index
.providers
.insert(provider.id.clone(), provider.clone());
Ok(provider.clone())
})
}
async fn create_provider(
&self,
provider: &StoredProviderCatalogProvider,
@@ -1587,6 +1638,102 @@ mod tests {
.expect("key should build")
}
#[tokio::test]
async fn scoped_provider_creation_is_atomic_and_rejects_stale_group_updates() {
use crate::repository::routing_profiles::InMemoryRoutingGroupRepository;
use aether_data_contracts::repository::routing_profiles::{
CreateRoutingGroupRecord, RoutingGroupReadRepository, RoutingGroupWriteRepository,
UpdateRoutingGroupRecord,
};
let groups = Arc::new(InMemoryRoutingGroupRepository::default());
for id in ["selected", "other", "disabled"] {
groups
.create_routing_group(CreateRoutingGroupRecord {
id: id.into(),
name: id.into(),
description: None,
enabled: id != "disabled",
is_system_default: id == "selected",
sort_order: 0,
config_json: json!({"disabled_providers": ["already-disabled"], "rules": []}),
version: 1,
created_at: 1,
updated_at: 1,
published_at: None,
})
.await
.unwrap();
}
let repository =
InMemoryProviderCatalogReadRepository::default().with_routing_groups(groups.clone());
let provider = sample_provider("new");
assert!(repository
.create_provider_in_routing_group(&provider, Some(0), "missing")
.await
.is_err());
assert!(repository.list_providers(false).await.unwrap().is_empty());
assert!(groups
.list_routing_groups()
.await
.unwrap()
.iter()
.all(|group| group.version == 1));
repository
.create_provider_in_routing_group(&provider, None, "selected")
.await
.unwrap();
for group in groups.list_routing_groups().await.unwrap() {
assert_eq!(group.version, if group.id == "selected" { 1 } else { 2 });
assert_eq!(
group.config_json["disabled_providers"],
if group.id == "selected" {
json!(["already-disabled"])
} else {
json!(["already-disabled", "new"])
}
);
}
let before = groups.list_routing_groups().await.unwrap();
assert!(repository
.create_provider_in_routing_group(&provider, Some(0), "other")
.await
.is_err());
assert_eq!(groups.list_routing_groups().await.unwrap(), before);
let stale = groups
.update_routing_group(
"other",
UpdateRoutingGroupRecord {
expected_version: Some(1),
config_json: Some(json!({"disabled_providers": []})),
..Default::default()
},
)
.await;
assert!(
matches!(stale, Err(DataLayerError::InvalidInput(message)) if message == "routing_group_version_conflict")
);
assert_eq!(groups.list_routing_groups().await.unwrap(), before);
let updated = groups
.update_routing_group(
"other",
UpdateRoutingGroupRecord {
expected_version: Some(2),
config_json: Some(json!({"disabled_providers": ["already-disabled"]})),
version: Some(2),
..Default::default()
},
)
.await
.unwrap()
.unwrap();
assert_eq!(updated.version, 3);
assert_eq!(
updated.config_json["disabled_providers"],
json!(["already-disabled"])
);
}
#[tokio::test]
async fn reads_provider_catalog_items_by_id() {
let repository = InMemoryProviderCatalogReadRepository::seed(
@@ -20,6 +20,58 @@ pub struct InMemoryRoutingGroupRepository {
}
impl InMemoryRoutingGroupRepository {
pub(crate) fn create_scoped_provider<T>(
&self,
selected_group_id: &str,
provider_id: &str,
create: impl FnOnce() -> Result<T, DataLayerError>,
) -> Result<T, DataLayerError> {
let mut groups = self.groups.write().expect("routing group repository lock");
if !groups.contains_key(selected_group_id) {
return Err(DataLayerError::InvalidInput(
"routing_group_not_found".to_string(),
));
}
let mut updated = groups.clone();
for group in updated.values_mut() {
let had_provider = group
.config_json
.get("disabled_providers")
.and_then(serde_json::Value::as_array)
.is_some_and(|disabled| {
disabled
.iter()
.any(|value| value.as_str() == Some(provider_id))
});
if (group.id == selected_group_id && !had_provider)
|| (group.id != selected_group_id && had_provider)
{
continue;
}
let object = group.config_json.as_object_mut().ok_or_else(|| {
DataLayerError::InvalidInput("routing group config must be an object".to_string())
})?;
let disabled = object
.entry("disabled_providers")
.or_insert_with(|| serde_json::json!([]));
let disabled = disabled.as_array_mut().ok_or_else(|| {
DataLayerError::InvalidInput("disabled_providers must be an array".to_string())
})?;
disabled.retain(|value| value.as_str() != Some(provider_id));
if group.id != selected_group_id {
disabled.push(serde_json::json!(provider_id));
}
group.version = group.version.saturating_add(1);
group.updated_at = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs() as i64;
}
let created = create()?;
*groups = updated;
Ok(created)
}
pub fn seed<I, B, V>(groups: I, bindings: B, versions: V) -> Self
where
I: IntoIterator<Item = StoredRoutingGroup>,