mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 02:17:46 +08:00
refactor(workspace): enforce layered crate boundaries
This commit is contained in:
@@ -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,692 @@
|
||||
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,
|
||||
}]);
|
||||
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,
|
||||
}]);
|
||||
|
||||
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()]),
|
||||
}]);
|
||||
|
||||
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("a.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;
|
||||
Reference in New Issue
Block a user