mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 10:57:03 +08:00
feat: unify provider scheduling workspace
This commit is contained in:
@@ -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