mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
refactor: 大规模模块拆分与代码精简,新增 ai-pipeline/data-contracts 独立 crate
- 新增 aether-ai-pipeline 和 aether-data-contracts crate,将 pipeline 逻辑与数据契约从 gateway 中解耦 - 重构 admin handlers:拆分单体模块为 auth/billing/endpoint/features/model/observability/provider/system 等独立子模块 - 合并 chat/cli 重复代码路径:精简 conversion、finalize、planner 中的 sync/chat/cli 分支 - 重构 scheduler/executor/data 层,引入 facade 模式降低模块间耦合 - 移除冗余的 intent 模块,将 plan_fallback/policy/stream_path/sync_path 迁移至 executor - 前端适配:调整 admin API 调用和 provider 模型测试对话框
This commit is contained in:
@@ -7,6 +7,7 @@ repository.workspace = true
|
||||
description = "Shared data access contracts and config for Aether Rust services"
|
||||
|
||||
[dependencies]
|
||||
aether-data-contracts.workspace = true
|
||||
aether-cache.workspace = true
|
||||
aether-wallet.workspace = true
|
||||
async-trait.workspace = true
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
ALTER TABLE provider_endpoints
|
||||
ADD COLUMN IF NOT EXISTS health_score DOUBLE PRECISION NOT NULL DEFAULT 1.0;
|
||||
@@ -1,5 +1,6 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::error::SqlxResultExt;
|
||||
use crate::postgres::{
|
||||
PostgresLeaseRunner, PostgresLeaseRunnerConfig, PostgresPool, PostgresPoolConfig,
|
||||
PostgresPoolFactory, PostgresTransactionRunner,
|
||||
@@ -304,10 +305,11 @@ impl PostgresBackend {
|
||||
let row = sqlx::query(FIND_SYSTEM_CONFIG_VALUE_SQL)
|
||||
.bind(key)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.map(|row| row.try_get("value"))
|
||||
.transpose()
|
||||
.map_err(Into::into)
|
||||
.map_postgres_err()
|
||||
}
|
||||
|
||||
pub async fn upsert_system_config_value(
|
||||
@@ -322,8 +324,9 @@ impl PostgresBackend {
|
||||
.bind(value)
|
||||
.bind(description)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
row.try_get("value").map_err(Into::into)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.try_get("value").map_postgres_err()
|
||||
}
|
||||
|
||||
pub async fn list_system_config_entries(
|
||||
@@ -331,19 +334,21 @@ impl PostgresBackend {
|
||||
) -> Result<Vec<StoredSystemConfigEntry>, DataLayerError> {
|
||||
let rows = sqlx::query(LIST_SYSTEM_CONFIG_ENTRIES_SQL)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.into_iter()
|
||||
.map(|row| {
|
||||
Ok(StoredSystemConfigEntry {
|
||||
key: row.try_get("key")?,
|
||||
value: row.try_get("value")?,
|
||||
description: row.try_get("description")?,
|
||||
key: row.try_get("key").map_postgres_err()?,
|
||||
value: row.try_get("value").map_postgres_err()?,
|
||||
description: row.try_get("description").map_postgres_err()?,
|
||||
updated_at_unix_secs: row
|
||||
.try_get::<Option<i64>, _>("updated_at_unix_secs")?
|
||||
.try_get::<Option<i64>, _>("updated_at_unix_secs")
|
||||
.map_postgres_err()?
|
||||
.map(|value| value.max(0) as u64),
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
.collect::<Result<Vec<_>, DataLayerError>>()
|
||||
}
|
||||
|
||||
pub async fn upsert_system_config_entry(
|
||||
@@ -358,13 +363,15 @@ impl PostgresBackend {
|
||||
.bind(value)
|
||||
.bind(description)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(StoredSystemConfigEntry {
|
||||
key: row.try_get("key")?,
|
||||
value: row.try_get("value")?,
|
||||
description: row.try_get("description")?,
|
||||
key: row.try_get("key").map_postgres_err()?,
|
||||
value: row.try_get("value").map_postgres_err()?,
|
||||
description: row.try_get("description").map_postgres_err()?,
|
||||
updated_at_unix_secs: row
|
||||
.try_get::<Option<i64>, _>("updated_at_unix_secs")?
|
||||
.try_get::<Option<i64>, _>("updated_at_unix_secs")
|
||||
.map_postgres_err()?
|
||||
.map(|value| value.max(0) as u64),
|
||||
})
|
||||
}
|
||||
@@ -373,19 +380,33 @@ impl PostgresBackend {
|
||||
let result = sqlx::query(DELETE_SYSTEM_CONFIG_VALUE_SQL)
|
||||
.bind(key)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(result.rows_affected() > 0)
|
||||
}
|
||||
|
||||
pub async fn read_admin_system_stats(&self) -> Result<AdminSystemStats, DataLayerError> {
|
||||
let row = sqlx::query(READ_ADMIN_SYSTEM_STATS_SQL)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(AdminSystemStats {
|
||||
total_users: row.try_get::<i64, _>("total_users")?.max(0) as u64,
|
||||
active_users: row.try_get::<i64, _>("active_users")?.max(0) as u64,
|
||||
total_api_keys: row.try_get::<i64, _>("total_api_keys")?.max(0) as u64,
|
||||
total_requests: row.try_get::<i64, _>("total_requests")?.max(0) as u64,
|
||||
total_users: row
|
||||
.try_get::<i64, _>("total_users")
|
||||
.map_postgres_err()?
|
||||
.max(0) as u64,
|
||||
active_users: row
|
||||
.try_get::<i64, _>("active_users")
|
||||
.map_postgres_err()?
|
||||
.max(0) as u64,
|
||||
total_api_keys: row
|
||||
.try_get::<i64, _>("total_api_keys")
|
||||
.map_postgres_err()?
|
||||
.max(0) as u64,
|
||||
total_requests: row
|
||||
.try_get::<i64, _>("total_requests")
|
||||
.map_postgres_err()?
|
||||
.max(0) as u64,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,20 +1,29 @@
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum DataLayerError {
|
||||
#[error("invalid configuration: {0}")]
|
||||
InvalidConfiguration(String),
|
||||
pub use aether_data_contracts::DataLayerError;
|
||||
|
||||
#[error("invalid input: {0}")]
|
||||
InvalidInput(String),
|
||||
|
||||
#[error("postgres error: {0}")]
|
||||
Postgres(#[from] sqlx::Error),
|
||||
|
||||
#[error("redis error: {0}")]
|
||||
Redis(#[from] redis::RedisError),
|
||||
|
||||
#[error("operation timed out: {0}")]
|
||||
TimedOut(String),
|
||||
|
||||
#[error("unexpected database value: {0}")]
|
||||
UnexpectedValue(String),
|
||||
pub(crate) fn postgres_error(error: impl std::fmt::Display) -> DataLayerError {
|
||||
DataLayerError::postgres(error)
|
||||
}
|
||||
|
||||
pub(crate) fn redis_error(error: impl std::fmt::Display) -> DataLayerError {
|
||||
DataLayerError::redis(error)
|
||||
}
|
||||
|
||||
pub(crate) trait SqlxResultExt<T> {
|
||||
fn map_postgres_err(self) -> Result<T, DataLayerError>;
|
||||
}
|
||||
|
||||
impl<T> SqlxResultExt<T> for Result<T, sqlx::Error> {
|
||||
fn map_postgres_err(self) -> Result<T, DataLayerError> {
|
||||
self.map_err(postgres_error)
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) trait RedisResultExt<T> {
|
||||
fn map_redis_err(self) -> Result<T, DataLayerError>;
|
||||
}
|
||||
|
||||
impl<T> RedisResultExt<T> for Result<T, redis::RedisError> {
|
||||
fn map_redis_err(self) -> Result<T, DataLayerError> {
|
||||
self.map_err(redis_error)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
use crate::error::SqlxResultExt;
|
||||
use crate::postgres::{DatabaseRecordId, PostgresTransactionOptions, PostgresTransactionRunner};
|
||||
use crate::DataLayerError;
|
||||
use futures_util::FutureExt;
|
||||
@@ -108,7 +109,8 @@ impl PostgresLeaseRunner {
|
||||
.bind(owner)
|
||||
.bind(lease_ms)
|
||||
.fetch_all(&mut **tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(rows.into_iter().map(DatabaseRecordId).collect())
|
||||
}
|
||||
.boxed()
|
||||
@@ -142,7 +144,8 @@ impl PostgresLeaseRunner {
|
||||
.bind(ids)
|
||||
.bind(owner)
|
||||
.fetch_all(&mut **tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(rows.into_iter().map(DatabaseRecordId).collect())
|
||||
}
|
||||
.boxed()
|
||||
@@ -186,7 +189,8 @@ impl PostgresLeaseRunner {
|
||||
.bind(owner)
|
||||
.bind(lease_ms)
|
||||
.fetch_all(&mut **tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(rows.into_iter().map(DatabaseRecordId).collect())
|
||||
}
|
||||
.boxed()
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use futures_util::future::BoxFuture;
|
||||
use sqlx::{Postgres, Transaction};
|
||||
|
||||
use crate::error::{postgres_error, SqlxResultExt};
|
||||
use crate::postgres::PostgresPool;
|
||||
use crate::DataLayerError;
|
||||
|
||||
@@ -70,9 +71,12 @@ impl PostgresTransactionRunner {
|
||||
) -> Result<PostgresTransaction, DataLayerError> {
|
||||
options.validate()?;
|
||||
|
||||
let mut tx = self.pool.begin().await?;
|
||||
let mut tx = self.pool.begin().await.map_postgres_err()?;
|
||||
for statement in build_transaction_setup_statements(options) {
|
||||
sqlx::query(statement.as_str()).execute(&mut *tx).await?;
|
||||
sqlx::query(statement.as_str())
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
}
|
||||
Ok(tx)
|
||||
}
|
||||
@@ -90,7 +94,7 @@ impl PostgresTransactionRunner {
|
||||
let mut tx = self.begin(options).await?;
|
||||
match f(&mut tx).await {
|
||||
Ok(value) => {
|
||||
tx.commit().await?;
|
||||
tx.commit().await.map_err(postgres_error)?;
|
||||
Ok(value)
|
||||
}
|
||||
Err(err) => {
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
use crate::error::RedisResultExt;
|
||||
use crate::redis::RedisKeyspace;
|
||||
use crate::DataLayerError;
|
||||
|
||||
@@ -44,7 +45,7 @@ impl RedisClientFactory {
|
||||
}
|
||||
|
||||
pub fn connect_lazy(&self) -> Result<RedisClient, DataLayerError> {
|
||||
Ok(RedisClient::open(self.config.url.clone())?)
|
||||
RedisClient::open(self.config.url.clone()).map_redis_err()
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use std::future::Future;
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::error::RedisResultExt;
|
||||
use crate::redis::{RedisClient, RedisKeyspace};
|
||||
use crate::DataLayerError;
|
||||
|
||||
@@ -79,13 +80,18 @@ impl RedisKvRunner {
|
||||
let resolved_ttl = ttl_seconds.unwrap_or(self.config.default_ttl_seconds);
|
||||
let namespaced_key = self.keyspace.key(key);
|
||||
self.run_with_timeout("redis kv setex", async {
|
||||
let mut connection = self.client.get_multiplexed_async_connection().await?;
|
||||
let mut connection = self
|
||||
.client
|
||||
.get_multiplexed_async_connection()
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
Ok(redis::cmd("SETEX")
|
||||
.arg(&namespaced_key)
|
||||
.arg(resolved_ttl)
|
||||
.arg(value)
|
||||
.query_async(&mut connection)
|
||||
.await?)
|
||||
.await
|
||||
.map_redis_err()?)
|
||||
})
|
||||
.await
|
||||
}
|
||||
@@ -93,11 +99,16 @@ impl RedisKvRunner {
|
||||
pub async fn del(&self, key: &str) -> Result<i64, DataLayerError> {
|
||||
let namespaced_key = self.keyspace.key(key);
|
||||
self.run_with_timeout("redis kv del", async {
|
||||
let mut connection = self.client.get_multiplexed_async_connection().await?;
|
||||
let mut connection = self
|
||||
.client
|
||||
.get_multiplexed_async_connection()
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
Ok(redis::cmd("DEL")
|
||||
.arg(&namespaced_key)
|
||||
.query_async(&mut connection)
|
||||
.await?)
|
||||
.await
|
||||
.map_redis_err()?)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
use std::future::Future;
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::error::RedisResultExt;
|
||||
use crate::redis::{RedisClient, RedisKeyspace};
|
||||
use crate::DataLayerError;
|
||||
use uuid::Uuid;
|
||||
@@ -92,7 +93,11 @@ impl RedisLockRunner {
|
||||
let token = format!("{owner}:{}", Uuid::new_v4());
|
||||
|
||||
self.run_with_timeout("redis lock acquire", async {
|
||||
let mut connection = self.client.get_multiplexed_async_connection().await?;
|
||||
let mut connection = self
|
||||
.client
|
||||
.get_multiplexed_async_connection()
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
let status = redis::cmd("SET")
|
||||
.arg(&key.0)
|
||||
.arg(&token)
|
||||
@@ -100,7 +105,8 @@ impl RedisLockRunner {
|
||||
.arg("PX")
|
||||
.arg(ttl_ms)
|
||||
.query_async::<Option<String>>(&mut connection)
|
||||
.await?;
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
|
||||
Ok(status.map(|_| RedisLockLease {
|
||||
key: key.clone(),
|
||||
@@ -115,7 +121,11 @@ impl RedisLockRunner {
|
||||
pub async fn release(&self, lease: &RedisLockLease) -> Result<bool, DataLayerError> {
|
||||
validate_lease(lease)?;
|
||||
self.run_with_timeout("redis lock release", async {
|
||||
let mut connection = self.client.get_multiplexed_async_connection().await?;
|
||||
let mut connection = self
|
||||
.client
|
||||
.get_multiplexed_async_connection()
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
let deleted = redis::Script::new(
|
||||
"if redis.call('get', KEYS[1]) == ARGV[1] then \
|
||||
return redis.call('del', KEYS[1]) \
|
||||
@@ -126,7 +136,8 @@ impl RedisLockRunner {
|
||||
.key(&lease.key.0)
|
||||
.arg(&lease.token)
|
||||
.invoke_async::<i32>(&mut connection)
|
||||
.await?;
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
Ok(deleted > 0)
|
||||
})
|
||||
.await
|
||||
@@ -141,7 +152,11 @@ impl RedisLockRunner {
|
||||
let ttl_ms = self.resolve_ttl_ms(ttl_ms)?;
|
||||
|
||||
self.run_with_timeout("redis lock renew", async {
|
||||
let mut connection = self.client.get_multiplexed_async_connection().await?;
|
||||
let mut connection = self
|
||||
.client
|
||||
.get_multiplexed_async_connection()
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
let renewed = redis::Script::new(
|
||||
"if redis.call('get', KEYS[1]) == ARGV[1] then \
|
||||
return redis.call('pexpire', KEYS[1], ARGV[2]) \
|
||||
@@ -153,7 +168,8 @@ impl RedisLockRunner {
|
||||
.arg(&lease.token)
|
||||
.arg(ttl_ms)
|
||||
.invoke_async::<i32>(&mut connection)
|
||||
.await?;
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
Ok(renewed > 0)
|
||||
})
|
||||
.await
|
||||
|
||||
@@ -6,6 +6,7 @@ use redis::from_redis_value;
|
||||
use redis::streams::StreamReadReply;
|
||||
use redis::Value as RedisValue;
|
||||
|
||||
use crate::error::{redis_error, RedisResultExt};
|
||||
use crate::redis::{RedisClient, RedisKeyspace};
|
||||
use crate::DataLayerError;
|
||||
|
||||
@@ -154,7 +155,11 @@ impl RedisStreamRunner {
|
||||
validate_stream_position(start_id)?;
|
||||
|
||||
self.run_with_timeout("redis stream ensure consumer group", async {
|
||||
let mut connection = self.client.get_multiplexed_async_connection().await?;
|
||||
let mut connection = self
|
||||
.client
|
||||
.get_multiplexed_async_connection()
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
let result = redis::cmd("XGROUP")
|
||||
.arg("CREATE")
|
||||
.arg(&stream.0)
|
||||
@@ -167,7 +172,7 @@ impl RedisStreamRunner {
|
||||
match result {
|
||||
Ok(_) => Ok(()),
|
||||
Err(err) if err.code() == Some("BUSYGROUP") => Ok(()),
|
||||
Err(err) => Err(DataLayerError::Redis(err)),
|
||||
Err(err) => Err(redis_error(err)),
|
||||
}
|
||||
})
|
||||
.await
|
||||
@@ -195,7 +200,11 @@ impl RedisStreamRunner {
|
||||
}
|
||||
|
||||
self.run_with_timeout("redis stream append", async {
|
||||
let mut connection = self.client.get_multiplexed_async_connection().await?;
|
||||
let mut connection = self
|
||||
.client
|
||||
.get_multiplexed_async_connection()
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
let mut command = redis::cmd("XADD");
|
||||
command.arg(&stream.0);
|
||||
if let Some(maxlen) = maxlen.filter(|value| *value > 0) {
|
||||
@@ -205,7 +214,10 @@ impl RedisStreamRunner {
|
||||
for (key, value) in fields {
|
||||
command.arg(key).arg(value);
|
||||
}
|
||||
Ok(command.query_async::<String>(&mut connection).await?)
|
||||
Ok(command
|
||||
.query_async::<String>(&mut connection)
|
||||
.await
|
||||
.map_redis_err()?)
|
||||
})
|
||||
.await
|
||||
}
|
||||
@@ -245,7 +257,11 @@ impl RedisStreamRunner {
|
||||
validate_consumer(consumer)?;
|
||||
|
||||
self.run_with_timeout("redis stream read group", async {
|
||||
let mut connection = self.client.get_multiplexed_async_connection().await?;
|
||||
let mut connection = self
|
||||
.client
|
||||
.get_multiplexed_async_connection()
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
let mut command = redis::cmd("XREADGROUP");
|
||||
command
|
||||
.arg("GROUP")
|
||||
@@ -260,7 +276,8 @@ impl RedisStreamRunner {
|
||||
|
||||
let reply = command
|
||||
.query_async::<StreamReadReply>(&mut connection)
|
||||
.await?;
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
|
||||
Ok(reply
|
||||
.keys
|
||||
@@ -296,13 +313,20 @@ impl RedisStreamRunner {
|
||||
}
|
||||
|
||||
self.run_with_timeout("redis stream ack", async {
|
||||
let mut connection = self.client.get_multiplexed_async_connection().await?;
|
||||
let mut connection = self
|
||||
.client
|
||||
.get_multiplexed_async_connection()
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
let mut command = redis::cmd("XACK");
|
||||
command.arg(&stream.0).arg(&group.0);
|
||||
for id in ids {
|
||||
command.arg(id);
|
||||
}
|
||||
Ok(command.query_async::<usize>(&mut connection).await?)
|
||||
Ok(command
|
||||
.query_async::<usize>(&mut connection)
|
||||
.await
|
||||
.map_redis_err()?)
|
||||
})
|
||||
.await
|
||||
}
|
||||
@@ -318,13 +342,20 @@ impl RedisStreamRunner {
|
||||
}
|
||||
|
||||
self.run_with_timeout("redis stream delete", async {
|
||||
let mut connection = self.client.get_multiplexed_async_connection().await?;
|
||||
let mut connection = self
|
||||
.client
|
||||
.get_multiplexed_async_connection()
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
let mut command = redis::cmd("XDEL");
|
||||
command.arg(&stream.0);
|
||||
for id in ids {
|
||||
command.arg(id);
|
||||
}
|
||||
Ok(command.query_async::<usize>(&mut connection).await?)
|
||||
Ok(command
|
||||
.query_async::<usize>(&mut connection)
|
||||
.await
|
||||
.map_redis_err()?)
|
||||
})
|
||||
.await
|
||||
}
|
||||
@@ -344,7 +375,11 @@ impl RedisStreamRunner {
|
||||
config.validate()?;
|
||||
|
||||
self.run_with_timeout("redis stream reclaim", async {
|
||||
let mut connection = self.client.get_multiplexed_async_connection().await?;
|
||||
let mut connection = self
|
||||
.client
|
||||
.get_multiplexed_async_connection()
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
let reply = redis::cmd("XAUTOCLAIM")
|
||||
.arg(&stream.0)
|
||||
.arg(&group.0)
|
||||
@@ -354,7 +389,8 @@ impl RedisStreamRunner {
|
||||
.arg("COUNT")
|
||||
.arg(config.count)
|
||||
.query_async::<RedisValue>(&mut connection)
|
||||
.await?;
|
||||
.await
|
||||
.map_redis_err()?;
|
||||
|
||||
parse_reclaim_result(reply)
|
||||
})
|
||||
|
||||
@@ -6,7 +6,7 @@ use super::types::{
|
||||
AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository,
|
||||
CreateAnnouncementRecord, StoredAnnouncement, StoredAnnouncementPage, UpdateAnnouncementRecord,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
|
||||
const FIND_ANNOUNCEMENT_BY_ID_SQL: &str = r#"
|
||||
SELECT
|
||||
@@ -199,7 +199,8 @@ impl AnnouncementReadRepository for SqlxAnnouncementReadRepository {
|
||||
let row = sqlx::query(FIND_ANNOUNCEMENT_BY_ID_SQL)
|
||||
.bind(announcement_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_announcement_row).transpose()
|
||||
}
|
||||
|
||||
@@ -212,8 +213,12 @@ impl AnnouncementReadRepository for SqlxAnnouncementReadRepository {
|
||||
.bind(query.active_only)
|
||||
.bind(now_unix_secs as f64)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
let total = total_row.try_get::<i64, _>("total")?.max(0) as u64;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let total = total_row
|
||||
.try_get::<i64, _>("total")
|
||||
.map_postgres_err()?
|
||||
.max(0) as u64;
|
||||
|
||||
let rows = sqlx::query(LIST_ANNOUNCEMENTS_SQL)
|
||||
.bind(query.active_only)
|
||||
@@ -221,7 +226,8 @@ impl AnnouncementReadRepository for SqlxAnnouncementReadRepository {
|
||||
.bind(query.offset as i64)
|
||||
.bind(query.limit as i64)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let items = rows
|
||||
.iter()
|
||||
.map(map_announcement_row)
|
||||
@@ -239,8 +245,9 @@ impl AnnouncementReadRepository for SqlxAnnouncementReadRepository {
|
||||
.bind(user_id)
|
||||
.bind(now_unix_secs as f64)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
Ok(row.try_get::<i64, _>("total")?.max(0) as u64)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(row.try_get::<i64, _>("total").map_postgres_err()?.max(0) as u64)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -262,7 +269,8 @@ impl AnnouncementWriteRepository for SqlxAnnouncementReadRepository {
|
||||
.bind(optional_datetime(record.start_time_unix_secs))
|
||||
.bind(optional_datetime(record.end_time_unix_secs))
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
map_announcement_row(&row)
|
||||
}
|
||||
|
||||
@@ -282,7 +290,8 @@ impl AnnouncementWriteRepository for SqlxAnnouncementReadRepository {
|
||||
.bind(optional_datetime(record.start_time_unix_secs))
|
||||
.bind(optional_datetime(record.end_time_unix_secs))
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_announcement_row).transpose()
|
||||
}
|
||||
|
||||
@@ -290,7 +299,8 @@ impl AnnouncementWriteRepository for SqlxAnnouncementReadRepository {
|
||||
let result = sqlx::query(DELETE_ANNOUNCEMENT_SQL)
|
||||
.bind(announcement_id)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(result.rows_affected() > 0)
|
||||
}
|
||||
|
||||
@@ -306,7 +316,8 @@ impl AnnouncementWriteRepository for SqlxAnnouncementReadRepository {
|
||||
.bind(announcement_id)
|
||||
.bind(read_at_unix_secs as f64)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(result.rows_affected() > 0)
|
||||
}
|
||||
}
|
||||
@@ -328,19 +339,19 @@ fn current_unix_secs() -> u64 {
|
||||
|
||||
fn map_announcement_row(row: &PgRow) -> Result<StoredAnnouncement, DataLayerError> {
|
||||
StoredAnnouncement::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("title")?,
|
||||
row.try_get("content")?,
|
||||
row.try_get("type")?,
|
||||
row.try_get("priority")?,
|
||||
row.try_get("is_active")?,
|
||||
row.try_get("is_pinned")?,
|
||||
row.try_get("author_id")?,
|
||||
row.try_get("author_username")?,
|
||||
row.try_get("start_time_unix_secs")?,
|
||||
row.try_get("end_time_unix_secs")?,
|
||||
row.try_get("created_at_unix_secs")?,
|
||||
row.try_get("updated_at_unix_secs")?,
|
||||
row.try_get("id").map_postgres_err()?,
|
||||
row.try_get("title").map_postgres_err()?,
|
||||
row.try_get("content").map_postgres_err()?,
|
||||
row.try_get("type").map_postgres_err()?,
|
||||
row.try_get("priority").map_postgres_err()?,
|
||||
row.try_get("is_active").map_postgres_err()?,
|
||||
row.try_get("is_pinned").map_postgres_err()?,
|
||||
row.try_get("author_id").map_postgres_err()?,
|
||||
row.try_get("author_username").map_postgres_err()?,
|
||||
row.try_get("start_time_unix_secs").map_postgres_err()?,
|
||||
row.try_get("end_time_unix_secs").map_postgres_err()?,
|
||||
row.try_get("created_at_unix_secs").map_postgres_err()?,
|
||||
row.try_get("updated_at_unix_secs").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -7,7 +7,10 @@ use super::types::{
|
||||
StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot,
|
||||
UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
use crate::{
|
||||
error::{postgres_error, SqlxResultExt},
|
||||
DataLayerError,
|
||||
};
|
||||
|
||||
const FIND_BY_KEY_HASH_SQL: &str = r#"
|
||||
SELECT
|
||||
@@ -455,9 +458,9 @@ UPDATE api_keys
|
||||
SET
|
||||
name = COALESCE($2, name),
|
||||
rate_limit = COALESCE($3, rate_limit),
|
||||
allowed_providers = CASE WHEN $4 THEN $5 ELSE allowed_providers END,
|
||||
allowed_api_formats = CASE WHEN $6 THEN $7 ELSE allowed_api_formats END,
|
||||
allowed_models = CASE WHEN $8 THEN $9 ELSE allowed_models END,
|
||||
allowed_providers = CASE WHEN $4 THEN $5::json ELSE allowed_providers END,
|
||||
allowed_api_formats = CASE WHEN $6 THEN $7::json ELSE allowed_api_formats END,
|
||||
allowed_models = CASE WHEN $8 THEN $9::json ELSE allowed_models END,
|
||||
updated_at = NOW()
|
||||
WHERE id = $1
|
||||
AND is_standalone = TRUE
|
||||
@@ -646,28 +649,25 @@ impl SqlxAuthApiKeySnapshotReadRepository {
|
||||
key: AuthApiKeyLookupKey<'_>,
|
||||
) -> Result<Option<StoredAuthApiKeySnapshot>, DataLayerError> {
|
||||
let row = match key {
|
||||
AuthApiKeyLookupKey::KeyHash(key_hash) => {
|
||||
sqlx::query(FIND_BY_KEY_HASH_SQL)
|
||||
.bind(key_hash)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?
|
||||
}
|
||||
AuthApiKeyLookupKey::ApiKeyId(api_key_id) => {
|
||||
sqlx::query(FIND_BY_API_KEY_ID_SQL)
|
||||
.bind(api_key_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?
|
||||
}
|
||||
AuthApiKeyLookupKey::KeyHash(key_hash) => sqlx::query(FIND_BY_KEY_HASH_SQL)
|
||||
.bind(key_hash)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?,
|
||||
AuthApiKeyLookupKey::ApiKeyId(api_key_id) => sqlx::query(FIND_BY_API_KEY_ID_SQL)
|
||||
.bind(api_key_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?,
|
||||
AuthApiKeyLookupKey::UserApiKeyIds {
|
||||
user_id,
|
||||
api_key_id,
|
||||
} => {
|
||||
sqlx::query(FIND_BY_USER_API_KEY_IDS_SQL)
|
||||
.bind(api_key_id)
|
||||
.bind(user_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?
|
||||
}
|
||||
} => sqlx::query(FIND_BY_USER_API_KEY_IDS_SQL)
|
||||
.bind(api_key_id)
|
||||
.bind(user_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?,
|
||||
};
|
||||
|
||||
row.as_ref().map(map_auth_api_key_snapshot_row).transpose()
|
||||
@@ -684,7 +684,8 @@ impl SqlxAuthApiKeySnapshotReadRepository {
|
||||
let rows = sqlx::query(LIST_BY_API_KEY_IDS_SQL)
|
||||
.bind(api_key_ids)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_auth_api_key_snapshot_row).collect()
|
||||
}
|
||||
|
||||
@@ -699,7 +700,8 @@ impl SqlxAuthApiKeySnapshotReadRepository {
|
||||
let rows = sqlx::query(LIST_EXPORT_BY_USER_IDS_SQL)
|
||||
.bind(user_ids)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_auth_api_key_export_row).collect()
|
||||
}
|
||||
|
||||
@@ -714,7 +716,8 @@ impl SqlxAuthApiKeySnapshotReadRepository {
|
||||
let rows = sqlx::query(LIST_EXPORT_BY_API_KEY_IDS_SQL)
|
||||
.bind(api_key_ids)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_auth_api_key_export_row).collect()
|
||||
}
|
||||
|
||||
@@ -731,10 +734,11 @@ impl SqlxAuthApiKeySnapshotReadRepository {
|
||||
.bind(user_ids)
|
||||
.bind(now_unix_secs as f64)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(AuthApiKeyExportSummary {
|
||||
total: row.try_get::<i64, _>("total")?.max(0) as u64,
|
||||
active: row.try_get::<i64, _>("active")?.max(0) as u64,
|
||||
total: row_get::<i64>(&row, "total")?.max(0) as u64,
|
||||
active: row_get::<i64>(&row, "active")?.max(0) as u64,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -745,10 +749,11 @@ impl SqlxAuthApiKeySnapshotReadRepository {
|
||||
let row = sqlx::query(SUMMARIZE_EXPORT_NON_STANDALONE_SQL)
|
||||
.bind(now_unix_secs as f64)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(AuthApiKeyExportSummary {
|
||||
total: row.try_get::<i64, _>("total")?.max(0) as u64,
|
||||
active: row.try_get::<i64, _>("active")?.max(0) as u64,
|
||||
total: row_get::<i64>(&row, "total")?.max(0) as u64,
|
||||
active: row_get::<i64>(&row, "active")?.max(0) as u64,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -757,7 +762,8 @@ impl SqlxAuthApiKeySnapshotReadRepository {
|
||||
) -> Result<Vec<StoredAuthApiKeyExportRecord>, DataLayerError> {
|
||||
let rows = sqlx::query(LIST_EXPORT_STANDALONE_SQL)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_auth_api_key_export_row).collect()
|
||||
}
|
||||
|
||||
@@ -774,7 +780,8 @@ impl SqlxAuthApiKeySnapshotReadRepository {
|
||||
.bind(skip)
|
||||
.bind(limit)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_auth_api_key_export_row).collect()
|
||||
}
|
||||
|
||||
@@ -785,8 +792,9 @@ impl SqlxAuthApiKeySnapshotReadRepository {
|
||||
let row = sqlx::query(COUNT_EXPORT_STANDALONE_SQL)
|
||||
.bind(is_active)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
Ok(row.try_get::<i64, _>("total")?.max(0) as u64)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(row_get::<i64>(&row, "total")?.max(0) as u64)
|
||||
}
|
||||
|
||||
pub async fn summarize_export_standalone_api_keys(
|
||||
@@ -796,10 +804,11 @@ impl SqlxAuthApiKeySnapshotReadRepository {
|
||||
let row = sqlx::query(SUMMARIZE_EXPORT_STANDALONE_SQL)
|
||||
.bind(now_unix_secs as f64)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(AuthApiKeyExportSummary {
|
||||
total: row.try_get::<i64, _>("total")?.max(0) as u64,
|
||||
active: row.try_get::<i64, _>("active")?.max(0) as u64,
|
||||
total: row_get::<i64>(&row, "total")?.max(0) as u64,
|
||||
active: row_get::<i64>(&row, "active")?.max(0) as u64,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -810,7 +819,8 @@ impl SqlxAuthApiKeySnapshotReadRepository {
|
||||
let row = sqlx::query(FIND_EXPORT_STANDALONE_BY_ID_SQL)
|
||||
.bind(api_key_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_auth_api_key_export_row).transpose()
|
||||
}
|
||||
}
|
||||
@@ -901,7 +911,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
|
||||
let result = sqlx::query(TOUCH_LAST_USED_AT_SQL)
|
||||
.bind(api_key_id)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(result.rows_affected() > 0)
|
||||
}
|
||||
|
||||
@@ -918,7 +929,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
|
||||
.bind(record.rate_limit)
|
||||
.bind(record.concurrent_limit)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_auth_api_key_export_row).transpose()
|
||||
}
|
||||
|
||||
@@ -953,7 +965,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
|
||||
.bind(record.rate_limit)
|
||||
.bind(record.concurrent_limit)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_auth_api_key_export_row).transpose()
|
||||
}
|
||||
|
||||
@@ -967,7 +980,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
|
||||
.bind(record.name)
|
||||
.bind(record.rate_limit)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_auth_api_key_export_row).transpose()
|
||||
}
|
||||
|
||||
@@ -1007,7 +1021,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
|
||||
.bind(record.allowed_models.is_some())
|
||||
.bind(allowed_models)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_auth_api_key_export_row).transpose()
|
||||
}
|
||||
|
||||
@@ -1022,7 +1037,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
|
||||
.bind(api_key_id)
|
||||
.bind(is_active)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_auth_api_key_export_row).transpose()
|
||||
}
|
||||
|
||||
@@ -1035,7 +1051,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
|
||||
.bind(api_key_id)
|
||||
.bind(is_active)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_auth_api_key_export_row).transpose()
|
||||
}
|
||||
|
||||
@@ -1050,7 +1067,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
|
||||
.bind(api_key_id)
|
||||
.bind(is_locked)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(result.rows_affected() > 0)
|
||||
}
|
||||
|
||||
@@ -1069,7 +1087,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
|
||||
.bind(api_key_id)
|
||||
.bind(allowed_providers)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_auth_api_key_export_row).transpose()
|
||||
}
|
||||
|
||||
@@ -1084,7 +1103,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
|
||||
.bind(api_key_id)
|
||||
.bind(force_capabilities)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_auth_api_key_export_row).transpose()
|
||||
}
|
||||
|
||||
@@ -1093,101 +1113,125 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut tx = self.pool.begin().await?;
|
||||
let mut tx = self.pool.begin().await.map_postgres_err()?;
|
||||
sqlx::query(NULL_USAGE_API_KEY_FK_SQL)
|
||||
.bind(api_key_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
sqlx::query(NULL_REQUEST_CANDIDATE_API_KEY_FK_SQL)
|
||||
.bind(api_key_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let result = sqlx::query(DELETE_USER_API_KEY_SQL)
|
||||
.bind(user_id)
|
||||
.bind(api_key_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
tx.commit().await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
tx.commit().await.map_err(postgres_error)?;
|
||||
Ok(result.rows_affected() > 0)
|
||||
}
|
||||
|
||||
async fn delete_standalone_api_key(&self, api_key_id: &str) -> Result<bool, DataLayerError> {
|
||||
let mut tx = self.pool.begin().await?;
|
||||
let mut tx = self.pool.begin().await.map_postgres_err()?;
|
||||
sqlx::query(NULL_USAGE_API_KEY_FK_SQL)
|
||||
.bind(api_key_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
sqlx::query(NULL_REQUEST_CANDIDATE_API_KEY_FK_SQL)
|
||||
.bind(api_key_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let result = sqlx::query(DELETE_STANDALONE_API_KEY_SQL)
|
||||
.bind(api_key_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
tx.commit().await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
tx.commit().await.map_err(postgres_error)?;
|
||||
Ok(result.rows_affected() > 0)
|
||||
}
|
||||
}
|
||||
|
||||
fn row_get<T>(row: &sqlx::postgres::PgRow, column: &str) -> Result<T, DataLayerError>
|
||||
where
|
||||
for<'r> T: sqlx::Decode<'r, sqlx::Postgres> + sqlx::Type<sqlx::Postgres>,
|
||||
{
|
||||
row.try_get(column).map_postgres_err()
|
||||
}
|
||||
|
||||
fn map_auth_api_key_snapshot_row(
|
||||
row: &sqlx::postgres::PgRow,
|
||||
) -> Result<StoredAuthApiKeySnapshot, DataLayerError> {
|
||||
let snapshot = StoredAuthApiKeySnapshot::new(
|
||||
row.try_get("user_id")?,
|
||||
row.try_get("username")?,
|
||||
row.try_get("email")?,
|
||||
row.try_get("user_role")?,
|
||||
row.try_get("user_auth_source")?,
|
||||
row.try_get("user_is_active")?,
|
||||
row.try_get("user_is_deleted")?,
|
||||
row.try_get("user_allowed_providers")?,
|
||||
row.try_get("user_allowed_api_formats")?,
|
||||
row.try_get("user_allowed_models")?,
|
||||
row.try_get("api_key_id")?,
|
||||
row.try_get("api_key_name")?,
|
||||
row.try_get("api_key_is_active")?,
|
||||
row.try_get("api_key_is_locked")?,
|
||||
row.try_get("api_key_is_standalone")?,
|
||||
row.try_get("api_key_rate_limit")?,
|
||||
row.try_get("api_key_concurrent_limit")?,
|
||||
row.try_get("api_key_expires_at_unix_secs")?,
|
||||
row.try_get("api_key_allowed_providers")?,
|
||||
row.try_get("api_key_allowed_api_formats")?,
|
||||
row.try_get("api_key_allowed_models")?,
|
||||
row_get(row, "user_id")?,
|
||||
row_get(row, "username")?,
|
||||
row_get(row, "email")?,
|
||||
row_get(row, "user_role")?,
|
||||
row_get(row, "user_auth_source")?,
|
||||
row_get(row, "user_is_active")?,
|
||||
row_get(row, "user_is_deleted")?,
|
||||
row_get(row, "user_allowed_providers")?,
|
||||
row_get(row, "user_allowed_api_formats")?,
|
||||
row_get(row, "user_allowed_models")?,
|
||||
row_get(row, "api_key_id")?,
|
||||
row_get(row, "api_key_name")?,
|
||||
row_get(row, "api_key_is_active")?,
|
||||
row_get(row, "api_key_is_locked")?,
|
||||
row_get(row, "api_key_is_standalone")?,
|
||||
row_get(row, "api_key_rate_limit")?,
|
||||
row_get(row, "api_key_concurrent_limit")?,
|
||||
row_get(row, "api_key_expires_at_unix_secs")?,
|
||||
row_get(row, "api_key_allowed_providers")?,
|
||||
row_get(row, "api_key_allowed_api_formats")?,
|
||||
row_get(row, "api_key_allowed_models")?,
|
||||
)?;
|
||||
Ok(snapshot.with_user_rate_limit(row.try_get("user_rate_limit")?))
|
||||
Ok(snapshot.with_user_rate_limit(row_get(row, "user_rate_limit")?))
|
||||
}
|
||||
|
||||
fn map_auth_api_key_export_row(
|
||||
row: &sqlx::postgres::PgRow,
|
||||
) -> Result<StoredAuthApiKeyExportRecord, DataLayerError> {
|
||||
StoredAuthApiKeyExportRecord::new(
|
||||
row.try_get("user_id")?,
|
||||
row.try_get("api_key_id")?,
|
||||
row.try_get("key_hash")?,
|
||||
row.try_get("key_encrypted")?,
|
||||
row.try_get("name")?,
|
||||
row.try_get("allowed_providers")?,
|
||||
row.try_get("allowed_api_formats")?,
|
||||
row.try_get("allowed_models")?,
|
||||
row.try_get("rate_limit")?,
|
||||
row.try_get("concurrent_limit")?,
|
||||
row.try_get("force_capabilities")?,
|
||||
row.try_get("is_active")?,
|
||||
row.try_get("expires_at_unix_secs")?,
|
||||
row.try_get("auto_delete_on_expiry")?,
|
||||
row.try_get::<i32, _>("total_requests")?.into(),
|
||||
row.try_get("total_cost_usd")?,
|
||||
row.try_get("is_standalone")?,
|
||||
row_get(row, "user_id")?,
|
||||
row_get(row, "api_key_id")?,
|
||||
row_get(row, "key_hash")?,
|
||||
row_get(row, "key_encrypted")?,
|
||||
row_get(row, "name")?,
|
||||
row_get(row, "allowed_providers")?,
|
||||
row_get(row, "allowed_api_formats")?,
|
||||
row_get(row, "allowed_models")?,
|
||||
row_get(row, "rate_limit")?,
|
||||
row_get(row, "concurrent_limit")?,
|
||||
row_get(row, "force_capabilities")?,
|
||||
row_get(row, "is_active")?,
|
||||
row_get(row, "expires_at_unix_secs")?,
|
||||
row_get(row, "auto_delete_on_expiry")?,
|
||||
row_get::<i32>(row, "total_requests")?.into(),
|
||||
row_get(row, "total_cost_usd")?,
|
||||
row_get(row, "is_standalone")?,
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::SqlxAuthApiKeySnapshotReadRepository;
|
||||
use super::{SqlxAuthApiKeySnapshotReadRepository, UPDATE_STANDALONE_API_KEY_BASIC_SQL};
|
||||
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
|
||||
|
||||
#[test]
|
||||
fn update_standalone_api_key_basic_sql_casts_json_case_values() {
|
||||
assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL
|
||||
.contains("allowed_providers = CASE WHEN $4 THEN $5::json ELSE allowed_providers END"));
|
||||
assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL.contains(
|
||||
"allowed_api_formats = CASE WHEN $6 THEN $7::json ELSE allowed_api_formats END"
|
||||
));
|
||||
assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL
|
||||
.contains("allowed_models = CASE WHEN $8 THEN $9::json ELSE allowed_models END"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn repository_constructs_from_lazy_pool() {
|
||||
let factory = PostgresPoolFactory::new(PostgresPoolConfig {
|
||||
|
||||
@@ -5,7 +5,7 @@ use super::types::{
|
||||
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
|
||||
StoredOAuthProviderModuleConfig,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
|
||||
const LIST_ENABLED_OAUTH_PROVIDERS_SQL: &str = r#"
|
||||
SELECT
|
||||
@@ -152,14 +152,16 @@ impl AuthModuleReadRepository for SqlxAuthModuleReadRepository {
|
||||
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
|
||||
let rows = sqlx::query(LIST_ENABLED_OAUTH_PROVIDERS_SQL)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_oauth_row).collect()
|
||||
}
|
||||
|
||||
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
|
||||
let row = sqlx::query(GET_LDAP_CONFIG_SQL)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_ldap_row).transpose()
|
||||
}
|
||||
}
|
||||
@@ -171,14 +173,16 @@ impl AuthModuleReadRepository for SqlxAuthModuleRepository {
|
||||
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
|
||||
let rows = sqlx::query(LIST_ENABLED_OAUTH_PROVIDERS_SQL)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_oauth_row).collect()
|
||||
}
|
||||
|
||||
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
|
||||
let row = sqlx::query(GET_LDAP_CONFIG_SQL)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_ldap_row).transpose()
|
||||
}
|
||||
}
|
||||
@@ -203,7 +207,8 @@ impl AuthModuleWriteRepository for SqlxAuthModuleRepository {
|
||||
.bind(config.use_starttls)
|
||||
.bind(config.connect_timeout)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
if let Some(row) = updated.as_ref() {
|
||||
return map_ldap_row(row).map(Some);
|
||||
}
|
||||
@@ -222,35 +227,36 @@ impl AuthModuleWriteRepository for SqlxAuthModuleRepository {
|
||||
.bind(config.use_starttls)
|
||||
.bind(config.connect_timeout)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
inserted.as_ref().map(map_ldap_row).transpose()
|
||||
}
|
||||
}
|
||||
|
||||
fn map_oauth_row(row: &PgRow) -> Result<StoredOAuthProviderModuleConfig, DataLayerError> {
|
||||
StoredOAuthProviderModuleConfig::new(
|
||||
row.try_get("provider_type")?,
|
||||
row.try_get("display_name")?,
|
||||
row.try_get("client_id")?,
|
||||
row.try_get("client_secret_encrypted")?,
|
||||
row.try_get("redirect_uri")?,
|
||||
row.try_get("provider_type").map_postgres_err()?,
|
||||
row.try_get("display_name").map_postgres_err()?,
|
||||
row.try_get("client_id").map_postgres_err()?,
|
||||
row.try_get("client_secret_encrypted").map_postgres_err()?,
|
||||
row.try_get("redirect_uri").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
fn map_ldap_row(row: &PgRow) -> Result<StoredLdapModuleConfig, DataLayerError> {
|
||||
Ok(StoredLdapModuleConfig {
|
||||
server_url: row.try_get("server_url")?,
|
||||
bind_dn: row.try_get("bind_dn")?,
|
||||
bind_password_encrypted: row.try_get("bind_password_encrypted")?,
|
||||
base_dn: row.try_get("base_dn")?,
|
||||
user_search_filter: row.try_get("user_search_filter")?,
|
||||
username_attr: row.try_get("username_attr")?,
|
||||
email_attr: row.try_get("email_attr")?,
|
||||
display_name_attr: row.try_get("display_name_attr")?,
|
||||
is_enabled: row.try_get("is_enabled")?,
|
||||
is_exclusive: row.try_get("is_exclusive")?,
|
||||
use_starttls: row.try_get("use_starttls")?,
|
||||
connect_timeout: row.try_get("connect_timeout")?,
|
||||
server_url: row.try_get("server_url").map_postgres_err()?,
|
||||
bind_dn: row.try_get("bind_dn").map_postgres_err()?,
|
||||
bind_password_encrypted: row.try_get("bind_password_encrypted").map_postgres_err()?,
|
||||
base_dn: row.try_get("base_dn").map_postgres_err()?,
|
||||
user_search_filter: row.try_get("user_search_filter").map_postgres_err()?,
|
||||
username_attr: row.try_get("username_attr").map_postgres_err()?,
|
||||
email_attr: row.try_get("email_attr").map_postgres_err()?,
|
||||
display_name_attr: row.try_get("display_name_attr").map_postgres_err()?,
|
||||
is_enabled: row.try_get("is_enabled").map_postgres_err()?,
|
||||
is_exclusive: row.try_get("is_exclusive").map_postgres_err()?,
|
||||
use_starttls: row.try_get("use_starttls").map_postgres_err()?,
|
||||
connect_timeout: row.try_get("connect_timeout").map_postgres_err()?,
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@ use std::sync::RwLock;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::types::{BillingReadRepository, StoredBillingModelContext};
|
||||
use super::{BillingReadRepository, StoredBillingModelContext};
|
||||
use crate::DataLayerError;
|
||||
|
||||
type BillingContextKey = (String, String, Option<String>);
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
mod memory;
|
||||
mod sql;
|
||||
mod types;
|
||||
|
||||
pub use memory::InMemoryBillingReadRepository;
|
||||
pub use sql::SqlxBillingReadRepository;
|
||||
pub use types::{
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use aether_data_contracts::repository::billing::{
|
||||
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingPresetApplyResult,
|
||||
AdminBillingRuleRecord, AdminBillingRuleWriteInput, BillingReadRepository,
|
||||
StoredBillingModelContext,
|
||||
};
|
||||
pub use memory::InMemoryBillingReadRepository;
|
||||
pub use sql::SqlxBillingReadRepository;
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{PgPool, Row};
|
||||
|
||||
use super::types::{BillingReadRepository, StoredBillingModelContext};
|
||||
use crate::DataLayerError;
|
||||
use super::{BillingReadRepository, StoredBillingModelContext};
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
|
||||
const FIND_MODEL_CONTEXT_SQL: &str = r#"
|
||||
SELECT
|
||||
@@ -58,7 +58,8 @@ impl SqlxBillingReadRepository {
|
||||
.bind(global_model_name)
|
||||
.bind(provider_api_key_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_row).transpose()
|
||||
}
|
||||
}
|
||||
@@ -77,22 +78,26 @@ impl BillingReadRepository for SqlxBillingReadRepository {
|
||||
|
||||
fn map_row(row: &sqlx::postgres::PgRow) -> Result<StoredBillingModelContext, DataLayerError> {
|
||||
StoredBillingModelContext::new(
|
||||
row.try_get("provider_id")?,
|
||||
row.try_get("provider_billing_type")?,
|
||||
row.try_get("provider_api_key_id")?,
|
||||
row.try_get("provider_api_key_rate_multipliers")?,
|
||||
row.try_get::<Option<i32>, _>("provider_api_key_cache_ttl_minutes")?
|
||||
row.try_get("provider_id").map_postgres_err()?,
|
||||
row.try_get("provider_billing_type").map_postgres_err()?,
|
||||
row.try_get("provider_api_key_id").map_postgres_err()?,
|
||||
row.try_get("provider_api_key_rate_multipliers")
|
||||
.map_postgres_err()?,
|
||||
row.try_get::<Option<i32>, _>("provider_api_key_cache_ttl_minutes")
|
||||
.map_postgres_err()?
|
||||
.map(i64::from),
|
||||
row.try_get("global_model_id")?,
|
||||
row.try_get("global_model_name")?,
|
||||
row.try_get("global_model_config")?,
|
||||
row.try_get("default_price_per_request")?,
|
||||
row.try_get("default_tiered_pricing")?,
|
||||
row.try_get("model_id")?,
|
||||
row.try_get("model_provider_model_name")?,
|
||||
row.try_get("model_config")?,
|
||||
row.try_get("model_price_per_request")?,
|
||||
row.try_get("model_tiered_pricing")?,
|
||||
row.try_get("global_model_id").map_postgres_err()?,
|
||||
row.try_get("global_model_name").map_postgres_err()?,
|
||||
row.try_get("global_model_config").map_postgres_err()?,
|
||||
row.try_get("default_price_per_request")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("default_tiered_pricing").map_postgres_err()?,
|
||||
row.try_get("model_id").map_postgres_err()?,
|
||||
row.try_get("model_provider_model_name")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("model_config").map_postgres_err()?,
|
||||
row.try_get("model_price_per_request").map_postgres_err()?,
|
||||
row.try_get("model_tiered_pricing").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -1,153 +0,0 @@
|
||||
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_secs: 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_secs: 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>,
|
||||
}
|
||||
|
||||
#[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>;
|
||||
}
|
||||
@@ -2,7 +2,7 @@ use std::sync::RwLock;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::types::{MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow};
|
||||
use super::{MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow};
|
||||
use crate::DataLayerError;
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
mod memory;
|
||||
mod sql;
|
||||
mod types;
|
||||
|
||||
pub use memory::InMemoryMinimalCandidateSelectionReadRepository;
|
||||
pub use sql::SqlxMinimalCandidateSelectionReadRepository;
|
||||
pub use types::{
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use aether_data_contracts::repository::candidate_selection::{
|
||||
MinimalCandidateSelectionReadRepository, MinimalCandidateSelectionRepository,
|
||||
StoredMinimalCandidateSelectionRow, StoredProviderModelMapping,
|
||||
};
|
||||
pub use memory::InMemoryMinimalCandidateSelectionReadRepository;
|
||||
pub use sql::SqlxMinimalCandidateSelectionReadRepository;
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{PgPool, Row};
|
||||
|
||||
use super::types::{
|
||||
use super::{
|
||||
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
|
||||
StoredProviderModelMapping,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
|
||||
const LIST_FOR_EXACT_API_FORMAT_SQL: &str = r#"
|
||||
SELECT
|
||||
@@ -172,7 +172,8 @@ impl SqlxMinimalCandidateSelectionReadRepository {
|
||||
let rows = sqlx::query(LIST_FOR_EXACT_API_FORMAT_SQL)
|
||||
.bind(api_format)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_candidate_selection_row).collect()
|
||||
}
|
||||
|
||||
@@ -185,7 +186,8 @@ impl SqlxMinimalCandidateSelectionReadRepository {
|
||||
.bind(api_format)
|
||||
.bind(global_model_name)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_candidate_selection_row).collect()
|
||||
}
|
||||
}
|
||||
@@ -212,46 +214,53 @@ fn map_candidate_selection_row(
|
||||
row: &sqlx::postgres::PgRow,
|
||||
) -> Result<StoredMinimalCandidateSelectionRow, DataLayerError> {
|
||||
Ok(StoredMinimalCandidateSelectionRow {
|
||||
provider_id: row.try_get("provider_id")?,
|
||||
provider_name: row.try_get("provider_name")?,
|
||||
provider_type: row.try_get("provider_type")?,
|
||||
provider_priority: row.try_get("provider_priority")?,
|
||||
provider_is_active: row.try_get("provider_is_active")?,
|
||||
endpoint_id: row.try_get("endpoint_id")?,
|
||||
endpoint_api_format: row.try_get("endpoint_api_format")?,
|
||||
endpoint_api_family: row.try_get("endpoint_api_family")?,
|
||||
endpoint_kind: row.try_get("endpoint_kind")?,
|
||||
endpoint_is_active: row.try_get("endpoint_is_active")?,
|
||||
key_id: row.try_get("key_id")?,
|
||||
key_name: row.try_get("key_name")?,
|
||||
key_auth_type: row.try_get("key_auth_type")?,
|
||||
key_is_active: row.try_get("key_is_active")?,
|
||||
provider_id: row.try_get("provider_id").map_postgres_err()?,
|
||||
provider_name: row.try_get("provider_name").map_postgres_err()?,
|
||||
provider_type: row.try_get("provider_type").map_postgres_err()?,
|
||||
provider_priority: row.try_get("provider_priority").map_postgres_err()?,
|
||||
provider_is_active: row.try_get("provider_is_active").map_postgres_err()?,
|
||||
endpoint_id: row.try_get("endpoint_id").map_postgres_err()?,
|
||||
endpoint_api_format: row.try_get("endpoint_api_format").map_postgres_err()?,
|
||||
endpoint_api_family: row.try_get("endpoint_api_family").map_postgres_err()?,
|
||||
endpoint_kind: row.try_get("endpoint_kind").map_postgres_err()?,
|
||||
endpoint_is_active: row.try_get("endpoint_is_active").map_postgres_err()?,
|
||||
key_id: row.try_get("key_id").map_postgres_err()?,
|
||||
key_name: row.try_get("key_name").map_postgres_err()?,
|
||||
key_auth_type: row.try_get("key_auth_type").map_postgres_err()?,
|
||||
key_is_active: row.try_get("key_is_active").map_postgres_err()?,
|
||||
key_api_formats: parse_string_list(
|
||||
row.try_get("key_api_formats")?,
|
||||
row.try_get("key_api_formats").map_postgres_err()?,
|
||||
"provider_api_keys.api_formats",
|
||||
)?,
|
||||
key_allowed_models: parse_string_list(
|
||||
row.try_get("key_allowed_models")?,
|
||||
row.try_get("key_allowed_models").map_postgres_err()?,
|
||||
"provider_api_keys.allowed_models",
|
||||
)?,
|
||||
key_capabilities: row.try_get("key_capabilities")?,
|
||||
key_internal_priority: row.try_get("key_internal_priority")?,
|
||||
key_global_priority_by_format: row.try_get("key_global_priority_by_format")?,
|
||||
model_id: row.try_get("model_id")?,
|
||||
global_model_id: row.try_get("global_model_id")?,
|
||||
global_model_name: row.try_get("global_model_name")?,
|
||||
key_capabilities: row.try_get("key_capabilities").map_postgres_err()?,
|
||||
key_internal_priority: row.try_get("key_internal_priority").map_postgres_err()?,
|
||||
key_global_priority_by_format: row
|
||||
.try_get("key_global_priority_by_format")
|
||||
.map_postgres_err()?,
|
||||
model_id: row.try_get("model_id").map_postgres_err()?,
|
||||
global_model_id: row.try_get("global_model_id").map_postgres_err()?,
|
||||
global_model_name: row.try_get("global_model_name").map_postgres_err()?,
|
||||
global_model_mappings: parse_string_list(
|
||||
row.try_get("global_model_mappings")?,
|
||||
row.try_get("global_model_mappings").map_postgres_err()?,
|
||||
"global_models.config.model_mappings",
|
||||
)?,
|
||||
global_model_supports_streaming: row.try_get("global_model_supports_streaming")?,
|
||||
model_provider_model_name: row.try_get("model_provider_model_name")?,
|
||||
global_model_supports_streaming: row
|
||||
.try_get("global_model_supports_streaming")
|
||||
.map_postgres_err()?,
|
||||
model_provider_model_name: row
|
||||
.try_get("model_provider_model_name")
|
||||
.map_postgres_err()?,
|
||||
model_provider_model_mappings: parse_provider_model_mappings(
|
||||
row.try_get("model_provider_model_mappings")?,
|
||||
row.try_get("model_provider_model_mappings")
|
||||
.map_postgres_err()?,
|
||||
)?,
|
||||
model_supports_streaming: row.try_get("model_supports_streaming")?,
|
||||
model_is_active: row.try_get("model_is_active")?,
|
||||
model_is_available: row.try_get("model_is_available")?,
|
||||
model_supports_streaming: row.try_get("model_supports_streaming").map_postgres_err()?,
|
||||
model_is_active: row.try_get("model_is_active").map_postgres_err()?,
|
||||
model_is_available: row.try_get("model_is_available").map_postgres_err()?,
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
@@ -1,144 +0,0 @@
|
||||
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>>,
|
||||
}
|
||||
|
||||
#[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,
|
||||
}
|
||||
|
||||
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 {
|
||||
let target = api_format.trim();
|
||||
match self.key_api_formats.as_deref() {
|
||||
None => true,
|
||||
Some(formats) => formats
|
||||
.iter()
|
||||
.any(|value| value.eq_ignore_ascii_case(target)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait MinimalCandidateSelectionReadRepository: Send + Sync {
|
||||
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>;
|
||||
}
|
||||
|
||||
pub trait MinimalCandidateSelectionRepository:
|
||||
MinimalCandidateSelectionReadRepository + Send + Sync
|
||||
{
|
||||
}
|
||||
|
||||
impl<T> MinimalCandidateSelectionRepository for T where
|
||||
T: MinimalCandidateSelectionReadRepository + Send + Sync
|
||||
{
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{StoredMinimalCandidateSelectionRow, StoredProviderModelMapping};
|
||||
|
||||
fn sample_row() -> StoredMinimalCandidateSelectionRow {
|
||||
StoredMinimalCandidateSelectionRow {
|
||||
provider_id: "provider-1".to_string(),
|
||||
provider_name: "OpenAI".to_string(),
|
||||
provider_type: "custom".to_string(),
|
||||
provider_priority: 10,
|
||||
provider_is_active: true,
|
||||
endpoint_id: "endpoint-1".to_string(),
|
||||
endpoint_api_format: "openai:chat".to_string(),
|
||||
endpoint_api_family: Some("openai".to_string()),
|
||||
endpoint_kind: Some("chat".to_string()),
|
||||
endpoint_is_active: true,
|
||||
key_id: "key-1".to_string(),
|
||||
key_name: "prod".to_string(),
|
||||
key_auth_type: "api_key".to_string(),
|
||||
key_is_active: true,
|
||||
key_api_formats: Some(vec!["openai:chat".to_string()]),
|
||||
key_allowed_models: None,
|
||||
key_capabilities: None,
|
||||
key_internal_priority: 50,
|
||||
key_global_priority_by_format: None,
|
||||
model_id: "model-1".to_string(),
|
||||
global_model_id: "global-model-1".to_string(),
|
||||
global_model_name: "gpt-4.1".to_string(),
|
||||
global_model_mappings: Some(vec!["gpt-4\\.1-.*".to_string()]),
|
||||
global_model_supports_streaming: Some(true),
|
||||
model_provider_model_name: "gpt-4.1-upstream".to_string(),
|
||||
model_provider_model_mappings: Some(vec![StoredProviderModelMapping {
|
||||
name: "gpt-4.1-canary".to_string(),
|
||||
priority: 1,
|
||||
api_formats: Some(vec!["openai:chat".to_string()]),
|
||||
}]),
|
||||
model_supports_streaming: None,
|
||||
model_is_active: true,
|
||||
model_is_available: true,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn defaults_streaming_support_to_true() {
|
||||
let mut row = sample_row();
|
||||
row.model_supports_streaming = None;
|
||||
row.global_model_supports_streaming = None;
|
||||
|
||||
assert!(row.supports_streaming());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn key_api_formats_none_means_support_all_formats() {
|
||||
let mut row = sample_row();
|
||||
row.key_api_formats = None;
|
||||
|
||||
assert!(row.key_supports_api_format("openai:chat"));
|
||||
assert!(row.key_supports_api_format("openai:responses"));
|
||||
}
|
||||
}
|
||||
@@ -3,7 +3,7 @@ use std::sync::RwLock;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::types::{
|
||||
use super::{
|
||||
PublicHealthStatusCount, PublicHealthTimelineBucket, RequestCandidateReadRepository,
|
||||
RequestCandidateStatus, RequestCandidateWriteRepository, StoredRequestCandidate,
|
||||
UpsertRequestCandidateRecord,
|
||||
|
||||
@@ -1,13 +1,13 @@
|
||||
mod memory;
|
||||
mod sql;
|
||||
mod types;
|
||||
|
||||
pub use memory::InMemoryRequestCandidateRepository;
|
||||
pub use sql::SqlxRequestCandidateReadRepository;
|
||||
pub use types::{
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use aether_data_contracts::repository::candidates::{
|
||||
build_decision_trace, derive_request_candidate_final_status, DecisionTrace,
|
||||
DecisionTraceCandidate, PublicHealthStatusCount, PublicHealthTimelineBucket,
|
||||
RequestCandidateFinalStatus, RequestCandidateReadRepository, RequestCandidateRepository,
|
||||
RequestCandidateStatus, RequestCandidateTrace, RequestCandidateWriteRepository,
|
||||
StoredRequestCandidate, UpsertRequestCandidateRecord,
|
||||
};
|
||||
pub use memory::InMemoryRequestCandidateRepository;
|
||||
pub use sql::SqlxRequestCandidateReadRepository;
|
||||
|
||||
@@ -3,13 +3,13 @@ use futures_util::future::BoxFuture;
|
||||
use sqlx::{PgPool, Row};
|
||||
use uuid::Uuid;
|
||||
|
||||
use super::types::{
|
||||
use super::{
|
||||
PublicHealthStatusCount, PublicHealthTimelineBucket, RequestCandidateReadRepository,
|
||||
RequestCandidateStatus, RequestCandidateWriteRepository, StoredRequestCandidate,
|
||||
UpsertRequestCandidateRecord,
|
||||
};
|
||||
use crate::postgres::PostgresTransactionRunner;
|
||||
use crate::DataLayerError;
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
|
||||
const LIST_BY_REQUEST_ID_SQL: &str = r#"
|
||||
SELECT
|
||||
@@ -309,7 +309,8 @@ impl SqlxRequestCandidateReadRepository {
|
||||
let rows = sqlx::query(LIST_BY_REQUEST_ID_SQL)
|
||||
.bind(request_id)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_request_candidate_row).collect()
|
||||
}
|
||||
|
||||
@@ -328,7 +329,8 @@ impl SqlxRequestCandidateReadRepository {
|
||||
))
|
||||
})?)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_request_candidate_row).collect()
|
||||
}
|
||||
|
||||
@@ -351,7 +353,8 @@ impl SqlxRequestCandidateReadRepository {
|
||||
.bind(provider_id)
|
||||
.bind(limit_value)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_request_candidate_row).collect()
|
||||
}
|
||||
|
||||
@@ -374,7 +377,8 @@ impl SqlxRequestCandidateReadRepository {
|
||||
))
|
||||
})?)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_request_candidate_row).collect()
|
||||
}
|
||||
|
||||
@@ -391,17 +395,18 @@ impl SqlxRequestCandidateReadRepository {
|
||||
.bind(endpoint_ids)
|
||||
.bind(since_unix_secs as f64)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
rows.iter()
|
||||
.map(|row| {
|
||||
let status = RequestCandidateStatus::from_database(
|
||||
row.try_get::<String, _>("status")?.as_str(),
|
||||
row_get::<String>(row, "status")?.as_str(),
|
||||
)?;
|
||||
Ok(PublicHealthStatusCount {
|
||||
endpoint_id: row.try_get("endpoint_id")?,
|
||||
endpoint_id: row_get(row, "endpoint_id")?,
|
||||
status,
|
||||
count: u64::try_from(row.try_get::<i64, _>("count")?).map_err(|_| {
|
||||
count: u64::try_from(row_get::<i64>(row, "count")?).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(
|
||||
"public health status count out of range".to_string(),
|
||||
)
|
||||
@@ -435,11 +440,12 @@ impl SqlxRequestCandidateReadRepository {
|
||||
.bind(until_unix_secs as f64)
|
||||
.bind(segment_seconds)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
rows.iter()
|
||||
.map(|row| {
|
||||
let raw_segment_idx = row.try_get::<i64, _>("segment_idx")?;
|
||||
let raw_segment_idx = row_get::<i64>(row, "segment_idx")?;
|
||||
let segment_idx = if raw_segment_idx < 0 {
|
||||
0
|
||||
} else {
|
||||
@@ -452,49 +458,53 @@ impl SqlxRequestCandidateReadRepository {
|
||||
.min(segments.saturating_sub(1));
|
||||
|
||||
Ok(PublicHealthTimelineBucket {
|
||||
endpoint_id: row.try_get("endpoint_id")?,
|
||||
endpoint_id: row_get(row, "endpoint_id")?,
|
||||
segment_idx,
|
||||
total_count: u64::try_from(row.try_get::<i64, _>("total_count")?).map_err(
|
||||
total_count: u64::try_from(row_get::<i64>(row, "total_count")?).map_err(
|
||||
|_| {
|
||||
DataLayerError::UnexpectedValue(
|
||||
"public health total_count out of range".to_string(),
|
||||
)
|
||||
},
|
||||
)?,
|
||||
success_count: u64::try_from(row.try_get::<i64, _>("success_count")?).map_err(
|
||||
success_count: u64::try_from(row_get::<i64>(row, "success_count")?).map_err(
|
||||
|_| {
|
||||
DataLayerError::UnexpectedValue(
|
||||
"public health success_count out of range".to_string(),
|
||||
)
|
||||
},
|
||||
)?,
|
||||
failed_count: u64::try_from(row.try_get::<i64, _>("failed_count")?).map_err(
|
||||
failed_count: u64::try_from(row_get::<i64>(row, "failed_count")?).map_err(
|
||||
|_| {
|
||||
DataLayerError::UnexpectedValue(
|
||||
"public health failed_count out of range".to_string(),
|
||||
)
|
||||
},
|
||||
)?,
|
||||
min_created_at_unix_secs: row
|
||||
.try_get::<Option<i64>, _>("min_created_at_unix_secs")?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"public health min_created_at_unix_secs out of range: {value}"
|
||||
))
|
||||
})
|
||||
min_created_at_unix_secs: row_get::<Option<i64>>(
|
||||
row,
|
||||
"min_created_at_unix_secs",
|
||||
)?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"public health min_created_at_unix_secs out of range: {value}"
|
||||
))
|
||||
})
|
||||
.transpose()?,
|
||||
max_created_at_unix_secs: row
|
||||
.try_get::<Option<i64>, _>("max_created_at_unix_secs")?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"public health max_created_at_unix_secs out of range: {value}"
|
||||
))
|
||||
})
|
||||
})
|
||||
.transpose()?,
|
||||
max_created_at_unix_secs: row_get::<Option<i64>>(
|
||||
row,
|
||||
"max_created_at_unix_secs",
|
||||
)?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"public health max_created_at_unix_secs out of range: {value}"
|
||||
))
|
||||
})
|
||||
.transpose()?,
|
||||
})
|
||||
.transpose()?,
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
@@ -538,7 +548,8 @@ impl SqlxRequestCandidateReadRepository {
|
||||
.bind(candidate.started_at_unix_secs.map(|value| value as f64))
|
||||
.bind(candidate.finished_at_unix_secs.map(|value| value as f64))
|
||||
.fetch_one(&mut **tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
map_request_candidate_row(&row)
|
||||
}) as BoxFuture<'_, Result<StoredRequestCandidate, DataLayerError>>
|
||||
})
|
||||
@@ -562,7 +573,8 @@ impl SqlxRequestCandidateReadRepository {
|
||||
))
|
||||
})?)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(result.rows_affected() as usize)
|
||||
}
|
||||
}
|
||||
@@ -648,36 +660,42 @@ impl RequestCandidateWriteRepository for SqlxRequestCandidateReadRepository {
|
||||
fn map_request_candidate_row(
|
||||
row: &sqlx::postgres::PgRow,
|
||||
) -> Result<StoredRequestCandidate, DataLayerError> {
|
||||
let status =
|
||||
RequestCandidateStatus::from_database(row.try_get::<String, _>("status")?.as_str())?;
|
||||
let status = RequestCandidateStatus::from_database(row_get::<String>(row, "status")?.as_str())?;
|
||||
StoredRequestCandidate::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("request_id")?,
|
||||
row.try_get("user_id")?,
|
||||
row.try_get("api_key_id")?,
|
||||
row.try_get("username")?,
|
||||
row.try_get("api_key_name")?,
|
||||
row.try_get("candidate_index")?,
|
||||
row.try_get("retry_index")?,
|
||||
row.try_get("provider_id")?,
|
||||
row.try_get("endpoint_id")?,
|
||||
row.try_get("key_id")?,
|
||||
row_get(row, "id")?,
|
||||
row_get(row, "request_id")?,
|
||||
row_get(row, "user_id")?,
|
||||
row_get(row, "api_key_id")?,
|
||||
row_get(row, "username")?,
|
||||
row_get(row, "api_key_name")?,
|
||||
row_get(row, "candidate_index")?,
|
||||
row_get(row, "retry_index")?,
|
||||
row_get(row, "provider_id")?,
|
||||
row_get(row, "endpoint_id")?,
|
||||
row_get(row, "key_id")?,
|
||||
status,
|
||||
row.try_get("skip_reason")?,
|
||||
row.try_get("is_cached")?,
|
||||
row.try_get("status_code")?,
|
||||
row.try_get("error_type")?,
|
||||
row.try_get("error_message")?,
|
||||
row.try_get("latency_ms")?,
|
||||
row.try_get("concurrent_requests")?,
|
||||
row.try_get("extra_data")?,
|
||||
row.try_get("required_capabilities")?,
|
||||
row.try_get("created_at_unix_secs")?,
|
||||
row.try_get("started_at_unix_secs")?,
|
||||
row.try_get("finished_at_unix_secs")?,
|
||||
row_get(row, "skip_reason")?,
|
||||
row_get(row, "is_cached")?,
|
||||
row_get(row, "status_code")?,
|
||||
row_get(row, "error_type")?,
|
||||
row_get(row, "error_message")?,
|
||||
row_get(row, "latency_ms")?,
|
||||
row_get(row, "concurrent_requests")?,
|
||||
row_get(row, "extra_data")?,
|
||||
row_get(row, "required_capabilities")?,
|
||||
row_get(row, "created_at_unix_secs")?,
|
||||
row_get(row, "started_at_unix_secs")?,
|
||||
row_get(row, "finished_at_unix_secs")?,
|
||||
)
|
||||
}
|
||||
|
||||
fn row_get<T>(row: &sqlx::postgres::PgRow, column: &str) -> Result<T, DataLayerError>
|
||||
where
|
||||
for<'r> T: sqlx::Decode<'r, sqlx::Postgres> + sqlx::Type<sqlx::Postgres>,
|
||||
{
|
||||
row.try_get(column).map_postgres_err()
|
||||
}
|
||||
|
||||
fn status_to_database(status: RequestCandidateStatus) -> &'static str {
|
||||
match status {
|
||||
RequestCandidateStatus::Available => "available",
|
||||
|
||||
@@ -1,825 +0,0 @@
|
||||
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_secs: Option<u64>) -> bool {
|
||||
match self {
|
||||
Self::Available | Self::Unused | Self::Skipped => false,
|
||||
Self::Pending => started_at_unix_secs.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_secs: u64,
|
||||
pub started_at_unix_secs: Option<u64>,
|
||||
pub finished_at_unix_secs: 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_secs: i64,
|
||||
started_at_unix_secs: Option<i64>,
|
||||
finished_at_unix_secs: 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_secs = u64::try_from(created_at_unix_secs).map_err(|_| {
|
||||
crate::DataLayerError::UnexpectedValue(format!(
|
||||
"invalid request_candidates.created_at_unix_secs: {created_at_unix_secs}"
|
||||
))
|
||||
})?;
|
||||
let started_at_unix_secs = started_at_unix_secs
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
crate::DataLayerError::UnexpectedValue(format!(
|
||||
"invalid request_candidates.started_at_unix_secs: {value}"
|
||||
))
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let finished_at_unix_secs = finished_at_unix_secs
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
crate::DataLayerError::UnexpectedValue(format!(
|
||||
"invalid request_candidates.finished_at_unix_secs: {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_secs,
|
||||
started_at_unix_secs,
|
||||
finished_at_unix_secs,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[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_secs)
|
||||
})
|
||||
.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
|
||||
|| matches!(candidate.status_code, Some(status_code) if (200..300).contains(&status_code))
|
||||
});
|
||||
if has_success {
|
||||
return RequestCandidateFinalStatus::Success;
|
||||
}
|
||||
|
||||
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_cancelled = candidates
|
||||
.iter()
|
||||
.any(|candidate| candidate.status == RequestCandidateStatus::Cancelled);
|
||||
let has_failed = candidates
|
||||
.iter()
|
||||
.any(|candidate| candidate.status == RequestCandidateStatus::Failed);
|
||||
if has_cancelled && !has_failed {
|
||||
return RequestCandidateFinalStatus::Cancelled;
|
||||
}
|
||||
|
||||
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 endpoint_api_format: Option<String>,
|
||||
pub endpoint_api_family: Option<String>,
|
||||
pub endpoint_kind: Option<String>,
|
||||
pub provider_key_name: Option<String>,
|
||||
pub provider_key_auth_type: Option<String>,
|
||||
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()),
|
||||
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()),
|
||||
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_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_secs: Option<u64>,
|
||||
pub max_created_at_unix_secs: 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_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_secs: Option<u64>,
|
||||
pub started_at_unix_secs: Option<u64>,
|
||||
pub finished_at_unix_secs: 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 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
|
||||
{
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
build_decision_trace, derive_request_candidate_final_status, DecisionTrace,
|
||||
DecisionTraceCandidate, RequestCandidateFinalStatus, RequestCandidateStatus,
|
||||
RequestCandidateTrace, StoredRequestCandidate, UpsertRequestCandidateRecord,
|
||||
};
|
||||
use crate::repository::provider_catalog::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn parses_status_from_database_text() {
|
||||
assert_eq!(
|
||||
RequestCandidateStatus::from_database("streaming").expect("status should parse"),
|
||||
RequestCandidateStatus::Streaming
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_database_status() {
|
||||
assert!(RequestCandidateStatus::from_database("mystery").is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_negative_candidate_index() {
|
||||
assert!(StoredRequestCandidate::new(
|
||||
"cand-1".to_string(),
|
||||
"req-1".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
-1,
|
||||
0,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
RequestCandidateStatus::Pending,
|
||||
None,
|
||||
false,
|
||||
Some(200),
|
||||
None,
|
||||
None,
|
||||
Some(10),
|
||||
Some(1),
|
||||
None,
|
||||
None,
|
||||
100,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
fn sample_candidate(
|
||||
id: &str,
|
||||
request_id: &str,
|
||||
candidate_index: i32,
|
||||
status: RequestCandidateStatus,
|
||||
started_at_unix_secs: Option<i64>,
|
||||
latency_ms: Option<i32>,
|
||||
status_code: Option<i32>,
|
||||
) -> StoredRequestCandidate {
|
||||
StoredRequestCandidate::new(
|
||||
id.to_string(),
|
||||
request_id.to_string(),
|
||||
Some("user-1".to_string()),
|
||||
Some("api-key-1".to_string()),
|
||||
Some("alice".to_string()),
|
||||
Some("default".to_string()),
|
||||
candidate_index,
|
||||
0,
|
||||
Some("provider-1".to_string()),
|
||||
Some("endpoint-1".to_string()),
|
||||
Some("provider-key-1".to_string()),
|
||||
status,
|
||||
None,
|
||||
false,
|
||||
status_code,
|
||||
None,
|
||||
None,
|
||||
latency_ms,
|
||||
Some(1),
|
||||
None,
|
||||
None,
|
||||
100 + i64::from(candidate_index),
|
||||
started_at_unix_secs,
|
||||
started_at_unix_secs.map(|value| value + 1),
|
||||
)
|
||||
.expect("candidate should build")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn derives_request_candidate_final_status_preferring_success() {
|
||||
let candidates = vec![sample_candidate(
|
||||
"cand-1",
|
||||
"req-1",
|
||||
0,
|
||||
RequestCandidateStatus::Success,
|
||||
Some(100),
|
||||
Some(25),
|
||||
Some(200),
|
||||
)];
|
||||
|
||||
assert_eq!(
|
||||
derive_request_candidate_final_status(&candidates),
|
||||
RequestCandidateFinalStatus::Success
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_candidate_trace_filters_attempted_rows() {
|
||||
let trace = RequestCandidateTrace::from_candidates(
|
||||
"req-1",
|
||||
vec![
|
||||
sample_candidate(
|
||||
"cand-1",
|
||||
"req-1",
|
||||
0,
|
||||
RequestCandidateStatus::Pending,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
),
|
||||
sample_candidate(
|
||||
"cand-2",
|
||||
"req-1",
|
||||
1,
|
||||
RequestCandidateStatus::Failed,
|
||||
Some(101),
|
||||
Some(33),
|
||||
Some(502),
|
||||
),
|
||||
],
|
||||
true,
|
||||
)
|
||||
.expect("trace should exist");
|
||||
|
||||
assert_eq!(trace.total_candidates, 1);
|
||||
assert_eq!(trace.candidates[0].id, "cand-2");
|
||||
assert_eq!(trace.final_status, RequestCandidateFinalStatus::Failed);
|
||||
assert_eq!(trace.total_latency_ms, 33);
|
||||
}
|
||||
|
||||
fn sample_provider() -> StoredProviderCatalogProvider {
|
||||
StoredProviderCatalogProvider::new(
|
||||
"provider-1".to_string(),
|
||||
"OpenAI".to_string(),
|
||||
Some("https://openai.com".to_string()),
|
||||
"custom".to_string(),
|
||||
)
|
||||
.expect("provider should build")
|
||||
}
|
||||
|
||||
fn sample_endpoint() -> StoredProviderCatalogEndpoint {
|
||||
StoredProviderCatalogEndpoint::new(
|
||||
"endpoint-1".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"openai:chat".to_string(),
|
||||
Some("openai".to_string()),
|
||||
Some("chat".to_string()),
|
||||
true,
|
||||
)
|
||||
.expect("endpoint should build")
|
||||
}
|
||||
|
||||
fn sample_key() -> StoredProviderCatalogKey {
|
||||
StoredProviderCatalogKey::new(
|
||||
"provider-key-1".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"prod-key".to_string(),
|
||||
"api_key".to_string(),
|
||||
Some(serde_json::json!({"cache_1h": true})),
|
||||
true,
|
||||
)
|
||||
.expect("key should build")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_decision_trace_enriches_candidate_with_provider_catalog_metadata() {
|
||||
let trace = RequestCandidateTrace::from_candidates(
|
||||
"req-1",
|
||||
vec![sample_candidate(
|
||||
"cand-1",
|
||||
"req-1",
|
||||
0,
|
||||
RequestCandidateStatus::Failed,
|
||||
Some(101),
|
||||
Some(37),
|
||||
Some(502),
|
||||
)],
|
||||
true,
|
||||
)
|
||||
.expect("trace should exist");
|
||||
|
||||
assert_eq!(
|
||||
build_decision_trace(
|
||||
trace,
|
||||
vec![sample_provider()],
|
||||
vec![sample_endpoint()],
|
||||
vec![sample_key()],
|
||||
),
|
||||
DecisionTrace {
|
||||
request_id: "req-1".to_string(),
|
||||
total_candidates: 1,
|
||||
final_status: RequestCandidateFinalStatus::Failed,
|
||||
total_latency_ms: 37,
|
||||
candidates: vec![DecisionTraceCandidate {
|
||||
candidate: sample_candidate(
|
||||
"cand-1",
|
||||
"req-1",
|
||||
0,
|
||||
RequestCandidateStatus::Failed,
|
||||
Some(101),
|
||||
Some(37),
|
||||
Some(502),
|
||||
),
|
||||
provider_name: Some("OpenAI".to_string()),
|
||||
provider_website: Some("https://openai.com".to_string()),
|
||||
provider_type: Some("custom".to_string()),
|
||||
endpoint_api_format: Some("openai:chat".to_string()),
|
||||
endpoint_api_family: Some("openai".to_string()),
|
||||
endpoint_kind: Some("chat".to_string()),
|
||||
provider_key_name: Some("prod-key".to_string()),
|
||||
provider_key_auth_type: Some("api_key".to_string()),
|
||||
provider_key_capabilities: Some(serde_json::json!({"cache_1h": true})),
|
||||
provider_key_is_active: Some(true),
|
||||
}],
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_negative_created_at() {
|
||||
assert!(StoredRequestCandidate::new(
|
||||
"cand-1".to_string(),
|
||||
"req-1".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
0,
|
||||
0,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
RequestCandidateStatus::Pending,
|
||||
None,
|
||||
false,
|
||||
Some(200),
|
||||
None,
|
||||
None,
|
||||
Some(10),
|
||||
Some(1),
|
||||
None,
|
||||
None,
|
||||
-1,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn pending_without_started_at_is_not_attempted() {
|
||||
assert!(!RequestCandidateStatus::Pending.is_attempted(None));
|
||||
assert!(RequestCandidateStatus::Pending.is_attempted(Some(1)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_upsert_payload() {
|
||||
assert!(UpsertRequestCandidateRecord {
|
||||
id: "".to_string(),
|
||||
request_id: "".to_string(),
|
||||
user_id: None,
|
||||
api_key_id: None,
|
||||
username: None,
|
||||
api_key_name: None,
|
||||
candidate_index: 0,
|
||||
retry_index: 0,
|
||||
provider_id: None,
|
||||
endpoint_id: None,
|
||||
key_id: None,
|
||||
status: RequestCandidateStatus::Available,
|
||||
skip_reason: None,
|
||||
is_cached: None,
|
||||
status_code: None,
|
||||
error_type: None,
|
||||
error_message: None,
|
||||
latency_ms: None,
|
||||
concurrent_requests: None,
|
||||
extra_data: None,
|
||||
required_capabilities: None,
|
||||
created_at_unix_secs: None,
|
||||
started_at_unix_secs: None,
|
||||
finished_at_unix_secs: None,
|
||||
}
|
||||
.validate()
|
||||
.is_err());
|
||||
}
|
||||
}
|
||||
@@ -6,7 +6,7 @@ use super::types::{
|
||||
GeminiFileMappingStats, GeminiFileMappingWriteRepository, StoredGeminiFileMapping,
|
||||
StoredGeminiFileMappingListPage, UpsertGeminiFileMappingRecord,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SqlxGeminiFileMappingRepository {
|
||||
@@ -20,25 +20,31 @@ impl SqlxGeminiFileMappingRepository {
|
||||
|
||||
fn map_row(row: &PgRow) -> Result<StoredGeminiFileMapping, DataLayerError> {
|
||||
Ok(StoredGeminiFileMapping {
|
||||
id: row.try_get("id")?,
|
||||
file_name: row.try_get("file_name")?,
|
||||
key_id: row.try_get("key_id")?,
|
||||
id: row.try_get("id").map_postgres_err()?,
|
||||
file_name: row.try_get("file_name").map_postgres_err()?,
|
||||
key_id: row.try_get("key_id").map_postgres_err()?,
|
||||
user_id: row.try_get("user_id").ok().flatten(),
|
||||
display_name: row.try_get("display_name").ok().flatten(),
|
||||
mime_type: row.try_get("mime_type").ok().flatten(),
|
||||
source_hash: row.try_get("source_hash").ok().flatten(),
|
||||
created_at_unix_secs: u64::try_from(row.try_get::<i64, _>("created_at_unix_secs")?)
|
||||
.map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(
|
||||
"gemini_file_mappings.created_at is invalid".to_string(),
|
||||
)
|
||||
})?,
|
||||
expires_at_unix_secs: u64::try_from(row.try_get::<i64, _>("expires_at_unix_secs")?)
|
||||
.map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(
|
||||
"gemini_file_mappings.expires_at is invalid".to_string(),
|
||||
)
|
||||
})?,
|
||||
created_at_unix_secs: u64::try_from(
|
||||
row.try_get::<i64, _>("created_at_unix_secs")
|
||||
.map_postgres_err()?,
|
||||
)
|
||||
.map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(
|
||||
"gemini_file_mappings.created_at is invalid".to_string(),
|
||||
)
|
||||
})?,
|
||||
expires_at_unix_secs: u64::try_from(
|
||||
row.try_get::<i64, _>("expires_at_unix_secs")
|
||||
.map_postgres_err()?,
|
||||
)
|
||||
.map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(
|
||||
"gemini_file_mappings.expires_at is invalid".to_string(),
|
||||
)
|
||||
})?,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -67,7 +73,8 @@ WHERE file_name = $1
|
||||
)
|
||||
.bind(file_name)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
match row {
|
||||
Some(row) => Ok(Some(Self::map_row(&row)?)),
|
||||
@@ -82,11 +89,13 @@ WHERE file_name = $1
|
||||
let total = build_list_count_query(query)
|
||||
.build_query_scalar::<i64>()
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let rows = build_list_rows_query(query)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(StoredGeminiFileMappingListPage {
|
||||
items: rows
|
||||
.iter()
|
||||
@@ -110,11 +119,20 @@ FROM gemini_file_mappings
|
||||
)
|
||||
.bind(now_unix_secs as f64)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
let total_mappings =
|
||||
usize::try_from(totals.try_get::<i64, _>("total_mappings")?).unwrap_or_default();
|
||||
let active_mappings =
|
||||
usize::try_from(totals.try_get::<i64, _>("active_mappings")?).unwrap_or_default();
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let total_mappings = usize::try_from(
|
||||
totals
|
||||
.try_get::<i64, _>("total_mappings")
|
||||
.map_postgres_err()?,
|
||||
)
|
||||
.unwrap_or_default();
|
||||
let active_mappings = usize::try_from(
|
||||
totals
|
||||
.try_get::<i64, _>("active_mappings")
|
||||
.map_postgres_err()?,
|
||||
)
|
||||
.unwrap_or_default();
|
||||
let by_mime_type_rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
@@ -128,7 +146,8 @@ ORDER BY mime_type ASC
|
||||
)
|
||||
.bind(now_unix_secs as f64)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(GeminiFileMappingStats {
|
||||
total_mappings,
|
||||
active_mappings,
|
||||
@@ -137,8 +156,9 @@ ORDER BY mime_type ASC
|
||||
.into_iter()
|
||||
.map(|row| {
|
||||
Ok(GeminiFileMappingMimeTypeCount {
|
||||
mime_type: row.try_get("mime_type")?,
|
||||
count: usize::try_from(row.try_get::<i64, _>("count")?).unwrap_or_default(),
|
||||
mime_type: row.try_get("mime_type").map_postgres_err()?,
|
||||
count: usize::try_from(row.try_get::<i64, _>("count").map_postgres_err()?)
|
||||
.unwrap_or_default(),
|
||||
})
|
||||
})
|
||||
.collect::<Result<Vec<_>, DataLayerError>>()?,
|
||||
@@ -197,7 +217,8 @@ RETURNING
|
||||
.bind(record.source_hash.clone())
|
||||
.bind(record.expires_at_unix_secs as f64)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
Self::map_row(&row)
|
||||
}
|
||||
@@ -211,7 +232,8 @@ WHERE file_name = $1
|
||||
)
|
||||
.bind(file_name)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
Ok(result.rows_affected() > 0)
|
||||
}
|
||||
@@ -238,7 +260,8 @@ RETURNING
|
||||
)
|
||||
.bind(mapping_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
match row {
|
||||
Some(row) => Ok(Some(Self::map_row(&row)?)),
|
||||
@@ -255,7 +278,8 @@ WHERE expires_at <= TO_TIMESTAMP($1::double precision)
|
||||
)
|
||||
.bind(now_unix_secs as f64)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
Ok(usize::try_from(result.rows_affected()).unwrap_or_default())
|
||||
}
|
||||
|
||||
@@ -2,7 +2,7 @@ use std::sync::RwLock;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::types::{
|
||||
use super::{
|
||||
AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord,
|
||||
GlobalModelReadRepository, GlobalModelWriteRepository, PublicCatalogModelListQuery,
|
||||
PublicCatalogModelSearchQuery, PublicGlobalModelQuery, StoredAdminGlobalModel,
|
||||
|
||||
@@ -1,10 +1,8 @@
|
||||
mod memory;
|
||||
mod sql;
|
||||
mod types;
|
||||
|
||||
pub use memory::InMemoryGlobalModelReadRepository;
|
||||
pub use sql::SqlxGlobalModelReadRepository;
|
||||
pub use types::{
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use aether_data_contracts::repository::global_models::{
|
||||
AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord,
|
||||
GlobalModelReadRepository, GlobalModelWriteRepository, PublicCatalogModelListQuery,
|
||||
PublicCatalogModelSearchQuery, PublicGlobalModelQuery, StoredAdminGlobalModel,
|
||||
@@ -12,3 +10,5 @@ pub use types::{
|
||||
StoredProviderModelStats, StoredPublicCatalogModel, StoredPublicGlobalModel,
|
||||
StoredPublicGlobalModelPage, UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
|
||||
};
|
||||
pub use memory::InMemoryGlobalModelReadRepository;
|
||||
pub use sql::SqlxGlobalModelReadRepository;
|
||||
|
||||
@@ -2,7 +2,7 @@ use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
use sqlx::{postgres::PgRow, PgPool, Postgres, QueryBuilder, Row};
|
||||
|
||||
use super::types::{
|
||||
use super::{
|
||||
AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord,
|
||||
GlobalModelReadRepository, GlobalModelWriteRepository, PublicCatalogModelListQuery,
|
||||
PublicCatalogModelSearchQuery, PublicGlobalModelQuery, StoredAdminGlobalModel,
|
||||
@@ -10,7 +10,7 @@ use super::types::{
|
||||
StoredProviderModelStats, StoredPublicCatalogModel, StoredPublicGlobalModel,
|
||||
StoredPublicGlobalModelPage, UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
|
||||
const LIST_PUBLIC_GLOBAL_MODELS_PREFIX: &str = r#"
|
||||
SELECT
|
||||
@@ -146,10 +146,15 @@ impl SqlxGlobalModelReadRepository {
|
||||
) -> Result<StoredPublicGlobalModelPage, DataLayerError> {
|
||||
let mut count_builder = QueryBuilder::<Postgres>::new(COUNT_PUBLIC_GLOBAL_MODELS_PREFIX);
|
||||
apply_public_model_filters(&mut count_builder, query);
|
||||
let count_row = count_builder.build().fetch_one(&self.pool).await?;
|
||||
let count_row = count_builder
|
||||
.build()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let total = count_row
|
||||
.try_get::<i64, _>("total")
|
||||
.map(|value| value.max(0) as usize)?;
|
||||
.map(|value| value.max(0) as usize)
|
||||
.map_postgres_err()?;
|
||||
|
||||
let mut list_builder = QueryBuilder::<Postgres>::new(LIST_PUBLIC_GLOBAL_MODELS_PREFIX);
|
||||
apply_public_model_filters(&mut list_builder, query);
|
||||
@@ -158,7 +163,11 @@ impl SqlxGlobalModelReadRepository {
|
||||
.push_bind(query.offset as i64)
|
||||
.push(" LIMIT ")
|
||||
.push_bind(query.limit as i64);
|
||||
let rows = list_builder.build().fetch_all(&self.pool).await?;
|
||||
let rows = list_builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let items = rows.iter().map(map_row).collect::<Result<Vec<_>, _>>()?;
|
||||
|
||||
Ok(StoredPublicGlobalModelPage { items, total })
|
||||
@@ -179,7 +188,8 @@ impl SqlxGlobalModelReadRepository {
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
rows.iter().map(map_provider_model_stats_row).collect()
|
||||
}
|
||||
@@ -199,7 +209,8 @@ impl SqlxGlobalModelReadRepository {
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
rows.iter()
|
||||
.map(map_provider_active_global_model_row)
|
||||
@@ -222,7 +233,11 @@ impl SqlxGlobalModelReadRepository {
|
||||
.push_bind(query.offset as i64)
|
||||
.push(" LIMIT ")
|
||||
.push_bind(query.limit as i64);
|
||||
let rows = builder.build().fetch_all(&self.pool).await?;
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_admin_provider_model_row).collect()
|
||||
}
|
||||
|
||||
@@ -232,10 +247,15 @@ impl SqlxGlobalModelReadRepository {
|
||||
) -> Result<StoredAdminGlobalModelPage, DataLayerError> {
|
||||
let mut count_builder = QueryBuilder::<Postgres>::new(COUNT_ADMIN_GLOBAL_MODELS_PREFIX);
|
||||
apply_admin_global_model_filters(&mut count_builder, query);
|
||||
let count_row = count_builder.build().fetch_one(&self.pool).await?;
|
||||
let count_row = count_builder
|
||||
.build()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let total = count_row
|
||||
.try_get::<i64, _>("total")
|
||||
.map(|value| value.max(0) as usize)?;
|
||||
.map(|value| value.max(0) as usize)
|
||||
.map_postgres_err()?;
|
||||
|
||||
let mut list_builder = QueryBuilder::<Postgres>::new(LIST_ADMIN_GLOBAL_MODELS_PREFIX);
|
||||
apply_admin_global_model_filters(&mut list_builder, query);
|
||||
@@ -244,7 +264,11 @@ impl SqlxGlobalModelReadRepository {
|
||||
.push_bind(query.offset as i64)
|
||||
.push(" LIMIT ")
|
||||
.push_bind(query.limit as i64);
|
||||
let rows = list_builder.build().fetch_all(&self.pool).await?;
|
||||
let rows = list_builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let items = rows
|
||||
.iter()
|
||||
.map(map_admin_global_model_row)
|
||||
@@ -292,7 +316,8 @@ LIMIT 1
|
||||
.bind(provider_id)
|
||||
.bind(model_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
row.as_ref().map(map_admin_provider_model_row).transpose()
|
||||
}
|
||||
@@ -336,7 +361,8 @@ ORDER BY gm.name ASC, m.created_at DESC, m.id ASC
|
||||
)
|
||||
.bind(provider_id)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
rows.iter().map(map_admin_provider_model_row).collect()
|
||||
}
|
||||
@@ -366,7 +392,8 @@ LIMIT 1
|
||||
)
|
||||
.bind(global_model_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
row.as_ref().map(map_admin_global_model_row).transpose()
|
||||
}
|
||||
@@ -396,7 +423,8 @@ LIMIT 1
|
||||
)
|
||||
.bind(model_name)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
row.as_ref().map(map_admin_global_model_row).transpose()
|
||||
}
|
||||
@@ -438,7 +466,8 @@ ORDER BY m.created_at DESC, m.id ASC
|
||||
)
|
||||
.bind(global_model_id)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
rows.iter().map(map_admin_provider_model_row).collect()
|
||||
}
|
||||
@@ -490,7 +519,8 @@ RETURNING id
|
||||
.bind(record.is_available)
|
||||
.bind(record.config.clone())
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
if inserted.is_none() {
|
||||
return Ok(None);
|
||||
@@ -543,7 +573,8 @@ RETURNING id
|
||||
.bind(record.is_available)
|
||||
.bind(record.config.clone())
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
if updated.is_none() {
|
||||
return Ok(None);
|
||||
@@ -569,7 +600,8 @@ RETURNING id
|
||||
.bind(provider_id)
|
||||
.bind(model_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
Ok(deleted.is_some())
|
||||
}
|
||||
@@ -605,7 +637,8 @@ RETURNING id
|
||||
.bind(record.supported_capabilities.clone())
|
||||
.bind(record.config.clone())
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
if inserted.is_none() {
|
||||
return Ok(None);
|
||||
@@ -641,7 +674,8 @@ RETURNING id
|
||||
.bind(record.supported_capabilities.clone())
|
||||
.bind(record.config.clone())
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
if updated.is_none() {
|
||||
return Ok(None);
|
||||
@@ -663,7 +697,8 @@ RETURNING id
|
||||
)
|
||||
.bind(global_model_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
Ok(deleted.is_some())
|
||||
}
|
||||
@@ -700,7 +735,8 @@ LIMIT 1
|
||||
)
|
||||
.bind(model_name)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
row.as_ref().map(map_row).transpose()
|
||||
}
|
||||
@@ -716,7 +752,11 @@ LIMIT 1
|
||||
.push_bind(query.offset as i64)
|
||||
.push(" LIMIT ")
|
||||
.push_bind(query.limit as i64);
|
||||
let rows = builder.build().fetch_all(&self.pool).await?;
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_public_catalog_model_row).collect()
|
||||
}
|
||||
|
||||
@@ -733,7 +773,11 @@ LIMIT 1
|
||||
builder
|
||||
.push(" ORDER BY p.provider_priority ASC, p.name ASC, COALESCE(gm.name, m.provider_model_name) ASC, m.id ASC LIMIT ")
|
||||
.push_bind(query.limit as i64);
|
||||
let rows = builder.build().fetch_all(&self.pool).await?;
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_public_catalog_model_row).collect()
|
||||
}
|
||||
|
||||
@@ -903,16 +947,18 @@ fn apply_admin_global_model_filters(
|
||||
}
|
||||
|
||||
fn map_row(row: &PgRow) -> Result<StoredPublicGlobalModel, DataLayerError> {
|
||||
let supported_capabilities: Option<Value> = row.try_get("supported_capabilities")?;
|
||||
let supported_capabilities: Option<Value> =
|
||||
row.try_get("supported_capabilities").map_postgres_err()?;
|
||||
StoredPublicGlobalModel::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("name")?,
|
||||
row.try_get("display_name")?,
|
||||
row.try_get("is_active")?,
|
||||
row.try_get("default_price_per_request")?,
|
||||
row.try_get("default_tiered_pricing")?,
|
||||
row.try_get("id").map_postgres_err()?,
|
||||
row.try_get("name").map_postgres_err()?,
|
||||
row.try_get("display_name").map_postgres_err()?,
|
||||
row.try_get("is_active").map_postgres_err()?,
|
||||
row.try_get("default_price_per_request")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("default_tiered_pricing").map_postgres_err()?,
|
||||
supported_capabilities,
|
||||
row.try_get("config")?,
|
||||
row.try_get("config").map_postgres_err()?,
|
||||
0,
|
||||
)
|
||||
}
|
||||
@@ -945,74 +991,87 @@ fn apply_public_catalog_model_filters(
|
||||
|
||||
fn map_public_catalog_model_row(row: &PgRow) -> Result<StoredPublicCatalogModel, DataLayerError> {
|
||||
StoredPublicCatalogModel::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("provider_id")?,
|
||||
row.try_get("provider_name")?,
|
||||
row.try_get("provider_model_name")?,
|
||||
row.try_get("name")?,
|
||||
row.try_get("display_name")?,
|
||||
row.try_get("description")?,
|
||||
row.try_get("icon_url")?,
|
||||
row.try_get("input_price_per_1m")?,
|
||||
row.try_get("output_price_per_1m")?,
|
||||
row.try_get("cache_creation_price_per_1m")?,
|
||||
row.try_get("cache_read_price_per_1m")?,
|
||||
row.try_get("supports_vision")?,
|
||||
row.try_get("supports_function_calling")?,
|
||||
row.try_get("supports_streaming")?,
|
||||
row.try_get("is_active")?,
|
||||
row.try_get("id").map_postgres_err()?,
|
||||
row.try_get("provider_id").map_postgres_err()?,
|
||||
row.try_get("provider_name").map_postgres_err()?,
|
||||
row.try_get("provider_model_name").map_postgres_err()?,
|
||||
row.try_get("name").map_postgres_err()?,
|
||||
row.try_get("display_name").map_postgres_err()?,
|
||||
row.try_get("description").map_postgres_err()?,
|
||||
row.try_get("icon_url").map_postgres_err()?,
|
||||
row.try_get("input_price_per_1m").map_postgres_err()?,
|
||||
row.try_get("output_price_per_1m").map_postgres_err()?,
|
||||
row.try_get("cache_creation_price_per_1m")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("cache_read_price_per_1m").map_postgres_err()?,
|
||||
row.try_get("supports_vision").map_postgres_err()?,
|
||||
row.try_get("supports_function_calling")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("supports_streaming").map_postgres_err()?,
|
||||
row.try_get("is_active").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
fn map_admin_provider_model_row(row: &PgRow) -> Result<StoredAdminProviderModel, DataLayerError> {
|
||||
let created_at_unix_secs = row
|
||||
.try_get::<Option<i64>, _>("created_at_unix_secs")?
|
||||
.try_get::<Option<i64>, _>("created_at_unix_secs")
|
||||
.map_postgres_err()?
|
||||
.map(|value| value.max(0) as u64);
|
||||
let updated_at_unix_secs = row
|
||||
.try_get::<Option<i64>, _>("updated_at_unix_secs")?
|
||||
.try_get::<Option<i64>, _>("updated_at_unix_secs")
|
||||
.map_postgres_err()?
|
||||
.map(|value| value.max(0) as u64);
|
||||
StoredAdminProviderModel::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("provider_id")?,
|
||||
row.try_get("global_model_id")?,
|
||||
row.try_get("provider_model_name")?,
|
||||
row.try_get("provider_model_mappings")?,
|
||||
row.try_get("price_per_request")?,
|
||||
row.try_get("tiered_pricing")?,
|
||||
row.try_get("supports_vision")?,
|
||||
row.try_get("supports_function_calling")?,
|
||||
row.try_get("supports_streaming")?,
|
||||
row.try_get("supports_extended_thinking")?,
|
||||
row.try_get("supports_image_generation")?,
|
||||
row.try_get("is_active")?,
|
||||
row.try_get("is_available")?,
|
||||
row.try_get("config")?,
|
||||
row.try_get("id").map_postgres_err()?,
|
||||
row.try_get("provider_id").map_postgres_err()?,
|
||||
row.try_get("global_model_id").map_postgres_err()?,
|
||||
row.try_get("provider_model_name").map_postgres_err()?,
|
||||
row.try_get("provider_model_mappings").map_postgres_err()?,
|
||||
row.try_get("price_per_request").map_postgres_err()?,
|
||||
row.try_get("tiered_pricing").map_postgres_err()?,
|
||||
row.try_get("supports_vision").map_postgres_err()?,
|
||||
row.try_get("supports_function_calling")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("supports_streaming").map_postgres_err()?,
|
||||
row.try_get("supports_extended_thinking")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("supports_image_generation")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("is_active").map_postgres_err()?,
|
||||
row.try_get("is_available").map_postgres_err()?,
|
||||
row.try_get("config").map_postgres_err()?,
|
||||
created_at_unix_secs,
|
||||
updated_at_unix_secs,
|
||||
row.try_get("global_model_name")?,
|
||||
row.try_get("global_model_display_name")?,
|
||||
row.try_get("global_model_default_price_per_request")?,
|
||||
row.try_get("global_model_default_tiered_pricing")?,
|
||||
row.try_get("global_model_config")?,
|
||||
row.try_get("global_model_name").map_postgres_err()?,
|
||||
row.try_get("global_model_display_name")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("global_model_default_price_per_request")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("global_model_default_tiered_pricing")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("global_model_config").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
fn map_admin_global_model_row(row: &PgRow) -> Result<StoredAdminGlobalModel, DataLayerError> {
|
||||
let created_at_unix_secs = row
|
||||
.try_get::<Option<i64>, _>("created_at_unix_secs")?
|
||||
.try_get::<Option<i64>, _>("created_at_unix_secs")
|
||||
.map_postgres_err()?
|
||||
.map(|value| value.max(0) as u64);
|
||||
let updated_at_unix_secs = row
|
||||
.try_get::<Option<i64>, _>("updated_at_unix_secs")?
|
||||
.try_get::<Option<i64>, _>("updated_at_unix_secs")
|
||||
.map_postgres_err()?
|
||||
.map(|value| value.max(0) as u64);
|
||||
StoredAdminGlobalModel::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("name")?,
|
||||
row.try_get("display_name")?,
|
||||
row.try_get("is_active")?,
|
||||
row.try_get("default_price_per_request")?,
|
||||
row.try_get("default_tiered_pricing")?,
|
||||
row.try_get("supported_capabilities")?,
|
||||
row.try_get("config")?,
|
||||
row.try_get("id").map_postgres_err()?,
|
||||
row.try_get("name").map_postgres_err()?,
|
||||
row.try_get("display_name").map_postgres_err()?,
|
||||
row.try_get("is_active").map_postgres_err()?,
|
||||
row.try_get("default_price_per_request")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("default_tiered_pricing").map_postgres_err()?,
|
||||
row.try_get("supported_capabilities").map_postgres_err()?,
|
||||
row.try_get("config").map_postgres_err()?,
|
||||
created_at_unix_secs,
|
||||
updated_at_unix_secs,
|
||||
)
|
||||
@@ -1034,9 +1093,9 @@ fn build_provider_id_list_query<'a>(
|
||||
|
||||
fn map_provider_model_stats_row(row: &PgRow) -> Result<StoredProviderModelStats, DataLayerError> {
|
||||
StoredProviderModelStats::new(
|
||||
row.try_get("provider_id")?,
|
||||
row.try_get("total_models")?,
|
||||
row.try_get("active_models")?,
|
||||
row.try_get("provider_id").map_postgres_err()?,
|
||||
row.try_get("total_models").map_postgres_err()?,
|
||||
row.try_get("active_models").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
@@ -1044,8 +1103,8 @@ fn map_provider_active_global_model_row(
|
||||
row: &PgRow,
|
||||
) -> Result<StoredProviderActiveGlobalModel, DataLayerError> {
|
||||
StoredProviderActiveGlobalModel::new(
|
||||
row.try_get("provider_id")?,
|
||||
row.try_get("global_model_id")?,
|
||||
row.try_get("provider_id").map_postgres_err()?,
|
||||
row.try_get("global_model_id").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -1,688 +0,0 @@
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
|
||||
#[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(),
|
||||
));
|
||||
}
|
||||
|
||||
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 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>,
|
||||
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,
|
||||
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 created_at_unix_secs: 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>,
|
||||
created_at_unix_secs: 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(),
|
||||
));
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
id,
|
||||
name,
|
||||
display_name,
|
||||
is_active,
|
||||
default_price_per_request,
|
||||
default_tiered_pricing,
|
||||
supported_capabilities,
|
||||
config,
|
||||
created_at_unix_secs,
|
||||
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_secs: 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_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_secs: 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_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(),
|
||||
));
|
||||
}
|
||||
|
||||
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_secs,
|
||||
updated_at_unix_secs,
|
||||
global_model_name,
|
||||
global_model_display_name,
|
||||
global_model_default_price_per_request,
|
||||
global_model_default_tiered_pricing,
|
||||
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(),
|
||||
));
|
||||
}
|
||||
|
||||
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>,
|
||||
}
|
||||
|
||||
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(),
|
||||
));
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
id,
|
||||
name,
|
||||
display_name,
|
||||
is_active,
|
||||
default_price_per_request,
|
||||
default_tiered_pricing,
|
||||
supported_capabilities,
|
||||
config,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[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>,
|
||||
}
|
||||
|
||||
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(),
|
||||
));
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
id,
|
||||
display_name,
|
||||
is_active,
|
||||
default_price_per_request,
|
||||
default_tiered_pricing,
|
||||
supported_capabilities,
|
||||
config,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[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>;
|
||||
}
|
||||
@@ -7,7 +7,7 @@ use super::types::{
|
||||
StoredManagementTokenListPage, StoredManagementTokenUserSummary, StoredManagementTokenWithUser,
|
||||
UpdateManagementTokenRecord,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
|
||||
const LIST_MANAGEMENT_TOKENS_SQL: &str = r#"
|
||||
SELECT
|
||||
@@ -217,8 +217,9 @@ impl ManagementTokenReadRepository for SqlxManagementTokenRepository {
|
||||
.bind(query.user_id.as_deref())
|
||||
.bind(query.is_active)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
let total = count_row.try_get::<i64, _>("total")?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let total = count_row.try_get::<i64, _>("total").map_postgres_err()?;
|
||||
|
||||
let rows = sqlx::query(LIST_MANAGEMENT_TOKENS_SQL)
|
||||
.bind(query.user_id.as_deref())
|
||||
@@ -226,7 +227,8 @@ impl ManagementTokenReadRepository for SqlxManagementTokenRepository {
|
||||
.bind(i64::try_from(query.offset).unwrap_or(i64::MAX))
|
||||
.bind(i64::try_from(query.limit).unwrap_or(i64::MAX))
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
Ok(StoredManagementTokenListPage {
|
||||
items: rows
|
||||
@@ -244,7 +246,8 @@ impl ManagementTokenReadRepository for SqlxManagementTokenRepository {
|
||||
let row = sqlx::query(GET_MANAGEMENT_TOKEN_WITH_USER_SQL)
|
||||
.bind(token_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_token_with_user_row).transpose()
|
||||
}
|
||||
}
|
||||
@@ -305,7 +308,8 @@ impl ManagementTokenWriteRepository for SqlxManagementTokenRepository {
|
||||
let result = sqlx::query(DELETE_MANAGEMENT_TOKEN_SQL)
|
||||
.bind(token_id)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(result.rows_affected() > 0)
|
||||
}
|
||||
|
||||
@@ -318,7 +322,8 @@ impl ManagementTokenWriteRepository for SqlxManagementTokenRepository {
|
||||
.bind(token_id)
|
||||
.bind(is_active)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_token_row).transpose()
|
||||
}
|
||||
|
||||
@@ -332,7 +337,8 @@ impl ManagementTokenWriteRepository for SqlxManagementTokenRepository {
|
||||
.bind(&mutation.token_hash)
|
||||
.bind(mutation.token_prefix.as_deref())
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_token_row).transpose()
|
||||
}
|
||||
}
|
||||
@@ -363,40 +369,40 @@ fn map_management_token_write_error(
|
||||
|
||||
match conflict {
|
||||
Some(detail) => DataLayerError::InvalidInput(detail),
|
||||
None => DataLayerError::Postgres(err),
|
||||
None => DataLayerError::Postgres(err.to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
fn map_token_row(row: &PgRow) -> Result<StoredManagementToken, DataLayerError> {
|
||||
Ok(StoredManagementToken::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("user_id")?,
|
||||
row.try_get("name")?,
|
||||
row.try_get("id").map_postgres_err()?,
|
||||
row.try_get("user_id").map_postgres_err()?,
|
||||
row.try_get("name").map_postgres_err()?,
|
||||
)?
|
||||
.with_display_fields(
|
||||
row.try_get("description")?,
|
||||
row.try_get("token_prefix")?,
|
||||
row.try_get("allowed_ips")?,
|
||||
row.try_get("description").map_postgres_err()?,
|
||||
row.try_get("token_prefix").map_postgres_err()?,
|
||||
row.try_get("allowed_ips").map_postgres_err()?,
|
||||
)
|
||||
.with_runtime_fields(
|
||||
optional_unix_secs(row.try_get("expires_at_unix_secs")?),
|
||||
optional_unix_secs(row.try_get("last_used_at_unix_secs")?),
|
||||
row.try_get("last_used_ip")?,
|
||||
u64::try_from(row.try_get::<i32, _>("usage_count")?).unwrap_or(0),
|
||||
row.try_get("is_active")?,
|
||||
optional_unix_secs(row.try_get("expires_at_unix_secs").map_postgres_err()?),
|
||||
optional_unix_secs(row.try_get("last_used_at_unix_secs").map_postgres_err()?),
|
||||
row.try_get("last_used_ip").map_postgres_err()?,
|
||||
u64::try_from(row.try_get::<i32, _>("usage_count").map_postgres_err()?).unwrap_or(0),
|
||||
row.try_get("is_active").map_postgres_err()?,
|
||||
)
|
||||
.with_timestamps(
|
||||
optional_unix_secs(row.try_get("created_at_unix_secs")?),
|
||||
optional_unix_secs(row.try_get("updated_at_unix_secs")?),
|
||||
optional_unix_secs(row.try_get("created_at_unix_secs").map_postgres_err()?),
|
||||
optional_unix_secs(row.try_get("updated_at_unix_secs").map_postgres_err()?),
|
||||
))
|
||||
}
|
||||
|
||||
fn map_user_summary_row(row: &PgRow) -> Result<StoredManagementTokenUserSummary, DataLayerError> {
|
||||
StoredManagementTokenUserSummary::new(
|
||||
row.try_get("user_row_id")?,
|
||||
row.try_get("user_email")?,
|
||||
row.try_get("user_username")?,
|
||||
row.try_get("user_role")?,
|
||||
row.try_get("user_row_id").map_postgres_err()?,
|
||||
row.try_get("user_email").map_postgres_err()?,
|
||||
row.try_get("user_username").map_postgres_err()?,
|
||||
row.try_get("user_role").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ use super::types::{
|
||||
OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig,
|
||||
UpsertOAuthProviderConfigRecord,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
|
||||
const LIST_OAUTH_PROVIDER_CONFIGS_SQL: &str = r#"
|
||||
SELECT
|
||||
@@ -183,7 +183,8 @@ impl OAuthProviderReadRepository for SqlxOAuthProviderRepository {
|
||||
) -> Result<Vec<StoredOAuthProviderConfig>, DataLayerError> {
|
||||
let rows = sqlx::query(LIST_OAUTH_PROVIDER_CONFIGS_SQL)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_oauth_provider_row).collect()
|
||||
}
|
||||
|
||||
@@ -194,7 +195,8 @@ impl OAuthProviderReadRepository for SqlxOAuthProviderRepository {
|
||||
let row = sqlx::query(GET_OAUTH_PROVIDER_CONFIG_SQL)
|
||||
.bind(provider_type)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_oauth_provider_row).transpose()
|
||||
}
|
||||
|
||||
@@ -207,7 +209,8 @@ impl OAuthProviderReadRepository for SqlxOAuthProviderRepository {
|
||||
.bind(provider_type)
|
||||
.bind(ldap_exclusive)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
usize::try_from(locked_count).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(
|
||||
"oauth_providers.locked_user_count is negative".to_string(),
|
||||
@@ -239,7 +242,8 @@ impl OAuthProviderWriteRepository for SqlxOAuthProviderRepository {
|
||||
.bind(record.extra_config.as_ref())
|
||||
.bind(record.is_enabled)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
map_oauth_provider_row(&row)
|
||||
}
|
||||
|
||||
@@ -250,7 +254,8 @@ impl OAuthProviderWriteRepository for SqlxOAuthProviderRepository {
|
||||
let result = sqlx::query(DELETE_OAUTH_PROVIDER_CONFIG_SQL)
|
||||
.bind(provider_type)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(result.rows_affected() > 0)
|
||||
}
|
||||
}
|
||||
@@ -275,51 +280,130 @@ fn parse_scopes(value: Option<serde_json::Value>) -> Result<Option<Vec<String>>,
|
||||
let Some(value) = value else {
|
||||
return Ok(None);
|
||||
};
|
||||
let serde_json::Value::Array(items) = value else {
|
||||
return Err(DataLayerError::UnexpectedValue(
|
||||
parse_scopes_value(&value)
|
||||
}
|
||||
|
||||
fn parse_scopes_value(value: &serde_json::Value) -> Result<Option<Vec<String>>, DataLayerError> {
|
||||
match value {
|
||||
serde_json::Value::Null => Ok(None),
|
||||
serde_json::Value::Array(items) => parse_scopes_array(items).map(Some),
|
||||
serde_json::Value::String(raw) => parse_embedded_scopes(raw),
|
||||
_ => Err(DataLayerError::UnexpectedValue(
|
||||
"oauth_providers.scopes is not a JSON array".to_string(),
|
||||
));
|
||||
};
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_embedded_scopes(raw: &str) -> Result<Option<Vec<String>>, DataLayerError> {
|
||||
let raw = raw.trim();
|
||||
if raw.is_empty() || raw.eq_ignore_ascii_case("null") {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
if let Ok(decoded) = serde_json::from_str::<serde_json::Value>(raw) {
|
||||
return parse_scopes_value(&decoded);
|
||||
}
|
||||
|
||||
Ok(Some(vec![raw.to_string()]))
|
||||
}
|
||||
|
||||
fn parse_scopes_array(items: &[serde_json::Value]) -> Result<Vec<String>, DataLayerError> {
|
||||
let mut scopes = Vec::with_capacity(items.len());
|
||||
for item in items {
|
||||
let serde_json::Value::String(scope) = item else {
|
||||
let Some(scope) = item.as_str() else {
|
||||
return Err(DataLayerError::UnexpectedValue(
|
||||
"oauth_providers.scopes contains non-string value".to_string(),
|
||||
));
|
||||
};
|
||||
scopes.push(scope);
|
||||
let scope = scope.trim();
|
||||
if !scope.is_empty() {
|
||||
scopes.push(scope.to_string());
|
||||
}
|
||||
}
|
||||
Ok(Some(scopes))
|
||||
Ok(scopes)
|
||||
}
|
||||
|
||||
fn map_oauth_provider_row(row: &PgRow) -> Result<StoredOAuthProviderConfig, DataLayerError> {
|
||||
Ok(StoredOAuthProviderConfig::new(
|
||||
row.try_get("provider_type")?,
|
||||
row.try_get("display_name")?,
|
||||
row.try_get("client_id")?,
|
||||
row.try_get("redirect_uri")?,
|
||||
row.try_get("frontend_callback_url")?,
|
||||
row.try_get("provider_type").map_postgres_err()?,
|
||||
row.try_get("display_name").map_postgres_err()?,
|
||||
row.try_get("client_id").map_postgres_err()?,
|
||||
row.try_get("redirect_uri").map_postgres_err()?,
|
||||
row.try_get("frontend_callback_url").map_postgres_err()?,
|
||||
)?
|
||||
.with_config_fields(
|
||||
row.try_get("client_secret_encrypted")?,
|
||||
row.try_get("authorization_url_override")?,
|
||||
row.try_get("token_url_override")?,
|
||||
row.try_get("userinfo_url_override")?,
|
||||
parse_scopes(row.try_get("scopes")?)?,
|
||||
row.try_get("attribute_mapping")?,
|
||||
row.try_get("extra_config")?,
|
||||
row.try_get("is_enabled")?,
|
||||
row.try_get("client_secret_encrypted").map_postgres_err()?,
|
||||
row.try_get("authorization_url_override")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("token_url_override").map_postgres_err()?,
|
||||
row.try_get("userinfo_url_override").map_postgres_err()?,
|
||||
parse_scopes(row.try_get("scopes").map_postgres_err()?)?,
|
||||
row.try_get("attribute_mapping").map_postgres_err()?,
|
||||
row.try_get("extra_config").map_postgres_err()?,
|
||||
row.try_get("is_enabled").map_postgres_err()?,
|
||||
)
|
||||
.with_timestamps(
|
||||
optional_unix_secs(row.try_get("created_at_unix_secs")?),
|
||||
optional_unix_secs(row.try_get("updated_at_unix_secs")?),
|
||||
optional_unix_secs(row.try_get("created_at_unix_secs").map_postgres_err()?),
|
||||
optional_unix_secs(row.try_get("updated_at_unix_secs").map_postgres_err()?),
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::SqlxOAuthProviderRepository;
|
||||
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
|
||||
use super::{parse_scopes, SqlxOAuthProviderRepository};
|
||||
use crate::{
|
||||
postgres::{PostgresPoolConfig, PostgresPoolFactory},
|
||||
DataLayerError,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn parse_scopes_accepts_json_arrays() {
|
||||
let scopes = parse_scopes(Some(serde_json::json!(["openid", " profile ", ""])))
|
||||
.expect("json array should parse");
|
||||
assert_eq!(
|
||||
scopes,
|
||||
Some(vec!["openid".to_string(), "profile".to_string()])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_scopes_accepts_stringified_json_arrays() {
|
||||
let scopes = parse_scopes(Some(serde_json::json!("[\"openid\", \" profile \", \"\"]")))
|
||||
.expect("stringified array should parse");
|
||||
assert_eq!(
|
||||
scopes,
|
||||
Some(vec!["openid".to_string(), "profile".to_string()])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_scopes_accepts_plain_strings_as_single_scope() {
|
||||
let scopes =
|
||||
parse_scopes(Some(serde_json::json!("openid"))).expect("plain string should parse");
|
||||
assert_eq!(scopes, Some(vec!["openid".to_string()]));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_scopes_rejects_non_string_items() {
|
||||
let err = parse_scopes(Some(serde_json::json!(["openid", 1])))
|
||||
.expect_err("non-string items should fail");
|
||||
assert!(matches!(
|
||||
err,
|
||||
DataLayerError::UnexpectedValue(ref message)
|
||||
if message == "oauth_providers.scopes contains non-string value"
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_scopes_rejects_non_array_objects() {
|
||||
let err = parse_scopes(Some(serde_json::json!({"scope": "openid"})))
|
||||
.expect_err("object should fail");
|
||||
assert!(matches!(
|
||||
err,
|
||||
DataLayerError::UnexpectedValue(ref message)
|
||||
if message == "oauth_providers.scopes is not a JSON array"
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn repository_constructs_from_lazy_pool() {
|
||||
|
||||
@@ -3,7 +3,7 @@ use std::sync::RwLock;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::types::{
|
||||
use super::{
|
||||
ProviderCatalogKeyListQuery, ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogKeyPage,
|
||||
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
mod memory;
|
||||
mod sql;
|
||||
mod types;
|
||||
|
||||
pub use memory::InMemoryProviderCatalogReadRepository;
|
||||
pub use sql::SqlxProviderCatalogReadRepository;
|
||||
pub use types::{
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyListQuery, ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogKeyPage,
|
||||
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
|
||||
};
|
||||
pub use memory::InMemoryProviderCatalogReadRepository;
|
||||
pub use sql::SqlxProviderCatalogReadRepository;
|
||||
|
||||
@@ -1,12 +1,15 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{postgres::PgRow, PgPool, Postgres, QueryBuilder, Row};
|
||||
|
||||
use super::types::{
|
||||
use super::{
|
||||
ProviderCatalogKeyListQuery, ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogKeyPage,
|
||||
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
use crate::{
|
||||
error::{postgres_error, SqlxResultExt},
|
||||
DataLayerError,
|
||||
};
|
||||
|
||||
const LIST_PROVIDERS_BY_IDS_PREFIX: &str = r#"
|
||||
SELECT
|
||||
@@ -271,7 +274,8 @@ impl SqlxProviderCatalogReadRepository {
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_provider_row).collect()
|
||||
}
|
||||
|
||||
@@ -312,7 +316,8 @@ ORDER BY provider_priority ASC, name ASC
|
||||
)
|
||||
.bind(active_only)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_provider_row).collect()
|
||||
}
|
||||
|
||||
@@ -334,17 +339,16 @@ ORDER BY provider_priority ASC, name ASC
|
||||
.await
|
||||
{
|
||||
Ok(rows) => rows,
|
||||
Err(error) if is_missing_endpoint_health_score_column(&error) => {
|
||||
build_list_query(
|
||||
LIST_ENDPOINTS_BY_IDS_PREFIX_LEGACY,
|
||||
endpoint_ids,
|
||||
" ORDER BY api_format ASC, id ASC",
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await?
|
||||
}
|
||||
Err(error) => return Err(error.into()),
|
||||
Err(error) if is_missing_endpoint_health_score_column(&error) => build_list_query(
|
||||
LIST_ENDPOINTS_BY_IDS_PREFIX_LEGACY,
|
||||
endpoint_ids,
|
||||
" ORDER BY api_format ASC, id ASC",
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?,
|
||||
Err(error) => return Err(postgres_error(error)),
|
||||
};
|
||||
rows.iter().map(map_endpoint_row).collect()
|
||||
}
|
||||
@@ -367,17 +371,16 @@ ORDER BY provider_priority ASC, name ASC
|
||||
.await
|
||||
{
|
||||
Ok(rows) => rows,
|
||||
Err(error) if is_missing_endpoint_health_score_column(&error) => {
|
||||
build_list_query(
|
||||
LIST_ENDPOINTS_BY_PROVIDER_IDS_PREFIX_LEGACY,
|
||||
provider_ids,
|
||||
" ORDER BY provider_id ASC, api_format ASC, id ASC",
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await?
|
||||
}
|
||||
Err(error) => return Err(error.into()),
|
||||
Err(error) if is_missing_endpoint_health_score_column(&error) => build_list_query(
|
||||
LIST_ENDPOINTS_BY_PROVIDER_IDS_PREFIX_LEGACY,
|
||||
provider_ids,
|
||||
" ORDER BY provider_id ASC, api_format ASC, id ASC",
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?,
|
||||
Err(error) => return Err(postgres_error(error)),
|
||||
};
|
||||
rows.iter().map(map_endpoint_row).collect()
|
||||
}
|
||||
@@ -397,7 +400,8 @@ ORDER BY provider_priority ASC, name ASC
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_key_row).collect()
|
||||
}
|
||||
|
||||
@@ -416,7 +420,8 @@ ORDER BY provider_priority ASC, name ASC
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_key_row).collect()
|
||||
}
|
||||
|
||||
@@ -462,8 +467,9 @@ WHERE provider_id = $1
|
||||
.bind(search_pattern.as_deref())
|
||||
.bind(query.is_active)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
let total = count_row.try_get::<i64, _>("total")?.max(0) as usize;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let total = row_get::<i64>(&count_row, "total")?.max(0) as usize;
|
||||
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
@@ -530,7 +536,8 @@ LIMIT $5
|
||||
.bind(offset)
|
||||
.bind(limit)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let items = rows
|
||||
.iter()
|
||||
.map(map_key_row)
|
||||
@@ -554,7 +561,8 @@ LIMIT $5
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_key_stats_row).collect()
|
||||
}
|
||||
|
||||
@@ -597,7 +605,8 @@ WHERE id = $1
|
||||
.bind(encrypted_auth_config)
|
||||
.bind(expires_at_unix_secs.map(|value| value as f64))
|
||||
.execute(&self.pool)
|
||||
.await?
|
||||
.await
|
||||
.map_postgres_err()?
|
||||
.rows_affected();
|
||||
|
||||
Ok(rows_affected > 0)
|
||||
@@ -634,7 +643,7 @@ WHERE id = $1
|
||||
));
|
||||
}
|
||||
|
||||
let mut tx = self.pool.begin().await?;
|
||||
let mut tx = self.pool.begin().await.map_postgres_err()?;
|
||||
|
||||
if let Some(target_priority) = shift_existing_priorities_from {
|
||||
sqlx::query(
|
||||
@@ -647,7 +656,8 @@ WHERE provider_priority IS NOT NULL
|
||||
)
|
||||
.bind(target_priority)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
}
|
||||
|
||||
sqlx::query(
|
||||
@@ -752,9 +762,10 @@ INSERT INTO providers (
|
||||
.bind(provider.created_at_unix_secs.map(|value| value as f64))
|
||||
.bind(provider.updated_at_unix_secs.map(|value| value as f64))
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
tx.commit().await?;
|
||||
tx.commit().await.map_err(postgres_error)?;
|
||||
|
||||
self.list_providers_by_ids(std::slice::from_ref(&provider.id))
|
||||
.await?
|
||||
@@ -871,7 +882,8 @@ WHERE id = $1
|
||||
.bind(&provider.config)
|
||||
.bind(provider.updated_at_unix_secs.map(|value| value as f64))
|
||||
.execute(&self.pool)
|
||||
.await?
|
||||
.await
|
||||
.map_postgres_err()?
|
||||
.rows_affected();
|
||||
|
||||
if rows_affected == 0 {
|
||||
@@ -908,7 +920,8 @@ WHERE id = $1
|
||||
)
|
||||
.bind(provider_id)
|
||||
.execute(&self.pool)
|
||||
.await?
|
||||
.await
|
||||
.map_postgres_err()?
|
||||
.rows_affected();
|
||||
|
||||
Ok(rows_affected > 0)
|
||||
@@ -926,26 +939,30 @@ WHERE id = $1
|
||||
));
|
||||
}
|
||||
|
||||
let mut tx = self.pool.begin().await?;
|
||||
let mut tx = self.pool.begin().await.map_postgres_err()?;
|
||||
|
||||
sqlx::query(
|
||||
"UPDATE user_preferences SET default_provider_id = NULL WHERE default_provider_id = $1",
|
||||
)
|
||||
.bind(provider_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
sqlx::query("UPDATE usage SET provider_id = NULL WHERE provider_id = $1")
|
||||
.bind(provider_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
sqlx::query("UPDATE video_tasks SET provider_id = NULL WHERE provider_id = $1")
|
||||
.bind(provider_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
sqlx::query("DELETE FROM request_candidates WHERE provider_id = $1")
|
||||
.bind(provider_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
for endpoint_id in endpoint_ids {
|
||||
sqlx::query(
|
||||
@@ -953,44 +970,52 @@ WHERE id = $1
|
||||
)
|
||||
.bind(endpoint_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
sqlx::query("UPDATE video_tasks SET endpoint_id = NULL WHERE endpoint_id = $1")
|
||||
.bind(endpoint_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
sqlx::query("DELETE FROM request_candidates WHERE endpoint_id = $1")
|
||||
.bind(endpoint_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
}
|
||||
|
||||
for key_id in key_ids {
|
||||
sqlx::query("DELETE FROM gemini_file_mappings WHERE key_id = $1")
|
||||
.bind(key_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
sqlx::query(
|
||||
"UPDATE usage SET provider_api_key_id = NULL WHERE provider_api_key_id = $1",
|
||||
)
|
||||
.bind(key_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
sqlx::query("UPDATE video_tasks SET key_id = NULL WHERE key_id = $1")
|
||||
.bind(key_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
}
|
||||
|
||||
sqlx::query("DELETE FROM api_key_provider_mappings WHERE provider_id = $1")
|
||||
.bind(provider_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
sqlx::query("DELETE FROM provider_usage_tracking WHERE provider_id = $1")
|
||||
.bind(provider_id)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
tx.commit().await?;
|
||||
tx.commit().await.map_err(postgres_error)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -1016,7 +1041,8 @@ WHERE id = $1
|
||||
)
|
||||
.bind(key_id)
|
||||
.execute(&self.pool)
|
||||
.await?
|
||||
.await
|
||||
.map_postgres_err()?
|
||||
.rows_affected();
|
||||
|
||||
Ok(rows_affected > 0)
|
||||
@@ -1217,7 +1243,8 @@ INSERT INTO provider_api_keys (
|
||||
.bind(key.created_at_unix_secs.map(|value| value as f64))
|
||||
.bind(key.updated_at_unix_secs.map(|value| value as f64))
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
self.list_keys_by_ids(std::slice::from_ref(&key.id))
|
||||
.await?
|
||||
@@ -1246,7 +1273,7 @@ INSERT INTO provider_api_keys (
|
||||
));
|
||||
}
|
||||
|
||||
sqlx::query(
|
||||
match sqlx::query(
|
||||
r#"
|
||||
INSERT INTO provider_endpoints (
|
||||
id,
|
||||
@@ -1311,7 +1338,77 @@ INSERT INTO provider_endpoints (
|
||||
.bind(endpoint.created_at_unix_secs.map(|value| value as f64))
|
||||
.bind(endpoint.updated_at_unix_secs.map(|value| value as f64))
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
{
|
||||
Ok(_) => {}
|
||||
Err(error) if is_missing_endpoint_health_score_column(&error) => {
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO provider_endpoints (
|
||||
id,
|
||||
provider_id,
|
||||
api_format,
|
||||
api_family,
|
||||
endpoint_kind,
|
||||
is_active,
|
||||
base_url,
|
||||
header_rules,
|
||||
body_rules,
|
||||
max_retries,
|
||||
custom_path,
|
||||
config,
|
||||
format_acceptance_config,
|
||||
proxy,
|
||||
created_at,
|
||||
updated_at
|
||||
) VALUES (
|
||||
$1,
|
||||
$2,
|
||||
$3,
|
||||
$4,
|
||||
$5,
|
||||
$6,
|
||||
$7,
|
||||
$8,
|
||||
$9,
|
||||
$10,
|
||||
$11,
|
||||
$12,
|
||||
$13,
|
||||
$14,
|
||||
CASE
|
||||
WHEN $15::double precision IS NULL THEN NOW()
|
||||
ELSE TO_TIMESTAMP($15::double precision)
|
||||
END,
|
||||
CASE
|
||||
WHEN $16::double precision IS NULL THEN NOW()
|
||||
ELSE TO_TIMESTAMP($16::double precision)
|
||||
END
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.bind(&endpoint.id)
|
||||
.bind(&endpoint.provider_id)
|
||||
.bind(&endpoint.api_format)
|
||||
.bind(&endpoint.api_family)
|
||||
.bind(&endpoint.endpoint_kind)
|
||||
.bind(endpoint.is_active)
|
||||
.bind(&endpoint.base_url)
|
||||
.bind(&endpoint.header_rules)
|
||||
.bind(&endpoint.body_rules)
|
||||
.bind(endpoint.max_retries)
|
||||
.bind(&endpoint.custom_path)
|
||||
.bind(&endpoint.config)
|
||||
.bind(&endpoint.format_acceptance_config)
|
||||
.bind(&endpoint.proxy)
|
||||
.bind(endpoint.created_at_unix_secs.map(|value| value as f64))
|
||||
.bind(endpoint.updated_at_unix_secs.map(|value| value as f64))
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
}
|
||||
Err(error) => return Err(postgres_error(error)),
|
||||
}
|
||||
|
||||
self.list_endpoints_by_ids(std::slice::from_ref(&endpoint.id))
|
||||
.await?
|
||||
@@ -1340,7 +1437,7 @@ INSERT INTO provider_endpoints (
|
||||
));
|
||||
}
|
||||
|
||||
let rows_affected = sqlx::query(
|
||||
let rows_affected = match sqlx::query(
|
||||
r#"
|
||||
UPDATE provider_endpoints
|
||||
SET
|
||||
@@ -1382,8 +1479,54 @@ WHERE id = $1
|
||||
.bind(&endpoint.proxy)
|
||||
.bind(endpoint.updated_at_unix_secs.map(|value| value as f64))
|
||||
.execute(&self.pool)
|
||||
.await?
|
||||
.rows_affected();
|
||||
.await
|
||||
{
|
||||
Ok(result) => result.rows_affected(),
|
||||
Err(error) if is_missing_endpoint_health_score_column(&error) => sqlx::query(
|
||||
r#"
|
||||
UPDATE provider_endpoints
|
||||
SET
|
||||
provider_id = $2,
|
||||
api_format = $3,
|
||||
api_family = $4,
|
||||
endpoint_kind = $5,
|
||||
is_active = $6,
|
||||
base_url = $7,
|
||||
header_rules = $8,
|
||||
body_rules = $9,
|
||||
max_retries = $10,
|
||||
custom_path = $11,
|
||||
config = $12,
|
||||
format_acceptance_config = $13,
|
||||
proxy = $14,
|
||||
updated_at = CASE
|
||||
WHEN $15::double precision IS NULL THEN NOW()
|
||||
ELSE TO_TIMESTAMP($15::double precision)
|
||||
END
|
||||
WHERE id = $1
|
||||
"#,
|
||||
)
|
||||
.bind(&endpoint.id)
|
||||
.bind(&endpoint.provider_id)
|
||||
.bind(&endpoint.api_format)
|
||||
.bind(&endpoint.api_family)
|
||||
.bind(&endpoint.endpoint_kind)
|
||||
.bind(endpoint.is_active)
|
||||
.bind(&endpoint.base_url)
|
||||
.bind(&endpoint.header_rules)
|
||||
.bind(&endpoint.body_rules)
|
||||
.bind(endpoint.max_retries)
|
||||
.bind(&endpoint.custom_path)
|
||||
.bind(&endpoint.config)
|
||||
.bind(&endpoint.format_acceptance_config)
|
||||
.bind(&endpoint.proxy)
|
||||
.bind(endpoint.updated_at_unix_secs.map(|value| value as f64))
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?
|
||||
.rows_affected(),
|
||||
Err(error) => return Err(postgres_error(error)),
|
||||
};
|
||||
|
||||
if rows_affected == 0 {
|
||||
return Err(DataLayerError::UnexpectedValue(format!(
|
||||
@@ -1419,7 +1562,8 @@ WHERE id = $1
|
||||
)
|
||||
.bind(endpoint_id)
|
||||
.execute(&self.pool)
|
||||
.await?
|
||||
.await
|
||||
.map_postgres_err()?
|
||||
.rows_affected();
|
||||
|
||||
Ok(rows_affected > 0)
|
||||
@@ -1511,7 +1655,8 @@ WHERE id = $1
|
||||
.bind(key.is_active)
|
||||
.bind(key.updated_at_unix_secs.map(|value| value as f64))
|
||||
.execute(&self.pool)
|
||||
.await?
|
||||
.await
|
||||
.map_postgres_err()?
|
||||
.rows_affected();
|
||||
|
||||
if rows_affected == 0 {
|
||||
@@ -1548,7 +1693,8 @@ WHERE id = $1
|
||||
)
|
||||
.bind(key_id)
|
||||
.execute(&self.pool)
|
||||
.await?
|
||||
.await
|
||||
.map_postgres_err()?
|
||||
.rows_affected();
|
||||
|
||||
Ok(rows_affected > 0)
|
||||
@@ -1583,7 +1729,8 @@ WHERE id = $1
|
||||
.bind(health_by_format)
|
||||
.bind(circuit_breaker_by_format)
|
||||
.execute(&self.pool)
|
||||
.await?
|
||||
.await
|
||||
.map_postgres_err()?
|
||||
.rows_affected();
|
||||
|
||||
Ok(rows_affected > 0)
|
||||
@@ -1769,9 +1916,15 @@ fn build_list_query<'a>(
|
||||
builder
|
||||
}
|
||||
|
||||
fn row_get<T>(row: &PgRow, column: &str) -> Result<T, DataLayerError>
|
||||
where
|
||||
for<'r> T: sqlx::Decode<'r, sqlx::Postgres> + sqlx::Type<sqlx::Postgres>,
|
||||
{
|
||||
row.try_get(column).map_postgres_err()
|
||||
}
|
||||
|
||||
fn map_provider_row(row: &PgRow) -> Result<StoredProviderCatalogProvider, DataLayerError> {
|
||||
let quota_reset_day = row
|
||||
.try_get::<Option<i32>, _>("quota_reset_day")?
|
||||
let quota_reset_day = row_get::<Option<i32>>(row, "quota_reset_day")?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -1780,8 +1933,7 @@ fn map_provider_row(row: &PgRow) -> Result<StoredProviderCatalogProvider, DataLa
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let created_at_unix_secs = row
|
||||
.try_get::<Option<i64>, _>("created_at_unix_secs")?
|
||||
let created_at_unix_secs = row_get::<Option<i64>>(row, "created_at_unix_secs")?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -1790,8 +1942,7 @@ fn map_provider_row(row: &PgRow) -> Result<StoredProviderCatalogProvider, DataLa
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let updated_at_unix_secs = row
|
||||
.try_get::<Option<i64>, _>("updated_at_unix_secs")?
|
||||
let updated_at_unix_secs = row_get::<Option<i64>>(row, "updated_at_unix_secs")?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -1801,40 +1952,37 @@ fn map_provider_row(row: &PgRow) -> Result<StoredProviderCatalogProvider, DataLa
|
||||
})
|
||||
.transpose()?;
|
||||
Ok(StoredProviderCatalogProvider::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("name")?,
|
||||
row.try_get("website")?,
|
||||
row.try_get("provider_type")?,
|
||||
row_get(row, "id")?,
|
||||
row_get(row, "name")?,
|
||||
row_get(row, "website")?,
|
||||
row_get(row, "provider_type")?,
|
||||
)?
|
||||
.with_description(row.try_get("description")?)
|
||||
.with_description(row_get(row, "description")?)
|
||||
.with_billing_fields(
|
||||
row.try_get("billing_type")?,
|
||||
row.try_get("monthly_quota_usd")?,
|
||||
row.try_get("monthly_used_usd")?,
|
||||
row_get(row, "billing_type")?,
|
||||
row_get(row, "monthly_quota_usd")?,
|
||||
row_get(row, "monthly_used_usd")?,
|
||||
quota_reset_day,
|
||||
row.try_get::<Option<i64>, _>("quota_last_reset_at_unix_secs")?
|
||||
.map(|value| value as u64),
|
||||
row.try_get::<Option<i64>, _>("quota_expires_at_unix_secs")?
|
||||
.map(|value| value as u64),
|
||||
row_get::<Option<i64>>(row, "quota_last_reset_at_unix_secs")?.map(|value| value as u64),
|
||||
row_get::<Option<i64>>(row, "quota_expires_at_unix_secs")?.map(|value| value as u64),
|
||||
)
|
||||
.with_routing_fields(row.try_get("provider_priority")?)
|
||||
.with_routing_fields(row_get(row, "provider_priority")?)
|
||||
.with_transport_fields(
|
||||
row.try_get("is_active")?,
|
||||
row.try_get("keep_priority_on_conversion")?,
|
||||
row.try_get("enable_format_conversion")?,
|
||||
row.try_get("concurrent_limit")?,
|
||||
row.try_get("max_retries")?,
|
||||
row.try_get("proxy")?,
|
||||
row.try_get("request_timeout")?,
|
||||
row.try_get("stream_first_byte_timeout")?,
|
||||
row.try_get("config")?,
|
||||
row_get(row, "is_active")?,
|
||||
row_get(row, "keep_priority_on_conversion")?,
|
||||
row_get(row, "enable_format_conversion")?,
|
||||
row_get(row, "concurrent_limit")?,
|
||||
row_get(row, "max_retries")?,
|
||||
row_get(row, "proxy")?,
|
||||
row_get(row, "request_timeout")?,
|
||||
row_get(row, "stream_first_byte_timeout")?,
|
||||
row_get(row, "config")?,
|
||||
)
|
||||
.with_timestamps(created_at_unix_secs, updated_at_unix_secs))
|
||||
}
|
||||
|
||||
fn map_endpoint_row(row: &PgRow) -> Result<StoredProviderCatalogEndpoint, DataLayerError> {
|
||||
let created_at_unix_secs = row
|
||||
.try_get::<Option<i64>, _>("created_at_unix_secs")?
|
||||
let created_at_unix_secs = row_get::<Option<i64>>(row, "created_at_unix_secs")?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -1843,8 +1991,7 @@ fn map_endpoint_row(row: &PgRow) -> Result<StoredProviderCatalogEndpoint, DataLa
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let updated_at_unix_secs = row
|
||||
.try_get::<Option<i64>, _>("updated_at_unix_secs")?
|
||||
let updated_at_unix_secs = row_get::<Option<i64>>(row, "updated_at_unix_secs")?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -1854,12 +2001,12 @@ fn map_endpoint_row(row: &PgRow) -> Result<StoredProviderCatalogEndpoint, DataLa
|
||||
})
|
||||
.transpose()?;
|
||||
StoredProviderCatalogEndpoint::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("provider_id")?,
|
||||
row.try_get("api_format")?,
|
||||
row.try_get("api_family")?,
|
||||
row.try_get("endpoint_kind")?,
|
||||
row.try_get("is_active")?,
|
||||
row_get(row, "id")?,
|
||||
row_get(row, "provider_id")?,
|
||||
row_get(row, "api_format")?,
|
||||
row_get(row, "api_family")?,
|
||||
row_get(row, "endpoint_kind")?,
|
||||
row_get(row, "is_active")?,
|
||||
)?
|
||||
.with_timestamps(created_at_unix_secs, updated_at_unix_secs)
|
||||
.with_health_score(
|
||||
@@ -1869,14 +2016,14 @@ fn map_endpoint_row(row: &PgRow) -> Result<StoredProviderCatalogEndpoint, DataLa
|
||||
.unwrap_or(1.0),
|
||||
)
|
||||
.with_transport_fields(
|
||||
row.try_get("base_url")?,
|
||||
row.try_get("header_rules")?,
|
||||
row.try_get("body_rules")?,
|
||||
row.try_get("max_retries")?,
|
||||
row.try_get("custom_path")?,
|
||||
row.try_get("config")?,
|
||||
row.try_get("format_acceptance_config")?,
|
||||
row.try_get("proxy")?,
|
||||
row_get(row, "base_url")?,
|
||||
row_get(row, "header_rules")?,
|
||||
row_get(row, "body_rules")?,
|
||||
row_get(row, "max_retries")?,
|
||||
row_get(row, "custom_path")?,
|
||||
row_get(row, "config")?,
|
||||
row_get(row, "format_acceptance_config")?,
|
||||
row_get(row, "proxy")?,
|
||||
)
|
||||
}
|
||||
|
||||
@@ -1893,15 +2040,14 @@ fn is_missing_endpoint_health_score_column(error: &sqlx::Error) -> bool {
|
||||
|
||||
fn map_key_stats_row(row: &PgRow) -> Result<StoredProviderCatalogKeyStats, DataLayerError> {
|
||||
StoredProviderCatalogKeyStats::new(
|
||||
row.try_get("provider_id")?,
|
||||
row.try_get("total_keys")?,
|
||||
row.try_get("active_keys")?,
|
||||
row_get(row, "provider_id")?,
|
||||
row_get(row, "total_keys")?,
|
||||
row_get(row, "active_keys")?,
|
||||
)
|
||||
}
|
||||
|
||||
fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError> {
|
||||
let rpm_limit = row
|
||||
.try_get::<Option<i32>, _>("rpm_limit")?
|
||||
let rpm_limit = row_get::<Option<i32>>(row, "rpm_limit")?
|
||||
.map(|value| {
|
||||
u32::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -1910,8 +2056,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let learned_rpm_limit = row
|
||||
.try_get::<Option<i32>, _>("learned_rpm_limit")?
|
||||
let learned_rpm_limit = row_get::<Option<i32>>(row, "learned_rpm_limit")?
|
||||
.map(|value| {
|
||||
u32::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -1920,8 +2065,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let concurrent_429_count = row
|
||||
.try_get::<Option<i32>, _>("concurrent_429_count")?
|
||||
let concurrent_429_count = row_get::<Option<i32>>(row, "concurrent_429_count")?
|
||||
.map(|value| {
|
||||
u32::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -1930,8 +2074,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let rpm_429_count = row
|
||||
.try_get::<Option<i32>, _>("rpm_429_count")?
|
||||
let rpm_429_count = row_get::<Option<i32>>(row, "rpm_429_count")?
|
||||
.map(|value| {
|
||||
u32::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -1940,8 +2083,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let request_count = row
|
||||
.try_get::<Option<i32>, _>("request_count")?
|
||||
let request_count = row_get::<Option<i32>>(row, "request_count")?
|
||||
.map(|value| {
|
||||
u32::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -1950,8 +2092,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let success_count = row
|
||||
.try_get::<Option<i32>, _>("success_count")?
|
||||
let success_count = row_get::<Option<i32>>(row, "success_count")?
|
||||
.map(|value| {
|
||||
u32::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -1960,8 +2101,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let error_count = row
|
||||
.try_get::<Option<i32>, _>("error_count")?
|
||||
let error_count = row_get::<Option<i32>>(row, "error_count")?
|
||||
.map(|value| {
|
||||
u32::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -1970,8 +2110,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let total_response_time_ms = row
|
||||
.try_get::<Option<i32>, _>("total_response_time_ms")?
|
||||
let total_response_time_ms = row_get::<Option<i32>>(row, "total_response_time_ms")?
|
||||
.map(|value| {
|
||||
u32::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -1980,28 +2119,27 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let last_probe_increase_at_unix_secs = row
|
||||
.try_get::<Option<i64>, _>("last_probe_increase_at_unix_secs")?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"invalid provider_api_keys.last_probe_increase_at_unix_secs: {value}"
|
||||
))
|
||||
let last_probe_increase_at_unix_secs =
|
||||
row_get::<Option<i64>>(row, "last_probe_increase_at_unix_secs")?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"invalid provider_api_keys.last_probe_increase_at_unix_secs: {value}"
|
||||
))
|
||||
})
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let last_models_fetch_at_unix_secs = row
|
||||
.try_get::<Option<i64>, _>("last_models_fetch_at_unix_secs")?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"invalid provider_api_keys.last_models_fetch_at_unix_secs: {value}"
|
||||
))
|
||||
.transpose()?;
|
||||
let last_models_fetch_at_unix_secs =
|
||||
row_get::<Option<i64>>(row, "last_models_fetch_at_unix_secs")?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"invalid provider_api_keys.last_models_fetch_at_unix_secs: {value}"
|
||||
))
|
||||
})
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let oauth_invalid_at_unix_secs = row
|
||||
.try_get::<Option<i64>, _>("oauth_invalid_at_unix_secs")?
|
||||
.transpose()?;
|
||||
let oauth_invalid_at_unix_secs = row_get::<Option<i64>>(row, "oauth_invalid_at_unix_secs")?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -2010,8 +2148,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let last_used_at_unix_secs = row
|
||||
.try_get::<Option<i64>, _>("last_used_at_unix_secs")?
|
||||
let last_used_at_unix_secs = row_get::<Option<i64>>(row, "last_used_at_unix_secs")?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -2020,8 +2157,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let created_at_unix_secs = row
|
||||
.try_get::<Option<i64>, _>("created_at_unix_secs")?
|
||||
let created_at_unix_secs = row_get::<Option<i64>>(row, "created_at_unix_secs")?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -2030,8 +2166,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let updated_at_unix_secs = row
|
||||
.try_get::<Option<i64>, _>("updated_at_unix_secs")?
|
||||
let updated_at_unix_secs = row_get::<Option<i64>>(row, "updated_at_unix_secs")?
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
@@ -2042,24 +2177,24 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
|
||||
.transpose()?;
|
||||
|
||||
StoredProviderCatalogKey::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("provider_id")?,
|
||||
row.try_get("name")?,
|
||||
row.try_get("auth_type")?,
|
||||
row.try_get("capabilities")?,
|
||||
row.try_get("is_active")?,
|
||||
row_get(row, "id")?,
|
||||
row_get(row, "provider_id")?,
|
||||
row_get(row, "name")?,
|
||||
row_get(row, "auth_type")?,
|
||||
row_get(row, "capabilities")?,
|
||||
row_get(row, "is_active")?,
|
||||
)?
|
||||
.with_transport_fields(
|
||||
row.try_get("api_formats")?,
|
||||
row.try_get("api_key")?,
|
||||
row.try_get("auth_config")?,
|
||||
row.try_get("rate_multipliers")?,
|
||||
row.try_get("global_priority_by_format")?,
|
||||
row.try_get("allowed_models")?,
|
||||
row.try_get::<Option<i64>, _>("expires_at_unix_secs")?
|
||||
row_get(row, "api_formats")?,
|
||||
row_get(row, "api_key")?,
|
||||
row_get(row, "auth_config")?,
|
||||
row_get(row, "rate_multipliers")?,
|
||||
row_get(row, "global_priority_by_format")?,
|
||||
row_get(row, "allowed_models")?,
|
||||
row_get::<Option<i64>>(row, "expires_at_unix_secs")?
|
||||
.and_then(|value| u64::try_from(value).ok()),
|
||||
row.try_get("proxy")?,
|
||||
row.try_get("fingerprint")?,
|
||||
row_get(row, "proxy")?,
|
||||
row_get(row, "fingerprint")?,
|
||||
)
|
||||
.map(|key| {
|
||||
let mut key = key
|
||||
|
||||
@@ -1,738 +0,0 @@
|
||||
use async_trait::async_trait;
|
||||
|
||||
#[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_secs: 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_secs: 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_secs: Option<u64>,
|
||||
updated_at_unix_secs: Option<u64>,
|
||||
) -> Self {
|
||||
self.created_at_unix_secs = created_at_unix_secs;
|
||||
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_secs: 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_secs: 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_secs: Option<u64>,
|
||||
updated_at_unix_secs: Option<u64>,
|
||||
) -> Self {
|
||||
self.created_at_unix_secs = created_at_unix_secs;
|
||||
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 encrypted_api_key: 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 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 request_count: Option<u32>,
|
||||
pub success_count: Option<u32>,
|
||||
pub error_count: Option<u32>,
|
||||
pub total_response_time_ms: Option<u32>,
|
||||
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_secs: 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>,
|
||||
}
|
||||
|
||||
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,
|
||||
encrypted_api_key: String::new(),
|
||||
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,
|
||||
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,
|
||||
request_count: None,
|
||||
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_secs: 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: 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> {
|
||||
if encrypted_api_key.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>,
|
||||
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.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<u32>,
|
||||
) -> Self {
|
||||
self.error_count = error_count;
|
||||
self.total_response_time_ms = total_response_time_ms;
|
||||
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
|
||||
}
|
||||
}
|
||||
|
||||
#[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,
|
||||
}
|
||||
|
||||
#[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 {
|
||||
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_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,
|
||||
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 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::{
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn rejects_empty_provider_name() {
|
||||
assert!(StoredProviderCatalogProvider::new(
|
||||
"provider-1".to_string(),
|
||||
"".to_string(),
|
||||
None,
|
||||
"custom".to_string(),
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_empty_endpoint_api_format() {
|
||||
assert!(StoredProviderCatalogEndpoint::new(
|
||||
"endpoint-1".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"".to_string(),
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_empty_endpoint_base_url() {
|
||||
let endpoint = StoredProviderCatalogEndpoint::new(
|
||||
"endpoint-1".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"openai:chat".to_string(),
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("endpoint should build");
|
||||
assert!(endpoint
|
||||
.with_transport_fields("".to_string(), None, None, None, None, None, None, None,)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_empty_key_auth_type() {
|
||||
assert!(StoredProviderCatalogKey::new(
|
||||
"key-1".to_string(),
|
||||
"provider-1".to_string(),
|
||||
"default".to_string(),
|
||||
"".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_empty_encrypted_api_key() {
|
||||
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!(key
|
||||
.with_transport_fields(
|
||||
None,
|
||||
"".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stores_key_rate_limit_fields() {
|
||||
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(100),
|
||||
Some(80),
|
||||
Some(2),
|
||||
Some(3),
|
||||
Some(1_700_000_000),
|
||||
Some(serde_json::json!([{"new_limit": 80}])),
|
||||
Some(120),
|
||||
Some(110),
|
||||
);
|
||||
|
||||
assert_eq!(key.rpm_limit, Some(100));
|
||||
assert_eq!(key.learned_rpm_limit, Some(80));
|
||||
assert_eq!(key.concurrent_429_count, Some(2));
|
||||
assert_eq!(key.rpm_429_count, Some(3));
|
||||
assert_eq!(key.last_429_at_unix_secs, Some(1_700_000_000));
|
||||
assert_eq!(key.request_count, Some(120));
|
||||
assert_eq!(key.success_count, Some(110));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stores_key_health_fields() {
|
||||
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_health_fields(
|
||||
Some(serde_json::json!({"openai:chat": {"health_score": 0.4}})),
|
||||
Some(serde_json::json!({"openai:chat": {"open": true}})),
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
key.health_by_format,
|
||||
Some(serde_json::json!({"openai:chat": {"health_score": 0.4}}))
|
||||
);
|
||||
assert_eq!(
|
||||
key.circuit_breaker_by_format,
|
||||
Some(serde_json::json!({"openai:chat": {"open": true}}))
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -5,7 +5,10 @@ use super::types::{
|
||||
normalize_proxy_metadata, ProxyNodeHeartbeatMutation, ProxyNodeReadRepository,
|
||||
ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, StoredProxyNode, StoredProxyNodeEvent,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
use crate::{
|
||||
error::{postgres_error, SqlxResultExt},
|
||||
DataLayerError,
|
||||
};
|
||||
|
||||
const FIND_PROXY_NODE_SQL: &str = r#"
|
||||
SELECT
|
||||
@@ -142,49 +145,58 @@ impl SqlxProxyNodeRepository {
|
||||
|
||||
fn row_to_stored(row: &PgRow) -> Result<StoredProxyNode, DataLayerError> {
|
||||
Ok(StoredProxyNode::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("name")?,
|
||||
row.try_get("ip")?,
|
||||
row.try_get("port")?,
|
||||
row.try_get("is_manual")?,
|
||||
row.try_get("status")?,
|
||||
row.try_get("heartbeat_interval")?,
|
||||
row.try_get("active_connections")?,
|
||||
row.try_get("total_requests")?,
|
||||
row.try_get("failed_requests")?,
|
||||
row.try_get("dns_failures")?,
|
||||
row.try_get("stream_errors")?,
|
||||
row.try_get("tunnel_mode")?,
|
||||
row.try_get("tunnel_connected")?,
|
||||
row.try_get("config_version")?,
|
||||
row.try_get("id").map_postgres_err()?,
|
||||
row.try_get("name").map_postgres_err()?,
|
||||
row.try_get("ip").map_postgres_err()?,
|
||||
row.try_get("port").map_postgres_err()?,
|
||||
row.try_get("is_manual").map_postgres_err()?,
|
||||
row.try_get("status").map_postgres_err()?,
|
||||
row.try_get("heartbeat_interval").map_postgres_err()?,
|
||||
row.try_get("active_connections").map_postgres_err()?,
|
||||
row.try_get("total_requests").map_postgres_err()?,
|
||||
row.try_get("failed_requests").map_postgres_err()?,
|
||||
row.try_get("dns_failures").map_postgres_err()?,
|
||||
row.try_get("stream_errors").map_postgres_err()?,
|
||||
row.try_get("tunnel_mode").map_postgres_err()?,
|
||||
row.try_get("tunnel_connected").map_postgres_err()?,
|
||||
row.try_get("config_version").map_postgres_err()?,
|
||||
)?
|
||||
.with_manual_proxy_fields(
|
||||
row.try_get("proxy_url")?,
|
||||
row.try_get("proxy_username")?,
|
||||
row.try_get("proxy_password")?,
|
||||
row.try_get("proxy_url").map_postgres_err()?,
|
||||
row.try_get("proxy_username").map_postgres_err()?,
|
||||
row.try_get("proxy_password").map_postgres_err()?,
|
||||
)
|
||||
.with_runtime_fields(
|
||||
row.try_get("region")?,
|
||||
row.try_get("registered_by")?,
|
||||
Self::optional_unix_secs(row.try_get("last_heartbeat_at_unix_secs")?),
|
||||
row.try_get("avg_latency_ms")?,
|
||||
row.try_get("proxy_metadata")?,
|
||||
row.try_get("hardware_info")?,
|
||||
row.try_get("estimated_max_concurrency")?,
|
||||
Self::optional_unix_secs(row.try_get("tunnel_connected_at_unix_secs")?),
|
||||
row.try_get("remote_config")?,
|
||||
Self::optional_unix_secs(row.try_get("created_at_unix_secs")?),
|
||||
Self::optional_unix_secs(row.try_get("updated_at_unix_secs")?),
|
||||
row.try_get("region").map_postgres_err()?,
|
||||
row.try_get("registered_by").map_postgres_err()?,
|
||||
Self::optional_unix_secs(
|
||||
row.try_get("last_heartbeat_at_unix_secs")
|
||||
.map_postgres_err()?,
|
||||
),
|
||||
row.try_get("avg_latency_ms").map_postgres_err()?,
|
||||
row.try_get("proxy_metadata").map_postgres_err()?,
|
||||
row.try_get("hardware_info").map_postgres_err()?,
|
||||
row.try_get("estimated_max_concurrency")
|
||||
.map_postgres_err()?,
|
||||
Self::optional_unix_secs(
|
||||
row.try_get("tunnel_connected_at_unix_secs")
|
||||
.map_postgres_err()?,
|
||||
),
|
||||
row.try_get("remote_config").map_postgres_err()?,
|
||||
Self::optional_unix_secs(row.try_get("created_at_unix_secs").map_postgres_err()?),
|
||||
Self::optional_unix_secs(row.try_get("updated_at_unix_secs").map_postgres_err()?),
|
||||
))
|
||||
}
|
||||
|
||||
fn row_to_event(row: &PgRow) -> Result<StoredProxyNodeEvent, DataLayerError> {
|
||||
Ok(StoredProxyNodeEvent {
|
||||
id: row.try_get("id")?,
|
||||
node_id: row.try_get("node_id")?,
|
||||
event_type: row.try_get("event_type")?,
|
||||
detail: row.try_get("detail")?,
|
||||
created_at_unix_secs: Self::optional_unix_secs(row.try_get("created_at_unix_secs")?),
|
||||
id: row.try_get("id").map_postgres_err()?,
|
||||
node_id: row.try_get("node_id").map_postgres_err()?,
|
||||
event_type: row.try_get("event_type").map_postgres_err()?,
|
||||
detail: row.try_get("detail").map_postgres_err()?,
|
||||
created_at_unix_secs: Self::optional_unix_secs(
|
||||
row.try_get("created_at_unix_secs").map_postgres_err()?,
|
||||
),
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -194,7 +206,8 @@ impl ProxyNodeReadRepository for SqlxProxyNodeRepository {
|
||||
async fn list_proxy_nodes(&self) -> Result<Vec<StoredProxyNode>, DataLayerError> {
|
||||
let rows = sqlx::query(LIST_PROXY_NODES_SQL)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(Self::row_to_stored).collect()
|
||||
}
|
||||
|
||||
@@ -205,7 +218,8 @@ impl ProxyNodeReadRepository for SqlxProxyNodeRepository {
|
||||
let row = sqlx::query(FIND_PROXY_NODE_SQL)
|
||||
.bind(node_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.map(|row| Self::row_to_stored(&row)).transpose()
|
||||
}
|
||||
|
||||
@@ -218,7 +232,8 @@ impl ProxyNodeReadRepository for SqlxProxyNodeRepository {
|
||||
.bind(node_id)
|
||||
.bind(i64::try_from(limit).unwrap_or(i64::MAX))
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(Self::row_to_event).collect()
|
||||
}
|
||||
}
|
||||
@@ -256,7 +271,8 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository {
|
||||
.bind(mutation.dns_failures_delta)
|
||||
.bind(mutation.stream_errors_delta)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
self.find_proxy_node(&mutation.node_id).await
|
||||
}
|
||||
@@ -283,7 +299,7 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository {
|
||||
)
|
||||
});
|
||||
|
||||
let mut tx = self.pool.begin().await?;
|
||||
let mut tx = self.pool.begin().await.map_postgres_err()?;
|
||||
|
||||
if existing
|
||||
.tunnel_connected_at_unix_secs
|
||||
@@ -305,8 +321,9 @@ VALUES (
|
||||
.bind(event_type)
|
||||
.bind(format!("[stale_ignored] {event_detail}"))
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
tx.commit().await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
tx.commit().await.map_err(postgres_error)?;
|
||||
return self.find_proxy_node(&mutation.node_id).await;
|
||||
}
|
||||
|
||||
@@ -334,7 +351,8 @@ WHERE id = $1
|
||||
.bind(mutation.connected)
|
||||
.bind(observed_at_unix_secs.map(|value| value as f64))
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
@@ -355,9 +373,10 @@ VALUES (
|
||||
.bind(event_detail)
|
||||
.bind(observed_at_unix_secs.map(|value| value as f64))
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
tx.commit().await?;
|
||||
tx.commit().await.map_err(postgres_error)?;
|
||||
self.find_proxy_node(&mutation.node_id).await
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,7 +3,7 @@ use std::sync::RwLock;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::types::{
|
||||
use super::{
|
||||
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
mod memory;
|
||||
mod sql;
|
||||
mod types;
|
||||
|
||||
pub use memory::InMemoryProviderQuotaRepository;
|
||||
pub use sql::SqlxProviderQuotaRepository;
|
||||
pub use types::{
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use aether_data_contracts::repository::quota::{
|
||||
ProviderQuotaReadRepository, ProviderQuotaRepository, ProviderQuotaWriteRepository,
|
||||
StoredProviderQuotaSnapshot,
|
||||
};
|
||||
pub use memory::InMemoryProviderQuotaRepository;
|
||||
pub use sql::SqlxProviderQuotaRepository;
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{PgPool, Row};
|
||||
|
||||
use super::types::{
|
||||
use super::{
|
||||
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
|
||||
const FIND_BY_PROVIDER_ID_SQL: &str = r#"
|
||||
SELECT
|
||||
@@ -56,7 +56,8 @@ impl ProviderQuotaReadRepository for SqlxProviderQuotaRepository {
|
||||
let row = sqlx::query(FIND_BY_PROVIDER_ID_SQL)
|
||||
.bind(provider_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_row).transpose()
|
||||
}
|
||||
}
|
||||
@@ -69,21 +70,24 @@ impl ProviderQuotaWriteRepository for SqlxProviderQuotaRepository {
|
||||
DataLayerError::InvalidInput("provider quota reset timestamp overflow".to_string())
|
||||
})?)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(result.rows_affected() as usize)
|
||||
}
|
||||
}
|
||||
|
||||
fn map_row(row: &sqlx::postgres::PgRow) -> Result<StoredProviderQuotaSnapshot, DataLayerError> {
|
||||
StoredProviderQuotaSnapshot::new(
|
||||
row.try_get("provider_id")?,
|
||||
row.try_get("billing_type")?,
|
||||
row.try_get("monthly_quota_usd")?,
|
||||
row.try_get("monthly_used_usd")?,
|
||||
row.try_get("quota_reset_day")?,
|
||||
row.try_get("quota_last_reset_at_unix_secs")?,
|
||||
row.try_get("quota_expires_at_unix_secs")?,
|
||||
row.try_get("is_active")?,
|
||||
row.try_get("provider_id").map_postgres_err()?,
|
||||
row.try_get("billing_type").map_postgres_err()?,
|
||||
row.try_get("monthly_quota_usd").map_postgres_err()?,
|
||||
row.try_get("monthly_used_usd").map_postgres_err()?,
|
||||
row.try_get("quota_reset_day").map_postgres_err()?,
|
||||
row.try_get("quota_last_reset_at_unix_secs")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("quota_expires_at_unix_secs")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("is_active").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -1,71 +0,0 @@
|
||||
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_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
|
||||
{
|
||||
}
|
||||
@@ -3,7 +3,7 @@ use std::sync::{Arc, RwLock};
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::types::{SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput};
|
||||
use super::{SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput};
|
||||
use crate::repository::wallet::{InMemoryWalletRepository, StoredWalletSnapshot};
|
||||
use crate::DataLayerError;
|
||||
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
mod memory;
|
||||
mod sql;
|
||||
mod types;
|
||||
|
||||
pub use memory::InMemorySettlementRepository;
|
||||
pub use sql::SqlxSettlementRepository;
|
||||
pub use types::{
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use aether_data_contracts::repository::settlement::{
|
||||
SettlementRepository, SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput,
|
||||
};
|
||||
pub use memory::InMemorySettlementRepository;
|
||||
pub use sql::SqlxSettlementRepository;
|
||||
|
||||
@@ -1,7 +1,8 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{PgPool, Row};
|
||||
|
||||
use super::types::{SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput};
|
||||
use super::{SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput};
|
||||
use crate::error::SqlxResultExt;
|
||||
use crate::postgres::PostgresTransactionRunner;
|
||||
use crate::DataLayerError;
|
||||
|
||||
@@ -56,31 +57,42 @@ FOR UPDATE
|
||||
)
|
||||
.bind(&input.request_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
let Some(usage_row) = row else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let current_billing_status: String = usage_row.try_get("billing_status")?;
|
||||
let current_billing_status: String =
|
||||
usage_row.try_get("billing_status").map_postgres_err()?;
|
||||
if current_billing_status == "settled" || current_billing_status == "void" {
|
||||
return Ok(Some(StoredUsageSettlement {
|
||||
request_id: usage_row.try_get("request_id")?,
|
||||
wallet_id: usage_row.try_get("wallet_id")?,
|
||||
request_id: usage_row.try_get("request_id").map_postgres_err()?,
|
||||
wallet_id: usage_row.try_get("wallet_id").map_postgres_err()?,
|
||||
billing_status: current_billing_status,
|
||||
wallet_balance_before: usage_row.try_get("wallet_balance_before")?,
|
||||
wallet_balance_after: usage_row.try_get("wallet_balance_after")?,
|
||||
wallet_balance_before: usage_row
|
||||
.try_get("wallet_balance_before")
|
||||
.map_postgres_err()?,
|
||||
wallet_balance_after: usage_row
|
||||
.try_get("wallet_balance_after")
|
||||
.map_postgres_err()?,
|
||||
wallet_recharge_balance_before: usage_row
|
||||
.try_get("wallet_recharge_balance_before")?,
|
||||
.try_get("wallet_recharge_balance_before")
|
||||
.map_postgres_err()?,
|
||||
wallet_recharge_balance_after: usage_row
|
||||
.try_get("wallet_recharge_balance_after")?,
|
||||
.try_get("wallet_recharge_balance_after")
|
||||
.map_postgres_err()?,
|
||||
wallet_gift_balance_before: usage_row
|
||||
.try_get("wallet_gift_balance_before")?,
|
||||
.try_get("wallet_gift_balance_before")
|
||||
.map_postgres_err()?,
|
||||
wallet_gift_balance_after: usage_row
|
||||
.try_get("wallet_gift_balance_after")?,
|
||||
.try_get("wallet_gift_balance_after")
|
||||
.map_postgres_err()?,
|
||||
provider_monthly_used_usd: None,
|
||||
finalized_at_unix_secs: usage_row
|
||||
.try_get::<Option<i64>, _>("finalized_at_unix_secs")?
|
||||
.try_get::<Option<i64>, _>("finalized_at_unix_secs")
|
||||
.map_postgres_err()?
|
||||
.map(|value| value as u64),
|
||||
}));
|
||||
}
|
||||
@@ -136,7 +148,8 @@ LIMIT 1
|
||||
)
|
||||
.bind(api_key_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await?
|
||||
.await
|
||||
.map_postgres_err()?
|
||||
} else {
|
||||
None
|
||||
};
|
||||
@@ -161,16 +174,20 @@ LIMIT 1
|
||||
)
|
||||
.bind(user_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await?
|
||||
.await
|
||||
.map_postgres_err()?
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
if let Some(wallet_row) = wallet_row {
|
||||
let wallet_id: String = wallet_row.try_get("id")?;
|
||||
let before_recharge: f64 = wallet_row.try_get("balance")?;
|
||||
let before_gift: f64 = wallet_row.try_get("gift_balance")?;
|
||||
let limit_mode: String = wallet_row.try_get("limit_mode")?;
|
||||
let wallet_id: String = wallet_row.try_get("id").map_postgres_err()?;
|
||||
let before_recharge: f64 =
|
||||
wallet_row.try_get("balance").map_postgres_err()?;
|
||||
let before_gift: f64 =
|
||||
wallet_row.try_get("gift_balance").map_postgres_err()?;
|
||||
let limit_mode: String =
|
||||
wallet_row.try_get("limit_mode").map_postgres_err()?;
|
||||
let before_total = before_recharge + before_gift;
|
||||
let mut after_recharge = before_recharge;
|
||||
let mut after_gift = before_gift;
|
||||
@@ -196,7 +213,8 @@ WHERE id = $1
|
||||
.bind(after_gift)
|
||||
.bind(input.total_cost_usd)
|
||||
.execute(&mut **tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
settlement.wallet_id = Some(wallet_id.clone());
|
||||
settlement.wallet_balance_before = Some(before_total);
|
||||
@@ -229,7 +247,8 @@ WHERE request_id = $1
|
||||
.bind(before_gift)
|
||||
.bind(after_gift)
|
||||
.execute(&mut **tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
}
|
||||
|
||||
if let Some(provider_id) = input
|
||||
@@ -250,7 +269,8 @@ RETURNING CAST(monthly_used_usd AS DOUBLE PRECISION) AS monthly_used_usd
|
||||
.bind(provider_id)
|
||||
.bind(input.actual_total_cost_usd)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
settlement.provider_monthly_used_usd =
|
||||
quota_row.and_then(|row| row.try_get("monthly_used_usd").ok());
|
||||
}
|
||||
@@ -261,7 +281,8 @@ RETURNING CAST(monthly_used_usd AS DOUBLE PRECISION) AS monthly_used_usd
|
||||
.bind(final_billing_status)
|
||||
.bind(finalized_at)
|
||||
.execute(&mut **tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
Ok(Some(settlement))
|
||||
})
|
||||
|
||||
@@ -1,83 +0,0 @@
|
||||
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>,
|
||||
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 {}
|
||||
|
||||
#[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,
|
||||
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());
|
||||
}
|
||||
}
|
||||
@@ -7,7 +7,7 @@ use super::types::{
|
||||
ShadowResultWriteRepository, StoredShadowResult, UpsertShadowResult,
|
||||
};
|
||||
use crate::postgres::PostgresTransactionRunner;
|
||||
use crate::DataLayerError;
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
|
||||
const FIND_BY_TRACE_FINGERPRINT_SQL: &str = r#"
|
||||
SELECT
|
||||
@@ -148,7 +148,8 @@ impl SqlxShadowResultRepository {
|
||||
.bind(trace_id)
|
||||
.bind(request_fingerprint)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_shadow_result_row).transpose()
|
||||
}
|
||||
|
||||
@@ -167,7 +168,8 @@ impl SqlxShadowResultRepository {
|
||||
))
|
||||
})?)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
rows.iter().map(map_shadow_result_row).collect()
|
||||
}
|
||||
@@ -193,7 +195,8 @@ impl SqlxShadowResultRepository {
|
||||
.bind(result.created_at_unix_secs as f64)
|
||||
.bind(result.updated_at_unix_secs as f64)
|
||||
.fetch_one(&mut **tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
map_shadow_result_row(&row)
|
||||
}) as BoxFuture<'_, Result<StoredShadowResult, DataLayerError>>
|
||||
})
|
||||
@@ -237,22 +240,25 @@ fn match_status_to_database(status: ShadowResultMatchStatus) -> &'static str {
|
||||
fn map_shadow_result_row(
|
||||
row: &sqlx::postgres::PgRow,
|
||||
) -> Result<StoredShadowResult, DataLayerError> {
|
||||
let match_status =
|
||||
ShadowResultMatchStatus::from_database(row.try_get::<String, _>("match_status")?.as_str())?;
|
||||
let match_status = ShadowResultMatchStatus::from_database(
|
||||
row.try_get::<String, _>("match_status")
|
||||
.map_postgres_err()?
|
||||
.as_str(),
|
||||
)?;
|
||||
StoredShadowResult::new(
|
||||
row.try_get("trace_id")?,
|
||||
row.try_get("request_fingerprint")?,
|
||||
row.try_get("request_id")?,
|
||||
row.try_get("route_family")?,
|
||||
row.try_get("route_kind")?,
|
||||
row.try_get("candidate_id")?,
|
||||
row.try_get("rust_result_digest")?,
|
||||
row.try_get("python_result_digest")?,
|
||||
row.try_get("trace_id").map_postgres_err()?,
|
||||
row.try_get("request_fingerprint").map_postgres_err()?,
|
||||
row.try_get("request_id").map_postgres_err()?,
|
||||
row.try_get("route_family").map_postgres_err()?,
|
||||
row.try_get("route_kind").map_postgres_err()?,
|
||||
row.try_get("candidate_id").map_postgres_err()?,
|
||||
row.try_get("rust_result_digest").map_postgres_err()?,
|
||||
row.try_get("python_result_digest").map_postgres_err()?,
|
||||
match_status,
|
||||
row.try_get("status_code")?,
|
||||
row.try_get("error_message")?,
|
||||
row.try_get("created_at_unix_secs")?,
|
||||
row.try_get("updated_at_unix_secs")?,
|
||||
row.try_get("status_code").map_postgres_err()?,
|
||||
row.try_get("error_message").map_postgres_err()?,
|
||||
row.try_get("created_at_unix_secs").map_postgres_err()?,
|
||||
row.try_get("updated_at_unix_secs").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@ use std::sync::RwLock;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::types::{
|
||||
use super::{
|
||||
StoredProviderUsageSummary, StoredProviderUsageWindow, StoredRequestUsageAudit,
|
||||
UpsertUsageRecord, UsageAuditListQuery, UsageReadRepository, UsageWriteRepository,
|
||||
};
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
mod memory;
|
||||
mod sql;
|
||||
mod types;
|
||||
|
||||
pub use memory::InMemoryUsageReadRepository;
|
||||
pub use sql::SqlxUsageReadRepository;
|
||||
pub use types::{
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use aether_data_contracts::repository::usage::{
|
||||
StoredProviderUsageSummary, StoredProviderUsageWindow, StoredRequestUsageAudit,
|
||||
UpsertUsageRecord, UsageAuditListQuery, UsageReadRepository, UsageRepository,
|
||||
UsageWriteRepository,
|
||||
};
|
||||
pub use memory::InMemoryUsageReadRepository;
|
||||
pub use sql::SqlxUsageReadRepository;
|
||||
|
||||
@@ -3,12 +3,12 @@ use futures_util::future::BoxFuture;
|
||||
use sqlx::{PgPool, Postgres, QueryBuilder, Row};
|
||||
use uuid::Uuid;
|
||||
|
||||
use super::types::{
|
||||
use super::{
|
||||
StoredProviderUsageSummary, StoredRequestUsageAudit, UpsertUsageRecord, UsageAuditListQuery,
|
||||
UsageReadRepository, UsageWriteRepository,
|
||||
};
|
||||
use crate::postgres::PostgresTransactionRunner;
|
||||
use crate::DataLayerError;
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
|
||||
const FIND_BY_REQUEST_ID_SQL: &str = r#"
|
||||
SELECT
|
||||
@@ -497,7 +497,8 @@ impl SqlxUsageReadRepository {
|
||||
let row = sqlx::query(FIND_BY_REQUEST_ID_SQL)
|
||||
.bind(request_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_usage_row).transpose()
|
||||
}
|
||||
|
||||
@@ -508,7 +509,8 @@ impl SqlxUsageReadRepository {
|
||||
let row = sqlx::query(FIND_BY_ID_SQL)
|
||||
.bind(id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_usage_row).transpose()
|
||||
}
|
||||
|
||||
@@ -521,14 +523,26 @@ impl SqlxUsageReadRepository {
|
||||
.bind(provider_id)
|
||||
.bind(since_unix_secs as f64)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
Ok(StoredProviderUsageSummary {
|
||||
total_requests: row.try_get::<i64, _>("total_requests")?.max(0) as u64,
|
||||
successful_requests: row.try_get::<i64, _>("successful_requests")?.max(0) as u64,
|
||||
failed_requests: row.try_get::<i64, _>("failed_requests")?.max(0) as u64,
|
||||
avg_response_time_ms: row.try_get::<f64, _>("avg_response_time_ms")?,
|
||||
total_cost_usd: row.try_get::<f64, _>("total_cost_usd")?,
|
||||
total_requests: row
|
||||
.try_get::<i64, _>("total_requests")
|
||||
.map_postgres_err()?
|
||||
.max(0) as u64,
|
||||
successful_requests: row
|
||||
.try_get::<i64, _>("successful_requests")
|
||||
.map_postgres_err()?
|
||||
.max(0) as u64,
|
||||
failed_requests: row
|
||||
.try_get::<i64, _>("failed_requests")
|
||||
.map_postgres_err()?
|
||||
.max(0) as u64,
|
||||
avg_response_time_ms: row
|
||||
.try_get::<f64, _>("avg_response_time_ms")
|
||||
.map_postgres_err()?,
|
||||
total_cost_usd: row.try_get::<f64, _>("total_cost_usd").map_postgres_err()?,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -573,7 +587,11 @@ impl SqlxUsageReadRepository {
|
||||
}
|
||||
|
||||
builder.push(" ORDER BY created_at ASC, request_id ASC");
|
||||
let rows = builder.build().fetch_all(&self.pool).await?;
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_usage_row).collect()
|
||||
}
|
||||
|
||||
@@ -593,7 +611,11 @@ impl SqlxUsageReadRepository {
|
||||
.push_bind(i64::try_from(limit).map_err(|_| {
|
||||
DataLayerError::InvalidInput(format!("invalid recent usage limit: {limit}"))
|
||||
})?);
|
||||
let rows = builder.build().fetch_all(&self.pool).await?;
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_usage_row).collect()
|
||||
}
|
||||
|
||||
@@ -608,12 +630,16 @@ impl SqlxUsageReadRepository {
|
||||
let rows = sqlx::query(SUMMARIZE_TOTAL_TOKENS_BY_API_KEY_IDS_SQL)
|
||||
.bind(api_key_ids)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
let mut totals = std::collections::BTreeMap::new();
|
||||
for row in rows {
|
||||
let api_key_id: String = row.try_get("api_key_id")?;
|
||||
let total_tokens = row.try_get::<i64, _>("total_tokens")?.max(0) as u64;
|
||||
let api_key_id: String = row.try_get("api_key_id").map_postgres_err()?;
|
||||
let total_tokens = row
|
||||
.try_get::<i64, _>("total_tokens")
|
||||
.map_postgres_err()?
|
||||
.max(0) as u64;
|
||||
totals.insert(api_key_id, total_tokens);
|
||||
}
|
||||
Ok(totals)
|
||||
@@ -689,7 +715,8 @@ impl SqlxUsageReadRepository {
|
||||
.bind(usage.finalized_at_unix_secs.map(|value| value as f64))
|
||||
.bind(usage.created_at_unix_secs.map(|value| value as f64))
|
||||
.fetch_one(&mut **tx)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
map_usage_row(&row)
|
||||
}) as BoxFuture<'_, Result<StoredRequestUsageAudit, DataLayerError>>
|
||||
})
|
||||
@@ -756,65 +783,71 @@ impl UsageWriteRepository for SqlxUsageReadRepository {
|
||||
|
||||
fn map_usage_row(row: &sqlx::postgres::PgRow) -> Result<StoredRequestUsageAudit, DataLayerError> {
|
||||
let mut usage = StoredRequestUsageAudit::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("request_id")?,
|
||||
row.try_get("user_id")?,
|
||||
row.try_get("api_key_id")?,
|
||||
row.try_get("username")?,
|
||||
row.try_get("api_key_name")?,
|
||||
row.try_get("provider_name")?,
|
||||
row.try_get("model")?,
|
||||
row.try_get("target_model")?,
|
||||
row.try_get("provider_id")?,
|
||||
row.try_get("provider_endpoint_id")?,
|
||||
row.try_get("provider_api_key_id")?,
|
||||
row.try_get("request_type")?,
|
||||
row.try_get("api_format")?,
|
||||
row.try_get("api_family")?,
|
||||
row.try_get("endpoint_kind")?,
|
||||
row.try_get("endpoint_api_format")?,
|
||||
row.try_get("provider_api_family")?,
|
||||
row.try_get("provider_endpoint_kind")?,
|
||||
row.try_get("has_format_conversion")?,
|
||||
row.try_get("is_stream")?,
|
||||
row.try_get("input_tokens")?,
|
||||
row.try_get("output_tokens")?,
|
||||
row.try_get("total_tokens")?,
|
||||
row.try_get("total_cost_usd")?,
|
||||
row.try_get("actual_total_cost_usd")?,
|
||||
row.try_get("status_code")?,
|
||||
row.try_get("error_message")?,
|
||||
row.try_get("error_category")?,
|
||||
row.try_get("response_time_ms")?,
|
||||
row.try_get("first_byte_time_ms")?,
|
||||
row.try_get("status")?,
|
||||
row.try_get("billing_status")?,
|
||||
row.try_get("created_at_unix_secs")?,
|
||||
row.try_get("updated_at_unix_secs")?,
|
||||
row.try_get("finalized_at_unix_secs")?,
|
||||
row.try_get("id").map_postgres_err()?,
|
||||
row.try_get("request_id").map_postgres_err()?,
|
||||
row.try_get("user_id").map_postgres_err()?,
|
||||
row.try_get("api_key_id").map_postgres_err()?,
|
||||
row.try_get("username").map_postgres_err()?,
|
||||
row.try_get("api_key_name").map_postgres_err()?,
|
||||
row.try_get("provider_name").map_postgres_err()?,
|
||||
row.try_get("model").map_postgres_err()?,
|
||||
row.try_get("target_model").map_postgres_err()?,
|
||||
row.try_get("provider_id").map_postgres_err()?,
|
||||
row.try_get("provider_endpoint_id").map_postgres_err()?,
|
||||
row.try_get("provider_api_key_id").map_postgres_err()?,
|
||||
row.try_get("request_type").map_postgres_err()?,
|
||||
row.try_get("api_format").map_postgres_err()?,
|
||||
row.try_get("api_family").map_postgres_err()?,
|
||||
row.try_get("endpoint_kind").map_postgres_err()?,
|
||||
row.try_get("endpoint_api_format").map_postgres_err()?,
|
||||
row.try_get("provider_api_family").map_postgres_err()?,
|
||||
row.try_get("provider_endpoint_kind").map_postgres_err()?,
|
||||
row.try_get("has_format_conversion").map_postgres_err()?,
|
||||
row.try_get("is_stream").map_postgres_err()?,
|
||||
row.try_get("input_tokens").map_postgres_err()?,
|
||||
row.try_get("output_tokens").map_postgres_err()?,
|
||||
row.try_get("total_tokens").map_postgres_err()?,
|
||||
row.try_get("total_cost_usd").map_postgres_err()?,
|
||||
row.try_get("actual_total_cost_usd").map_postgres_err()?,
|
||||
row.try_get("status_code").map_postgres_err()?,
|
||||
row.try_get("error_message").map_postgres_err()?,
|
||||
row.try_get("error_category").map_postgres_err()?,
|
||||
row.try_get("response_time_ms").map_postgres_err()?,
|
||||
row.try_get("first_byte_time_ms").map_postgres_err()?,
|
||||
row.try_get("status").map_postgres_err()?,
|
||||
row.try_get("billing_status").map_postgres_err()?,
|
||||
row.try_get("created_at_unix_secs").map_postgres_err()?,
|
||||
row.try_get("updated_at_unix_secs").map_postgres_err()?,
|
||||
row.try_get("finalized_at_unix_secs").map_postgres_err()?,
|
||||
)?;
|
||||
usage.cache_creation_input_tokens = row
|
||||
.try_get::<Option<i32>, _>("cache_creation_input_tokens")?
|
||||
.try_get::<Option<i32>, _>("cache_creation_input_tokens")
|
||||
.map_postgres_err()?
|
||||
.map(|value| to_u64(value, "usage.cache_creation_input_tokens"))
|
||||
.transpose()?
|
||||
.unwrap_or_default();
|
||||
usage.cache_read_input_tokens = row
|
||||
.try_get::<Option<i32>, _>("cache_read_input_tokens")?
|
||||
.try_get::<Option<i32>, _>("cache_read_input_tokens")
|
||||
.map_postgres_err()?
|
||||
.map(|value| to_u64(value, "usage.cache_read_input_tokens"))
|
||||
.transpose()?
|
||||
.unwrap_or_default();
|
||||
usage.cache_creation_cost_usd = row.try_get::<f64, _>("cache_creation_cost_usd")?;
|
||||
usage.cache_read_cost_usd = row.try_get::<f64, _>("cache_read_cost_usd")?;
|
||||
usage.output_price_per_1m = row.try_get("output_price_per_1m")?;
|
||||
usage.request_headers = row.try_get("request_headers")?;
|
||||
usage.request_body = row.try_get("request_body")?;
|
||||
usage.provider_request_headers = row.try_get("provider_request_headers")?;
|
||||
usage.provider_request_body = row.try_get("provider_request_body")?;
|
||||
usage.response_headers = row.try_get("response_headers")?;
|
||||
usage.response_body = row.try_get("response_body")?;
|
||||
usage.client_response_headers = row.try_get("client_response_headers")?;
|
||||
usage.client_response_body = row.try_get("client_response_body")?;
|
||||
usage.request_metadata = row.try_get("request_metadata")?;
|
||||
usage.cache_creation_cost_usd = row
|
||||
.try_get::<f64, _>("cache_creation_cost_usd")
|
||||
.map_postgres_err()?;
|
||||
usage.cache_read_cost_usd = row
|
||||
.try_get::<f64, _>("cache_read_cost_usd")
|
||||
.map_postgres_err()?;
|
||||
usage.output_price_per_1m = row.try_get("output_price_per_1m").map_postgres_err()?;
|
||||
usage.request_headers = row.try_get("request_headers").map_postgres_err()?;
|
||||
usage.request_body = row.try_get("request_body").map_postgres_err()?;
|
||||
usage.provider_request_headers = row.try_get("provider_request_headers").map_postgres_err()?;
|
||||
usage.provider_request_body = row.try_get("provider_request_body").map_postgres_err()?;
|
||||
usage.response_headers = row.try_get("response_headers").map_postgres_err()?;
|
||||
usage.response_body = row.try_get("response_body").map_postgres_err()?;
|
||||
usage.client_response_headers = row.try_get("client_response_headers").map_postgres_err()?;
|
||||
usage.client_response_body = row.try_get("client_response_body").map_postgres_err()?;
|
||||
usage.request_metadata = row.try_get("request_metadata").map_postgres_err()?;
|
||||
Ok(usage)
|
||||
}
|
||||
|
||||
|
||||
@@ -1,639 +0,0 @@
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredRequestUsageAudit {
|
||||
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 provider_name: String,
|
||||
pub model: String,
|
||||
pub target_model: Option<String>,
|
||||
pub provider_id: Option<String>,
|
||||
pub provider_endpoint_id: Option<String>,
|
||||
pub provider_api_key_id: Option<String>,
|
||||
pub request_type: Option<String>,
|
||||
pub api_format: Option<String>,
|
||||
pub api_family: Option<String>,
|
||||
pub endpoint_kind: Option<String>,
|
||||
pub endpoint_api_format: Option<String>,
|
||||
pub provider_api_family: Option<String>,
|
||||
pub provider_endpoint_kind: Option<String>,
|
||||
pub has_format_conversion: bool,
|
||||
pub is_stream: bool,
|
||||
pub input_tokens: u64,
|
||||
pub output_tokens: u64,
|
||||
pub total_tokens: u64,
|
||||
pub cache_creation_input_tokens: u64,
|
||||
pub cache_read_input_tokens: u64,
|
||||
pub cache_creation_cost_usd: f64,
|
||||
pub cache_read_cost_usd: f64,
|
||||
pub output_price_per_1m: Option<f64>,
|
||||
pub total_cost_usd: f64,
|
||||
pub actual_total_cost_usd: f64,
|
||||
pub status_code: Option<u16>,
|
||||
pub error_message: Option<String>,
|
||||
pub error_category: Option<String>,
|
||||
pub response_time_ms: Option<u64>,
|
||||
pub first_byte_time_ms: Option<u64>,
|
||||
pub status: String,
|
||||
pub billing_status: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub request_headers: Option<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub request_body: Option<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub provider_request_headers: Option<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub provider_request_body: Option<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub response_headers: Option<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub response_body: Option<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub client_response_headers: Option<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub client_response_body: Option<Value>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub request_metadata: Option<Value>,
|
||||
pub created_at_unix_secs: u64,
|
||||
pub updated_at_unix_secs: u64,
|
||||
pub finalized_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
impl StoredRequestUsageAudit {
|
||||
#[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>,
|
||||
provider_name: String,
|
||||
model: String,
|
||||
target_model: Option<String>,
|
||||
provider_id: Option<String>,
|
||||
provider_endpoint_id: Option<String>,
|
||||
provider_api_key_id: Option<String>,
|
||||
request_type: Option<String>,
|
||||
api_format: Option<String>,
|
||||
api_family: Option<String>,
|
||||
endpoint_kind: Option<String>,
|
||||
endpoint_api_format: Option<String>,
|
||||
provider_api_family: Option<String>,
|
||||
provider_endpoint_kind: Option<String>,
|
||||
has_format_conversion: bool,
|
||||
is_stream: bool,
|
||||
input_tokens: i32,
|
||||
output_tokens: i32,
|
||||
total_tokens: i32,
|
||||
total_cost_usd: f64,
|
||||
actual_total_cost_usd: f64,
|
||||
status_code: Option<i32>,
|
||||
error_message: Option<String>,
|
||||
error_category: Option<String>,
|
||||
response_time_ms: Option<i32>,
|
||||
first_byte_time_ms: Option<i32>,
|
||||
status: String,
|
||||
billing_status: String,
|
||||
created_at_unix_secs: i64,
|
||||
updated_at_unix_secs: i64,
|
||||
finalized_at_unix_secs: Option<i64>,
|
||||
) -> Result<Self, crate::DataLayerError> {
|
||||
if request_id.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"usage.request_id is empty".to_string(),
|
||||
));
|
||||
}
|
||||
if provider_name.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"usage.provider_name is empty".to_string(),
|
||||
));
|
||||
}
|
||||
if model.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"usage.model is empty".to_string(),
|
||||
));
|
||||
}
|
||||
if status.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"usage.status is empty".to_string(),
|
||||
));
|
||||
}
|
||||
if billing_status.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"usage.billing_status is empty".to_string(),
|
||||
));
|
||||
}
|
||||
if !total_cost_usd.is_finite() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"usage.total_cost_usd is not finite".to_string(),
|
||||
));
|
||||
}
|
||||
if !actual_total_cost_usd.is_finite() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"usage.actual_total_cost_usd is not finite".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
id,
|
||||
request_id,
|
||||
user_id,
|
||||
api_key_id,
|
||||
username,
|
||||
api_key_name,
|
||||
provider_name,
|
||||
model,
|
||||
target_model,
|
||||
provider_id,
|
||||
provider_endpoint_id,
|
||||
provider_api_key_id,
|
||||
request_type,
|
||||
api_format,
|
||||
api_family,
|
||||
endpoint_kind,
|
||||
endpoint_api_format,
|
||||
provider_api_family,
|
||||
provider_endpoint_kind,
|
||||
has_format_conversion,
|
||||
is_stream,
|
||||
input_tokens: parse_u64(input_tokens, "usage.input_tokens")?,
|
||||
output_tokens: parse_u64(output_tokens, "usage.output_tokens")?,
|
||||
total_tokens: parse_u64(total_tokens, "usage.total_tokens")?,
|
||||
cache_creation_input_tokens: 0,
|
||||
cache_read_input_tokens: 0,
|
||||
cache_creation_cost_usd: 0.0,
|
||||
cache_read_cost_usd: 0.0,
|
||||
output_price_per_1m: None,
|
||||
total_cost_usd,
|
||||
actual_total_cost_usd,
|
||||
status_code: parse_u16(status_code, "usage.status_code")?,
|
||||
error_message,
|
||||
error_category,
|
||||
response_time_ms: parse_optional_u64(response_time_ms, "usage.response_time_ms")?,
|
||||
first_byte_time_ms: parse_optional_u64(first_byte_time_ms, "usage.first_byte_time_ms")?,
|
||||
status,
|
||||
billing_status,
|
||||
request_headers: None,
|
||||
request_body: None,
|
||||
provider_request_headers: None,
|
||||
provider_request_body: None,
|
||||
response_headers: None,
|
||||
response_body: None,
|
||||
client_response_headers: None,
|
||||
client_response_body: None,
|
||||
request_metadata: None,
|
||||
created_at_unix_secs: parse_timestamp(
|
||||
created_at_unix_secs,
|
||||
"usage.created_at_unix_secs",
|
||||
)?,
|
||||
updated_at_unix_secs: parse_timestamp(
|
||||
updated_at_unix_secs,
|
||||
"usage.updated_at_unix_secs",
|
||||
)?,
|
||||
finalized_at_unix_secs: finalized_at_unix_secs
|
||||
.map(|value| parse_timestamp(value, "usage.finalized_at_unix_secs"))
|
||||
.transpose()?,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn with_cache_input_tokens(
|
||||
mut self,
|
||||
cache_creation_input_tokens: u64,
|
||||
cache_read_input_tokens: u64,
|
||||
) -> Self {
|
||||
self.cache_creation_input_tokens = cache_creation_input_tokens;
|
||||
self.cache_read_input_tokens = cache_read_input_tokens;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredProviderUsageWindow {
|
||||
pub provider_id: String,
|
||||
pub window_start_unix_secs: u64,
|
||||
pub total_requests: u64,
|
||||
pub successful_requests: u64,
|
||||
pub failed_requests: u64,
|
||||
pub avg_response_time_ms: f64,
|
||||
pub total_cost_usd: f64,
|
||||
}
|
||||
|
||||
impl StoredProviderUsageWindow {
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn new(
|
||||
provider_id: String,
|
||||
window_start_unix_secs: i64,
|
||||
total_requests: i64,
|
||||
successful_requests: i64,
|
||||
failed_requests: i64,
|
||||
avg_response_time_ms: f64,
|
||||
total_cost_usd: f64,
|
||||
) -> Result<Self, crate::DataLayerError> {
|
||||
if provider_id.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"provider usage window provider_id is empty".to_string(),
|
||||
));
|
||||
}
|
||||
if !avg_response_time_ms.is_finite() || !total_cost_usd.is_finite() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"provider usage window value is not finite".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
provider_id,
|
||||
window_start_unix_secs: parse_timestamp(
|
||||
window_start_unix_secs,
|
||||
"provider_usage_tracking.window_start_unix_secs",
|
||||
)?,
|
||||
total_requests: parse_timestamp(
|
||||
total_requests,
|
||||
"provider_usage_tracking.total_requests",
|
||||
)?,
|
||||
successful_requests: parse_timestamp(
|
||||
successful_requests,
|
||||
"provider_usage_tracking.successful_requests",
|
||||
)?,
|
||||
failed_requests: parse_timestamp(
|
||||
failed_requests,
|
||||
"provider_usage_tracking.failed_requests",
|
||||
)?,
|
||||
avg_response_time_ms,
|
||||
total_cost_usd,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Default, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredProviderUsageSummary {
|
||||
pub total_requests: u64,
|
||||
pub successful_requests: u64,
|
||||
pub failed_requests: u64,
|
||||
pub avg_response_time_ms: f64,
|
||||
pub total_cost_usd: f64,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
|
||||
pub struct UsageAuditListQuery {
|
||||
pub created_from_unix_secs: Option<u64>,
|
||||
pub created_until_unix_secs: Option<u64>,
|
||||
pub user_id: Option<String>,
|
||||
pub provider_name: Option<String>,
|
||||
pub model: Option<String>,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait UsageReadRepository: Send + Sync {
|
||||
async fn find_by_id(
|
||||
&self,
|
||||
id: &str,
|
||||
) -> Result<Option<StoredRequestUsageAudit>, crate::DataLayerError>;
|
||||
|
||||
async fn find_by_request_id(
|
||||
&self,
|
||||
request_id: &str,
|
||||
) -> Result<Option<StoredRequestUsageAudit>, crate::DataLayerError>;
|
||||
|
||||
async fn list_usage_audits(
|
||||
&self,
|
||||
query: &UsageAuditListQuery,
|
||||
) -> Result<Vec<StoredRequestUsageAudit>, crate::DataLayerError>;
|
||||
|
||||
async fn list_recent_usage_audits(
|
||||
&self,
|
||||
user_id: Option<&str>,
|
||||
limit: usize,
|
||||
) -> Result<Vec<StoredRequestUsageAudit>, crate::DataLayerError>;
|
||||
|
||||
async fn summarize_total_tokens_by_api_key_ids(
|
||||
&self,
|
||||
api_key_ids: &[String],
|
||||
) -> Result<std::collections::BTreeMap<String, u64>, crate::DataLayerError>;
|
||||
|
||||
async fn summarize_provider_usage_since(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
since_unix_secs: u64,
|
||||
) -> Result<StoredProviderUsageSummary, crate::DataLayerError>;
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct UpsertUsageRecord {
|
||||
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 provider_name: String,
|
||||
pub model: String,
|
||||
pub target_model: Option<String>,
|
||||
pub provider_id: Option<String>,
|
||||
pub provider_endpoint_id: Option<String>,
|
||||
pub provider_api_key_id: Option<String>,
|
||||
pub request_type: Option<String>,
|
||||
pub api_format: Option<String>,
|
||||
pub api_family: Option<String>,
|
||||
pub endpoint_kind: Option<String>,
|
||||
pub endpoint_api_format: Option<String>,
|
||||
pub provider_api_family: Option<String>,
|
||||
pub provider_endpoint_kind: Option<String>,
|
||||
pub has_format_conversion: Option<bool>,
|
||||
pub is_stream: Option<bool>,
|
||||
pub input_tokens: Option<u64>,
|
||||
pub output_tokens: Option<u64>,
|
||||
pub total_tokens: Option<u64>,
|
||||
pub cache_creation_input_tokens: Option<u64>,
|
||||
pub cache_read_input_tokens: Option<u64>,
|
||||
pub cache_creation_cost_usd: Option<f64>,
|
||||
pub cache_read_cost_usd: Option<f64>,
|
||||
pub output_price_per_1m: Option<f64>,
|
||||
pub total_cost_usd: Option<f64>,
|
||||
pub actual_total_cost_usd: Option<f64>,
|
||||
pub status_code: Option<u16>,
|
||||
pub error_message: Option<String>,
|
||||
pub error_category: Option<String>,
|
||||
pub response_time_ms: Option<u64>,
|
||||
pub first_byte_time_ms: Option<u64>,
|
||||
pub status: String,
|
||||
pub billing_status: String,
|
||||
pub request_headers: Option<Value>,
|
||||
pub request_body: Option<Value>,
|
||||
pub provider_request_headers: Option<Value>,
|
||||
pub provider_request_body: Option<Value>,
|
||||
pub response_headers: Option<Value>,
|
||||
pub response_body: Option<Value>,
|
||||
pub client_response_headers: Option<Value>,
|
||||
pub client_response_body: Option<Value>,
|
||||
pub request_metadata: Option<Value>,
|
||||
pub finalized_at_unix_secs: Option<u64>,
|
||||
pub created_at_unix_secs: Option<u64>,
|
||||
pub updated_at_unix_secs: u64,
|
||||
}
|
||||
|
||||
impl UpsertUsageRecord {
|
||||
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
|
||||
if self.request_id.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert request_id cannot be empty".to_string(),
|
||||
));
|
||||
}
|
||||
if self.provider_name.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert provider_name cannot be empty".to_string(),
|
||||
));
|
||||
}
|
||||
if self.model.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert model cannot be empty".to_string(),
|
||||
));
|
||||
}
|
||||
if self.status.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert status cannot be empty".to_string(),
|
||||
));
|
||||
}
|
||||
if self.billing_status.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert billing_status cannot be empty".to_string(),
|
||||
));
|
||||
}
|
||||
if let Some(value) = self.total_cost_usd {
|
||||
if !value.is_finite() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert total_cost_usd must be finite".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
if let Some(value) = self.cache_creation_cost_usd {
|
||||
if !value.is_finite() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert cache_creation_cost_usd must be finite".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
if let Some(value) = self.cache_read_cost_usd {
|
||||
if !value.is_finite() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert cache_read_cost_usd must be finite".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
if let Some(value) = self.output_price_per_1m {
|
||||
if !value.is_finite() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert output_price_per_1m must be finite".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
if let Some(value) = self.actual_total_cost_usd {
|
||||
if !value.is_finite() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert actual_total_cost_usd must be finite".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait UsageWriteRepository: Send + Sync {
|
||||
async fn upsert(
|
||||
&self,
|
||||
usage: UpsertUsageRecord,
|
||||
) -> Result<StoredRequestUsageAudit, crate::DataLayerError>;
|
||||
}
|
||||
|
||||
pub trait UsageRepository: UsageReadRepository + UsageWriteRepository + Send + Sync {}
|
||||
|
||||
impl<T> UsageRepository for T where T: UsageReadRepository + UsageWriteRepository + Send + Sync {}
|
||||
|
||||
fn parse_u64(value: i32, field_name: &str) -> Result<u64, crate::DataLayerError> {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
crate::DataLayerError::UnexpectedValue(format!("invalid {field_name}: {value}"))
|
||||
})
|
||||
}
|
||||
|
||||
fn parse_optional_u64(
|
||||
value: Option<i32>,
|
||||
field_name: &str,
|
||||
) -> Result<Option<u64>, crate::DataLayerError> {
|
||||
value
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
crate::DataLayerError::UnexpectedValue(format!("invalid {field_name}: {value}"))
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn parse_u16(value: Option<i32>, field_name: &str) -> Result<Option<u16>, crate::DataLayerError> {
|
||||
value
|
||||
.map(|value| {
|
||||
u16::try_from(value).map_err(|_| {
|
||||
crate::DataLayerError::UnexpectedValue(format!("invalid {field_name}: {value}"))
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn parse_timestamp(value: i64, field_name: &str) -> Result<u64, crate::DataLayerError> {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
crate::DataLayerError::UnexpectedValue(format!("invalid {field_name}: {value}"))
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{StoredRequestUsageAudit, UpsertUsageRecord};
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn rejects_empty_request_id() {
|
||||
assert!(StoredRequestUsageAudit::new(
|
||||
"usage-1".to_string(),
|
||||
"".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
"OpenAI".to_string(),
|
||||
"gpt-4.1".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
Some("chat".to_string()),
|
||||
Some("openai:chat".to_string()),
|
||||
Some("openai".to_string()),
|
||||
Some("chat".to_string()),
|
||||
Some("openai:chat".to_string()),
|
||||
Some("openai".to_string()),
|
||||
Some("chat".to_string()),
|
||||
false,
|
||||
false,
|
||||
10,
|
||||
20,
|
||||
30,
|
||||
0.1,
|
||||
0.1,
|
||||
Some(200),
|
||||
None,
|
||||
None,
|
||||
Some(120),
|
||||
Some(80),
|
||||
"completed".to_string(),
|
||||
"settled".to_string(),
|
||||
100,
|
||||
101,
|
||||
Some(102),
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_negative_token_count() {
|
||||
assert!(StoredRequestUsageAudit::new(
|
||||
"usage-1".to_string(),
|
||||
"req-1".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
"OpenAI".to_string(),
|
||||
"gpt-4.1".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
Some("chat".to_string()),
|
||||
Some("openai:chat".to_string()),
|
||||
Some("openai".to_string()),
|
||||
Some("chat".to_string()),
|
||||
Some("openai:chat".to_string()),
|
||||
Some("openai".to_string()),
|
||||
Some("chat".to_string()),
|
||||
false,
|
||||
false,
|
||||
-1,
|
||||
20,
|
||||
30,
|
||||
0.1,
|
||||
0.1,
|
||||
Some(200),
|
||||
None,
|
||||
None,
|
||||
Some(120),
|
||||
Some(80),
|
||||
"completed".to_string(),
|
||||
"settled".to_string(),
|
||||
100,
|
||||
101,
|
||||
Some(102),
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_upsert_payload() {
|
||||
let record = UpsertUsageRecord {
|
||||
request_id: "".to_string(),
|
||||
user_id: None,
|
||||
api_key_id: None,
|
||||
username: None,
|
||||
api_key_name: None,
|
||||
provider_name: "openai".to_string(),
|
||||
model: "gpt-5".to_string(),
|
||||
target_model: None,
|
||||
provider_id: None,
|
||||
provider_endpoint_id: None,
|
||||
provider_api_key_id: None,
|
||||
request_type: Some("chat".to_string()),
|
||||
api_format: Some("openai:chat".to_string()),
|
||||
api_family: Some("openai".to_string()),
|
||||
endpoint_kind: Some("chat".to_string()),
|
||||
endpoint_api_format: Some("openai:chat".to_string()),
|
||||
provider_api_family: Some("openai".to_string()),
|
||||
provider_endpoint_kind: Some("chat".to_string()),
|
||||
has_format_conversion: Some(false),
|
||||
is_stream: Some(false),
|
||||
input_tokens: Some(10),
|
||||
output_tokens: Some(20),
|
||||
total_tokens: Some(30),
|
||||
cache_creation_input_tokens: None,
|
||||
cache_read_input_tokens: None,
|
||||
cache_creation_cost_usd: None,
|
||||
cache_read_cost_usd: None,
|
||||
output_price_per_1m: None,
|
||||
total_cost_usd: None,
|
||||
actual_total_cost_usd: None,
|
||||
status_code: Some(200),
|
||||
error_message: None,
|
||||
error_category: None,
|
||||
response_time_ms: Some(120),
|
||||
first_byte_time_ms: None,
|
||||
status: "completed".to_string(),
|
||||
billing_status: "pending".to_string(),
|
||||
request_headers: Some(json!({"authorization": "Bearer test"})),
|
||||
request_body: Some(json!({"model": "gpt-5"})),
|
||||
provider_request_headers: None,
|
||||
provider_request_body: None,
|
||||
response_headers: None,
|
||||
response_body: None,
|
||||
client_response_headers: None,
|
||||
client_response_body: None,
|
||||
request_metadata: None,
|
||||
finalized_at_unix_secs: None,
|
||||
created_at_unix_secs: Some(100),
|
||||
updated_at_unix_secs: 101,
|
||||
};
|
||||
|
||||
assert!(record.validate().is_err());
|
||||
}
|
||||
}
|
||||
@@ -5,7 +5,7 @@ use super::types::{
|
||||
StoredUserAuthRecord, StoredUserExportRow, StoredUserSummary, UserExportListQuery,
|
||||
UserExportSummary, UserReadRepository,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
|
||||
const LIST_USERS_BY_IDS_SQL: &str = r#"
|
||||
SELECT
|
||||
@@ -192,7 +192,8 @@ impl SqlxUserReadRepository {
|
||||
let rows = sqlx::query(LIST_USERS_BY_IDS_SQL)
|
||||
.bind(user_ids)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_user_row).collect()
|
||||
}
|
||||
|
||||
@@ -201,14 +202,16 @@ impl SqlxUserReadRepository {
|
||||
) -> Result<Vec<StoredUserExportRow>, DataLayerError> {
|
||||
let rows = sqlx::query(LIST_NON_ADMIN_EXPORT_USERS_SQL)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_user_export_row).collect()
|
||||
}
|
||||
|
||||
pub async fn list_export_users(&self) -> Result<Vec<StoredUserExportRow>, DataLayerError> {
|
||||
let rows = sqlx::query(LIST_EXPORT_USERS_SQL)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_user_export_row).collect()
|
||||
}
|
||||
|
||||
@@ -237,17 +240,22 @@ impl SqlxUserReadRepository {
|
||||
DataLayerError::InvalidInput(format!("invalid user export limit: {}", query.limit))
|
||||
})?);
|
||||
|
||||
let rows = builder.build().fetch_all(&self.pool).await?;
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_user_export_row).collect()
|
||||
}
|
||||
|
||||
pub async fn summarize_export_users(&self) -> Result<UserExportSummary, DataLayerError> {
|
||||
let row = sqlx::query(SUMMARIZE_EXPORT_USERS_SQL)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(UserExportSummary {
|
||||
total: row.try_get::<i64, _>("total")?.max(0) as u64,
|
||||
active: row.try_get::<i64, _>("active")?.max(0) as u64,
|
||||
total: row.try_get::<i64, _>("total").map_postgres_err()?.max(0) as u64,
|
||||
active: row.try_get::<i64, _>("active").map_postgres_err()?.max(0) as u64,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -258,7 +266,8 @@ impl SqlxUserReadRepository {
|
||||
let row = sqlx::query(FIND_EXPORT_USER_BY_ID_SQL)
|
||||
.bind(user_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_user_export_row).transpose()
|
||||
}
|
||||
|
||||
@@ -273,7 +282,8 @@ impl SqlxUserReadRepository {
|
||||
let rows = sqlx::query(LIST_USER_AUTH_BY_IDS_SQL)
|
||||
.bind(user_ids)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_user_auth_row).collect()
|
||||
}
|
||||
|
||||
@@ -284,7 +294,8 @@ impl SqlxUserReadRepository {
|
||||
let row = sqlx::query(FIND_USER_AUTH_BY_ID_SQL)
|
||||
.bind(user_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_user_auth_row).transpose()
|
||||
}
|
||||
|
||||
@@ -295,56 +306,58 @@ impl SqlxUserReadRepository {
|
||||
let row = sqlx::query(FIND_USER_AUTH_BY_IDENTIFIER_SQL)
|
||||
.bind(identifier)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_user_auth_row).transpose()
|
||||
}
|
||||
}
|
||||
|
||||
fn map_user_row(row: &sqlx::postgres::PgRow) -> Result<StoredUserSummary, DataLayerError> {
|
||||
StoredUserSummary::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("username")?,
|
||||
row.try_get("email")?,
|
||||
row.try_get("role")?,
|
||||
row.try_get("is_active")?,
|
||||
row.try_get("is_deleted")?,
|
||||
row.try_get("id").map_postgres_err()?,
|
||||
row.try_get("username").map_postgres_err()?,
|
||||
row.try_get("email").map_postgres_err()?,
|
||||
row.try_get("role").map_postgres_err()?,
|
||||
row.try_get("is_active").map_postgres_err()?,
|
||||
row.try_get("is_deleted").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
fn map_user_export_row(row: &sqlx::postgres::PgRow) -> Result<StoredUserExportRow, DataLayerError> {
|
||||
StoredUserExportRow::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("email")?,
|
||||
row.try_get("email_verified")?,
|
||||
row.try_get("username")?,
|
||||
row.try_get("password_hash")?,
|
||||
row.try_get("role")?,
|
||||
row.try_get("auth_source")?,
|
||||
row.try_get("allowed_providers")?,
|
||||
row.try_get("allowed_api_formats")?,
|
||||
row.try_get("allowed_models")?,
|
||||
row.try_get("rate_limit")?,
|
||||
row.try_get("model_capability_settings")?,
|
||||
row.try_get("is_active")?,
|
||||
row.try_get("id").map_postgres_err()?,
|
||||
row.try_get("email").map_postgres_err()?,
|
||||
row.try_get("email_verified").map_postgres_err()?,
|
||||
row.try_get("username").map_postgres_err()?,
|
||||
row.try_get("password_hash").map_postgres_err()?,
|
||||
row.try_get("role").map_postgres_err()?,
|
||||
row.try_get("auth_source").map_postgres_err()?,
|
||||
row.try_get("allowed_providers").map_postgres_err()?,
|
||||
row.try_get("allowed_api_formats").map_postgres_err()?,
|
||||
row.try_get("allowed_models").map_postgres_err()?,
|
||||
row.try_get("rate_limit").map_postgres_err()?,
|
||||
row.try_get("model_capability_settings")
|
||||
.map_postgres_err()?,
|
||||
row.try_get("is_active").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
fn map_user_auth_row(row: &sqlx::postgres::PgRow) -> Result<StoredUserAuthRecord, DataLayerError> {
|
||||
StoredUserAuthRecord::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("email")?,
|
||||
row.try_get("email_verified")?,
|
||||
row.try_get("username")?,
|
||||
row.try_get("password_hash")?,
|
||||
row.try_get("role")?,
|
||||
row.try_get("auth_source")?,
|
||||
row.try_get("allowed_providers")?,
|
||||
row.try_get("allowed_api_formats")?,
|
||||
row.try_get("allowed_models")?,
|
||||
row.try_get("is_active")?,
|
||||
row.try_get("is_deleted")?,
|
||||
row.try_get("created_at")?,
|
||||
row.try_get("last_login_at")?,
|
||||
row.try_get("id").map_postgres_err()?,
|
||||
row.try_get("email").map_postgres_err()?,
|
||||
row.try_get("email_verified").map_postgres_err()?,
|
||||
row.try_get("username").map_postgres_err()?,
|
||||
row.try_get("password_hash").map_postgres_err()?,
|
||||
row.try_get("role").map_postgres_err()?,
|
||||
row.try_get("auth_source").map_postgres_err()?,
|
||||
row.try_get("allowed_providers").map_postgres_err()?,
|
||||
row.try_get("allowed_api_formats").map_postgres_err()?,
|
||||
row.try_get("allowed_models").map_postgres_err()?,
|
||||
row.try_get("is_active").map_postgres_err()?,
|
||||
row.try_get("is_deleted").map_postgres_err()?,
|
||||
row.try_get("created_at").map_postgres_err()?,
|
||||
row.try_get("last_login_at").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@ use std::sync::RwLock;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::types::{
|
||||
use super::{
|
||||
StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount,
|
||||
VideoTaskQueryFilter, VideoTaskReadRepository, VideoTaskStatus, VideoTaskStatusCount,
|
||||
VideoTaskWriteRepository,
|
||||
@@ -140,9 +140,9 @@ impl VideoTaskReadRepository for InMemoryVideoTaskRepository {
|
||||
.filter(|task| {
|
||||
matches!(
|
||||
task.status,
|
||||
super::types::VideoTaskStatus::Submitted
|
||||
| super::types::VideoTaskStatus::Queued
|
||||
| super::types::VideoTaskStatus::Processing
|
||||
super::VideoTaskStatus::Submitted
|
||||
| super::VideoTaskStatus::Queued
|
||||
| super::VideoTaskStatus::Processing
|
||||
) && task.poll_count < task.max_poll_count
|
||||
&& task
|
||||
.next_poll_at_unix_secs
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
mod memory;
|
||||
mod sql;
|
||||
mod types;
|
||||
|
||||
pub use memory::InMemoryVideoTaskRepository;
|
||||
pub use sql::{SqlxVideoTaskReadRepository, SqlxVideoTaskRepository};
|
||||
pub use types::{
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use aether_data_contracts::repository::video_tasks::{
|
||||
StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount,
|
||||
VideoTaskQueryFilter, VideoTaskReadRepository, VideoTaskRepository, VideoTaskStatus,
|
||||
VideoTaskStatusCount, VideoTaskWriteRepository,
|
||||
};
|
||||
pub use memory::InMemoryVideoTaskRepository;
|
||||
pub use sql::{SqlxVideoTaskReadRepository, SqlxVideoTaskRepository};
|
||||
|
||||
@@ -2,6 +2,7 @@ use sqlx::{PgPool, Postgres, QueryBuilder, Row};
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use crate::error::SqlxResultExt;
|
||||
use crate::repository::video_tasks::{
|
||||
StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount,
|
||||
VideoTaskQueryFilter, VideoTaskReadRepository, VideoTaskStatus, VideoTaskStatusCount,
|
||||
@@ -302,7 +303,8 @@ impl SqlxVideoTaskRepository {
|
||||
let row = sqlx::query(&sql)
|
||||
.bind(id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_video_task_row).transpose()
|
||||
}
|
||||
|
||||
@@ -314,7 +316,8 @@ impl SqlxVideoTaskRepository {
|
||||
let row = sqlx::query(&sql)
|
||||
.bind(short_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_video_task_row).transpose()
|
||||
}
|
||||
|
||||
@@ -328,7 +331,8 @@ impl SqlxVideoTaskRepository {
|
||||
.bind(user_id)
|
||||
.bind(external_task_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_video_task_row).transpose()
|
||||
}
|
||||
|
||||
@@ -345,7 +349,8 @@ impl SqlxVideoTaskRepository {
|
||||
DataLayerError::UnexpectedValue(format!("invalid active task limit: {limit}"))
|
||||
})?)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
rows.iter().map(map_video_task_row).collect()
|
||||
}
|
||||
@@ -368,7 +373,8 @@ impl SqlxVideoTaskRepository {
|
||||
DataLayerError::UnexpectedValue(format!("invalid due task limit: {limit}"))
|
||||
})?)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
rows.iter().map(map_video_task_row).collect()
|
||||
}
|
||||
@@ -398,7 +404,11 @@ impl SqlxVideoTaskRepository {
|
||||
builder.push("\nLIMIT ");
|
||||
builder.push_bind(limit);
|
||||
|
||||
let rows = builder.build().fetch_all(&self.pool).await?;
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_video_task_row).collect()
|
||||
}
|
||||
|
||||
@@ -407,8 +417,12 @@ impl SqlxVideoTaskRepository {
|
||||
QueryBuilder::<Postgres>::new("SELECT COUNT(id) AS total FROM video_tasks");
|
||||
push_video_task_filter(&mut builder, filter, None);
|
||||
|
||||
let row = builder.build().fetch_one(&self.pool).await?;
|
||||
let total = row.try_get::<i64, _>("total")?;
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let total = row.try_get::<i64, _>("total").map_postgres_err()?;
|
||||
u64::try_from(total)
|
||||
.map_err(|_| DataLayerError::UnexpectedValue(format!("invalid count result: {total}")))
|
||||
}
|
||||
@@ -422,12 +436,19 @@ impl SqlxVideoTaskRepository {
|
||||
push_video_task_filter(&mut builder, filter, None);
|
||||
builder.push("\nGROUP BY status\nORDER BY status ASC");
|
||||
|
||||
let rows = builder.build().fetch_all(&self.pool).await?;
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.into_iter()
|
||||
.map(|row| {
|
||||
let status =
|
||||
VideoTaskStatus::from_database(row.try_get::<String, _>("status")?.as_str())?;
|
||||
let total = row.try_get::<i64, _>("total")?;
|
||||
let status = VideoTaskStatus::from_database(
|
||||
row.try_get::<String, _>("status")
|
||||
.map_postgres_err()?
|
||||
.as_str(),
|
||||
)?;
|
||||
let total = row.try_get::<i64, _>("total").map_postgres_err()?;
|
||||
Ok(VideoTaskStatusCount {
|
||||
status,
|
||||
count: u64::try_from(total).map_err(|_| {
|
||||
@@ -452,8 +473,12 @@ impl SqlxVideoTaskRepository {
|
||||
push_sql_clause(&mut builder, &mut has_where, "user_id IS NOT NULL");
|
||||
push_sql_clause(&mut builder, &mut has_where, "user_id <> ''");
|
||||
|
||||
let row = builder.build().fetch_one(&self.pool).await?;
|
||||
let total = row.try_get::<i64, _>("total")?;
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let total = row.try_get::<i64, _>("total").map_postgres_err()?;
|
||||
u64::try_from(total)
|
||||
.map_err(|_| DataLayerError::UnexpectedValue(format!("invalid count result: {total}")))
|
||||
}
|
||||
@@ -479,11 +504,15 @@ impl SqlxVideoTaskRepository {
|
||||
builder.push("\nGROUP BY model\nORDER BY total DESC, model ASC\nLIMIT ");
|
||||
builder.push_bind(limit);
|
||||
|
||||
let rows = builder.build().fetch_all(&self.pool).await?;
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.into_iter()
|
||||
.map(|row| {
|
||||
let model = row.try_get::<String, _>("model")?;
|
||||
let total = row.try_get::<i64, _>("total")?;
|
||||
let model = row.try_get::<String, _>("model").map_postgres_err()?;
|
||||
let total = row.try_get::<i64, _>("total").map_postgres_err()?;
|
||||
Ok(VideoTaskModelCount {
|
||||
model,
|
||||
count: u64::try_from(total).map_err(|_| {
|
||||
@@ -505,8 +534,12 @@ impl SqlxVideoTaskRepository {
|
||||
QueryBuilder::<Postgres>::new("SELECT COUNT(id) AS total FROM video_tasks");
|
||||
push_video_task_filter(&mut builder, filter, Some(created_since_unix_secs));
|
||||
|
||||
let row = builder.build().fetch_one(&self.pool).await?;
|
||||
let total = row.try_get::<i64, _>("total")?;
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let total = row.try_get::<i64, _>("total").map_postgres_err()?;
|
||||
u64::try_from(total)
|
||||
.map_err(|_| DataLayerError::UnexpectedValue(format!("invalid count result: {total}")))
|
||||
}
|
||||
@@ -571,7 +604,8 @@ impl SqlxVideoTaskRepository {
|
||||
.bind(task.completed_at_unix_secs.map(|value| value as f64))
|
||||
.bind(task.updated_at_unix_secs as f64)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
map_video_task_row(&row)
|
||||
}
|
||||
@@ -640,7 +674,8 @@ impl SqlxVideoTaskRepository {
|
||||
.bind(task.updated_at_unix_secs as f64)
|
||||
.bind(vec!["pending", "submitted", "queued", "processing"])
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
row.as_ref().map(map_video_task_row).transpose()
|
||||
}
|
||||
@@ -665,7 +700,8 @@ impl SqlxVideoTaskRepository {
|
||||
.bind(claim_until_unix_secs as f64)
|
||||
.bind(now_unix_secs as f64)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
let mut tasks = rows
|
||||
.iter()
|
||||
@@ -846,45 +882,49 @@ fn map_status_for_database(status: VideoTaskStatus) -> &'static str {
|
||||
}
|
||||
|
||||
fn map_video_task_row(row: &sqlx::postgres::PgRow) -> Result<StoredVideoTask, DataLayerError> {
|
||||
let status = VideoTaskStatus::from_database(row.try_get::<String, _>("status")?.as_str())?;
|
||||
let status = VideoTaskStatus::from_database(
|
||||
row.try_get::<String, _>("status")
|
||||
.map_postgres_err()?
|
||||
.as_str(),
|
||||
)?;
|
||||
StoredVideoTask::new(
|
||||
row.try_get("id")?,
|
||||
row.try_get("short_id")?,
|
||||
row.try_get("request_id")?,
|
||||
row.try_get("user_id")?,
|
||||
row.try_get("api_key_id")?,
|
||||
row.try_get("username")?,
|
||||
row.try_get("api_key_name")?,
|
||||
row.try_get("external_task_id")?,
|
||||
row.try_get("provider_id")?,
|
||||
row.try_get("endpoint_id")?,
|
||||
row.try_get("key_id")?,
|
||||
row.try_get("client_api_format")?,
|
||||
row.try_get("provider_api_format")?,
|
||||
row.try_get("format_converted")?,
|
||||
row.try_get("model")?,
|
||||
row.try_get("prompt")?,
|
||||
row.try_get("original_request_body")?,
|
||||
row.try_get("duration_seconds")?,
|
||||
row.try_get("resolution")?,
|
||||
row.try_get("aspect_ratio")?,
|
||||
row.try_get("size")?,
|
||||
row.try_get("id").map_postgres_err()?,
|
||||
row.try_get("short_id").map_postgres_err()?,
|
||||
row.try_get("request_id").map_postgres_err()?,
|
||||
row.try_get("user_id").map_postgres_err()?,
|
||||
row.try_get("api_key_id").map_postgres_err()?,
|
||||
row.try_get("username").map_postgres_err()?,
|
||||
row.try_get("api_key_name").map_postgres_err()?,
|
||||
row.try_get("external_task_id").map_postgres_err()?,
|
||||
row.try_get("provider_id").map_postgres_err()?,
|
||||
row.try_get("endpoint_id").map_postgres_err()?,
|
||||
row.try_get("key_id").map_postgres_err()?,
|
||||
row.try_get("client_api_format").map_postgres_err()?,
|
||||
row.try_get("provider_api_format").map_postgres_err()?,
|
||||
row.try_get("format_converted").map_postgres_err()?,
|
||||
row.try_get("model").map_postgres_err()?,
|
||||
row.try_get("prompt").map_postgres_err()?,
|
||||
row.try_get("original_request_body").map_postgres_err()?,
|
||||
row.try_get("duration_seconds").map_postgres_err()?,
|
||||
row.try_get("resolution").map_postgres_err()?,
|
||||
row.try_get("aspect_ratio").map_postgres_err()?,
|
||||
row.try_get("size").map_postgres_err()?,
|
||||
status,
|
||||
row.try_get("progress_percent")?,
|
||||
row.try_get("progress_message")?,
|
||||
row.try_get("retry_count")?,
|
||||
row.try_get("poll_interval_seconds")?,
|
||||
row.try_get("next_poll_at_unix_secs")?,
|
||||
row.try_get("poll_count")?,
|
||||
row.try_get("max_poll_count")?,
|
||||
row.try_get("created_at_unix_secs")?,
|
||||
row.try_get("submitted_at_unix_secs")?,
|
||||
row.try_get("completed_at_unix_secs")?,
|
||||
row.try_get("updated_at_unix_secs")?,
|
||||
row.try_get("error_code")?,
|
||||
row.try_get("error_message")?,
|
||||
row.try_get("video_url")?,
|
||||
row.try_get("request_metadata")?,
|
||||
row.try_get("progress_percent").map_postgres_err()?,
|
||||
row.try_get("progress_message").map_postgres_err()?,
|
||||
row.try_get("retry_count").map_postgres_err()?,
|
||||
row.try_get("poll_interval_seconds").map_postgres_err()?,
|
||||
row.try_get("next_poll_at_unix_secs").map_postgres_err()?,
|
||||
row.try_get("poll_count").map_postgres_err()?,
|
||||
row.try_get("max_poll_count").map_postgres_err()?,
|
||||
row.try_get("created_at_unix_secs").map_postgres_err()?,
|
||||
row.try_get("submitted_at_unix_secs").map_postgres_err()?,
|
||||
row.try_get("completed_at_unix_secs").map_postgres_err()?,
|
||||
row.try_get("updated_at_unix_secs").map_postgres_err()?,
|
||||
row.try_get("error_code").map_postgres_err()?,
|
||||
row.try_get("error_message").map_postgres_err()?,
|
||||
row.try_get("video_url").map_postgres_err()?,
|
||||
row.try_get("request_metadata").map_postgres_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -1,611 +0,0 @@
|
||||
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_secs: 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_secs: 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_secs = u64::try_from(created_at_unix_secs).map_err(|_| {
|
||||
crate::DataLayerError::UnexpectedValue(format!(
|
||||
"invalid created_at_unix_secs: {created_at_unix_secs}"
|
||||
))
|
||||
})?;
|
||||
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_secs,
|
||||
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_secs: 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_secs: self.created_at_unix_secs,
|
||||
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_secs: task.created_at_unix_secs,
|
||||
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 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());
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user