mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 10:27:46 +08:00
feat: unify provider scheduling workspace
This commit is contained in:
@@ -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>,
|
||||
|
||||
Reference in New Issue
Block a user