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

# Conflicts:
#	apps/aether-gateway/src/handlers/admin/request/provider/tasks.rs
#	frontend/src/api/endpoints/pool.ts
This commit is contained in:
MMEXA
2026-07-16 23:43:04 +08:00
1257 changed files with 80521 additions and 35495 deletions
@@ -0,0 +1,320 @@
use std::fmt;
use std::str::FromStr;
use crate::DataLayerError;
pub const DEFAULT_SQLITE_DATABASE_URL: &str = "sqlite://./data/aether.db";
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
pub struct PostgresPoolConfig {
pub database_url: String,
pub min_connections: u32,
pub max_connections: u32,
pub acquire_timeout_ms: u64,
pub idle_timeout_ms: u64,
pub max_lifetime_ms: u64,
pub statement_cache_capacity: usize,
pub require_ssl: bool,
}
impl Default for PostgresPoolConfig {
fn default() -> Self {
Self {
database_url: String::new(),
min_connections: 4,
max_connections: 20,
acquire_timeout_ms: 10_000,
idle_timeout_ms: 30_000,
max_lifetime_ms: 30 * 60_000,
statement_cache_capacity: 100,
require_ssl: false,
}
}
}
impl PostgresPoolConfig {
pub fn validate(&self) -> Result<(), DataLayerError> {
if self.database_url.trim().is_empty() {
return Err(DataLayerError::InvalidConfiguration(
"postgres database_url cannot be empty".to_string(),
));
}
if self.min_connections > self.max_connections {
return Err(DataLayerError::InvalidConfiguration(
"postgres min_connections cannot exceed max_connections".to_string(),
));
}
if self.statement_cache_capacity == 0 {
return Err(DataLayerError::InvalidConfiguration(
"postgres statement_cache_capacity must be positive".to_string(),
));
}
Ok(())
}
}
#[derive(
Debug, Copy, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq, Hash, PartialOrd, Ord,
)]
#[serde(rename_all = "snake_case")]
pub enum DatabaseDriver {
Sqlite,
Mysql,
Postgres,
}
impl DatabaseDriver {
pub const fn as_str(self) -> &'static str {
match self {
Self::Sqlite => "sqlite",
Self::Mysql => "mysql",
Self::Postgres => "postgres",
}
}
pub fn from_database_url(url: &str) -> Option<Self> {
let scheme = url.split_once(':')?.0.to_ascii_lowercase();
match scheme.as_str() {
"sqlite" => Some(Self::Sqlite),
"mysql" | "mariadb" => Some(Self::Mysql),
"postgres" | "postgresql" => Some(Self::Postgres),
_ => None,
}
}
}
impl fmt::Display for DatabaseDriver {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
impl FromStr for DatabaseDriver {
type Err = DataLayerError;
fn from_str(value: &str) -> Result<Self, Self::Err> {
match value.trim().to_ascii_lowercase().as_str() {
"sqlite" => Ok(Self::Sqlite),
"mysql" | "mariadb" => Ok(Self::Mysql),
"postgres" | "postgresql" => Ok(Self::Postgres),
other => Err(DataLayerError::InvalidConfiguration(format!(
"unsupported database driver '{other}'; expected sqlite, mysql, or postgres"
))),
}
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
pub struct SqlPoolConfig {
pub min_connections: u32,
pub max_connections: u32,
pub acquire_timeout_ms: u64,
pub idle_timeout_ms: u64,
pub max_lifetime_ms: u64,
pub statement_cache_capacity: usize,
pub require_ssl: bool,
}
impl Default for SqlPoolConfig {
fn default() -> Self {
Self {
min_connections: 1,
max_connections: 20,
acquire_timeout_ms: 10_000,
idle_timeout_ms: 30_000,
max_lifetime_ms: 30 * 60_000,
statement_cache_capacity: 100,
require_ssl: false,
}
}
}
impl SqlPoolConfig {
pub fn validate(&self, driver: DatabaseDriver) -> Result<(), DataLayerError> {
if self.min_connections > self.max_connections {
return Err(DataLayerError::InvalidConfiguration(format!(
"{driver} min_connections cannot exceed max_connections"
)));
}
if self.statement_cache_capacity == 0 {
return Err(DataLayerError::InvalidConfiguration(format!(
"{driver} statement_cache_capacity must be positive"
)));
}
if driver == DatabaseDriver::Sqlite && self.require_ssl {
return Err(DataLayerError::InvalidConfiguration(
"sqlite database does not support require_ssl".to_string(),
));
}
Ok(())
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
pub struct SqlDatabaseConfig {
pub driver: DatabaseDriver,
pub url: String,
pub pool: SqlPoolConfig,
}
impl SqlDatabaseConfig {
pub fn new(
driver: DatabaseDriver,
url: impl Into<String>,
pool: SqlPoolConfig,
) -> Result<Self, DataLayerError> {
let config = Self {
driver,
url: url.into(),
pool,
};
config.validate()?;
Ok(config)
}
pub fn sqlite_default() -> Self {
Self {
driver: DatabaseDriver::Sqlite,
url: DEFAULT_SQLITE_DATABASE_URL.to_string(),
pool: SqlPoolConfig::default(),
}
}
pub fn validate(&self) -> Result<(), DataLayerError> {
if self.url.trim().is_empty() {
return Err(DataLayerError::InvalidConfiguration(format!(
"{} database url cannot be empty",
self.driver
)));
}
if let Some(url_driver) = DatabaseDriver::from_database_url(&self.url) {
if url_driver != self.driver {
return Err(DataLayerError::InvalidConfiguration(format!(
"database driver '{}' does not match url scheme '{}'",
self.driver, url_driver
)));
}
}
self.pool.validate(self.driver)
}
pub fn from_postgres_config(postgres: PostgresPoolConfig) -> Self {
Self {
driver: DatabaseDriver::Postgres,
url: postgres.database_url,
pool: SqlPoolConfig {
min_connections: postgres.min_connections,
max_connections: postgres.max_connections,
acquire_timeout_ms: postgres.acquire_timeout_ms,
idle_timeout_ms: postgres.idle_timeout_ms,
max_lifetime_ms: postgres.max_lifetime_ms,
statement_cache_capacity: postgres.statement_cache_capacity,
require_ssl: postgres.require_ssl,
},
}
}
pub fn to_postgres_config(&self) -> Result<PostgresPoolConfig, DataLayerError> {
if self.driver != DatabaseDriver::Postgres {
return Err(DataLayerError::InvalidConfiguration(format!(
"cannot build postgres pool config from {} database config",
self.driver
)));
}
Ok(PostgresPoolConfig {
database_url: self.url.clone(),
min_connections: self.pool.min_connections,
max_connections: self.pool.max_connections,
acquire_timeout_ms: self.pool.acquire_timeout_ms,
idle_timeout_ms: self.pool.idle_timeout_ms,
max_lifetime_ms: self.pool.max_lifetime_ms,
statement_cache_capacity: self.pool.statement_cache_capacity,
require_ssl: self.pool.require_ssl,
})
}
}
impl From<PostgresPoolConfig> for SqlDatabaseConfig {
fn from(value: PostgresPoolConfig) -> Self {
Self::from_postgres_config(value)
}
}
#[cfg(test)]
mod tests {
use super::{
DatabaseDriver, PostgresPoolConfig, SqlDatabaseConfig, SqlPoolConfig,
DEFAULT_SQLITE_DATABASE_URL,
};
#[test]
fn parses_database_driver_aliases() {
assert_eq!(
"sqlite".parse::<DatabaseDriver>().unwrap(),
DatabaseDriver::Sqlite
);
assert_eq!(
"mariadb".parse::<DatabaseDriver>().unwrap(),
DatabaseDriver::Mysql
);
assert_eq!(
"postgresql".parse::<DatabaseDriver>().unwrap(),
DatabaseDriver::Postgres
);
assert!("oracle".parse::<DatabaseDriver>().is_err());
}
#[test]
fn infers_driver_from_database_url_scheme() {
assert_eq!(
DatabaseDriver::from_database_url("sqlite://./data/aether.db"),
Some(DatabaseDriver::Sqlite)
);
assert_eq!(
DatabaseDriver::from_database_url("mysql://localhost/aether"),
Some(DatabaseDriver::Mysql)
);
assert_eq!(
DatabaseDriver::from_database_url("postgres://localhost/aether"),
Some(DatabaseDriver::Postgres)
);
}
#[test]
fn validates_driver_url_mismatch() {
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Mysql,
url: "postgres://localhost/aether".to_string(),
pool: SqlPoolConfig::default(),
};
assert!(config.validate().is_err());
}
#[test]
fn builds_legacy_postgres_config_round_trip() {
let postgres = PostgresPoolConfig {
database_url: "postgres://localhost/aether".to_string(),
min_connections: 2,
max_connections: 8,
acquire_timeout_ms: 1_500,
idle_timeout_ms: 5_000,
max_lifetime_ms: 30_000,
statement_cache_capacity: 64,
require_ssl: true,
};
let database = SqlDatabaseConfig::from_postgres_config(postgres.clone());
assert_eq!(database.driver, DatabaseDriver::Postgres);
assert_eq!(database.to_postgres_config().unwrap(), postgres);
}
#[test]
fn sqlite_default_uses_local_database_path() {
let database = SqlDatabaseConfig::sqlite_default();
assert_eq!(database.driver, DatabaseDriver::Sqlite);
assert_eq!(database.url, DEFAULT_SQLITE_DATABASE_URL);
}
}
+37
View File
@@ -0,0 +1,37 @@
#[derive(Debug, thiserror::Error)]
pub enum DataLayerError {
#[error("invalid configuration: {0}")]
InvalidConfiguration(String),
#[error("invalid input: {0}")]
InvalidInput(String),
#[error("postgres error: {0}")]
Postgres(String),
#[error("redis error: {0}")]
Redis(String),
#[error("sql error: {0}")]
Sql(String),
#[error("operation timed out: {0}")]
TimedOut(String),
#[error("unexpected database value: {0}")]
UnexpectedValue(String),
}
impl DataLayerError {
pub fn postgres(error: impl std::fmt::Display) -> Self {
Self::Postgres(error.to_string())
}
pub fn redis(error: impl std::fmt::Display) -> Self {
Self::Redis(error.to_string())
}
pub fn sql(error: impl std::fmt::Display) -> Self {
Self::Sql(error.to_string())
}
}
+11
View File
@@ -0,0 +1,11 @@
pub mod database;
mod error;
pub mod migration;
pub mod repository;
pub use database::{
DatabaseDriver, PostgresPoolConfig, SqlDatabaseConfig, SqlPoolConfig,
DEFAULT_SQLITE_DATABASE_URL,
};
pub use error::DataLayerError;
pub use migration::PendingMigrationInfo;
@@ -0,0 +1,7 @@
//! Database lifecycle contracts shared by driver adapters and the data facade.
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PendingMigrationInfo {
pub version: i64,
pub description: String,
}
@@ -0,0 +1,246 @@
use async_trait::async_trait;
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredAnnouncement {
pub id: String,
pub title: String,
pub content: String,
pub kind: String,
pub priority: i32,
pub is_active: bool,
pub is_pinned: bool,
pub requires_ack: bool,
pub author_id: Option<String>,
pub author_username: Option<String>,
pub start_time_unix_secs: Option<u64>,
pub end_time_unix_secs: Option<u64>,
pub created_at_unix_ms: u64,
pub updated_at_unix_secs: u64,
}
impl StoredAnnouncement {
#[allow(clippy::too_many_arguments)]
pub fn new(
id: String,
title: String,
content: String,
kind: String,
priority: i32,
is_active: bool,
is_pinned: bool,
requires_ack: bool,
author_id: Option<String>,
author_username: Option<String>,
start_time_unix_secs: Option<i64>,
end_time_unix_secs: Option<i64>,
created_at_unix_ms: i64,
updated_at_unix_secs: i64,
) -> Result<Self, crate::DataLayerError> {
if id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"announcements.id is empty".to_string(),
));
}
if title.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"announcements.title is empty".to_string(),
));
}
if content.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"announcements.content is empty".to_string(),
));
}
if kind.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"announcements.type is empty".to_string(),
));
}
Ok(Self {
id,
title,
content,
kind,
priority,
is_active,
is_pinned,
requires_ack,
author_id,
author_username,
start_time_unix_secs: start_time_unix_secs
.map(|value| parse_timestamp(value, "announcements.start_time"))
.transpose()?,
end_time_unix_secs: end_time_unix_secs
.map(|value| parse_timestamp(value, "announcements.end_time"))
.transpose()?,
created_at_unix_ms: parse_timestamp(created_at_unix_ms, "announcements.created_at")?,
updated_at_unix_secs: parse_timestamp(
updated_at_unix_secs,
"announcements.updated_at",
)?,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
pub struct AnnouncementListQuery {
pub active_only: bool,
pub offset: usize,
pub limit: usize,
pub now_unix_secs: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
pub struct StoredAnnouncementPage {
pub items: Vec<StoredAnnouncement>,
pub total: u64,
}
fn parse_timestamp(value: i64, field: &str) -> Result<u64, crate::DataLayerError> {
u64::try_from(value).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!("{field} is negative: {value}"))
})
}
#[async_trait]
pub trait AnnouncementReadRepository: Send + Sync {
async fn find_by_id(
&self,
announcement_id: &str,
) -> Result<Option<StoredAnnouncement>, crate::DataLayerError>;
async fn list_announcements(
&self,
query: &AnnouncementListQuery,
) -> Result<StoredAnnouncementPage, crate::DataLayerError>;
async fn count_unread_active_announcements(
&self,
user_id: &str,
now_unix_secs: u64,
) -> Result<u64, crate::DataLayerError>;
async fn list_required_unread_active_announcements(
&self,
user_id: &str,
now_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredAnnouncement>, crate::DataLayerError>;
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct CreateAnnouncementRecord {
pub title: String,
pub content: String,
pub kind: String,
pub priority: i32,
pub is_pinned: bool,
pub requires_ack: bool,
pub author_id: String,
pub start_time_unix_secs: Option<u64>,
pub end_time_unix_secs: Option<u64>,
}
impl CreateAnnouncementRecord {
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
if self.title.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"announcement title cannot be empty".to_string(),
));
}
if self.content.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"announcement content cannot be empty".to_string(),
));
}
if self.kind.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"announcement type cannot be empty".to_string(),
));
}
if self.author_id.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"announcement author_id cannot be empty".to_string(),
));
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct UpdateAnnouncementRecord {
pub announcement_id: String,
pub title: Option<String>,
pub content: Option<String>,
pub kind: Option<String>,
pub priority: Option<i32>,
pub is_active: Option<bool>,
pub is_pinned: Option<bool>,
pub requires_ack: Option<bool>,
pub start_time_unix_secs: Option<u64>,
pub end_time_unix_secs: Option<u64>,
}
impl UpdateAnnouncementRecord {
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
if self.announcement_id.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"announcement_id cannot be empty".to_string(),
));
}
if self
.title
.as_deref()
.is_some_and(|value| value.trim().is_empty())
{
return Err(crate::DataLayerError::InvalidInput(
"announcement title cannot be empty".to_string(),
));
}
if self
.content
.as_deref()
.is_some_and(|value| value.trim().is_empty())
{
return Err(crate::DataLayerError::InvalidInput(
"announcement content cannot be empty".to_string(),
));
}
if self
.kind
.as_deref()
.is_some_and(|value| value.trim().is_empty())
{
return Err(crate::DataLayerError::InvalidInput(
"announcement type cannot be empty".to_string(),
));
}
Ok(())
}
}
#[async_trait]
pub trait AnnouncementWriteRepository: Send + Sync {
async fn create_announcement(
&self,
record: CreateAnnouncementRecord,
) -> Result<StoredAnnouncement, crate::DataLayerError>;
async fn update_announcement(
&self,
record: UpdateAnnouncementRecord,
) -> Result<Option<StoredAnnouncement>, crate::DataLayerError>;
async fn delete_announcement(
&self,
announcement_id: &str,
) -> Result<bool, crate::DataLayerError>;
async fn mark_announcement_as_read(
&self,
user_id: &str,
announcement_id: &str,
read_at_unix_secs: u64,
) -> Result<bool, crate::DataLayerError>;
}
@@ -0,0 +1,136 @@
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use serde_json::Value;
pub const SUSPICIOUS_EVENT_TYPES: &[&str] = &[
"suspicious_activity",
"unauthorized_access",
"login_failed",
"request_rate_limited",
];
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AuditLogListQuery {
pub cutoff_unix_secs: u64,
pub username_pattern: Option<String>,
pub event_type: Option<String>,
pub limit: usize,
pub offset: usize,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredAdminAuditLog {
pub id: String,
pub event_type: String,
pub user_id: Option<String>,
pub user_email: Option<String>,
pub user_username: Option<String>,
pub description: Option<String>,
pub ip_address: Option<String>,
pub status_code: Option<i32>,
pub error_message: Option<String>,
pub metadata: Option<Value>,
pub created_at_unix_secs: u64,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredSuspiciousActivity {
pub id: String,
pub event_type: String,
pub user_id: Option<String>,
pub description: Option<String>,
pub ip_address: Option<String>,
pub metadata: Option<Value>,
pub created_at_unix_secs: u64,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredUserAuditLog {
pub id: String,
pub event_type: String,
pub description: Option<String>,
pub ip_address: Option<String>,
pub status_code: Option<i32>,
pub created_at_unix_secs: u64,
}
fn unix_secs_to_rfc3339(secs: u64) -> Option<String> {
DateTime::<Utc>::from_timestamp(secs.min(i64::MAX as u64) as i64, 0)
.map(|value| value.to_rfc3339())
}
impl StoredAdminAuditLog {
pub fn created_at_rfc3339(&self) -> Option<String> {
unix_secs_to_rfc3339(self.created_at_unix_secs)
}
}
impl StoredSuspiciousActivity {
pub fn created_at_rfc3339(&self) -> Option<String> {
unix_secs_to_rfc3339(self.created_at_unix_secs)
}
}
impl StoredUserAuditLog {
pub fn created_at_rfc3339(&self) -> Option<String> {
unix_secs_to_rfc3339(self.created_at_unix_secs)
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct StoredAdminAuditLogPage {
pub items: Vec<StoredAdminAuditLog>,
pub total: u64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct StoredUserAuditLogPage {
pub items: Vec<StoredUserAuditLog>,
pub total: u64,
}
#[async_trait]
pub trait AuditLogReadRepository: Send + Sync {
async fn list_admin_audit_logs(
&self,
query: &AuditLogListQuery,
) -> Result<StoredAdminAuditLogPage, crate::DataLayerError>;
async fn list_admin_suspicious_activities(
&self,
cutoff_unix_secs: u64,
) -> Result<Vec<StoredSuspiciousActivity>, crate::DataLayerError>;
async fn read_admin_user_behavior_event_counts(
&self,
user_id: &str,
cutoff_unix_secs: u64,
) -> Result<std::collections::BTreeMap<String, u64>, crate::DataLayerError>;
async fn list_user_audit_logs(
&self,
user_id: &str,
query: &AuditLogListQuery,
) -> Result<StoredUserAuditLogPage, crate::DataLayerError>;
async fn delete_audit_logs_before(
&self,
cutoff_unix_secs: u64,
limit: usize,
) -> Result<usize, crate::DataLayerError>;
}
pub fn optional_json_from_text(
value: Option<String>,
) -> Result<Option<Value>, crate::DataLayerError> {
value
.filter(|raw| !raw.trim().is_empty())
.map(|raw| {
serde_json::from_str(&raw).map_err(|err| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid audit log metadata json: {err}"
))
})
})
.transpose()
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,73 @@
use async_trait::async_trait;
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredOAuthProviderModuleConfig {
pub provider_type: String,
pub display_name: String,
pub client_id: String,
pub client_secret_encrypted: Option<String>,
pub redirect_uri: String,
}
impl StoredOAuthProviderModuleConfig {
pub fn new(
provider_type: String,
display_name: String,
client_id: String,
client_secret_encrypted: Option<String>,
redirect_uri: String,
) -> Result<Self, crate::DataLayerError> {
if provider_type.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"oauth_providers.provider_type is empty".to_string(),
));
}
if display_name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"oauth_providers.display_name is empty".to_string(),
));
}
Ok(Self {
provider_type,
display_name,
client_id,
client_secret_encrypted,
redirect_uri,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredLdapModuleConfig {
pub server_url: String,
pub bind_dn: String,
pub bind_password_encrypted: Option<String>,
pub base_dn: String,
pub user_search_filter: Option<String>,
pub username_attr: Option<String>,
pub email_attr: Option<String>,
pub display_name_attr: Option<String>,
pub is_enabled: bool,
pub is_exclusive: bool,
pub use_starttls: bool,
pub connect_timeout: Option<i32>,
}
#[async_trait]
pub trait AuthModuleReadRepository: Send + Sync {
async fn list_enabled_oauth_providers(
&self,
) -> Result<Vec<StoredOAuthProviderModuleConfig>, crate::DataLayerError>;
async fn get_ldap_config(
&self,
) -> Result<Option<StoredLdapModuleConfig>, crate::DataLayerError>;
}
#[async_trait]
pub trait AuthModuleWriteRepository: Send + Sync {
async fn upsert_ldap_config(
&self,
config: &StoredLdapModuleConfig,
) -> Result<Option<StoredLdapModuleConfig>, crate::DataLayerError>;
}
@@ -0,0 +1,8 @@
mod types;
pub use types::{
BackgroundTaskKind, BackgroundTaskListQuery, BackgroundTaskReadRepository,
BackgroundTaskRepository, BackgroundTaskStatus, BackgroundTaskSummary,
BackgroundTaskWriteRepository, StoredBackgroundTaskEvent, StoredBackgroundTaskRun,
StoredBackgroundTaskRunPage, UpsertBackgroundTaskEvent, UpsertBackgroundTaskRun,
};
@@ -0,0 +1,290 @@
use std::collections::BTreeMap;
use async_trait::async_trait;
use serde_json::Value;
#[derive(
Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize,
)]
pub enum BackgroundTaskKind {
Scheduled,
Daemon,
OnDemand,
FireAndForget,
}
impl BackgroundTaskKind {
pub fn as_database(self) -> &'static str {
match self {
Self::Scheduled => "scheduled",
Self::Daemon => "daemon",
Self::OnDemand => "on_demand",
Self::FireAndForget => "fire_and_forget",
}
}
pub fn from_database(value: &str) -> Result<Self, crate::DataLayerError> {
match value.trim().to_ascii_lowercase().as_str() {
"scheduled" => Ok(Self::Scheduled),
"daemon" => Ok(Self::Daemon),
"on_demand" => Ok(Self::OnDemand),
"fire_and_forget" => Ok(Self::FireAndForget),
other => Err(crate::DataLayerError::UnexpectedValue(format!(
"unsupported background_tasks.kind: {other}"
))),
}
}
}
#[derive(
Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize,
)]
pub enum BackgroundTaskStatus {
Queued,
Running,
Retrying,
Succeeded,
Failed,
Cancelled,
Skipped,
}
impl BackgroundTaskStatus {
pub fn as_database(self) -> &'static str {
match self {
Self::Queued => "queued",
Self::Running => "running",
Self::Retrying => "retrying",
Self::Succeeded => "succeeded",
Self::Failed => "failed",
Self::Cancelled => "cancelled",
Self::Skipped => "skipped",
}
}
pub fn from_database(value: &str) -> Result<Self, crate::DataLayerError> {
match value.trim().to_ascii_lowercase().as_str() {
"queued" => Ok(Self::Queued),
"running" => Ok(Self::Running),
"retrying" => Ok(Self::Retrying),
"succeeded" => Ok(Self::Succeeded),
"failed" => Ok(Self::Failed),
"cancelled" => Ok(Self::Cancelled),
"skipped" => Ok(Self::Skipped),
other => Err(crate::DataLayerError::UnexpectedValue(format!(
"unsupported background_tasks.status: {other}"
))),
}
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredBackgroundTaskRun {
pub id: String,
pub task_key: String,
pub kind: BackgroundTaskKind,
pub trigger: String,
pub status: BackgroundTaskStatus,
pub attempt: u32,
pub max_attempts: u32,
pub owner_instance: Option<String>,
pub progress_percent: u16,
pub progress_message: Option<String>,
pub payload_json: Option<Value>,
pub result_json: Option<Value>,
pub error_message: Option<String>,
pub cancel_requested: bool,
pub created_by: Option<String>,
pub created_at_unix_secs: u64,
pub started_at_unix_secs: Option<u64>,
pub finished_at_unix_secs: Option<u64>,
pub updated_at_unix_secs: u64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct UpsertBackgroundTaskRun {
pub id: String,
pub task_key: String,
pub kind: BackgroundTaskKind,
pub trigger: String,
pub status: BackgroundTaskStatus,
pub attempt: u32,
pub max_attempts: u32,
pub owner_instance: Option<String>,
pub progress_percent: u16,
pub progress_message: Option<String>,
pub payload_json: Option<Value>,
pub result_json: Option<Value>,
pub error_message: Option<String>,
pub cancel_requested: bool,
pub created_by: Option<String>,
pub created_at_unix_secs: u64,
pub started_at_unix_secs: Option<u64>,
pub finished_at_unix_secs: Option<u64>,
pub updated_at_unix_secs: u64,
}
impl UpsertBackgroundTaskRun {
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
if self.id.trim().is_empty()
|| self.task_key.trim().is_empty()
|| self.trigger.trim().is_empty()
{
return Err(crate::DataLayerError::UnexpectedValue(
"background task run identity is empty".to_string(),
));
}
if self.progress_percent > 100 {
return Err(crate::DataLayerError::UnexpectedValue(format!(
"background task progress_percent out of range: {}",
self.progress_percent
)));
}
Ok(())
}
pub fn into_stored(self) -> StoredBackgroundTaskRun {
StoredBackgroundTaskRun {
id: self.id,
task_key: self.task_key,
kind: self.kind,
trigger: self.trigger,
status: self.status,
attempt: self.attempt,
max_attempts: self.max_attempts,
owner_instance: self.owner_instance,
progress_percent: self.progress_percent,
progress_message: self.progress_message,
payload_json: self.payload_json,
result_json: self.result_json,
error_message: self.error_message,
cancel_requested: self.cancel_requested,
created_by: self.created_by,
created_at_unix_secs: self.created_at_unix_secs,
started_at_unix_secs: self.started_at_unix_secs,
finished_at_unix_secs: self.finished_at_unix_secs,
updated_at_unix_secs: self.updated_at_unix_secs,
}
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredBackgroundTaskEvent {
pub id: String,
pub run_id: String,
pub event_type: String,
pub message: String,
pub payload_json: Option<Value>,
pub created_at_unix_secs: u64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct UpsertBackgroundTaskEvent {
pub id: String,
pub run_id: String,
pub event_type: String,
pub message: String,
pub payload_json: Option<Value>,
pub created_at_unix_secs: u64,
}
impl UpsertBackgroundTaskEvent {
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
if self.id.trim().is_empty()
|| self.run_id.trim().is_empty()
|| self.event_type.trim().is_empty()
|| self.message.trim().is_empty()
{
return Err(crate::DataLayerError::UnexpectedValue(
"background task event identity is empty".to_string(),
));
}
Ok(())
}
pub fn into_stored(self) -> StoredBackgroundTaskEvent {
StoredBackgroundTaskEvent {
id: self.id,
run_id: self.run_id,
event_type: self.event_type,
message: self.message,
payload_json: self.payload_json,
created_at_unix_secs: self.created_at_unix_secs,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
pub struct BackgroundTaskListQuery {
pub task_key_substring: Option<String>,
pub kind: Option<BackgroundTaskKind>,
pub status: Option<BackgroundTaskStatus>,
pub trigger: Option<String>,
pub offset: usize,
pub limit: usize,
}
#[derive(Debug, Clone, PartialEq, Default, serde::Serialize, serde::Deserialize)]
pub struct StoredBackgroundTaskRunPage {
pub items: Vec<StoredBackgroundTaskRun>,
pub total: usize,
}
#[derive(Debug, Clone, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
pub struct BackgroundTaskSummary {
pub total: u64,
pub running_count: u64,
pub by_status: BTreeMap<String, u64>,
pub by_kind: BTreeMap<String, u64>,
}
#[async_trait]
pub trait BackgroundTaskReadRepository: Send + Sync {
async fn find_run(
&self,
run_id: &str,
) -> Result<Option<StoredBackgroundTaskRun>, crate::DataLayerError>;
async fn list_runs(
&self,
query: &BackgroundTaskListQuery,
) -> Result<StoredBackgroundTaskRunPage, crate::DataLayerError>;
async fn list_events(
&self,
run_id: &str,
offset: usize,
limit: usize,
) -> Result<Vec<StoredBackgroundTaskEvent>, crate::DataLayerError>;
async fn summarize_runs(&self) -> Result<BackgroundTaskSummary, crate::DataLayerError>;
}
#[async_trait]
pub trait BackgroundTaskWriteRepository: Send + Sync {
async fn upsert_run(
&self,
run: UpsertBackgroundTaskRun,
) -> Result<StoredBackgroundTaskRun, crate::DataLayerError>;
async fn request_cancel(
&self,
run_id: &str,
updated_at_unix_secs: u64,
) -> Result<bool, crate::DataLayerError>;
async fn upsert_event(
&self,
event: UpsertBackgroundTaskEvent,
) -> Result<StoredBackgroundTaskEvent, crate::DataLayerError>;
}
pub trait BackgroundTaskRepository:
BackgroundTaskReadRepository + BackgroundTaskWriteRepository + Send + Sync
{
}
impl<T> BackgroundTaskRepository for T where
T: BackgroundTaskReadRepository + BackgroundTaskWriteRepository + Send + Sync
{
}
@@ -0,0 +1,9 @@
mod types;
pub use types::{
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput,
BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, PaymentGatewayConfigRecord,
PaymentGatewayConfigWriteInput, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord,
UserPlanEntitlementRecord,
};
@@ -0,0 +1,446 @@
use async_trait::async_trait;
use serde_json::Value;
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredBillingModelContext {
pub provider_id: String,
pub provider_billing_type: Option<String>,
pub provider_api_key_id: Option<String>,
pub provider_api_key_rate_multipliers: Option<Value>,
pub provider_api_key_cache_ttl_minutes: Option<i64>,
pub global_model_id: String,
pub global_model_name: String,
pub global_model_config: Option<Value>,
pub default_price_per_request: Option<f64>,
pub default_tiered_pricing: Option<Value>,
pub model_id: Option<String>,
pub model_provider_model_name: Option<String>,
pub model_config: Option<Value>,
pub model_price_per_request: Option<f64>,
pub model_tiered_pricing: Option<Value>,
}
impl StoredBillingModelContext {
#[allow(clippy::too_many_arguments)]
pub fn new(
provider_id: String,
provider_billing_type: Option<String>,
provider_api_key_id: Option<String>,
provider_api_key_rate_multipliers: Option<Value>,
provider_api_key_cache_ttl_minutes: Option<i64>,
global_model_id: String,
global_model_name: String,
global_model_config: Option<Value>,
default_price_per_request: Option<f64>,
default_tiered_pricing: Option<Value>,
model_id: Option<String>,
model_provider_model_name: Option<String>,
model_config: Option<Value>,
model_price_per_request: Option<f64>,
model_tiered_pricing: Option<Value>,
) -> Result<Self, crate::DataLayerError> {
if provider_id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"billing.provider_id is empty".to_string(),
));
}
if global_model_id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"billing.global_model_id is empty".to_string(),
));
}
if global_model_name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"billing.global_model_name is empty".to_string(),
));
}
Ok(Self {
provider_id,
provider_billing_type,
provider_api_key_id,
provider_api_key_rate_multipliers,
provider_api_key_cache_ttl_minutes,
global_model_id,
global_model_name,
global_model_config,
default_price_per_request,
default_tiered_pricing,
model_id,
model_provider_model_name,
model_config,
model_price_per_request,
model_tiered_pricing,
})
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct AdminBillingRuleRecord {
pub id: String,
pub name: String,
pub task_type: String,
pub global_model_id: Option<String>,
pub model_id: Option<String>,
pub expression: String,
pub variables: Value,
pub dimension_mappings: Value,
pub is_enabled: bool,
pub created_at_unix_ms: u64,
pub updated_at_unix_secs: u64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct AdminBillingRuleWriteInput {
pub name: String,
pub task_type: String,
pub global_model_id: Option<String>,
pub model_id: Option<String>,
pub expression: String,
pub variables: Value,
pub dimension_mappings: Value,
pub is_enabled: bool,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct AdminBillingCollectorRecord {
pub id: String,
pub api_format: String,
pub task_type: String,
pub dimension_name: String,
pub source_type: String,
pub source_path: Option<String>,
pub value_type: String,
pub transform_expression: Option<String>,
pub default_value: Option<String>,
pub priority: i32,
pub is_enabled: bool,
pub created_at_unix_ms: u64,
pub updated_at_unix_secs: u64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct AdminBillingCollectorWriteInput {
pub api_format: String,
pub task_type: String,
pub dimension_name: String,
pub source_type: String,
pub source_path: Option<String>,
pub value_type: String,
pub transform_expression: Option<String>,
pub default_value: Option<String>,
pub priority: i32,
pub is_enabled: bool,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct AdminBillingPresetApplyResult {
pub preset: String,
pub mode: String,
pub created: u64,
pub updated: u64,
pub skipped: u64,
pub errors: Vec<String>,
}
#[derive(Debug, Clone, PartialEq)]
pub enum AdminBillingMutationOutcome<T> {
Applied(T),
NotFound,
Invalid(String),
Unavailable,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct PaymentGatewayConfigRecord {
pub provider: String,
pub enabled: bool,
pub endpoint_url: String,
pub callback_base_url: Option<String>,
pub merchant_id: String,
pub merchant_key_encrypted: Option<String>,
pub pay_currency: String,
pub usd_exchange_rate: f64,
pub min_recharge_usd: f64,
pub channels_json: Value,
pub created_at_unix_secs: u64,
pub updated_at_unix_secs: u64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct PaymentGatewayConfigWriteInput {
pub provider: String,
pub enabled: bool,
pub endpoint_url: String,
pub callback_base_url: Option<String>,
pub merchant_id: String,
pub merchant_key_encrypted: Option<String>,
pub preserve_existing_secret: bool,
pub pay_currency: String,
pub usd_exchange_rate: f64,
pub min_recharge_usd: f64,
pub channels_json: Value,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct BillingPlanRecord {
pub id: String,
pub title: String,
pub description: Option<String>,
pub price_amount: f64,
pub price_currency: String,
pub duration_unit: String,
pub duration_value: i64,
pub enabled: bool,
pub sort_order: i64,
pub max_active_per_user: i64,
pub purchase_limit_scope: String,
pub entitlements_json: Value,
pub created_at_unix_secs: u64,
pub updated_at_unix_secs: u64,
}
#[derive(Debug, Clone, PartialEq)]
pub struct BillingPlanWriteInput {
pub title: String,
pub description: Option<String>,
pub price_amount: f64,
pub price_currency: String,
pub duration_unit: String,
pub duration_value: i64,
pub enabled: bool,
pub sort_order: i64,
pub max_active_per_user: i64,
pub purchase_limit_scope: String,
pub entitlements_json: Value,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct UserPlanEntitlementRecord {
pub id: String,
pub user_id: String,
pub plan_id: String,
pub payment_order_id: String,
pub status: String,
pub starts_at_unix_secs: u64,
pub expires_at_unix_secs: u64,
pub entitlements_snapshot: Value,
pub created_at_unix_secs: u64,
pub updated_at_unix_secs: u64,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct UserDailyQuotaAvailabilityRecord {
pub has_active_daily_quota: bool,
pub total_quota_usd: f64,
pub used_usd: f64,
pub remaining_usd: f64,
pub allow_wallet_overage: bool,
}
#[async_trait]
pub trait BillingReadRepository: Send + Sync {
async fn find_model_context(
&self,
provider_id: &str,
provider_api_key_id: Option<&str>,
global_model_name: &str,
) -> Result<Option<StoredBillingModelContext>, crate::DataLayerError>;
async fn find_model_context_by_model_id(
&self,
provider_id: &str,
provider_api_key_id: Option<&str>,
model_id: &str,
) -> Result<Option<StoredBillingModelContext>, crate::DataLayerError> {
let _ = (provider_id, provider_api_key_id, model_id);
Ok(None)
}
async fn admin_billing_enabled_default_value_exists(
&self,
api_format: &str,
task_type: &str,
dimension_name: &str,
existing_id: Option<&str>,
) -> Result<Option<bool>, crate::DataLayerError> {
let _ = (api_format, task_type, dimension_name, existing_id);
Ok(None)
}
async fn create_admin_billing_rule(
&self,
input: &AdminBillingRuleWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingRuleRecord>, crate::DataLayerError> {
let _ = input;
Ok(AdminBillingMutationOutcome::Unavailable)
}
async fn list_admin_billing_rules(
&self,
task_type: Option<&str>,
is_enabled: Option<bool>,
page: u32,
page_size: u32,
) -> Result<Option<(Vec<AdminBillingRuleRecord>, u64)>, crate::DataLayerError> {
let _ = (task_type, is_enabled, page, page_size);
Ok(None)
}
async fn find_admin_billing_rule(
&self,
rule_id: &str,
) -> Result<Option<AdminBillingRuleRecord>, crate::DataLayerError> {
let _ = rule_id;
Ok(None)
}
async fn update_admin_billing_rule(
&self,
rule_id: &str,
input: &AdminBillingRuleWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingRuleRecord>, crate::DataLayerError> {
let _ = (rule_id, input);
Ok(AdminBillingMutationOutcome::Unavailable)
}
async fn create_admin_billing_collector(
&self,
input: &AdminBillingCollectorWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingCollectorRecord>, crate::DataLayerError>
{
let _ = input;
Ok(AdminBillingMutationOutcome::Unavailable)
}
async fn list_admin_billing_collectors(
&self,
api_format: Option<&str>,
task_type: Option<&str>,
dimension_name: Option<&str>,
is_enabled: Option<bool>,
page: u32,
page_size: u32,
) -> Result<Option<(Vec<AdminBillingCollectorRecord>, u64)>, crate::DataLayerError> {
let _ = (
api_format,
task_type,
dimension_name,
is_enabled,
page,
page_size,
);
Ok(None)
}
async fn find_admin_billing_collector(
&self,
collector_id: &str,
) -> Result<Option<AdminBillingCollectorRecord>, crate::DataLayerError> {
let _ = collector_id;
Ok(None)
}
async fn update_admin_billing_collector(
&self,
collector_id: &str,
input: &AdminBillingCollectorWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingCollectorRecord>, crate::DataLayerError>
{
let _ = (collector_id, input);
Ok(AdminBillingMutationOutcome::Unavailable)
}
async fn apply_admin_billing_preset(
&self,
preset: &str,
mode: &str,
collectors: &[AdminBillingCollectorWriteInput],
) -> Result<AdminBillingMutationOutcome<AdminBillingPresetApplyResult>, crate::DataLayerError>
{
let _ = (preset, mode, collectors);
Ok(AdminBillingMutationOutcome::Unavailable)
}
async fn find_payment_gateway_config(
&self,
provider: &str,
) -> Result<Option<PaymentGatewayConfigRecord>, crate::DataLayerError> {
let _ = provider;
Ok(None)
}
async fn upsert_payment_gateway_config(
&self,
input: &PaymentGatewayConfigWriteInput,
) -> Result<AdminBillingMutationOutcome<PaymentGatewayConfigRecord>, crate::DataLayerError>
{
let _ = input;
Ok(AdminBillingMutationOutcome::Unavailable)
}
async fn list_billing_plans(
&self,
include_disabled: bool,
) -> Result<Option<Vec<BillingPlanRecord>>, crate::DataLayerError> {
let _ = include_disabled;
Ok(None)
}
async fn find_billing_plan(
&self,
plan_id: &str,
) -> Result<Option<BillingPlanRecord>, crate::DataLayerError> {
let _ = plan_id;
Ok(None)
}
async fn create_billing_plan(
&self,
input: &BillingPlanWriteInput,
) -> Result<AdminBillingMutationOutcome<BillingPlanRecord>, crate::DataLayerError> {
let _ = input;
Ok(AdminBillingMutationOutcome::Unavailable)
}
async fn update_billing_plan(
&self,
plan_id: &str,
input: &BillingPlanWriteInput,
) -> Result<AdminBillingMutationOutcome<BillingPlanRecord>, crate::DataLayerError> {
let _ = (plan_id, input);
Ok(AdminBillingMutationOutcome::Unavailable)
}
async fn set_billing_plan_enabled(
&self,
plan_id: &str,
enabled: bool,
) -> Result<AdminBillingMutationOutcome<BillingPlanRecord>, crate::DataLayerError> {
let _ = (plan_id, enabled);
Ok(AdminBillingMutationOutcome::Unavailable)
}
async fn delete_billing_plan(
&self,
plan_id: &str,
) -> Result<AdminBillingMutationOutcome<()>, crate::DataLayerError> {
let _ = plan_id;
Ok(AdminBillingMutationOutcome::Unavailable)
}
async fn list_user_plan_entitlements(
&self,
user_id: &str,
) -> Result<Option<Vec<UserPlanEntitlementRecord>>, crate::DataLayerError> {
let _ = user_id;
Ok(None)
}
async fn find_user_daily_quota_availability(
&self,
user_id: &str,
) -> Result<Option<UserDailyQuotaAvailabilityRecord>, crate::DataLayerError> {
let _ = user_id;
Ok(None)
}
}
@@ -0,0 +1,8 @@
mod types;
pub use types::{
MinimalCandidateSelectionReadRepository, MinimalCandidateSelectionRepository,
StoredMinimalCandidateSelectionRow, StoredPoolKeyCandidateOrder,
StoredPoolKeyCandidateRowsByKeyIdsQuery, StoredPoolKeyCandidateRowsQuery,
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
};
@@ -0,0 +1,158 @@
use async_trait::async_trait;
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredProviderModelMapping {
pub name: String,
pub priority: i32,
pub api_formats: Option<Vec<String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub endpoint_ids: Option<Vec<String>>,
/// Optional request-operation scope. An omitted scope applies to every
/// operation supported by the selected API format.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub operations: Option<Vec<String>>,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredMinimalCandidateSelectionRow {
pub provider_id: String,
pub provider_name: String,
pub provider_type: String,
pub provider_priority: i32,
pub provider_is_active: bool,
pub endpoint_id: String,
pub endpoint_api_format: String,
pub endpoint_api_family: Option<String>,
pub endpoint_kind: Option<String>,
pub endpoint_is_active: bool,
pub key_id: String,
pub key_name: String,
pub key_auth_type: String,
pub key_is_active: bool,
pub key_api_formats: Option<Vec<String>>,
pub key_allowed_models: Option<Vec<String>>,
pub key_capabilities: Option<serde_json::Value>,
pub key_internal_priority: i32,
pub key_global_priority_by_format: Option<serde_json::Value>,
pub model_id: String,
pub global_model_id: String,
pub global_model_name: String,
pub global_model_mappings: Option<Vec<String>>,
pub global_model_supports_streaming: Option<bool>,
pub model_provider_model_name: String,
pub model_provider_model_mappings: Option<Vec<StoredProviderModelMapping>>,
pub model_supports_streaming: Option<bool>,
pub model_is_active: bool,
pub model_is_available: bool,
}
#[derive(Debug, Clone, Default, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub enum StoredPoolKeyCandidateOrder {
#[default]
InternalPriority,
Lru,
CacheAffinity,
SingleAccount,
LoadBalance {
seed: String,
},
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredPoolKeyCandidateRowsQuery {
pub api_format: String,
pub provider_id: String,
pub endpoint_id: String,
pub model_id: String,
pub selected_provider_model_name: String,
#[serde(default)]
pub order: StoredPoolKeyCandidateOrder,
pub offset: u32,
pub limit: u32,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredPoolKeyCandidateRowsByKeyIdsQuery {
pub api_format: String,
pub provider_id: String,
pub endpoint_id: String,
pub model_id: String,
pub selected_provider_model_name: String,
pub key_ids: Vec<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredRequestedModelCandidateRowsQuery {
pub api_format: String,
pub requested_model_name: String,
pub offset: u32,
pub limit: u32,
}
impl StoredMinimalCandidateSelectionRow {
pub fn supports_streaming(&self) -> bool {
self.model_supports_streaming
.or(self.global_model_supports_streaming)
.unwrap_or(true)
}
pub fn key_supports_api_format(&self, api_format: &str) -> bool {
match self.key_api_formats.as_deref() {
None => true,
Some(formats) => formats
.iter()
.any(|value| api_format_permission_covers(value, api_format)),
}
}
}
fn api_format_permission_covers(allowed: &str, requested: &str) -> bool {
aether_ai_formats::api_format_permission_covers(allowed, requested)
}
#[async_trait]
pub trait MinimalCandidateSelectionReadRepository: Send + Sync {
fn clear_local_cache(&self) {}
async fn list_for_exact_api_format(
&self,
api_format: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
async fn list_for_exact_api_format_and_global_model(
&self,
api_format: &str,
global_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
async fn list_for_exact_api_format_and_requested_model(
&self,
api_format: &str,
requested_model_name: &str,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
async fn list_for_exact_api_format_and_requested_model_page(
&self,
query: &StoredRequestedModelCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
async fn list_pool_key_rows_for_group(
&self,
query: &StoredPoolKeyCandidateRowsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
async fn list_pool_key_rows_for_group_key_ids(
&self,
query: &StoredPoolKeyCandidateRowsByKeyIdsQuery,
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
}
pub trait MinimalCandidateSelectionRepository:
MinimalCandidateSelectionReadRepository + Send + Sync
{
}
impl<T> MinimalCandidateSelectionRepository for T where
T: MinimalCandidateSelectionReadRepository + Send + Sync
{
}
@@ -0,0 +1,10 @@
mod types;
pub use types::{
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,
};
@@ -0,0 +1,657 @@
use std::collections::BTreeMap;
use async_trait::async_trait;
use crate::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum RequestCandidateStatus {
Available,
Unused,
Pending,
Streaming,
Success,
Failed,
Cancelled,
Skipped,
}
impl RequestCandidateStatus {
pub fn from_database(value: &str) -> Result<Self, crate::DataLayerError> {
match value.trim().to_ascii_lowercase().as_str() {
"available" => Ok(Self::Available),
"unused" => Ok(Self::Unused),
"pending" => Ok(Self::Pending),
"streaming" => Ok(Self::Streaming),
"success" => Ok(Self::Success),
"failed" => Ok(Self::Failed),
"cancelled" => Ok(Self::Cancelled),
"skipped" => Ok(Self::Skipped),
other => Err(crate::DataLayerError::UnexpectedValue(format!(
"unsupported request_candidates.status: {other}"
))),
}
}
pub fn is_attempted(self, started_at_unix_ms: Option<u64>) -> bool {
match self {
Self::Available | Self::Unused | Self::Skipped => false,
Self::Pending => started_at_unix_ms.is_some(),
Self::Streaming | Self::Success | Self::Failed | Self::Cancelled => true,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredRequestCandidate {
pub id: String,
pub request_id: String,
pub user_id: Option<String>,
pub api_key_id: Option<String>,
pub username: Option<String>,
pub api_key_name: Option<String>,
pub candidate_index: u32,
pub retry_index: u32,
pub provider_id: Option<String>,
pub endpoint_id: Option<String>,
pub key_id: Option<String>,
pub status: RequestCandidateStatus,
pub skip_reason: Option<String>,
pub is_cached: bool,
pub status_code: Option<u16>,
pub error_type: Option<String>,
pub error_message: Option<String>,
pub latency_ms: Option<u64>,
pub concurrent_requests: Option<u32>,
pub extra_data: Option<serde_json::Value>,
pub required_capabilities: Option<serde_json::Value>,
pub created_at_unix_ms: u64,
pub started_at_unix_ms: Option<u64>,
pub finished_at_unix_ms: Option<u64>,
}
impl StoredRequestCandidate {
#[allow(clippy::too_many_arguments)]
pub fn new(
id: String,
request_id: String,
user_id: Option<String>,
api_key_id: Option<String>,
username: Option<String>,
api_key_name: Option<String>,
candidate_index: i32,
retry_index: i32,
provider_id: Option<String>,
endpoint_id: Option<String>,
key_id: Option<String>,
status: RequestCandidateStatus,
skip_reason: Option<String>,
is_cached: bool,
status_code: Option<i32>,
error_type: Option<String>,
error_message: Option<String>,
latency_ms: Option<i32>,
concurrent_requests: Option<i32>,
extra_data: Option<serde_json::Value>,
required_capabilities: Option<serde_json::Value>,
created_at_unix_ms: i64,
started_at_unix_ms: Option<i64>,
finished_at_unix_ms: Option<i64>,
) -> Result<Self, crate::DataLayerError> {
let candidate_index = u32::try_from(candidate_index).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid request_candidates.candidate_index: {candidate_index}"
))
})?;
let retry_index = u32::try_from(retry_index).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid request_candidates.retry_index: {retry_index}"
))
})?;
let status_code = status_code
.map(|value| {
u16::try_from(value).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid request_candidates.status_code: {value}"
))
})
})
.transpose()?;
let latency_ms = latency_ms
.map(|value| {
u64::try_from(value).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid request_candidates.latency_ms: {value}"
))
})
})
.transpose()?;
let concurrent_requests = concurrent_requests
.map(|value| {
u32::try_from(value).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid request_candidates.concurrent_requests: {value}"
))
})
})
.transpose()?;
let created_at_unix_ms = u64::try_from(created_at_unix_ms).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid request_candidates.created_at_unix_ms: {created_at_unix_ms}"
))
})?;
let started_at_unix_ms = started_at_unix_ms
.map(|value| {
u64::try_from(value).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid request_candidates.started_at_unix_ms: {value}"
))
})
})
.transpose()?;
let finished_at_unix_ms = finished_at_unix_ms
.map(|value| {
u64::try_from(value).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid request_candidates.finished_at_unix_ms: {value}"
))
})
})
.transpose()?;
Ok(Self {
id,
request_id,
user_id,
api_key_id,
username,
api_key_name,
candidate_index,
retry_index,
provider_id,
endpoint_id,
key_id,
status,
skip_reason,
is_cached,
status_code,
error_type,
error_message,
latency_ms,
concurrent_requests,
extra_data,
required_capabilities,
created_at_unix_ms,
started_at_unix_ms,
finished_at_unix_ms,
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum RequestCandidateFinalStatus {
Success,
Failed,
Cancelled,
Streaming,
Pending,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct RequestCandidateTrace {
pub request_id: String,
pub total_candidates: usize,
pub final_status: RequestCandidateFinalStatus,
pub total_latency_ms: u64,
pub candidates: Vec<StoredRequestCandidate>,
}
impl RequestCandidateTrace {
pub fn from_candidates(
request_id: impl Into<String>,
all_candidates: Vec<StoredRequestCandidate>,
attempted_only: bool,
) -> Option<Self> {
if all_candidates.is_empty() {
return None;
}
let candidates = if attempted_only {
all_candidates
.iter()
.filter(|candidate| candidate.status.is_attempted(candidate.started_at_unix_ms))
.cloned()
.collect::<Vec<_>>()
} else {
all_candidates.clone()
};
let total_latency_ms = candidates
.iter()
.filter(|candidate| {
matches!(
candidate.status,
RequestCandidateStatus::Success
| RequestCandidateStatus::Failed
| RequestCandidateStatus::Cancelled
) && candidate.latency_ms.is_some()
})
.map(|candidate| candidate.latency_ms.unwrap_or(0))
.sum();
let final_status_source = if attempted_only && candidates.is_empty() {
&all_candidates
} else {
&candidates
};
Some(Self {
request_id: request_id.into(),
total_candidates: candidates.len(),
final_status: derive_request_candidate_final_status(final_status_source),
total_latency_ms,
candidates,
})
}
}
pub fn derive_request_candidate_final_status(
candidates: &[StoredRequestCandidate],
) -> RequestCandidateFinalStatus {
let has_success = candidates
.iter()
.any(|candidate| candidate.status == RequestCandidateStatus::Success);
if has_success {
return RequestCandidateFinalStatus::Success;
}
let has_failed = candidates
.iter()
.any(|candidate| candidate.status == RequestCandidateStatus::Failed);
if has_failed {
return RequestCandidateFinalStatus::Failed;
}
let has_cancelled = candidates
.iter()
.any(|candidate| candidate.status == RequestCandidateStatus::Cancelled);
if has_cancelled {
return RequestCandidateFinalStatus::Cancelled;
}
if candidates
.iter()
.any(|candidate| candidate.status == RequestCandidateStatus::Streaming)
{
return RequestCandidateFinalStatus::Streaming;
}
if candidates
.iter()
.any(|candidate| candidate.status == RequestCandidateStatus::Pending)
{
return RequestCandidateFinalStatus::Pending;
}
let has_legacy_success_status_code = candidates
.iter()
.any(|candidate| matches!(candidate.status_code, Some(status_code) if (200..300).contains(&status_code)));
if has_legacy_success_status_code {
return RequestCandidateFinalStatus::Success;
}
RequestCandidateFinalStatus::Failed
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct DecisionTraceCandidate {
#[serde(flatten)]
pub candidate: StoredRequestCandidate,
pub provider_name: Option<String>,
pub provider_website: Option<String>,
pub provider_type: Option<String>,
pub provider_priority: Option<i32>,
pub provider_keep_priority_on_conversion: Option<bool>,
pub provider_enable_format_conversion: Option<bool>,
pub endpoint_api_format: Option<String>,
pub endpoint_api_family: Option<String>,
pub endpoint_kind: Option<String>,
pub endpoint_format_acceptance_config: Option<serde_json::Value>,
pub provider_key_name: Option<String>,
pub provider_key_auth_type: Option<String>,
pub provider_key_api_formats: Option<serde_json::Value>,
pub provider_key_internal_priority: Option<i32>,
pub provider_key_global_priority_by_format: Option<serde_json::Value>,
pub provider_key_capabilities: Option<serde_json::Value>,
pub provider_key_is_active: Option<bool>,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct DecisionTrace {
pub request_id: String,
pub total_candidates: usize,
pub final_status: RequestCandidateFinalStatus,
pub total_latency_ms: u64,
pub candidates: Vec<DecisionTraceCandidate>,
}
pub fn build_decision_trace(
trace: RequestCandidateTrace,
providers: Vec<StoredProviderCatalogProvider>,
endpoints: Vec<StoredProviderCatalogEndpoint>,
keys: Vec<StoredProviderCatalogKey>,
) -> DecisionTrace {
let provider_map = providers
.into_iter()
.map(|item| (item.id.clone(), item))
.collect::<BTreeMap<_, _>>();
let endpoint_map = endpoints
.into_iter()
.map(|item| (item.id.clone(), item))
.collect::<BTreeMap<_, _>>();
let key_map = keys
.into_iter()
.map(|item| (item.id.clone(), item))
.collect::<BTreeMap<_, _>>();
DecisionTrace {
request_id: trace.request_id,
total_candidates: trace.total_candidates,
final_status: trace.final_status,
total_latency_ms: trace.total_latency_ms,
candidates: trace
.candidates
.into_iter()
.map(|candidate| {
enrich_decision_trace_candidate(candidate, &provider_map, &endpoint_map, &key_map)
})
.collect(),
}
}
fn enrich_decision_trace_candidate(
candidate: StoredRequestCandidate,
provider_map: &BTreeMap<String, StoredProviderCatalogProvider>,
endpoint_map: &BTreeMap<String, StoredProviderCatalogEndpoint>,
key_map: &BTreeMap<String, StoredProviderCatalogKey>,
) -> DecisionTraceCandidate {
let provider = candidate
.provider_id
.as_ref()
.and_then(|provider_id| provider_map.get(provider_id));
let endpoint = candidate
.endpoint_id
.as_ref()
.and_then(|endpoint_id| endpoint_map.get(endpoint_id));
let provider_key = candidate
.key_id
.as_ref()
.and_then(|key_id| key_map.get(key_id));
DecisionTraceCandidate {
provider_name: provider.map(|item| item.name.clone()),
provider_website: provider.and_then(|item| item.website.clone()),
provider_type: provider.map(|item| item.provider_type.clone()),
provider_priority: provider.map(|item| item.provider_priority),
provider_keep_priority_on_conversion: provider.map(|item| item.keep_priority_on_conversion),
provider_enable_format_conversion: provider.map(|item| item.enable_format_conversion),
endpoint_api_format: endpoint.map(|item| item.api_format.clone()),
endpoint_api_family: endpoint.and_then(|item| item.api_family.clone()),
endpoint_kind: endpoint.and_then(|item| item.endpoint_kind.clone()),
endpoint_format_acceptance_config: endpoint
.and_then(|item| item.format_acceptance_config.clone()),
provider_key_name: provider_key
.map(|item| item.name.clone())
.or_else(|| candidate.api_key_name.clone()),
provider_key_auth_type: provider_key.map(|item| item.auth_type.clone()),
provider_key_api_formats: provider_key.and_then(|item| item.api_formats.clone()),
provider_key_internal_priority: provider_key.map(|item| item.internal_priority),
provider_key_global_priority_by_format: provider_key
.and_then(|item| item.global_priority_by_format.clone()),
provider_key_capabilities: provider_key.and_then(|item| item.capabilities.clone()),
provider_key_is_active: provider_key.map(|item| item.is_active),
candidate,
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct PublicHealthStatusCount {
pub endpoint_id: String,
pub status: RequestCandidateStatus,
pub count: u64,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct PublicHealthTimelineBucket {
pub endpoint_id: String,
pub segment_idx: u32,
pub total_count: u64,
pub success_count: u64,
pub failed_count: u64,
pub min_created_at_unix_ms: Option<u64>,
pub max_created_at_unix_ms: Option<u64>,
}
#[async_trait]
pub trait RequestCandidateReadRepository: Send + Sync {
async fn list_by_request_id(
&self,
request_id: &str,
) -> Result<Vec<StoredRequestCandidate>, crate::DataLayerError>;
async fn list_attempted_by_request_id(
&self,
request_id: &str,
) -> Result<Vec<StoredRequestCandidate>, crate::DataLayerError> {
Ok(self
.list_by_request_id(request_id)
.await?
.into_iter()
.filter(|candidate| candidate.status.is_attempted(candidate.started_at_unix_ms))
.collect())
}
async fn list_recent(
&self,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, crate::DataLayerError>;
async fn list_by_provider_id(
&self,
provider_id: &str,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, crate::DataLayerError>;
async fn list_finalized_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
since_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, crate::DataLayerError>;
async fn count_finalized_statuses_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
since_unix_secs: u64,
) -> Result<Vec<PublicHealthStatusCount>, crate::DataLayerError>;
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>, crate::DataLayerError>;
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct UpsertRequestCandidateRecord {
pub id: String,
pub request_id: String,
pub user_id: Option<String>,
pub api_key_id: Option<String>,
pub username: Option<String>,
pub api_key_name: Option<String>,
pub candidate_index: u32,
pub retry_index: u32,
pub provider_id: Option<String>,
pub endpoint_id: Option<String>,
pub key_id: Option<String>,
pub status: RequestCandidateStatus,
pub skip_reason: Option<String>,
pub is_cached: Option<bool>,
pub status_code: Option<u16>,
pub error_type: Option<String>,
pub error_message: Option<String>,
pub latency_ms: Option<u64>,
pub concurrent_requests: Option<u32>,
pub extra_data: Option<serde_json::Value>,
pub required_capabilities: Option<serde_json::Value>,
pub created_at_unix_ms: Option<u64>,
pub started_at_unix_ms: Option<u64>,
pub finished_at_unix_ms: Option<u64>,
}
impl UpsertRequestCandidateRecord {
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
if self.id.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"request candidate upsert id cannot be empty".to_string(),
));
}
if self.request_id.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"request candidate upsert request_id cannot be empty".to_string(),
));
}
Ok(())
}
}
#[async_trait]
pub trait RequestCandidateWriteRepository: Send + Sync {
async fn upsert(
&self,
candidate: UpsertRequestCandidateRecord,
) -> Result<StoredRequestCandidate, crate::DataLayerError>;
async fn upsert_many(
&self,
candidates: Vec<UpsertRequestCandidateRecord>,
) -> Result<usize, crate::DataLayerError> {
let mut persisted = 0usize;
for candidate in candidates {
self.upsert(candidate).await?;
persisted = persisted.saturating_add(1);
}
Ok(persisted)
}
async fn delete_created_before(
&self,
created_before_unix_secs: u64,
limit: usize,
) -> Result<usize, crate::DataLayerError>;
}
pub trait RequestCandidateRepository:
RequestCandidateReadRepository + RequestCandidateWriteRepository + Send + Sync
{
}
impl<T> RequestCandidateRepository for T where
T: RequestCandidateReadRepository + RequestCandidateWriteRepository + Send + Sync
{
}
pub fn request_candidate_lifecycle_would_regress(
existing: RequestCandidateStatus,
incoming: RequestCandidateStatus,
) -> bool {
matches!(
existing,
RequestCandidateStatus::Success
| RequestCandidateStatus::Failed
| RequestCandidateStatus::Cancelled
| RequestCandidateStatus::Skipped
) && matches!(
incoming,
RequestCandidateStatus::Available
| RequestCandidateStatus::Unused
| RequestCandidateStatus::Pending
| RequestCandidateStatus::Streaming
) || existing == RequestCandidateStatus::Streaming
&& incoming == RequestCandidateStatus::Pending
}
#[cfg(test)]
mod tests {
use super::{
derive_request_candidate_final_status, RequestCandidateFinalStatus, RequestCandidateStatus,
StoredRequestCandidate,
};
fn candidate(
id: &str,
status: RequestCandidateStatus,
status_code: Option<i32>,
) -> StoredRequestCandidate {
StoredRequestCandidate::new(
id.to_string(),
"req-1".to_string(),
None,
None,
None,
None,
0,
0,
None,
None,
None,
status,
None,
false,
status_code,
None,
None,
Some(100),
None,
None,
None,
1_700_000_000_000,
Some(1_700_000_000_000),
Some(1_700_000_000_100),
)
.expect("candidate should build")
}
#[test]
fn failed_candidate_with_http_200_stays_final_failed() {
let candidates = vec![candidate(
"cand-1",
RequestCandidateStatus::Failed,
Some(200),
)];
assert_eq!(
derive_request_candidate_final_status(&candidates),
RequestCandidateFinalStatus::Failed
);
}
#[test]
fn explicit_success_candidate_still_wins_after_failed_attempt() {
let candidates = vec![
candidate("cand-1", RequestCandidateStatus::Failed, Some(503)),
candidate("cand-2", RequestCandidateStatus::Success, Some(200)),
];
assert_eq!(
derive_request_candidate_final_status(&candidates),
RequestCandidateFinalStatus::Success
);
}
}
@@ -0,0 +1,165 @@
use async_trait::async_trait;
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct GeminiFileMappingListQuery {
pub include_expired: bool,
pub search: Option<String>,
pub offset: usize,
pub limit: usize,
pub now_unix_secs: u64,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StoredGeminiFileMappingListPage {
pub items: Vec<StoredGeminiFileMapping>,
pub total: usize,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct GeminiFileMappingMimeTypeCount {
pub mime_type: String,
pub count: usize,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct GeminiFileMappingStats {
pub total_mappings: usize,
pub active_mappings: usize,
pub expired_mappings: usize,
pub by_mime_type: Vec<GeminiFileMappingMimeTypeCount>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StoredGeminiFileMapping {
pub id: String,
pub file_name: String,
pub key_id: String,
pub user_id: Option<String>,
pub display_name: Option<String>,
pub mime_type: Option<String>,
pub source_hash: Option<String>,
pub created_at_unix_ms: u64,
pub expires_at_unix_secs: u64,
}
impl StoredGeminiFileMapping {
pub fn new(
id: String,
file_name: String,
key_id: String,
created_at_unix_ms: i64,
expires_at_unix_secs: i64,
) -> Result<Self, crate::DataLayerError> {
if file_name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"gemini_file_mappings.file_name is empty".to_string(),
));
}
if key_id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"gemini_file_mappings.key_id is empty".to_string(),
));
}
let created_at_unix_ms = u64::try_from(created_at_unix_ms).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid gemini_file_mappings.created_at: {created_at_unix_ms}"
))
})?;
let expires_at_unix_secs = u64::try_from(expires_at_unix_secs).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid gemini_file_mappings.expires_at: {expires_at_unix_secs}"
))
})?;
Ok(Self {
id,
file_name,
key_id,
user_id: None,
display_name: None,
mime_type: None,
source_hash: None,
created_at_unix_ms,
expires_at_unix_secs,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct UpsertGeminiFileMappingRecord {
pub id: String,
pub file_name: String,
pub key_id: String,
pub user_id: Option<String>,
pub display_name: Option<String>,
pub mime_type: Option<String>,
pub source_hash: Option<String>,
pub expires_at_unix_secs: u64,
}
impl UpsertGeminiFileMappingRecord {
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
if self.file_name.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"gemini_file_mappings.file_name is empty".to_string(),
));
}
if self.key_id.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"gemini_file_mappings.key_id is empty".to_string(),
));
}
if self.expires_at_unix_secs == 0 {
return Err(crate::DataLayerError::InvalidInput(
"gemini_file_mappings.expires_at is empty".to_string(),
));
}
Ok(())
}
}
#[async_trait]
pub trait GeminiFileMappingReadRepository: Send + Sync {
async fn find_by_file_name(
&self,
file_name: &str,
) -> Result<Option<StoredGeminiFileMapping>, crate::DataLayerError>;
async fn list_mappings(
&self,
query: &GeminiFileMappingListQuery,
) -> Result<StoredGeminiFileMappingListPage, crate::DataLayerError>;
async fn summarize_mappings(
&self,
now_unix_secs: u64,
) -> Result<GeminiFileMappingStats, crate::DataLayerError>;
}
#[async_trait]
pub trait GeminiFileMappingWriteRepository: Send + Sync {
async fn upsert(
&self,
record: UpsertGeminiFileMappingRecord,
) -> Result<StoredGeminiFileMapping, crate::DataLayerError>;
async fn delete_by_file_name(&self, file_name: &str) -> Result<bool, crate::DataLayerError>;
async fn delete_by_id(
&self,
mapping_id: &str,
) -> Result<Option<StoredGeminiFileMapping>, crate::DataLayerError>;
async fn delete_expired_before(
&self,
now_unix_secs: u64,
) -> Result<usize, crate::DataLayerError>;
}
pub trait GeminiFileMappingRepository:
GeminiFileMappingReadRepository + GeminiFileMappingWriteRepository
{
}
impl<T> GeminiFileMappingRepository for T where
T: GeminiFileMappingReadRepository + GeminiFileMappingWriteRepository
{
}
@@ -0,0 +1,13 @@
mod snapshot;
mod types;
pub use snapshot::GlobalModelSnapshot;
pub use types::{
metadata_supports_embedding, AdminGlobalModelListQuery, AdminProviderModelListQuery,
CreateAdminGlobalModelRecord, GlobalModelReadRepository, GlobalModelWriteRepository,
PublicCatalogModelListQuery, PublicCatalogModelSearchQuery, PublicGlobalModelQuery,
StoredAdminGlobalModel, StoredAdminGlobalModelPage, StoredAdminProviderModel,
StoredProviderActiveGlobalModel, StoredProviderModelStats, StoredPublicCatalogModel,
StoredPublicGlobalModel, StoredPublicGlobalModelPage, UpdateAdminGlobalModelRecord,
UpsertAdminProviderModelRecord,
};
@@ -0,0 +1,352 @@
use std::collections::BTreeSet;
use super::{
AdminGlobalModelListQuery, AdminProviderModelListQuery, PublicCatalogModelListQuery,
PublicCatalogModelSearchQuery, PublicGlobalModelQuery, StoredAdminGlobalModel,
StoredAdminGlobalModelPage, StoredAdminProviderModel, StoredProviderActiveGlobalModel,
StoredProviderModelStats, StoredPublicCatalogModel, StoredPublicGlobalModel,
StoredPublicGlobalModelPage,
};
/// Immutable input for the shared global-model read policy.
///
/// Database adapters can load one consistent view and apply exactly the same
/// filtering, sorting, pagination, and enrichment rules as the memory adapter
/// without depending on the `aether-data` facade.
#[derive(Debug, Clone, Default)]
pub struct GlobalModelSnapshot {
public_global_models: Vec<StoredPublicGlobalModel>,
admin_global_models: Vec<StoredAdminGlobalModel>,
public_catalog_models: Vec<StoredPublicCatalogModel>,
admin_provider_models: Vec<StoredAdminProviderModel>,
provider_model_stats: Vec<StoredProviderModelStats>,
active_global_model_refs: Vec<StoredProviderActiveGlobalModel>,
}
impl GlobalModelSnapshot {
pub fn seed<I>(items: I) -> Self
where
I: IntoIterator<Item = StoredPublicGlobalModel>,
{
Self {
public_global_models: items.into_iter().collect(),
..Self::default()
}
}
pub fn with_admin_global_models<I>(mut self, items: I) -> Self
where
I: IntoIterator<Item = StoredAdminGlobalModel>,
{
self.admin_global_models = items.into_iter().collect();
self
}
pub fn with_public_catalog_models<I>(mut self, items: I) -> Self
where
I: IntoIterator<Item = StoredPublicCatalogModel>,
{
self.public_catalog_models = items.into_iter().collect();
self
}
pub fn with_admin_provider_models<I>(mut self, items: I) -> Self
where
I: IntoIterator<Item = StoredAdminProviderModel>,
{
self.admin_provider_models = items.into_iter().collect();
self
}
pub fn with_provider_model_stats<I>(mut self, items: I) -> Self
where
I: IntoIterator<Item = StoredProviderModelStats>,
{
self.provider_model_stats = items.into_iter().collect();
self
}
pub fn with_active_global_model_refs<I>(mut self, items: I) -> Self
where
I: IntoIterator<Item = StoredProviderActiveGlobalModel>,
{
self.active_global_model_refs = items.into_iter().collect();
self
}
pub fn list_public_models(
&self,
query: &PublicGlobalModelQuery,
) -> StoredPublicGlobalModelPage {
let search = normalized_optional_search(query.search.as_deref());
let mut filtered = self
.public_global_models
.iter()
.filter(|item| match query.is_active {
Some(is_active) => item.is_active == is_active,
None => item.is_active,
})
.filter(|item| {
let Some(search) = search.as_deref() else {
return true;
};
item.name.to_ascii_lowercase().contains(search)
|| item
.display_name
.as_deref()
.is_some_and(|value| value.to_ascii_lowercase().contains(search))
})
.cloned()
.collect::<Vec<_>>();
filtered.sort_by(|left, right| left.name.cmp(&right.name));
let total = filtered.len();
let items = filtered
.into_iter()
.skip(query.offset)
.take(query.limit)
.collect();
StoredPublicGlobalModelPage { items, total }
}
pub fn get_public_model_by_name(&self, model_name: &str) -> Option<StoredPublicGlobalModel> {
self.public_global_models
.iter()
.find(|item| item.is_active && item.name == model_name)
.cloned()
}
pub fn list_public_catalog_models(
&self,
query: &PublicCatalogModelListQuery,
) -> Vec<StoredPublicCatalogModel> {
let provider_id = normalized_optional_value(query.provider_id.as_deref());
let mut filtered = self
.public_catalog_models
.iter()
.filter(|item| item.is_active)
.filter(|item| provider_id.is_none_or(|value| item.provider_id == value))
.cloned()
.collect::<Vec<_>>();
sort_public_catalog_models(&mut filtered);
filtered
.into_iter()
.skip(query.offset)
.take(query.limit)
.collect()
}
pub fn search_public_catalog_models(
&self,
query: &PublicCatalogModelSearchQuery,
) -> Vec<StoredPublicCatalogModel> {
let provider_id = normalized_optional_value(query.provider_id.as_deref());
let search = query.search.trim().to_ascii_lowercase();
let mut filtered = self
.public_catalog_models
.iter()
.filter(|item| item.is_active)
.filter(|item| provider_id.is_none_or(|value| item.provider_id == value))
.filter(|item| {
item.provider_model_name
.to_ascii_lowercase()
.contains(&search)
|| item.name.to_ascii_lowercase().contains(&search)
|| item.display_name.to_ascii_lowercase().contains(&search)
})
.cloned()
.collect::<Vec<_>>();
sort_public_catalog_models(&mut filtered);
filtered.truncate(query.limit);
filtered
}
pub fn list_admin_global_models(
&self,
query: &AdminGlobalModelListQuery,
) -> StoredAdminGlobalModelPage {
let search = normalized_optional_search(query.search.as_deref());
let mut filtered = self
.admin_global_models
.iter()
.filter(|item| query.is_active.is_none_or(|value| item.is_active == value))
.filter(|item| {
let Some(search) = search.as_deref() else {
return true;
};
item.name.to_ascii_lowercase().contains(search)
|| item.display_name.to_ascii_lowercase().contains(search)
})
.cloned()
.collect::<Vec<_>>();
filtered.sort_by(|left, right| left.name.cmp(&right.name));
let total = filtered.len();
let items = filtered
.into_iter()
.skip(query.offset)
.take(query.limit)
.map(|item| self.enrich_admin_global_model(&item))
.collect();
StoredAdminGlobalModelPage { items, total }
}
pub fn list_admin_provider_models(
&self,
query: &AdminProviderModelListQuery,
) -> Vec<StoredAdminProviderModel> {
let mut filtered = self
.admin_provider_models
.iter()
.filter(|item| item.provider_id == query.provider_id)
.filter(|item| query.is_active.is_none_or(|value| item.is_active == value))
.cloned()
.collect::<Vec<_>>();
sort_admin_provider_models(&mut filtered);
filtered
.into_iter()
.skip(query.offset)
.take(query.limit)
.collect()
}
pub fn get_admin_provider_model(
&self,
provider_id: &str,
model_id: &str,
) -> Option<StoredAdminProviderModel> {
self.admin_provider_models
.iter()
.find(|item| item.provider_id == provider_id && item.id == model_id)
.cloned()
}
pub fn list_admin_provider_available_source_models(
&self,
provider_id: &str,
) -> Vec<StoredAdminProviderModel> {
let mut filtered = self
.admin_provider_models
.iter()
.filter(|item| item.provider_id == provider_id && item.is_active)
.filter(|item| {
self.admin_global_models
.iter()
.find(|global| global.id == item.global_model_id)
.is_some_and(|global| global.is_active)
})
.cloned()
.collect::<Vec<_>>();
filtered.sort_by(|left, right| {
left.global_model_name
.cmp(&right.global_model_name)
.then_with(|| right.created_at_unix_ms.cmp(&left.created_at_unix_ms))
.then_with(|| left.id.cmp(&right.id))
});
filtered
}
pub fn get_admin_global_model_by_id(
&self,
global_model_id: &str,
) -> Option<StoredAdminGlobalModel> {
self.admin_global_models
.iter()
.find(|item| item.id == global_model_id)
.map(|item| self.enrich_admin_global_model(item))
}
pub fn get_admin_global_model_by_name(
&self,
model_name: &str,
) -> Option<StoredAdminGlobalModel> {
self.admin_global_models
.iter()
.find(|item| item.name == model_name)
.map(|item| self.enrich_admin_global_model(item))
}
pub fn list_admin_provider_models_by_global_model_id(
&self,
global_model_id: &str,
) -> Vec<StoredAdminProviderModel> {
let mut filtered = self
.admin_provider_models
.iter()
.filter(|item| item.global_model_id == global_model_id)
.cloned()
.collect::<Vec<_>>();
sort_admin_provider_models(&mut filtered);
filtered
}
pub fn list_provider_model_stats(
&self,
provider_ids: &[String],
) -> Vec<StoredProviderModelStats> {
let provider_ids = provider_ids.iter().collect::<BTreeSet<_>>();
self.provider_model_stats
.iter()
.filter(|item| provider_ids.contains(&item.provider_id))
.cloned()
.collect()
}
pub fn list_active_global_model_ids_by_provider_ids(
&self,
provider_ids: &[String],
) -> Vec<StoredProviderActiveGlobalModel> {
let provider_ids = provider_ids.iter().collect::<BTreeSet<_>>();
self.active_global_model_refs
.iter()
.filter(|item| provider_ids.contains(&item.provider_id))
.cloned()
.collect()
}
fn enrich_admin_global_model(&self, item: &StoredAdminGlobalModel) -> StoredAdminGlobalModel {
let mut enriched = item.clone();
let providers = self
.admin_provider_models
.iter()
.filter(|model| model.global_model_id == item.id)
.map(|model| model.provider_id.as_str())
.collect::<BTreeSet<_>>();
let active_providers = self
.admin_provider_models
.iter()
.filter(|model| {
model.global_model_id == item.id && model.is_active && model.is_available
})
.map(|model| model.provider_id.as_str())
.collect::<BTreeSet<_>>();
enriched.provider_count = providers.len() as u64;
enriched.active_provider_count = active_providers.len() as u64;
enriched
}
}
fn normalized_optional_value(value: Option<&str>) -> Option<&str> {
value.map(str::trim).filter(|value| !value.is_empty())
}
fn normalized_optional_search(value: Option<&str>) -> Option<String> {
normalized_optional_value(value).map(str::to_ascii_lowercase)
}
fn sort_public_catalog_models(items: &mut [StoredPublicCatalogModel]) {
items.sort_by(|left, right| {
left.provider_name
.cmp(&right.provider_name)
.then_with(|| left.name.cmp(&right.name))
.then_with(|| left.id.cmp(&right.id))
});
}
fn sort_admin_provider_models(items: &mut [StoredAdminProviderModel]) {
items.sort_by(|left, right| {
right
.created_at_unix_ms
.unwrap_or_default()
.cmp(&left.created_at_unix_ms.unwrap_or_default())
.then_with(|| left.id.cmp(&right.id))
});
}
@@ -0,0 +1,988 @@
use async_trait::async_trait;
use serde_json::Value;
const EMBEDDING_CAPABILITY: &str = "embedding";
const EMBEDDING_API_FORMATS: &[&str] = &[
"openai:embedding",
"gemini:embedding",
"jina:embedding",
"doubao:embedding",
"aliyun:multimodal_embedding",
"/v1/embeddings",
"/jina/v1/embeddings",
];
fn validate_optional_price(
field_name: &str,
value: Option<f64>,
) -> Result<(), crate::DataLayerError> {
if value.is_some_and(|price| !price.is_finite() || price < 0.0) {
return Err(crate::DataLayerError::UnexpectedValue(format!(
"{field_name} must be a non-negative finite number"
)));
}
Ok(())
}
fn validate_embedding_global_billing(
default_price_per_request: Option<f64>,
default_tiered_pricing: Option<&Value>,
supported_capabilities: Option<&Value>,
config: Option<&Value>,
) -> Result<(), crate::DataLayerError> {
validate_optional_price(
"global_models.default_price_per_request",
default_price_per_request,
)?;
if !has_embedding_metadata(supported_capabilities, config) {
return Ok(());
}
if has_request_or_input_token_pricing(default_price_per_request, default_tiered_pricing) {
return Ok(());
}
Err(crate::DataLayerError::UnexpectedValue(
"embedding global model requires default_price_per_request or default_tiered_pricing.tiers[].input_price_per_1m".to_string(),
))
}
fn validate_provider_model_pricing(
price_per_request: Option<f64>,
) -> Result<(), crate::DataLayerError> {
validate_optional_price("models.price_per_request", price_per_request)
}
fn has_request_or_input_token_pricing(
price_per_request: Option<f64>,
tiered_pricing: Option<&Value>,
) -> bool {
price_per_request.is_some_and(|price| price.is_finite() && price >= 0.0)
|| tiered_pricing
.and_then(|value| value.get("tiers"))
.and_then(Value::as_array)
.is_some_and(|tiers| {
tiers.iter().any(|tier| {
tier.get("input_price_per_1m")
.and_then(Value::as_f64)
.is_some_and(|price| price.is_finite() && price >= 0.0)
})
})
}
fn has_embedding_metadata(supported_capabilities: Option<&Value>, config: Option<&Value>) -> bool {
metadata_supports_embedding(supported_capabilities, config, None).unwrap_or(false)
}
/// Derives embedding support from global capabilities and both global/model metadata.
///
/// `Some(false)` is intentional: callers use this value to distinguish a completed
/// metadata decision from an absent database column.
pub fn metadata_supports_embedding(
supported_capabilities: Option<&Value>,
global_config: Option<&Value>,
model_config: Option<&Value>,
) -> Option<bool> {
Some(
supported_capabilities.is_some_and(value_contains_embedding_capability)
|| global_config.is_some_and(value_contains_embedding_metadata)
|| model_config.is_some_and(value_contains_embedding_metadata),
)
}
fn value_contains_embedding_capability(value: &Value) -> bool {
match value {
Value::String(value) => value.trim().eq_ignore_ascii_case(EMBEDDING_CAPABILITY),
Value::Array(values) => values.iter().any(value_contains_embedding_capability),
Value::Object(object) => {
object
.get(EMBEDDING_CAPABILITY)
.and_then(Value::as_bool)
.unwrap_or(false)
|| [
"capability",
"model_type",
"type",
"task_type",
"request_type",
]
.iter()
.any(|key| {
object
.get(*key)
.is_some_and(value_contains_embedding_capability)
})
|| ["capabilities", "supported_capabilities"]
.iter()
.any(|key| {
object
.get(*key)
.is_some_and(value_contains_embedding_capability)
})
}
_ => false,
}
}
fn value_contains_embedding_metadata(value: &Value) -> bool {
match value {
Value::String(value) => {
value.trim().eq_ignore_ascii_case(EMBEDDING_CAPABILITY)
|| is_known_embedding_api_format(value)
}
Value::Array(values) => values.iter().any(value_contains_embedding_metadata),
Value::Object(object) => {
value_contains_embedding_capability(value)
|| ["api_format", "client_api_format", "provider_api_format"]
.iter()
.any(|key| {
object
.get(*key)
.and_then(Value::as_str)
.is_some_and(is_known_embedding_api_format)
})
|| ["api_formats", "client_api_formats", "provider_api_formats"]
.iter()
.any(|key| {
object
.get(*key)
.is_some_and(value_contains_embedding_metadata)
})
}
_ => false,
}
}
fn is_known_embedding_api_format(value: &str) -> bool {
let normalized = value.trim().to_ascii_lowercase();
EMBEDDING_API_FORMATS
.iter()
.any(|api_format| normalized == *api_format || normalized.ends_with(*api_format))
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredPublicGlobalModel {
pub id: String,
pub name: String,
pub display_name: Option<String>,
pub is_active: bool,
pub default_price_per_request: Option<f64>,
pub default_tiered_pricing: Option<Value>,
pub supported_capabilities: Option<Value>,
pub config: Option<Value>,
pub usage_count: u64,
}
impl StoredPublicGlobalModel {
#[allow(clippy::too_many_arguments)]
pub fn new(
id: String,
name: String,
display_name: Option<String>,
is_active: bool,
default_price_per_request: Option<f64>,
default_tiered_pricing: Option<Value>,
supported_capabilities: Option<Value>,
config: Option<Value>,
usage_count: u64,
) -> Result<Self, crate::DataLayerError> {
if id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"global_models.id is empty".to_string(),
));
}
if name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"global_models.name is empty".to_string(),
));
}
validate_embedding_global_billing(
default_price_per_request,
default_tiered_pricing.as_ref(),
supported_capabilities.as_ref(),
config.as_ref(),
)?;
Ok(Self {
id,
name,
display_name,
is_active,
default_price_per_request,
default_tiered_pricing,
supported_capabilities,
config,
usage_count,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct PublicGlobalModelQuery {
pub offset: usize,
pub limit: usize,
pub is_active: Option<bool>,
pub search: Option<String>,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredPublicCatalogModel {
pub id: String,
pub provider_id: String,
pub provider_name: String,
pub provider_model_name: String,
pub name: String,
pub display_name: String,
pub description: Option<String>,
pub icon_url: Option<String>,
pub input_price_per_1m: Option<f64>,
pub output_price_per_1m: Option<f64>,
pub cache_creation_price_per_1m: Option<f64>,
pub cache_read_price_per_1m: Option<f64>,
pub supports_vision: Option<bool>,
pub supports_function_calling: Option<bool>,
pub supports_streaming: Option<bool>,
pub supports_embedding: Option<bool>,
pub is_active: bool,
}
impl StoredPublicCatalogModel {
#[allow(clippy::too_many_arguments)]
pub fn new(
id: String,
provider_id: String,
provider_name: String,
provider_model_name: String,
name: String,
display_name: String,
description: Option<String>,
icon_url: Option<String>,
input_price_per_1m: Option<f64>,
output_price_per_1m: Option<f64>,
cache_creation_price_per_1m: Option<f64>,
cache_read_price_per_1m: Option<f64>,
supports_vision: Option<bool>,
supports_function_calling: Option<bool>,
supports_streaming: Option<bool>,
supports_embedding: Option<bool>,
is_active: bool,
) -> Result<Self, crate::DataLayerError> {
if id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"models.id is empty".to_string(),
));
}
if provider_id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"models.provider_id is empty".to_string(),
));
}
if provider_name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"providers.name is empty".to_string(),
));
}
if provider_model_name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"models.provider_model_name is empty".to_string(),
));
}
if name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"public model name is empty".to_string(),
));
}
if display_name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"public model display_name is empty".to_string(),
));
}
Ok(Self {
id,
provider_id,
provider_name,
provider_model_name,
name,
display_name,
description,
icon_url,
input_price_per_1m,
output_price_per_1m,
cache_creation_price_per_1m,
cache_read_price_per_1m,
supports_vision,
supports_function_calling,
supports_streaming,
supports_embedding,
is_active,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct PublicCatalogModelListQuery {
pub provider_id: Option<String>,
pub offset: usize,
pub limit: usize,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PublicCatalogModelSearchQuery {
pub search: String,
pub provider_id: Option<String>,
pub limit: usize,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AdminProviderModelListQuery {
pub provider_id: String,
pub is_active: Option<bool>,
pub offset: usize,
pub limit: usize,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredAdminGlobalModel {
pub id: String,
pub name: String,
pub display_name: String,
pub is_active: bool,
pub default_price_per_request: Option<f64>,
pub default_tiered_pricing: Option<Value>,
pub supported_capabilities: Option<Value>,
pub config: Option<Value>,
pub provider_count: u64,
pub active_provider_count: u64,
pub usage_count: u64,
pub created_at_unix_ms: Option<u64>,
pub updated_at_unix_secs: Option<u64>,
}
impl StoredAdminGlobalModel {
#[allow(clippy::too_many_arguments)]
pub fn new(
id: String,
name: String,
display_name: String,
is_active: bool,
default_price_per_request: Option<f64>,
default_tiered_pricing: Option<Value>,
supported_capabilities: Option<Value>,
config: Option<Value>,
provider_count: u64,
active_provider_count: u64,
usage_count: u64,
created_at_unix_ms: Option<u64>,
updated_at_unix_secs: Option<u64>,
) -> Result<Self, crate::DataLayerError> {
if id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"global_models.id is empty".to_string(),
));
}
if name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"global_models.name is empty".to_string(),
));
}
if display_name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"global_models.display_name is empty".to_string(),
));
}
validate_embedding_global_billing(
default_price_per_request,
default_tiered_pricing.as_ref(),
supported_capabilities.as_ref(),
config.as_ref(),
)?;
Ok(Self {
id,
name,
display_name,
is_active,
default_price_per_request,
default_tiered_pricing,
supported_capabilities,
config,
provider_count,
active_provider_count,
usage_count,
created_at_unix_ms,
updated_at_unix_secs,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct AdminGlobalModelListQuery {
pub offset: usize,
pub limit: usize,
pub is_active: Option<bool>,
pub search: Option<String>,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredAdminProviderModel {
pub id: String,
pub provider_id: String,
pub global_model_id: String,
pub provider_model_name: String,
pub provider_model_mappings: Option<Value>,
pub price_per_request: Option<f64>,
pub tiered_pricing: Option<Value>,
pub supports_vision: Option<bool>,
pub supports_function_calling: Option<bool>,
pub supports_streaming: Option<bool>,
pub supports_extended_thinking: Option<bool>,
pub supports_image_generation: Option<bool>,
pub is_active: bool,
pub is_available: bool,
pub config: Option<Value>,
pub created_at_unix_ms: Option<u64>,
pub updated_at_unix_secs: Option<u64>,
pub global_model_name: Option<String>,
pub global_model_display_name: Option<String>,
pub global_model_default_price_per_request: Option<f64>,
pub global_model_default_tiered_pricing: Option<Value>,
pub global_model_supported_capabilities: Option<Value>,
pub global_model_config: Option<Value>,
}
impl StoredAdminProviderModel {
#[allow(clippy::too_many_arguments)]
pub fn new(
id: String,
provider_id: String,
global_model_id: String,
provider_model_name: String,
provider_model_mappings: Option<Value>,
price_per_request: Option<f64>,
tiered_pricing: Option<Value>,
supports_vision: Option<bool>,
supports_function_calling: Option<bool>,
supports_streaming: Option<bool>,
supports_extended_thinking: Option<bool>,
supports_image_generation: Option<bool>,
is_active: bool,
is_available: bool,
config: Option<Value>,
created_at_unix_ms: Option<u64>,
updated_at_unix_secs: Option<u64>,
global_model_name: Option<String>,
global_model_display_name: Option<String>,
global_model_default_price_per_request: Option<f64>,
global_model_default_tiered_pricing: Option<Value>,
global_model_supported_capabilities: Option<Value>,
global_model_config: Option<Value>,
) -> Result<Self, crate::DataLayerError> {
if id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"models.id is empty".to_string(),
));
}
if provider_id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"models.provider_id is empty".to_string(),
));
}
if global_model_id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"models.global_model_id is empty".to_string(),
));
}
if provider_model_name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"models.provider_model_name is empty".to_string(),
));
}
validate_provider_model_pricing(price_per_request)?;
Ok(Self {
id,
provider_id,
global_model_id,
provider_model_name,
provider_model_mappings,
price_per_request,
tiered_pricing,
supports_vision,
supports_function_calling,
supports_streaming,
supports_extended_thinking,
supports_image_generation,
is_active,
is_available,
config,
created_at_unix_ms,
updated_at_unix_secs,
global_model_name,
global_model_display_name,
global_model_default_price_per_request,
global_model_default_tiered_pricing,
global_model_supported_capabilities,
global_model_config,
})
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct UpsertAdminProviderModelRecord {
pub id: String,
pub provider_id: String,
pub global_model_id: String,
pub provider_model_name: String,
pub provider_model_mappings: Option<Value>,
pub price_per_request: Option<f64>,
pub tiered_pricing: Option<Value>,
pub supports_vision: Option<bool>,
pub supports_function_calling: Option<bool>,
pub supports_streaming: Option<bool>,
pub supports_extended_thinking: Option<bool>,
pub supports_image_generation: Option<bool>,
pub is_active: bool,
pub is_available: bool,
pub config: Option<Value>,
}
impl UpsertAdminProviderModelRecord {
#[allow(clippy::too_many_arguments)]
pub fn new(
id: String,
provider_id: String,
global_model_id: String,
provider_model_name: String,
provider_model_mappings: Option<Value>,
price_per_request: Option<f64>,
tiered_pricing: Option<Value>,
supports_vision: Option<bool>,
supports_function_calling: Option<bool>,
supports_streaming: Option<bool>,
supports_extended_thinking: Option<bool>,
supports_image_generation: Option<bool>,
is_active: bool,
is_available: bool,
config: Option<Value>,
) -> Result<Self, crate::DataLayerError> {
if id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"models.id is empty".to_string(),
));
}
if provider_id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"models.provider_id is empty".to_string(),
));
}
if global_model_id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"models.global_model_id is empty".to_string(),
));
}
if provider_model_name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"models.provider_model_name is empty".to_string(),
));
}
validate_provider_model_pricing(price_per_request)?;
Ok(Self {
id,
provider_id,
global_model_id,
provider_model_name,
provider_model_mappings,
price_per_request,
tiered_pricing,
supports_vision,
supports_function_calling,
supports_streaming,
supports_extended_thinking,
supports_image_generation,
is_active,
is_available,
config,
})
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct CreateAdminGlobalModelRecord {
pub id: String,
pub name: String,
pub display_name: String,
pub is_active: bool,
pub default_price_per_request: Option<f64>,
pub default_tiered_pricing: Option<Value>,
pub supported_capabilities: Option<Value>,
pub config: Option<Value>,
#[serde(default)]
pub usage_count: Option<u64>,
}
impl CreateAdminGlobalModelRecord {
#[allow(clippy::too_many_arguments)]
pub fn new(
id: String,
name: String,
display_name: String,
is_active: bool,
default_price_per_request: Option<f64>,
default_tiered_pricing: Option<Value>,
supported_capabilities: Option<Value>,
config: Option<Value>,
) -> Result<Self, crate::DataLayerError> {
if id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"global_models.id is empty".to_string(),
));
}
if name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"global_models.name is empty".to_string(),
));
}
if display_name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"global_models.display_name is empty".to_string(),
));
}
validate_embedding_global_billing(
default_price_per_request,
default_tiered_pricing.as_ref(),
supported_capabilities.as_ref(),
config.as_ref(),
)?;
Ok(Self {
id,
name,
display_name,
is_active,
default_price_per_request,
default_tiered_pricing,
supported_capabilities,
config,
usage_count: None,
})
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct UpdateAdminGlobalModelRecord {
pub id: String,
pub display_name: String,
pub is_active: bool,
pub default_price_per_request: Option<f64>,
pub default_tiered_pricing: Option<Value>,
pub supported_capabilities: Option<Value>,
pub config: Option<Value>,
#[serde(default)]
pub usage_count: Option<u64>,
}
impl UpdateAdminGlobalModelRecord {
#[allow(clippy::too_many_arguments)]
pub fn new(
id: String,
display_name: String,
is_active: bool,
default_price_per_request: Option<f64>,
default_tiered_pricing: Option<Value>,
supported_capabilities: Option<Value>,
config: Option<Value>,
) -> Result<Self, crate::DataLayerError> {
if id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"global_models.id is empty".to_string(),
));
}
if display_name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"global_models.display_name is empty".to_string(),
));
}
validate_embedding_global_billing(
default_price_per_request,
default_tiered_pricing.as_ref(),
supported_capabilities.as_ref(),
config.as_ref(),
)?;
Ok(Self {
id,
display_name,
is_active,
default_price_per_request,
default_tiered_pricing,
supported_capabilities,
config,
usage_count: None,
})
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredPublicGlobalModelPage {
pub items: Vec<StoredPublicGlobalModel>,
pub total: usize,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredAdminGlobalModelPage {
pub items: Vec<StoredAdminGlobalModel>,
pub total: usize,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredProviderModelStats {
pub provider_id: String,
pub total_models: u64,
pub active_models: u64,
}
impl StoredProviderModelStats {
pub fn new(
provider_id: String,
total_models: i64,
active_models: i64,
) -> Result<Self, crate::DataLayerError> {
if provider_id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"provider model stats provider_id is empty".to_string(),
));
}
if total_models < 0 || active_models < 0 {
return Err(crate::DataLayerError::UnexpectedValue(
"provider model stats count is negative".to_string(),
));
}
Ok(Self {
provider_id,
total_models: total_models as u64,
active_models: active_models as u64,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredProviderActiveGlobalModel {
pub provider_id: String,
pub global_model_id: String,
}
impl StoredProviderActiveGlobalModel {
pub fn new(
provider_id: String,
global_model_id: String,
) -> Result<Self, crate::DataLayerError> {
if provider_id.trim().is_empty() || global_model_id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"provider active global model identity is empty".to_string(),
));
}
Ok(Self {
provider_id,
global_model_id,
})
}
}
#[async_trait]
pub trait GlobalModelReadRepository: Send + Sync {
async fn list_public_models(
&self,
query: &PublicGlobalModelQuery,
) -> Result<StoredPublicGlobalModelPage, crate::DataLayerError>;
async fn get_public_model_by_name(
&self,
model_name: &str,
) -> Result<Option<StoredPublicGlobalModel>, crate::DataLayerError>;
async fn list_public_catalog_models(
&self,
query: &PublicCatalogModelListQuery,
) -> Result<Vec<StoredPublicCatalogModel>, crate::DataLayerError>;
async fn search_public_catalog_models(
&self,
query: &PublicCatalogModelSearchQuery,
) -> Result<Vec<StoredPublicCatalogModel>, crate::DataLayerError>;
async fn list_admin_global_models(
&self,
query: &AdminGlobalModelListQuery,
) -> Result<StoredAdminGlobalModelPage, crate::DataLayerError>;
async fn list_admin_provider_models(
&self,
query: &AdminProviderModelListQuery,
) -> Result<Vec<StoredAdminProviderModel>, crate::DataLayerError>;
async fn list_admin_provider_available_source_models(
&self,
provider_id: &str,
) -> Result<Vec<StoredAdminProviderModel>, crate::DataLayerError>;
async fn get_admin_provider_model(
&self,
provider_id: &str,
model_id: &str,
) -> Result<Option<StoredAdminProviderModel>, crate::DataLayerError>;
async fn get_admin_global_model_by_id(
&self,
global_model_id: &str,
) -> Result<Option<StoredAdminGlobalModel>, crate::DataLayerError>;
async fn get_admin_global_model_by_name(
&self,
model_name: &str,
) -> Result<Option<StoredAdminGlobalModel>, crate::DataLayerError>;
async fn list_admin_provider_models_by_global_model_id(
&self,
global_model_id: &str,
) -> Result<Vec<StoredAdminProviderModel>, crate::DataLayerError>;
async fn list_provider_model_stats(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderModelStats>, crate::DataLayerError>;
async fn list_active_global_model_ids_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderActiveGlobalModel>, crate::DataLayerError>;
}
#[async_trait]
pub trait GlobalModelWriteRepository: Send + Sync {
async fn create_admin_provider_model(
&self,
record: &UpsertAdminProviderModelRecord,
) -> Result<Option<StoredAdminProviderModel>, crate::DataLayerError>;
async fn update_admin_provider_model(
&self,
record: &UpsertAdminProviderModelRecord,
) -> Result<Option<StoredAdminProviderModel>, crate::DataLayerError>;
async fn delete_admin_provider_model(
&self,
provider_id: &str,
model_id: &str,
) -> Result<bool, crate::DataLayerError>;
async fn create_admin_global_model(
&self,
record: &CreateAdminGlobalModelRecord,
) -> Result<Option<StoredAdminGlobalModel>, crate::DataLayerError>;
async fn update_admin_global_model(
&self,
record: &UpdateAdminGlobalModelRecord,
) -> Result<Option<StoredAdminGlobalModel>, crate::DataLayerError>;
async fn delete_admin_global_model(
&self,
global_model_id: &str,
) -> Result<bool, crate::DataLayerError>;
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::{CreateAdminGlobalModelRecord, UpsertAdminProviderModelRecord};
#[test]
fn embedding_missing_billing_config_rejected() {
let err = 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 model without explicit billing should be rejected");
assert!(err.to_string().contains(
"embedding global model requires default_price_per_request or default_tiered_pricing.tiers[].input_price_per_1m"
));
}
#[test]
fn embedding_api_format_config_requires_billing_config() {
let err = CreateAdminGlobalModelRecord::new(
"gm-embedding".to_string(),
"jina-embeddings-v3".to_string(),
"Jina Embeddings v3".to_string(),
true,
None,
Some(json!({"tiers": []})),
None,
Some(json!({"api_formats": ["jina:embedding"]})),
)
.expect_err("embedding API format without price should be rejected");
assert!(err.to_string().contains("embedding global model requires"));
}
#[test]
fn embedding_input_token_or_request_pricing_is_accepted() {
CreateAdminGlobalModelRecord::new(
"gm-embedding-input".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!({"dimensions": 1536})),
)
.expect("input-token pricing should satisfy embedding billing");
CreateAdminGlobalModelRecord::new(
"gm-embedding-request".to_string(),
"custom-embedding".to_string(),
"Custom Embedding".to_string(),
true,
Some(0.0),
None,
Some(json!(["embedding"])),
None,
)
.expect("explicit request pricing should satisfy embedding billing");
}
#[test]
fn provider_model_negative_request_price_rejected() {
let err = UpsertAdminProviderModelRecord::new(
"model-1".to_string(),
"provider-1".to_string(),
"global-model-1".to_string(),
"text-embedding-3-small".to_string(),
None,
Some(-0.01),
None,
None,
None,
None,
None,
None,
true,
true,
None,
)
.expect_err("negative provider model request price should be rejected");
assert!(err
.to_string()
.contains("models.price_per_request must be a non-negative finite number"));
}
}
@@ -0,0 +1,381 @@
use async_trait::async_trait;
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredManagementTokenUserSummary {
pub id: String,
pub email: Option<String>,
pub username: String,
pub role: String,
}
impl StoredManagementTokenUserSummary {
pub fn new(
id: String,
email: Option<String>,
username: String,
role: String,
) -> Result<Self, crate::DataLayerError> {
if id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"users.id is empty".to_string(),
));
}
if username.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"users.username is empty".to_string(),
));
}
if role.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"users.role is empty".to_string(),
));
}
Ok(Self {
id,
email,
username,
role,
})
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredManagementToken {
pub id: String,
pub user_id: String,
pub name: String,
pub description: Option<String>,
pub token_prefix: Option<String>,
pub allowed_ips: Option<serde_json::Value>,
pub permissions: Option<serde_json::Value>,
pub expires_at_unix_secs: Option<u64>,
pub last_used_at_unix_secs: Option<u64>,
pub last_used_ip: Option<String>,
pub usage_count: u64,
pub is_active: bool,
pub created_at_unix_ms: Option<u64>,
pub updated_at_unix_secs: Option<u64>,
}
impl StoredManagementToken {
pub fn new(id: String, user_id: String, name: String) -> Result<Self, crate::DataLayerError> {
if id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"management_tokens.id is empty".to_string(),
));
}
if user_id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"management_tokens.user_id is empty".to_string(),
));
}
if name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"management_tokens.name is empty".to_string(),
));
}
Ok(Self {
id,
user_id,
name,
description: None,
token_prefix: None,
allowed_ips: None,
permissions: None,
expires_at_unix_secs: None,
last_used_at_unix_secs: None,
last_used_ip: None,
usage_count: 0,
is_active: true,
created_at_unix_ms: None,
updated_at_unix_secs: None,
})
}
pub fn with_display_fields(
mut self,
description: Option<String>,
token_prefix: Option<String>,
allowed_ips: Option<serde_json::Value>,
) -> Self {
self.description = description;
self.token_prefix = token_prefix;
self.allowed_ips = allowed_ips;
self
}
pub fn with_permissions(mut self, permissions: Option<serde_json::Value>) -> Self {
self.permissions = permissions;
self
}
pub fn with_runtime_fields(
mut self,
expires_at_unix_secs: Option<u64>,
last_used_at_unix_secs: Option<u64>,
last_used_ip: Option<String>,
usage_count: u64,
is_active: bool,
) -> Self {
self.expires_at_unix_secs = expires_at_unix_secs;
self.last_used_at_unix_secs = last_used_at_unix_secs;
self.last_used_ip = last_used_ip;
self.usage_count = usage_count;
self.is_active = is_active;
self
}
pub fn with_timestamps(
mut self,
created_at_unix_ms: Option<u64>,
updated_at_unix_secs: Option<u64>,
) -> Self {
self.created_at_unix_ms = created_at_unix_ms;
self.updated_at_unix_secs = updated_at_unix_secs;
self
}
pub fn token_display(&self) -> String {
self.token_prefix
.as_deref()
.map(|prefix| format!("{prefix}...****"))
.unwrap_or_else(|| "ae-****".to_string())
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredManagementTokenWithUser {
pub token: StoredManagementToken,
pub user: StoredManagementTokenUserSummary,
}
impl StoredManagementTokenWithUser {
pub fn new(token: StoredManagementToken, user: StoredManagementTokenUserSummary) -> Self {
Self { token, user }
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct ManagementTokenListQuery {
pub user_id: Option<String>,
pub is_active: Option<bool>,
pub offset: usize,
pub limit: usize,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct CreateManagementTokenRecord {
pub id: String,
pub user_id: String,
pub user: StoredManagementTokenUserSummary,
pub token_hash: String,
pub token_prefix: Option<String>,
pub name: String,
pub description: Option<String>,
pub allowed_ips: Option<serde_json::Value>,
pub permissions: Option<serde_json::Value>,
pub expires_at_unix_secs: Option<u64>,
pub is_active: bool,
}
impl CreateManagementTokenRecord {
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
if self.id.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"token_id is required".to_string(),
));
}
if self.user_id.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"user_id is required".to_string(),
));
}
if self.user.id != self.user_id {
return Err(crate::DataLayerError::InvalidInput(
"management token user summary does not match user_id".to_string(),
));
}
if self.token_hash.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"token_hash is required".to_string(),
));
}
if self.name.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"name is required".to_string(),
));
}
if let Some(allowed_ips) = &self.allowed_ips {
let Some(items) = allowed_ips.as_array() else {
return Err(crate::DataLayerError::InvalidInput(
"IP 限制规则必须是数组".to_string(),
));
};
if items.is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"IP 限制规则不能为空".to_string(),
));
}
if items.iter().any(|value| value.as_str().is_none()) {
return Err(crate::DataLayerError::InvalidInput(
"IP 限制规则只能包含字符串".to_string(),
));
}
}
validate_management_token_permissions(self.permissions.as_ref())?;
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct UpdateManagementTokenRecord {
pub token_id: String,
pub name: Option<String>,
pub description: Option<String>,
pub clear_description: bool,
pub allowed_ips: Option<serde_json::Value>,
pub clear_allowed_ips: bool,
pub permissions: Option<serde_json::Value>,
pub expires_at_unix_secs: Option<u64>,
pub clear_expires_at: bool,
pub is_active: Option<bool>,
}
impl UpdateManagementTokenRecord {
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
if self.token_id.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"token_id is required".to_string(),
));
}
if let Some(name) = &self.name {
if name.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"name must not be empty".to_string(),
));
}
}
if let Some(allowed_ips) = &self.allowed_ips {
let Some(items) = allowed_ips.as_array() else {
return Err(crate::DataLayerError::InvalidInput(
"IP 限制规则必须是数组".to_string(),
));
};
if items.is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"IP 限制规则不能为空".to_string(),
));
}
if items.iter().any(|value| value.as_str().is_none()) {
return Err(crate::DataLayerError::InvalidInput(
"IP 限制规则只能包含字符串".to_string(),
));
}
}
validate_management_token_permissions(self.permissions.as_ref())?;
Ok(())
}
}
fn validate_management_token_permissions(
permissions: Option<&serde_json::Value>,
) -> Result<(), crate::DataLayerError> {
let Some(permissions) = permissions else {
return Ok(());
};
let Some(items) = permissions.as_array() else {
return Err(crate::DataLayerError::InvalidInput(
"permissions must be an array".to_string(),
));
};
if items.is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"permissions must not be empty".to_string(),
));
}
if items.iter().any(|value| value.as_str().is_none()) {
return Err(crate::DataLayerError::InvalidInput(
"permissions must contain only strings".to_string(),
));
}
Ok(())
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct RegenerateManagementTokenSecret {
pub token_id: String,
pub token_hash: String,
pub token_prefix: Option<String>,
}
impl RegenerateManagementTokenSecret {
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
if self.token_id.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"token_id is required".to_string(),
));
}
if self.token_hash.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"token_hash is required".to_string(),
));
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredManagementTokenListPage {
pub items: Vec<StoredManagementTokenWithUser>,
pub total: usize,
}
#[async_trait]
pub trait ManagementTokenReadRepository: Send + Sync {
async fn list_management_tokens(
&self,
query: &ManagementTokenListQuery,
) -> Result<StoredManagementTokenListPage, crate::DataLayerError>;
async fn get_management_token_with_user(
&self,
token_id: &str,
) -> Result<Option<StoredManagementTokenWithUser>, crate::DataLayerError>;
async fn get_management_token_with_user_by_hash(
&self,
token_hash: &str,
) -> Result<Option<StoredManagementTokenWithUser>, crate::DataLayerError>;
}
#[async_trait]
pub trait ManagementTokenWriteRepository: Send + Sync {
async fn create_management_token(
&self,
record: &CreateManagementTokenRecord,
) -> Result<StoredManagementToken, crate::DataLayerError>;
async fn update_management_token(
&self,
record: &UpdateManagementTokenRecord,
) -> Result<Option<StoredManagementToken>, crate::DataLayerError>;
async fn delete_management_token(&self, token_id: &str) -> Result<bool, crate::DataLayerError>;
async fn set_management_token_active(
&self,
token_id: &str,
is_active: bool,
) -> Result<Option<StoredManagementToken>, crate::DataLayerError>;
async fn regenerate_management_token_secret(
&self,
mutation: &RegenerateManagementTokenSecret,
) -> Result<Option<StoredManagementToken>, crate::DataLayerError>;
async fn record_management_token_usage(
&self,
token_id: &str,
last_used_ip: Option<&str>,
) -> Result<Option<StoredManagementToken>, crate::DataLayerError>;
}
@@ -0,0 +1,22 @@
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 proxy_nodes;
pub mod quota;
pub mod routing_profiles;
pub mod settlement;
pub mod usage;
pub mod users;
pub mod video_tasks;
pub mod wallet;
@@ -0,0 +1,235 @@
use async_trait::async_trait;
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredOAuthProviderConfig {
pub provider_type: String,
pub display_name: String,
pub client_id: String,
pub client_secret_encrypted: Option<String>,
pub authorization_url_override: Option<String>,
pub token_url_override: Option<String>,
pub userinfo_url_override: Option<String>,
pub scopes: Option<Vec<String>>,
pub redirect_uri: String,
pub frontend_callback_url: String,
pub attribute_mapping: Option<serde_json::Value>,
pub extra_config: Option<serde_json::Value>,
pub icon_url: Option<String>,
pub is_enabled: bool,
pub created_at_unix_ms: Option<u64>,
pub updated_at_unix_secs: Option<u64>,
}
impl StoredOAuthProviderConfig {
pub fn new(
provider_type: String,
display_name: String,
client_id: String,
redirect_uri: String,
frontend_callback_url: String,
) -> Result<Self, crate::DataLayerError> {
if provider_type.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"oauth_providers.provider_type is empty".to_string(),
));
}
if display_name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"oauth_providers.display_name is empty".to_string(),
));
}
if client_id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"oauth_providers.client_id is empty".to_string(),
));
}
if redirect_uri.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"oauth_providers.redirect_uri is empty".to_string(),
));
}
if frontend_callback_url.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"oauth_providers.frontend_callback_url is empty".to_string(),
));
}
Ok(Self {
provider_type,
display_name,
client_id,
client_secret_encrypted: None,
authorization_url_override: None,
token_url_override: None,
userinfo_url_override: None,
scopes: None,
redirect_uri,
frontend_callback_url,
attribute_mapping: None,
extra_config: None,
icon_url: None,
is_enabled: false,
created_at_unix_ms: None,
updated_at_unix_secs: None,
})
}
#[allow(clippy::too_many_arguments)]
pub fn with_config_fields(
mut self,
client_secret_encrypted: Option<String>,
authorization_url_override: Option<String>,
token_url_override: Option<String>,
userinfo_url_override: Option<String>,
scopes: Option<Vec<String>>,
attribute_mapping: Option<serde_json::Value>,
extra_config: Option<serde_json::Value>,
icon_url: Option<String>,
is_enabled: bool,
) -> Self {
self.client_secret_encrypted = client_secret_encrypted;
self.authorization_url_override = authorization_url_override;
self.token_url_override = token_url_override;
self.userinfo_url_override = userinfo_url_override;
self.scopes = scopes;
self.attribute_mapping = attribute_mapping;
self.extra_config = extra_config;
self.icon_url = icon_url;
self.is_enabled = is_enabled;
self
}
pub fn with_timestamps(
mut self,
created_at_unix_ms: Option<u64>,
updated_at_unix_secs: Option<u64>,
) -> Self {
self.created_at_unix_ms = created_at_unix_ms;
self.updated_at_unix_secs = updated_at_unix_secs;
self
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize, Default)]
pub enum EncryptedSecretUpdate {
#[default]
Preserve,
Clear,
Set(String),
}
impl EncryptedSecretUpdate {
pub fn mode_name(&self) -> &'static str {
match self {
Self::Preserve => "preserve",
Self::Clear => "clear",
Self::Set(_) => "set",
}
}
pub fn value(&self) -> Option<&str> {
match self {
Self::Set(value) => Some(value.as_str()),
Self::Preserve | Self::Clear => None,
}
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct UpsertOAuthProviderConfigRecord {
pub provider_type: String,
pub display_name: String,
pub client_id: String,
pub client_secret_encrypted: EncryptedSecretUpdate,
pub authorization_url_override: Option<String>,
pub token_url_override: Option<String>,
pub userinfo_url_override: Option<String>,
pub scopes: Option<Vec<String>>,
pub redirect_uri: String,
pub frontend_callback_url: String,
pub attribute_mapping: Option<serde_json::Value>,
pub extra_config: Option<serde_json::Value>,
pub icon_url: Option<String>,
pub is_enabled: bool,
}
impl UpsertOAuthProviderConfigRecord {
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
if self.provider_type.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"provider_type is required".to_string(),
));
}
if self.display_name.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"display_name is required".to_string(),
));
}
if self.client_id.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"client_id is required".to_string(),
));
}
if self.redirect_uri.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"redirect_uri is required".to_string(),
));
}
if self.frontend_callback_url.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"frontend_callback_url is required".to_string(),
));
}
if let Some(scopes) = &self.scopes {
for scope in scopes {
if scope.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"scopes must not contain empty values".to_string(),
));
}
}
}
Ok(())
}
}
#[async_trait]
pub trait OAuthProviderReadRepository: Send + Sync {
async fn list_oauth_provider_configs(
&self,
) -> Result<Vec<StoredOAuthProviderConfig>, crate::DataLayerError>;
async fn get_oauth_provider_config(
&self,
provider_type: &str,
) -> Result<Option<StoredOAuthProviderConfig>, crate::DataLayerError>;
async fn count_locked_users_if_provider_disabled(
&self,
provider_type: &str,
ldap_exclusive: bool,
) -> Result<usize, crate::DataLayerError>;
}
#[async_trait]
pub trait OAuthProviderWriteRepository: Send + Sync {
async fn upsert_oauth_provider_config(
&self,
record: &UpsertOAuthProviderConfigRecord,
) -> Result<StoredOAuthProviderConfig, crate::DataLayerError>;
async fn delete_oauth_provider_config(
&self,
provider_type: &str,
) -> Result<bool, crate::DataLayerError>;
}
pub trait OAuthProviderRepository:
OAuthProviderReadRepository + OAuthProviderWriteRepository + Send + Sync
{
}
impl<T> OAuthProviderRepository for T where
T: OAuthProviderReadRepository + OAuthProviderWriteRepository + Send + Sync
{
}
@@ -0,0 +1,48 @@
pub fn merge_score_reason_patch(
mut current: serde_json::Value,
patch: Option<serde_json::Value>,
) -> serde_json::Value {
let Some(patch) = patch else {
return current;
};
match (current.as_object_mut(), patch) {
(Some(current), serde_json::Value::Object(patch)) => {
for (key, value) in patch {
current.insert(key, value);
}
serde_json::Value::Object(current.clone())
}
(_, patch) => patch,
}
}
pub fn score_with_delta(score: f64, delta_basis_points: Option<i32>) -> f64 {
let delta = delta_basis_points.unwrap_or_default() as f64 / 10_000.0;
(score + delta).clamp(0.0, 1.0)
}
pub fn i64_from_u64(value: u64, field: &str) -> Result<i64, crate::DataLayerError> {
i64::try_from(value).map_err(|_| {
crate::DataLayerError::InvalidInput(format!("{field} exceeds signed 64-bit range"))
})
}
pub fn i64_opt_from_u64(
value: Option<u64>,
field: &str,
) -> Result<Option<i64>, crate::DataLayerError> {
value.map(|value| i64_from_u64(value, field)).transpose()
}
pub fn u64_from_i64(value: i64, field: &str) -> Result<u64, crate::DataLayerError> {
u64::try_from(value).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!("{field} is negative: {value}"))
})
}
pub fn u64_opt_from_i64(
value: Option<i64>,
field: &str,
) -> Result<Option<u64>, crate::DataLayerError> {
value.map(|value| u64_from_i64(value, field)).transpose()
}
@@ -0,0 +1,16 @@
mod helpers;
mod types;
pub use helpers::{
i64_from_u64, i64_opt_from_u64, merge_score_reason_patch, score_with_delta, u64_from_i64,
u64_opt_from_i64,
};
pub use types::{
GetPoolMemberScoresByIdsQuery, ListPoolMemberProbeCandidatesQuery, ListPoolMemberScoresQuery,
ListRankedPoolMembersQuery, PoolMemberHardState, PoolMemberIdentity, PoolMemberProbeAttempt,
PoolMemberProbeResult, PoolMemberProbeStatus, PoolMemberScheduleFeedback,
PoolMemberScoreRepository, PoolMemberScoreWriteRepository, PoolScoreReadRepository,
PoolScoreScope, StoredPoolMemberScore, UpsertPoolMemberScore, POOL_KIND_PROVIDER_KEY_POOL,
POOL_MEMBER_KIND_PROVIDER_API_KEY, POOL_SCORE_CAPABILITY_ACCOUNT,
POOL_SCORE_CAPABILITY_API_FORMAT, POOL_SCORE_SCOPE_KIND_ACCOUNT, POOL_SCORE_SCOPE_KIND_MODEL,
};
@@ -0,0 +1,361 @@
use async_trait::async_trait;
pub const POOL_KIND_PROVIDER_KEY_POOL: &str = "provider_key_pool";
pub const POOL_MEMBER_KIND_PROVIDER_API_KEY: &str = "provider_api_key";
pub const POOL_SCORE_CAPABILITY_ACCOUNT: &str = "account";
pub const POOL_SCORE_SCOPE_KIND_ACCOUNT: &str = "account";
pub const POOL_SCORE_CAPABILITY_API_FORMAT: &str = POOL_SCORE_CAPABILITY_ACCOUNT;
pub const POOL_SCORE_SCOPE_KIND_MODEL: &str = POOL_SCORE_SCOPE_KIND_ACCOUNT;
#[derive(
Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize,
)]
#[serde(rename_all = "snake_case")]
pub enum PoolMemberHardState {
Available,
Unknown,
Cooldown,
QuotaExhausted,
AuthInvalid,
Banned,
Inactive,
}
impl PoolMemberHardState {
pub fn as_database(self) -> &'static str {
match self {
Self::Available => "available",
Self::Unknown => "unknown",
Self::Cooldown => "cooldown",
Self::QuotaExhausted => "quota_exhausted",
Self::AuthInvalid => "auth_invalid",
Self::Banned => "banned",
Self::Inactive => "inactive",
}
}
pub fn from_database(value: &str) -> Result<Self, crate::DataLayerError> {
match value.trim().to_ascii_lowercase().as_str() {
"available" => Ok(Self::Available),
"unknown" => Ok(Self::Unknown),
"cooldown" => Ok(Self::Cooldown),
"quota_exhausted" => Ok(Self::QuotaExhausted),
"auth_invalid" => Ok(Self::AuthInvalid),
"banned" => Ok(Self::Banned),
"inactive" => Ok(Self::Inactive),
other => Err(crate::DataLayerError::UnexpectedValue(format!(
"unknown pool member hard_state: {other}"
))),
}
}
pub fn schedulable(self) -> bool {
matches!(self, Self::Available | Self::Unknown)
}
}
#[derive(
Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize,
)]
#[serde(rename_all = "snake_case")]
pub enum PoolMemberProbeStatus {
Never,
Ok,
Failed,
Stale,
InProgress,
}
impl PoolMemberProbeStatus {
pub fn as_database(self) -> &'static str {
match self {
Self::Never => "never",
Self::Ok => "ok",
Self::Failed => "failed",
Self::Stale => "stale",
Self::InProgress => "in_progress",
}
}
pub fn from_database(value: &str) -> Result<Self, crate::DataLayerError> {
match value.trim().to_ascii_lowercase().as_str() {
"never" => Ok(Self::Never),
"ok" => Ok(Self::Ok),
"failed" => Ok(Self::Failed),
"stale" => Ok(Self::Stale),
"in_progress" => Ok(Self::InProgress),
other => Err(crate::DataLayerError::UnexpectedValue(format!(
"unknown pool member probe_status: {other}"
))),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct PoolScoreScope {
pub capability: String,
pub scope_kind: String,
pub scope_id: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct PoolMemberIdentity {
pub pool_kind: String,
pub pool_id: String,
pub member_kind: String,
pub member_id: String,
}
impl PoolMemberIdentity {
pub fn provider_api_key(provider_id: impl Into<String>, key_id: impl Into<String>) -> Self {
Self {
pool_kind: POOL_KIND_PROVIDER_KEY_POOL.to_string(),
pool_id: provider_id.into(),
member_kind: POOL_MEMBER_KIND_PROVIDER_API_KEY.to_string(),
member_id: key_id.into(),
}
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredPoolMemberScore {
pub id: String,
pub pool_kind: String,
pub pool_id: String,
pub member_kind: String,
pub member_id: String,
pub capability: String,
pub scope_kind: String,
pub scope_id: Option<String>,
pub score: f64,
pub hard_state: PoolMemberHardState,
pub score_version: u64,
pub score_reason: serde_json::Value,
pub last_ranked_at: Option<u64>,
pub last_scheduled_at: Option<u64>,
pub last_success_at: Option<u64>,
pub last_failure_at: Option<u64>,
pub failure_count: u64,
pub last_probe_attempt_at: Option<u64>,
pub last_probe_success_at: Option<u64>,
pub last_probe_failure_at: Option<u64>,
pub probe_failure_count: u64,
pub probe_status: PoolMemberProbeStatus,
pub updated_at: u64,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct UpsertPoolMemberScore {
pub id: String,
pub identity: PoolMemberIdentity,
pub scope: PoolScoreScope,
pub score: f64,
pub hard_state: PoolMemberHardState,
pub score_version: u64,
pub score_reason: serde_json::Value,
pub last_ranked_at: Option<u64>,
pub last_scheduled_at: Option<u64>,
pub last_success_at: Option<u64>,
pub last_failure_at: Option<u64>,
pub failure_count: u64,
pub last_probe_attempt_at: Option<u64>,
pub last_probe_success_at: Option<u64>,
pub last_probe_failure_at: Option<u64>,
pub probe_failure_count: u64,
pub probe_status: PoolMemberProbeStatus,
pub updated_at: u64,
}
impl UpsertPoolMemberScore {
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
validate_non_empty(&self.id, "pool_member_scores.id")?;
validate_non_empty(&self.identity.pool_kind, "pool_member_scores.pool_kind")?;
validate_non_empty(&self.identity.pool_id, "pool_member_scores.pool_id")?;
validate_non_empty(&self.identity.member_kind, "pool_member_scores.member_kind")?;
validate_non_empty(&self.identity.member_id, "pool_member_scores.member_id")?;
validate_non_empty(&self.scope.capability, "pool_member_scores.capability")?;
validate_non_empty(&self.scope.scope_kind, "pool_member_scores.scope_kind")?;
if !self.score.is_finite() {
return Err(crate::DataLayerError::InvalidInput(
"pool_member_scores.score must be finite".to_string(),
));
}
Ok(())
}
pub fn into_stored(self) -> StoredPoolMemberScore {
StoredPoolMemberScore {
id: self.id,
pool_kind: self.identity.pool_kind,
pool_id: self.identity.pool_id,
member_kind: self.identity.member_kind,
member_id: self.identity.member_id,
capability: self.scope.capability,
scope_kind: self.scope.scope_kind,
scope_id: self.scope.scope_id,
score: self.score,
hard_state: self.hard_state,
score_version: self.score_version,
score_reason: self.score_reason,
last_ranked_at: self.last_ranked_at,
last_scheduled_at: self.last_scheduled_at,
last_success_at: self.last_success_at,
last_failure_at: self.last_failure_at,
failure_count: self.failure_count,
last_probe_attempt_at: self.last_probe_attempt_at,
last_probe_success_at: self.last_probe_success_at,
last_probe_failure_at: self.last_probe_failure_at,
probe_failure_count: self.probe_failure_count,
probe_status: self.probe_status,
updated_at: self.updated_at,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct ListRankedPoolMembersQuery {
pub pool_kind: String,
pub pool_id: String,
pub capability: String,
pub scope_kind: String,
pub scope_id: Option<String>,
pub hard_states: Vec<PoolMemberHardState>,
pub probe_statuses: Option<Vec<PoolMemberProbeStatus>>,
pub offset: usize,
pub limit: usize,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct ListPoolMemberScoresQuery {
pub pool_kind: String,
pub pool_id: String,
pub capability: Option<String>,
pub scope_kind: Option<String>,
pub scope_id: Option<String>,
pub hard_states: Vec<PoolMemberHardState>,
pub probe_statuses: Option<Vec<PoolMemberProbeStatus>>,
pub offset: usize,
pub limit: usize,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct ListPoolMemberProbeCandidatesQuery {
pub pool_kind: String,
pub pool_id: String,
pub capability: Option<String>,
pub stale_before_unix_secs: u64,
pub limit: usize,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct GetPoolMemberScoresByIdsQuery {
pub ids: Vec<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct PoolMemberProbeResult {
pub identity: PoolMemberIdentity,
pub scope: Option<PoolScoreScope>,
pub attempted_at: u64,
pub succeeded: bool,
pub hard_state: Option<PoolMemberHardState>,
pub probe_status: PoolMemberProbeStatus,
pub score_reason_patch: Option<serde_json::Value>,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct PoolMemberProbeAttempt {
pub identity: PoolMemberIdentity,
pub scope: Option<PoolScoreScope>,
pub attempted_at: u64,
pub score_reason_patch: Option<serde_json::Value>,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct PoolMemberScheduleFeedback {
pub identity: PoolMemberIdentity,
pub scope: Option<PoolScoreScope>,
pub scheduled_at: u64,
pub succeeded: Option<bool>,
pub hard_state: Option<PoolMemberHardState>,
pub score_delta: Option<i32>,
pub score_reason_patch: Option<serde_json::Value>,
}
#[async_trait]
pub trait PoolScoreReadRepository: Send + Sync {
async fn list_ranked_pool_members(
&self,
query: &ListRankedPoolMembersQuery,
) -> Result<Vec<StoredPoolMemberScore>, crate::DataLayerError>;
async fn list_pool_member_scores(
&self,
query: &ListPoolMemberScoresQuery,
) -> Result<Vec<StoredPoolMemberScore>, crate::DataLayerError>;
async fn list_pool_member_probe_candidates(
&self,
query: &ListPoolMemberProbeCandidatesQuery,
) -> Result<Vec<StoredPoolMemberScore>, crate::DataLayerError>;
async fn get_pool_member_scores_by_ids(
&self,
query: &GetPoolMemberScoresByIdsQuery,
) -> Result<Vec<StoredPoolMemberScore>, crate::DataLayerError>;
}
#[async_trait]
pub trait PoolMemberScoreWriteRepository: Send + Sync {
async fn upsert_pool_member_score(
&self,
score: UpsertPoolMemberScore,
) -> Result<StoredPoolMemberScore, crate::DataLayerError>;
async fn mark_pool_member_probe_in_progress(
&self,
attempt: PoolMemberProbeAttempt,
) -> Result<usize, crate::DataLayerError>;
async fn record_pool_member_probe_result(
&self,
result: PoolMemberProbeResult,
) -> Result<usize, crate::DataLayerError>;
async fn record_pool_member_schedule_feedback(
&self,
feedback: PoolMemberScheduleFeedback,
) -> Result<usize, crate::DataLayerError>;
async fn mark_pool_member_hard_state(
&self,
identity: &PoolMemberIdentity,
scope: Option<&PoolScoreScope>,
hard_state: PoolMemberHardState,
updated_at: u64,
) -> Result<usize, crate::DataLayerError>;
async fn delete_pool_member_scores_for_member(
&self,
identity: &PoolMemberIdentity,
) -> Result<usize, crate::DataLayerError>;
}
pub trait PoolMemberScoreRepository:
PoolScoreReadRepository + PoolMemberScoreWriteRepository + Send + Sync
{
}
impl<T> PoolMemberScoreRepository for T where
T: PoolScoreReadRepository + PoolMemberScoreWriteRepository + Send + Sync
{
}
fn validate_non_empty(value: &str, field: &str) -> Result<(), crate::DataLayerError> {
if value.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(format!(
"{field} is empty"
)));
}
Ok(())
}
@@ -0,0 +1,11 @@
mod snapshot;
mod types;
pub use snapshot::ProviderCatalogSnapshot;
pub use types::{
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, ProviderCatalogReadRepository,
ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
};
@@ -0,0 +1,253 @@
use std::cmp::Ordering;
use std::collections::BTreeMap;
use super::{
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, StoredProviderCatalogEndpoint,
StoredProviderCatalogKey, StoredProviderCatalogKeyMaintenanceSummary,
StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
};
use crate::DataLayerError;
/// Immutable provider-catalog view used by memory and driver adapters.
#[derive(Debug, Clone, Default)]
pub struct ProviderCatalogSnapshot {
providers: BTreeMap<String, StoredProviderCatalogProvider>,
endpoints: BTreeMap<String, StoredProviderCatalogEndpoint>,
keys: BTreeMap<String, StoredProviderCatalogKey>,
}
impl ProviderCatalogSnapshot {
pub fn new(
providers: Vec<StoredProviderCatalogProvider>,
endpoints: Vec<StoredProviderCatalogEndpoint>,
keys: Vec<StoredProviderCatalogKey>,
) -> Self {
Self {
providers: providers
.into_iter()
.map(|item| (item.id.clone(), item))
.collect(),
endpoints: endpoints
.into_iter()
.map(|item| (item.id.clone(), item))
.collect(),
keys: keys
.into_iter()
.map(|item| (item.id.clone(), item))
.collect(),
}
}
pub fn list_providers(&self, active_only: bool) -> Vec<StoredProviderCatalogProvider> {
let mut providers = self
.providers
.values()
.filter(|provider| !active_only || provider.is_active)
.cloned()
.collect::<Vec<_>>();
providers.sort_by(|left, right| {
left.provider_priority
.cmp(&right.provider_priority)
.then(left.name.cmp(&right.name))
});
providers
}
pub fn list_providers_by_ids(
&self,
provider_ids: &[String],
) -> Vec<StoredProviderCatalogProvider> {
provider_ids
.iter()
.filter_map(|id| self.providers.get(id).cloned())
.collect()
}
pub fn list_endpoints_by_ids(
&self,
endpoint_ids: &[String],
) -> Vec<StoredProviderCatalogEndpoint> {
endpoint_ids
.iter()
.filter_map(|id| self.endpoints.get(id).cloned())
.collect()
}
pub fn list_endpoints_by_provider_ids(
&self,
provider_ids: &[String],
) -> Vec<StoredProviderCatalogEndpoint> {
let mut endpoints = self
.endpoints
.values()
.filter(|endpoint| provider_ids.contains(&endpoint.provider_id))
.cloned()
.collect::<Vec<_>>();
endpoints.sort_by(|left, right| {
left.provider_id
.cmp(&right.provider_id)
.then(left.api_format.cmp(&right.api_format))
.then(left.id.cmp(&right.id))
});
endpoints
}
pub fn list_keys_by_ids(&self, key_ids: &[String]) -> Vec<StoredProviderCatalogKey> {
key_ids
.iter()
.filter_map(|id| self.keys.get(id).cloned())
.collect()
}
pub fn list_keys_by_provider_ids(
&self,
provider_ids: &[String],
) -> Vec<StoredProviderCatalogKey> {
let mut keys = self
.keys
.values()
.filter(|key| provider_ids.contains(&key.provider_id))
.cloned()
.collect::<Vec<_>>();
keys.sort_by(|left, right| {
left.provider_id
.cmp(&right.provider_id)
.then(left.name.cmp(&right.name))
.then(left.id.cmp(&right.id))
});
keys
}
pub fn list_key_maintenance_summaries_by_provider_ids(
&self,
provider_ids: &[String],
) -> Vec<StoredProviderCatalogKeyMaintenanceSummary> {
let mut keys = self
.keys
.values()
.filter(|key| provider_ids.contains(&key.provider_id))
.map(|key| StoredProviderCatalogKeyMaintenanceSummary {
id: key.id.clone(),
provider_id: key.provider_id.clone(),
is_active: key.is_active,
upstream_metadata: key.upstream_metadata.clone(),
})
.collect::<Vec<_>>();
keys.sort_by(|left, right| {
left.provider_id
.cmp(&right.provider_id)
.then(left.id.cmp(&right.id))
});
keys
}
pub fn list_keys_page(
&self,
query: &ProviderCatalogKeyListQuery,
) -> StoredProviderCatalogKeyPage {
let mut keys = self
.keys
.values()
.filter(|key| key.provider_id == query.provider_id)
.filter(|key| {
query.search.as_ref().is_none_or(|keyword| {
let keyword = keyword.trim().to_ascii_lowercase();
keyword.is_empty()
|| key.name.to_ascii_lowercase().contains(&keyword)
|| key.id.to_ascii_lowercase().contains(&keyword)
})
})
.filter(|key| query.is_active.is_none_or(|value| key.is_active == value))
.cloned()
.collect::<Vec<_>>();
sort_key_page(&mut keys, query.order.clone());
let total = keys.len();
let items = keys
.into_iter()
.skip(query.offset)
.take(query.limit)
.collect();
StoredProviderCatalogKeyPage { items, total }
}
pub fn list_key_stats_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogKeyStats>, DataLayerError> {
let mut stats = provider_ids
.iter()
.map(|provider_id| {
let total_keys = self
.keys
.values()
.filter(|key| &key.provider_id == provider_id)
.count() as i64;
let active_keys = self
.keys
.values()
.filter(|key| &key.provider_id == provider_id && key.is_active)
.count() as i64;
StoredProviderCatalogKeyStats::new(provider_id.clone(), total_keys, active_keys)
})
.collect::<Result<Vec<_>, _>>()?;
stats.retain(|item| item.total_keys > 0);
Ok(stats)
}
}
fn compare_optional_u64_null_last(
left: Option<u64>,
right: Option<u64>,
descending: bool,
) -> Ordering {
match (left, right) {
(Some(left), Some(right)) if descending => right.cmp(&left),
(Some(left), Some(right)) => left.cmp(&right),
(Some(_), None) => Ordering::Less,
(None, Some(_)) => Ordering::Greater,
(None, None) => Ordering::Equal,
}
}
fn sort_key_page(items: &mut [StoredProviderCatalogKey], order: ProviderCatalogKeyListOrder) {
items.sort_by(|left, right| match order {
ProviderCatalogKeyListOrder::Name => left
.internal_priority
.cmp(&right.internal_priority)
.then(left.name.cmp(&right.name))
.then(left.id.cmp(&right.id)),
ProviderCatalogKeyListOrder::CreatedAt => left
.internal_priority
.cmp(&right.internal_priority)
.then(
left.created_at_unix_ms
.unwrap_or_default()
.cmp(&right.created_at_unix_ms.unwrap_or_default()),
)
.then(left.id.cmp(&right.id)),
ProviderCatalogKeyListOrder::CreatedAtAsc => {
compare_optional_u64_null_last(left.created_at_unix_ms, right.created_at_unix_ms, false)
.then(left.name.cmp(&right.name))
.then(left.id.cmp(&right.id))
}
ProviderCatalogKeyListOrder::CreatedAtDesc => {
compare_optional_u64_null_last(left.created_at_unix_ms, right.created_at_unix_ms, true)
.then(left.name.cmp(&right.name))
.then(left.id.cmp(&right.id))
}
ProviderCatalogKeyListOrder::LastUsedAtAsc => compare_optional_u64_null_last(
left.last_used_at_unix_secs,
right.last_used_at_unix_secs,
false,
)
.then(left.name.cmp(&right.name))
.then(left.id.cmp(&right.id)),
ProviderCatalogKeyListOrder::LastUsedAtDesc => compare_optional_u64_null_last(
left.last_used_at_unix_secs,
right.last_used_at_unix_secs,
true,
)
.then(left.name.cmp(&right.name))
.then(left.id.cmp(&right.id)),
});
}
@@ -0,0 +1,804 @@
use async_trait::async_trait;
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct ProviderCatalogUpstreamMetadataNamespaceUpdate {
pub namespace: String,
pub value: serde_json::Value,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredProviderCatalogProvider {
pub id: String,
pub name: String,
pub description: Option<String>,
pub website: Option<String>,
pub provider_type: String,
pub billing_type: Option<String>,
pub monthly_quota_usd: Option<f64>,
pub monthly_used_usd: Option<f64>,
pub quota_reset_day: Option<u64>,
pub quota_last_reset_at_unix_secs: Option<u64>,
pub quota_expires_at_unix_secs: Option<u64>,
pub provider_priority: i32,
pub is_active: bool,
pub keep_priority_on_conversion: bool,
pub enable_format_conversion: bool,
pub concurrent_limit: Option<i32>,
pub max_retries: Option<i32>,
pub proxy: Option<serde_json::Value>,
pub request_timeout_secs: Option<f64>,
pub stream_first_byte_timeout_secs: Option<f64>,
pub config: Option<serde_json::Value>,
pub created_at_unix_ms: Option<u64>,
pub updated_at_unix_secs: Option<u64>,
}
impl StoredProviderCatalogProvider {
pub fn new(
id: String,
name: String,
website: Option<String>,
provider_type: String,
) -> Result<Self, crate::DataLayerError> {
if name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"providers.name is empty".to_string(),
));
}
if provider_type.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"providers.provider_type is empty".to_string(),
));
}
Ok(Self {
id,
name,
description: None,
website,
provider_type,
billing_type: None,
monthly_quota_usd: None,
monthly_used_usd: None,
quota_reset_day: None,
quota_last_reset_at_unix_secs: None,
quota_expires_at_unix_secs: None,
provider_priority: 0,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
created_at_unix_ms: None,
updated_at_unix_secs: None,
})
}
#[allow(clippy::too_many_arguments)]
pub fn with_transport_fields(
mut self,
is_active: bool,
keep_priority_on_conversion: bool,
enable_format_conversion: bool,
concurrent_limit: Option<i32>,
max_retries: Option<i32>,
proxy: Option<serde_json::Value>,
request_timeout_secs: Option<f64>,
stream_first_byte_timeout_secs: Option<f64>,
config: Option<serde_json::Value>,
) -> Self {
self.is_active = is_active;
self.keep_priority_on_conversion = keep_priority_on_conversion;
self.enable_format_conversion = enable_format_conversion;
self.concurrent_limit = concurrent_limit;
self.max_retries = max_retries;
self.proxy = proxy;
self.request_timeout_secs = request_timeout_secs;
self.stream_first_byte_timeout_secs = stream_first_byte_timeout_secs;
self.config = config;
self
}
pub fn with_description(mut self, description: Option<String>) -> Self {
self.description = description;
self
}
#[allow(clippy::too_many_arguments)]
pub fn with_billing_fields(
mut self,
billing_type: Option<String>,
monthly_quota_usd: Option<f64>,
monthly_used_usd: Option<f64>,
quota_reset_day: Option<u64>,
quota_last_reset_at_unix_secs: Option<u64>,
quota_expires_at_unix_secs: Option<u64>,
) -> Self {
self.billing_type = billing_type;
self.monthly_quota_usd = monthly_quota_usd;
self.monthly_used_usd = monthly_used_usd;
self.quota_reset_day = quota_reset_day;
self.quota_last_reset_at_unix_secs = quota_last_reset_at_unix_secs;
self.quota_expires_at_unix_secs = quota_expires_at_unix_secs;
self
}
pub fn with_routing_fields(mut self, provider_priority: i32) -> Self {
self.provider_priority = provider_priority;
self
}
pub fn with_timestamps(
mut self,
created_at_unix_ms: Option<u64>,
updated_at_unix_secs: Option<u64>,
) -> Self {
self.created_at_unix_ms = created_at_unix_ms;
self.updated_at_unix_secs = updated_at_unix_secs;
self
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredProviderCatalogEndpoint {
pub id: String,
pub provider_id: String,
pub api_format: String,
pub api_family: Option<String>,
pub endpoint_kind: Option<String>,
pub is_active: bool,
pub health_score: f64,
pub base_url: String,
pub header_rules: Option<serde_json::Value>,
pub body_rules: Option<serde_json::Value>,
pub max_retries: Option<i32>,
pub custom_path: Option<String>,
pub config: Option<serde_json::Value>,
pub format_acceptance_config: Option<serde_json::Value>,
pub proxy: Option<serde_json::Value>,
pub created_at_unix_ms: Option<u64>,
pub updated_at_unix_secs: Option<u64>,
}
impl StoredProviderCatalogEndpoint {
pub fn new(
id: String,
provider_id: String,
api_format: String,
api_family: Option<String>,
endpoint_kind: Option<String>,
is_active: bool,
) -> Result<Self, crate::DataLayerError> {
if api_format.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"provider_endpoints.api_format is empty".to_string(),
));
}
Ok(Self {
id,
provider_id,
api_format,
api_family,
endpoint_kind,
is_active,
health_score: 1.0,
base_url: String::new(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
created_at_unix_ms: None,
updated_at_unix_secs: None,
})
}
#[allow(clippy::too_many_arguments)]
pub fn with_transport_fields(
mut self,
base_url: String,
header_rules: Option<serde_json::Value>,
body_rules: Option<serde_json::Value>,
max_retries: Option<i32>,
custom_path: Option<String>,
config: Option<serde_json::Value>,
format_acceptance_config: Option<serde_json::Value>,
proxy: Option<serde_json::Value>,
) -> Result<Self, crate::DataLayerError> {
if base_url.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"provider_endpoints.base_url is empty".to_string(),
));
}
self.base_url = base_url;
self.header_rules = header_rules;
self.body_rules = body_rules;
self.max_retries = max_retries;
self.custom_path = custom_path;
self.config = config;
self.format_acceptance_config = format_acceptance_config;
self.proxy = proxy;
Ok(self)
}
pub fn with_health_score(mut self, health_score: f64) -> Self {
self.health_score = health_score;
self
}
pub fn with_timestamps(
mut self,
created_at_unix_ms: Option<u64>,
updated_at_unix_secs: Option<u64>,
) -> Self {
self.created_at_unix_ms = created_at_unix_ms;
self.updated_at_unix_secs = updated_at_unix_secs;
self
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredProviderCatalogKey {
pub id: String,
pub provider_id: String,
pub name: String,
pub auth_type: String,
pub capabilities: Option<serde_json::Value>,
pub is_active: bool,
pub api_formats: Option<serde_json::Value>,
pub auth_type_by_format: Option<serde_json::Value>,
pub allow_auth_channel_mismatch_formats: Option<serde_json::Value>,
pub encrypted_api_key: Option<String>,
pub encrypted_auth_config: Option<String>,
pub note: Option<String>,
pub internal_priority: i32,
pub rate_multipliers: Option<serde_json::Value>,
pub global_priority_by_format: Option<serde_json::Value>,
pub allowed_models: Option<serde_json::Value>,
pub expires_at_unix_secs: Option<u64>,
pub cache_ttl_minutes: i32,
pub max_probe_interval_minutes: i32,
pub proxy: Option<serde_json::Value>,
pub fingerprint: Option<serde_json::Value>,
pub rpm_limit: Option<u32>,
pub concurrent_limit: Option<i32>,
pub learned_rpm_limit: Option<u32>,
pub concurrent_429_count: Option<u32>,
pub rpm_429_count: Option<u32>,
pub last_429_at_unix_secs: Option<u64>,
pub last_429_type: Option<String>,
pub adjustment_history: Option<serde_json::Value>,
pub utilization_samples: Option<serde_json::Value>,
pub last_probe_increase_at_unix_secs: Option<u64>,
pub last_rpm_peak: Option<u32>,
pub request_count: Option<u32>,
pub total_tokens: u64,
pub total_cost_usd: f64,
pub success_count: Option<u32>,
pub error_count: Option<u32>,
pub total_response_time_ms: Option<u64>,
pub last_used_at_unix_secs: Option<u64>,
pub auto_fetch_models: bool,
pub last_models_fetch_at_unix_secs: Option<u64>,
pub last_models_fetch_error: Option<String>,
pub locked_models: Option<serde_json::Value>,
pub model_include_patterns: Option<serde_json::Value>,
pub model_exclude_patterns: Option<serde_json::Value>,
pub upstream_metadata: Option<serde_json::Value>,
pub oauth_invalid_at_unix_secs: Option<u64>,
pub oauth_invalid_reason: Option<String>,
pub status_snapshot: Option<serde_json::Value>,
pub created_at_unix_ms: Option<u64>,
pub updated_at_unix_secs: Option<u64>,
pub health_by_format: Option<serde_json::Value>,
pub circuit_breaker_by_format: Option<serde_json::Value>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct StoredProviderCatalogKeyMaintenanceSummary {
pub id: String,
pub provider_id: String,
pub is_active: bool,
pub upstream_metadata: Option<serde_json::Value>,
}
impl StoredProviderCatalogKey {
pub fn new(
id: String,
provider_id: String,
name: String,
auth_type: String,
capabilities: Option<serde_json::Value>,
is_active: bool,
) -> Result<Self, crate::DataLayerError> {
if name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"provider_api_keys.name is empty".to_string(),
));
}
if auth_type.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"provider_api_keys.auth_type is empty".to_string(),
));
}
Ok(Self {
id,
provider_id,
name,
auth_type,
capabilities,
is_active,
api_formats: None,
auth_type_by_format: None,
allow_auth_channel_mismatch_formats: None,
encrypted_api_key: None,
encrypted_auth_config: None,
note: None,
internal_priority: 50,
rate_multipliers: None,
global_priority_by_format: None,
allowed_models: None,
expires_at_unix_secs: None,
cache_ttl_minutes: 5,
max_probe_interval_minutes: 32,
proxy: None,
fingerprint: None,
rpm_limit: None,
concurrent_limit: None,
learned_rpm_limit: None,
concurrent_429_count: None,
rpm_429_count: None,
last_429_at_unix_secs: None,
last_429_type: None,
adjustment_history: None,
utilization_samples: None,
last_probe_increase_at_unix_secs: None,
last_rpm_peak: None,
request_count: None,
total_tokens: 0,
total_cost_usd: 0.0,
success_count: None,
error_count: None,
total_response_time_ms: None,
last_used_at_unix_secs: None,
auto_fetch_models: false,
last_models_fetch_at_unix_secs: None,
last_models_fetch_error: None,
locked_models: None,
model_include_patterns: None,
model_exclude_patterns: None,
upstream_metadata: None,
oauth_invalid_at_unix_secs: None,
oauth_invalid_reason: None,
status_snapshot: None,
created_at_unix_ms: None,
updated_at_unix_secs: None,
health_by_format: None,
circuit_breaker_by_format: None,
})
}
#[allow(clippy::too_many_arguments)]
pub fn with_transport_fields(
mut self,
api_formats: Option<serde_json::Value>,
encrypted_api_key: impl Into<Option<String>>,
encrypted_auth_config: Option<String>,
rate_multipliers: Option<serde_json::Value>,
global_priority_by_format: Option<serde_json::Value>,
allowed_models: Option<serde_json::Value>,
expires_at_unix_secs: Option<u64>,
proxy: Option<serde_json::Value>,
fingerprint: Option<serde_json::Value>,
) -> Result<Self, crate::DataLayerError> {
let encrypted_api_key = encrypted_api_key.into();
if encrypted_api_key
.as_deref()
.is_some_and(|value| value.trim().is_empty())
{
return Err(crate::DataLayerError::UnexpectedValue(
"provider_api_keys.api_key is empty".to_string(),
));
}
self.api_formats = api_formats;
self.encrypted_api_key = encrypted_api_key;
self.encrypted_auth_config = encrypted_auth_config;
self.rate_multipliers = rate_multipliers;
self.global_priority_by_format = global_priority_by_format;
self.allowed_models = allowed_models;
self.expires_at_unix_secs = expires_at_unix_secs;
self.proxy = proxy;
self.fingerprint = fingerprint;
Ok(self)
}
#[allow(clippy::too_many_arguments)]
pub fn with_rate_limit_fields(
mut self,
rpm_limit: Option<u32>,
concurrent_limit: Option<i32>,
learned_rpm_limit: Option<u32>,
concurrent_429_count: Option<u32>,
rpm_429_count: Option<u32>,
last_429_at_unix_secs: Option<u64>,
adjustment_history: Option<serde_json::Value>,
request_count: Option<u32>,
success_count: Option<u32>,
) -> Self {
self.rpm_limit = rpm_limit;
self.concurrent_limit = concurrent_limit;
self.learned_rpm_limit = learned_rpm_limit;
self.concurrent_429_count = concurrent_429_count;
self.rpm_429_count = rpm_429_count;
self.last_429_at_unix_secs = last_429_at_unix_secs;
self.adjustment_history = adjustment_history;
self.request_count = request_count;
self.success_count = success_count;
self
}
pub fn with_usage_fields(
mut self,
error_count: Option<u32>,
total_response_time_ms: Option<u64>,
) -> Self {
self.error_count = error_count;
self.total_response_time_ms = total_response_time_ms;
self
}
pub fn with_usage_totals(mut self, total_tokens: u64, total_cost_usd: f64) -> Self {
self.total_tokens = total_tokens;
self.total_cost_usd = if total_cost_usd.is_finite() {
total_cost_usd
} else {
0.0
};
self
}
pub fn with_health_fields(
mut self,
health_by_format: Option<serde_json::Value>,
circuit_breaker_by_format: Option<serde_json::Value>,
) -> Self {
self.health_by_format = health_by_format;
self.circuit_breaker_by_format = circuit_breaker_by_format;
self
}
}
#[cfg(test)]
mod transport_tests {
use super::StoredProviderCatalogKey;
#[test]
fn provider_catalog_key_defaults_concurrent_limit_to_none() {
let key = StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"default".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("key should build");
assert_eq!(key.concurrent_limit, None);
}
#[test]
fn provider_catalog_key_rate_limit_builder_sets_concurrent_limit() {
let key = StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"default".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("key should build")
.with_rate_limit_fields(
Some(120),
Some(3),
None,
None,
None,
None,
None,
None,
None,
);
assert_eq!(key.rpm_limit, Some(120));
assert_eq!(key.concurrent_limit, Some(3));
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub enum ProviderCatalogKeyListOrder {
#[default]
Name,
CreatedAt,
CreatedAtAsc,
CreatedAtDesc,
LastUsedAtAsc,
LastUsedAtDesc,
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct ProviderCatalogKeyListQuery {
pub provider_id: String,
pub search: Option<String>,
pub is_active: Option<bool>,
pub offset: usize,
pub limit: usize,
pub order: ProviderCatalogKeyListOrder,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredProviderCatalogKeyPage {
pub items: Vec<StoredProviderCatalogKey>,
pub total: usize,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredProviderCatalogKeyStats {
pub provider_id: String,
pub total_keys: u64,
pub active_keys: u64,
}
impl StoredProviderCatalogKeyStats {
pub fn new(
provider_id: String,
total_keys: i64,
active_keys: i64,
) -> Result<Self, crate::DataLayerError> {
if provider_id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"provider key stats provider_id is empty".to_string(),
));
}
if total_keys < 0 || active_keys < 0 {
return Err(crate::DataLayerError::UnexpectedValue(
"provider key stats count is negative".to_string(),
));
}
Ok(Self {
provider_id,
total_keys: total_keys as u64,
active_keys: active_keys as u64,
})
}
}
#[async_trait]
pub trait ProviderCatalogReadRepository: Send + Sync {
fn clear_local_cache(&self) {}
async fn list_providers(
&self,
active_only: bool,
) -> Result<Vec<StoredProviderCatalogProvider>, crate::DataLayerError>;
async fn list_providers_by_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogProvider>, crate::DataLayerError>;
async fn list_endpoints_by_ids(
&self,
endpoint_ids: &[String],
) -> Result<Vec<StoredProviderCatalogEndpoint>, crate::DataLayerError>;
async fn list_endpoints_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogEndpoint>, crate::DataLayerError>;
async fn list_keys_by_ids(
&self,
key_ids: &[String],
) -> Result<Vec<StoredProviderCatalogKey>, crate::DataLayerError>;
async fn list_keys_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogKey>, crate::DataLayerError>;
async fn list_key_summaries_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogKey>, crate::DataLayerError>;
async fn list_key_maintenance_summaries_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogKeyMaintenanceSummary>, crate::DataLayerError>;
async fn list_keys_page(
&self,
query: &ProviderCatalogKeyListQuery,
) -> Result<StoredProviderCatalogKeyPage, crate::DataLayerError>;
async fn list_key_stats_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogKeyStats>, crate::DataLayerError>;
}
#[async_trait]
pub trait ProviderCatalogWriteRepository: Send + Sync {
async fn create_provider(
&self,
provider: &StoredProviderCatalogProvider,
shift_existing_priorities_from: Option<i32>,
) -> Result<StoredProviderCatalogProvider, crate::DataLayerError>;
async fn update_provider(
&self,
provider: &StoredProviderCatalogProvider,
) -> Result<StoredProviderCatalogProvider, crate::DataLayerError>;
async fn delete_provider(&self, provider_id: &str) -> Result<bool, crate::DataLayerError>;
async fn cleanup_deleted_provider_refs(
&self,
provider_id: &str,
provider_deleted: bool,
endpoint_ids: &[String],
key_ids: &[String],
) -> Result<(), crate::DataLayerError>;
async fn create_endpoint(
&self,
endpoint: &StoredProviderCatalogEndpoint,
) -> Result<StoredProviderCatalogEndpoint, crate::DataLayerError>;
async fn update_endpoint(
&self,
endpoint: &StoredProviderCatalogEndpoint,
) -> Result<StoredProviderCatalogEndpoint, crate::DataLayerError>;
async fn delete_endpoint(&self, endpoint_id: &str) -> Result<bool, crate::DataLayerError>;
async fn create_key(
&self,
key: &StoredProviderCatalogKey,
) -> Result<StoredProviderCatalogKey, crate::DataLayerError>;
async fn update_key(
&self,
key: &StoredProviderCatalogKey,
) -> Result<StoredProviderCatalogKey, crate::DataLayerError>;
async fn update_keys(
&self,
keys: &[StoredProviderCatalogKey],
) -> Result<Vec<StoredProviderCatalogKey>, crate::DataLayerError>;
async fn update_key_upstream_metadata(
&self,
key_id: &str,
upstream_metadata: Option<&serde_json::Value>,
updated_at_unix_secs: Option<u64>,
) -> Result<bool, crate::DataLayerError>;
async fn upsert_key_upstream_metadata_namespace(
&self,
key_id: &str,
namespace: &str,
value: &serde_json::Value,
updated_at_unix_secs: Option<u64>,
) -> Result<bool, crate::DataLayerError>;
async fn update_key_model_fetch_state(
&self,
key_id: &str,
allowed_models: Option<&serde_json::Value>,
last_models_fetch_at_unix_secs: Option<u64>,
last_models_fetch_error: Option<&str>,
updated_at_unix_secs: Option<u64>,
) -> Result<bool, crate::DataLayerError>;
async fn update_key_model_fetch_success(
&self,
key_id: &str,
allowed_models: Option<&serde_json::Value>,
last_models_fetch_at_unix_secs: u64,
upstream_metadata_updates: &[ProviderCatalogUpstreamMetadataNamespaceUpdate],
updated_at_unix_secs: Option<u64>,
) -> Result<bool, crate::DataLayerError>;
async fn delete_key(&self, key_id: &str) -> Result<bool, crate::DataLayerError>;
async fn clear_key_oauth_invalid_marker(
&self,
key_id: &str,
) -> Result<bool, crate::DataLayerError>;
async fn update_key_oauth_credentials(
&self,
key_id: &str,
encrypted_api_key: &str,
encrypted_auth_config: Option<&str>,
expires_at_unix_secs: Option<u64>,
) -> Result<bool, crate::DataLayerError>;
async fn update_key_health_state(
&self,
key_id: &str,
is_active: bool,
health_by_format: Option<&serde_json::Value>,
circuit_breaker_by_format: Option<&serde_json::Value>,
) -> Result<bool, crate::DataLayerError>;
}
#[cfg(test)]
mod tests {
use super::StoredProviderCatalogKey;
fn sample_key() -> StoredProviderCatalogKey {
StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"key".to_string(),
"service_account".to_string(),
None,
true,
)
.expect("key should build")
}
#[test]
fn transport_fields_allow_null_encrypted_api_key() {
let key = sample_key()
.with_transport_fields(
None,
None::<String>,
None,
None,
None,
None,
None,
None,
None,
)
.expect("null api key should be accepted");
assert_eq!(key.encrypted_api_key, None);
}
#[test]
fn transport_fields_reject_empty_encrypted_api_key_string() {
let err = sample_key()
.with_transport_fields(
None,
Some(" ".to_string()),
None,
None,
None,
None,
None,
None,
None,
)
.expect_err("empty api key string should be rejected");
assert!(err
.to_string()
.contains("provider_api_keys.api_key is empty"));
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,6 @@
mod types;
pub use types::{
ProviderQuotaReadRepository, ProviderQuotaRepository, ProviderQuotaWriteRepository,
StoredProviderQuotaSnapshot,
};
@@ -0,0 +1,76 @@
use async_trait::async_trait;
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredProviderQuotaSnapshot {
pub provider_id: String,
pub billing_type: String,
pub monthly_quota_usd: Option<f64>,
pub monthly_used_usd: f64,
pub quota_reset_day: Option<u64>,
pub quota_last_reset_at_unix_secs: Option<u64>,
pub quota_expires_at_unix_secs: Option<u64>,
pub is_active: bool,
}
impl StoredProviderQuotaSnapshot {
#[allow(clippy::too_many_arguments)]
pub fn new(
provider_id: String,
billing_type: String,
monthly_quota_usd: Option<f64>,
monthly_used_usd: f64,
quota_reset_day: Option<i32>,
quota_last_reset_at_unix_secs: Option<i64>,
quota_expires_at_unix_secs: Option<i64>,
is_active: bool,
) -> Result<Self, crate::DataLayerError> {
if provider_id.trim().is_empty() || billing_type.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"provider quota identity is empty".to_string(),
));
}
if !monthly_used_usd.is_finite() || monthly_quota_usd.is_some_and(|v| !v.is_finite()) {
return Err(crate::DataLayerError::UnexpectedValue(
"provider quota value is not finite".to_string(),
));
}
Ok(Self {
provider_id,
billing_type,
monthly_quota_usd,
monthly_used_usd,
quota_reset_day: quota_reset_day.map(|value| value as u64),
quota_last_reset_at_unix_secs: quota_last_reset_at_unix_secs.map(|value| value as u64),
quota_expires_at_unix_secs: quota_expires_at_unix_secs.map(|value| value as u64),
is_active,
})
}
}
#[async_trait]
pub trait ProviderQuotaReadRepository: Send + Sync {
async fn find_by_provider_id(
&self,
provider_id: &str,
) -> Result<Option<StoredProviderQuotaSnapshot>, crate::DataLayerError>;
async fn find_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderQuotaSnapshot>, crate::DataLayerError>;
}
#[async_trait]
pub trait ProviderQuotaWriteRepository: Send + Sync {
async fn reset_due(&self, now_unix_secs: u64) -> Result<usize, crate::DataLayerError>;
}
pub trait ProviderQuotaRepository:
ProviderQuotaReadRepository + ProviderQuotaWriteRepository + Send + Sync
{
}
impl<T> ProviderQuotaRepository for T where
T: ProviderQuotaReadRepository + ProviderQuotaWriteRepository + Send + Sync
{
}
@@ -0,0 +1,11 @@
mod types;
pub use aether_routing_core::{RoutingGroupBindingSubject, RoutingGroupConfig, RoutingGroupRecord};
pub use types::{
apply_binding_patch, apply_group_patch, binding_subject_from_database,
binding_subject_to_database, CreateRoutingGroupBindingRecord, CreateRoutingGroupRecord,
CreateRoutingGroupVersionRecord, RoutingGroupBindingQuery, RoutingGroupLookupKey,
RoutingGroupReadRepository, RoutingGroupWriteRepository, StoredRoutingGroup,
StoredRoutingGroupBinding, StoredRoutingGroupVersion, UpdateRoutingGroupBindingRecord,
UpdateRoutingGroupRecord,
};
@@ -0,0 +1,333 @@
use async_trait::async_trait;
use serde_json::Value;
use aether_routing_core::RoutingGroupBindingSubject;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StoredRoutingGroup {
pub id: String,
pub name: String,
pub description: Option<String>,
pub enabled: bool,
pub is_system_default: bool,
pub config_json: Value,
pub version: i64,
pub created_at: i64,
pub updated_at: i64,
pub published_at: Option<i64>,
}
impl StoredRoutingGroup {
pub fn new(record: CreateRoutingGroupRecord) -> Result<Self, crate::DataLayerError> {
validate_non_empty(&record.id, "routing_groups.id")?;
validate_non_empty(&record.name, "routing_groups.name")?;
validate_config_object(&record.config_json, "routing_groups.config_json")?;
Ok(Self {
id: record.id,
name: record.name,
description: record.description,
enabled: record.enabled,
is_system_default: record.is_system_default,
config_json: record.config_json,
version: record.version.max(1),
created_at: record.created_at,
updated_at: record.updated_at,
published_at: record.published_at,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CreateRoutingGroupRecord {
pub id: String,
pub name: String,
pub description: Option<String>,
pub enabled: bool,
pub is_system_default: bool,
pub config_json: Value,
pub version: i64,
pub created_at: i64,
pub updated_at: i64,
pub published_at: Option<i64>,
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct UpdateRoutingGroupRecord {
pub name: Option<String>,
pub description: Option<Option<String>>,
pub enabled: Option<bool>,
pub is_system_default: Option<bool>,
pub config_json: Option<Value>,
pub version: Option<i64>,
pub updated_at: i64,
pub published_at: Option<Option<i64>>,
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct UpdateRoutingGroupBindingRecord {
pub group_id: Option<String>,
pub subject_type: Option<RoutingGroupBindingSubject>,
pub subject_id: Option<String>,
pub is_default: Option<bool>,
pub allow_explicit_select: Option<bool>,
pub updated_at: i64,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StoredRoutingGroupBinding {
pub id: String,
pub group_id: String,
pub subject_type: RoutingGroupBindingSubject,
pub subject_id: String,
pub is_default: bool,
pub allow_explicit_select: bool,
pub created_at: i64,
pub updated_at: i64,
}
impl StoredRoutingGroupBinding {
pub fn new(record: CreateRoutingGroupBindingRecord) -> Result<Self, crate::DataLayerError> {
validate_non_empty(&record.id, "routing_group_bindings.id")?;
validate_non_empty(&record.group_id, "routing_group_bindings.group_id")?;
validate_non_empty(&record.subject_id, "routing_group_bindings.subject_id")?;
Ok(Self {
id: record.id,
group_id: record.group_id,
subject_type: record.subject_type,
subject_id: record.subject_id,
is_default: record.is_default,
allow_explicit_select: record.allow_explicit_select,
created_at: record.created_at,
updated_at: record.updated_at,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CreateRoutingGroupBindingRecord {
pub id: String,
pub group_id: String,
pub subject_type: RoutingGroupBindingSubject,
pub subject_id: String,
pub is_default: bool,
pub allow_explicit_select: bool,
pub created_at: i64,
pub updated_at: i64,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StoredRoutingGroupVersion {
pub id: String,
pub group_id: String,
pub version: i64,
pub config_json: Value,
pub created_at: i64,
pub created_by: Option<String>,
}
impl StoredRoutingGroupVersion {
pub fn new(record: CreateRoutingGroupVersionRecord) -> Result<Self, crate::DataLayerError> {
validate_non_empty(&record.id, "routing_group_versions.id")?;
validate_non_empty(&record.group_id, "routing_group_versions.group_id")?;
validate_config_object(&record.config_json, "routing_group_versions.config_json")?;
Ok(Self {
id: record.id,
group_id: record.group_id,
version: record.version.max(1),
config_json: record.config_json,
created_at: record.created_at,
created_by: record.created_by,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CreateRoutingGroupVersionRecord {
pub id: String,
pub group_id: String,
pub version: i64,
pub config_json: Value,
pub created_at: i64,
pub created_by: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum RoutingGroupLookupKey<'a> {
Id(&'a str),
Name(&'a str),
SystemDefault,
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct RoutingGroupBindingQuery {
pub group_id: Option<String>,
pub subject_type: Option<RoutingGroupBindingSubject>,
pub subject_id: Option<String>,
}
#[async_trait]
pub trait RoutingGroupReadRepository: Send + Sync {
fn clear_local_cache(&self) {}
async fn list_routing_groups(&self) -> Result<Vec<StoredRoutingGroup>, crate::DataLayerError>;
async fn find_routing_group(
&self,
lookup: RoutingGroupLookupKey<'_>,
) -> Result<Option<StoredRoutingGroup>, crate::DataLayerError>;
async fn list_routing_group_bindings(
&self,
query: &RoutingGroupBindingQuery,
) -> Result<Vec<StoredRoutingGroupBinding>, crate::DataLayerError>;
async fn list_routing_group_versions(
&self,
group_id: &str,
) -> Result<Vec<StoredRoutingGroupVersion>, crate::DataLayerError>;
}
#[async_trait]
pub trait RoutingGroupWriteRepository: Send + Sync {
async fn create_routing_group(
&self,
record: CreateRoutingGroupRecord,
) -> Result<StoredRoutingGroup, crate::DataLayerError>;
async fn update_routing_group(
&self,
id: &str,
patch: UpdateRoutingGroupRecord,
) -> Result<Option<StoredRoutingGroup>, crate::DataLayerError>;
async fn delete_routing_group(&self, id: &str) -> Result<bool, crate::DataLayerError>;
async fn create_routing_group_binding(
&self,
record: CreateRoutingGroupBindingRecord,
) -> Result<StoredRoutingGroupBinding, crate::DataLayerError>;
async fn delete_routing_group_binding(&self, id: &str) -> Result<bool, crate::DataLayerError>;
async fn update_routing_group_binding(
&self,
id: &str,
patch: UpdateRoutingGroupBindingRecord,
) -> Result<Option<StoredRoutingGroupBinding>, crate::DataLayerError>;
async fn create_routing_group_version(
&self,
record: CreateRoutingGroupVersionRecord,
) -> Result<StoredRoutingGroupVersion, crate::DataLayerError>;
}
pub fn apply_group_patch(
group: &mut StoredRoutingGroup,
patch: UpdateRoutingGroupRecord,
) -> Result<(), crate::DataLayerError> {
if let Some(name) = patch.name {
if name.trim().is_empty() {
return Err(crate::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(crate::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(())
}
pub fn apply_binding_patch(
binding: &mut StoredRoutingGroupBinding,
patch: UpdateRoutingGroupBindingRecord,
) -> Result<(), crate::DataLayerError> {
if let Some(group_id) = patch.group_id {
if group_id.trim().is_empty() {
return Err(crate::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(crate::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(())
}
pub fn binding_subject_to_database(subject: RoutingGroupBindingSubject) -> &'static str {
match subject {
RoutingGroupBindingSubject::User => "user",
RoutingGroupBindingSubject::ApiKey => "api_key",
RoutingGroupBindingSubject::UserGroup => "user_group",
}
}
pub fn binding_subject_from_database(
value: String,
) -> Result<RoutingGroupBindingSubject, crate::DataLayerError> {
match value.as_str() {
"user" => Ok(RoutingGroupBindingSubject::User),
"api_key" => Ok(RoutingGroupBindingSubject::ApiKey),
"user_group" => Ok(RoutingGroupBindingSubject::UserGroup),
_ => Err(crate::DataLayerError::UnexpectedValue(format!(
"invalid routing_group_bindings.subject_type: {value}"
))),
}
}
fn validate_non_empty(value: &str, field: &str) -> Result<(), crate::DataLayerError> {
if value.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(format!(
"{field} is empty"
)));
}
Ok(())
}
fn validate_config_object(value: &Value, field: &str) -> Result<(), crate::DataLayerError> {
if !value.is_object() {
return Err(crate::DataLayerError::InvalidInput(format!(
"{field} must be a JSON object"
)));
}
Ok(())
}
@@ -0,0 +1,7 @@
mod types;
pub use types::{
finite_wallet_available_usd, plan_finite_wallet_debit, settlement_billable_cost_usd,
settlement_billing_status_for_usage_status, SettlementRepository, SettlementWriteRepository,
StoredUsageSettlement, UsageSettlementInput, WalletDebitPlan, SETTLEMENT_EPSILON_USD,
};
@@ -0,0 +1,136 @@
use async_trait::async_trait;
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct UsageSettlementInput {
pub request_id: String,
pub user_id: Option<String>,
pub api_key_id: Option<String>,
#[serde(default)]
pub api_key_is_standalone: bool,
pub provider_id: Option<String>,
pub status: String,
pub billing_status: String,
pub total_cost_usd: f64,
pub actual_total_cost_usd: f64,
pub finalized_at_unix_secs: Option<u64>,
}
impl UsageSettlementInput {
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
if self.request_id.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"settlement request_id cannot be empty".to_string(),
));
}
if self.status.trim().is_empty() || self.billing_status.trim().is_empty() {
return Err(crate::DataLayerError::InvalidInput(
"settlement status cannot be empty".to_string(),
));
}
if !self.total_cost_usd.is_finite() || !self.actual_total_cost_usd.is_finite() {
return Err(crate::DataLayerError::InvalidInput(
"settlement cost must be finite".to_string(),
));
}
Ok(())
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredUsageSettlement {
pub request_id: String,
pub wallet_id: Option<String>,
pub billing_status: String,
pub wallet_balance_before: Option<f64>,
pub wallet_balance_after: Option<f64>,
pub wallet_recharge_balance_before: Option<f64>,
pub wallet_recharge_balance_after: Option<f64>,
pub wallet_gift_balance_before: Option<f64>,
pub wallet_gift_balance_after: Option<f64>,
pub provider_monthly_used_usd: Option<f64>,
pub finalized_at_unix_secs: Option<u64>,
}
#[async_trait]
pub trait SettlementWriteRepository: Send + Sync {
async fn settle_usage(
&self,
input: UsageSettlementInput,
) -> Result<Option<StoredUsageSettlement>, crate::DataLayerError>;
}
pub trait SettlementRepository: SettlementWriteRepository + Send + Sync {}
impl<T> SettlementRepository for T where T: SettlementWriteRepository + Send + Sync {}
pub const SETTLEMENT_EPSILON_USD: f64 = 0.000_000_01;
#[derive(Debug, Clone, Copy)]
pub struct WalletDebitPlan {
pub recharge_deduction: f64,
pub gift_deduction: f64,
pub recharge_overdraft: f64,
}
impl WalletDebitPlan {
pub fn after_balances(self, recharge_balance: f64, gift_balance: f64) -> (f64, f64) {
(
recharge_balance - self.recharge_deduction - self.recharge_overdraft,
gift_balance - self.gift_deduction,
)
}
}
pub fn finite_wallet_available_usd(recharge_balance: f64, gift_balance: f64) -> f64 {
recharge_balance.max(0.0) + gift_balance.max(0.0)
}
pub fn plan_finite_wallet_debit(
recharge_balance: f64,
gift_balance: f64,
requested_usd: f64,
) -> WalletDebitPlan {
let requested_usd = requested_usd.max(0.0);
let recharge_deduction = recharge_balance.max(0.0).min(requested_usd);
let after_recharge_remaining = (requested_usd - recharge_deduction).max(0.0);
let gift_deduction = gift_balance.max(0.0).min(after_recharge_remaining);
let recharge_overdraft = (after_recharge_remaining - gift_deduction).max(0.0);
WalletDebitPlan {
recharge_deduction,
gift_deduction,
recharge_overdraft,
}
}
pub fn settlement_billing_status_for_usage_status(status: &str) -> &'static str {
match status {
"completed" | "cancelled" => "settled",
_ => "void",
}
}
pub fn settlement_billable_cost_usd(input: &UsageSettlementInput) -> f64 {
input.actual_total_cost_usd.max(0.0)
}
#[cfg(test)]
mod tests {
use super::UsageSettlementInput;
#[test]
fn rejects_invalid_settlement_input() {
let input = UsageSettlementInput {
request_id: "".to_string(),
user_id: None,
api_key_id: None,
api_key_is_standalone: false,
provider_id: None,
status: "completed".to_string(),
billing_status: "pending".to_string(),
total_cost_usd: 0.1,
actual_total_cost_usd: 0.1,
finalized_at_unix_secs: None,
};
assert!(input.validate().is_err());
}
}
@@ -0,0 +1,38 @@
mod policy;
mod types;
pub use policy::*;
pub use types::{
extract_provider_actual_service_tier_from_response,
extract_provider_cache_ttl_minutes_from_metadata, extract_provider_reasoning_effort_from_body,
extract_provider_service_tier_from_body, normalize_provider_service_tier, parse_usage_body_ref,
resolve_provider_cache_ttl_minutes, usage_body_ref, usage_request_metadata_client_family,
ApiKeyLastUsedDelta, ManagementTokenCounterDelta, PendingUsageCleanupSummary,
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, UsageBodyCaptureResult, UsageBodyCaptureState,
UsageBodyCaptureStorage, UsageBodyField, UsageBreakdownGroupBy, UsageBreakdownSummaryQuery,
UsageCacheAffinityHitSummaryQuery, UsageCacheAffinityIntervalGroupBy,
UsageCacheAffinityIntervalQuery, UsageCacheHitSummaryQuery, UsageCleanupExecutionMode,
UsageCleanupPreviewCounts, UsageCleanupSummary, UsageCleanupTargets, UsageCleanupWindow,
UsageCostSavingsSummaryQuery, UsageCounterFlushSummary, UsageCounterHealthSnapshot,
UsageCounterPendingHealthSnapshot, UsageDailyHeatmapQuery, UsageDashboardDailyBreakdownQuery,
UsageDashboardProviderCountsQuery, UsageDashboardSummaryQuery, UsageErrorDistributionQuery,
UsageLeaderboardGroupBy, UsageLeaderboardQuery, UsageMonitoringErrorCountQuery,
UsageMonitoringErrorListQuery, UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery,
UsageReadRepository, UsageRepository, UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity,
UsageTimeSeriesQuery, UsageWriteRepository, PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY,
PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, PROVIDER_REASONING_EFFORT_METADATA_KEY,
PROVIDER_SERVICE_TIER_METADATA_KEY,
};
@@ -0,0 +1,327 @@
use super::{StoredRequestUsageAudit, UpsertUsageRecord};
#[derive(Debug, Clone, PartialEq, Default)]
pub struct ApiKeyUsageContribution {
pub api_key_id: String,
pub total_requests: i64,
pub total_tokens: i64,
pub total_cost_usd: f64,
pub last_used_at_unix_secs: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Default)]
pub struct ApiKeyUsageDelta {
pub total_requests: i64,
pub total_tokens: i64,
pub total_cost_usd: f64,
pub candidate_last_used_at_unix_secs: Option<u64>,
pub removed_last_used_at_unix_secs: Option<u64>,
}
impl ApiKeyUsageDelta {
pub fn between(before: &ApiKeyUsageContribution, after: &ApiKeyUsageContribution) -> Self {
Self {
total_requests: after.total_requests - before.total_requests,
total_tokens: after.total_tokens - before.total_tokens,
total_cost_usd: after.total_cost_usd - before.total_cost_usd,
candidate_last_used_at_unix_secs: newer_last_used_at(
before.last_used_at_unix_secs,
after.last_used_at_unix_secs,
),
removed_last_used_at_unix_secs: None,
}
}
pub fn addition(after: &ApiKeyUsageContribution) -> Self {
Self {
total_requests: after.total_requests,
total_tokens: after.total_tokens,
total_cost_usd: after.total_cost_usd,
candidate_last_used_at_unix_secs: after.last_used_at_unix_secs,
removed_last_used_at_unix_secs: None,
}
}
pub fn removal(before: &ApiKeyUsageContribution) -> Self {
Self {
total_requests: -before.total_requests,
total_tokens: -before.total_tokens,
total_cost_usd: -before.total_cost_usd,
candidate_last_used_at_unix_secs: None,
removed_last_used_at_unix_secs: before.last_used_at_unix_secs,
}
}
pub fn is_noop(&self) -> bool {
self.total_requests == 0
&& self.total_tokens == 0
&& self.total_cost_usd == 0.0
&& self.candidate_last_used_at_unix_secs.is_none()
&& self.removed_last_used_at_unix_secs.is_none()
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct ModelUsageContribution {
pub model: String,
pub request_count: i64,
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct ModelUsageDelta {
pub request_count: i64,
}
impl ModelUsageDelta {
pub fn between(before: &ModelUsageContribution, after: &ModelUsageContribution) -> Self {
Self {
request_count: after.request_count - before.request_count,
}
}
pub fn addition(after: &ModelUsageContribution) -> Self {
Self {
request_count: after.request_count,
}
}
pub fn removal(before: &ModelUsageContribution) -> Self {
Self {
request_count: -before.request_count,
}
}
pub fn is_noop(&self) -> bool {
self.request_count == 0
}
}
#[derive(Debug, Clone, PartialEq, Default)]
pub struct ProviderApiKeyUsageContribution {
pub key_id: String,
pub request_count: i64,
pub success_count: i64,
pub error_count: i64,
pub total_tokens: i64,
pub total_cost_usd: f64,
pub total_response_time_ms: i64,
pub last_used_at_unix_secs: Option<u64>,
pub usage_created_at_unix_secs: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Default)]
pub struct ProviderApiKeyUsageDelta {
pub request_count: i64,
pub success_count: i64,
pub error_count: i64,
pub total_tokens: i64,
pub total_cost_usd: f64,
pub total_response_time_ms: i64,
pub candidate_last_used_at_unix_secs: Option<u64>,
pub removed_last_used_at_unix_secs: Option<u64>,
pub usage_created_at_unix_secs: Option<u64>,
}
impl ProviderApiKeyUsageDelta {
pub fn between(
before: &ProviderApiKeyUsageContribution,
after: &ProviderApiKeyUsageContribution,
) -> Self {
Self {
request_count: after.request_count - before.request_count,
success_count: after.success_count - before.success_count,
error_count: after.error_count - before.error_count,
total_tokens: after.total_tokens - before.total_tokens,
total_cost_usd: after.total_cost_usd - before.total_cost_usd,
total_response_time_ms: after.total_response_time_ms - before.total_response_time_ms,
candidate_last_used_at_unix_secs: newer_last_used_at(
before.last_used_at_unix_secs,
after.last_used_at_unix_secs,
),
removed_last_used_at_unix_secs: None,
usage_created_at_unix_secs: after.usage_created_at_unix_secs,
}
}
pub fn addition(after: &ProviderApiKeyUsageContribution) -> Self {
Self {
request_count: after.request_count,
success_count: after.success_count,
error_count: after.error_count,
total_tokens: after.total_tokens,
total_cost_usd: after.total_cost_usd,
total_response_time_ms: after.total_response_time_ms,
candidate_last_used_at_unix_secs: after.last_used_at_unix_secs,
removed_last_used_at_unix_secs: None,
usage_created_at_unix_secs: after.usage_created_at_unix_secs,
}
}
pub fn removal(before: &ProviderApiKeyUsageContribution) -> Self {
Self {
request_count: -before.request_count,
success_count: -before.success_count,
error_count: -before.error_count,
total_tokens: -before.total_tokens,
total_cost_usd: -before.total_cost_usd,
total_response_time_ms: -before.total_response_time_ms,
candidate_last_used_at_unix_secs: None,
removed_last_used_at_unix_secs: before.last_used_at_unix_secs,
usage_created_at_unix_secs: before.usage_created_at_unix_secs,
}
}
pub fn is_noop(&self) -> bool {
self.request_count == 0
&& self.success_count == 0
&& self.error_count == 0
&& self.total_tokens == 0
&& self.total_cost_usd == 0.0
&& self.total_response_time_ms == 0
&& self.candidate_last_used_at_unix_secs.is_none()
&& self.removed_last_used_at_unix_secs.is_none()
}
}
pub fn incoming_usage_can_recover_terminal_failure(
incoming_status: &str,
incoming_billing_status: &str,
) -> bool {
incoming_billing_status == "pending" && incoming_status == "completed"
}
pub fn usage_can_recover_terminal_failure(
existing_status: &str,
existing_billing_status: &str,
incoming_status: &str,
incoming_billing_status: &str,
) -> bool {
existing_billing_status == "void"
&& matches!(existing_status, "failed" | "cancelled")
&& incoming_usage_can_recover_terminal_failure(incoming_status, incoming_billing_status)
}
pub fn strip_deprecated_usage_display_fields(mut usage: UpsertUsageRecord) -> UpsertUsageRecord {
usage.username = None;
usage.api_key_name = None;
usage
}
pub fn provider_api_key_usage_is_success(
status: &str,
status_code: Option<u16>,
error_message: Option<&str>,
) -> bool {
matches!(
status,
"completed" | "success" | "ok" | "billed" | "settled"
) && status_code.is_none_or(|code| code < 400)
&& error_message.is_none_or(|value| value.trim().is_empty())
}
pub fn provider_api_key_usage_is_error(
status: &str,
status_code: Option<u16>,
error_message: Option<&str>,
) -> bool {
!matches!(status, "pending" | "streaming")
&& !provider_api_key_usage_is_success(status, status_code, error_message)
}
pub fn provider_api_key_usage_contribution(
usage: &StoredRequestUsageAudit,
) -> Option<ProviderApiKeyUsageContribution> {
let key_id = usage
.provider_api_key_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())?
.to_string();
let is_in_flight = matches!(usage.status.as_str(), "pending" | "streaming");
let is_success = provider_api_key_usage_is_success(
usage.status.as_str(),
usage.status_code,
usage.error_message.as_deref(),
);
let is_error = provider_api_key_usage_is_error(
usage.status.as_str(),
usage.status_code,
usage.error_message.as_deref(),
);
Some(ProviderApiKeyUsageContribution {
key_id,
request_count: 1,
success_count: i64::from(is_success),
error_count: i64::from(is_error),
total_tokens: if is_in_flight {
0
} else {
i64::try_from(usage.total_tokens).unwrap_or(i64::MAX)
},
total_cost_usd: if is_in_flight {
0.0
} else if usage.total_cost_usd.is_finite() {
usage.total_cost_usd.max(0.0)
} else {
0.0
},
total_response_time_ms: if is_success {
usage
.response_time_ms
.and_then(|value| i64::try_from(value).ok())
.unwrap_or_default()
} else {
0
},
last_used_at_unix_secs: Some(usage.created_at_unix_ms),
usage_created_at_unix_secs: Some(usage.created_at_unix_ms),
})
}
pub fn model_usage_contribution(usage: &StoredRequestUsageAudit) -> Option<ModelUsageContribution> {
if matches!(usage.status.as_str(), "pending" | "streaming") {
return None;
}
let model = usage.model.trim();
if model.is_empty() {
return None;
}
Some(ModelUsageContribution {
model: model.to_string(),
request_count: 1,
})
}
pub fn api_key_usage_contribution(
usage: &StoredRequestUsageAudit,
) -> Option<ApiKeyUsageContribution> {
if matches!(usage.status.as_str(), "pending" | "streaming") {
return None;
}
let api_key_id = usage
.api_key_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())?
.to_string();
Some(ApiKeyUsageContribution {
api_key_id,
total_requests: 1,
total_tokens: i64::try_from(usage.total_tokens).unwrap_or(i64::MAX),
total_cost_usd: if usage.total_cost_usd.is_finite() {
usage.total_cost_usd.max(0.0)
} else {
0.0
},
last_used_at_unix_secs: Some(usage.created_at_unix_ms),
})
}
fn newer_last_used_at(before: Option<u64>, after: Option<u64>) -> Option<u64> {
match (before, after) {
(Some(before), Some(after)) if after > before => Some(after),
(None, Some(after)) => Some(after),
_ => None,
}
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,7 @@
mod types;
pub use types::{
StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount,
VideoTaskQueryFilter, VideoTaskReadRepository, VideoTaskRepository, VideoTaskStatus,
VideoTaskStatusCount, VideoTaskWriteRepository,
};
@@ -0,0 +1,618 @@
use async_trait::async_trait;
use serde_json::Value;
#[derive(
Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize,
)]
pub enum VideoTaskStatus {
Pending,
Submitted,
Queued,
Processing,
Completed,
Failed,
Cancelled,
Expired,
Deleted,
}
impl VideoTaskStatus {
pub fn from_database(value: &str) -> Result<Self, crate::DataLayerError> {
match value.trim().to_ascii_lowercase().as_str() {
"pending" => Ok(Self::Pending),
"submitted" => Ok(Self::Submitted),
"queued" => Ok(Self::Queued),
"processing" => Ok(Self::Processing),
"completed" => Ok(Self::Completed),
"failed" => Ok(Self::Failed),
"cancelled" => Ok(Self::Cancelled),
"expired" => Ok(Self::Expired),
"deleted" => Ok(Self::Deleted),
other => Err(crate::DataLayerError::UnexpectedValue(format!(
"unsupported video_tasks.status: {other}"
))),
}
}
pub fn is_active(self) -> bool {
matches!(
self,
Self::Pending | Self::Submitted | Self::Queued | Self::Processing
)
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredVideoTask {
pub id: String,
pub short_id: Option<String>,
pub request_id: String,
pub user_id: Option<String>,
pub api_key_id: Option<String>,
pub username: Option<String>,
pub api_key_name: Option<String>,
pub external_task_id: Option<String>,
pub provider_id: Option<String>,
pub endpoint_id: Option<String>,
pub key_id: Option<String>,
pub client_api_format: Option<String>,
pub provider_api_format: Option<String>,
pub format_converted: bool,
pub model: Option<String>,
pub prompt: Option<String>,
pub original_request_body: Option<Value>,
pub duration_seconds: Option<u32>,
pub resolution: Option<String>,
pub aspect_ratio: Option<String>,
pub size: Option<String>,
pub status: VideoTaskStatus,
pub progress_percent: u16,
pub progress_message: Option<String>,
pub retry_count: u32,
pub poll_interval_seconds: u32,
pub next_poll_at_unix_secs: Option<u64>,
pub poll_count: u32,
pub max_poll_count: u32,
pub created_at_unix_ms: u64,
pub submitted_at_unix_secs: Option<u64>,
pub completed_at_unix_secs: Option<u64>,
pub updated_at_unix_secs: u64,
pub error_code: Option<String>,
pub error_message: Option<String>,
pub video_url: Option<String>,
pub request_metadata: Option<Value>,
}
impl StoredVideoTask {
#[allow(clippy::too_many_arguments)]
pub fn new(
id: String,
short_id: Option<String>,
request_id: String,
user_id: Option<String>,
api_key_id: Option<String>,
username: Option<String>,
api_key_name: Option<String>,
external_task_id: Option<String>,
provider_id: Option<String>,
endpoint_id: Option<String>,
key_id: Option<String>,
client_api_format: Option<String>,
provider_api_format: Option<String>,
format_converted: bool,
model: Option<String>,
prompt: Option<String>,
original_request_body: Option<Value>,
duration_seconds: Option<i32>,
resolution: Option<String>,
aspect_ratio: Option<String>,
size: Option<String>,
status: VideoTaskStatus,
progress_percent: i32,
progress_message: Option<String>,
retry_count: i32,
poll_interval_seconds: i32,
next_poll_at_unix_secs: Option<i64>,
poll_count: i32,
max_poll_count: i32,
created_at_unix_ms: i64,
submitted_at_unix_secs: Option<i64>,
completed_at_unix_secs: Option<i64>,
updated_at_unix_secs: i64,
error_code: Option<String>,
error_message: Option<String>,
video_url: Option<String>,
request_metadata: Option<Value>,
) -> Result<Self, crate::DataLayerError> {
let progress_percent = u16::try_from(progress_percent).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid progress_percent: {progress_percent}"
))
})?;
let retry_count = u32::try_from(retry_count).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!("invalid retry_count: {retry_count}"))
})?;
let poll_interval_seconds = u32::try_from(poll_interval_seconds).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid poll_interval_seconds: {poll_interval_seconds}"
))
})?;
let next_poll_at_unix_secs =
coerce_optional_unix_secs(next_poll_at_unix_secs, "next_poll_at_unix_secs")?;
let poll_count = u32::try_from(poll_count).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!("invalid poll_count: {poll_count}"))
})?;
let max_poll_count = u32::try_from(max_poll_count).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid max_poll_count: {max_poll_count}"
))
})?;
let created_at_unix_ms = u64::try_from(created_at_unix_ms).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid created_at_unix_ms: {created_at_unix_ms}"
))
})?;
let submitted_at_unix_secs =
coerce_optional_unix_secs(submitted_at_unix_secs, "submitted_at_unix_secs")?;
let completed_at_unix_secs =
coerce_optional_unix_secs(completed_at_unix_secs, "completed_at_unix_secs")?;
let updated_at_unix_secs = u64::try_from(updated_at_unix_secs).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!(
"invalid updated_at_unix_secs: {updated_at_unix_secs}"
))
})?;
let duration_seconds = match duration_seconds {
Some(value) => Some(u32::try_from(value).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!("invalid duration_seconds: {value}"))
})?),
None => None,
};
Ok(Self {
id,
short_id,
request_id,
user_id,
api_key_id,
username,
api_key_name,
external_task_id,
provider_id,
endpoint_id,
key_id,
client_api_format,
provider_api_format,
format_converted,
model,
prompt,
original_request_body,
duration_seconds,
resolution,
aspect_ratio,
size,
status,
progress_percent,
progress_message,
retry_count,
poll_interval_seconds,
next_poll_at_unix_secs,
poll_count,
max_poll_count,
created_at_unix_ms,
submitted_at_unix_secs,
completed_at_unix_secs,
updated_at_unix_secs,
error_code,
error_message,
video_url,
request_metadata,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct UpsertVideoTask {
pub id: String,
pub short_id: Option<String>,
pub request_id: String,
pub user_id: Option<String>,
pub api_key_id: Option<String>,
pub username: Option<String>,
pub api_key_name: Option<String>,
pub external_task_id: Option<String>,
pub provider_id: Option<String>,
pub endpoint_id: Option<String>,
pub key_id: Option<String>,
pub client_api_format: Option<String>,
pub provider_api_format: Option<String>,
pub format_converted: bool,
pub model: Option<String>,
pub prompt: Option<String>,
pub original_request_body: Option<Value>,
pub duration_seconds: Option<u32>,
pub resolution: Option<String>,
pub aspect_ratio: Option<String>,
pub size: Option<String>,
pub status: VideoTaskStatus,
pub progress_percent: u16,
pub progress_message: Option<String>,
pub retry_count: u32,
pub poll_interval_seconds: u32,
pub next_poll_at_unix_secs: Option<u64>,
pub poll_count: u32,
pub max_poll_count: u32,
pub created_at_unix_ms: u64,
pub submitted_at_unix_secs: Option<u64>,
pub completed_at_unix_secs: Option<u64>,
pub updated_at_unix_secs: u64,
pub error_code: Option<String>,
pub error_message: Option<String>,
pub video_url: Option<String>,
pub request_metadata: Option<Value>,
}
impl UpsertVideoTask {
pub fn into_stored(self) -> StoredVideoTask {
StoredVideoTask {
id: self.id,
short_id: self.short_id,
request_id: self.request_id,
user_id: self.user_id,
api_key_id: self.api_key_id,
username: self.username,
api_key_name: self.api_key_name,
external_task_id: self.external_task_id,
provider_id: self.provider_id,
endpoint_id: self.endpoint_id,
key_id: self.key_id,
client_api_format: self.client_api_format,
provider_api_format: self.provider_api_format,
format_converted: self.format_converted,
model: self.model,
prompt: self.prompt,
original_request_body: self.original_request_body,
duration_seconds: self.duration_seconds,
resolution: self.resolution,
aspect_ratio: self.aspect_ratio,
size: self.size,
status: self.status,
progress_percent: self.progress_percent,
progress_message: self.progress_message,
retry_count: self.retry_count,
poll_interval_seconds: self.poll_interval_seconds,
next_poll_at_unix_secs: self.next_poll_at_unix_secs,
poll_count: self.poll_count,
max_poll_count: self.max_poll_count,
created_at_unix_ms: self.created_at_unix_ms,
submitted_at_unix_secs: self.submitted_at_unix_secs,
completed_at_unix_secs: self.completed_at_unix_secs,
updated_at_unix_secs: self.updated_at_unix_secs,
error_code: self.error_code,
error_message: self.error_message,
video_url: self.video_url,
request_metadata: self.request_metadata,
}
}
}
impl From<StoredVideoTask> for UpsertVideoTask {
fn from(task: StoredVideoTask) -> Self {
Self {
id: task.id,
short_id: task.short_id,
request_id: task.request_id,
user_id: task.user_id,
api_key_id: task.api_key_id,
username: task.username,
api_key_name: task.api_key_name,
external_task_id: task.external_task_id,
provider_id: task.provider_id,
endpoint_id: task.endpoint_id,
key_id: task.key_id,
client_api_format: task.client_api_format,
provider_api_format: task.provider_api_format,
format_converted: task.format_converted,
model: task.model,
prompt: task.prompt,
original_request_body: task.original_request_body,
duration_seconds: task.duration_seconds,
resolution: task.resolution,
aspect_ratio: task.aspect_ratio,
size: task.size,
status: task.status,
progress_percent: task.progress_percent,
progress_message: task.progress_message,
retry_count: task.retry_count,
poll_interval_seconds: task.poll_interval_seconds,
next_poll_at_unix_secs: task.next_poll_at_unix_secs,
poll_count: task.poll_count,
max_poll_count: task.max_poll_count,
created_at_unix_ms: task.created_at_unix_ms,
submitted_at_unix_secs: task.submitted_at_unix_secs,
completed_at_unix_secs: task.completed_at_unix_secs,
updated_at_unix_secs: task.updated_at_unix_secs,
error_code: task.error_code,
error_message: task.error_message,
video_url: task.video_url,
request_metadata: task.request_metadata,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum VideoTaskLookupKey<'a> {
Id(&'a str),
ShortId(&'a str),
UserExternal {
user_id: &'a str,
external_task_id: &'a str,
},
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct VideoTaskQueryFilter {
pub user_id: Option<String>,
pub status: Option<VideoTaskStatus>,
pub model_substring: Option<String>,
pub client_api_format: Option<String>,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct VideoTaskStatusCount {
pub status: VideoTaskStatus,
pub count: u64,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct VideoTaskModelCount {
pub model: String,
pub count: u64,
}
#[async_trait]
pub trait VideoTaskReadRepository: Send + Sync {
async fn find(
&self,
key: VideoTaskLookupKey<'_>,
) -> Result<Option<StoredVideoTask>, crate::DataLayerError>;
async fn list_active(
&self,
limit: usize,
) -> Result<Vec<StoredVideoTask>, crate::DataLayerError>;
async fn list_due(
&self,
now_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredVideoTask>, crate::DataLayerError>;
async fn list_page(
&self,
filter: &VideoTaskQueryFilter,
offset: usize,
limit: usize,
) -> Result<Vec<StoredVideoTask>, crate::DataLayerError>;
async fn list_page_summary(
&self,
filter: &VideoTaskQueryFilter,
offset: usize,
limit: usize,
) -> Result<Vec<StoredVideoTask>, crate::DataLayerError>;
async fn count(&self, filter: &VideoTaskQueryFilter) -> Result<u64, crate::DataLayerError>;
async fn count_by_status(
&self,
filter: &VideoTaskQueryFilter,
) -> Result<Vec<VideoTaskStatusCount>, crate::DataLayerError>;
async fn count_distinct_users(
&self,
filter: &VideoTaskQueryFilter,
) -> Result<u64, crate::DataLayerError>;
async fn top_models(
&self,
filter: &VideoTaskQueryFilter,
limit: usize,
) -> Result<Vec<VideoTaskModelCount>, crate::DataLayerError>;
async fn count_created_since(
&self,
filter: &VideoTaskQueryFilter,
created_since_unix_secs: u64,
) -> Result<u64, crate::DataLayerError>;
}
#[async_trait]
pub trait VideoTaskWriteRepository: Send + Sync {
async fn upsert(&self, task: UpsertVideoTask)
-> Result<StoredVideoTask, crate::DataLayerError>;
async fn update_if_active(
&self,
task: UpsertVideoTask,
) -> Result<Option<StoredVideoTask>, crate::DataLayerError>;
async fn claim_due(
&self,
now_unix_secs: u64,
claim_until_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredVideoTask>, crate::DataLayerError>;
}
pub trait VideoTaskRepository:
VideoTaskReadRepository + VideoTaskWriteRepository + Send + Sync
{
}
impl<T> VideoTaskRepository for T where
T: VideoTaskReadRepository + VideoTaskWriteRepository + Send + Sync
{
}
fn coerce_optional_unix_secs(
value: Option<i64>,
field: &str,
) -> Result<Option<u64>, crate::DataLayerError> {
match value {
Some(value) => Ok(Some(u64::try_from(value).map_err(|_| {
crate::DataLayerError::UnexpectedValue(format!("invalid {field}: {value}"))
})?)),
None => Ok(None),
}
}
#[cfg(test)]
mod tests {
use super::{StoredVideoTask, VideoTaskStatus};
#[allow(clippy::type_complexity)]
fn base_new_args() -> (
String,
Option<String>,
String,
Option<String>,
Option<String>,
Option<String>,
Option<String>,
Option<String>,
Option<String>,
Option<String>,
Option<String>,
Option<String>,
Option<String>,
bool,
Option<String>,
Option<String>,
Option<serde_json::Value>,
Option<i32>,
Option<String>,
Option<String>,
Option<String>,
VideoTaskStatus,
i32,
Option<String>,
i32,
i32,
Option<i64>,
i32,
i32,
i64,
Option<i64>,
Option<i64>,
i64,
Option<String>,
Option<String>,
Option<String>,
Option<serde_json::Value>,
) {
(
"task-1".to_string(),
None,
"request-1".to_string(),
None,
None,
None,
None,
None,
None,
None,
None,
None,
None,
false,
None,
None,
None,
None,
None,
None,
None,
VideoTaskStatus::Submitted,
10,
None,
0,
10,
Some(1),
0,
360,
1,
None,
None,
1,
None,
None,
None,
None,
)
}
#[test]
fn parses_status_from_database_text() {
assert_eq!(
VideoTaskStatus::from_database("processing").expect("status should parse"),
VideoTaskStatus::Processing
);
}
#[test]
fn rejects_invalid_database_status() {
assert!(VideoTaskStatus::from_database("mystery").is_err());
}
#[test]
fn rejects_invalid_numeric_fields() {
let mut args = base_new_args();
args.22 = -1;
assert!(StoredVideoTask::new(
args.0, args.1, args.2, args.3, args.4, args.5, args.6, args.7, args.8, args.9,
args.10, args.11, args.12, args.13, args.14, args.15, args.16, args.17, args.18,
args.19, args.20, args.21, args.22, args.23, args.24, args.25, args.26, args.27,
args.28, args.29, args.30, args.31, args.32, args.33, args.34, args.35, args.36,
)
.is_err());
}
#[test]
fn rejects_negative_updated_at_values() {
let mut args = base_new_args();
args.32 = -1;
assert!(StoredVideoTask::new(
args.0, args.1, args.2, args.3, args.4, args.5, args.6, args.7, args.8, args.9,
args.10, args.11, args.12, args.13, args.14, args.15, args.16, args.17, args.18,
args.19, args.20, args.21, args.22, args.23, args.24, args.25, args.26, args.27,
args.28, args.29, args.30, args.31, args.32, args.33, args.34, args.35, args.36,
)
.is_err());
}
#[test]
fn rejects_negative_created_at_values() {
let mut args = base_new_args();
args.29 = -1;
assert!(StoredVideoTask::new(
args.0, args.1, args.2, args.3, args.4, args.5, args.6, args.7, args.8, args.9,
args.10, args.11, args.12, args.13, args.14, args.15, args.16, args.17, args.18,
args.19, args.20, args.21, args.22, args.23, args.24, args.25, args.26, args.27,
args.28, args.29, args.30, args.31, args.32, args.33, args.34, args.35, args.36,
)
.is_err());
}
#[test]
fn rejects_negative_optional_completed_at_values() {
let mut args = base_new_args();
args.31 = Some(-1);
assert!(StoredVideoTask::new(
args.0, args.1, args.2, args.3, args.4, args.5, args.6, args.7, args.8, args.9,
args.10, args.11, args.12, args.13, args.14, args.15, args.16, args.17, args.18,
args.19, args.20, args.21, args.22, args.23, args.24, args.25, args.26, args.27,
args.28, args.29, args.30, args.31, args.32, args.33, args.34, args.35, args.36,
)
.is_err());
}
}
@@ -0,0 +1,5 @@
mod snapshot;
mod types;
pub use snapshot::{WalletReadSeed, WalletReadSnapshot};
pub use types::*;
@@ -0,0 +1,443 @@
use std::collections::{BTreeMap, BTreeSet};
use super::{
AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery,
AdminWalletLedgerQuery, AdminWalletListQuery, AdminWalletRefundRequestListQuery,
StoredAdminPaymentCallback, StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder,
StoredAdminPaymentOrderPage, StoredAdminRedeemCode, StoredAdminRedeemCodeBatch,
StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, StoredAdminWalletLedgerPage,
StoredAdminWalletListItem, StoredAdminWalletListPage, StoredAdminWalletRefund,
StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestItem,
StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction,
StoredAdminWalletTransactionPage, StoredWalletSnapshot, WalletLookupKey,
};
#[derive(Debug, Clone, Default)]
pub struct WalletReadSeed {
pub wallets: Vec<StoredWalletSnapshot>,
pub payment_orders: Vec<StoredAdminPaymentOrder>,
pub payment_callbacks: Vec<StoredAdminPaymentCallback>,
pub wallet_transactions: Vec<StoredAdminWalletTransaction>,
pub refunds: Vec<StoredAdminWalletRefund>,
pub redeem_batches: Vec<StoredAdminRedeemCodeBatch>,
pub redeem_codes: Vec<StoredAdminRedeemCode>,
}
/// Immutable wallet read model shared by memory and SQL adapters.
#[derive(Debug, Clone, Default)]
pub struct WalletReadSnapshot {
wallets: BTreeMap<String, StoredWalletSnapshot>,
payment_orders: BTreeMap<String, StoredAdminPaymentOrder>,
payment_callbacks: BTreeMap<String, StoredAdminPaymentCallback>,
wallet_transactions: BTreeMap<String, StoredAdminWalletTransaction>,
refunds: BTreeMap<String, StoredAdminWalletRefund>,
redeem_batches: BTreeMap<String, StoredAdminRedeemCodeBatch>,
redeem_codes: BTreeMap<String, StoredAdminRedeemCode>,
}
impl WalletReadSnapshot {
pub fn new(seed: WalletReadSeed) -> Self {
Self {
wallets: by_id(seed.wallets, |item| &item.id),
payment_orders: by_id(seed.payment_orders, |item| &item.id),
payment_callbacks: by_id(seed.payment_callbacks, |item| &item.id),
wallet_transactions: by_id(seed.wallet_transactions, |item| &item.id),
refunds: by_id(seed.refunds, |item| &item.id),
redeem_batches: by_id(seed.redeem_batches, |item| &item.id),
redeem_codes: by_id(seed.redeem_codes, |item| &item.id),
}
}
pub fn find(&self, key: WalletLookupKey<'_>) -> Option<StoredWalletSnapshot> {
match key {
WalletLookupKey::WalletId(wallet_id) => self.wallets.get(wallet_id).cloned(),
WalletLookupKey::UserId(user_id) => self
.wallets
.values()
.find(|wallet| wallet.user_id.as_deref() == Some(user_id))
.cloned(),
WalletLookupKey::ApiKeyId(api_key_id) => self
.wallets
.values()
.find(|wallet| wallet.api_key_id.as_deref() == Some(api_key_id))
.cloned(),
}
}
pub fn list_wallets_by_user_ids(&self, user_ids: &[String]) -> Vec<StoredWalletSnapshot> {
let ids = user_ids.iter().map(String::as_str).collect::<BTreeSet<_>>();
self.wallets
.values()
.filter(|wallet| wallet.user_id.as_deref().is_some_and(|id| ids.contains(id)))
.cloned()
.collect()
}
pub fn list_wallets_by_api_key_ids(&self, api_key_ids: &[String]) -> Vec<StoredWalletSnapshot> {
let ids = api_key_ids
.iter()
.map(String::as_str)
.collect::<BTreeSet<_>>();
self.wallets
.values()
.filter(|wallet| {
wallet
.api_key_id
.as_deref()
.is_some_and(|id| ids.contains(id))
})
.cloned()
.collect()
}
pub fn list_admin_wallets(&self, query: &AdminWalletListQuery) -> StoredAdminWalletListPage {
let mut items = self
.wallets
.values()
.filter(|wallet| {
query
.status
.as_deref()
.is_none_or(|expected| wallet.status == expected)
})
.filter(|wallet| match query.owner_type.as_deref() {
Some("user") => wallet.user_id.is_some(),
Some("api_key") => wallet.api_key_id.is_some(),
_ => true,
})
.map(|wallet| StoredAdminWalletListItem {
id: wallet.id.clone(),
user_id: wallet.user_id.clone(),
api_key_id: wallet.api_key_id.clone(),
balance: wallet.balance,
gift_balance: wallet.gift_balance,
limit_mode: wallet.limit_mode.clone(),
currency: wallet.currency.clone(),
status: wallet.status.clone(),
total_recharged: wallet.total_recharged,
total_consumed: wallet.total_consumed,
total_refunded: wallet.total_refunded,
total_adjusted: wallet.total_adjusted,
user_name: None,
api_key_name: None,
created_at_unix_ms: None,
updated_at_unix_secs: Some(wallet.updated_at_unix_secs),
})
.collect::<Vec<_>>();
items.sort_by(|left, right| {
right
.updated_at_unix_secs
.cmp(&left.updated_at_unix_secs)
.then_with(|| right.id.cmp(&left.id))
});
let total = items.len() as u64;
let items = items
.into_iter()
.skip(query.offset)
.take(query.limit)
.collect();
StoredAdminWalletListPage { items, total }
}
pub fn list_admin_wallet_ledger(
&self,
_query: &AdminWalletLedgerQuery,
) -> StoredAdminWalletLedgerPage {
StoredAdminWalletLedgerPage::default()
}
pub fn list_admin_wallet_refund_requests(
&self,
query: &AdminWalletRefundRequestListQuery,
) -> StoredAdminWalletRefundRequestPage {
let mut items = self
.refunds
.values()
.filter(|refund| {
query
.status
.as_deref()
.is_none_or(|expected| refund.status == expected)
})
.filter_map(|refund| {
let wallet = self.wallets.get(&refund.wallet_id)?;
Some(StoredAdminWalletRefundRequestItem {
id: refund.id.clone(),
refund_no: refund.refund_no.clone(),
wallet_id: refund.wallet_id.clone(),
user_id: refund.user_id.clone(),
payment_order_id: refund.payment_order_id.clone(),
source_type: refund.source_type.clone(),
source_id: refund.source_id.clone(),
refund_mode: refund.refund_mode.clone(),
amount_usd: refund.amount_usd,
status: refund.status.clone(),
reason: refund.reason.clone(),
failure_reason: refund.failure_reason.clone(),
gateway_refund_id: refund.gateway_refund_id.clone(),
payout_method: refund.payout_method.clone(),
payout_reference: refund.payout_reference.clone(),
payout_proof: refund.payout_proof.clone(),
requested_by: refund.requested_by.clone(),
approved_by: refund.approved_by.clone(),
processed_by: refund.processed_by.clone(),
wallet_user_id: wallet.user_id.clone(),
wallet_user_name: None,
wallet_api_key_id: wallet.api_key_id.clone(),
api_key_name: None,
wallet_status: wallet.status.clone(),
created_at_unix_ms: Some(refund.created_at_unix_ms),
updated_at_unix_secs: Some(refund.updated_at_unix_secs),
processed_at_unix_secs: refund.processed_at_unix_secs,
completed_at_unix_secs: refund.completed_at_unix_secs,
})
})
.collect::<Vec<_>>();
items.sort_by_key(|item| std::cmp::Reverse(item.created_at_unix_ms));
page(items, query.offset, query.limit, |items, total| {
StoredAdminWalletRefundRequestPage { items, total }
})
}
pub fn list_admin_wallet_transactions(
&self,
wallet_id: &str,
limit: usize,
offset: usize,
) -> StoredAdminWalletTransactionPage {
let mut items = self
.wallet_transactions
.values()
.filter(|item| item.wallet_id == wallet_id)
.cloned()
.collect::<Vec<_>>();
items.sort_by_key(|item| std::cmp::Reverse(item.created_at_unix_ms));
page(items, offset, limit, |items, total| {
StoredAdminWalletTransactionPage { items, total }
})
}
pub fn list_admin_wallet_refunds(
&self,
wallet_id: &str,
limit: usize,
offset: usize,
) -> StoredAdminWalletRefundPage {
let mut items = self
.refunds
.values()
.filter(|item| item.wallet_id == wallet_id)
.cloned()
.collect::<Vec<_>>();
items.sort_by_key(|item| std::cmp::Reverse(item.created_at_unix_ms));
page(items, offset, limit, |items, total| {
StoredAdminWalletRefundPage { items, total }
})
}
pub fn list_admin_payment_orders(
&self,
query: &AdminPaymentOrderListQuery,
now_unix_secs: u64,
) -> StoredAdminPaymentOrderPage {
let mut items = self
.payment_orders
.values()
.filter(|order| {
query.status.as_deref().is_none_or(|expected| {
let effective = if order.status == "pending"
&& order
.expires_at_unix_secs
.is_some_and(|value| value < now_unix_secs)
{
"expired"
} else {
order.status.as_str()
};
effective == expected
}) && query
.payment_method
.as_deref()
.is_none_or(|expected| order.payment_method == expected)
})
.cloned()
.collect::<Vec<_>>();
items.sort_by_key(|item| std::cmp::Reverse(item.created_at_unix_ms));
page(items, query.offset, query.limit, |items, total| {
StoredAdminPaymentOrderPage { items, total }
})
}
pub fn find_admin_payment_order(&self, order_id: &str) -> Option<StoredAdminPaymentOrder> {
self.payment_orders.get(order_id).cloned()
}
pub fn list_wallet_payment_orders_by_user_id(
&self,
user_id: &str,
limit: usize,
offset: usize,
) -> StoredAdminPaymentOrderPage {
let mut items = self
.payment_orders
.values()
.filter(|order| order.user_id.as_deref() == Some(user_id))
.cloned()
.collect::<Vec<_>>();
items.sort_by_key(|item| std::cmp::Reverse(item.created_at_unix_ms));
page(items, offset, limit, |items, total| {
StoredAdminPaymentOrderPage { items, total }
})
}
pub fn count_pending_refunds_by_user_id(&self, user_id: &str) -> u64 {
const STATUSES: &[&str] = &["pending_approval", "approved", "processing"];
self.refunds
.values()
.filter(|item| {
item.user_id.as_deref() == Some(user_id) && STATUSES.contains(&item.status.as_str())
})
.count() as u64
}
pub fn count_pending_payment_orders_by_user_id(&self, user_id: &str) -> u64 {
const STATUSES: &[&str] = &["pending", "paid"];
self.payment_orders
.values()
.filter(|item| {
item.user_id.as_deref() == Some(user_id) && STATUSES.contains(&item.status.as_str())
})
.count() as u64
}
pub fn find_wallet_payment_order_by_user_id(
&self,
user_id: &str,
order_id: &str,
) -> Option<StoredAdminPaymentOrder> {
self.payment_orders
.get(order_id)
.filter(|order| order.user_id.as_deref() == Some(user_id))
.cloned()
}
pub fn find_pending_plan_purchase_order_by_user_id(
&self,
user_id: &str,
product_id: &str,
now_unix_secs: u64,
) -> Option<StoredAdminPaymentOrder> {
self.payment_orders
.values()
.filter(|order| {
order.user_id.as_deref() == Some(user_id)
&& order.status == "pending"
&& order
.expires_at_unix_secs
.is_some_and(|expires_at| expires_at > now_unix_secs)
&& order.gateway_response.as_ref().is_some_and(|response| {
response
.get("order_kind")
.and_then(serde_json::Value::as_str)
== Some("plan_purchase")
&& response
.get("product_id")
.and_then(serde_json::Value::as_str)
== Some(product_id)
})
})
.max_by_key(|order| order.created_at_unix_ms)
.cloned()
}
pub fn find_wallet_refund(
&self,
wallet_id: &str,
refund_id: &str,
) -> Option<StoredAdminWalletRefund> {
self.refunds
.get(refund_id)
.filter(|refund| refund.wallet_id == wallet_id)
.cloned()
}
pub fn list_admin_payment_callbacks(
&self,
payment_method: Option<&str>,
limit: usize,
offset: usize,
) -> StoredAdminPaymentCallbackPage {
let mut items = self
.payment_callbacks
.values()
.filter(|item| payment_method.is_none_or(|value| item.payment_method == value))
.cloned()
.collect::<Vec<_>>();
items.sort_by_key(|item| std::cmp::Reverse(item.created_at_unix_ms));
page(items, offset, limit, |items, total| {
StoredAdminPaymentCallbackPage { items, total }
})
}
pub fn list_admin_redeem_code_batches(
&self,
query: &AdminRedeemCodeBatchListQuery,
) -> StoredAdminRedeemCodeBatchPage {
let mut items = self
.redeem_batches
.values()
.filter(|item| {
query
.status
.as_deref()
.is_none_or(|value| item.status == value)
})
.cloned()
.collect::<Vec<_>>();
items.sort_by_key(|item| std::cmp::Reverse(item.created_at_unix_ms));
page(items, query.offset, query.limit, |items, total| {
StoredAdminRedeemCodeBatchPage { items, total }
})
}
pub fn find_admin_redeem_code_batch(
&self,
batch_id: &str,
) -> Option<StoredAdminRedeemCodeBatch> {
self.redeem_batches.get(batch_id).cloned()
}
pub fn list_admin_redeem_codes(
&self,
query: &AdminRedeemCodeListQuery,
) -> StoredAdminRedeemCodePage {
let mut items = self
.redeem_codes
.values()
.filter(|item| item.batch_id == query.batch_id)
.filter(|item| {
query
.status
.as_deref()
.is_none_or(|value| item.status == value)
})
.cloned()
.collect::<Vec<_>>();
items.sort_by_key(|item| std::cmp::Reverse(item.created_at_unix_ms));
page(items, query.offset, query.limit, |items, total| {
StoredAdminRedeemCodePage { items, total }
})
}
}
fn by_id<T>(items: Vec<T>, id: impl Fn(&T) -> &str) -> BTreeMap<String, T> {
items
.into_iter()
.map(|item| (id(&item).to_string(), item))
.collect()
}
fn page<T, P>(items: Vec<T>, offset: usize, limit: usize, build: impl Fn(Vec<T>, u64) -> P) -> P {
let total = items.len() as u64;
build(items.into_iter().skip(offset).take(limit).collect(), total)
}
File diff suppressed because it is too large Load Diff