Merge branch 'fawney19:main' into main

This commit is contained in:
ZheFox
2026-07-28 13:56:34 +08:00
committed by GitHub
117 changed files with 10467 additions and 1253 deletions
@@ -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()
}
}