mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-13 06:30:20 +08:00
Merge branch 'fawney19:main' into main
This commit is contained in:
@@ -3,9 +3,10 @@ use std::sync::RwLock;
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::{
|
||||
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
|
||||
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
|
||||
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
|
||||
MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery,
|
||||
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder,
|
||||
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
|
||||
StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
|
||||
@@ -61,6 +62,19 @@ impl MinimalCandidateSelectionReadRepository for InMemoryMinimalCandidateSelecti
|
||||
Ok(rows)
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_page(
|
||||
&self,
|
||||
query: &StoredApiFormatCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Ok(self
|
||||
.list_for_exact_api_format(&query.api_format)
|
||||
.await?
|
||||
.into_iter()
|
||||
.skip(query.offset as usize)
|
||||
.take(query.limit as usize)
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_global_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
@@ -351,8 +365,9 @@ fn key_auth_channel_matches(row: &StoredMinimalCandidateSelectionRow, api_format
|
||||
mod tests {
|
||||
use super::InMemoryMinimalCandidateSelectionReadRepository;
|
||||
use crate::repository::candidate_selection::{
|
||||
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
|
||||
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
|
||||
MinimalCandidateSelectionReadRepository, StoredApiFormatCandidateRowsQuery,
|
||||
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder,
|
||||
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
|
||||
StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
|
||||
@@ -697,6 +712,27 @@ mod tests {
|
||||
assert_eq!(rows[1].provider_id, "provider-2");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn lists_exact_api_format_in_stable_pages() {
|
||||
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||
sample_row("provider-3", "openai:chat", "gpt-4.1", 30),
|
||||
sample_row("provider-1", "openai:chat", "gpt-4.1", 10),
|
||||
sample_row("provider-2", "openai:chat", "gpt-4.1", 20),
|
||||
]);
|
||||
|
||||
let page = repository
|
||||
.list_for_exact_api_format_page(&StoredApiFormatCandidateRowsQuery {
|
||||
api_format: "openai:chat".to_string(),
|
||||
offset: 1,
|
||||
limit: 1,
|
||||
})
|
||||
.await
|
||||
.expect("API-format page should load");
|
||||
|
||||
assert_eq!(page.len(), 1);
|
||||
assert_eq!(page[0].provider_id, "provider-2");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn list_pool_key_rows_for_group_returns_requested_page_only() {
|
||||
let mut rows = Vec::new();
|
||||
|
||||
@@ -3,9 +3,10 @@ mod memory;
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use aether_data_contracts::repository::candidate_selection::{
|
||||
MinimalCandidateSelectionReadRepository, MinimalCandidateSelectionRepository,
|
||||
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder,
|
||||
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
|
||||
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
|
||||
StoredApiFormatCandidateRowsQuery, StoredMinimalCandidateSelectionRow,
|
||||
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
|
||||
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
|
||||
StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
#[cfg(feature = "mysql")]
|
||||
pub use aether_data_mysql::MysqlMinimalCandidateSelectionReadRepository;
|
||||
|
||||
@@ -72,6 +72,13 @@ pub fn build_provider_oauth_batch_task_status_payload(
|
||||
"submitted" | "processing" | "completed" | "failed" => raw_status,
|
||||
_ => "failed",
|
||||
};
|
||||
let import_kind = match state.get("import_kind").and_then(serde_json::Value::as_str) {
|
||||
Some("oauth_batch" | "agent_identity" | "cookie_authorize") => state
|
||||
.get("import_kind")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or_default(),
|
||||
_ => "",
|
||||
};
|
||||
let error_samples = state
|
||||
.get("error_samples")
|
||||
.and_then(serde_json::Value::as_array)
|
||||
@@ -94,6 +101,7 @@ pub fn build_provider_oauth_batch_task_status_payload(
|
||||
.get("provider_type")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or_default(),
|
||||
"import_kind": import_kind,
|
||||
"status": normalized_status,
|
||||
"total": state.get("total").and_then(serde_json::Value::as_i64).unwrap_or(0),
|
||||
"processed": state.get("processed").and_then(serde_json::Value::as_i64).unwrap_or(0),
|
||||
@@ -169,6 +177,7 @@ mod tests {
|
||||
let input = json!({
|
||||
"task_id": "task-123",
|
||||
"provider_type": "codex",
|
||||
"import_kind": "oauth_batch",
|
||||
"status": "weird",
|
||||
"total": 4,
|
||||
"processed": 2,
|
||||
@@ -194,6 +203,10 @@ mod tests {
|
||||
payload.get("provider_id").and_then(|v| v.as_str()),
|
||||
Some("provider-123")
|
||||
);
|
||||
assert_eq!(
|
||||
payload.get("import_kind").and_then(|v| v.as_str()),
|
||||
Some("oauth_batch")
|
||||
);
|
||||
assert_eq!(
|
||||
payload.get("status").and_then(|v| v.as_str()),
|
||||
Some("failed")
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::sync::RwLock;
|
||||
|
||||
use aether_data_contracts::repository::routing_profiles::{apply_binding_patch, apply_group_patch};
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::{
|
||||
@@ -149,10 +150,13 @@ impl RoutingGroupWriteRepository for InMemoryRoutingGroupRepository {
|
||||
record: CreateRoutingGroupRecord,
|
||||
) -> Result<StoredRoutingGroup, DataLayerError> {
|
||||
let group = StoredRoutingGroup::new(record)?;
|
||||
self.groups
|
||||
.write()
|
||||
.expect("routing group repository lock")
|
||||
.insert(group.id.clone(), group.clone());
|
||||
let mut groups = self.groups.write().expect("routing group repository lock");
|
||||
if group.is_system_default {
|
||||
for existing in groups.values_mut() {
|
||||
existing.is_system_default = false;
|
||||
}
|
||||
}
|
||||
groups.insert(group.id.clone(), group.clone());
|
||||
Ok(group)
|
||||
}
|
||||
|
||||
@@ -162,42 +166,17 @@ impl RoutingGroupWriteRepository for InMemoryRoutingGroupRepository {
|
||||
patch: UpdateRoutingGroupRecord,
|
||||
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
|
||||
let mut groups = self.groups.write().expect("routing group repository lock");
|
||||
let Some(group) = groups.get_mut(id) else {
|
||||
let Some(mut group) = groups.get(id).cloned() else {
|
||||
return Ok(None);
|
||||
};
|
||||
if let Some(name) = patch.name {
|
||||
if name.trim().is_empty() {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"routing_groups.name is empty".to_string(),
|
||||
));
|
||||
apply_group_patch(&mut group, patch)?;
|
||||
if group.is_system_default {
|
||||
for existing in groups.values_mut() {
|
||||
existing.is_system_default = false;
|
||||
}
|
||||
group.name = name;
|
||||
}
|
||||
if let Some(description) = patch.description {
|
||||
group.description = description;
|
||||
}
|
||||
if let Some(enabled) = patch.enabled {
|
||||
group.enabled = enabled;
|
||||
}
|
||||
if let Some(is_system_default) = patch.is_system_default {
|
||||
group.is_system_default = is_system_default;
|
||||
}
|
||||
if let Some(config_json) = patch.config_json {
|
||||
if !config_json.is_object() {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"routing_groups.config_json must be a JSON object".to_string(),
|
||||
));
|
||||
}
|
||||
group.config_json = config_json;
|
||||
}
|
||||
if let Some(version) = patch.version {
|
||||
group.version = version.max(1);
|
||||
}
|
||||
if let Some(published_at) = patch.published_at {
|
||||
group.published_at = published_at;
|
||||
}
|
||||
group.updated_at = patch.updated_at;
|
||||
Ok(Some(group.clone()))
|
||||
groups.insert(id.to_string(), group.clone());
|
||||
Ok(Some(group))
|
||||
}
|
||||
|
||||
async fn delete_routing_group(&self, id: &str) -> Result<bool, DataLayerError> {
|
||||
@@ -214,10 +193,19 @@ impl RoutingGroupWriteRepository for InMemoryRoutingGroupRepository {
|
||||
record: CreateRoutingGroupBindingRecord,
|
||||
) -> Result<StoredRoutingGroupBinding, DataLayerError> {
|
||||
let binding = StoredRoutingGroupBinding::new(record)?;
|
||||
self.bindings
|
||||
let mut bindings = self
|
||||
.bindings
|
||||
.write()
|
||||
.expect("routing group binding repository lock")
|
||||
.insert(binding.id.clone(), binding.clone());
|
||||
.expect("routing group binding repository lock");
|
||||
if binding.is_default {
|
||||
for existing in bindings.values_mut().filter(|existing| {
|
||||
existing.subject_type == binding.subject_type
|
||||
&& existing.subject_id == binding.subject_id
|
||||
}) {
|
||||
existing.is_default = false;
|
||||
}
|
||||
}
|
||||
bindings.insert(binding.id.clone(), binding.clone());
|
||||
Ok(binding)
|
||||
}
|
||||
|
||||
@@ -239,36 +227,20 @@ impl RoutingGroupWriteRepository for InMemoryRoutingGroupRepository {
|
||||
.bindings
|
||||
.write()
|
||||
.expect("routing group binding repository lock");
|
||||
let Some(binding) = bindings.get_mut(id) else {
|
||||
let Some(mut binding) = bindings.get(id).cloned() else {
|
||||
return Ok(None);
|
||||
};
|
||||
if let Some(group_id) = patch.group_id {
|
||||
if group_id.trim().is_empty() {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"routing_group_bindings.group_id is empty".to_string(),
|
||||
));
|
||||
apply_binding_patch(&mut binding, patch)?;
|
||||
if binding.is_default {
|
||||
for existing in bindings.values_mut().filter(|existing| {
|
||||
existing.subject_type == binding.subject_type
|
||||
&& existing.subject_id == binding.subject_id
|
||||
}) {
|
||||
existing.is_default = false;
|
||||
}
|
||||
binding.group_id = group_id;
|
||||
}
|
||||
if let Some(subject_type) = patch.subject_type {
|
||||
binding.subject_type = subject_type;
|
||||
}
|
||||
if let Some(subject_id) = patch.subject_id {
|
||||
if subject_id.trim().is_empty() {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"routing_group_bindings.subject_id is empty".to_string(),
|
||||
));
|
||||
}
|
||||
binding.subject_id = subject_id;
|
||||
}
|
||||
if let Some(is_default) = patch.is_default {
|
||||
binding.is_default = is_default;
|
||||
}
|
||||
if let Some(allow_explicit_select) = patch.allow_explicit_select {
|
||||
binding.allow_explicit_select = allow_explicit_select;
|
||||
}
|
||||
binding.updated_at = patch.updated_at;
|
||||
Ok(Some(binding.clone()))
|
||||
bindings.insert(id.to_string(), binding.clone());
|
||||
Ok(Some(binding))
|
||||
}
|
||||
|
||||
async fn create_routing_group_version(
|
||||
@@ -347,4 +319,153 @@ mod tests {
|
||||
1
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn keeps_system_and_subject_defaults_unique() {
|
||||
let repository = InMemoryRoutingGroupRepository::default();
|
||||
for (id, is_system_default) in [("group-1", true), ("group-2", true), ("group-3", false)] {
|
||||
repository
|
||||
.create_routing_group(group_record(id, is_system_default))
|
||||
.await
|
||||
.expect("group should store");
|
||||
}
|
||||
|
||||
assert_eq!(system_default_ids(&repository).await, vec!["group-2"]);
|
||||
|
||||
repository
|
||||
.update_routing_group(
|
||||
"group-1",
|
||||
UpdateRoutingGroupRecord {
|
||||
is_system_default: Some(true),
|
||||
updated_at: 2,
|
||||
..UpdateRoutingGroupRecord::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("group should update");
|
||||
assert_eq!(system_default_ids(&repository).await, vec!["group-1"]);
|
||||
|
||||
repository
|
||||
.create_routing_group_binding(binding_record("binding-1", "group-1", "subject-1", true))
|
||||
.await
|
||||
.expect("binding should store");
|
||||
repository
|
||||
.create_routing_group_binding(binding_record("binding-2", "group-2", "subject-1", true))
|
||||
.await
|
||||
.expect("binding should store");
|
||||
repository
|
||||
.create_routing_group_binding(binding_record("binding-3", "group-3", "subject-2", true))
|
||||
.await
|
||||
.expect("binding should store");
|
||||
|
||||
assert_eq!(
|
||||
default_binding_ids(&repository, "subject-1").await,
|
||||
vec!["binding-2"]
|
||||
);
|
||||
assert_eq!(
|
||||
default_binding_ids(&repository, "subject-2").await,
|
||||
vec!["binding-3"]
|
||||
);
|
||||
|
||||
repository
|
||||
.update_routing_group_binding(
|
||||
"binding-1",
|
||||
UpdateRoutingGroupBindingRecord {
|
||||
is_default: Some(true),
|
||||
updated_at: 2,
|
||||
..UpdateRoutingGroupBindingRecord::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("binding should update");
|
||||
assert_eq!(
|
||||
default_binding_ids(&repository, "subject-1").await,
|
||||
vec!["binding-1"]
|
||||
);
|
||||
assert_eq!(
|
||||
default_binding_ids(&repository, "subject-2").await,
|
||||
vec!["binding-3"]
|
||||
);
|
||||
|
||||
repository
|
||||
.update_routing_group_binding(
|
||||
"binding-3",
|
||||
UpdateRoutingGroupBindingRecord {
|
||||
subject_id: Some("subject-1".to_string()),
|
||||
updated_at: 3,
|
||||
..UpdateRoutingGroupBindingRecord::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("binding should move");
|
||||
assert_eq!(
|
||||
default_binding_ids(&repository, "subject-1").await,
|
||||
vec!["binding-3"]
|
||||
);
|
||||
assert!(default_binding_ids(&repository, "subject-2")
|
||||
.await
|
||||
.is_empty());
|
||||
}
|
||||
|
||||
fn group_record(id: &str, is_system_default: bool) -> CreateRoutingGroupRecord {
|
||||
CreateRoutingGroupRecord {
|
||||
id: id.to_string(),
|
||||
name: id.to_string(),
|
||||
description: None,
|
||||
enabled: true,
|
||||
is_system_default,
|
||||
config_json: json!({}),
|
||||
version: 1,
|
||||
created_at: 1,
|
||||
updated_at: 1,
|
||||
published_at: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn binding_record(
|
||||
id: &str,
|
||||
group_id: &str,
|
||||
subject_id: &str,
|
||||
is_default: bool,
|
||||
) -> CreateRoutingGroupBindingRecord {
|
||||
CreateRoutingGroupBindingRecord {
|
||||
id: id.to_string(),
|
||||
group_id: group_id.to_string(),
|
||||
subject_type: RoutingGroupBindingSubject::ApiKey,
|
||||
subject_id: subject_id.to_string(),
|
||||
is_default,
|
||||
allow_explicit_select: true,
|
||||
created_at: 1,
|
||||
updated_at: 1,
|
||||
}
|
||||
}
|
||||
|
||||
async fn system_default_ids(repository: &InMemoryRoutingGroupRepository) -> Vec<String> {
|
||||
repository
|
||||
.list_routing_groups()
|
||||
.await
|
||||
.expect("groups should list")
|
||||
.into_iter()
|
||||
.filter(|group| group.is_system_default)
|
||||
.map(|group| group.id)
|
||||
.collect()
|
||||
}
|
||||
|
||||
async fn default_binding_ids(
|
||||
repository: &InMemoryRoutingGroupRepository,
|
||||
subject_id: &str,
|
||||
) -> Vec<String> {
|
||||
repository
|
||||
.list_routing_group_bindings(&RoutingGroupBindingQuery {
|
||||
group_id: None,
|
||||
subject_type: Some(RoutingGroupBindingSubject::ApiKey),
|
||||
subject_id: Some(subject_id.to_string()),
|
||||
})
|
||||
.await
|
||||
.expect("bindings should list")
|
||||
.into_iter()
|
||||
.filter(|binding| binding.is_default)
|
||||
.map(|binding| binding.id)
|
||||
.collect()
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user