Merge remote-tracking branch 'origin/main' into codex/pool-key-bulk-management-20260714

# Conflicts:
#	apps/aether-gateway/src/handlers/admin/request/provider/tasks.rs
#	frontend/src/api/endpoints/pool.ts
This commit is contained in:
MMEXA
2026-07-16 23:43:04 +08:00
1257 changed files with 80521 additions and 35495 deletions
@@ -0,0 +1,437 @@
use std::collections::BTreeSet;
use std::sync::RwLock;
use std::time::{SystemTime, UNIX_EPOCH};
use async_trait::async_trait;
use uuid::Uuid;
use crate::DataLayerError;
use aether_data_contracts::repository::announcements::{
AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository,
CreateAnnouncementRecord, StoredAnnouncement, StoredAnnouncementPage, UpdateAnnouncementRecord,
};
#[derive(Debug, Default)]
pub struct InMemoryAnnouncementReadRepository {
announcements: RwLock<Vec<StoredAnnouncement>>,
announcement_reads: RwLock<BTreeSet<(String, String)>>,
}
impl InMemoryAnnouncementReadRepository {
pub fn seed<I>(announcements: I) -> Self
where
I: IntoIterator<Item = StoredAnnouncement>,
{
Self::seed_with_reads(announcements, std::iter::empty::<(String, String)>())
}
pub fn seed_with_reads<I, J>(announcements: I, reads: J) -> Self
where
I: IntoIterator<Item = StoredAnnouncement>,
J: IntoIterator<Item = (String, String)>,
{
Self {
announcements: RwLock::new(announcements.into_iter().collect()),
announcement_reads: RwLock::new(reads.into_iter().collect()),
}
}
fn now_unix_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_secs()
}
}
#[async_trait]
impl AnnouncementReadRepository for InMemoryAnnouncementReadRepository {
async fn find_by_id(
&self,
announcement_id: &str,
) -> Result<Option<StoredAnnouncement>, DataLayerError> {
Ok(self
.announcements
.read()
.expect("announcement repository lock")
.iter()
.find(|announcement| announcement.id == announcement_id)
.cloned())
}
async fn list_announcements(
&self,
query: &AnnouncementListQuery,
) -> Result<StoredAnnouncementPage, DataLayerError> {
let now_unix_secs = query.now_unix_secs.unwrap_or_else(Self::now_unix_secs);
let announcements = self
.announcements
.read()
.expect("announcement repository lock");
let mut items: Vec<_> = announcements
.iter()
.filter(|announcement| {
if !query.active_only {
return true;
}
announcement.is_active
&& announcement
.start_time_unix_secs
.is_none_or(|value| value <= now_unix_secs)
&& announcement
.end_time_unix_secs
.is_none_or(|value| value >= now_unix_secs)
})
.cloned()
.collect();
items.sort_by(|left, right| {
right
.is_pinned
.cmp(&left.is_pinned)
.then_with(|| right.priority.cmp(&left.priority))
.then_with(|| right.created_at_unix_ms.cmp(&left.created_at_unix_ms))
.then_with(|| left.id.cmp(&right.id))
});
let total = items.len() as u64;
let items = items
.into_iter()
.skip(query.offset)
.take(query.limit)
.collect();
Ok(StoredAnnouncementPage { items, total })
}
async fn count_unread_active_announcements(
&self,
user_id: &str,
now_unix_secs: u64,
) -> Result<u64, DataLayerError> {
let announcements = self
.announcements
.read()
.expect("announcement repository lock");
let reads = self
.announcement_reads
.read()
.expect("announcement reads repository lock");
let total = announcements
.iter()
.filter(|announcement| {
announcement.is_active
&& announcement
.start_time_unix_secs
.is_none_or(|value| value <= now_unix_secs)
&& announcement
.end_time_unix_secs
.is_none_or(|value| value >= now_unix_secs)
&& !reads.contains(&(user_id.to_string(), announcement.id.clone()))
})
.count() as u64;
Ok(total)
}
async fn list_required_unread_active_announcements(
&self,
user_id: &str,
now_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredAnnouncement>, DataLayerError> {
let announcements = self
.announcements
.read()
.expect("announcement repository lock");
let reads = self
.announcement_reads
.read()
.expect("announcement reads repository lock");
let mut items = announcements
.iter()
.filter(|announcement| {
announcement.requires_ack
&& announcement.is_active
&& announcement
.start_time_unix_secs
.is_none_or(|value| value <= now_unix_secs)
&& announcement
.end_time_unix_secs
.is_none_or(|value| value >= now_unix_secs)
&& !reads.contains(&(user_id.to_string(), announcement.id.clone()))
})
.cloned()
.collect::<Vec<_>>();
items.sort_by(|left, right| {
right
.is_pinned
.cmp(&left.is_pinned)
.then_with(|| right.priority.cmp(&left.priority))
.then_with(|| right.created_at_unix_ms.cmp(&left.created_at_unix_ms))
.then_with(|| left.id.cmp(&right.id))
});
items.truncate(limit);
Ok(items)
}
}
#[async_trait]
impl AnnouncementWriteRepository for InMemoryAnnouncementReadRepository {
async fn create_announcement(
&self,
record: CreateAnnouncementRecord,
) -> Result<StoredAnnouncement, DataLayerError> {
record.validate()?;
let now_unix_secs = Self::now_unix_secs();
let announcement = StoredAnnouncement::new(
Uuid::new_v4().to_string(),
record.title,
record.content,
record.kind,
record.priority,
true,
record.is_pinned,
record.requires_ack,
Some(record.author_id),
None,
record.start_time_unix_secs.map(|value| value as i64),
record.end_time_unix_secs.map(|value| value as i64),
now_unix_secs as i64,
now_unix_secs as i64,
)?;
self.announcements
.write()
.expect("announcement repository lock")
.push(announcement.clone());
Ok(announcement)
}
async fn update_announcement(
&self,
record: UpdateAnnouncementRecord,
) -> Result<Option<StoredAnnouncement>, DataLayerError> {
record.validate()?;
let mut announcements = self
.announcements
.write()
.expect("announcement repository lock");
let Some(announcement) = announcements
.iter_mut()
.find(|announcement| announcement.id == record.announcement_id)
else {
return Ok(None);
};
if let Some(title) = record.title {
announcement.title = title;
}
if let Some(content) = record.content {
announcement.content = content;
}
if let Some(kind) = record.kind {
announcement.kind = kind;
}
if let Some(priority) = record.priority {
announcement.priority = priority;
}
if let Some(is_active) = record.is_active {
announcement.is_active = is_active;
}
if let Some(is_pinned) = record.is_pinned {
announcement.is_pinned = is_pinned;
}
if let Some(requires_ack) = record.requires_ack {
announcement.requires_ack = requires_ack;
}
if let Some(start_time_unix_secs) = record.start_time_unix_secs {
announcement.start_time_unix_secs = Some(start_time_unix_secs);
}
if let Some(end_time_unix_secs) = record.end_time_unix_secs {
announcement.end_time_unix_secs = Some(end_time_unix_secs);
}
announcement.updated_at_unix_secs = Self::now_unix_secs();
Ok(Some(announcement.clone()))
}
async fn delete_announcement(&self, announcement_id: &str) -> Result<bool, DataLayerError> {
let mut announcements = self
.announcements
.write()
.expect("announcement repository lock");
let original_len = announcements.len();
announcements.retain(|announcement| announcement.id != announcement_id);
let deleted = announcements.len() != original_len;
if deleted {
self.announcement_reads
.write()
.expect("announcement reads repository lock")
.retain(|(_, read_announcement_id)| read_announcement_id != announcement_id);
}
Ok(deleted)
}
async fn mark_announcement_as_read(
&self,
user_id: &str,
announcement_id: &str,
_read_at_unix_secs: u64,
) -> Result<bool, DataLayerError> {
let inserted = self
.announcement_reads
.write()
.expect("announcement reads repository lock")
.insert((user_id.to_string(), announcement_id.to_string()));
Ok(inserted)
}
}
#[cfg(test)]
mod tests {
use super::InMemoryAnnouncementReadRepository;
use crate::repository::announcements::{
AnnouncementReadRepository, AnnouncementWriteRepository, CreateAnnouncementRecord,
StoredAnnouncement, UpdateAnnouncementRecord,
};
#[tokio::test]
async fn reads_seeded_announcements() {
let repository = InMemoryAnnouncementReadRepository::seed(vec![StoredAnnouncement::new(
"announcement-1".to_string(),
"系统维护".to_string(),
"今天晚些时候维护".to_string(),
"maintenance".to_string(),
10,
true,
true,
false,
Some("admin-1".to_string()),
Some("admin".to_string()),
None,
None,
1_711_000_000,
1_711_000_100,
)
.expect("announcement should build")]);
let announcement = repository
.find_by_id("announcement-1")
.await
.expect("announcement should load")
.expect("announcement should exist");
assert_eq!(announcement.title, "系统维护");
assert_eq!(announcement.author_username.as_deref(), Some("admin"));
}
#[tokio::test]
async fn mutates_seeded_announcements() {
let repository = InMemoryAnnouncementReadRepository::seed(vec![]);
let created = repository
.create_announcement(CreateAnnouncementRecord {
title: "系统维护".to_string(),
content: "今天晚些时候维护".to_string(),
kind: "maintenance".to_string(),
priority: 10,
is_pinned: true,
requires_ack: false,
author_id: "admin-1".to_string(),
start_time_unix_secs: None,
end_time_unix_secs: None,
})
.await
.expect("create should succeed");
assert_eq!(created.kind, "maintenance");
let updated = repository
.update_announcement(UpdateAnnouncementRecord {
announcement_id: created.id.clone(),
title: Some("系统升级".to_string()),
content: None,
kind: Some("important".to_string()),
priority: Some(99),
is_active: Some(false),
is_pinned: Some(false),
requires_ack: Some(true),
start_time_unix_secs: None,
end_time_unix_secs: None,
})
.await
.expect("update should succeed")
.expect("announcement should exist");
assert_eq!(updated.title, "系统升级");
assert_eq!(updated.kind, "important");
assert!(!updated.is_active);
let deleted = repository
.delete_announcement(&created.id)
.await
.expect("delete should succeed");
assert!(deleted);
}
#[tokio::test]
async fn tracks_user_announcement_read_state() {
let repository = InMemoryAnnouncementReadRepository::seed_with_reads(
vec![
StoredAnnouncement::new(
"announcement-1".to_string(),
"系统维护".to_string(),
"今天晚些时候维护".to_string(),
"maintenance".to_string(),
10,
true,
true,
false,
Some("admin-1".to_string()),
Some("admin".to_string()),
None,
None,
1_711_000_000,
1_711_000_100,
)
.expect("announcement should build"),
StoredAnnouncement::new(
"announcement-2".to_string(),
"系统升级".to_string(),
"升级说明".to_string(),
"info".to_string(),
5,
true,
false,
false,
Some("admin-1".to_string()),
Some("admin".to_string()),
None,
None,
1_711_000_000,
1_711_000_100,
)
.expect("announcement should build"),
],
[("user-1".to_string(), "announcement-1".to_string())],
);
let unread = repository
.count_unread_active_announcements("user-1", 1_711_000_200)
.await
.expect("count should succeed");
assert_eq!(unread, 1);
let inserted = repository
.mark_announcement_as_read("user-1", "announcement-2", 1_711_000_300)
.await
.expect("mark read should succeed");
assert!(inserted);
let unread = repository
.count_unread_active_announcements("user-1", 1_711_000_200)
.await
.expect("count should succeed");
assert_eq!(unread, 0);
}
}
@@ -0,0 +1,13 @@
mod memory;
pub use aether_data_contracts::repository::announcements::{
AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository,
CreateAnnouncementRecord, StoredAnnouncement, StoredAnnouncementPage, UpdateAnnouncementRecord,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlAnnouncementRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxAnnouncementReadRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteAnnouncementRepository;
pub use memory::InMemoryAnnouncementReadRepository;
@@ -0,0 +1,17 @@
mod types;
#[cfg(test)]
mod tests;
pub use aether_data_contracts::repository::audit::{
optional_json_from_text, AuditLogListQuery, AuditLogReadRepository, StoredAdminAuditLog,
StoredAdminAuditLogPage, StoredSuspiciousActivity, StoredUserAuditLog, StoredUserAuditLogPage,
SUSPICIOUS_EVENT_TYPES,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlAuditLogReadRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::PostgresAuditLogReadRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteAuditLogReadRepository;
pub use types::{read_request_audit_bundle, RequestAuditBundle, RequestAuditReader};
@@ -0,0 +1,225 @@
use std::sync::atomic::{AtomicUsize, Ordering};
use async_trait::async_trait;
use super::{read_request_audit_bundle, RequestAuditReader};
use crate::repository::auth::{ResolvedAuthApiKeySnapshot, StoredAuthApiKeySnapshot};
use crate::repository::candidates::{
DecisionTrace, DecisionTraceCandidate, RequestCandidateFinalStatus, RequestCandidateStatus,
StoredRequestCandidate,
};
use crate::repository::usage::StoredRequestUsageAudit;
use crate::DataLayerError;
#[derive(Default)]
struct FakeRequestAuditReader {
usage: Option<StoredRequestUsageAudit>,
decision_trace: Option<DecisionTrace>,
auth_snapshot: Option<ResolvedAuthApiKeySnapshot>,
auth_snapshot_reads: AtomicUsize,
}
#[async_trait]
impl RequestAuditReader for FakeRequestAuditReader {
async fn find_request_usage_audit_by_request_id(
&self,
_request_id: &str,
) -> Result<Option<StoredRequestUsageAudit>, DataLayerError> {
Ok(self.usage.clone())
}
async fn read_request_decision_trace(
&self,
_request_id: &str,
_attempted_only: bool,
) -> Result<Option<DecisionTrace>, DataLayerError> {
Ok(self.decision_trace.clone())
}
async fn read_resolved_auth_api_key_snapshot(
&self,
_user_id: &str,
_api_key_id: &str,
_now_unix_secs: u64,
) -> Result<Option<ResolvedAuthApiKeySnapshot>, DataLayerError> {
self.auth_snapshot_reads.fetch_add(1, Ordering::Relaxed);
Ok(self.auth_snapshot.clone())
}
}
#[tokio::test]
async fn read_request_audit_bundle_resolves_usage_trace_and_auth_snapshot() {
let state = FakeRequestAuditReader {
usage: Some(sample_usage("req-audit-1")),
decision_trace: Some(sample_decision_trace("req-audit-1")),
auth_snapshot: Some(sample_resolved_auth_snapshot("user-1", "api-key-1")),
auth_snapshot_reads: AtomicUsize::new(0),
};
let bundle = read_request_audit_bundle(&state, "req-audit-1", true, 123)
.await
.expect("bundle should read")
.expect("bundle should exist");
assert_eq!(bundle.request_id, "req-audit-1");
assert_eq!(
bundle
.usage
.as_ref()
.map(|usage| usage.provider_name.as_str()),
Some("OpenAI")
);
assert_eq!(
bundle
.decision_trace
.as_ref()
.map(|trace| trace.total_candidates),
Some(1)
);
assert_eq!(
bundle
.auth_snapshot
.as_ref()
.map(|snapshot| snapshot.api_key_id.as_str()),
Some("api-key-1")
);
assert_eq!(state.auth_snapshot_reads.load(Ordering::Relaxed), 1);
}
#[tokio::test]
async fn read_request_audit_bundle_returns_none_when_all_sources_are_empty() {
let state = FakeRequestAuditReader::default();
let bundle = read_request_audit_bundle(&state, "req-audit-empty", false, 123)
.await
.expect("bundle should read");
assert!(bundle.is_none());
assert_eq!(state.auth_snapshot_reads.load(Ordering::Relaxed), 0);
}
fn sample_usage(request_id: &str) -> StoredRequestUsageAudit {
StoredRequestUsageAudit::new(
"usage-1".to_string(),
request_id.to_string(),
Some("user-1".to_string()),
Some("api-key-1".to_string()),
Some("alice".to_string()),
Some("default".to_string()),
"OpenAI".to_string(),
"gpt-4.1".to_string(),
None,
Some("provider-1".to_string()),
Some("endpoint-1".to_string()),
Some("provider-key-1".to_string()),
Some("chat".to_string()),
Some("openai:chat".to_string()),
Some("openai".to_string()),
Some("chat".to_string()),
Some("openai:chat".to_string()),
Some("openai".to_string()),
Some("chat".to_string()),
false,
false,
120,
40,
160,
0.24,
0.36,
Some(200),
None,
None,
Some(450),
Some(120),
"completed".to_string(),
"settled".to_string(),
100,
101,
Some(102),
)
.expect("usage should build")
}
fn sample_decision_trace(request_id: &str) -> DecisionTrace {
let candidate = StoredRequestCandidate::new(
"cand-1".to_string(),
request_id.to_string(),
Some("user-1".to_string()),
Some("api-key-1".to_string()),
Some("alice".to_string()),
Some("default".to_string()),
0,
0,
Some("provider-1".to_string()),
Some("endpoint-1".to_string()),
Some("provider-key-1".to_string()),
RequestCandidateStatus::Success,
None,
false,
Some(200),
None,
None,
Some(37),
None,
None,
None,
100,
Some(101),
Some(102),
)
.expect("candidate should build");
DecisionTrace {
request_id: request_id.to_string(),
total_candidates: 1,
final_status: RequestCandidateFinalStatus::Success,
total_latency_ms: 37,
candidates: vec![DecisionTraceCandidate {
candidate,
provider_name: Some("OpenAI".to_string()),
provider_website: None,
provider_type: Some("custom".to_string()),
provider_priority: Some(0),
provider_keep_priority_on_conversion: Some(false),
provider_enable_format_conversion: Some(false),
endpoint_api_format: Some("openai:chat".to_string()),
endpoint_api_family: Some("openai".to_string()),
endpoint_kind: Some("chat".to_string()),
endpoint_format_acceptance_config: None,
provider_key_name: Some("prod".to_string()),
provider_key_auth_type: Some("api_key".to_string()),
provider_key_api_formats: None,
provider_key_internal_priority: Some(10),
provider_key_global_priority_by_format: None,
provider_key_capabilities: None,
provider_key_is_active: Some(true),
}],
}
}
fn sample_resolved_auth_snapshot(user_id: &str, api_key_id: &str) -> ResolvedAuthApiKeySnapshot {
let stored = StoredAuthApiKeySnapshot::new(
user_id.to_string(),
"alice".to_string(),
Some("[email protected]".to_string()),
"user".to_string(),
"local".to_string(),
true,
false,
Some(serde_json::json!(["openai"])),
Some(serde_json::json!(["openai:chat"])),
Some(serde_json::json!(["gpt-4.1"])),
api_key_id.to_string(),
Some("default".to_string()),
true,
false,
false,
Some(60),
Some(5),
Some(4_102_444_800),
Some(serde_json::json!(["openai"])),
Some(serde_json::json!(["openai:chat"])),
Some(serde_json::json!(["gpt-4.1"])),
)
.expect("auth snapshot should build");
ResolvedAuthApiKeySnapshot::from_stored(stored, 123)
}
@@ -0,0 +1,76 @@
use async_trait::async_trait;
use crate::repository::auth::ResolvedAuthApiKeySnapshot;
use crate::repository::candidates::DecisionTrace;
use crate::repository::usage::StoredRequestUsageAudit;
use crate::DataLayerError;
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct RequestAuditBundle {
pub request_id: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub usage: Option<StoredRequestUsageAudit>,
#[serde(skip_serializing_if = "Option::is_none")]
pub decision_trace: Option<DecisionTrace>,
#[serde(skip_serializing_if = "Option::is_none")]
pub auth_snapshot: Option<ResolvedAuthApiKeySnapshot>,
}
#[async_trait]
pub trait RequestAuditReader {
async fn find_request_usage_audit_by_request_id(
&self,
request_id: &str,
) -> Result<Option<StoredRequestUsageAudit>, DataLayerError>;
async fn read_request_decision_trace(
&self,
request_id: &str,
attempted_only: bool,
) -> Result<Option<DecisionTrace>, DataLayerError>;
async fn read_resolved_auth_api_key_snapshot(
&self,
user_id: &str,
api_key_id: &str,
now_unix_secs: u64,
) -> Result<Option<ResolvedAuthApiKeySnapshot>, DataLayerError>;
}
pub async fn read_request_audit_bundle(
state: &impl RequestAuditReader,
request_id: &str,
attempted_only: bool,
now_unix_secs: u64,
) -> Result<Option<RequestAuditBundle>, DataLayerError> {
let usage = state
.find_request_usage_audit_by_request_id(request_id)
.await?;
let decision_trace = state
.read_request_decision_trace(request_id, attempted_only)
.await?;
let auth_snapshot = if let Some(usage) = usage.as_ref() {
match (usage.user_id.as_deref(), usage.api_key_id.as_deref()) {
(Some(user_id), Some(api_key_id)) => {
state
.read_resolved_auth_api_key_snapshot(user_id, api_key_id, now_unix_secs)
.await?
}
_ => None,
}
} else {
None
};
if usage.is_none() && decision_trace.is_none() && auth_snapshot.is_none() {
return Ok(None);
}
Ok(Some(RequestAuditBundle {
request_id: request_id.to_string(),
usage,
decision_trace,
auth_snapshot,
}))
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,18 @@
mod memory;
pub use aether_data_contracts::repository::auth::{
read_resolved_auth_api_key_snapshot, read_resolved_auth_api_key_snapshot_by_key_hash,
read_resolved_auth_api_key_snapshot_by_user_api_key_ids, AuthApiKeyExportSummary,
AuthApiKeyLookupKey, AuthApiKeyReadRepository, AuthApiKeyWriteRepository, AuthRepository,
CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord, ResolvedAuthApiKeySnapshot,
ResolvedAuthApiKeySnapshotReader, StandaloneApiKeyExportListQuery,
StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, UpdateStandaloneApiKeyBasicRecord,
UpdateUserApiKeyBasicRecord,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlAuthApiKeyReadRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxAuthApiKeySnapshotReadRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteAuthApiKeyReadRepository;
pub use memory::InMemoryAuthApiKeySnapshotRepository;
@@ -0,0 +1,114 @@
use std::sync::RwLock;
use async_trait::async_trait;
use super::{
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
StoredOAuthProviderModuleConfig,
};
use crate::DataLayerError;
#[derive(Debug, Default)]
pub struct InMemoryAuthModuleReadRepository {
oauth_providers: RwLock<Vec<StoredOAuthProviderModuleConfig>>,
ldap_config: RwLock<Option<StoredLdapModuleConfig>>,
}
impl InMemoryAuthModuleReadRepository {
pub fn seed<I>(oauth_providers: I, ldap_config: Option<StoredLdapModuleConfig>) -> Self
where
I: IntoIterator<Item = StoredOAuthProviderModuleConfig>,
{
Self {
oauth_providers: RwLock::new(oauth_providers.into_iter().collect()),
ldap_config: RwLock::new(ldap_config),
}
}
}
#[async_trait]
impl AuthModuleReadRepository for InMemoryAuthModuleReadRepository {
async fn list_enabled_oauth_providers(
&self,
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
Ok(self
.oauth_providers
.read()
.expect("auth module oauth provider repository lock")
.clone())
}
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
Ok(self
.ldap_config
.read()
.expect("auth module ldap repository lock")
.clone())
}
}
#[async_trait]
impl AuthModuleWriteRepository for InMemoryAuthModuleReadRepository {
async fn upsert_ldap_config(
&self,
config: &StoredLdapModuleConfig,
) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
self.ldap_config
.write()
.expect("auth module ldap repository lock")
.replace(config.clone());
Ok(Some(config.clone()))
}
}
#[cfg(test)]
mod tests {
use super::InMemoryAuthModuleReadRepository;
use crate::repository::auth_modules::{
AuthModuleReadRepository, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig,
};
#[tokio::test]
async fn reads_seeded_auth_module_configs() {
let repository = InMemoryAuthModuleReadRepository::seed(
vec![StoredOAuthProviderModuleConfig::new(
"linuxdo".to_string(),
"Linux DO".to_string(),
"client-id".to_string(),
Some("encrypted".to_string()),
"https://example.com/callback".to_string(),
)
.expect("oauth provider should build")],
Some(StoredLdapModuleConfig {
server_url: "ldaps://ldap.example.com".to_string(),
bind_dn: "cn=admin,dc=example,dc=com".to_string(),
bind_password_encrypted: Some("encrypted-password".to_string()),
base_dn: "dc=example,dc=com".to_string(),
user_search_filter: Some("(uid={username})".to_string()),
username_attr: Some("uid".to_string()),
email_attr: Some("mail".to_string()),
display_name_attr: Some("displayName".to_string()),
is_enabled: true,
is_exclusive: false,
use_starttls: true,
connect_timeout: Some(10),
}),
);
let oauth = repository
.list_enabled_oauth_providers()
.await
.expect("oauth providers should load");
let ldap = repository
.get_ldap_config()
.await
.expect("ldap config should load");
assert_eq!(oauth.len(), 1);
assert_eq!(oauth[0].provider_type, "linuxdo");
assert_eq!(
ldap.expect("ldap config should exist").server_url,
"ldaps://ldap.example.com"
);
}
}
@@ -0,0 +1,13 @@
mod memory;
pub use aether_data_contracts::repository::auth_modules::{
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
StoredOAuthProviderModuleConfig,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::{MysqlAuthModuleReadRepository, MysqlAuthModuleRepository};
#[cfg(feature = "postgres")]
pub use aether_data_postgres::{SqlxAuthModuleReadRepository, SqlxAuthModuleRepository};
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::{SqliteAuthModuleReadRepository, SqliteAuthModuleRepository};
pub use memory::InMemoryAuthModuleReadRepository;
@@ -0,0 +1,218 @@
use std::collections::{BTreeMap, BTreeSet};
use std::sync::RwLock;
use async_trait::async_trait;
use super::{
BackgroundTaskListQuery, BackgroundTaskReadRepository, BackgroundTaskStatus,
BackgroundTaskSummary, BackgroundTaskWriteRepository, StoredBackgroundTaskEvent,
StoredBackgroundTaskRun, StoredBackgroundTaskRunPage, UpsertBackgroundTaskEvent,
UpsertBackgroundTaskRun,
};
use crate::DataLayerError;
#[derive(Debug, Default)]
struct InMemoryBackgroundTaskIndex {
runs: BTreeMap<String, StoredBackgroundTaskRun>,
events_by_run: BTreeMap<String, Vec<StoredBackgroundTaskEvent>>,
}
#[derive(Debug, Default)]
pub struct InMemoryBackgroundTaskRepository {
index: RwLock<InMemoryBackgroundTaskIndex>,
}
impl InMemoryBackgroundTaskRepository {
fn matches_filter(run: &StoredBackgroundTaskRun, query: &BackgroundTaskListQuery) -> bool {
if let Some(kind) = query.kind {
if run.kind != kind {
return false;
}
}
if let Some(status) = query.status {
if run.status != status {
return false;
}
}
if let Some(trigger) = query.trigger.as_deref() {
if run.trigger != trigger {
return false;
}
}
if let Some(task_key_substring) = query.task_key_substring.as_deref() {
let needle = task_key_substring.to_ascii_lowercase();
if !run.task_key.to_ascii_lowercase().contains(&needle) {
return false;
}
}
true
}
pub fn seed_runs<I>(runs: I) -> Self
where
I: IntoIterator<Item = StoredBackgroundTaskRun>,
{
let mut index = InMemoryBackgroundTaskIndex::default();
for run in runs {
index.runs.insert(run.id.clone(), run);
}
Self {
index: RwLock::new(index),
}
}
}
#[async_trait]
impl BackgroundTaskReadRepository for InMemoryBackgroundTaskRepository {
async fn find_run(
&self,
run_id: &str,
) -> Result<Option<StoredBackgroundTaskRun>, DataLayerError> {
Ok(self
.index
.read()
.expect("background task repository lock")
.runs
.get(run_id)
.cloned())
}
async fn list_runs(
&self,
query: &BackgroundTaskListQuery,
) -> Result<StoredBackgroundTaskRunPage, DataLayerError> {
let mut items = self
.index
.read()
.expect("background task repository lock")
.runs
.values()
.filter(|run| Self::matches_filter(run, query))
.cloned()
.collect::<Vec<_>>();
items.sort_by(|left, right| {
right
.created_at_unix_secs
.cmp(&left.created_at_unix_secs)
.then_with(|| right.updated_at_unix_secs.cmp(&left.updated_at_unix_secs))
});
let total = items.len();
let limit = query.limit.max(1);
let items = items
.into_iter()
.skip(query.offset)
.take(limit)
.collect::<Vec<_>>();
Ok(StoredBackgroundTaskRunPage { items, total })
}
async fn list_events(
&self,
run_id: &str,
offset: usize,
limit: usize,
) -> Result<Vec<StoredBackgroundTaskEvent>, DataLayerError> {
let Some(events) = self
.index
.read()
.expect("background task repository lock")
.events_by_run
.get(run_id)
.cloned()
else {
return Ok(Vec::new());
};
let limit = limit.max(1);
Ok(events.into_iter().skip(offset).take(limit).collect())
}
async fn summarize_runs(&self) -> Result<BackgroundTaskSummary, DataLayerError> {
let runs = self
.index
.read()
.expect("background task repository lock")
.runs
.values()
.cloned()
.collect::<Vec<_>>();
let mut by_status = BTreeMap::new();
let mut by_kind = BTreeMap::new();
let mut running_count = 0_u64;
for run in runs {
*by_status
.entry(run.status.as_database().to_string())
.or_insert(0) += 1;
*by_kind
.entry(run.kind.as_database().to_string())
.or_insert(0) += 1;
if run.status == BackgroundTaskStatus::Running {
running_count += 1;
}
}
let total = by_status.values().copied().sum();
Ok(BackgroundTaskSummary {
total,
running_count,
by_status,
by_kind,
})
}
}
#[async_trait]
impl BackgroundTaskWriteRepository for InMemoryBackgroundTaskRepository {
async fn upsert_run(
&self,
run: UpsertBackgroundTaskRun,
) -> Result<StoredBackgroundTaskRun, DataLayerError> {
run.validate()?;
let stored = run.into_stored();
self.index
.write()
.expect("background task repository lock")
.runs
.insert(stored.id.clone(), stored.clone());
Ok(stored)
}
async fn request_cancel(
&self,
run_id: &str,
updated_at_unix_secs: u64,
) -> Result<bool, DataLayerError> {
let mut guard = self.index.write().expect("background task repository lock");
let Some(run) = guard.runs.get_mut(run_id) else {
return Ok(false);
};
run.cancel_requested = true;
run.updated_at_unix_secs = updated_at_unix_secs;
Ok(true)
}
async fn upsert_event(
&self,
event: UpsertBackgroundTaskEvent,
) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
event.validate()?;
let stored = event.into_stored();
let mut guard = self.index.write().expect("background task repository lock");
let entries = guard
.events_by_run
.entry(stored.run_id.clone())
.or_default();
if let Some(position) = entries.iter().position(|value| value.id == stored.id) {
entries[position] = stored.clone();
} else {
entries.push(stored.clone());
}
let mut seen = BTreeSet::new();
entries.retain(|entry| seen.insert(entry.id.clone()));
entries.sort_by(|left, right| {
left.created_at_unix_secs
.cmp(&right.created_at_unix_secs)
.then_with(|| left.id.cmp(&right.id))
});
Ok(stored)
}
}
@@ -0,0 +1,17 @@
mod memory;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::background_tasks::{
BackgroundTaskKind, BackgroundTaskListQuery, BackgroundTaskReadRepository,
BackgroundTaskRepository, BackgroundTaskStatus, BackgroundTaskSummary,
BackgroundTaskWriteRepository, StoredBackgroundTaskEvent, StoredBackgroundTaskRun,
StoredBackgroundTaskRunPage, UpsertBackgroundTaskEvent, UpsertBackgroundTaskRun,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlBackgroundTaskRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxBackgroundTaskRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteBackgroundTaskRepository;
pub use memory::InMemoryBackgroundTaskRepository;
@@ -0,0 +1,560 @@
use std::collections::BTreeMap;
use std::sync::RwLock;
use async_trait::async_trait;
use super::{
AdminBillingMutationOutcome, BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository,
PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, StoredBillingModelContext,
UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord,
};
use crate::DataLayerError;
type BillingContextKey = (String, String, Option<String>);
type BillingContextMap = BTreeMap<BillingContextKey, StoredBillingModelContext>;
#[derive(Debug, Default)]
pub struct InMemoryBillingReadRepository {
by_key: RwLock<BillingContextMap>,
gateway_configs_by_provider: RwLock<BTreeMap<String, PaymentGatewayConfigRecord>>,
billing_plans_by_id: RwLock<BTreeMap<String, BillingPlanRecord>>,
entitlements_by_id: RwLock<BTreeMap<String, UserPlanEntitlementRecord>>,
}
impl InMemoryBillingReadRepository {
pub fn seed<I>(items: I) -> Self
where
I: IntoIterator<Item = StoredBillingModelContext>,
{
let mut by_key = BTreeMap::new();
for item in items {
by_key.insert(
(
item.provider_id.clone(),
item.global_model_name.clone(),
item.provider_api_key_id.clone(),
),
item,
);
}
Self {
by_key: RwLock::new(by_key),
gateway_configs_by_provider: RwLock::new(BTreeMap::new()),
billing_plans_by_id: RwLock::new(BTreeMap::new()),
entitlements_by_id: RwLock::new(BTreeMap::new()),
}
}
}
fn current_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn billing_plan_from_input(
id: String,
input: &BillingPlanWriteInput,
created_at: u64,
) -> BillingPlanRecord {
BillingPlanRecord {
id,
title: input.title.clone(),
description: input.description.clone(),
price_amount: input.price_amount,
price_currency: input.price_currency.clone(),
duration_unit: input.duration_unit.clone(),
duration_value: input.duration_value,
enabled: input.enabled,
sort_order: input.sort_order,
max_active_per_user: input.max_active_per_user,
purchase_limit_scope: input.purchase_limit_scope.clone(),
entitlements_json: input.entitlements_json.clone(),
created_at_unix_secs: created_at,
updated_at_unix_secs: current_unix_secs(),
}
}
fn daily_quota_availability_from_entitlements(
entitlements: impl IntoIterator<Item = UserPlanEntitlementRecord>,
now: u64,
) -> UserDailyQuotaAvailabilityRecord {
let mut has_active_daily_quota = false;
let mut total_quota_usd = 0.0;
let used_usd = 0.0;
let mut remaining_usd = 0.0;
let mut allow_wallet_overage = true;
for entitlement in entitlements {
if entitlement.status != "active"
|| entitlement.starts_at_unix_secs > now
|| entitlement.expires_at_unix_secs <= now
{
continue;
}
let Some(items) = entitlement.entitlements_snapshot.as_array() else {
continue;
};
for item in items {
if item.get("type").and_then(serde_json::Value::as_str) != Some("daily_quota") {
continue;
}
let daily_quota_usd = item
.get("daily_quota_usd")
.and_then(serde_json::Value::as_f64)
.unwrap_or(0.0);
if !daily_quota_usd.is_finite() || daily_quota_usd <= 0.0 {
continue;
}
has_active_daily_quota = true;
total_quota_usd += daily_quota_usd;
remaining_usd += daily_quota_usd;
allow_wallet_overage &= item
.get("allow_wallet_overage")
.and_then(serde_json::Value::as_bool)
.unwrap_or(false);
}
}
UserDailyQuotaAvailabilityRecord {
has_active_daily_quota,
total_quota_usd,
used_usd,
remaining_usd,
allow_wallet_overage,
}
}
#[async_trait]
impl BillingReadRepository for InMemoryBillingReadRepository {
async fn find_model_context(
&self,
provider_id: &str,
provider_api_key_id: Option<&str>,
global_model_name: &str,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
let by_key = self.by_key.read().expect("billing repository lock");
if let Some(value) = find_context_by_provider_model_name(
&by_key,
provider_id,
provider_api_key_id,
global_model_name,
) {
return Ok(Some(value));
}
let key = (
provider_id.to_string(),
global_model_name.to_string(),
provider_api_key_id.map(ToOwned::to_owned),
);
if let Some(value) = by_key.get(&key) {
return Ok(Some(value.clone()));
}
if let Some(value) = by_key
.get(&(provider_id.to_string(), global_model_name.to_string(), None))
.cloned()
{
return Ok(Some(value));
}
Ok(by_key
.iter()
.find(|((stored_provider_id, stored_model_name, _), _)| {
stored_provider_id == provider_id && stored_model_name == global_model_name
})
.map(|(_, value)| value.clone()))
}
async fn find_model_context_by_model_id(
&self,
provider_id: &str,
provider_api_key_id: Option<&str>,
model_id: &str,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
let by_key = self.by_key.read().expect("billing repository lock");
if let Some(value) =
find_context_by_model_id_and_key(&by_key, provider_id, provider_api_key_id, model_id)
{
return Ok(Some(value));
}
if let Some(value) = find_context_by_model_id_and_key(&by_key, provider_id, None, model_id)
{
return Ok(Some(value));
}
Ok(by_key
.iter()
.find(|((stored_provider_id, _, _), value)| {
stored_provider_id == provider_id && value.model_id.as_deref() == Some(model_id)
})
.map(|(_, value)| value.clone()))
}
async fn find_payment_gateway_config(
&self,
provider: &str,
) -> Result<Option<PaymentGatewayConfigRecord>, DataLayerError> {
Ok(self
.gateway_configs_by_provider
.read()
.expect("billing repository lock")
.get(&provider.trim().to_ascii_lowercase())
.cloned())
}
async fn upsert_payment_gateway_config(
&self,
input: &PaymentGatewayConfigWriteInput,
) -> Result<AdminBillingMutationOutcome<PaymentGatewayConfigRecord>, DataLayerError> {
let provider = input.provider.trim().to_ascii_lowercase();
let now = current_unix_secs();
let mut configs = self
.gateway_configs_by_provider
.write()
.expect("billing repository lock");
let created_at = configs
.get(&provider)
.map(|value| value.created_at_unix_secs)
.unwrap_or(now);
let merchant_key_encrypted = if input.preserve_existing_secret {
configs
.get(&provider)
.and_then(|value| value.merchant_key_encrypted.clone())
} else {
input.merchant_key_encrypted.clone()
};
let record = PaymentGatewayConfigRecord {
provider: provider.clone(),
enabled: input.enabled,
endpoint_url: input.endpoint_url.clone(),
callback_base_url: input.callback_base_url.clone(),
merchant_id: input.merchant_id.clone(),
merchant_key_encrypted,
pay_currency: input.pay_currency.clone(),
usd_exchange_rate: input.usd_exchange_rate,
min_recharge_usd: input.min_recharge_usd,
channels_json: input.channels_json.clone(),
created_at_unix_secs: created_at,
updated_at_unix_secs: now,
};
configs.insert(provider, record.clone());
Ok(AdminBillingMutationOutcome::Applied(record))
}
async fn list_billing_plans(
&self,
include_disabled: bool,
) -> Result<Option<Vec<BillingPlanRecord>>, DataLayerError> {
let mut items = self
.billing_plans_by_id
.read()
.expect("billing repository lock")
.values()
.filter(|item| include_disabled || item.enabled)
.cloned()
.collect::<Vec<_>>();
items.sort_by(|left, right| {
left.sort_order
.cmp(&right.sort_order)
.then_with(|| left.price_amount.total_cmp(&right.price_amount))
.then_with(|| left.id.cmp(&right.id))
});
Ok(Some(items))
}
async fn find_billing_plan(
&self,
plan_id: &str,
) -> Result<Option<BillingPlanRecord>, DataLayerError> {
Ok(self
.billing_plans_by_id
.read()
.expect("billing repository lock")
.get(plan_id)
.cloned())
}
async fn create_billing_plan(
&self,
input: &BillingPlanWriteInput,
) -> Result<AdminBillingMutationOutcome<BillingPlanRecord>, DataLayerError> {
let id = uuid::Uuid::new_v4().to_string();
let record = billing_plan_from_input(id.clone(), input, current_unix_secs());
self.billing_plans_by_id
.write()
.expect("billing repository lock")
.insert(id, record.clone());
Ok(AdminBillingMutationOutcome::Applied(record))
}
async fn update_billing_plan(
&self,
plan_id: &str,
input: &BillingPlanWriteInput,
) -> Result<AdminBillingMutationOutcome<BillingPlanRecord>, DataLayerError> {
let mut plans = self
.billing_plans_by_id
.write()
.expect("billing repository lock");
let Some(existing) = plans.get(plan_id).cloned() else {
return Ok(AdminBillingMutationOutcome::NotFound);
};
let record =
billing_plan_from_input(plan_id.to_string(), input, existing.created_at_unix_secs);
plans.insert(plan_id.to_string(), record.clone());
Ok(AdminBillingMutationOutcome::Applied(record))
}
async fn set_billing_plan_enabled(
&self,
plan_id: &str,
enabled: bool,
) -> Result<AdminBillingMutationOutcome<BillingPlanRecord>, DataLayerError> {
let mut plans = self
.billing_plans_by_id
.write()
.expect("billing repository lock");
let Some(record) = plans.get_mut(plan_id) else {
return Ok(AdminBillingMutationOutcome::NotFound);
};
record.enabled = enabled;
record.updated_at_unix_secs = current_unix_secs();
Ok(AdminBillingMutationOutcome::Applied(record.clone()))
}
async fn delete_billing_plan(
&self,
plan_id: &str,
) -> Result<AdminBillingMutationOutcome<()>, DataLayerError> {
let mut plans = self
.billing_plans_by_id
.write()
.expect("billing repository lock");
if !plans.contains_key(plan_id) {
return Ok(AdminBillingMutationOutcome::NotFound);
}
let has_entitlements = self
.entitlements_by_id
.read()
.expect("billing repository lock")
.values()
.any(|item| item.plan_id == plan_id);
if has_entitlements {
return Ok(AdminBillingMutationOutcome::Invalid(
"套餐已有订单或权益,不能删除,请停用该套餐".to_string(),
));
}
plans.remove(plan_id);
Ok(AdminBillingMutationOutcome::Applied(()))
}
async fn list_user_plan_entitlements(
&self,
user_id: &str,
) -> Result<Option<Vec<UserPlanEntitlementRecord>>, DataLayerError> {
let now = current_unix_secs();
let mut items = self
.entitlements_by_id
.read()
.expect("billing repository lock")
.values()
.filter(|item| {
item.user_id == user_id
&& item.status == "active"
&& item.expires_at_unix_secs > now
})
.cloned()
.collect::<Vec<_>>();
items.sort_by_key(|item| item.expires_at_unix_secs);
Ok(Some(items))
}
async fn find_user_daily_quota_availability(
&self,
user_id: &str,
) -> Result<Option<UserDailyQuotaAvailabilityRecord>, DataLayerError> {
let now = current_unix_secs();
let entitlements = self
.entitlements_by_id
.read()
.expect("billing repository lock")
.values()
.filter(|item| item.user_id == user_id)
.cloned()
.collect::<Vec<_>>();
Ok(Some(daily_quota_availability_from_entitlements(
entitlements,
now,
)))
}
}
fn find_context_by_provider_model_name(
by_key: &BillingContextMap,
provider_id: &str,
provider_api_key_id: Option<&str>,
provider_model_name: &str,
) -> Option<StoredBillingModelContext> {
find_context_by_provider_model_name_and_key(
by_key,
provider_id,
provider_api_key_id,
provider_model_name,
)
.or_else(|| {
find_context_by_provider_model_name_and_key(by_key, provider_id, None, provider_model_name)
})
.or_else(|| {
by_key
.iter()
.find(|((stored_provider_id, _, _), value)| {
stored_provider_id == provider_id
&& value.model_provider_model_name.as_deref() == Some(provider_model_name)
})
.map(|(_, value)| value.clone())
})
}
fn find_context_by_provider_model_name_and_key(
by_key: &BillingContextMap,
provider_id: &str,
provider_api_key_id: Option<&str>,
provider_model_name: &str,
) -> Option<StoredBillingModelContext> {
by_key
.iter()
.find(|((stored_provider_id, _, stored_key_id), value)| {
stored_provider_id == provider_id
&& stored_key_id.as_deref() == provider_api_key_id
&& value.model_provider_model_name.as_deref() == Some(provider_model_name)
})
.map(|(_, value)| value.clone())
}
fn find_context_by_model_id_and_key(
by_key: &BillingContextMap,
provider_id: &str,
provider_api_key_id: Option<&str>,
model_id: &str,
) -> Option<StoredBillingModelContext> {
by_key
.iter()
.find(|((stored_provider_id, _, stored_key_id), value)| {
stored_provider_id == provider_id
&& stored_key_id.as_deref() == provider_api_key_id
&& value.model_id.as_deref() == Some(model_id)
})
.map(|(_, value)| value.clone())
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::InMemoryBillingReadRepository;
use crate::repository::billing::{BillingReadRepository, StoredBillingModelContext};
fn sample_context() -> StoredBillingModelContext {
StoredBillingModelContext::new(
"provider-1".to_string(),
Some("pay_as_you_go".to_string()),
Some("key-1".to_string()),
Some(json!({"openai:chat": 0.8})),
Some(60),
"global-model-1".to_string(),
"gpt-5".to_string(),
Some(json!({"streaming": true})),
Some(0.02),
Some(json!({"tiers":[{"up_to":null,"input_price_per_1m":3.0,"output_price_per_1m":15.0}]})),
Some("model-1".to_string()),
Some("gpt-5-upstream".to_string()),
None,
Some(0.01),
None,
)
.expect("billing context should build")
}
#[tokio::test]
async fn falls_back_to_provider_without_key_scope() {
let repository = InMemoryBillingReadRepository::seed(vec![sample_context()]);
let stored = repository
.find_model_context("provider-1", Some("key-2"), "gpt-5")
.await
.expect("lookup should succeed")
.expect("context should exist");
assert_eq!(stored.provider_id, "provider-1");
assert_eq!(stored.global_model_name, "gpt-5");
}
#[tokio::test]
async fn resolves_by_provider_model_name_before_global_name_collision() {
let global_named_context = StoredBillingModelContext::new(
"provider-1".to_string(),
Some("pay_as_you_go".to_string()),
Some("key-1".to_string()),
None,
Some(60),
"global-model-blank".to_string(),
"claude-sonnet-4-6".to_string(),
None,
None,
None,
Some("model-blank".to_string()),
Some("blank-upstream".to_string()),
None,
None,
None,
)
.expect("blank billing context should build");
let provider_priced_context = StoredBillingModelContext::new(
"provider-1".to_string(),
Some("pay_as_you_go".to_string()),
Some("key-1".to_string()),
None,
Some(60),
"global-model-priced".to_string(),
"claude-opus-4-6".to_string(),
None,
None,
Some(json!({"tiers":[{"up_to":null,"input_price_per_1m":3.0,"output_price_per_1m":15.0}]})),
Some("model-priced".to_string()),
Some("claude-sonnet-4-6".to_string()),
None,
None,
None,
)
.expect("priced billing context should build");
let repository = InMemoryBillingReadRepository::seed(vec![
global_named_context,
provider_priced_context,
]);
let stored = repository
.find_model_context("provider-1", Some("key-1"), "claude-sonnet-4-6")
.await
.expect("lookup should succeed")
.expect("context should exist");
assert_eq!(stored.global_model_name, "claude-opus-4-6");
assert_eq!(
stored.model_provider_model_name.as_deref(),
Some("claude-sonnet-4-6")
);
assert!(stored.default_tiered_pricing.is_some());
}
#[tokio::test]
async fn resolves_by_model_id() {
let repository = InMemoryBillingReadRepository::seed(vec![sample_context()]);
let stored = repository
.find_model_context_by_model_id("provider-1", Some("key-1"), "model-1")
.await
.expect("lookup should succeed")
.expect("context should exist");
assert_eq!(stored.global_model_name, "gpt-5");
assert_eq!(
stored.model_provider_model_name.as_deref(),
Some("gpt-5-upstream")
);
}
}
@@ -0,0 +1,9 @@
mod memory;
pub use aether_data_contracts::repository::billing::*;
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlBillingReadRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxBillingReadRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteBillingReadRepository;
pub use memory::InMemoryBillingReadRepository;
@@ -0,0 +1,695 @@
use std::sync::RwLock;
use async_trait::async_trait;
use super::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
StoredPoolKeyCandidateRowsQuery, StoredRequestedModelCandidateRowsQuery,
};
use crate::DataLayerError;
#[derive(Debug, Default)]
pub struct InMemoryMinimalCandidateSelectionReadRepository {
rows: RwLock<Vec<StoredMinimalCandidateSelectionRow>>,
}
impl InMemoryMinimalCandidateSelectionReadRepository {
pub fn seed<I>(rows: I) -> Self
where
I: IntoIterator<Item = StoredMinimalCandidateSelectionRow>,
{
Self {
rows: RwLock::new(rows.into_iter().collect()),
}
}
}
#[async_trait]
impl MinimalCandidateSelectionReadRepository for InMemoryMinimalCandidateSelectionReadRepository {
async fn list_for_exact_api_format(
&self,
api_format: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let api_format = api_format.trim();
let mut rows = self
.rows
.read()
.expect("candidate selection repository lock")
.iter()
.filter(|row| {
row.provider_is_active
&& row.endpoint_is_active
&& row.key_is_active
&& row.model_is_active
&& row.model_is_available
&& api_format_matches(&row.endpoint_api_format, api_format)
&& row.key_supports_api_format(api_format)
&& key_auth_channel_matches(row, api_format)
})
.cloned()
.collect::<Vec<_>>();
rows.sort_by(|left, right| {
left.provider_priority
.cmp(&right.provider_priority)
.then(left.key_internal_priority.cmp(&right.key_internal_priority))
.then(left.provider_id.cmp(&right.provider_id))
.then(left.endpoint_id.cmp(&right.endpoint_id))
.then(left.key_id.cmp(&right.key_id))
.then(left.model_id.cmp(&right.model_id))
});
Ok(rows)
}
async fn list_for_exact_api_format_and_global_model(
&self,
api_format: &str,
global_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let rows = self.list_for_exact_api_format(api_format).await?;
Ok(rows
.into_iter()
.filter(|row| row.global_model_name == global_model_name)
.collect())
}
async fn list_for_exact_api_format_and_requested_model(
&self,
api_format: &str,
requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
self.list_for_exact_api_format_and_requested_model_page(
&StoredRequestedModelCandidateRowsQuery {
api_format: api_format.to_string(),
requested_model_name: requested_model_name.to_string(),
offset: 0,
limit: u32::MAX,
},
)
.await
}
async fn list_for_exact_api_format_and_requested_model_page(
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let rows = self.list_for_exact_api_format(&query.api_format).await?;
let mut rows = rows
.into_iter()
.filter(|row| {
row_matches_requested_model(row, &query.requested_model_name, &query.api_format)
})
.collect::<Vec<_>>();
rows.sort_by(|left, right| {
left.global_model_name
.cmp(&right.global_model_name)
.then(left.provider_priority.cmp(&right.provider_priority))
.then(left.key_internal_priority.cmp(&right.key_internal_priority))
.then(left.provider_id.cmp(&right.provider_id))
.then(left.endpoint_id.cmp(&right.endpoint_id))
.then(left.key_id.cmp(&right.key_id))
.then(left.model_id.cmp(&right.model_id))
});
Ok(rows
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.collect())
}
async fn list_pool_key_rows_for_group(
&self,
query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
let mut rows = self
.list_for_exact_api_format(&query.api_format)
.await?
.into_iter()
.filter(|row| {
row.provider_id == query.provider_id
&& row.endpoint_id == query.endpoint_id
&& row.model_id == query.model_id
})
.collect::<Vec<_>>();
sort_pool_key_rows(&mut rows, &query.order);
Ok(rows
.into_iter()
.skip(query.offset as usize)
.take(query.limit as usize)
.collect())
}
async fn list_pool_key_rows_for_group_key_ids(
&self,
query: &StoredPoolKeyCandidateRowsByKeyIdsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
if query.key_ids.is_empty() {
return Ok(Vec::new());
}
let key_order = query
.key_ids
.iter()
.enumerate()
.map(|(index, key_id)| (key_id.as_str(), index))
.collect::<std::collections::BTreeMap<_, _>>();
let mut rows = self
.list_for_exact_api_format(&query.api_format)
.await?
.into_iter()
.filter(|row| {
row.provider_id == query.provider_id
&& row.endpoint_id == query.endpoint_id
&& row.model_id == query.model_id
&& key_order.contains_key(row.key_id.as_str())
})
.collect::<Vec<_>>();
rows.sort_by(|left, right| {
key_order
.get(left.key_id.as_str())
.cmp(&key_order.get(right.key_id.as_str()))
.then(left.key_id.cmp(&right.key_id))
});
Ok(rows)
}
}
fn sort_pool_key_rows(
rows: &mut [StoredMinimalCandidateSelectionRow],
order: &StoredPoolKeyCandidateOrder,
) {
rows.sort_by(|left, right| match order {
StoredPoolKeyCandidateOrder::LoadBalance { seed } => {
stable_pool_key_hash(seed.as_str(), left.key_id.as_str())
.cmp(&stable_pool_key_hash(seed.as_str(), right.key_id.as_str()))
.then(left.key_id.cmp(&right.key_id))
}
_ => left
.key_internal_priority
.cmp(&right.key_internal_priority)
.then(left.key_id.cmp(&right.key_id)),
});
}
fn stable_pool_key_hash(seed: &str, key_id: &str) -> u64 {
let mut hash = 0xcbf29ce484222325u64;
for byte in seed
.as_bytes()
.iter()
.copied()
.chain(std::iter::once(b':'))
.chain(key_id.as_bytes().iter().copied())
{
hash ^= u64::from(byte);
hash = hash.wrapping_mul(0x100000001b3);
}
hash
}
fn normalize_api_format(value: &str) -> String {
aether_ai_formats::normalize_api_format_alias(value)
}
fn api_format_matches(left: &str, right: &str) -> bool {
aether_ai_formats::api_format_alias_matches(left, right)
}
fn row_matches_requested_model(
row: &StoredMinimalCandidateSelectionRow,
requested_model_name: &str,
api_format: &str,
) -> bool {
(row_has_available_provider_model(row, api_format)
&& row.global_model_name == requested_model_name)
|| (row_default_provider_model_name_available(row, api_format)
&& row.model_provider_model_name == requested_model_name)
|| row
.model_provider_model_mappings
.as_ref()
.is_some_and(|mappings| {
mappings.iter().any(|mapping| {
mapping.api_formats.as_ref().is_none_or(|formats| {
formats
.iter()
.any(|value| api_format_scope_covers(value, api_format))
}) && mapping.endpoint_ids.as_ref().is_none_or(|endpoint_ids| {
endpoint_ids
.iter()
.any(|endpoint_id| endpoint_id == &row.endpoint_id)
}) && mapping.name == requested_model_name
})
})
}
fn row_has_available_provider_model(
row: &StoredMinimalCandidateSelectionRow,
api_format: &str,
) -> bool {
row_mapping_matches_scope(row, api_format)
|| row_default_provider_model_name_available(row, api_format)
}
fn row_default_provider_model_name_available(
row: &StoredMinimalCandidateSelectionRow,
api_format: &str,
) -> bool {
let Some(mappings) = row.model_provider_model_mappings.as_ref() else {
return true;
};
let mut has_explicit_default_mapping = false;
for mapping in mappings {
if mapping.name != row.model_provider_model_name {
continue;
}
has_explicit_default_mapping = true;
if mapping_scope_matches(mapping, row, api_format) {
return true;
}
}
!has_explicit_default_mapping
}
fn row_mapping_matches_scope(row: &StoredMinimalCandidateSelectionRow, api_format: &str) -> bool {
row.model_provider_model_mappings
.as_ref()
.is_some_and(|mappings| {
mappings
.iter()
.any(|mapping| mapping_scope_matches(mapping, row, api_format))
})
}
fn mapping_scope_matches(
mapping: &super::StoredProviderModelMapping,
row: &StoredMinimalCandidateSelectionRow,
api_format: &str,
) -> bool {
mapping.api_formats.as_ref().is_none_or(|formats| {
formats
.iter()
.any(|value| api_format_scope_covers(value, api_format))
}) && mapping.endpoint_ids.as_ref().is_none_or(|endpoint_ids| {
endpoint_ids
.iter()
.any(|endpoint_id| endpoint_id == &row.endpoint_id)
})
}
fn api_format_scope_covers(allowed: &str, requested: &str) -> bool {
aether_ai_formats::api_format_permission_covers(allowed, requested)
}
fn key_auth_channel_matches(row: &StoredMinimalCandidateSelectionRow, api_format: &str) -> bool {
let provider_type = row.provider_type.trim().to_ascii_lowercase();
let auth_type = row.key_auth_type.trim().to_ascii_lowercase();
let api_format = normalize_api_format(api_format);
match provider_type.as_str() {
"codex" => {
auth_type == "oauth"
&& matches!(
api_format.as_str(),
"openai:responses"
| "openai:responses:compact"
| "openai:search"
| "openai:image"
)
}
"chatgpt_web" => {
matches!(auth_type.as_str(), "oauth" | "bearer") && api_format == "openai:image"
}
"claude_code" => auth_type == "oauth" && api_format == "claude:messages",
"kiro" => {
matches!(auth_type.as_str(), "oauth" | "bearer") && api_format == "claude:messages"
}
"gemini_cli" | "antigravity" => {
auth_type == "oauth" && api_format == "gemini:generate_content"
}
"grok" => {
auth_type == "oauth"
&& matches!(
api_format.as_str(),
"openai:chat" | "openai:responses" | "claude:messages" | "openai:image"
)
}
"windsurf" => {
matches!(auth_type.as_str(), "oauth" | "api_key" | "bearer")
&& api_format == "openai:chat"
}
"vertex_ai" => {
(auth_type == "api_key"
&& matches!(
api_format.as_str(),
"gemini:generate_content" | "gemini:embedding"
))
|| (matches!(auth_type.as_str(), "service_account" | "vertex_ai")
&& matches!(
api_format.as_str(),
"claude:messages" | "gemini:generate_content" | "gemini:embedding"
))
}
_ => auth_type != "oauth",
}
}
#[cfg(test)]
mod tests {
use super::InMemoryMinimalCandidateSelectionReadRepository;
use crate::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
StoredRequestedModelCandidateRowsQuery,
};
fn sample_row(
provider_id: &str,
api_format: &str,
global_model_name: &str,
provider_priority: i32,
) -> StoredMinimalCandidateSelectionRow {
StoredMinimalCandidateSelectionRow {
provider_id: provider_id.to_string(),
provider_name: provider_id.to_string(),
provider_type: "custom".to_string(),
provider_priority,
provider_is_active: true,
endpoint_id: format!("endpoint-{provider_id}"),
endpoint_api_format: api_format.to_string(),
endpoint_api_family: Some("openai".to_string()),
endpoint_kind: Some("chat".to_string()),
endpoint_is_active: true,
key_id: format!("key-{provider_id}"),
key_name: "prod".to_string(),
key_auth_type: "api_key".to_string(),
key_is_active: true,
key_api_formats: Some(vec![api_format.to_string()]),
key_allowed_models: None,
key_capabilities: None,
key_internal_priority: 50,
key_global_priority_by_format: None,
model_id: format!("model-{provider_id}"),
global_model_id: "global-model-1".to_string(),
global_model_name: global_model_name.to_string(),
global_model_mappings: None,
global_model_supports_streaming: Some(true),
model_provider_model_name: global_model_name.to_string(),
model_provider_model_mappings: None,
model_supports_streaming: None,
model_is_active: true,
model_is_available: true,
}
}
#[tokio::test]
async fn filters_by_exact_api_format_and_global_model() {
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
sample_row("provider-2", "openai:chat", "gpt-4.1", 20),
sample_row("provider-1", "openai:chat", "gpt-4.1", 10),
sample_row("provider-3", "openai:responses", "gpt-4.1", 5),
sample_row("provider-4", "openai:chat", "gpt-4.1-mini", 1),
]);
let rows = repository
.list_for_exact_api_format_and_global_model("openai:chat", "gpt-4.1")
.await
.expect("list should succeed");
assert_eq!(rows.len(), 2);
assert_eq!(rows[0].provider_id, "provider-1");
assert_eq!(rows[1].provider_id, "provider-2");
}
#[tokio::test]
async fn filters_by_exact_api_format_and_requested_model_aliases() {
let mut mapped = sample_row("provider-1", "openai:chat", "gpt-4.1", 10);
mapped.model_provider_model_name = "provider-gpt-4.1".to_string();
mapped.model_provider_model_mappings = Some(vec![StoredProviderModelMapping {
name: "alias-gpt-4.1".to_string(),
priority: 0,
api_formats: Some(vec!["openai:chat".to_string()]),
endpoint_ids: None,
operations: None,
}]);
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
mapped,
sample_row("provider-2", "openai:chat", "gpt-4.1-mini", 20),
]);
let rows = repository
.list_for_exact_api_format_and_requested_model("openai:chat", "alias-gpt-4.1")
.await
.expect("list should succeed");
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].provider_id, "provider-1");
}
#[tokio::test]
async fn search_uses_responses_key_and_model_permissions_with_exact_endpoint_identity() {
let mut search = sample_row(
"provider-search",
"openai:search",
"global-search-model",
10,
);
search.provider_type = "codex".to_string();
search.key_auth_type = "oauth".to_string();
search.key_api_formats = Some(vec!["openai:responses".to_string()]);
search.model_provider_model_name = "upstream-search-model".to_string();
search.model_provider_model_mappings = Some(vec![StoredProviderModelMapping {
name: "gpt-5.6-sol".to_string(),
priority: 0,
api_formats: Some(vec!["openai:responses".to_string()]),
endpoint_ids: None,
operations: None,
}]);
let mut responses = search.clone();
responses.provider_id = "provider-responses".to_string();
responses.endpoint_id = "endpoint-responses".to_string();
responses.endpoint_api_format = "openai:responses".to_string();
responses.key_id = "key-responses".to_string();
responses.model_id = "model-responses".to_string();
let repository =
InMemoryMinimalCandidateSelectionReadRepository::seed(vec![responses, search]);
let rows = repository
.list_for_exact_api_format_and_requested_model("openai:search", "gpt-5.6-sol")
.await
.expect("Search candidate should load");
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].endpoint_api_format, "openai:search");
assert_eq!(
rows[0].key_api_formats,
Some(vec!["openai:responses".to_string()])
);
}
#[tokio::test]
async fn includes_grok_oauth_rows_for_chat_models() {
let mut row = sample_row(
"provider-grok",
"openai:chat",
"grok-4.20-0309-non-reasoning",
10,
);
row.provider_type = "grok".to_string();
row.provider_name = "grok".to_string();
row.key_auth_type = "oauth".to_string();
row.key_api_formats = Some(vec![
"openai:chat".to_string(),
"openai:responses".to_string(),
"claude:messages".to_string(),
"openai:image".to_string(),
]);
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![row]);
let rows = repository
.list_for_exact_api_format("openai:chat")
.await
.expect("list should succeed");
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].provider_type, "grok");
assert_eq!(rows[0].global_model_name, "grok-4.20-0309-non-reasoning");
}
#[tokio::test]
async fn requested_model_filter_respects_endpoint_scoped_default_mapping() {
let mut selected = sample_row("provider-1", "openai:chat", "deepseek-v4-pro", 10);
selected.endpoint_id = "endpoint-openai".to_string();
selected.model_provider_model_name = "deepseek-v4-pro".to_string();
selected.model_provider_model_mappings = Some(vec![StoredProviderModelMapping {
name: "deepseek-v4-pro".to_string(),
priority: 1,
api_formats: None,
endpoint_ids: Some(vec!["endpoint-openai".to_string()]),
operations: None,
}]);
let mut scoped_out = selected.clone();
scoped_out.provider_id = "provider-2".to_string();
scoped_out.endpoint_id = "endpoint-claude".to_string();
scoped_out.endpoint_api_format = "claude:messages".to_string();
scoped_out.key_id = "key-provider-2".to_string();
scoped_out.key_api_formats = Some(vec!["claude:messages".to_string()]);
let repository =
InMemoryMinimalCandidateSelectionReadRepository::seed(vec![scoped_out, selected]);
let rows = repository
.list_for_exact_api_format_and_requested_model("claude:messages", "deepseek-v4-pro")
.await
.expect("list should succeed");
assert!(rows.is_empty());
let rows = repository
.list_for_exact_api_format_and_requested_model("openai:chat", "deepseek-v4-pro")
.await
.expect("list should succeed");
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].endpoint_id, "endpoint-openai");
}
#[tokio::test]
async fn requested_model_page_returns_requested_slice_only() {
let mut rows = Vec::new();
for index in 0..5 {
let mut row = sample_row(&format!("provider-{index}"), "openai:chat", "gpt-5", index);
row.key_internal_priority = index;
rows.push(row);
}
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(rows);
let page = repository
.list_for_exact_api_format_and_requested_model_page(
&StoredRequestedModelCandidateRowsQuery {
api_format: "openai:chat".to_string(),
requested_model_name: "gpt-5".to_string(),
offset: 2,
limit: 2,
},
)
.await
.expect("page should load");
assert_eq!(
page.iter()
.map(|row| row.provider_id.as_str())
.collect::<Vec<_>>(),
vec!["provider-2", "provider-3"]
);
}
#[tokio::test]
async fn allows_chatgpt_web_oauth_and_bearer_for_openai_image_only() {
let mut oauth = sample_row("chatgpt-web-oauth", "openai:image", "gpt-image-2", 10);
oauth.provider_type = "chatgpt_web".to_string();
oauth.key_auth_type = "oauth".to_string();
let mut bearer = sample_row("chatgpt-web-bearer", "openai:image", "gpt-image-2", 20);
bearer.provider_type = "chatgpt_web".to_string();
bearer.key_auth_type = "bearer".to_string();
let mut api_key = sample_row("chatgpt-web-api-key", "openai:image", "gpt-image-2", 30);
api_key.provider_type = "chatgpt_web".to_string();
api_key.key_auth_type = "api_key".to_string();
let mut responses = sample_row("chatgpt-web-responses", "openai:responses", "gpt-5", 40);
responses.provider_type = "chatgpt_web".to_string();
responses.key_auth_type = "oauth".to_string();
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
oauth, bearer, api_key, responses,
]);
let rows = repository
.list_for_exact_api_format_and_requested_model("openai:image", "gpt-image-2")
.await
.expect("list should succeed");
assert_eq!(
rows.iter()
.map(|row| row.provider_id.as_str())
.collect::<Vec<_>>(),
vec!["chatgpt-web-oauth", "chatgpt-web-bearer"]
);
}
#[tokio::test]
async fn allows_windsurf_managed_keys_for_openai_chat_only() {
let mut oauth = sample_row("windsurf-oauth", "openai:chat", "gpt-5", 10);
oauth.provider_type = "windsurf".to_string();
oauth.key_auth_type = "oauth".to_string();
let mut api_key = sample_row("windsurf-api-key", "openai:chat", "gpt-5", 20);
api_key.provider_type = "windsurf".to_string();
api_key.key_auth_type = "api_key".to_string();
let mut responses = sample_row("windsurf-responses", "openai:responses", "gpt-5", 30);
responses.provider_type = "windsurf".to_string();
responses.key_auth_type = "oauth".to_string();
let repository =
InMemoryMinimalCandidateSelectionReadRepository::seed(vec![oauth, api_key, responses]);
let rows = repository
.list_for_exact_api_format_and_requested_model("openai:chat", "gpt-5")
.await
.expect("list should succeed");
assert_eq!(
rows.iter()
.map(|row| row.provider_id.as_str())
.collect::<Vec<_>>(),
vec!["windsurf-oauth", "windsurf-api-key"]
);
}
#[tokio::test]
async fn filters_by_exact_api_format_only() {
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
sample_row("provider-2", "openai:chat", "gpt-4.1", 20),
sample_row("provider-1", "openai:chat", "gpt-4.1-mini", 10),
sample_row("provider-3", "openai:responses", "gpt-4.1", 5),
]);
let rows = repository
.list_for_exact_api_format("openai:chat")
.await
.expect("list should succeed");
assert_eq!(rows.len(), 2);
assert_eq!(rows[0].provider_id, "provider-1");
assert_eq!(rows[1].provider_id, "provider-2");
}
#[tokio::test]
async fn list_pool_key_rows_for_group_returns_requested_page_only() {
let mut rows = Vec::new();
for index in 0..5 {
let mut row = sample_row("provider-pool", "openai:chat", "gpt-5", 10);
row.endpoint_id = "endpoint-pool".to_string();
row.model_id = "model-pool".to_string();
row.key_id = format!("key-{index}");
row.key_internal_priority = index;
rows.push(row);
}
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(rows);
let page = repository
.list_pool_key_rows_for_group(&StoredPoolKeyCandidateRowsQuery {
api_format: "openai:chat".to_string(),
provider_id: "provider-pool".to_string(),
endpoint_id: "endpoint-pool".to_string(),
model_id: "model-pool".to_string(),
selected_provider_model_name: "gpt-5".to_string(),
order: StoredPoolKeyCandidateOrder::InternalPriority,
offset: 2,
limit: 2,
})
.await
.expect("pool key page should load");
assert_eq!(
page.iter()
.map(|row| row.key_id.as_str())
.collect::<Vec<_>>(),
vec!["key-2", "key-3"]
);
}
}
@@ -0,0 +1,16 @@
mod memory;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, MinimalCandidateSelectionRepository,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder,
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlMinimalCandidateSelectionReadRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxMinimalCandidateSelectionReadRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteMinimalCandidateSelectionReadRepository;
pub use memory::InMemoryMinimalCandidateSelectionReadRepository;
@@ -0,0 +1,767 @@
use std::collections::{BTreeMap, BTreeSet};
use std::sync::RwLock;
use async_trait::async_trait;
use super::{
request_candidate_lifecycle_would_regress, PublicHealthStatusCount, PublicHealthTimelineBucket,
RequestCandidateReadRepository, RequestCandidateStatus, RequestCandidateWriteRepository,
StoredRequestCandidate, UpsertRequestCandidateRecord,
};
use crate::DataLayerError;
fn merge_extra_data(
existing: Option<serde_json::Value>,
overlay: Option<serde_json::Value>,
) -> Option<serde_json::Value> {
match (existing, overlay) {
(
Some(serde_json::Value::Object(mut existing_object)),
Some(serde_json::Value::Object(overlay_object)),
) => {
existing_object.extend(overlay_object);
Some(serde_json::Value::Object(existing_object))
}
(_existing, Some(overlay)) => Some(overlay),
(existing, None) => existing,
}
}
#[derive(Debug, Default)]
pub struct InMemoryRequestCandidateRepository {
by_id: RwLock<BTreeMap<String, StoredRequestCandidate>>,
}
impl InMemoryRequestCandidateRepository {
pub fn seed<I>(items: I) -> Self
where
I: IntoIterator<Item = StoredRequestCandidate>,
{
let mut by_id = BTreeMap::new();
for item in items {
by_id.insert(item.id.clone(), item);
}
Self {
by_id: RwLock::new(by_id),
}
}
}
#[async_trait]
impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
async fn list_by_request_id(
&self,
request_id: &str,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
let mut rows = self
.by_id
.read()
.expect("request candidate repository lock")
.values()
.filter(|row| row.request_id == request_id)
.cloned()
.collect::<Vec<_>>();
rows.sort_by(|left, right| {
left.candidate_index
.cmp(&right.candidate_index)
.then(left.retry_index.cmp(&right.retry_index))
.then(left.created_at_unix_ms.cmp(&right.created_at_unix_ms))
});
Ok(rows)
}
async fn list_recent(
&self,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let mut rows = self
.by_id
.read()
.expect("request candidate repository lock")
.values()
.cloned()
.collect::<Vec<_>>();
rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms));
rows.truncate(limit);
Ok(rows)
}
async fn list_by_provider_id(
&self,
provider_id: &str,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let mut rows = self
.by_id
.read()
.expect("request candidate repository lock")
.values()
.filter(|row| row.provider_id.as_deref() == Some(provider_id))
.cloned()
.collect::<Vec<_>>();
rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms));
rows.truncate(limit);
Ok(rows)
}
async fn list_finalized_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
since_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
if endpoint_ids.is_empty() || limit == 0 {
return Ok(Vec::new());
}
let endpoint_ids = endpoint_ids.iter().cloned().collect::<BTreeSet<_>>();
let mut rows = self
.by_id
.read()
.expect("request candidate repository lock")
.values()
.filter(|row| {
row.endpoint_id
.as_ref()
.is_some_and(|endpoint_id| endpoint_ids.contains(endpoint_id))
&& row.created_at_unix_ms >= since_unix_secs * 1000
&& matches!(
row.status,
RequestCandidateStatus::Success
| RequestCandidateStatus::Failed
| RequestCandidateStatus::Skipped
)
})
.cloned()
.collect::<Vec<_>>();
rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms));
rows.truncate(limit);
Ok(rows)
}
async fn count_finalized_statuses_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
since_unix_secs: u64,
) -> Result<Vec<PublicHealthStatusCount>, DataLayerError> {
if endpoint_ids.is_empty() {
return Ok(Vec::new());
}
let endpoint_ids = endpoint_ids.iter().cloned().collect::<BTreeSet<_>>();
let mut counts = BTreeMap::<(String, &'static str), u64>::new();
for row in self
.by_id
.read()
.expect("request candidate repository lock")
.values()
{
let Some(endpoint_id) = row.endpoint_id.as_ref() else {
continue;
};
if !endpoint_ids.contains(endpoint_id)
|| row.created_at_unix_ms < since_unix_secs * 1000
{
continue;
}
if !matches!(
row.status,
RequestCandidateStatus::Success
| RequestCandidateStatus::Failed
| RequestCandidateStatus::Skipped
) {
continue;
}
let status_key = match row.status {
RequestCandidateStatus::Success => "success",
RequestCandidateStatus::Failed => "failed",
RequestCandidateStatus::Skipped => "skipped",
_ => continue,
};
*counts.entry((endpoint_id.clone(), status_key)).or_insert(0) += 1;
}
Ok(counts
.into_iter()
.map(|((endpoint_id, status_key), count)| {
let status = match status_key {
"success" => RequestCandidateStatus::Success,
"failed" => RequestCandidateStatus::Failed,
"skipped" => RequestCandidateStatus::Skipped,
_ => unreachable!("filtered status should stay finalized"),
};
PublicHealthStatusCount {
endpoint_id,
status,
count,
}
})
.collect())
}
async fn aggregate_finalized_timeline_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
since_unix_secs: u64,
until_unix_secs: u64,
segments: u32,
) -> Result<Vec<PublicHealthTimelineBucket>, DataLayerError> {
if endpoint_ids.is_empty() || segments == 0 || until_unix_secs < since_unix_secs {
return Ok(Vec::new());
}
let endpoint_ids = endpoint_ids.iter().cloned().collect::<BTreeSet<_>>();
let span_ms = until_unix_secs.saturating_sub(since_unix_secs) * 1000;
let mut buckets = BTreeMap::<(String, u32), PublicHealthTimelineBucket>::new();
for row in self
.by_id
.read()
.expect("request candidate repository lock")
.values()
{
let Some(endpoint_id) = row.endpoint_id.as_ref() else {
continue;
};
if !endpoint_ids.contains(endpoint_id)
|| row.created_at_unix_ms < since_unix_secs * 1000
|| row.created_at_unix_ms > until_unix_secs * 1000
{
continue;
}
if !matches!(
row.status,
RequestCandidateStatus::Success
| RequestCandidateStatus::Failed
| RequestCandidateStatus::Skipped
) {
continue;
}
let segment_idx = if span_ms == 0 {
0
} else {
let offset = row
.created_at_unix_ms
.saturating_sub(since_unix_secs * 1000);
let idx = ((offset as u128) * (segments as u128) / (span_ms as u128)) as u32;
idx.min(segments.saturating_sub(1))
};
let bucket = buckets
.entry((endpoint_id.clone(), segment_idx))
.or_insert_with(|| PublicHealthTimelineBucket {
endpoint_id: endpoint_id.clone(),
segment_idx,
total_count: 0,
success_count: 0,
failed_count: 0,
min_created_at_unix_ms: None,
max_created_at_unix_ms: None,
});
bucket.total_count += 1;
if row.status == RequestCandidateStatus::Success {
bucket.success_count += 1;
} else if row.status == RequestCandidateStatus::Failed {
bucket.failed_count += 1;
}
bucket.min_created_at_unix_ms = Some(
bucket
.min_created_at_unix_ms
.map(|value| value.min(row.created_at_unix_ms))
.unwrap_or(row.created_at_unix_ms),
);
bucket.max_created_at_unix_ms = Some(
bucket
.max_created_at_unix_ms
.map(|value| value.max(row.created_at_unix_ms))
.unwrap_or(row.created_at_unix_ms),
);
}
Ok(buckets.into_values().collect())
}
}
#[async_trait]
impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
async fn upsert(
&self,
candidate: UpsertRequestCandidateRecord,
) -> Result<StoredRequestCandidate, DataLayerError> {
candidate.validate()?;
let mut by_id = self
.by_id
.write()
.expect("request candidate repository lock");
let existing = by_id
.values()
.find(|row| {
row.request_id == candidate.request_id
&& row.candidate_index == candidate.candidate_index
&& row.retry_index == candidate.retry_index
})
.cloned();
let preserve_existing_lifecycle = existing.as_ref().is_some_and(|row| {
request_candidate_lifecycle_would_regress(row.status, candidate.status)
});
let merged_status = if preserve_existing_lifecycle {
existing
.as_ref()
.map(|row| row.status)
.unwrap_or(candidate.status)
} else {
candidate.status
};
let created_at_unix_ms = existing
.as_ref()
.map(|row| row.created_at_unix_ms)
.or(candidate.created_at_unix_ms)
.or(candidate.started_at_unix_ms)
.or(candidate.finished_at_unix_ms)
.unwrap_or_default();
let stored = StoredRequestCandidate {
id: existing
.as_ref()
.map(|row| row.id.clone())
.unwrap_or_else(|| candidate.id.clone()),
request_id: candidate.request_id.clone(),
user_id: candidate
.user_id
.or_else(|| existing.as_ref().and_then(|row| row.user_id.clone())),
api_key_id: candidate
.api_key_id
.or_else(|| existing.as_ref().and_then(|row| row.api_key_id.clone())),
username: candidate
.username
.or_else(|| existing.as_ref().and_then(|row| row.username.clone())),
api_key_name: candidate
.api_key_name
.or_else(|| existing.as_ref().and_then(|row| row.api_key_name.clone())),
candidate_index: candidate.candidate_index,
retry_index: candidate.retry_index,
provider_id: candidate
.provider_id
.or_else(|| existing.as_ref().and_then(|row| row.provider_id.clone())),
endpoint_id: candidate
.endpoint_id
.or_else(|| existing.as_ref().and_then(|row| row.endpoint_id.clone())),
key_id: candidate
.key_id
.or_else(|| existing.as_ref().and_then(|row| row.key_id.clone())),
status: merged_status,
skip_reason: candidate
.skip_reason
.or_else(|| existing.as_ref().and_then(|row| row.skip_reason.clone())),
is_cached: candidate
.is_cached
.unwrap_or_else(|| existing.as_ref().map(|row| row.is_cached).unwrap_or(false)),
status_code: if preserve_existing_lifecycle {
existing.as_ref().and_then(|row| row.status_code)
} else {
candidate
.status_code
.or_else(|| existing.as_ref().and_then(|row| row.status_code))
},
error_type: if preserve_existing_lifecycle {
existing.as_ref().and_then(|row| row.error_type.clone())
} else {
candidate
.error_type
.or_else(|| existing.as_ref().and_then(|row| row.error_type.clone()))
},
error_message: if preserve_existing_lifecycle {
existing.as_ref().and_then(|row| row.error_message.clone())
} else {
candidate
.error_message
.or_else(|| existing.as_ref().and_then(|row| row.error_message.clone()))
},
latency_ms: if preserve_existing_lifecycle {
existing.as_ref().and_then(|row| row.latency_ms)
} else {
candidate
.latency_ms
.or_else(|| existing.as_ref().and_then(|row| row.latency_ms))
},
concurrent_requests: candidate
.concurrent_requests
.or_else(|| existing.as_ref().and_then(|row| row.concurrent_requests)),
extra_data: merge_extra_data(
existing.as_ref().and_then(|row| row.extra_data.clone()),
candidate.extra_data,
),
required_capabilities: candidate.required_capabilities.or_else(|| {
existing
.as_ref()
.and_then(|row| row.required_capabilities.clone())
}),
created_at_unix_ms,
started_at_unix_ms: candidate
.started_at_unix_ms
.or_else(|| existing.as_ref().and_then(|row| row.started_at_unix_ms)),
finished_at_unix_ms: if preserve_existing_lifecycle {
existing.as_ref().and_then(|row| row.finished_at_unix_ms)
} else {
candidate
.finished_at_unix_ms
.or_else(|| existing.as_ref().and_then(|row| row.finished_at_unix_ms))
},
};
by_id.insert(stored.id.clone(), stored.clone());
Ok(stored)
}
async fn delete_created_before(
&self,
created_before_unix_secs: u64,
limit: usize,
) -> Result<usize, DataLayerError> {
if limit == 0 {
return Ok(0);
}
let mut by_id = self
.by_id
.write()
.expect("request candidate repository lock");
let mut ids = by_id
.values()
.filter(|row| row.created_at_unix_ms < created_before_unix_secs * 1000)
.map(|row| (row.created_at_unix_ms, row.id.clone()))
.collect::<Vec<_>>();
ids.sort();
let mut deleted = 0usize;
for (_, id) in ids.into_iter().take(limit) {
if by_id.remove(&id).is_some() {
deleted += 1;
}
}
Ok(deleted)
}
}
#[cfg(test)]
mod tests {
use super::InMemoryRequestCandidateRepository;
use crate::repository::candidates::{
RequestCandidateReadRepository, RequestCandidateStatus, RequestCandidateWriteRepository,
StoredRequestCandidate, UpsertRequestCandidateRecord,
};
use serde_json::json;
fn sample_candidate(
id: &str,
request_id: &str,
created_at_unix_ms: i64,
) -> StoredRequestCandidate {
StoredRequestCandidate::new(
id.to_string(),
request_id.to_string(),
Some("user-1".to_string()),
Some("api-key-1".to_string()),
Some("alice".to_string()),
Some("default".to_string()),
0,
0,
Some("provider-1".to_string()),
Some("endpoint-1".to_string()),
Some("key-1".to_string()),
RequestCandidateStatus::Success,
None,
false,
Some(200),
None,
None,
Some(10),
Some(1),
None,
None,
created_at_unix_ms,
Some(created_at_unix_ms),
Some(created_at_unix_ms + 1),
)
.expect("candidate should build")
}
#[tokio::test]
async fn lists_request_candidates_by_request_id_in_candidate_order() {
let repository = InMemoryRequestCandidateRepository::seed(vec![
sample_candidate("cand-2", "req-1", 200),
sample_candidate("cand-1", "req-1", 100),
sample_candidate("cand-3", "req-2", 300),
]);
let rows = repository
.list_by_request_id("req-1")
.await
.expect("list should succeed");
assert_eq!(rows.len(), 2);
assert_eq!(rows[0].request_id, "req-1");
assert_eq!(rows[1].request_id, "req-1");
}
#[tokio::test]
async fn lists_recent_request_candidates_in_descending_created_order() {
let repository = InMemoryRequestCandidateRepository::seed(vec![
sample_candidate("cand-1", "req-1", 100),
sample_candidate("cand-2", "req-2", 200),
]);
let rows = repository
.list_recent(10)
.await
.expect("list recent should succeed");
assert_eq!(rows.len(), 2);
assert_eq!(rows[0].id, "cand-2");
assert_eq!(rows[1].id, "cand-1");
}
#[tokio::test]
async fn aggregates_finalized_health_data_by_endpoint_ids() {
let repository = InMemoryRequestCandidateRepository::seed(vec![
sample_candidate("cand-1", "req-1", 100_000),
sample_candidate("cand-2", "req-2", 200_000),
]);
let counts = repository
.count_finalized_statuses_by_endpoint_ids_since(&["endpoint-1".to_string()], 0)
.await
.expect("count should succeed");
assert_eq!(counts.len(), 1);
assert_eq!(counts[0].endpoint_id, "endpoint-1");
assert_eq!(counts[0].status, RequestCandidateStatus::Success);
assert_eq!(counts[0].count, 2);
let timeline = repository
.aggregate_finalized_timeline_by_endpoint_ids_since(
&["endpoint-1".to_string()],
0,
300,
3,
)
.await
.expect("timeline should succeed");
assert_eq!(timeline.len(), 2);
let attempts = repository
.list_finalized_by_endpoint_ids_since(&["endpoint-1".to_string()], 0, 1)
.await
.expect("attempt list should succeed");
assert_eq!(attempts.len(), 1);
assert_eq!(attempts[0].id, "cand-2");
}
#[tokio::test]
async fn upsert_writes_and_updates_request_candidate() {
let repository = InMemoryRequestCandidateRepository::default();
let created = repository
.upsert(UpsertRequestCandidateRecord {
id: "cand-1".to_string(),
request_id: "req-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("api-key-1".to_string()),
username: Some("alice".to_string()),
api_key_name: Some("default".to_string()),
candidate_index: 0,
retry_index: 0,
provider_id: Some("provider-1".to_string()),
endpoint_id: Some("endpoint-1".to_string()),
key_id: Some("key-1".to_string()),
status: RequestCandidateStatus::Available,
skip_reason: None,
is_cached: Some(false),
status_code: None,
error_type: None,
error_message: None,
latency_ms: None,
concurrent_requests: None,
extra_data: Some(json!({
"execution_strategy": "local_cross_format",
"provider_name": "primary",
})),
required_capabilities: None,
created_at_unix_ms: Some(100),
started_at_unix_ms: None,
finished_at_unix_ms: None,
})
.await
.expect("create should succeed");
assert_eq!(created.id, "cand-1");
assert_eq!(created.status, RequestCandidateStatus::Available);
let updated = repository
.upsert(UpsertRequestCandidateRecord {
id: "cand-1-replacement".to_string(),
request_id: "req-1".to_string(),
user_id: None,
api_key_id: None,
username: None,
api_key_name: None,
candidate_index: 0,
retry_index: 0,
provider_id: None,
endpoint_id: None,
key_id: None,
status: RequestCandidateStatus::Success,
skip_reason: None,
is_cached: None,
status_code: Some(200),
error_type: None,
error_message: None,
latency_ms: Some(25),
concurrent_requests: Some(2),
extra_data: Some(json!({
"provider_api_format": "openai:responses",
"provider_name": "updated",
})),
required_capabilities: None,
created_at_unix_ms: None,
started_at_unix_ms: Some(101),
finished_at_unix_ms: Some(102),
})
.await
.expect("update should succeed");
assert_eq!(updated.id, "cand-1");
assert_eq!(updated.status, RequestCandidateStatus::Success);
assert_eq!(updated.status_code, Some(200));
assert_eq!(updated.latency_ms, Some(25));
assert_eq!(
updated
.extra_data
.as_ref()
.and_then(|value| value.get("execution_strategy")),
Some(&json!("local_cross_format"))
);
assert_eq!(
updated
.extra_data
.as_ref()
.and_then(|value| value.get("provider_api_format")),
Some(&json!("openai:responses"))
);
assert_eq!(
updated
.extra_data
.as_ref()
.and_then(|value| value.get("provider_name")),
Some(&json!("updated"))
);
assert_eq!(updated.started_at_unix_ms, Some(101));
}
#[tokio::test]
async fn upsert_keeps_terminal_candidate_state_when_streaming_arrives_late() {
let existing = StoredRequestCandidate::new(
"cand-1".to_string(),
"req-1".to_string(),
Some("user-1".to_string()),
Some("api-key-1".to_string()),
Some("alice".to_string()),
Some("default".to_string()),
0,
0,
Some("provider-1".to_string()),
Some("endpoint-1".to_string()),
Some("key-1".to_string()),
RequestCandidateStatus::Failed,
None,
false,
Some(503),
Some("upstream_error".to_string()),
Some("retryable upstream failure".to_string()),
Some(45),
Some(1),
Some(json!({"terminal": true})),
None,
100,
Some(101),
Some(145),
)
.expect("candidate should build");
let repository = InMemoryRequestCandidateRepository::seed(vec![existing]);
let updated = repository
.upsert(UpsertRequestCandidateRecord {
id: "cand-1-late".to_string(),
request_id: "req-1".to_string(),
user_id: None,
api_key_id: None,
username: None,
api_key_name: None,
candidate_index: 0,
retry_index: 0,
provider_id: None,
endpoint_id: None,
key_id: None,
status: RequestCandidateStatus::Streaming,
skip_reason: None,
is_cached: None,
status_code: Some(200),
error_type: None,
error_message: None,
latency_ms: Some(9_999),
concurrent_requests: Some(2),
extra_data: Some(json!({"late": true})),
required_capabilities: None,
created_at_unix_ms: None,
started_at_unix_ms: Some(102),
finished_at_unix_ms: None,
})
.await
.expect("late update should succeed");
assert_eq!(updated.id, "cand-1");
assert_eq!(updated.status, RequestCandidateStatus::Failed);
assert_eq!(updated.status_code, Some(503));
assert_eq!(updated.error_type.as_deref(), Some("upstream_error"));
assert_eq!(
updated.error_message.as_deref(),
Some("retryable upstream failure")
);
assert_eq!(updated.latency_ms, Some(45));
assert_eq!(updated.concurrent_requests, Some(2));
assert_eq!(updated.finished_at_unix_ms, Some(145));
assert_eq!(
updated.extra_data,
Some(json!({"terminal": true, "late": true}))
);
}
#[tokio::test]
async fn delete_created_before_removes_oldest_matching_rows_up_to_limit() {
let repository = InMemoryRequestCandidateRepository::seed(vec![
sample_candidate("cand-1", "req-1", 100),
sample_candidate("cand-2", "req-2", 200),
sample_candidate("cand-3", "req-3", 400),
]);
let deleted = repository
.delete_created_before(350, 1)
.await
.expect("delete should succeed");
assert_eq!(deleted, 1);
let rows = repository
.list_recent(10)
.await
.expect("list recent should succeed");
assert_eq!(rows.len(), 2);
assert_eq!(rows[0].id, "cand-3");
assert_eq!(rows[1].id, "cand-2");
}
}
@@ -0,0 +1,18 @@
mod memory;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::candidates::{
build_decision_trace, derive_request_candidate_final_status,
request_candidate_lifecycle_would_regress, DecisionTrace, DecisionTraceCandidate,
PublicHealthStatusCount, PublicHealthTimelineBucket, RequestCandidateFinalStatus,
RequestCandidateReadRepository, RequestCandidateRepository, RequestCandidateStatus,
RequestCandidateTrace, RequestCandidateWriteRepository, StoredRequestCandidate,
UpsertRequestCandidateRecord,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlRequestCandidateRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxRequestCandidateReadRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteRequestCandidateRepository;
pub use memory::InMemoryRequestCandidateRepository;
@@ -0,0 +1,321 @@
use std::collections::BTreeMap;
use std::sync::RwLock;
use std::time::{SystemTime, UNIX_EPOCH};
use async_trait::async_trait;
use crate::DataLayerError;
use aether_data_contracts::repository::gemini_file_mappings::{
GeminiFileMappingListQuery, GeminiFileMappingMimeTypeCount, GeminiFileMappingReadRepository,
GeminiFileMappingStats, GeminiFileMappingWriteRepository, StoredGeminiFileMapping,
StoredGeminiFileMappingListPage, UpsertGeminiFileMappingRecord,
};
#[derive(Default)]
pub struct InMemoryGeminiFileMappingRepository {
by_file: RwLock<BTreeMap<String, StoredGeminiFileMapping>>,
}
impl InMemoryGeminiFileMappingRepository {
pub fn seed<I>(items: I) -> Self
where
I: IntoIterator<Item = StoredGeminiFileMapping>,
{
let mut by_file = BTreeMap::new();
for item in items {
by_file.insert(item.file_name.clone(), item);
}
Self {
by_file: RwLock::new(by_file),
}
}
}
#[async_trait]
impl GeminiFileMappingReadRepository for InMemoryGeminiFileMappingRepository {
async fn find_by_file_name(
&self,
file_name: &str,
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
let guard = self.by_file.read().expect("gemini mapping repository lock");
Ok(guard.get(file_name).cloned())
}
async fn list_mappings(
&self,
query: &GeminiFileMappingListQuery,
) -> Result<StoredGeminiFileMappingListPage, DataLayerError> {
let guard = self.by_file.read().expect("gemini mapping repository lock");
let search = query
.search
.as_deref()
.map(|value| value.to_ascii_lowercase());
let mut items = guard
.values()
.filter(|item| query.include_expired || item.expires_at_unix_secs > query.now_unix_secs)
.filter(|item| {
search.as_deref().is_none_or(|needle| {
item.file_name.to_ascii_lowercase().contains(needle)
|| item
.display_name
.as_deref()
.map(|value| value.to_ascii_lowercase().contains(needle))
.unwrap_or(false)
})
})
.cloned()
.collect::<Vec<_>>();
items.sort_by(|left, right| {
right
.created_at_unix_ms
.cmp(&left.created_at_unix_ms)
.then_with(|| left.file_name.cmp(&right.file_name))
});
let total = items.len();
let page_items = items
.into_iter()
.skip(query.offset)
.take(query.limit)
.collect::<Vec<_>>();
Ok(StoredGeminiFileMappingListPage {
items: page_items,
total,
})
}
async fn summarize_mappings(
&self,
now_unix_secs: u64,
) -> Result<GeminiFileMappingStats, DataLayerError> {
let guard = self.by_file.read().expect("gemini mapping repository lock");
let total_mappings = guard.len();
let active_items = guard
.values()
.filter(|item| item.expires_at_unix_secs > now_unix_secs)
.collect::<Vec<_>>();
let active_mappings = active_items.len();
let expired_mappings = total_mappings.saturating_sub(active_mappings);
let mut by_mime_type = BTreeMap::<String, usize>::new();
for item in active_items {
let mime_type = item
.mime_type
.clone()
.filter(|value| !value.trim().is_empty())
.unwrap_or_else(|| "unknown".to_string());
*by_mime_type.entry(mime_type).or_default() += 1;
}
Ok(GeminiFileMappingStats {
total_mappings,
active_mappings,
expired_mappings,
by_mime_type: by_mime_type
.into_iter()
.map(|(mime_type, count)| GeminiFileMappingMimeTypeCount { mime_type, count })
.collect(),
})
}
}
#[async_trait]
impl GeminiFileMappingWriteRepository for InMemoryGeminiFileMappingRepository {
async fn upsert(
&self,
record: UpsertGeminiFileMappingRecord,
) -> Result<StoredGeminiFileMapping, DataLayerError> {
record.validate()?;
let mut guard = self
.by_file
.write()
.expect("gemini mapping repository lock");
let created_at_unix_ms = guard
.get(&record.file_name)
.map(|existing| existing.created_at_unix_ms)
.unwrap_or_else(current_unix_secs);
let mapping = StoredGeminiFileMapping {
id: record.id.clone(),
file_name: record.file_name.clone(),
key_id: record.key_id.clone(),
user_id: record.user_id.clone(),
display_name: record.display_name.clone(),
mime_type: record.mime_type.clone(),
source_hash: record.source_hash.clone(),
created_at_unix_ms,
expires_at_unix_secs: record.expires_at_unix_secs,
};
guard.insert(record.file_name.clone(), mapping.clone());
Ok(mapping)
}
async fn delete_by_file_name(&self, file_name: &str) -> Result<bool, DataLayerError> {
let mut guard = self
.by_file
.write()
.expect("gemini mapping repository lock");
Ok(guard.remove(file_name).is_some())
}
async fn delete_by_id(
&self,
mapping_id: &str,
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
let mut guard = self
.by_file
.write()
.expect("gemini mapping repository lock");
let Some(file_name) = guard
.iter()
.find_map(|(file_name, item)| (item.id == mapping_id).then(|| file_name.clone()))
else {
return Ok(None);
};
Ok(guard.remove(&file_name))
}
async fn delete_expired_before(&self, now_unix_secs: u64) -> Result<usize, DataLayerError> {
let mut guard = self
.by_file
.write()
.expect("gemini mapping repository lock");
let before = guard.len();
guard.retain(|_, item| item.expires_at_unix_secs > now_unix_secs);
Ok(before.saturating_sub(guard.len()))
}
}
fn current_unix_secs() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|duration| duration.as_secs())
.unwrap_or_default()
}
#[cfg(test)]
mod tests {
use crate::repository::gemini_file_mappings::{
GeminiFileMappingListQuery, GeminiFileMappingReadRepository,
GeminiFileMappingWriteRepository,
};
use super::{InMemoryGeminiFileMappingRepository, UpsertGeminiFileMappingRecord};
use crate::DataLayerError;
fn sample_record(id: &str, file_name: &str) -> UpsertGeminiFileMappingRecord {
UpsertGeminiFileMappingRecord {
id: id.to_string(),
file_name: file_name.to_string(),
key_id: "key-1".to_string(),
user_id: Some("user-1".to_string()),
display_name: Some("display".to_string()),
mime_type: Some("image/png".to_string()),
source_hash: Some("hash-1".to_string()),
expires_at_unix_secs: 4_102_444_800,
}
}
#[tokio::test]
async fn upsert_and_find() -> Result<(), DataLayerError> {
let repo = InMemoryGeminiFileMappingRepository::default();
let record = sample_record("id-1", "files/abc");
let stored = repo.upsert(record.clone()).await?;
assert_eq!(stored.file_name, "files/abc");
let fetched = repo.find_by_file_name("files/abc").await?;
assert_eq!(fetched.unwrap().key_id, "key-1");
Ok(())
}
#[tokio::test]
async fn delete_removes_entry() -> Result<(), DataLayerError> {
let repo = InMemoryGeminiFileMappingRepository::default();
let record = sample_record("id-2", "files/def");
let _stored = repo.upsert(record).await?;
assert!(repo.find_by_file_name("files/def").await?.is_some());
assert!(repo.delete_by_file_name("files/def").await?);
assert!(repo.find_by_file_name("files/def").await?.is_none());
Ok(())
}
#[tokio::test]
async fn upsert_preserves_created_at_when_replacing_existing_file_name(
) -> Result<(), DataLayerError> {
let repo = InMemoryGeminiFileMappingRepository::default();
let first = repo.upsert(sample_record("id-1", "files/same")).await?;
let replaced = repo.upsert(sample_record("id-2", "files/same")).await?;
assert_eq!(replaced.created_at_unix_ms, first.created_at_unix_ms);
assert_eq!(replaced.id, "id-2");
Ok(())
}
#[tokio::test]
async fn list_and_summarize_mappings() -> Result<(), DataLayerError> {
let repo = InMemoryGeminiFileMappingRepository::seed(vec![
repo_item("id-1", "files/alpha", "image/png", 10, 200),
repo_item("id-2", "files/beta", "video/mp4", 20, 50),
repo_item("id-3", "files/gamma", "", 30, 220),
]);
let page = repo
.list_mappings(&GeminiFileMappingListQuery {
include_expired: false,
search: Some("ga".to_string()),
offset: 0,
limit: 10,
now_unix_secs: 100,
})
.await?;
assert_eq!(page.total, 1);
assert_eq!(page.items[0].id, "id-3");
let stats = repo.summarize_mappings(100).await?;
assert_eq!(stats.total_mappings, 3);
assert_eq!(stats.active_mappings, 2);
assert_eq!(stats.expired_mappings, 1);
assert_eq!(stats.by_mime_type.len(), 2);
assert_eq!(stats.by_mime_type[0].mime_type, "image/png");
assert_eq!(stats.by_mime_type[0].count, 1);
assert_eq!(stats.by_mime_type[1].mime_type, "unknown");
assert_eq!(stats.by_mime_type[1].count, 1);
Ok(())
}
#[tokio::test]
async fn delete_by_id_and_cleanup_expired() -> Result<(), DataLayerError> {
let repo = InMemoryGeminiFileMappingRepository::seed(vec![
repo_item("id-1", "files/alpha", "image/png", 10, 200),
repo_item("id-2", "files/beta", "video/mp4", 20, 50),
]);
let deleted = repo.delete_by_id("id-1").await?;
assert_eq!(
deleted.as_ref().map(|item| item.file_name.as_str()),
Some("files/alpha")
);
assert!(repo.find_by_file_name("files/alpha").await?.is_none());
let deleted_count = repo.delete_expired_before(100).await?;
assert_eq!(deleted_count, 1);
assert!(repo.find_by_file_name("files/beta").await?.is_none());
Ok(())
}
fn repo_item(
id: &str,
file_name: &str,
mime_type: &str,
created_at_unix_ms: u64,
expires_at_unix_secs: u64,
) -> crate::repository::gemini_file_mappings::StoredGeminiFileMapping {
crate::repository::gemini_file_mappings::StoredGeminiFileMapping {
id: id.to_string(),
file_name: file_name.to_string(),
key_id: "key-1".to_string(),
user_id: Some("user-1".to_string()),
display_name: Some(format!("display-{id}")),
mime_type: (!mime_type.is_empty()).then(|| mime_type.to_string()),
source_hash: Some(format!("hash-{id}")),
created_at_unix_ms,
expires_at_unix_secs,
}
}
}
@@ -0,0 +1,29 @@
pub mod memory;
#[cfg(feature = "mysql")]
pub mod mysql {
pub use aether_data_mysql::MysqlGeminiFileMappingRepository;
}
#[cfg(feature = "postgres")]
pub mod postgres {
pub use aether_data_postgres::SqlxGeminiFileMappingRepository;
}
#[cfg(feature = "sqlite")]
pub mod sqlite {
pub use aether_data_sqlite::SqliteGeminiFileMappingRepository;
}
pub mod types {
pub use aether_data_contracts::repository::gemini_file_mappings::*;
}
pub use aether_data_contracts::repository::gemini_file_mappings::{
GeminiFileMappingListQuery, GeminiFileMappingMimeTypeCount, GeminiFileMappingReadRepository,
GeminiFileMappingRepository, GeminiFileMappingStats, GeminiFileMappingWriteRepository,
StoredGeminiFileMapping, StoredGeminiFileMappingListPage, UpsertGeminiFileMappingRecord,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlGeminiFileMappingRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxGeminiFileMappingRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteGeminiFileMappingRepository;
pub use memory::InMemoryGeminiFileMappingRepository;
@@ -0,0 +1,728 @@
use std::sync::RwLock;
use async_trait::async_trait;
use super::{
AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord,
GlobalModelReadRepository, GlobalModelSnapshot, GlobalModelWriteRepository,
PublicCatalogModelListQuery, PublicCatalogModelSearchQuery, PublicGlobalModelQuery,
StoredAdminGlobalModel, StoredAdminGlobalModelPage, StoredAdminProviderModel,
StoredProviderActiveGlobalModel, StoredProviderModelStats, StoredPublicCatalogModel,
StoredPublicGlobalModel, StoredPublicGlobalModelPage, UpdateAdminGlobalModelRecord,
UpsertAdminProviderModelRecord,
};
use crate::DataLayerError;
#[derive(Debug, Default)]
pub struct InMemoryGlobalModelReadRepository {
items: RwLock<Vec<StoredPublicGlobalModel>>,
admin_global_model_items: RwLock<Vec<StoredAdminGlobalModel>>,
public_catalog_items: RwLock<Vec<StoredPublicCatalogModel>>,
admin_provider_model_items: RwLock<Vec<StoredAdminProviderModel>>,
provider_model_stats: RwLock<Vec<StoredProviderModelStats>>,
active_global_model_refs: RwLock<Vec<StoredProviderActiveGlobalModel>>,
}
impl InMemoryGlobalModelReadRepository {
pub fn seed<I>(items: I) -> Self
where
I: IntoIterator<Item = StoredPublicGlobalModel>,
{
Self {
items: RwLock::new(items.into_iter().collect()),
admin_global_model_items: RwLock::new(Vec::new()),
public_catalog_items: RwLock::new(Vec::new()),
admin_provider_model_items: RwLock::new(Vec::new()),
provider_model_stats: RwLock::new(Vec::new()),
active_global_model_refs: RwLock::new(Vec::new()),
}
}
pub fn with_public_catalog_models<I>(self, items: I) -> Self
where
I: IntoIterator<Item = StoredPublicCatalogModel>,
{
*self
.public_catalog_items
.write()
.expect("public catalog model repository lock") = items.into_iter().collect();
self
}
pub fn with_provider_model_stats<I>(self, items: I) -> Self
where
I: IntoIterator<Item = StoredProviderModelStats>,
{
*self
.provider_model_stats
.write()
.expect("provider model stats repository lock") = items.into_iter().collect();
self
}
pub fn with_admin_provider_models<I>(self, items: I) -> Self
where
I: IntoIterator<Item = StoredAdminProviderModel>,
{
*self
.admin_provider_model_items
.write()
.expect("admin provider model repository lock") = items.into_iter().collect();
self
}
pub fn with_active_global_model_refs<I>(self, items: I) -> Self
where
I: IntoIterator<Item = StoredProviderActiveGlobalModel>,
{
*self
.active_global_model_refs
.write()
.expect("active global model repository lock") = items.into_iter().collect();
self
}
pub fn with_admin_global_models<I>(self, items: I) -> Self
where
I: IntoIterator<Item = StoredAdminGlobalModel>,
{
*self
.admin_global_model_items
.write()
.expect("admin global model repository lock") = items.into_iter().collect();
self
}
fn snapshot(&self) -> GlobalModelSnapshot {
GlobalModelSnapshot::seed(
self.items
.read()
.expect("global model repository lock")
.clone(),
)
.with_admin_global_models(
self.admin_global_model_items
.read()
.expect("admin global model repository lock")
.clone(),
)
.with_public_catalog_models(
self.public_catalog_items
.read()
.expect("public catalog model repository lock")
.clone(),
)
.with_admin_provider_models(
self.admin_provider_model_items
.read()
.expect("admin provider model repository lock")
.clone(),
)
.with_provider_model_stats(
self.provider_model_stats
.read()
.expect("provider model stats repository lock")
.clone(),
)
.with_active_global_model_refs(
self.active_global_model_refs
.read()
.expect("active global model repository lock")
.clone(),
)
}
}
#[async_trait]
impl GlobalModelReadRepository for InMemoryGlobalModelReadRepository {
async fn list_public_models(
&self,
query: &PublicGlobalModelQuery,
) -> Result<StoredPublicGlobalModelPage, DataLayerError> {
Ok(self.snapshot().list_public_models(query))
}
async fn get_public_model_by_name(
&self,
model_name: &str,
) -> Result<Option<StoredPublicGlobalModel>, DataLayerError> {
Ok(self.snapshot().get_public_model_by_name(model_name))
}
async fn list_public_catalog_models(
&self,
query: &PublicCatalogModelListQuery,
) -> Result<Vec<StoredPublicCatalogModel>, DataLayerError> {
Ok(self.snapshot().list_public_catalog_models(query))
}
async fn search_public_catalog_models(
&self,
query: &PublicCatalogModelSearchQuery,
) -> Result<Vec<StoredPublicCatalogModel>, DataLayerError> {
Ok(self.snapshot().search_public_catalog_models(query))
}
async fn list_admin_global_models(
&self,
query: &AdminGlobalModelListQuery,
) -> Result<StoredAdminGlobalModelPage, DataLayerError> {
Ok(self.snapshot().list_admin_global_models(query))
}
async fn list_admin_provider_models(
&self,
query: &AdminProviderModelListQuery,
) -> Result<Vec<StoredAdminProviderModel>, DataLayerError> {
Ok(self.snapshot().list_admin_provider_models(query))
}
async fn get_admin_provider_model(
&self,
provider_id: &str,
model_id: &str,
) -> Result<Option<StoredAdminProviderModel>, DataLayerError> {
Ok(self
.snapshot()
.get_admin_provider_model(provider_id, model_id))
}
async fn list_admin_provider_available_source_models(
&self,
provider_id: &str,
) -> Result<Vec<StoredAdminProviderModel>, DataLayerError> {
Ok(self
.snapshot()
.list_admin_provider_available_source_models(provider_id))
}
async fn get_admin_global_model_by_id(
&self,
global_model_id: &str,
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
Ok(self
.snapshot()
.get_admin_global_model_by_id(global_model_id))
}
async fn get_admin_global_model_by_name(
&self,
model_name: &str,
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
Ok(self.snapshot().get_admin_global_model_by_name(model_name))
}
async fn list_admin_provider_models_by_global_model_id(
&self,
global_model_id: &str,
) -> Result<Vec<StoredAdminProviderModel>, DataLayerError> {
Ok(self
.snapshot()
.list_admin_provider_models_by_global_model_id(global_model_id))
}
async fn list_provider_model_stats(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderModelStats>, DataLayerError> {
Ok(self.snapshot().list_provider_model_stats(provider_ids))
}
async fn list_active_global_model_ids_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderActiveGlobalModel>, DataLayerError> {
Ok(self
.snapshot()
.list_active_global_model_ids_by_provider_ids(provider_ids))
}
}
#[async_trait]
impl GlobalModelWriteRepository for InMemoryGlobalModelReadRepository {
async fn create_admin_provider_model(
&self,
record: &UpsertAdminProviderModelRecord,
) -> Result<Option<StoredAdminProviderModel>, DataLayerError> {
let global_model = self
.get_admin_global_model_by_id(&record.global_model_id)
.await?
.ok_or_else(|| DataLayerError::UnexpectedValue("global model not found".to_string()))?;
let stored = StoredAdminProviderModel::new(
record.id.clone(),
record.provider_id.clone(),
record.global_model_id.clone(),
record.provider_model_name.clone(),
record.provider_model_mappings.clone(),
record.price_per_request,
record.tiered_pricing.clone(),
record.supports_vision,
record.supports_function_calling,
record.supports_streaming,
record.supports_extended_thinking,
record.supports_image_generation,
record.is_active,
record.is_available,
record.config.clone(),
Some(1_711_000_000),
Some(1_711_000_000),
Some(global_model.name.clone()),
Some(global_model.display_name.clone()),
global_model.default_price_per_request,
global_model.default_tiered_pricing.clone(),
global_model.supported_capabilities.clone(),
global_model.config.clone(),
)?;
self.admin_provider_model_items
.write()
.expect("admin provider model repository lock")
.push(stored.clone());
Ok(Some(stored))
}
async fn update_admin_provider_model(
&self,
record: &UpsertAdminProviderModelRecord,
) -> Result<Option<StoredAdminProviderModel>, DataLayerError> {
let global_model = self
.get_admin_global_model_by_id(&record.global_model_id)
.await?
.ok_or_else(|| DataLayerError::UnexpectedValue("global model not found".to_string()))?;
let mut items = self
.admin_provider_model_items
.write()
.expect("admin provider model repository lock");
let Some(existing) = items
.iter_mut()
.find(|item| item.id == record.id && item.provider_id == record.provider_id)
else {
return Ok(None);
};
existing.global_model_id = record.global_model_id.clone();
existing.provider_model_name = record.provider_model_name.clone();
existing.provider_model_mappings = record.provider_model_mappings.clone();
existing.price_per_request = record.price_per_request;
existing.tiered_pricing = record.tiered_pricing.clone();
existing.supports_vision = record.supports_vision;
existing.supports_function_calling = record.supports_function_calling;
existing.supports_streaming = record.supports_streaming;
existing.supports_extended_thinking = record.supports_extended_thinking;
existing.supports_image_generation = record.supports_image_generation;
existing.is_active = record.is_active;
existing.is_available = record.is_available;
existing.config = record.config.clone();
existing.updated_at_unix_secs = Some(1_711_000_100);
existing.global_model_name = Some(global_model.name.clone());
existing.global_model_display_name = Some(global_model.display_name.clone());
existing.global_model_default_price_per_request = global_model.default_price_per_request;
existing.global_model_default_tiered_pricing = global_model.default_tiered_pricing.clone();
existing.global_model_supported_capabilities = global_model.supported_capabilities.clone();
existing.global_model_config = global_model.config.clone();
Ok(Some(existing.clone()))
}
async fn delete_admin_provider_model(
&self,
provider_id: &str,
model_id: &str,
) -> Result<bool, DataLayerError> {
let mut items = self
.admin_provider_model_items
.write()
.expect("admin provider model repository lock");
let original_len = items.len();
items.retain(|item| !(item.provider_id == provider_id && item.id == model_id));
Ok(items.len() != original_len)
}
async fn create_admin_global_model(
&self,
record: &CreateAdminGlobalModelRecord,
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
let stored = StoredAdminGlobalModel::new(
record.id.clone(),
record.name.clone(),
record.display_name.clone(),
record.is_active,
record.default_price_per_request,
record.default_tiered_pricing.clone(),
record.supported_capabilities.clone(),
record.config.clone(),
0,
0,
record.usage_count.unwrap_or(0),
Some(1_711_000_000),
Some(1_711_000_000),
)?;
self.admin_global_model_items
.write()
.expect("admin global model repository lock")
.push(stored);
self.get_admin_global_model_by_id(&record.id).await
}
async fn update_admin_global_model(
&self,
record: &UpdateAdminGlobalModelRecord,
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
{
let mut items = self
.admin_global_model_items
.write()
.expect("admin global model repository lock");
let Some(existing) = items.iter_mut().find(|item| item.id == record.id) else {
return Ok(None);
};
existing.display_name = record.display_name.clone();
existing.is_active = record.is_active;
existing.default_price_per_request = record.default_price_per_request;
existing.default_tiered_pricing = record.default_tiered_pricing.clone();
existing.supported_capabilities = record.supported_capabilities.clone();
existing.config = record.config.clone();
if let Some(usage_count) = record.usage_count {
existing.usage_count = usage_count;
}
existing.updated_at_unix_secs = Some(1_711_000_100);
}
self.get_admin_global_model_by_id(&record.id).await
}
async fn delete_admin_global_model(
&self,
global_model_id: &str,
) -> Result<bool, DataLayerError> {
let mut globals = self
.admin_global_model_items
.write()
.expect("admin global model repository lock");
let original_len = globals.len();
globals.retain(|item| item.id != global_model_id);
drop(globals);
self.admin_provider_model_items
.write()
.expect("admin provider model repository lock")
.retain(|item| item.global_model_id != global_model_id);
Ok(original_len
!= self
.admin_global_model_items
.read()
.expect("admin global model repository lock")
.len())
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::InMemoryGlobalModelReadRepository;
use crate::repository::global_models::{
CreateAdminGlobalModelRecord, GlobalModelReadRepository, GlobalModelWriteRepository,
PublicCatalogModelListQuery, PublicCatalogModelSearchQuery, PublicGlobalModelQuery,
StoredPublicCatalogModel, StoredPublicGlobalModel,
};
fn sample_model(
id: &str,
name: &str,
display_name: &str,
is_active: bool,
) -> StoredPublicGlobalModel {
StoredPublicGlobalModel::new(
id.to_string(),
name.to_string(),
Some(display_name.to_string()),
is_active,
Some(0.02),
Some(json!({"tiers":[{"up_to": null, "input_price_per_1m": 3.0, "output_price_per_1m": 15.0}]})),
Some(json!(["vision"])),
Some(json!({"family": "test"})),
0,
)
.expect("global model should build")
}
fn sample_public_catalog_model(
id: &str,
provider_id: &str,
provider_name: &str,
provider_model_name: &str,
name: &str,
display_name: &str,
) -> StoredPublicCatalogModel {
StoredPublicCatalogModel::new(
id.to_string(),
provider_id.to_string(),
provider_name.to_string(),
provider_model_name.to_string(),
name.to_string(),
display_name.to_string(),
Some(format!("{display_name} description")),
Some(format!("https://cdn.example/{name}.png")),
Some(3.0),
Some(15.0),
Some(1.5),
Some(0.3),
Some(true),
Some(true),
Some(true),
Some(false),
true,
)
.expect("public catalog model should build")
}
#[tokio::test]
async fn embedding_model_metadata_roundtrip() {
let repository =
InMemoryGlobalModelReadRepository::seed(Vec::<StoredPublicGlobalModel>::new());
let record = CreateAdminGlobalModelRecord::new(
"gm-embedding".to_string(),
"text-embedding-3-small".to_string(),
"Text Embedding 3 Small".to_string(),
true,
None,
Some(json!({"tiers":[{"up_to":null,"input_price_per_1m":0.02}]})),
Some(json!(["embedding"])),
Some(json!({
"api_formats": ["openai:embedding"],
"dimensions": 1536
})),
)
.expect("embedding global model should validate");
repository
.create_admin_global_model(&record)
.await
.expect("embedding global model should persist")
.expect("embedding global model should be returned");
let stored = repository
.get_admin_global_model_by_name("text-embedding-3-small")
.await
.expect("embedding global model should read")
.expect("embedding global model should exist");
assert_eq!(stored.supported_capabilities, Some(json!(["embedding"])));
assert_eq!(
stored
.config
.as_ref()
.and_then(|value| value.get("dimensions")),
Some(&json!(1536))
);
assert_eq!(
stored
.default_tiered_pricing
.as_ref()
.and_then(|value| value.get("tiers"))
.and_then(serde_json::Value::as_array)
.and_then(|tiers| tiers.first())
.and_then(|tier| tier.get("input_price_per_1m"))
.and_then(serde_json::Value::as_f64),
Some(0.02)
);
}
#[tokio::test]
async fn embedding_missing_billing_config_rejected() {
let error = CreateAdminGlobalModelRecord::new(
"gm-embedding".to_string(),
"text-embedding-3-small".to_string(),
"Text Embedding 3 Small".to_string(),
true,
None,
None,
Some(json!(["embedding"])),
None,
)
.expect_err("embedding metadata without billing should fail closed");
assert!(error
.to_string()
.contains("embedding global model requires"));
}
#[tokio::test]
async fn defaults_to_active_models_only() {
let repository = InMemoryGlobalModelReadRepository::seed(vec![
sample_model("gm-1", "claude-sonnet-4-5", "Claude Sonnet 4.5", true),
sample_model("gm-2", "legacy-model", "Legacy Model", false),
]);
let page = repository
.list_public_models(&PublicGlobalModelQuery {
offset: 0,
limit: 50,
is_active: None,
search: None,
})
.await
.expect("list should succeed");
assert_eq!(page.total, 1);
assert_eq!(page.items[0].name, "claude-sonnet-4-5");
}
#[tokio::test]
async fn search_matches_name_and_display_name() {
let repository = InMemoryGlobalModelReadRepository::seed(vec![
sample_model("gm-1", "gpt-5", "GPT 5", true),
sample_model("gm-2", "claude-sonnet-4-5", "Claude Sonnet 4.5", true),
]);
let page = repository
.list_public_models(&PublicGlobalModelQuery {
offset: 0,
limit: 50,
is_active: None,
search: Some("sonnet".to_string()),
})
.await
.expect("list should succeed");
assert_eq!(page.total, 1);
assert_eq!(page.items[0].name, "claude-sonnet-4-5");
}
#[tokio::test]
async fn get_public_model_by_name_only_returns_active_exact_match() {
let repository = InMemoryGlobalModelReadRepository::seed(vec![
sample_model("gm-1", "gpt-5", "GPT 5", true),
sample_model("gm-2", "gpt-5-old", "GPT 5 Old", false),
]);
let model = repository
.get_public_model_by_name("gpt-5")
.await
.expect("lookup should succeed");
assert_eq!(model.expect("model should exist").name, "gpt-5");
let missing = repository
.get_public_model_by_name("gpt-5-old")
.await
.expect("lookup should succeed");
assert!(missing.is_none());
}
#[tokio::test]
async fn lists_public_catalog_models_with_provider_filter() {
let repository =
InMemoryGlobalModelReadRepository::seed(Vec::<StoredPublicGlobalModel>::new())
.with_public_catalog_models(vec![
sample_public_catalog_model(
"model-1",
"provider-openai",
"openai",
"gpt-5-preview",
"gpt-5",
"GPT 5",
),
sample_public_catalog_model(
"model-2",
"provider-claude",
"claude",
"claude-3-7-sonnet",
"claude-3-7-sonnet",
"Claude 3.7 Sonnet",
),
]);
let items = repository
.list_public_catalog_models(&PublicCatalogModelListQuery {
provider_id: Some("provider-openai".to_string()),
offset: 0,
limit: 50,
})
.await
.expect("list should succeed");
assert_eq!(items.len(), 1);
assert_eq!(items[0].provider_id, "provider-openai");
assert_eq!(items[0].name, "gpt-5");
}
#[tokio::test]
async fn public_catalog_preserves_embedding_capability_without_contaminating_chat_models() {
let mut embedding_model = sample_public_catalog_model(
"model-embedding",
"provider-openai",
"openai",
"text-embedding-3-small",
"text-embedding-3-small",
"Text Embedding 3 Small",
);
embedding_model.supports_embedding = Some(true);
embedding_model.supports_streaming = Some(false);
let chat_model = sample_public_catalog_model(
"model-chat",
"provider-openai",
"openai",
"gpt-5-upstream",
"gpt-5",
"GPT 5",
);
let repository =
InMemoryGlobalModelReadRepository::seed(Vec::<StoredPublicGlobalModel>::new())
.with_public_catalog_models(vec![embedding_model, chat_model]);
let items = repository
.list_public_catalog_models(&PublicCatalogModelListQuery {
provider_id: Some("provider-openai".to_string()),
offset: 0,
limit: 50,
})
.await
.expect("catalog should list");
let embedding = items
.iter()
.find(|item| item.name == "text-embedding-3-small")
.expect("embedding model should be listed");
let chat = items
.iter()
.find(|item| item.name == "gpt-5")
.expect("chat model should be listed");
assert_eq!(embedding.supports_embedding, Some(true));
assert_eq!(embedding.supports_streaming, Some(false));
assert_eq!(chat.supports_embedding, Some(false));
}
#[tokio::test]
async fn searches_public_catalog_models_by_provider_and_display_name() {
let repository =
InMemoryGlobalModelReadRepository::seed(Vec::<StoredPublicGlobalModel>::new())
.with_public_catalog_models(vec![
sample_public_catalog_model(
"model-1",
"provider-openai",
"openai",
"gpt-5-preview",
"gpt-5",
"GPT 5",
),
sample_public_catalog_model(
"model-2",
"provider-claude",
"claude",
"claude-3-7-sonnet",
"claude-3-7-sonnet",
"Claude 3.7 Sonnet",
),
]);
let items = repository
.search_public_catalog_models(&PublicCatalogModelSearchQuery {
search: "sonnet".to_string(),
provider_id: Some("provider-claude".to_string()),
limit: 20,
})
.await
.expect("search should succeed");
assert_eq!(items.len(), 1);
assert_eq!(items[0].provider_name, "claude");
assert_eq!(items[0].display_name, "Claude 3.7 Sonnet");
}
}
@@ -0,0 +1,19 @@
mod memory;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::global_models::{
metadata_supports_embedding, AdminGlobalModelListQuery, AdminProviderModelListQuery,
CreateAdminGlobalModelRecord, GlobalModelReadRepository, GlobalModelSnapshot,
GlobalModelWriteRepository, PublicCatalogModelListQuery, PublicCatalogModelSearchQuery,
PublicGlobalModelQuery, StoredAdminGlobalModel, StoredAdminGlobalModelPage,
StoredAdminProviderModel, StoredProviderActiveGlobalModel, StoredProviderModelStats,
StoredPublicCatalogModel, StoredPublicGlobalModel, StoredPublicGlobalModelPage,
UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlGlobalModelReadRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxGlobalModelReadRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteGlobalModelReadRepository;
pub use memory::InMemoryGlobalModelReadRepository;
@@ -0,0 +1,459 @@
use std::collections::BTreeMap;
use std::sync::RwLock;
use std::time::{SystemTime, UNIX_EPOCH};
use async_trait::async_trait;
use crate::DataLayerError;
use aether_data_contracts::repository::management_tokens::{
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken,
StoredManagementTokenListPage, StoredManagementTokenWithUser, UpdateManagementTokenRecord,
};
#[derive(Debug, Default)]
pub struct InMemoryManagementTokenRepository {
items: RwLock<Vec<StoredManagementTokenWithUser>>,
hashes: RwLock<BTreeMap<String, String>>,
}
impl InMemoryManagementTokenRepository {
pub fn seed<I>(items: I) -> Self
where
I: IntoIterator<Item = StoredManagementTokenWithUser>,
{
Self {
items: RwLock::new(items.into_iter().collect()),
hashes: RwLock::new(BTreeMap::new()),
}
}
pub fn seed_with_hashes<I, J>(items: I, hashes: J) -> Self
where
I: IntoIterator<Item = StoredManagementTokenWithUser>,
J: IntoIterator<Item = (String, String)>,
{
Self {
items: RwLock::new(items.into_iter().collect()),
hashes: RwLock::new(hashes.into_iter().collect()),
}
}
fn now_unix_secs() -> Option<u64> {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
}
fn remove_hash_for_token(hashes: &mut BTreeMap<String, String>, token_id: &str) {
hashes.retain(|_, existing_token_id| existing_token_id != token_id);
}
}
#[async_trait]
impl ManagementTokenReadRepository for InMemoryManagementTokenRepository {
async fn list_management_tokens(
&self,
query: &ManagementTokenListQuery,
) -> Result<StoredManagementTokenListPage, DataLayerError> {
let items = self.items.read().expect("management token repository lock");
let mut filtered = items
.iter()
.filter(|item| match query.user_id.as_deref() {
Some(user_id) => item.token.user_id == user_id,
None => true,
})
.filter(|item| match query.is_active {
Some(is_active) => item.token.is_active == is_active,
None => true,
})
.cloned()
.collect::<Vec<_>>();
filtered.sort_by(|left, right| {
right
.token
.created_at_unix_ms
.cmp(&left.token.created_at_unix_ms)
.then_with(|| right.token.id.cmp(&left.token.id))
});
let total = filtered.len();
let items = filtered
.into_iter()
.skip(query.offset)
.take(query.limit)
.collect();
Ok(StoredManagementTokenListPage { items, total })
}
async fn get_management_token_with_user(
&self,
token_id: &str,
) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> {
let items = self.items.read().expect("management token repository lock");
Ok(items.iter().find(|item| item.token.id == token_id).cloned())
}
async fn get_management_token_with_user_by_hash(
&self,
token_hash: &str,
) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> {
let token_id = {
let hashes = self
.hashes
.read()
.expect("management token repository lock");
hashes.get(token_hash).cloned()
};
let Some(token_id) = token_id else {
return Ok(None);
};
let items = self.items.read().expect("management token repository lock");
Ok(items.iter().find(|item| item.token.id == token_id).cloned())
}
}
#[async_trait]
impl ManagementTokenWriteRepository for InMemoryManagementTokenRepository {
async fn create_management_token(
&self,
record: &CreateManagementTokenRecord,
) -> Result<StoredManagementToken, DataLayerError> {
record.validate()?;
let mut items = self
.items
.write()
.expect("management token repository lock");
let mut hashes = self
.hashes
.write()
.expect("management token repository lock");
if items
.iter()
.any(|item| item.token.user_id == record.user_id && item.token.name == record.name)
{
return Err(DataLayerError::InvalidInput(format!(
"已存在名为 '{}' 的 Token",
record.name
)));
}
let now = Self::now_unix_secs();
let token = StoredManagementToken::new(
record.id.clone(),
record.user_id.clone(),
record.name.clone(),
)?
.with_display_fields(
record.description.clone(),
record.token_prefix.clone(),
record.allowed_ips.clone(),
)
.with_permissions(record.permissions.clone())
.with_runtime_fields(record.expires_at_unix_secs, None, None, 0, record.is_active)
.with_timestamps(now, now);
items.push(StoredManagementTokenWithUser::new(
token.clone(),
record.user.clone(),
));
hashes.insert(record.token_hash.clone(), record.id.clone());
Ok(token)
}
async fn update_management_token(
&self,
record: &UpdateManagementTokenRecord,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
record.validate()?;
let mut items = self
.items
.write()
.expect("management token repository lock");
let Some(index) = items
.iter()
.position(|item| item.token.id == record.token_id)
else {
return Ok(None);
};
if let Some(name) = &record.name {
if items.iter().enumerate().any(|(position, item)| {
position != index
&& item.token.user_id == items[index].token.user_id
&& item.token.name == *name
}) {
return Err(DataLayerError::InvalidInput(format!(
"已存在名为 '{}' 的 Token",
name
)));
}
items[index].token.name = name.clone();
}
if record.clear_description {
items[index].token.description = None;
} else if let Some(description) = &record.description {
items[index].token.description = Some(description.clone());
}
if record.clear_allowed_ips {
items[index].token.allowed_ips = None;
} else if let Some(allowed_ips) = &record.allowed_ips {
items[index].token.allowed_ips = Some(allowed_ips.clone());
}
if let Some(permissions) = &record.permissions {
items[index].token.permissions = Some(permissions.clone());
}
if record.clear_expires_at {
items[index].token.expires_at_unix_secs = None;
} else if let Some(expires_at_unix_secs) = record.expires_at_unix_secs {
items[index].token.expires_at_unix_secs = Some(expires_at_unix_secs);
}
if let Some(is_active) = record.is_active {
items[index].token.is_active = is_active;
}
items[index].token.updated_at_unix_secs = Self::now_unix_secs();
Ok(Some(items[index].token.clone()))
}
async fn delete_management_token(&self, token_id: &str) -> Result<bool, DataLayerError> {
let mut items = self
.items
.write()
.expect("management token repository lock");
let mut hashes = self
.hashes
.write()
.expect("management token repository lock");
let original_len = items.len();
items.retain(|item| item.token.id != token_id);
Self::remove_hash_for_token(&mut hashes, token_id);
Ok(items.len() != original_len)
}
async fn set_management_token_active(
&self,
token_id: &str,
is_active: bool,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
let mut items = self
.items
.write()
.expect("management token repository lock");
let Some(item) = items.iter_mut().find(|item| item.token.id == token_id) else {
return Ok(None);
};
item.token.is_active = is_active;
item.token.updated_at_unix_secs = Self::now_unix_secs();
Ok(Some(item.token.clone()))
}
async fn regenerate_management_token_secret(
&self,
mutation: &RegenerateManagementTokenSecret,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
mutation.validate()?;
let mut items = self
.items
.write()
.expect("management token repository lock");
let mut hashes = self
.hashes
.write()
.expect("management token repository lock");
let Some(item) = items
.iter_mut()
.find(|item| item.token.id == mutation.token_id)
else {
return Ok(None);
};
Self::remove_hash_for_token(&mut hashes, &mutation.token_id);
hashes.insert(mutation.token_hash.clone(), mutation.token_id.clone());
item.token.token_prefix = mutation.token_prefix.clone();
item.token.updated_at_unix_secs = Self::now_unix_secs();
Ok(Some(item.token.clone()))
}
async fn record_management_token_usage(
&self,
token_id: &str,
last_used_ip: Option<&str>,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
let mut items = self
.items
.write()
.expect("management token repository lock");
let Some(item) = items.iter_mut().find(|item| item.token.id == token_id) else {
return Ok(None);
};
item.token.last_used_at_unix_secs = Self::now_unix_secs();
item.token.last_used_ip = last_used_ip.map(ToOwned::to_owned);
item.token.usage_count = item.token.usage_count.saturating_add(1);
item.token.updated_at_unix_secs = Self::now_unix_secs();
Ok(Some(item.token.clone()))
}
}
#[cfg(test)]
mod tests {
use super::InMemoryManagementTokenRepository;
use crate::repository::management_tokens::{
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken,
StoredManagementTokenUserSummary, StoredManagementTokenWithUser,
UpdateManagementTokenRecord,
};
fn sample_token(id: &str, user_id: &str, is_active: bool) -> StoredManagementTokenWithUser {
let token = StoredManagementToken::new(id.to_string(), user_id.to_string(), id.to_string())
.expect("token should build")
.with_runtime_fields(None, None, None, 2, is_active)
.with_timestamps(Some(1_700_000_000), Some(1_700_000_100));
let user = StoredManagementTokenUserSummary::new(
user_id.to_string(),
Some(format!("{user_id}@example.com")),
format!("{user_id}-name"),
"admin".to_string(),
)
.expect("user should build");
StoredManagementTokenWithUser::new(token, user)
}
#[tokio::test]
async fn lists_filters_and_mutates_management_tokens() {
let repository = InMemoryManagementTokenRepository::seed_with_hashes(
vec![
sample_token("token-1", "user-1", true),
sample_token("token-2", "user-2", false),
],
vec![
("hash-1".to_string(), "token-1".to_string()),
("hash-2".to_string(), "token-2".to_string()),
],
);
let page = repository
.list_management_tokens(&ManagementTokenListQuery {
user_id: None,
is_active: Some(true),
offset: 0,
limit: 10,
})
.await
.expect("list should succeed");
assert_eq!(page.total, 1);
assert_eq!(page.items[0].token.id, "token-1");
let toggled = repository
.set_management_token_active("token-2", true)
.await
.expect("toggle should succeed")
.expect("token should exist");
assert!(toggled.is_active);
let created = repository
.create_management_token(&CreateManagementTokenRecord {
id: "token-3".to_string(),
user_id: "user-1".to_string(),
user: StoredManagementTokenUserSummary::new(
"user-1".to_string(),
Some("[email protected]".to_string()),
"user-1-name".to_string(),
"user".to_string(),
)
.expect("user should build"),
token_hash: "hash-3".to_string(),
token_prefix: Some("ae_1234".to_string()),
name: "created".to_string(),
description: Some("created token".to_string()),
allowed_ips: Some(serde_json::json!(["127.0.0.1"])),
permissions: Some(serde_json::json!(["admin:usage:read"])),
expires_at_unix_secs: Some(1_800_000_000),
is_active: true,
})
.await
.expect("create should succeed");
assert_eq!(created.name, "created");
assert_eq!(
created.permissions,
Some(serde_json::json!(["admin:usage:read"]))
);
let updated = repository
.update_management_token(&UpdateManagementTokenRecord {
token_id: "token-3".to_string(),
name: Some("renamed".to_string()),
description: None,
clear_description: true,
allowed_ips: Some(serde_json::json!(["10.0.0.1"])),
clear_allowed_ips: false,
permissions: Some(serde_json::json!(["admin:usage:read", "admin:usage:write"])),
expires_at_unix_secs: None,
clear_expires_at: true,
is_active: Some(false),
})
.await
.expect("update should succeed")
.expect("token should exist");
assert_eq!(updated.name, "renamed");
assert_eq!(updated.description, None);
assert_eq!(updated.allowed_ips, Some(serde_json::json!(["10.0.0.1"])));
assert_eq!(
updated.permissions,
Some(serde_json::json!(["admin:usage:read", "admin:usage:write"]))
);
assert_eq!(updated.expires_at_unix_secs, None);
assert!(!updated.is_active);
let regenerated = repository
.regenerate_management_token_secret(&RegenerateManagementTokenSecret {
token_id: "token-3".to_string(),
token_hash: "hash-3b".to_string(),
token_prefix: Some("ae_5678".to_string()),
})
.await
.expect("regenerate should succeed")
.expect("token should exist");
assert_eq!(regenerated.token_prefix.as_deref(), Some("ae_5678"));
let by_hash = repository
.get_management_token_with_user_by_hash("hash-3b")
.await
.expect("lookup by hash should succeed")
.expect("token should exist");
assert_eq!(by_hash.token.id, "token-3");
assert_eq!(
by_hash.token.permissions,
Some(serde_json::json!(["admin:usage:read", "admin:usage:write"]))
);
let used = repository
.record_management_token_usage("token-3", Some("127.0.0.1"))
.await
.expect("usage update should succeed")
.expect("token should exist");
assert_eq!(used.last_used_ip.as_deref(), Some("127.0.0.1"));
assert_eq!(used.usage_count, 1);
let deleted = repository
.delete_management_token("token-1")
.await
.expect("delete should succeed");
assert!(deleted);
let deleted_by_hash = repository
.get_management_token_with_user_by_hash("hash-1")
.await
.expect("hash lookup should succeed");
assert!(deleted_by_hash.is_none());
}
}
@@ -0,0 +1,15 @@
mod memory;
pub use aether_data_contracts::repository::management_tokens::{
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken,
StoredManagementTokenListPage, StoredManagementTokenUserSummary, StoredManagementTokenWithUser,
UpdateManagementTokenRecord,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlManagementTokenRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxManagementTokenRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteManagementTokenRepository;
pub use memory::InMemoryManagementTokenRepository;
@@ -0,0 +1,31 @@
//! Runtime repository facade and in-memory implementations.
//!
//! Repository contracts and shared DTOs are re-exported from
//! `aether-data-contracts`. Selected database implementations are re-exported
//! from the adapter crates to preserve existing import paths; concrete SQL does
//! not belong in this facade.
pub mod announcements;
pub mod audit;
pub mod auth;
pub mod auth_modules;
pub mod background_tasks;
pub mod billing;
pub mod candidate_selection;
pub mod candidates;
pub mod gemini_file_mappings;
pub mod global_models;
pub mod management_tokens;
pub mod oauth_providers;
pub mod pool_scores;
pub mod provider_catalog;
pub mod provider_oauth;
pub mod proxy_nodes;
pub mod quota;
pub mod routing_profiles;
pub mod settlement;
pub mod system;
pub mod usage;
pub mod users;
pub mod video_tasks;
pub mod wallet;
@@ -0,0 +1,198 @@
use std::collections::BTreeMap;
use std::sync::RwLock;
use std::time::{SystemTime, UNIX_EPOCH};
use async_trait::async_trait;
use crate::DataLayerError;
use aether_data_contracts::repository::oauth_providers::{
EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderWriteRepository,
StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord,
};
#[derive(Debug, Default)]
pub struct InMemoryOAuthProviderRepository {
items: RwLock<BTreeMap<String, StoredOAuthProviderConfig>>,
}
impl InMemoryOAuthProviderRepository {
pub fn seed<I>(items: I) -> Self
where
I: IntoIterator<Item = StoredOAuthProviderConfig>,
{
let items = items
.into_iter()
.map(|item| (item.provider_type.clone(), item))
.collect();
Self {
items: RwLock::new(items),
}
}
fn now_unix_secs() -> Option<u64> {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.ok()
.map(|duration| duration.as_secs())
}
}
#[async_trait]
impl OAuthProviderReadRepository for InMemoryOAuthProviderRepository {
async fn list_oauth_provider_configs(
&self,
) -> Result<Vec<StoredOAuthProviderConfig>, DataLayerError> {
let items = self.items.read().expect("oauth provider repository lock");
Ok(items.values().cloned().collect())
}
async fn get_oauth_provider_config(
&self,
provider_type: &str,
) -> Result<Option<StoredOAuthProviderConfig>, DataLayerError> {
let items = self.items.read().expect("oauth provider repository lock");
Ok(items.get(provider_type).cloned())
}
async fn count_locked_users_if_provider_disabled(
&self,
_provider_type: &str,
_ldap_exclusive: bool,
) -> Result<usize, DataLayerError> {
Ok(0)
}
}
#[async_trait]
impl OAuthProviderWriteRepository for InMemoryOAuthProviderRepository {
async fn upsert_oauth_provider_config(
&self,
record: &UpsertOAuthProviderConfigRecord,
) -> Result<StoredOAuthProviderConfig, DataLayerError> {
record.validate()?;
let mut items = self.items.write().expect("oauth provider repository lock");
let now = Self::now_unix_secs();
let existing = items.get(&record.provider_type).cloned();
let created_at = existing
.as_ref()
.and_then(|item| item.created_at_unix_ms)
.or(now);
let client_secret_encrypted = match (&record.client_secret_encrypted, existing.as_ref()) {
(EncryptedSecretUpdate::Preserve, Some(item)) => item.client_secret_encrypted.clone(),
(EncryptedSecretUpdate::Preserve, None) => None,
(EncryptedSecretUpdate::Clear, _) => None,
(EncryptedSecretUpdate::Set(value), _) => Some(value.clone()),
};
let item = StoredOAuthProviderConfig::new(
record.provider_type.clone(),
record.display_name.clone(),
record.client_id.clone(),
record.redirect_uri.clone(),
record.frontend_callback_url.clone(),
)?
.with_config_fields(
client_secret_encrypted,
record.authorization_url_override.clone(),
record.token_url_override.clone(),
record.userinfo_url_override.clone(),
record.scopes.clone(),
record.attribute_mapping.clone(),
record.extra_config.clone(),
record.icon_url.clone(),
record.is_enabled,
)
.with_timestamps(created_at, now);
items.insert(record.provider_type.clone(), item.clone());
Ok(item)
}
async fn delete_oauth_provider_config(
&self,
provider_type: &str,
) -> Result<bool, DataLayerError> {
let mut items = self.items.write().expect("oauth provider repository lock");
Ok(items.remove(provider_type).is_some())
}
}
#[cfg(test)]
mod tests {
use super::InMemoryOAuthProviderRepository;
use crate::repository::oauth_providers::{
EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderWriteRepository,
StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord,
};
fn sample_provider(provider_type: &str) -> StoredOAuthProviderConfig {
StoredOAuthProviderConfig::new(
provider_type.to_string(),
format!("{provider_type} display"),
format!("{provider_type}-client"),
format!("https://{provider_type}.example.com/redirect"),
"https://frontend.example.com/auth/callback".to_string(),
)
.expect("provider should build")
}
fn sample_upsert(provider_type: &str) -> UpsertOAuthProviderConfigRecord {
UpsertOAuthProviderConfigRecord {
provider_type: provider_type.to_string(),
display_name: format!("{provider_type} display"),
client_id: format!("{provider_type}-client"),
client_secret_encrypted: EncryptedSecretUpdate::Preserve,
authorization_url_override: Some(format!("https://{provider_type}.example.com/auth")),
token_url_override: Some(format!("https://{provider_type}.example.com/token")),
userinfo_url_override: None,
scopes: Some(vec!["openid".to_string(), "profile".to_string()]),
redirect_uri: format!("https://{provider_type}.example.com/redirect"),
frontend_callback_url: "https://frontend.example.com/auth/callback".to_string(),
attribute_mapping: Some(serde_json::json!({"email": "email"})),
extra_config: Some(serde_json::json!({"team": true})),
icon_url: None,
is_enabled: true,
}
}
#[tokio::test]
async fn reads_and_mutates_oauth_provider_configs() {
let repository = InMemoryOAuthProviderRepository::seed(vec![
sample_provider("linuxdo"),
sample_provider("github"),
]);
let listed = repository
.list_oauth_provider_configs()
.await
.expect("list should succeed");
assert_eq!(listed.len(), 2);
assert_eq!(listed[0].provider_type, "github");
assert_eq!(listed[1].provider_type, "linuxdo");
let created = repository
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
client_secret_encrypted: EncryptedSecretUpdate::Set("secret-1".to_string()),
..sample_upsert("google")
})
.await
.expect("create should succeed");
assert_eq!(created.client_secret_encrypted.as_deref(), Some("secret-1"));
let updated = repository
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
client_secret_encrypted: EncryptedSecretUpdate::Clear,
..sample_upsert("google")
})
.await
.expect("update should succeed");
assert!(updated.client_secret_encrypted.is_none());
let deleted = repository
.delete_oauth_provider_config("google")
.await
.expect("delete should succeed");
assert!(deleted);
}
}
@@ -0,0 +1,13 @@
mod memory;
pub use aether_data_contracts::repository::oauth_providers::{
EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderRepository,
OAuthProviderWriteRepository, StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlOAuthProviderRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxOAuthProviderRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteOAuthProviderRepository;
pub use memory::InMemoryOAuthProviderRepository;
@@ -0,0 +1,513 @@
use std::collections::BTreeMap;
use std::sync::RwLock;
use async_trait::async_trait;
use super::{
score_with_delta, GetPoolMemberScoresByIdsQuery, ListPoolMemberProbeCandidatesQuery,
ListPoolMemberScoresQuery, ListRankedPoolMembersQuery, PoolMemberHardState, PoolMemberIdentity,
PoolMemberProbeAttempt, PoolMemberProbeResult, PoolMemberProbeStatus,
PoolMemberScheduleFeedback, PoolMemberScoreWriteRepository, PoolScoreReadRepository,
PoolScoreScope, StoredPoolMemberScore, UpsertPoolMemberScore,
};
use crate::repository::pool_scores::merge_score_reason_patch;
use crate::DataLayerError;
#[derive(Debug, Default)]
pub struct InMemoryPoolMemberScoreRepository {
scores: RwLock<BTreeMap<String, StoredPoolMemberScore>>,
}
impl InMemoryPoolMemberScoreRepository {
pub fn seed<I>(scores: I) -> Self
where
I: IntoIterator<Item = StoredPoolMemberScore>,
{
Self {
scores: RwLock::new(
scores
.into_iter()
.map(|score| (score.id.clone(), score))
.collect(),
),
}
}
fn matches_identity(score: &StoredPoolMemberScore, identity: &PoolMemberIdentity) -> bool {
score.pool_kind == identity.pool_kind
&& score.pool_id == identity.pool_id
&& score.member_kind == identity.member_kind
&& score.member_id == identity.member_id
}
fn matches_scope(score: &StoredPoolMemberScore, scope: Option<&PoolScoreScope>) -> bool {
let Some(scope) = scope else {
return true;
};
score.capability == scope.capability
&& score.scope_kind == scope.scope_kind
&& score.scope_id == scope.scope_id
}
fn sort_ranked(scores: &mut [StoredPoolMemberScore]) {
scores.sort_by(|left, right| {
right
.score
.partial_cmp(&left.score)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| {
right
.last_ranked_at
.unwrap_or(0)
.cmp(&left.last_ranked_at.unwrap_or(0))
})
.then_with(|| left.member_id.cmp(&right.member_id))
.then_with(|| left.id.cmp(&right.id))
});
}
fn sort_probe(scores: &mut [StoredPoolMemberScore]) {
scores.sort_by(|left, right| {
probe_priority(left)
.cmp(&probe_priority(right))
.then_with(|| right.probe_failure_count.cmp(&left.probe_failure_count))
.then_with(|| {
left.last_probe_success_at
.unwrap_or(0)
.cmp(&right.last_probe_success_at.unwrap_or(0))
})
.then_with(|| {
left.last_scheduled_at
.unwrap_or(0)
.cmp(&right.last_scheduled_at.unwrap_or(0))
.reverse()
})
.then_with(|| left.member_id.cmp(&right.member_id))
});
}
}
#[async_trait]
impl PoolScoreReadRepository for InMemoryPoolMemberScoreRepository {
async fn list_ranked_pool_members(
&self,
query: &ListRankedPoolMembersQuery,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
let hard_states = query
.hard_states
.iter()
.copied()
.collect::<std::collections::BTreeSet<_>>();
let probe_statuses = query.probe_statuses.as_ref().map(|items| {
items
.iter()
.copied()
.collect::<std::collections::BTreeSet<_>>()
});
let mut scores = self
.scores
.read()
.expect("pool member score repository lock")
.values()
.filter(|score| {
score.pool_kind == query.pool_kind
&& score.pool_id == query.pool_id
&& score.capability == query.capability
&& score.scope_kind == query.scope_kind
&& score.scope_id == query.scope_id
&& (hard_states.is_empty() || hard_states.contains(&score.hard_state))
&& probe_statuses
.as_ref()
.is_none_or(|statuses| statuses.contains(&score.probe_status))
})
.cloned()
.collect::<Vec<_>>();
Self::sort_ranked(&mut scores);
Ok(scores
.into_iter()
.skip(query.offset)
.take(query.limit.max(1))
.collect())
}
async fn list_pool_member_scores(
&self,
query: &ListPoolMemberScoresQuery,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
let hard_states = query
.hard_states
.iter()
.copied()
.collect::<std::collections::BTreeSet<_>>();
let probe_statuses = query.probe_statuses.as_ref().map(|items| {
items
.iter()
.copied()
.collect::<std::collections::BTreeSet<_>>()
});
let mut scores = self
.scores
.read()
.expect("pool member score repository lock")
.values()
.filter(|score| {
score.pool_kind == query.pool_kind
&& score.pool_id == query.pool_id
&& query
.capability
.as_ref()
.is_none_or(|capability| score.capability == *capability)
&& query
.scope_kind
.as_ref()
.is_none_or(|scope_kind| score.scope_kind == *scope_kind)
&& query
.scope_id
.as_ref()
.is_none_or(|scope_id| score.scope_id.as_ref() == Some(scope_id))
&& (hard_states.is_empty() || hard_states.contains(&score.hard_state))
&& probe_statuses
.as_ref()
.is_none_or(|statuses| statuses.contains(&score.probe_status))
})
.cloned()
.collect::<Vec<_>>();
Self::sort_ranked(&mut scores);
Ok(scores
.into_iter()
.skip(query.offset)
.take(query.limit.max(1))
.collect())
}
async fn list_pool_member_probe_candidates(
&self,
query: &ListPoolMemberProbeCandidatesQuery,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
let mut scores = self
.scores
.read()
.expect("pool member score repository lock")
.values()
.filter(|score| {
score.pool_kind == query.pool_kind
&& score.pool_id == query.pool_id
&& query
.capability
.as_ref()
.is_none_or(|capability| score.capability == *capability)
&& matches!(
score.hard_state,
PoolMemberHardState::Available
| PoolMemberHardState::Unknown
| PoolMemberHardState::Cooldown
| PoolMemberHardState::QuotaExhausted
)
&& match score.probe_status {
PoolMemberProbeStatus::Never
| PoolMemberProbeStatus::Failed
| PoolMemberProbeStatus::Stale => true,
PoolMemberProbeStatus::Ok => score
.last_probe_success_at
.is_none_or(|ts| ts <= query.stale_before_unix_secs),
PoolMemberProbeStatus::InProgress => score
.last_probe_attempt_at
.is_none_or(|ts| ts <= query.stale_before_unix_secs),
}
})
.cloned()
.collect::<Vec<_>>();
Self::sort_probe(&mut scores);
Ok(scores.into_iter().take(query.limit.max(1)).collect())
}
async fn get_pool_member_scores_by_ids(
&self,
query: &GetPoolMemberScoresByIdsQuery,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
let ids = query
.ids
.iter()
.cloned()
.collect::<std::collections::BTreeSet<_>>();
let scores = self
.scores
.read()
.expect("pool member score repository lock")
.values()
.filter(|score| ids.contains(&score.id))
.cloned()
.collect::<Vec<_>>();
Ok(scores)
}
}
#[async_trait]
impl PoolMemberScoreWriteRepository for InMemoryPoolMemberScoreRepository {
async fn upsert_pool_member_score(
&self,
score: UpsertPoolMemberScore,
) -> Result<StoredPoolMemberScore, DataLayerError> {
score.validate()?;
let stored = score.into_stored();
self.scores
.write()
.expect("pool member score repository lock")
.insert(stored.id.clone(), stored.clone());
Ok(stored)
}
async fn mark_pool_member_probe_in_progress(
&self,
attempt: PoolMemberProbeAttempt,
) -> Result<usize, DataLayerError> {
let mut updated = 0;
let mut guard = self
.scores
.write()
.expect("pool member score repository lock");
for score in guard.values_mut() {
if !Self::matches_identity(score, &attempt.identity)
|| !Self::matches_scope(score, attempt.scope.as_ref())
{
continue;
}
score.last_probe_attempt_at = Some(attempt.attempted_at);
score.probe_status = PoolMemberProbeStatus::InProgress;
score.score_reason = merge_score_reason_patch(
score.score_reason.clone(),
attempt.score_reason_patch.clone(),
);
score.updated_at = attempt.attempted_at;
updated += 1;
}
Ok(updated)
}
async fn record_pool_member_probe_result(
&self,
result: PoolMemberProbeResult,
) -> Result<usize, DataLayerError> {
let mut updated = 0;
let mut guard = self
.scores
.write()
.expect("pool member score repository lock");
for score in guard.values_mut() {
if !Self::matches_identity(score, &result.identity)
|| !Self::matches_scope(score, result.scope.as_ref())
{
continue;
}
score.last_probe_attempt_at = Some(result.attempted_at);
score.probe_status = result.probe_status;
if result.succeeded {
score.last_probe_success_at = Some(result.attempted_at);
score.probe_failure_count = 0;
} else {
score.last_probe_failure_at = Some(result.attempted_at);
score.probe_failure_count = score.probe_failure_count.saturating_add(1);
}
if let Some(hard_state) = result.hard_state {
score.hard_state = hard_state;
}
score.score_reason = merge_score_reason_patch(
score.score_reason.clone(),
result.score_reason_patch.clone(),
);
score.updated_at = result.attempted_at;
updated += 1;
}
Ok(updated)
}
async fn record_pool_member_schedule_feedback(
&self,
feedback: PoolMemberScheduleFeedback,
) -> Result<usize, DataLayerError> {
let mut updated = 0;
let mut guard = self
.scores
.write()
.expect("pool member score repository lock");
for score in guard.values_mut() {
if !Self::matches_identity(score, &feedback.identity)
|| !Self::matches_scope(score, feedback.scope.as_ref())
{
continue;
}
score.last_scheduled_at = Some(feedback.scheduled_at);
match feedback.succeeded {
Some(true) => {
score.last_success_at = Some(feedback.scheduled_at);
}
Some(false) => {
score.last_failure_at = Some(feedback.scheduled_at);
score.failure_count = score.failure_count.saturating_add(1);
}
None => {}
}
if let Some(hard_state) = feedback.hard_state {
score.hard_state = hard_state;
}
score.score = score_with_delta(score.score, feedback.score_delta);
score.score_reason = merge_score_reason_patch(
score.score_reason.clone(),
feedback.score_reason_patch.clone(),
);
score.updated_at = feedback.scheduled_at;
updated += 1;
}
Ok(updated)
}
async fn mark_pool_member_hard_state(
&self,
identity: &PoolMemberIdentity,
scope: Option<&PoolScoreScope>,
hard_state: PoolMemberHardState,
updated_at: u64,
) -> Result<usize, DataLayerError> {
let mut updated = 0;
let mut guard = self
.scores
.write()
.expect("pool member score repository lock");
for score in guard.values_mut() {
if Self::matches_identity(score, identity) && Self::matches_scope(score, scope) {
score.hard_state = hard_state;
score.updated_at = updated_at;
updated += 1;
}
}
Ok(updated)
}
async fn delete_pool_member_scores_for_member(
&self,
identity: &PoolMemberIdentity,
) -> Result<usize, DataLayerError> {
let mut guard = self
.scores
.write()
.expect("pool member score repository lock");
let before = guard.len();
guard.retain(|_, score| !Self::matches_identity(score, identity));
Ok(before.saturating_sub(guard.len()))
}
}
fn probe_priority(score: &StoredPoolMemberScore) -> u8 {
if score.last_scheduled_at.is_some() && score.probe_status != PoolMemberProbeStatus::Ok {
return 0;
}
match score.hard_state {
PoolMemberHardState::QuotaExhausted => 1,
PoolMemberHardState::Unknown => 2,
_ if score.probe_status == PoolMemberProbeStatus::Stale => 3,
_ => 4,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::repository::pool_scores::{
POOL_KIND_PROVIDER_KEY_POOL, POOL_MEMBER_KIND_PROVIDER_API_KEY,
POOL_SCORE_CAPABILITY_API_FORMAT, POOL_SCORE_SCOPE_KIND_MODEL,
};
fn score(id: &str, member_id: &str, value: f64) -> StoredPoolMemberScore {
StoredPoolMemberScore {
id: id.to_string(),
pool_kind: POOL_KIND_PROVIDER_KEY_POOL.to_string(),
pool_id: "provider-1".to_string(),
member_kind: POOL_MEMBER_KIND_PROVIDER_API_KEY.to_string(),
member_id: member_id.to_string(),
capability: POOL_SCORE_CAPABILITY_API_FORMAT.to_string(),
scope_kind: POOL_SCORE_SCOPE_KIND_MODEL.to_string(),
scope_id: Some("model-1".to_string()),
score: value,
hard_state: PoolMemberHardState::Available,
score_version: 1,
score_reason: serde_json::json!({}),
last_ranked_at: Some(1),
last_scheduled_at: None,
last_success_at: None,
last_failure_at: None,
failure_count: 0,
last_probe_attempt_at: None,
last_probe_success_at: None,
last_probe_failure_at: None,
probe_failure_count: 0,
probe_status: PoolMemberProbeStatus::Never,
updated_at: 1,
}
}
#[tokio::test]
async fn lists_ranked_members_by_score() {
let repository = InMemoryPoolMemberScoreRepository::seed(vec![
score("score-1", "key-1", 0.2),
score("score-2", "key-2", 0.9),
]);
let rows = repository
.list_ranked_pool_members(&ListRankedPoolMembersQuery {
pool_kind: POOL_KIND_PROVIDER_KEY_POOL.to_string(),
pool_id: "provider-1".to_string(),
capability: POOL_SCORE_CAPABILITY_API_FORMAT.to_string(),
scope_kind: POOL_SCORE_SCOPE_KIND_MODEL.to_string(),
scope_id: Some("model-1".to_string()),
hard_states: vec![PoolMemberHardState::Available],
probe_statuses: None,
offset: 0,
limit: 10,
})
.await
.expect("list should succeed");
assert_eq!(
rows.into_iter()
.map(|row| row.member_id)
.collect::<Vec<_>>(),
vec!["key-2".to_string(), "key-1".to_string()]
);
}
#[tokio::test]
async fn marks_probe_in_progress_without_incrementing_failure_count() {
let repository =
InMemoryPoolMemberScoreRepository::seed(vec![score("score-1", "key-1", 0.2)]);
let updated = repository
.mark_pool_member_probe_in_progress(PoolMemberProbeAttempt {
identity: PoolMemberIdentity::provider_api_key("provider-1", "key-1"),
scope: None,
attempted_at: 100,
score_reason_patch: Some(serde_json::json!({ "last_probe": "in_progress" })),
})
.await
.expect("mark should succeed");
assert_eq!(updated, 1);
let rows = repository
.list_pool_member_scores(&ListPoolMemberScoresQuery {
pool_kind: POOL_KIND_PROVIDER_KEY_POOL.to_string(),
pool_id: "provider-1".to_string(),
capability: None,
scope_kind: None,
scope_id: None,
hard_states: Vec::new(),
probe_statuses: Some(vec![PoolMemberProbeStatus::InProgress]),
offset: 0,
limit: 10,
})
.await
.expect("list should succeed");
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].last_probe_attempt_at, Some(100));
assert_eq!(rows[0].probe_failure_count, 0);
}
}
@@ -0,0 +1,11 @@
pub use aether_data_contracts::repository::pool_scores::*;
mod memory;
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlPoolMemberScoreRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::PostgresPoolMemberScoreRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqlitePoolMemberScoreRepository;
pub use memory::InMemoryPoolMemberScoreRepository;
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,17 @@
mod memory;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, ProviderCatalogReadRepository,
ProviderCatalogSnapshot, ProviderCatalogUpstreamMetadataNamespaceUpdate,
ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlProviderCatalogReadRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxProviderCatalogReadRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteProviderCatalogReadRepository;
pub use memory::InMemoryProviderCatalogReadRepository;
@@ -0,0 +1,219 @@
const KIRO_DEVICE_AUTH_SESSION_PREFIX: &str = "device_auth_session:";
const PROVIDER_OAUTH_BATCH_TASK_PREFIX: &str = "provider_oauth_batch_task:";
const PROVIDER_OAUTH_STATE_PREFIX: &str = "provider_oauth_state:";
pub const KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS: u64 = 60;
pub const PROVIDER_OAUTH_BATCH_TASK_TTL_SECS: u64 = 24 * 60 * 60;
pub const PROVIDER_OAUTH_STATE_TTL_SECS: u64 = 600;
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct StoredAdminProviderOAuthDeviceSession {
pub provider_id: String,
pub region: String,
pub client_id: String,
pub client_secret: String,
pub device_code: String,
#[serde(default)]
pub auth_type: Option<String>,
#[serde(default)]
pub social_provider: Option<String>,
#[serde(default)]
pub code_verifier: Option<String>,
#[serde(default)]
pub redirect_uri: Option<String>,
#[serde(default)]
pub machine_id: Option<String>,
pub interval: u64,
pub expires_at_unix_secs: u64,
pub status: String,
pub proxy_node_id: Option<String>,
pub created_at_unix_ms: u64,
pub key_id: Option<String>,
pub email: Option<String>,
pub replaced: bool,
pub error_msg: Option<String>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct StoredAdminProviderOAuthState {
pub key_id: String,
pub provider_id: String,
pub provider_type: String,
pub pkce_verifier: Option<String>,
}
pub fn provider_oauth_device_session_storage_key(session_id: &str) -> String {
format!("{KIRO_DEVICE_AUTH_SESSION_PREFIX}{session_id}")
}
pub fn provider_oauth_state_storage_key(nonce: &str) -> String {
format!("{PROVIDER_OAUTH_STATE_PREFIX}{nonce}")
}
pub fn provider_oauth_batch_task_storage_key(task_id: &str) -> String {
format!("{PROVIDER_OAUTH_BATCH_TASK_PREFIX}{task_id}")
}
pub fn build_provider_oauth_batch_task_status_payload(
provider_id: &str,
state: &serde_json::Map<String, serde_json::Value>,
) -> serde_json::Value {
let now_unix_secs = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|duration| duration.as_secs())
.unwrap_or(0);
let raw_status = state
.get("status")
.and_then(serde_json::Value::as_str)
.unwrap_or("failed");
let normalized_status = match raw_status {
"submitted" | "processing" | "completed" | "failed" => raw_status,
_ => "failed",
};
let error_samples = state
.get("error_samples")
.and_then(serde_json::Value::as_array)
.map(|items| {
items
.iter()
.filter(|item| item.is_object())
.cloned()
.collect::<Vec<_>>()
})
.unwrap_or_default();
serde_json::json!({
"task_id": state
.get("task_id")
.and_then(serde_json::Value::as_str)
.unwrap_or_default(),
"provider_id": provider_id,
"provider_type": state
.get("provider_type")
.and_then(serde_json::Value::as_str)
.unwrap_or_default(),
"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),
"success": state.get("success").and_then(serde_json::Value::as_i64).unwrap_or(0),
"failed": state.get("failed").and_then(serde_json::Value::as_i64).unwrap_or(0),
"created_count": state
.get("created_count")
.and_then(serde_json::Value::as_i64)
.unwrap_or(0),
"replaced_count": state
.get("replaced_count")
.and_then(serde_json::Value::as_i64)
.unwrap_or(0),
"progress_percent": state
.get("progress_percent")
.and_then(serde_json::Value::as_i64)
.unwrap_or(0)
.clamp(0, 100),
"message": state.get("message").cloned().unwrap_or(serde_json::Value::Null),
"error": state.get("error").cloned().unwrap_or(serde_json::Value::Null),
"error_samples": error_samples,
"created_at": state
.get("created_at")
.and_then(serde_json::Value::as_u64)
.unwrap_or(now_unix_secs),
"started_at": state.get("started_at").cloned().unwrap_or(serde_json::Value::Null),
"finished_at": state
.get("finished_at")
.cloned()
.unwrap_or(serde_json::Value::Null),
"updated_at": state
.get("updated_at")
.and_then(serde_json::Value::as_u64)
.unwrap_or(now_unix_secs),
})
}
#[cfg(test)]
mod tests {
use super::{
build_provider_oauth_batch_task_status_payload, provider_oauth_batch_task_storage_key,
provider_oauth_device_session_storage_key, provider_oauth_state_storage_key,
KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS, PROVIDER_OAUTH_BATCH_TASK_TTL_SECS,
PROVIDER_OAUTH_STATE_TTL_SECS,
};
use serde_json::json;
#[test]
fn builds_provider_oauth_storage_keys_with_expected_prefixes() {
assert_eq!(
provider_oauth_device_session_storage_key("session-123"),
"device_auth_session:session-123"
);
assert_eq!(
provider_oauth_state_storage_key("nonce-123"),
"provider_oauth_state:nonce-123"
);
assert_eq!(
provider_oauth_batch_task_storage_key("task-123"),
"provider_oauth_batch_task:task-123"
);
}
#[test]
fn provider_oauth_storage_ttls_match_gateway_expectations() {
assert_eq!(KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS, 60);
assert_eq!(PROVIDER_OAUTH_BATCH_TASK_TTL_SECS, 24 * 60 * 60);
assert_eq!(PROVIDER_OAUTH_STATE_TTL_SECS, 600);
}
#[test]
fn batch_task_status_payload_normalizes_status_and_clamps_progress() {
let input = json!({
"task_id": "task-123",
"provider_type": "codex",
"status": "weird",
"total": 4,
"processed": 2,
"success": 1,
"failed": 1,
"created_count": 0,
"replaced_count": 1,
"progress_percent": 999,
"error_samples": [
{"detail": "x"},
"skip-me"
],
"created_at": 1u64,
"updated_at": 2u64
});
let payload = build_provider_oauth_batch_task_status_payload(
"provider-123",
input.as_object().expect("input should be object"),
);
assert_eq!(
payload.get("provider_id").and_then(|v| v.as_str()),
Some("provider-123")
);
assert_eq!(
payload.get("status").and_then(|v| v.as_str()),
Some("failed")
);
assert_eq!(
payload.get("progress_percent").and_then(|v| v.as_i64()),
Some(100)
);
assert_eq!(
payload.get("created_count").and_then(|v| v.as_i64()),
Some(0)
);
assert_eq!(
payload.get("replaced_count").and_then(|v| v.as_i64()),
Some(1)
);
assert_eq!(
payload
.get("error_samples")
.and_then(|v| v.as_array())
.map(Vec::len),
Some(1)
);
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,32 @@
mod memory;
pub use aether_data_contracts::repository::proxy_nodes::*;
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlProxyNodeReadRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxProxyNodeRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteProxyNodeReadRepository;
pub use memory::InMemoryProxyNodeRepository;
pub fn log_reported_tunnel_error_event(
node_id: &str,
event: &TunnelErrorEventRecord,
received_at_unix_secs: u64,
) {
tracing::warn!(
event_name = "proxy_tunnel_error_reported",
source = "heartbeat",
node_id = %node_id,
category = %event.category,
message = %event.message,
severity = ?event.severity,
component = ?event.component,
summary = ?event.summary,
operator_action = ?event.operator_action,
error_reported_at_unix_secs = event.timestamp_unix_secs,
error_reported_at_unix_ms = ?event.timestamp_unix_ms,
report_received_at_unix_secs = received_at_unix_secs,
"proxy reported tunnel error via heartbeat"
);
}
@@ -0,0 +1,151 @@
use std::collections::BTreeMap;
use std::sync::RwLock;
use async_trait::async_trait;
use super::{
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
};
use crate::DataLayerError;
use aether_wallet::{ProviderBillingType, ProviderQuotaSnapshot};
#[derive(Debug, Default)]
pub struct InMemoryProviderQuotaRepository {
by_provider_id: RwLock<BTreeMap<String, StoredProviderQuotaSnapshot>>,
}
impl InMemoryProviderQuotaRepository {
pub fn seed<I>(items: I) -> Self
where
I: IntoIterator<Item = StoredProviderQuotaSnapshot>,
{
let mut by_provider_id = BTreeMap::new();
for item in items {
by_provider_id.insert(item.provider_id.clone(), item);
}
Self {
by_provider_id: RwLock::new(by_provider_id),
}
}
}
#[async_trait]
impl ProviderQuotaReadRepository for InMemoryProviderQuotaRepository {
async fn find_by_provider_id(
&self,
provider_id: &str,
) -> Result<Option<StoredProviderQuotaSnapshot>, DataLayerError> {
Ok(self
.by_provider_id
.read()
.expect("quota repository lock")
.get(provider_id)
.cloned())
}
async fn find_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderQuotaSnapshot>, DataLayerError> {
let quotas = self.by_provider_id.read().expect("quota repository lock");
Ok(provider_ids
.iter()
.filter_map(|provider_id| quotas.get(provider_id).cloned())
.collect())
}
}
#[async_trait]
impl ProviderQuotaWriteRepository for InMemoryProviderQuotaRepository {
async fn reset_due(&self, now_unix_secs: u64) -> Result<usize, DataLayerError> {
let mut count = 0usize;
let mut quotas = self.by_provider_id.write().expect("quota repository lock");
for quota in quotas.values_mut() {
let snapshot = ProviderQuotaSnapshot {
provider_id: quota.provider_id.clone(),
billing_type: ProviderBillingType::parse(&quota.billing_type),
monthly_quota_usd: quota.monthly_quota_usd,
monthly_used_usd: quota.monthly_used_usd,
quota_reset_day: quota.quota_reset_day,
quota_last_reset_at_unix_secs: quota.quota_last_reset_at_unix_secs,
quota_expires_at_unix_secs: quota.quota_expires_at_unix_secs,
is_active: quota.is_active,
};
if snapshot.should_reset(now_unix_secs) {
quota.monthly_used_usd = 0.0;
quota.quota_last_reset_at_unix_secs = Some(now_unix_secs);
count += 1;
}
}
Ok(count)
}
}
#[cfg(test)]
mod tests {
use super::InMemoryProviderQuotaRepository;
use crate::repository::quota::{
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
};
fn sample_quota() -> StoredProviderQuotaSnapshot {
StoredProviderQuotaSnapshot::new(
"provider-1".to_string(),
"monthly_quota".to_string(),
Some(20.0),
5.0,
Some(7),
Some(1_000),
None,
true,
)
.expect("quota should build")
}
#[tokio::test]
async fn resets_due_monthly_quota() {
let repository = InMemoryProviderQuotaRepository::seed(vec![sample_quota()]);
let reset = repository
.reset_due(1_000 + 7 * 24 * 60 * 60)
.await
.expect("reset should succeed");
assert_eq!(reset, 1);
let stored = repository
.find_by_provider_id("provider-1")
.await
.expect("lookup should succeed")
.expect("quota should exist");
assert_eq!(stored.monthly_used_usd, 0.0);
}
#[tokio::test]
async fn finds_quotas_by_provider_ids() {
let repository = InMemoryProviderQuotaRepository::seed(vec![
sample_quota(),
StoredProviderQuotaSnapshot::new(
"provider-2".to_string(),
"payg".to_string(),
None,
1.5,
None,
None,
None,
true,
)
.expect("quota should build"),
]);
let stored = repository
.find_by_provider_ids(&[
"provider-2".to_string(),
"missing".to_string(),
"provider-1".to_string(),
])
.await
.expect("lookup should succeed");
assert_eq!(stored.len(), 2);
assert_eq!(stored[0].provider_id, "provider-2");
assert_eq!(stored[1].provider_id, "provider-1");
}
}
@@ -0,0 +1,14 @@
mod memory;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::quota::{
ProviderQuotaReadRepository, ProviderQuotaRepository, ProviderQuotaWriteRepository,
StoredProviderQuotaSnapshot,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlProviderQuotaRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxProviderQuotaRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteProviderQuotaRepository;
pub use memory::InMemoryProviderQuotaRepository;
@@ -0,0 +1,342 @@
use std::collections::BTreeMap;
use std::sync::RwLock;
use async_trait::async_trait;
use super::{
CreateRoutingGroupBindingRecord, CreateRoutingGroupRecord, CreateRoutingGroupVersionRecord,
RoutingGroupBindingQuery, RoutingGroupLookupKey, RoutingGroupReadRepository,
RoutingGroupWriteRepository, StoredRoutingGroup, StoredRoutingGroupBinding,
StoredRoutingGroupVersion, UpdateRoutingGroupBindingRecord, UpdateRoutingGroupRecord,
};
use crate::DataLayerError;
#[derive(Debug, Default)]
pub struct InMemoryRoutingGroupRepository {
groups: RwLock<BTreeMap<String, StoredRoutingGroup>>,
bindings: RwLock<BTreeMap<String, StoredRoutingGroupBinding>>,
versions: RwLock<BTreeMap<String, StoredRoutingGroupVersion>>,
}
impl InMemoryRoutingGroupRepository {
pub fn seed<I, B, V>(groups: I, bindings: B, versions: V) -> Self
where
I: IntoIterator<Item = StoredRoutingGroup>,
B: IntoIterator<Item = StoredRoutingGroupBinding>,
V: IntoIterator<Item = StoredRoutingGroupVersion>,
{
Self {
groups: RwLock::new(
groups
.into_iter()
.map(|item| (item.id.clone(), item))
.collect(),
),
bindings: RwLock::new(
bindings
.into_iter()
.map(|item| (item.id.clone(), item))
.collect(),
),
versions: RwLock::new(
versions
.into_iter()
.map(|item| (item.id.clone(), item))
.collect(),
),
}
}
}
#[async_trait]
impl RoutingGroupReadRepository for InMemoryRoutingGroupRepository {
async fn list_routing_groups(&self) -> Result<Vec<StoredRoutingGroup>, DataLayerError> {
let mut groups = self
.groups
.read()
.expect("routing group repository lock")
.values()
.cloned()
.collect::<Vec<_>>();
groups.sort_by(|left, right| left.name.cmp(&right.name).then(left.id.cmp(&right.id)));
Ok(groups)
}
async fn find_routing_group(
&self,
lookup: RoutingGroupLookupKey<'_>,
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
let groups = self.groups.read().expect("routing group repository lock");
Ok(match lookup {
RoutingGroupLookupKey::Id(id) => groups.get(id).cloned(),
RoutingGroupLookupKey::Name(name) => {
groups.values().find(|group| group.name == name).cloned()
}
RoutingGroupLookupKey::SystemDefault => groups
.values()
.find(|group| group.is_system_default && group.enabled)
.cloned(),
})
}
async fn list_routing_group_bindings(
&self,
query: &RoutingGroupBindingQuery,
) -> Result<Vec<StoredRoutingGroupBinding>, DataLayerError> {
let mut rows = self
.bindings
.read()
.expect("routing group binding repository lock")
.values()
.filter(|row| {
query
.group_id
.as_ref()
.is_none_or(|group_id| &row.group_id == group_id)
&& query
.subject_type
.as_ref()
.is_none_or(|subject_type| &row.subject_type == subject_type)
&& query
.subject_id
.as_ref()
.is_none_or(|subject_id| &row.subject_id == subject_id)
})
.cloned()
.collect::<Vec<_>>();
rows.sort_by(|left, right| {
left.created_at
.cmp(&right.created_at)
.then(left.id.cmp(&right.id))
});
Ok(rows)
}
async fn list_routing_group_versions(
&self,
group_id: &str,
) -> Result<Vec<StoredRoutingGroupVersion>, DataLayerError> {
let mut rows = self
.versions
.read()
.expect("routing group version repository lock")
.values()
.filter(|row| row.group_id == group_id)
.cloned()
.collect::<Vec<_>>();
rows.sort_by(|left, right| {
right
.version
.cmp(&left.version)
.then(right.created_at.cmp(&left.created_at))
});
Ok(rows)
}
}
#[async_trait]
impl RoutingGroupWriteRepository for InMemoryRoutingGroupRepository {
async fn create_routing_group(
&self,
record: CreateRoutingGroupRecord,
) -> Result<StoredRoutingGroup, DataLayerError> {
let group = StoredRoutingGroup::new(record)?;
self.groups
.write()
.expect("routing group repository lock")
.insert(group.id.clone(), group.clone());
Ok(group)
}
async fn update_routing_group(
&self,
id: &str,
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 {
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(),
));
}
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()))
}
async fn delete_routing_group(&self, id: &str) -> Result<bool, DataLayerError> {
Ok(self
.groups
.write()
.expect("routing group repository lock")
.remove(id)
.is_some())
}
async fn create_routing_group_binding(
&self,
record: CreateRoutingGroupBindingRecord,
) -> Result<StoredRoutingGroupBinding, DataLayerError> {
let binding = StoredRoutingGroupBinding::new(record)?;
self.bindings
.write()
.expect("routing group binding repository lock")
.insert(binding.id.clone(), binding.clone());
Ok(binding)
}
async fn delete_routing_group_binding(&self, id: &str) -> Result<bool, DataLayerError> {
Ok(self
.bindings
.write()
.expect("routing group binding repository lock")
.remove(id)
.is_some())
}
async fn update_routing_group_binding(
&self,
id: &str,
patch: UpdateRoutingGroupBindingRecord,
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
let mut bindings = self
.bindings
.write()
.expect("routing group binding repository lock");
let Some(binding) = bindings.get_mut(id) 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(),
));
}
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()))
}
async fn create_routing_group_version(
&self,
record: CreateRoutingGroupVersionRecord,
) -> Result<StoredRoutingGroupVersion, DataLayerError> {
let version = StoredRoutingGroupVersion::new(record)?;
self.versions
.write()
.expect("routing group version repository lock")
.insert(version.id.clone(), version.clone());
Ok(version)
}
}
#[cfg(test)]
mod tests {
use aether_data_contracts::repository::routing_profiles::RoutingGroupBindingSubject;
use serde_json::json;
use super::*;
#[tokio::test]
async fn stores_groups_bindings_and_versions() {
let repository = InMemoryRoutingGroupRepository::default();
let group = repository
.create_routing_group(CreateRoutingGroupRecord {
id: "group-1".to_string(),
name: "default".to_string(),
description: None,
enabled: true,
is_system_default: true,
config_json: json!({}),
version: 1,
created_at: 1,
updated_at: 1,
published_at: None,
})
.await
.expect("group should store");
assert_eq!(
repository
.find_routing_group(RoutingGroupLookupKey::SystemDefault)
.await
.unwrap()
.as_ref()
.map(|group| group.id.as_str()),
Some(group.id.as_str())
);
repository
.create_routing_group_binding(CreateRoutingGroupBindingRecord {
id: "binding-1".to_string(),
group_id: "group-1".to_string(),
subject_type: RoutingGroupBindingSubject::ApiKey,
subject_id: "api-key-1".to_string(),
is_default: true,
allow_explicit_select: true,
created_at: 1,
updated_at: 1,
})
.await
.unwrap();
assert_eq!(
repository
.list_routing_group_bindings(&RoutingGroupBindingQuery {
subject_type: Some(RoutingGroupBindingSubject::ApiKey),
subject_id: Some("api-key-1".to_string()),
group_id: None,
})
.await
.unwrap()
.len(),
1
);
}
}
@@ -0,0 +1,16 @@
mod memory;
pub(crate) use aether_data_contracts::repository::routing_profiles::{
CreateRoutingGroupBindingRecord, CreateRoutingGroupRecord, CreateRoutingGroupVersionRecord,
RoutingGroupBindingQuery, RoutingGroupLookupKey, RoutingGroupReadRepository,
RoutingGroupWriteRepository, StoredRoutingGroup, StoredRoutingGroupBinding,
StoredRoutingGroupVersion, UpdateRoutingGroupBindingRecord, UpdateRoutingGroupRecord,
};
pub use memory::InMemoryRoutingGroupRepository;
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlRoutingGroupRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::PostgresRoutingGroupRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteRoutingGroupRepository;
@@ -0,0 +1,424 @@
use std::collections::BTreeMap;
use std::sync::{Arc, RwLock};
use async_trait::async_trait;
use super::{
plan_finite_wallet_debit, settlement_billable_cost_usd,
settlement_billing_status_for_usage_status, SettlementWriteRepository, StoredUsageSettlement,
UsageSettlementInput, SETTLEMENT_EPSILON_USD,
};
use crate::repository::wallet::{InMemoryWalletRepository, StoredWalletSnapshot};
use crate::DataLayerError;
#[derive(Debug)]
enum InMemorySettlementWalletStore {
Owned(RwLock<BTreeMap<String, StoredWalletSnapshot>>),
Shared(Arc<InMemoryWalletRepository>),
}
impl Default for InMemorySettlementWalletStore {
fn default() -> Self {
Self::Owned(RwLock::new(BTreeMap::new()))
}
}
impl InMemorySettlementWalletStore {
fn seeded<I>(items: I) -> Self
where
I: IntoIterator<Item = StoredWalletSnapshot>,
{
let mut wallets_by_id = BTreeMap::new();
for item in items {
wallets_by_id.insert(item.id.clone(), item);
}
Self::Owned(RwLock::new(wallets_by_id))
}
fn with_mut<R>(&self, f: impl FnOnce(&mut BTreeMap<String, StoredWalletSnapshot>) -> R) -> R {
match self {
Self::Owned(wallets_by_id) => {
let mut wallets = wallets_by_id.write().expect("settlement repo lock");
f(&mut wallets)
}
Self::Shared(repository) => repository.with_wallets_mut(f),
}
}
}
#[derive(Debug, Default)]
pub struct InMemorySettlementRepository {
wallets: InMemorySettlementWalletStore,
provider_monthly_used: RwLock<BTreeMap<String, f64>>,
settlements: RwLock<BTreeMap<String, StoredUsageSettlement>>,
}
impl InMemorySettlementRepository {
pub fn seed<I>(items: I) -> Self
where
I: IntoIterator<Item = StoredWalletSnapshot>,
{
Self {
wallets: InMemorySettlementWalletStore::seeded(items),
provider_monthly_used: RwLock::new(BTreeMap::new()),
settlements: RwLock::new(BTreeMap::new()),
}
}
pub fn from_wallet_repository(wallet_repository: Arc<InMemoryWalletRepository>) -> Self {
Self {
wallets: InMemorySettlementWalletStore::Shared(wallet_repository),
provider_monthly_used: RwLock::new(BTreeMap::new()),
settlements: RwLock::new(BTreeMap::new()),
}
}
}
#[async_trait]
impl SettlementWriteRepository for InMemorySettlementRepository {
async fn settle_usage(
&self,
input: UsageSettlementInput,
) -> Result<Option<StoredUsageSettlement>, DataLayerError> {
input.validate()?;
if input.billing_status != "pending" {
let existing = self
.settlements
.read()
.expect("settlement snapshot lock")
.get(&input.request_id)
.cloned();
return Ok(Some(existing.unwrap_or(StoredUsageSettlement {
request_id: input.request_id,
wallet_id: None,
billing_status: input.billing_status,
wallet_balance_before: None,
wallet_balance_after: None,
wallet_recharge_balance_before: None,
wallet_recharge_balance_after: None,
wallet_gift_balance_before: None,
wallet_gift_balance_after: None,
provider_monthly_used_usd: None,
finalized_at_unix_secs: input.finalized_at_unix_secs,
})));
}
let mut final_billing_status =
settlement_billing_status_for_usage_status(&input.status).to_string();
let billable_cost_usd = settlement_billable_cost_usd(&input);
let mut settlement = self.wallets.with_mut(|wallets| {
let wallet_id = input
.api_key_id
.as_deref()
.and_then(|api_key_id| {
wallets
.values()
.find(|wallet| wallet.api_key_id.as_deref() == Some(api_key_id))
.map(|wallet| wallet.id.clone())
})
.or_else(|| {
if input.api_key_is_standalone {
return None;
}
input.user_id.as_deref().and_then(|user_id| {
wallets
.values()
.find(|wallet| wallet.user_id.as_deref() == Some(user_id))
.map(|wallet| wallet.id.clone())
})
});
let wallet = wallet_id
.as_deref()
.and_then(|wallet_id| wallets.get_mut(wallet_id));
let mut settlement = StoredUsageSettlement {
request_id: input.request_id.clone(),
wallet_id: None,
billing_status: final_billing_status.to_string(),
wallet_balance_before: None,
wallet_balance_after: None,
wallet_recharge_balance_before: None,
wallet_recharge_balance_after: None,
wallet_gift_balance_before: None,
wallet_gift_balance_after: None,
provider_monthly_used_usd: None,
finalized_at_unix_secs: input.finalized_at_unix_secs,
};
if let Some(wallet) = wallet {
let before_recharge = wallet.balance;
let before_gift = wallet.gift_balance;
let before_total = before_recharge + before_gift;
settlement.wallet_id = Some(wallet.id.clone());
settlement.wallet_balance_before = Some(before_total);
settlement.wallet_recharge_balance_before = Some(before_recharge);
settlement.wallet_gift_balance_before = Some(before_gift);
if final_billing_status == "settled" {
if wallet.limit_mode.eq_ignore_ascii_case("unlimited") {
wallet.total_consumed += billable_cost_usd;
} else {
let debit_plan = plan_finite_wallet_debit(
before_recharge,
before_gift,
billable_cost_usd,
);
(wallet.balance, wallet.gift_balance) =
debit_plan.after_balances(before_recharge, before_gift);
wallet.total_consumed += billable_cost_usd;
}
}
settlement.wallet_recharge_balance_after = Some(wallet.balance);
settlement.wallet_gift_balance_after = Some(wallet.gift_balance);
settlement.wallet_balance_after = Some(wallet.balance + wallet.gift_balance);
} else if final_billing_status == "settled"
&& billable_cost_usd > SETTLEMENT_EPSILON_USD
{
final_billing_status = "insufficient_quota".to_string();
settlement.billing_status = final_billing_status.clone();
}
settlement
});
if final_billing_status == "settled" {
if let Some(provider_id) = input.provider_id {
let mut quotas = self
.provider_monthly_used
.write()
.expect("provider quota lock");
let value = quotas.entry(provider_id).or_insert(0.0);
*value += input.actual_total_cost_usd;
settlement.provider_monthly_used_usd = Some(*value);
}
}
self.settlements
.write()
.expect("settlement snapshot lock")
.insert(settlement.request_id.clone(), settlement.clone());
Ok(Some(settlement))
}
}
#[cfg(test)]
mod tests {
use super::InMemorySettlementRepository;
use crate::repository::settlement::{SettlementWriteRepository, UsageSettlementInput};
use crate::repository::wallet::StoredWalletSnapshot;
fn sample_wallet() -> StoredWalletSnapshot {
StoredWalletSnapshot::new(
"wallet-1".to_string(),
Some("user-1".to_string()),
Some("key-1".to_string()),
10.0,
2.0,
"finite".to_string(),
"USD".to_string(),
"active".to_string(),
0.0,
0.0,
0.0,
0.0,
100,
)
.expect("wallet should build")
}
fn sample_user_wallet(wallet_id: &str, user_id: &str) -> StoredWalletSnapshot {
StoredWalletSnapshot::new(
wallet_id.to_string(),
Some(user_id.to_string()),
None,
10.0,
2.0,
"finite".to_string(),
"USD".to_string(),
"active".to_string(),
0.0,
0.0,
0.0,
0.0,
100,
)
.expect("wallet should build")
}
#[tokio::test]
async fn settles_usage_against_wallet_and_provider_quota() {
let repository = InMemorySettlementRepository::seed(vec![sample_wallet()]);
let settlement = repository
.settle_usage(UsageSettlementInput {
request_id: "req-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("key-1".to_string()),
api_key_is_standalone: false,
provider_id: Some("provider-1".to_string()),
status: "completed".to_string(),
billing_status: "pending".to_string(),
total_cost_usd: 3.0,
actual_total_cost_usd: 6.0,
finalized_at_unix_secs: Some(200),
})
.await
.expect("settlement should succeed")
.expect("settlement should exist");
assert_eq!(settlement.billing_status, "settled");
assert_eq!(settlement.wallet_balance_before, Some(12.0));
assert_eq!(settlement.wallet_balance_after, Some(6.0));
assert_eq!(settlement.provider_monthly_used_usd, Some(6.0));
}
#[tokio::test]
async fn normal_key_settlement_falls_back_to_user_wallet() {
let repository =
InMemorySettlementRepository::seed(vec![sample_user_wallet("wallet-user-1", "user-1")]);
let settlement = repository
.settle_usage(UsageSettlementInput {
request_id: "req-user-wallet".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("normal-key-without-wallet".to_string()),
api_key_is_standalone: false,
provider_id: None,
status: "completed".to_string(),
billing_status: "pending".to_string(),
total_cost_usd: 3.0,
actual_total_cost_usd: 6.0,
finalized_at_unix_secs: Some(200),
})
.await
.expect("settlement should succeed")
.expect("settlement should exist");
assert_eq!(settlement.wallet_id.as_deref(), Some("wallet-user-1"));
assert_eq!(settlement.wallet_balance_before, Some(12.0));
assert_eq!(settlement.wallet_balance_after, Some(6.0));
}
#[tokio::test]
async fn settles_cancelled_usage_against_wallet_and_provider_quota() {
let repository = InMemorySettlementRepository::seed(vec![sample_wallet()]);
let settlement = repository
.settle_usage(UsageSettlementInput {
request_id: "req-cancelled".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("key-1".to_string()),
api_key_is_standalone: false,
provider_id: Some("provider-1".to_string()),
status: "cancelled".to_string(),
billing_status: "pending".to_string(),
total_cost_usd: 3.0,
actual_total_cost_usd: 6.0,
finalized_at_unix_secs: Some(200),
})
.await
.expect("settlement should succeed")
.expect("settlement should exist");
assert_eq!(settlement.billing_status, "settled");
assert_eq!(settlement.wallet_balance_before, Some(12.0));
assert_eq!(settlement.wallet_balance_after, Some(6.0));
assert_eq!(settlement.provider_monthly_used_usd, Some(6.0));
}
#[tokio::test]
async fn standalone_key_settlement_never_falls_back_to_owner_wallet() {
let repository = InMemorySettlementRepository::seed(vec![sample_user_wallet(
"wallet-admin-owner",
"admin-owner",
)]);
let settlement = repository
.settle_usage(UsageSettlementInput {
request_id: "req-standalone-no-key-wallet".to_string(),
user_id: Some("admin-owner".to_string()),
api_key_id: Some("standalone-key-without-wallet".to_string()),
api_key_is_standalone: true,
provider_id: None,
status: "completed".to_string(),
billing_status: "pending".to_string(),
total_cost_usd: 3.0,
actual_total_cost_usd: 1.5,
finalized_at_unix_secs: Some(200),
})
.await
.expect("settlement should succeed")
.expect("settlement should exist");
assert_eq!(settlement.billing_status, "insufficient_quota");
assert_eq!(settlement.wallet_id, None);
assert_eq!(settlement.wallet_balance_before, None);
assert_eq!(settlement.wallet_balance_after, None);
}
#[tokio::test]
async fn finite_wallet_insufficient_balance_overdraws_and_settles() {
let repository = InMemorySettlementRepository::seed(vec![sample_wallet()]);
let settlement = repository
.settle_usage(UsageSettlementInput {
request_id: "req-insufficient-wallet".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("key-1".to_string()),
api_key_is_standalone: false,
provider_id: Some("provider-1".to_string()),
status: "completed".to_string(),
billing_status: "pending".to_string(),
total_cost_usd: 3.0,
actual_total_cost_usd: 15.0,
finalized_at_unix_secs: Some(200),
})
.await
.expect("settlement should succeed")
.expect("settlement should exist");
assert_eq!(settlement.billing_status, "settled");
assert_eq!(settlement.wallet_balance_before, Some(12.0));
assert_eq!(settlement.wallet_balance_after, Some(-3.0));
assert_eq!(settlement.wallet_recharge_balance_after, Some(-3.0));
assert_eq!(settlement.wallet_gift_balance_after, Some(0.0));
assert_eq!(settlement.provider_monthly_used_usd, Some(15.0));
}
#[tokio::test]
async fn returns_stored_snapshot_when_usage_is_already_finalized() {
let repository = InMemorySettlementRepository::seed(vec![sample_wallet()]);
let settled = repository
.settle_usage(UsageSettlementInput {
request_id: "req-2".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("key-1".to_string()),
api_key_is_standalone: false,
provider_id: Some("provider-1".to_string()),
status: "completed".to_string(),
billing_status: "pending".to_string(),
total_cost_usd: 2.0,
actual_total_cost_usd: 1.0,
finalized_at_unix_secs: Some(250),
})
.await
.expect("settlement should succeed")
.expect("settlement should exist");
let replay = repository
.settle_usage(UsageSettlementInput {
request_id: "req-2".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("key-1".to_string()),
api_key_is_standalone: false,
provider_id: Some("provider-1".to_string()),
status: "completed".to_string(),
billing_status: "settled".to_string(),
total_cost_usd: 2.0,
actual_total_cost_usd: 1.0,
finalized_at_unix_secs: Some(250),
})
.await
.expect("replay should succeed")
.expect("snapshot should exist");
assert_eq!(replay, settled);
}
}
@@ -0,0 +1,27 @@
mod memory;
pub use aether_data_contracts::repository::settlement::*;
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlSettlementRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxSettlementRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteSettlementRepository;
pub use memory::InMemorySettlementRepository;
#[cfg(test)]
mod tests {
use aether_data_contracts::repository::settlement::settlement_billing_status_for_usage_status;
#[test]
fn cancelled_usage_status_is_billable() {
assert_eq!(
settlement_billing_status_for_usage_status("completed"),
"settled"
);
assert_eq!(
settlement_billing_status_for_usage_status("cancelled"),
"settled"
);
assert_eq!(settlement_billing_status_for_usage_status("failed"), "void");
}
}
@@ -0,0 +1,132 @@
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredSystemConfigEntry {
pub key: String,
pub value: serde_json::Value,
pub description: Option<String>,
pub updated_at_unix_secs: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct AdminSecurityBlacklistEntry {
pub ip_address: String,
pub reason: String,
pub ttl_seconds: Option<i64>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize, Default)]
pub struct AdminSystemStats {
pub total_users: u64,
pub active_users: u64,
pub total_api_keys: u64,
pub total_requests: u64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AdminSystemUsageAggregateImportMode {
Skip,
Overwrite,
Error,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct AdminSystemStatsDailyAggregate {
pub date_unix_secs: u64,
pub total_requests: u64,
pub success_requests: u64,
pub error_requests: u64,
pub input_tokens: u64,
pub output_tokens: u64,
pub cache_creation_tokens: u64,
pub cache_read_tokens: u64,
pub total_cost: f64,
pub actual_total_cost: f64,
pub is_complete: bool,
pub aggregated_at_unix_secs: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct AdminSystemStatsUserDailyAggregate {
pub user_id: String,
pub username: Option<String>,
pub date_unix_secs: u64,
pub total_requests: u64,
pub success_requests: u64,
pub error_requests: u64,
pub input_tokens: u64,
pub output_tokens: u64,
pub cache_creation_tokens: u64,
pub cache_read_tokens: u64,
pub total_cost: f64,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct AdminSystemStatsDailyApiKeyAggregate {
pub api_key_id: String,
pub api_key_name: Option<String>,
pub date_unix_secs: u64,
pub total_requests: u64,
pub success_requests: u64,
pub error_requests: u64,
pub input_tokens: u64,
pub output_tokens: u64,
pub cache_creation_tokens: u64,
pub cache_read_tokens: u64,
pub total_cost: f64,
}
#[derive(Debug, Clone, Default, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct AdminSystemUsageAggregateSnapshot {
#[serde(default)]
pub stats_daily: Vec<AdminSystemStatsDailyAggregate>,
#[serde(default)]
pub stats_user_daily: Vec<AdminSystemStatsUserDailyAggregate>,
#[serde(default)]
pub stats_daily_api_key: Vec<AdminSystemStatsDailyApiKeyAggregate>,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct AdminSystemUsageAggregateImportCounter {
pub created: u64,
pub updated: u64,
pub skipped: u64,
}
#[derive(Debug, Clone, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct AdminSystemUsageAggregateImportSummary {
pub stats_daily: AdminSystemUsageAggregateImportCounter,
pub stats_user_daily: AdminSystemUsageAggregateImportCounter,
pub stats_daily_api_key: AdminSystemUsageAggregateImportCounter,
pub skipped_unmapped_user_daily: u64,
pub skipped_unmapped_api_key_daily: u64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AdminSystemPurgeTarget {
Config,
Users,
Usage,
AuditLogs,
RequestBodies,
Stats,
}
#[derive(Debug, Clone, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
pub struct AdminSystemPurgeSummary {
pub affected: std::collections::BTreeMap<String, u64>,
}
impl AdminSystemPurgeSummary {
pub fn add(&mut self, key: impl Into<String>, count: u64) {
*self.affected.entry(key.into()).or_insert(0) += count;
}
pub fn merge(&mut self, other: &Self) {
for (key, count) in &other.affected {
self.add(key.clone(), *count);
}
}
pub fn total(&self) -> u64 {
self.affected.values().copied().sum()
}
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,880 @@
#[cfg(feature = "mysql")]
macro_rules! impl_materialized_usage_read_repository {
($repository:ty) => {
#[async_trait::async_trait]
impl $crate::repository::usage::UsageReadRepository for $repository {
async fn find_by_id(
&self,
id: &str,
) -> Result<
Option<$crate::repository::usage::StoredRequestUsageAudit>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::find_by_id(&repository, id).await
}
async fn list_by_ids(
&self,
ids: &[String],
) -> Result<
Vec<$crate::repository::usage::StoredRequestUsageAudit>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::list_by_ids(&repository, ids).await
}
async fn find_by_request_id(
&self,
request_id: &str,
) -> Result<
Option<$crate::repository::usage::StoredRequestUsageAudit>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::find_by_request_id(&repository, request_id).await
}
async fn resolve_body_ref(
&self,
body_ref: &str,
) -> Result<Option<serde_json::Value>, $crate::DataLayerError> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::resolve_body_ref(&repository, body_ref).await
}
async fn list_usage_audits(
&self,
query: &$crate::repository::usage::UsageAuditListQuery,
) -> Result<
Vec<$crate::repository::usage::StoredRequestUsageAudit>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::list_usage_audits(&repository, query).await
}
async fn count_usage_audits(
&self,
query: &$crate::repository::usage::UsageAuditListQuery,
) -> Result<u64, $crate::DataLayerError> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::count_usage_audits(&repository, query).await
}
async fn list_usage_audits_by_keyword_search(
&self,
query: &$crate::repository::usage::UsageAuditKeywordSearchQuery,
) -> Result<
Vec<$crate::repository::usage::StoredRequestUsageAudit>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::list_usage_audits_by_keyword_search(&repository, query).await
}
async fn count_usage_audits_by_keyword_search(
&self,
query: &$crate::repository::usage::UsageAuditKeywordSearchQuery,
) -> Result<u64, $crate::DataLayerError> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::count_usage_audits_by_keyword_search(&repository, query).await
}
async fn aggregate_usage_audits(
&self,
query: &$crate::repository::usage::UsageAuditAggregationQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageAuditAggregation>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::aggregate_usage_audits(&repository, query).await
}
async fn summarize_usage_audits(
&self,
query: &$crate::repository::usage::UsageAuditSummaryQuery,
) -> Result<
$crate::repository::usage::StoredUsageAuditSummary,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_audits(&repository, query).await
}
async fn summarize_usage_totals_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<$crate::repository::usage::StoredUsageUserTotals>, $crate::DataLayerError>
{
<$repository>::summarize_usage_totals_by_user_ids(self, user_ids).await
}
async fn summarize_usage_cache_hit_summary(
&self,
query: &$crate::repository::usage::UsageCacheHitSummaryQuery,
) -> Result<
$crate::repository::usage::StoredUsageCacheHitSummary,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_cache_hit_summary(&repository, query).await
}
async fn summarize_usage_settled_cost(
&self,
query: &$crate::repository::usage::UsageSettledCostSummaryQuery,
) -> Result<
$crate::repository::usage::StoredUsageSettledCostSummary,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_settled_cost(&repository, query).await
}
async fn summarize_usage_cache_affinity_hit_summary(
&self,
query: &$crate::repository::usage::UsageCacheAffinityHitSummaryQuery,
) -> Result<
$crate::repository::usage::StoredUsageCacheAffinityHitSummary,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_cache_affinity_hit_summary(&repository, query).await
}
async fn list_usage_cache_affinity_intervals(
&self,
query: &$crate::repository::usage::UsageCacheAffinityIntervalQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageCacheAffinityIntervalRow>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::list_usage_cache_affinity_intervals(&repository, query).await
}
async fn summarize_dashboard_usage(
&self,
query: &$crate::repository::usage::UsageDashboardSummaryQuery,
) -> Result<
$crate::repository::usage::StoredUsageDashboardSummary,
$crate::DataLayerError,
> {
<$repository>::summarize_dashboard_usage(self, query).await
}
async fn list_dashboard_daily_breakdown(
&self,
query: &$crate::repository::usage::UsageDashboardDailyBreakdownQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageDashboardDailyBreakdownRow>,
$crate::DataLayerError,
> {
<$repository>::list_dashboard_daily_breakdown(self, query).await
}
async fn summarize_dashboard_provider_counts(
&self,
query: &$crate::repository::usage::UsageDashboardProviderCountsQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageDashboardProviderCount>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_dashboard_provider_counts(&repository, query).await
}
async fn summarize_usage_breakdown(
&self,
query: &$crate::repository::usage::UsageBreakdownSummaryQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageBreakdownSummaryRow>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_breakdown(&repository, query).await
}
async fn count_monitoring_usage_errors(
&self,
query: &$crate::repository::usage::UsageMonitoringErrorCountQuery,
) -> Result<u64, $crate::DataLayerError> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::count_monitoring_usage_errors(&repository, query).await
}
async fn list_monitoring_usage_errors(
&self,
query: &$crate::repository::usage::UsageMonitoringErrorListQuery,
) -> Result<
Vec<$crate::repository::usage::StoredRequestUsageAudit>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::list_monitoring_usage_errors(&repository, query).await
}
async fn summarize_usage_error_distribution(
&self,
query: &$crate::repository::usage::UsageErrorDistributionQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageErrorDistributionRow>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_error_distribution(&repository, query).await
}
async fn summarize_usage_performance_percentiles(
&self,
query: &$crate::repository::usage::UsagePerformancePercentilesQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsagePerformancePercentilesRow>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_performance_percentiles(&repository, query).await
}
async fn summarize_usage_provider_performance(
&self,
query: &$crate::repository::usage::UsageProviderPerformanceQuery,
) -> Result<
$crate::repository::usage::StoredUsageProviderPerformance,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_provider_performance(&repository, query).await
}
async fn summarize_usage_cost_savings(
&self,
query: &$crate::repository::usage::UsageCostSavingsSummaryQuery,
) -> Result<
$crate::repository::usage::StoredUsageCostSavingsSummary,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_cost_savings(&repository, query).await
}
async fn summarize_usage_time_series(
&self,
query: &$crate::repository::usage::UsageTimeSeriesQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageTimeSeriesBucket>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_time_series(&repository, query).await
}
async fn summarize_usage_leaderboard(
&self,
query: &$crate::repository::usage::UsageLeaderboardQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageLeaderboardSummary>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_leaderboard(&repository, query).await
}
async fn list_recent_usage_audits(
&self,
user_id: Option<&str>,
limit: usize,
) -> Result<
Vec<$crate::repository::usage::StoredRequestUsageAudit>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::list_recent_usage_audits(&repository, user_id, limit).await
}
async fn summarize_total_tokens_by_api_key_ids(
&self,
api_key_ids: &[String],
) -> Result<std::collections::BTreeMap<String, u64>, $crate::DataLayerError> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_total_tokens_by_api_key_ids(&repository, api_key_ids).await
}
async fn summarize_usage_by_provider_api_key_ids(
&self,
provider_api_key_ids: &[String],
) -> Result<
std::collections::BTreeMap<
String,
$crate::repository::usage::StoredProviderApiKeyUsageSummary,
>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_by_provider_api_key_ids(&repository, provider_api_key_ids).await
}
async fn summarize_usage_by_provider_api_key_windows(
&self,
requests: &[$crate::repository::usage::ProviderApiKeyWindowUsageRequest],
) -> Result<
Vec<$crate::repository::usage::StoredProviderApiKeyWindowUsageSummary>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_by_provider_api_key_windows(&repository, requests).await
}
async fn summarize_provider_usage_since(
&self,
provider_id: &str,
since_unix_secs: u64,
) -> Result<
$crate::repository::usage::StoredProviderUsageSummary,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_provider_usage_since(&repository, provider_id, since_unix_secs).await
}
async fn summarize_usage_daily_heatmap(
&self,
query: &$crate::repository::usage::UsageDailyHeatmapQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageDailySummary>,
$crate::DataLayerError,
> {
<$repository>::summarize_usage_daily_heatmap(self, query).await
}
}
};
}
mod memory;
#[cfg(feature = "mysql")]
mod mysql;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::usage::{
api_key_usage_contribution, incoming_usage_can_recover_terminal_failure,
model_usage_contribution, provider_api_key_usage_contribution, provider_api_key_usage_is_error,
provider_api_key_usage_is_success, strip_deprecated_usage_display_fields,
usage_can_recover_terminal_failure, usage_request_metadata_client_family, ApiKeyLastUsedDelta,
ApiKeyUsageContribution, ApiKeyUsageDelta, ManagementTokenCounterDelta, ModelUsageContribution,
ModelUsageDelta, PendingUsageCleanupSummary, ProviderApiKeyUsageContribution,
ProviderApiKeyUsageDelta, ProviderApiKeyWindowUsageRequest, ProxyNodeCounterDelta,
StoredProviderApiKeyUsageSummary, StoredProviderApiKeyWindowUsageSummary,
StoredProviderUsageSummary, StoredProviderUsageWindow, StoredRequestUsageAudit,
StoredUsageAuditAggregation, StoredUsageAuditSummary, StoredUsageBreakdownSummaryRow,
StoredUsageCacheAffinityHitSummary, StoredUsageCacheAffinityIntervalRow,
StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary, StoredUsageDailySummary,
StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardProviderCount,
StoredUsageDashboardStatsSummary, StoredUsageDashboardSummary, StoredUsageErrorDistributionRow,
StoredUsageLeaderboardSummary, StoredUsagePerformancePercentilesRow,
StoredUsageProviderPerformance, StoredUsageProviderPerformanceProviderRow,
StoredUsageProviderPerformanceSummary, StoredUsageProviderPerformanceTimelineRow,
StoredUsageSettledCostSummary, StoredUsageTimeSeriesBucket, StoredUsageUserTotals,
UpsertUsageRecord, UsageAuditAggregationGroupBy, UsageAuditAggregationQuery,
UsageAuditKeywordSearchQuery, UsageAuditListQuery, UsageAuditSummaryQuery,
UsageBreakdownGroupBy, UsageBreakdownSummaryQuery, UsageCacheAffinityHitSummaryQuery,
UsageCacheAffinityIntervalGroupBy, UsageCacheAffinityIntervalQuery, UsageCacheHitSummaryQuery,
UsageCleanupPreviewCounts, UsageCleanupSummary, UsageCleanupWindow,
UsageCostSavingsSummaryQuery, UsageCounterFlushSummary, UsageCounterHealthSnapshot,
UsageCounterPendingHealthSnapshot, UsageDailyHeatmapQuery, UsageDashboardDailyBreakdownQuery,
UsageDashboardProviderCountsQuery, UsageDashboardSummaryQuery, UsageErrorDistributionQuery,
UsageLeaderboardGroupBy, UsageLeaderboardQuery, UsageMonitoringErrorCountQuery,
UsageMonitoringErrorListQuery, UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery,
UsageReadRepository, UsageRepository, UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity,
UsageTimeSeriesQuery, UsageWriteRepository,
};
#[cfg(feature = "postgres")]
pub mod cleanup {
pub use aether_data_postgres::cleanup::*;
}
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxUsageReadRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::{SqliteUsageReadRepository, SqliteUsageWriteRepository};
pub use memory::InMemoryUsageReadRepository;
#[cfg(feature = "mysql")]
pub use mysql::{MysqlUsageReadRepository, MysqlUsageWriteRepository};
#[cfg(test)]
mod tests {
use super::{
api_key_usage_contribution, incoming_usage_can_recover_terminal_failure,
model_usage_contribution, provider_api_key_usage_contribution,
provider_api_key_usage_is_error, provider_api_key_usage_is_success,
strip_deprecated_usage_display_fields, usage_can_recover_terminal_failure,
ApiKeyUsageDelta, ModelUsageDelta, ProviderApiKeyUsageDelta, StoredRequestUsageAudit,
UpsertUsageRecord,
};
#[test]
fn strip_deprecated_usage_display_fields_clears_legacy_display_columns() {
let usage = strip_deprecated_usage_display_fields(UpsertUsageRecord {
request_id: "req-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("key-1".to_string()),
username: Some("alice".to_string()),
api_key_name: Some("default".to_string()),
provider_name: "OpenAI".to_string(),
model: "gpt-5".to_string(),
target_model: None,
provider_id: None,
provider_endpoint_id: None,
provider_api_key_id: None,
request_type: Some("chat".to_string()),
api_format: Some("openai:chat".to_string()),
api_family: Some("openai".to_string()),
endpoint_kind: Some("chat".to_string()),
endpoint_api_format: Some("openai:chat".to_string()),
provider_api_family: Some("openai".to_string()),
provider_endpoint_kind: Some("chat".to_string()),
has_format_conversion: Some(false),
is_stream: Some(false),
input_tokens: Some(10),
output_tokens: Some(20),
total_tokens: Some(30),
cache_creation_input_tokens: None,
cache_creation_ephemeral_5m_input_tokens: None,
cache_creation_ephemeral_1h_input_tokens: None,
cache_read_input_tokens: None,
cache_creation_cost_usd: None,
cache_read_cost_usd: None,
output_price_per_1m: None,
total_cost_usd: Some(0.25),
actual_total_cost_usd: Some(0.15),
status_code: Some(200),
error_message: None,
error_category: None,
response_time_ms: Some(120),
first_byte_time_ms: Some(40),
status: "completed".to_string(),
billing_status: "pending".to_string(),
request_headers: None,
request_body: None,
request_body_ref: None,
provider_request_headers: None,
provider_request_body: None,
provider_request_body_ref: None,
response_headers: None,
response_body: None,
response_body_ref: None,
client_response_headers: None,
client_response_body: None,
client_response_body_ref: None,
request_body_state: None,
provider_request_body_state: None,
response_body_state: None,
client_response_body_state: None,
candidate_id: None,
candidate_index: None,
key_name: None,
planner_kind: None,
route_family: None,
route_kind: None,
execution_path: None,
local_execution_runtime_miss_reason: None,
request_metadata: None,
finalized_at_unix_secs: None,
created_at_unix_ms: Some(100),
updated_at_unix_secs: 101,
});
assert_eq!(usage.user_id.as_deref(), Some("user-1"));
assert_eq!(usage.api_key_id.as_deref(), Some("key-1"));
assert_eq!(usage.username, None);
assert_eq!(usage.api_key_name, None);
assert_eq!(usage.provider_name, "OpenAI");
assert_eq!(usage.model, "gpt-5");
}
#[test]
fn incoming_usage_recovery_requires_completed_state() {
assert!(incoming_usage_can_recover_terminal_failure(
"completed",
"pending",
));
assert!(!incoming_usage_can_recover_terminal_failure(
"streaming",
"pending",
));
assert!(!incoming_usage_can_recover_terminal_failure(
"pending", "pending",
));
assert!(!incoming_usage_can_recover_terminal_failure(
"failed", "void",
));
assert!(!incoming_usage_can_recover_terminal_failure(
"completed",
"settled",
));
}
#[test]
fn usage_recovery_requires_void_failure_and_completed_state() {
assert!(usage_can_recover_terminal_failure(
"failed",
"void",
"completed",
"pending",
));
assert!(usage_can_recover_terminal_failure(
"cancelled",
"void",
"completed",
"pending",
));
assert!(!usage_can_recover_terminal_failure(
"failed",
"void",
"streaming",
"pending",
));
assert!(!usage_can_recover_terminal_failure(
"failed", "void", "pending", "pending",
));
assert!(!usage_can_recover_terminal_failure(
"completed",
"pending",
"completed",
"pending",
));
assert!(!usage_can_recover_terminal_failure(
"failed", "void", "failed", "void",
));
}
#[test]
fn provider_key_usage_success_requires_clean_terminal_success() {
assert!(provider_api_key_usage_is_success(
"completed",
Some(200),
None
));
assert!(!provider_api_key_usage_is_success(
"completed",
Some(500),
None
));
assert!(!provider_api_key_usage_is_success(
"completed",
Some(200),
Some("boom")
));
assert!(!provider_api_key_usage_is_success(
"streaming",
Some(200),
None
));
}
#[test]
fn provider_key_usage_error_ignores_pending_states() {
assert!(provider_api_key_usage_is_error(
"failed",
Some(500),
Some("boom")
));
assert!(provider_api_key_usage_is_error(
"completed",
Some(200),
Some("boom")
));
assert!(!provider_api_key_usage_is_error("pending", None, None));
assert!(!provider_api_key_usage_is_error("streaming", None, None));
}
#[test]
fn provider_key_usage_contribution_tracks_success_response_time() {
let usage = StoredRequestUsageAudit::new(
"usage-1".to_string(),
"request-1".to_string(),
None,
None,
None,
None,
"OpenAI".to_string(),
"gpt-5".to_string(),
None,
Some("provider-1".to_string()),
None,
Some("provider-key-1".to_string()),
None,
None,
None,
None,
None,
None,
None,
false,
false,
12,
8,
20,
0.25,
0.25,
Some(200),
None,
None,
Some(120),
None,
"completed".to_string(),
"settled".to_string(),
123,
124,
Some(125),
)
.expect("usage should build");
let contribution =
provider_api_key_usage_contribution(&usage).expect("contribution should exist");
assert_eq!(contribution.key_id, "provider-key-1");
assert_eq!(contribution.request_count, 1);
assert_eq!(contribution.success_count, 1);
assert_eq!(contribution.error_count, 0);
assert_eq!(contribution.total_tokens, 20);
assert_eq!(contribution.total_cost_usd, 0.25);
assert_eq!(contribution.total_response_time_ms, 120);
assert_eq!(contribution.last_used_at_unix_secs, Some(123));
}
#[test]
fn api_key_usage_contribution_tracks_request_totals() {
let usage = StoredRequestUsageAudit::new(
"usage-1".to_string(),
"request-1".to_string(),
Some("user-1".to_string()),
Some("api-key-1".to_string()),
None,
None,
"OpenAI".to_string(),
"gpt-5".to_string(),
None,
Some("provider-1".to_string()),
None,
None,
None,
None,
None,
None,
None,
None,
None,
false,
false,
12,
8,
20,
0.25,
0.25,
Some(200),
None,
None,
Some(120),
None,
"completed".to_string(),
"settled".to_string(),
123,
124,
Some(125),
)
.expect("usage should build");
let contribution = api_key_usage_contribution(&usage).expect("contribution should exist");
assert_eq!(contribution.api_key_id, "api-key-1");
assert_eq!(contribution.total_requests, 1);
assert_eq!(contribution.total_tokens, 20);
assert_eq!(contribution.total_cost_usd, 0.25);
assert_eq!(contribution.last_used_at_unix_secs, Some(123));
let mut streaming = usage.clone();
streaming.status = "streaming".to_string();
assert!(api_key_usage_contribution(&streaming).is_none());
let mut pending = usage;
pending.status = "pending".to_string();
assert!(api_key_usage_contribution(&pending).is_none());
}
#[test]
fn provider_api_key_usage_contribution_counts_in_flight_requests_once() {
let usage = StoredRequestUsageAudit::new(
"usage-1".to_string(),
"request-1".to_string(),
Some("user-1".to_string()),
Some("api-key-1".to_string()),
None,
None,
"OpenAI".to_string(),
"gpt-5".to_string(),
None,
Some("provider-1".to_string()),
None,
Some("provider-key-1".to_string()),
None,
None,
None,
None,
None,
None,
None,
false,
false,
12,
8,
20,
0.25,
0.25,
Some(200),
None,
None,
Some(120),
None,
"completed".to_string(),
"settled".to_string(),
123,
124,
Some(125),
)
.expect("usage should build");
assert!(provider_api_key_usage_contribution(&usage).is_some());
let mut streaming = usage.clone();
streaming.status = "streaming".to_string();
let streaming_contribution =
provider_api_key_usage_contribution(&streaming).expect("streaming should count");
assert_eq!(streaming_contribution.request_count, 1);
assert_eq!(streaming_contribution.success_count, 0);
assert_eq!(streaming_contribution.error_count, 0);
assert_eq!(streaming_contribution.total_tokens, 0);
assert_eq!(streaming_contribution.total_cost_usd, 0.0);
assert_eq!(streaming_contribution.total_response_time_ms, 0);
let mut pending = usage.clone();
pending.status = "pending".to_string();
let pending_contribution =
provider_api_key_usage_contribution(&pending).expect("pending should count");
assert_eq!(pending_contribution.request_count, 1);
assert_eq!(pending_contribution.success_count, 0);
assert_eq!(pending_contribution.error_count, 0);
assert_eq!(pending_contribution.total_tokens, 0);
assert_eq!(pending_contribution.total_cost_usd, 0.0);
assert_eq!(pending_contribution.total_response_time_ms, 0);
let terminal_contribution =
provider_api_key_usage_contribution(&usage).expect("terminal should count");
let delta =
ProviderApiKeyUsageDelta::between(&pending_contribution, &terminal_contribution);
assert_eq!(delta.request_count, 0);
assert_eq!(delta.success_count, 1);
assert_eq!(delta.error_count, 0);
assert_eq!(delta.total_tokens, 20);
assert_eq!(delta.total_cost_usd, 0.25);
assert_eq!(delta.total_response_time_ms, 120);
}
#[test]
fn usage_delta_between_does_not_emit_duplicate_last_used_candidate() {
let api_key_contribution = super::ApiKeyUsageContribution {
api_key_id: "api-key-1".to_string(),
total_requests: 1,
total_tokens: 20,
total_cost_usd: 0.25,
last_used_at_unix_secs: Some(123),
};
assert!(ApiKeyUsageDelta::between(&api_key_contribution, &api_key_contribution).is_noop());
let provider_contribution = super::ProviderApiKeyUsageContribution {
key_id: "provider-key-1".to_string(),
request_count: 1,
success_count: 1,
error_count: 0,
total_tokens: 20,
total_cost_usd: 0.25,
total_response_time_ms: 120,
last_used_at_unix_secs: Some(123),
usage_created_at_unix_secs: Some(123),
};
assert!(
ProviderApiKeyUsageDelta::between(&provider_contribution, &provider_contribution,)
.is_noop()
);
}
#[test]
fn model_usage_contribution_tracks_terminal_requests_only() {
let completed = StoredRequestUsageAudit::new(
"usage-1".to_string(),
"request-1".to_string(),
Some("user-1".to_string()),
Some("api-key-1".to_string()),
None,
None,
"OpenAI".to_string(),
" gpt-5.5 ".to_string(),
None,
Some("provider-1".to_string()),
None,
None,
None,
None,
None,
None,
None,
None,
None,
false,
false,
12,
8,
20,
0.25,
0.25,
Some(200),
None,
None,
Some(120),
None,
"completed".to_string(),
"settled".to_string(),
123,
124,
Some(125),
)
.expect("usage should build");
let contribution =
model_usage_contribution(&completed).expect("completed usage should count");
assert_eq!(contribution.model, "gpt-5.5");
assert_eq!(contribution.request_count, 1);
let mut streaming = completed.clone();
streaming.status = "streaming".to_string();
assert!(model_usage_contribution(&streaming).is_none());
let mut pending = completed;
pending.status = "pending".to_string();
assert!(model_usage_contribution(&pending).is_none());
}
#[test]
fn model_usage_delta_handles_model_changes() {
let before = super::ModelUsageContribution {
model: "gpt-5.4".to_string(),
request_count: 1,
};
let after = super::ModelUsageContribution {
model: "gpt-5.5".to_string(),
request_count: 1,
};
assert_eq!(ModelUsageDelta::removal(&before).request_count, -1);
assert_eq!(ModelUsageDelta::addition(&after).request_count, 1);
assert!(ModelUsageDelta::between(&before, &before).is_noop());
}
}
@@ -0,0 +1,78 @@
use aether_data_contracts::repository::usage::{
StoredUsageDailySummary, StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardSummary,
StoredUsageUserTotals, UsageDailyHeatmapQuery, UsageDashboardDailyBreakdownQuery,
UsageDashboardSummaryQuery, UsageReadRepository,
};
use super::InMemoryUsageReadRepository;
use crate::driver::mysql::MysqlPool;
use crate::DataLayerError;
pub use aether_data_mysql::MysqlUsageWriteRepository;
#[derive(Debug, Clone)]
pub struct MysqlUsageReadRepository {
storage: aether_data_mysql::MysqlUsageStorage,
}
impl MysqlUsageReadRepository {
pub fn new(pool: MysqlPool) -> Self {
Self {
storage: aether_data_mysql::MysqlUsageStorage::new(pool),
}
}
async fn materialize_read_model(&self) -> Result<InMemoryUsageReadRepository, DataLayerError> {
Ok(InMemoryUsageReadRepository::seed(
self.storage.load_usage_records().await?,
))
}
async fn summarize_usage_daily_heatmap(
&self,
query: &UsageDailyHeatmapQuery,
) -> Result<Vec<StoredUsageDailySummary>, DataLayerError> {
self.storage.summarize_usage_daily_heatmap(query).await
}
async fn summarize_usage_totals_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUsageUserTotals>, DataLayerError> {
self.storage
.summarize_usage_totals_by_user_ids(user_ids)
.await
}
async fn summarize_dashboard_usage(
&self,
query: &UsageDashboardSummaryQuery,
) -> Result<StoredUsageDashboardSummary, DataLayerError> {
if let Some(summary) = self
.storage
.summarize_dashboard_usage_from_daily_aggregates(query)
.await?
{
return Ok(summary);
}
let repository = self.materialize_read_model().await?;
repository.summarize_dashboard_usage(query).await
}
async fn list_dashboard_daily_breakdown(
&self,
query: &UsageDashboardDailyBreakdownQuery,
) -> Result<Vec<StoredUsageDashboardDailyBreakdownRow>, DataLayerError> {
let rows = self
.storage
.list_dashboard_daily_breakdown_from_daily_aggregates(query)
.await?;
if !rows.is_empty() {
return Ok(rows);
}
let repository = self.materialize_read_model().await?;
repository.list_dashboard_daily_breakdown(query).await
}
}
impl_materialized_usage_read_repository!(MysqlUsageReadRepository);
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,16 @@
mod memory;
pub use aether_data_contracts::repository::users::{
normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord,
StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership,
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSortBy,
UserExportSortOrder, UserExportSummary, UserReadRepository,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlUserReadRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxUserReadRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteUserReadRepository;
pub use memory::InMemoryUserReadRepository;
@@ -0,0 +1,685 @@
use std::collections::BTreeMap;
use std::sync::RwLock;
use async_trait::async_trait;
use super::{
StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount,
VideoTaskQueryFilter, VideoTaskReadRepository, VideoTaskStatus, VideoTaskStatusCount,
VideoTaskWriteRepository,
};
use crate::DataLayerError;
#[derive(Debug, Default)]
struct MemoryVideoTaskIndex {
by_id: BTreeMap<String, StoredVideoTask>,
short_to_id: BTreeMap<String, String>,
user_external_to_id: BTreeMap<(String, String), String>,
}
#[derive(Debug, Default)]
pub struct InMemoryVideoTaskRepository {
index: RwLock<MemoryVideoTaskIndex>,
}
impl InMemoryVideoTaskRepository {
fn store_locked(index: &mut MemoryVideoTaskIndex, task: StoredVideoTask) -> StoredVideoTask {
if let Some(previous) = index.by_id.insert(task.id.clone(), task.clone()) {
if let Some(short_id) = previous.short_id {
index.short_to_id.remove(&short_id);
}
if let (Some(user_id), Some(external_task_id)) =
(previous.user_id, previous.external_task_id)
{
index
.user_external_to_id
.remove(&(user_id, external_task_id));
}
}
if let Some(short_id) = &task.short_id {
index.short_to_id.insert(short_id.clone(), task.id.clone());
}
if let (Some(user_id), Some(external_task_id)) = (&task.user_id, &task.external_task_id) {
index
.user_external_to_id
.insert((user_id.clone(), external_task_id.clone()), task.id.clone());
}
task
}
fn matches_filter(task: &StoredVideoTask, filter: &VideoTaskQueryFilter) -> bool {
if let Some(user_id) = filter.user_id.as_deref() {
if task.user_id.as_deref() != Some(user_id) {
return false;
}
}
if let Some(status) = filter.status {
if task.status != status {
return false;
}
}
if let Some(model_substring) = filter.model_substring.as_deref() {
let needle = model_substring.trim().to_ascii_lowercase();
let Some(model) = task.model.as_deref() else {
return false;
};
if !model.to_ascii_lowercase().contains(&needle) {
return false;
}
}
if let Some(client_api_format) = filter.client_api_format.as_deref() {
if task.client_api_format.as_deref() != Some(client_api_format) {
return false;
}
}
true
}
}
#[async_trait]
impl VideoTaskReadRepository for InMemoryVideoTaskRepository {
async fn find(
&self,
key: VideoTaskLookupKey<'_>,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
let index = self.index.read().expect("video task repository lock");
Ok(match key {
VideoTaskLookupKey::Id(id) => index.by_id.get(id).cloned(),
VideoTaskLookupKey::ShortId(short_id) => index
.short_to_id
.get(short_id)
.and_then(|id| index.by_id.get(id))
.cloned(),
VideoTaskLookupKey::UserExternal {
user_id,
external_task_id,
} => index
.user_external_to_id
.get(&(user_id.to_string(), external_task_id.to_string()))
.and_then(|id| index.by_id.get(id))
.cloned(),
})
}
async fn list_active(&self, limit: usize) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let mut tasks = self
.index
.read()
.expect("video task repository lock")
.by_id
.values()
.filter(|task| task.status.is_active())
.cloned()
.collect::<Vec<_>>();
tasks.sort_by_key(|entry| std::cmp::Reverse(entry.updated_at_unix_secs));
tasks.truncate(limit);
Ok(tasks)
}
async fn list_due(
&self,
now_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let mut tasks = self
.index
.read()
.expect("video task repository lock")
.by_id
.values()
.filter(|task| {
matches!(
task.status,
super::VideoTaskStatus::Submitted
| super::VideoTaskStatus::Queued
| super::VideoTaskStatus::Processing
) && task.poll_count < task.max_poll_count
&& task
.next_poll_at_unix_secs
.is_some_and(|value| value <= now_unix_secs)
})
.cloned()
.collect::<Vec<_>>();
tasks.sort_by(|left, right| {
left.next_poll_at_unix_secs
.cmp(&right.next_poll_at_unix_secs)
.then_with(|| left.updated_at_unix_secs.cmp(&right.updated_at_unix_secs))
});
tasks.truncate(limit);
Ok(tasks)
}
async fn list_page(
&self,
filter: &VideoTaskQueryFilter,
offset: usize,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let mut tasks = self
.index
.read()
.expect("video task repository lock")
.by_id
.values()
.filter(|task| Self::matches_filter(task, filter))
.cloned()
.collect::<Vec<_>>();
tasks.sort_by(|left, right| {
right
.created_at_unix_ms
.cmp(&left.created_at_unix_ms)
.then_with(|| right.updated_at_unix_secs.cmp(&left.updated_at_unix_secs))
});
Ok(tasks.into_iter().skip(offset).take(limit).collect())
}
async fn list_page_summary(
&self,
filter: &VideoTaskQueryFilter,
offset: usize,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
Self::list_page(self, filter, offset, limit).await
}
async fn count(&self, filter: &VideoTaskQueryFilter) -> Result<u64, DataLayerError> {
Ok(self
.index
.read()
.expect("video task repository lock")
.by_id
.values()
.filter(|task| Self::matches_filter(task, filter))
.count() as u64)
}
async fn count_by_status(
&self,
filter: &VideoTaskQueryFilter,
) -> Result<Vec<VideoTaskStatusCount>, DataLayerError> {
let mut counts = BTreeMap::<VideoTaskStatus, u64>::new();
for task in self
.index
.read()
.expect("video task repository lock")
.by_id
.values()
.filter(|task| Self::matches_filter(task, filter))
{
*counts.entry(task.status).or_default() += 1;
}
Ok(counts
.into_iter()
.map(|(status, count)| VideoTaskStatusCount { status, count })
.collect())
}
async fn count_distinct_users(
&self,
filter: &VideoTaskQueryFilter,
) -> Result<u64, DataLayerError> {
let index = self.index.read().expect("video task repository lock");
let users = index
.by_id
.values()
.filter(|task| Self::matches_filter(task, filter))
.filter_map(|task| task.user_id.as_deref())
.map(str::trim)
.filter(|value| !value.is_empty())
.collect::<std::collections::BTreeSet<_>>();
Ok(users.len() as u64)
}
async fn top_models(
&self,
filter: &VideoTaskQueryFilter,
limit: usize,
) -> Result<Vec<VideoTaskModelCount>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let mut counts = BTreeMap::<String, u64>::new();
for task in self
.index
.read()
.expect("video task repository lock")
.by_id
.values()
.filter(|task| Self::matches_filter(task, filter))
{
let Some(model) = task.model.as_deref() else {
continue;
};
if model.trim().is_empty() {
continue;
}
*counts.entry(model.to_string()).or_default() += 1;
}
let mut models = counts
.into_iter()
.map(|(model, count)| VideoTaskModelCount { model, count })
.collect::<Vec<_>>();
models.sort_by(|left, right| {
right
.count
.cmp(&left.count)
.then_with(|| left.model.cmp(&right.model))
});
models.truncate(limit);
Ok(models)
}
async fn count_created_since(
&self,
filter: &VideoTaskQueryFilter,
created_since_unix_secs: u64,
) -> Result<u64, DataLayerError> {
Ok(self
.index
.read()
.expect("video task repository lock")
.by_id
.values()
.filter(|task| {
Self::matches_filter(task, filter)
&& task.created_at_unix_ms >= created_since_unix_secs
})
.count() as u64)
}
}
#[async_trait]
impl VideoTaskWriteRepository for InMemoryVideoTaskRepository {
async fn upsert(&self, task: UpsertVideoTask) -> Result<StoredVideoTask, DataLayerError> {
let mut index = self.index.write().expect("video task repository lock");
Ok(Self::store_locked(&mut index, task.into_stored()))
}
async fn update_if_active(
&self,
task: UpsertVideoTask,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
let mut index = self.index.write().expect("video task repository lock");
let Some(existing) = index.by_id.get(&task.id) else {
return Ok(None);
};
if !existing.status.is_active() {
return Ok(None);
}
Ok(Some(Self::store_locked(&mut index, task.into_stored())))
}
async fn claim_due(
&self,
now_unix_secs: u64,
claim_until_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let mut index = self.index.write().expect("video task repository lock");
let mut due_ids = index
.by_id
.values()
.filter(|task| {
matches!(
task.status,
VideoTaskStatus::Submitted
| VideoTaskStatus::Queued
| VideoTaskStatus::Processing
) && task.poll_count < task.max_poll_count
&& task
.next_poll_at_unix_secs
.is_some_and(|value| value <= now_unix_secs)
})
.map(|task| task.id.clone())
.collect::<Vec<_>>();
due_ids.sort_by(|left_id, right_id| {
let left = index.by_id.get(left_id).expect("task should exist");
let right = index.by_id.get(right_id).expect("task should exist");
left.next_poll_at_unix_secs
.cmp(&right.next_poll_at_unix_secs)
.then_with(|| left.updated_at_unix_secs.cmp(&right.updated_at_unix_secs))
});
due_ids.truncate(limit);
let mut claimed = Vec::with_capacity(due_ids.len());
for id in due_ids {
let Some(task) = index.by_id.get_mut(&id) else {
continue;
};
task.next_poll_at_unix_secs = Some(claim_until_unix_secs);
task.updated_at_unix_secs = now_unix_secs.max(task.updated_at_unix_secs);
claimed.push(task.clone());
}
Ok(claimed)
}
}
#[cfg(test)]
mod tests {
use super::InMemoryVideoTaskRepository;
use crate::repository::video_tasks::{
UpsertVideoTask, VideoTaskLookupKey, VideoTaskQueryFilter, VideoTaskReadRepository,
VideoTaskStatus, VideoTaskWriteRepository,
};
fn sample_task(
id: &str,
status: VideoTaskStatus,
updated_at_unix_secs: u64,
) -> UpsertVideoTask {
UpsertVideoTask {
id: id.to_string(),
short_id: Some(format!("short-{id}")),
request_id: format!("request-{id}"),
user_id: Some("user-1".to_string()),
api_key_id: Some("api-key-1".to_string()),
username: Some("user".to_string()),
api_key_name: Some("primary".to_string()),
external_task_id: Some(format!("ext-{id}")),
provider_id: Some("provider-1".to_string()),
endpoint_id: Some("endpoint-1".to_string()),
key_id: Some("provider-key-1".to_string()),
client_api_format: Some("openai:video".to_string()),
provider_api_format: Some("openai:video".to_string()),
format_converted: false,
model: Some("sora-2".to_string()),
prompt: Some("hello".to_string()),
original_request_body: Some(serde_json::json!({"prompt": "hello"})),
duration_seconds: Some(4),
resolution: Some("720p".to_string()),
aspect_ratio: Some("16:9".to_string()),
size: Some("1280x720".to_string()),
status,
progress_percent: 0,
progress_message: None,
retry_count: 0,
poll_interval_seconds: 10,
next_poll_at_unix_secs: Some(updated_at_unix_secs),
poll_count: 0,
max_poll_count: 360,
created_at_unix_ms: updated_at_unix_secs.saturating_sub(10),
submitted_at_unix_secs: Some(updated_at_unix_secs.saturating_sub(10)),
completed_at_unix_secs: None,
updated_at_unix_secs,
error_code: None,
error_message: None,
video_url: None,
request_metadata: None,
}
}
#[tokio::test]
async fn reads_task_by_all_supported_lookup_keys() {
let repo = InMemoryVideoTaskRepository::default();
repo.upsert(sample_task("task-1", VideoTaskStatus::Submitted, 100))
.await
.expect("upsert should succeed");
assert!(repo
.find(VideoTaskLookupKey::Id("task-1"))
.await
.expect("find by id should succeed")
.is_some());
assert!(repo
.find(VideoTaskLookupKey::ShortId("short-task-1"))
.await
.expect("find by short id should succeed")
.is_some());
assert!(repo
.find(VideoTaskLookupKey::UserExternal {
user_id: "user-1",
external_task_id: "ext-task-1",
})
.await
.expect("find by user/external should succeed")
.is_some());
}
#[tokio::test]
async fn list_active_only_returns_active_tasks_in_descending_update_order() {
let repo = InMemoryVideoTaskRepository::default();
repo.upsert(sample_task("task-1", VideoTaskStatus::Completed, 100))
.await
.expect("upsert should succeed");
repo.upsert(sample_task("task-2", VideoTaskStatus::Processing, 200))
.await
.expect("upsert should succeed");
repo.upsert(sample_task("task-3", VideoTaskStatus::Queued, 150))
.await
.expect("upsert should succeed");
let active = repo
.list_active(10)
.await
.expect("list active should succeed");
assert_eq!(active.len(), 2);
assert_eq!(active[0].id, "task-2");
assert_eq!(active[1].id, "task-3");
}
#[tokio::test]
async fn upsert_replaces_secondary_indexes() {
let repo = InMemoryVideoTaskRepository::default();
repo.upsert(sample_task("task-1", VideoTaskStatus::Submitted, 100))
.await
.expect("upsert should succeed");
repo.upsert(UpsertVideoTask {
id: "task-1".to_string(),
short_id: Some("short-task-1b".to_string()),
request_id: "request-task-1b".to_string(),
user_id: Some("user-2".to_string()),
api_key_id: Some("api-key-2".to_string()),
username: Some("user-2".to_string()),
api_key_name: Some("secondary".to_string()),
external_task_id: Some("ext-task-1b".to_string()),
provider_id: Some("provider-2".to_string()),
endpoint_id: Some("endpoint-2".to_string()),
key_id: Some("provider-key-2".to_string()),
client_api_format: Some("gemini:video".to_string()),
provider_api_format: Some("gemini:video".to_string()),
format_converted: false,
model: Some("veo-3".to_string()),
prompt: Some("remix".to_string()),
original_request_body: Some(serde_json::json!({"prompt": "remix"})),
duration_seconds: Some(8),
resolution: Some("1080p".to_string()),
aspect_ratio: Some("16:9".to_string()),
size: Some("720p".to_string()),
status: VideoTaskStatus::Processing,
progress_percent: 50,
progress_message: Some("processing".to_string()),
retry_count: 1,
poll_interval_seconds: 10,
next_poll_at_unix_secs: Some(200),
poll_count: 2,
max_poll_count: 360,
created_at_unix_ms: 150,
submitted_at_unix_secs: Some(150),
completed_at_unix_secs: None,
updated_at_unix_secs: 200,
error_code: None,
error_message: None,
video_url: None,
request_metadata: None,
})
.await
.expect("upsert should succeed");
assert!(repo
.find(VideoTaskLookupKey::ShortId("short-task-1"))
.await
.expect("find should succeed")
.is_none());
assert!(repo
.find(VideoTaskLookupKey::UserExternal {
user_id: "user-1",
external_task_id: "ext-task-1",
})
.await
.expect("find should succeed")
.is_none());
assert!(repo
.find(VideoTaskLookupKey::ShortId("short-task-1b"))
.await
.expect("find should succeed")
.is_some());
}
#[tokio::test]
async fn list_due_returns_due_active_tasks_in_next_poll_order() {
let repo = InMemoryVideoTaskRepository::default();
repo.upsert(sample_task("task-1", VideoTaskStatus::Submitted, 300))
.await
.expect("upsert should succeed");
repo.upsert(sample_task("task-2", VideoTaskStatus::Processing, 100))
.await
.expect("upsert should succeed");
repo.upsert(UpsertVideoTask {
next_poll_at_unix_secs: Some(500),
..sample_task("task-3", VideoTaskStatus::Queued, 200)
})
.await
.expect("upsert should succeed");
let due = repo
.list_due(300, 10)
.await
.expect("list due should succeed");
assert_eq!(due.len(), 2);
assert_eq!(due[0].id, "task-2");
assert_eq!(due[1].id, "task-1");
}
#[tokio::test]
async fn update_if_active_skips_terminal_tasks() {
let repo = InMemoryVideoTaskRepository::default();
repo.upsert(sample_task("task-1", VideoTaskStatus::Completed, 100))
.await
.expect("upsert should succeed");
let updated = repo
.update_if_active(UpsertVideoTask {
progress_percent: 100,
..sample_task("task-1", VideoTaskStatus::Completed, 200)
})
.await
.expect("update should succeed");
assert!(updated.is_none());
}
#[tokio::test]
async fn list_page_and_stats_apply_filters() {
let repo = InMemoryVideoTaskRepository::default();
repo.upsert(sample_task("task-1", VideoTaskStatus::Submitted, 100))
.await
.expect("upsert should succeed");
repo.upsert(UpsertVideoTask {
model: Some("veo-3-fast".to_string()),
user_id: Some("user-2".to_string()),
client_api_format: Some("gemini:video".to_string()),
created_at_unix_ms: 250,
updated_at_unix_secs: 250,
..sample_task("task-2", VideoTaskStatus::Completed, 250)
})
.await
.expect("upsert should succeed");
repo.upsert(UpsertVideoTask {
model: Some("veo-3-fast".to_string()),
user_id: Some("user-2".to_string()),
client_api_format: Some("gemini:video".to_string()),
created_at_unix_ms: 260,
updated_at_unix_secs: 260,
..sample_task("task-3", VideoTaskStatus::Completed, 260)
})
.await
.expect("upsert should succeed");
let filter = VideoTaskQueryFilter {
user_id: Some("user-2".to_string()),
status: Some(VideoTaskStatus::Completed),
model_substring: Some("veo".to_string()),
client_api_format: Some("gemini:video".to_string()),
};
let page = repo
.list_page(&filter, 0, 10)
.await
.expect("list page should succeed");
assert_eq!(page.len(), 2);
assert_eq!(page[0].id, "task-3");
assert_eq!(page[1].id, "task-2");
let count = repo.count(&filter).await.expect("count should succeed");
assert_eq!(count, 2);
let by_status = repo
.count_by_status(&filter)
.await
.expect("status count should succeed");
assert_eq!(by_status.len(), 1);
assert_eq!(by_status[0].status, VideoTaskStatus::Completed);
assert_eq!(by_status[0].count, 2);
let top_models = repo
.top_models(&filter, 10)
.await
.expect("top models should succeed");
assert_eq!(top_models.len(), 1);
assert_eq!(top_models[0].model, "veo-3-fast");
assert_eq!(top_models[0].count, 2);
let today_count = repo
.count_created_since(&filter, 255)
.await
.expect("today count should succeed");
assert_eq!(today_count, 1);
}
#[tokio::test]
async fn claim_due_advances_claimed_tasks_until_claim_deadline() {
let repo = InMemoryVideoTaskRepository::default();
repo.upsert(sample_task("task-1", VideoTaskStatus::Submitted, 100))
.await
.expect("upsert should succeed");
repo.upsert(sample_task("task-2", VideoTaskStatus::Processing, 90))
.await
.expect("upsert should succeed");
let claimed = repo
.claim_due(100, 130, 1)
.await
.expect("claim should succeed");
assert_eq!(claimed.len(), 1);
assert_eq!(claimed[0].id, "task-2");
assert_eq!(claimed[0].next_poll_at_unix_secs, Some(130));
let remaining = repo
.list_due(100, 10)
.await
.expect("list due should succeed");
assert_eq!(remaining.len(), 1);
assert_eq!(remaining[0].id, "task-1");
}
}
@@ -0,0 +1,15 @@
mod memory;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::video_tasks::{
StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount,
VideoTaskQueryFilter, VideoTaskReadRepository, VideoTaskRepository, VideoTaskStatus,
VideoTaskStatusCount, VideoTaskWriteRepository,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlVideoTaskRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::{SqlxVideoTaskReadRepository, SqlxVideoTaskRepository};
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteVideoTaskRepository;
pub use memory::InMemoryVideoTaskRepository;
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,35 @@
mod memory;
pub use aether_data_contracts::repository::wallet::{
redeem_code_credits_recharge_balance, redeem_code_payment_method,
redeem_code_refundable_amount, AdjustWalletBalanceInput, AdminPaymentCallbackRecord,
AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery,
AdminWalletLedgerQuery, AdminWalletListQuery, AdminWalletPaymentOrderRecord,
AdminWalletRefundRecord, AdminWalletRefundRequestListQuery, AdminWalletTransactionRecord,
CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput,
CreateAdminRedeemCodeBatchResult, CreateManualWalletRechargeInput,
CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput,
CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput,
CreateWalletRefundRequestOutcome, CreatedAdminRedeemCodePlaintext,
CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput,
DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput,
ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome,
RedeemWalletCodeInput, RedeemWalletCodeOutcome, StoredAdminPaymentCallback,
StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage,
StoredAdminRedeemCode, StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage,
StoredAdminRedeemCodePage, StoredAdminWalletLedgerItem, StoredAdminWalletLedgerPage,
StoredAdminWalletListItem, StoredAdminWalletListPage, StoredAdminWalletRefund,
StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestItem,
StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction,
StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger,
StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, WalletLookupKey, WalletMutationOutcome,
WalletReadRepository, WalletReadSeed, WalletReadSnapshot, WalletRepository,
WalletWriteRepository,
};
#[cfg(feature = "mysql")]
pub use aether_data_mysql::MysqlWalletReadRepository;
#[cfg(feature = "postgres")]
pub use aether_data_postgres::SqlxWalletRepository;
#[cfg(feature = "sqlite")]
pub use aether_data_sqlite::SqliteWalletReadRepository;
pub use memory::InMemoryWalletRepository;