mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-07 01:47:47 +08:00
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:
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
@@ -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
Reference in New Issue
Block a user