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
@@ -1179,6 +1179,16 @@ WHERE id = $1
&self,
provider: &StoredProviderCatalogProvider,
shift_existing_priorities_from: Option<i32>,
) -> Result<StoredProviderCatalogProvider, DataLayerError> {
self.create_provider_with_routing_group(provider, shift_existing_priorities_from, None)
.await
}
async fn create_provider_with_routing_group(
&self,
provider: &StoredProviderCatalogProvider,
shift_existing_priorities_from: Option<i32>,
routing_group_id: Option<&str>,
) -> Result<StoredProviderCatalogProvider, DataLayerError> {
if provider.id.trim().is_empty() {
return Err(DataLayerError::InvalidInput(
@@ -1208,6 +1218,40 @@ WHERE id = $1
let mut tx = self.pool.begin().await.map_postgres_err()?;
if let Some(group_id) = routing_group_id {
// Group edits use the same lock: validation, exclusions, and provider
// creation are committed together, including concurrent deletions.
sqlx::query("LOCK TABLE routing_groups IN SHARE ROW EXCLUSIVE MODE")
.execute(&mut *tx)
.await
.map_postgres_err()?;
let exists: bool =
sqlx::query_scalar("SELECT EXISTS(SELECT 1 FROM routing_groups WHERE id = $1)")
.bind(group_id)
.fetch_one(&mut *tx)
.await
.map_postgres_err()?;
if !exists {
return Err(DataLayerError::InvalidInput(
"routing_group_not_found".to_string(),
));
}
sqlx::query(r#"
UPDATE routing_groups
SET config_json = jsonb_set(config_json::jsonb, '{disabled_providers}',
CASE WHEN id = $1
THEN COALESCE(config_json::jsonb -> 'disabled_providers', '[]'::jsonb) - $2::text
ELSE (COALESCE(config_json::jsonb -> 'disabled_providers', '[]'::jsonb) - $2::text) || jsonb_build_array($2::text)
END),
version = version + 1,
updated_at = EXTRACT(EPOCH FROM NOW())::bigint
WHERE (id <> $1 AND NOT (COALESCE(config_json::jsonb -> 'disabled_providers', '[]'::jsonb) ? $2::text))
OR (id = $1 AND (COALESCE(config_json::jsonb -> 'disabled_providers', '[]'::jsonb) ? $2::text))
"#)
.bind(group_id).bind(&provider.id)
.execute(&mut *tx).await.map_postgres_err()?;
}
if let Some(target_priority) = shift_existing_priorities_from {
sqlx::query(
r#"
@@ -3015,6 +3059,20 @@ impl ProviderCatalogReadRepository for SqlxProviderCatalogReadRepository {
#[async_trait]
impl ProviderCatalogWriteRepository for SqlxProviderCatalogReadRepository {
async fn create_provider_in_routing_group(
&self,
provider: &StoredProviderCatalogProvider,
shift_existing_priorities_from: Option<i32>,
routing_group_id: &str,
) -> Result<StoredProviderCatalogProvider, DataLayerError> {
self.create_provider_with_routing_group(
provider,
shift_existing_priorities_from,
Some(routing_group_id),
)
.await
}
async fn create_provider(
&self,
provider: &StoredProviderCatalogProvider,
@@ -1463,6 +1463,19 @@ pub trait ProviderCatalogReadRepository: Send + Sync {
#[async_trait]
pub trait ProviderCatalogWriteRepository: Send + Sync {
/// Create a provider and exclude it from every other existing routing group
/// in one transaction. Implementations must fail closed if unsupported.
async fn create_provider_in_routing_group(
&self,
_provider: &StoredProviderCatalogProvider,
_shift_existing_priorities_from: Option<i32>,
_routing_group_id: &str,
) -> Result<StoredProviderCatalogProvider, crate::DataLayerError> {
Err(crate::DataLayerError::InvalidConfiguration(
"atomic provider creation in a routing group is not supported".to_string(),
))
}
async fn create_provider(
&self,
provider: &StoredProviderCatalogProvider,
@@ -58,6 +58,7 @@ pub struct CreateRoutingGroupRecord {
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct UpdateRoutingGroupRecord {
pub expected_version: Option<i64>,
pub name: Option<String>,
pub description: Option<Option<String>>,
pub enabled: Option<bool>,
@@ -244,6 +245,16 @@ pub fn apply_group_patch(
group: &mut StoredRoutingGroup,
patch: UpdateRoutingGroupRecord,
) -> Result<(), crate::DataLayerError> {
if patch
.expected_version
.is_some_and(|version| version != group.version)
{
return Err(crate::DataLayerError::InvalidInput(
"routing_group_version_conflict".to_string(),
));
}
let previous_version = group.version;
let config_changed = patch.config_json.is_some();
if let Some(name) = patch.name {
if name.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
@@ -273,7 +284,13 @@ pub fn apply_group_patch(
group.config_json = config_json;
}
if let Some(version) = patch.version {
group.version = version.max(1);
group.version = if config_changed {
version.max(previous_version.saturating_add(1))
} else {
version.max(previous_version)
};
} else if config_changed {
group.version = previous_version.saturating_add(1);
}
if let Some(published_at) = patch.published_at {
group.published_at = published_at;
@@ -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>,