mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-10 11:19:50 +08:00
feat: 扩展 Rust gateway 全功能模块,新增 billing/crypto/wallet crate 及完整数据层
- 新增 aether-billing、aether-crypto、aether-wallet 独立 crate - aether-data 扩展 repository 层:announcements、auth_modules、billing、 candidate_selection、gemini_file_mappings、global_models、management_tokens、 oauth_providers、proxy_nodes、quota、users、wallet 等模块 - aether-gateway 新增 api/auth/billing/control/middleware/scheduler/usage/ video_tasks/hooks/maintenance/model_fetch/provider_transport 等功能模块 - 重构 executor decision 和 gateway state 为模块目录结构 - 新增 gateway router、frontdoor 路由层及对应测试 - Python 侧 API 路由重构,新增 compat/support 模块 - 前端 Logo 组件更新及 Provider 管理页面调整
This commit is contained in:
@@ -0,0 +1,196 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::sync::RwLock;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::types::{
|
||||
EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderWriteRepository,
|
||||
StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
pub struct InMemoryOAuthProviderRepository {
|
||||
items: RwLock<BTreeMap<String, StoredOAuthProviderConfig>>,
|
||||
}
|
||||
|
||||
impl InMemoryOAuthProviderRepository {
|
||||
pub fn seed<I>(items: I) -> Self
|
||||
where
|
||||
I: IntoIterator<Item = StoredOAuthProviderConfig>,
|
||||
{
|
||||
let items = items
|
||||
.into_iter()
|
||||
.map(|item| (item.provider_type.clone(), item))
|
||||
.collect();
|
||||
Self {
|
||||
items: RwLock::new(items),
|
||||
}
|
||||
}
|
||||
|
||||
fn now_unix_secs() -> Option<u64> {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.ok()
|
||||
.map(|duration| duration.as_secs())
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl OAuthProviderReadRepository for InMemoryOAuthProviderRepository {
|
||||
async fn list_oauth_provider_configs(
|
||||
&self,
|
||||
) -> Result<Vec<StoredOAuthProviderConfig>, DataLayerError> {
|
||||
let items = self.items.read().expect("oauth provider repository lock");
|
||||
Ok(items.values().cloned().collect())
|
||||
}
|
||||
|
||||
async fn get_oauth_provider_config(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
) -> Result<Option<StoredOAuthProviderConfig>, DataLayerError> {
|
||||
let items = self.items.read().expect("oauth provider repository lock");
|
||||
Ok(items.get(provider_type).cloned())
|
||||
}
|
||||
|
||||
async fn count_locked_users_if_provider_disabled(
|
||||
&self,
|
||||
_provider_type: &str,
|
||||
_ldap_exclusive: bool,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
Ok(0)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl OAuthProviderWriteRepository for InMemoryOAuthProviderRepository {
|
||||
async fn upsert_oauth_provider_config(
|
||||
&self,
|
||||
record: &UpsertOAuthProviderConfigRecord,
|
||||
) -> Result<StoredOAuthProviderConfig, DataLayerError> {
|
||||
record.validate()?;
|
||||
|
||||
let mut items = self.items.write().expect("oauth provider repository lock");
|
||||
let now = Self::now_unix_secs();
|
||||
let existing = items.get(&record.provider_type).cloned();
|
||||
let created_at = existing
|
||||
.as_ref()
|
||||
.and_then(|item| item.created_at_unix_secs)
|
||||
.or(now);
|
||||
let client_secret_encrypted = match (&record.client_secret_encrypted, existing.as_ref()) {
|
||||
(EncryptedSecretUpdate::Preserve, Some(item)) => item.client_secret_encrypted.clone(),
|
||||
(EncryptedSecretUpdate::Preserve, None) => None,
|
||||
(EncryptedSecretUpdate::Clear, _) => None,
|
||||
(EncryptedSecretUpdate::Set(value), _) => Some(value.clone()),
|
||||
};
|
||||
|
||||
let item = StoredOAuthProviderConfig::new(
|
||||
record.provider_type.clone(),
|
||||
record.display_name.clone(),
|
||||
record.client_id.clone(),
|
||||
record.redirect_uri.clone(),
|
||||
record.frontend_callback_url.clone(),
|
||||
)?
|
||||
.with_config_fields(
|
||||
client_secret_encrypted,
|
||||
record.authorization_url_override.clone(),
|
||||
record.token_url_override.clone(),
|
||||
record.userinfo_url_override.clone(),
|
||||
record.scopes.clone(),
|
||||
record.attribute_mapping.clone(),
|
||||
record.extra_config.clone(),
|
||||
record.is_enabled,
|
||||
)
|
||||
.with_timestamps(created_at, now);
|
||||
|
||||
items.insert(record.provider_type.clone(), item.clone());
|
||||
Ok(item)
|
||||
}
|
||||
|
||||
async fn delete_oauth_provider_config(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut items = self.items.write().expect("oauth provider repository lock");
|
||||
Ok(items.remove(provider_type).is_some())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::InMemoryOAuthProviderRepository;
|
||||
use crate::repository::oauth_providers::{
|
||||
EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderWriteRepository,
|
||||
StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord,
|
||||
};
|
||||
|
||||
fn sample_provider(provider_type: &str) -> StoredOAuthProviderConfig {
|
||||
StoredOAuthProviderConfig::new(
|
||||
provider_type.to_string(),
|
||||
format!("{provider_type} display"),
|
||||
format!("{provider_type}-client"),
|
||||
format!("https://{provider_type}.example.com/redirect"),
|
||||
"https://frontend.example.com/auth/callback".to_string(),
|
||||
)
|
||||
.expect("provider should build")
|
||||
}
|
||||
|
||||
fn sample_upsert(provider_type: &str) -> UpsertOAuthProviderConfigRecord {
|
||||
UpsertOAuthProviderConfigRecord {
|
||||
provider_type: provider_type.to_string(),
|
||||
display_name: format!("{provider_type} display"),
|
||||
client_id: format!("{provider_type}-client"),
|
||||
client_secret_encrypted: EncryptedSecretUpdate::Preserve,
|
||||
authorization_url_override: Some(format!("https://{provider_type}.example.com/auth")),
|
||||
token_url_override: Some(format!("https://{provider_type}.example.com/token")),
|
||||
userinfo_url_override: None,
|
||||
scopes: Some(vec!["openid".to_string(), "profile".to_string()]),
|
||||
redirect_uri: format!("https://{provider_type}.example.com/redirect"),
|
||||
frontend_callback_url: "https://frontend.example.com/auth/callback".to_string(),
|
||||
attribute_mapping: Some(serde_json::json!({"email": "email"})),
|
||||
extra_config: Some(serde_json::json!({"team": true})),
|
||||
is_enabled: true,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn reads_and_mutates_oauth_provider_configs() {
|
||||
let repository = InMemoryOAuthProviderRepository::seed(vec![
|
||||
sample_provider("linuxdo"),
|
||||
sample_provider("github"),
|
||||
]);
|
||||
|
||||
let listed = repository
|
||||
.list_oauth_provider_configs()
|
||||
.await
|
||||
.expect("list should succeed");
|
||||
assert_eq!(listed.len(), 2);
|
||||
assert_eq!(listed[0].provider_type, "github");
|
||||
assert_eq!(listed[1].provider_type, "linuxdo");
|
||||
|
||||
let created = repository
|
||||
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
|
||||
client_secret_encrypted: EncryptedSecretUpdate::Set("secret-1".to_string()),
|
||||
..sample_upsert("google")
|
||||
})
|
||||
.await
|
||||
.expect("create should succeed");
|
||||
assert_eq!(created.client_secret_encrypted.as_deref(), Some("secret-1"));
|
||||
|
||||
let updated = repository
|
||||
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
|
||||
client_secret_encrypted: EncryptedSecretUpdate::Clear,
|
||||
..sample_upsert("google")
|
||||
})
|
||||
.await
|
||||
.expect("update should succeed");
|
||||
assert!(updated.client_secret_encrypted.is_none());
|
||||
|
||||
let deleted = repository
|
||||
.delete_oauth_provider_config("google")
|
||||
.await
|
||||
.expect("delete should succeed");
|
||||
assert!(deleted);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,10 @@
|
||||
mod memory;
|
||||
mod sql;
|
||||
mod types;
|
||||
|
||||
pub use memory::InMemoryOAuthProviderRepository;
|
||||
pub use sql::SqlxOAuthProviderRepository;
|
||||
pub use types::{
|
||||
EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderRepository,
|
||||
OAuthProviderWriteRepository, StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord,
|
||||
};
|
||||
@@ -0,0 +1,341 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{postgres::PgRow, PgPool, Row};
|
||||
|
||||
use super::types::{
|
||||
OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig,
|
||||
UpsertOAuthProviderConfigRecord,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
|
||||
const LIST_OAUTH_PROVIDER_CONFIGS_SQL: &str = r#"
|
||||
SELECT
|
||||
provider_type,
|
||||
display_name,
|
||||
client_id,
|
||||
client_secret_encrypted,
|
||||
authorization_url_override,
|
||||
token_url_override,
|
||||
userinfo_url_override,
|
||||
scopes,
|
||||
redirect_uri,
|
||||
frontend_callback_url,
|
||||
attribute_mapping,
|
||||
extra_config,
|
||||
is_enabled,
|
||||
EXTRACT(EPOCH FROM created_at)::bigint AS created_at_unix_secs,
|
||||
EXTRACT(EPOCH FROM updated_at)::bigint AS updated_at_unix_secs
|
||||
FROM oauth_providers
|
||||
ORDER BY provider_type ASC
|
||||
"#;
|
||||
|
||||
const GET_OAUTH_PROVIDER_CONFIG_SQL: &str = r#"
|
||||
SELECT
|
||||
provider_type,
|
||||
display_name,
|
||||
client_id,
|
||||
client_secret_encrypted,
|
||||
authorization_url_override,
|
||||
token_url_override,
|
||||
userinfo_url_override,
|
||||
scopes,
|
||||
redirect_uri,
|
||||
frontend_callback_url,
|
||||
attribute_mapping,
|
||||
extra_config,
|
||||
is_enabled,
|
||||
EXTRACT(EPOCH FROM created_at)::bigint AS created_at_unix_secs,
|
||||
EXTRACT(EPOCH FROM updated_at)::bigint AS updated_at_unix_secs
|
||||
FROM oauth_providers
|
||||
WHERE provider_type = $1
|
||||
LIMIT 1
|
||||
"#;
|
||||
|
||||
const COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL: &str = r#"
|
||||
WITH affected_users AS (
|
||||
SELECT DISTINCT
|
||||
users.id,
|
||||
users.auth_source,
|
||||
users.role,
|
||||
(
|
||||
SELECT COUNT(*)
|
||||
FROM user_oauth_links other_links
|
||||
JOIN oauth_providers other_provider
|
||||
ON other_links.provider_type = other_provider.provider_type
|
||||
WHERE other_links.user_id = users.id
|
||||
AND other_links.provider_type <> $1
|
||||
AND other_provider.is_enabled IS TRUE
|
||||
) AS other_enabled_count
|
||||
FROM users
|
||||
JOIN user_oauth_links
|
||||
ON users.id = user_oauth_links.user_id
|
||||
WHERE users.is_active IS TRUE
|
||||
AND users.is_deleted IS FALSE
|
||||
AND user_oauth_links.provider_type = $1
|
||||
)
|
||||
SELECT COUNT(*)::bigint AS locked_count
|
||||
FROM affected_users
|
||||
WHERE (
|
||||
auth_source = 'oauth'
|
||||
AND other_enabled_count = 0
|
||||
) OR (
|
||||
$2::boolean IS TRUE
|
||||
AND auth_source = 'local'
|
||||
AND role <> 'admin'
|
||||
AND other_enabled_count = 0
|
||||
)
|
||||
"#;
|
||||
|
||||
const UPSERT_OAUTH_PROVIDER_CONFIG_SQL: &str = r#"
|
||||
INSERT INTO oauth_providers (
|
||||
provider_type,
|
||||
display_name,
|
||||
client_id,
|
||||
client_secret_encrypted,
|
||||
authorization_url_override,
|
||||
token_url_override,
|
||||
userinfo_url_override,
|
||||
scopes,
|
||||
redirect_uri,
|
||||
frontend_callback_url,
|
||||
attribute_mapping,
|
||||
extra_config,
|
||||
is_enabled,
|
||||
created_at,
|
||||
updated_at
|
||||
)
|
||||
VALUES (
|
||||
$1,
|
||||
$2,
|
||||
$3,
|
||||
CASE $4
|
||||
WHEN 'set' THEN $5
|
||||
WHEN 'clear' THEN NULL
|
||||
ELSE NULL
|
||||
END,
|
||||
$6,
|
||||
$7,
|
||||
$8,
|
||||
$9,
|
||||
$10,
|
||||
$11,
|
||||
$12,
|
||||
$13,
|
||||
$14,
|
||||
NOW(),
|
||||
NOW()
|
||||
)
|
||||
ON CONFLICT (provider_type) DO UPDATE
|
||||
SET display_name = EXCLUDED.display_name,
|
||||
client_id = EXCLUDED.client_id,
|
||||
client_secret_encrypted = CASE $4
|
||||
WHEN 'set' THEN $5
|
||||
WHEN 'clear' THEN NULL
|
||||
ELSE oauth_providers.client_secret_encrypted
|
||||
END,
|
||||
authorization_url_override = EXCLUDED.authorization_url_override,
|
||||
token_url_override = EXCLUDED.token_url_override,
|
||||
userinfo_url_override = EXCLUDED.userinfo_url_override,
|
||||
scopes = EXCLUDED.scopes,
|
||||
redirect_uri = EXCLUDED.redirect_uri,
|
||||
frontend_callback_url = EXCLUDED.frontend_callback_url,
|
||||
attribute_mapping = EXCLUDED.attribute_mapping,
|
||||
extra_config = EXCLUDED.extra_config,
|
||||
is_enabled = EXCLUDED.is_enabled,
|
||||
updated_at = NOW()
|
||||
RETURNING
|
||||
provider_type,
|
||||
display_name,
|
||||
client_id,
|
||||
client_secret_encrypted,
|
||||
authorization_url_override,
|
||||
token_url_override,
|
||||
userinfo_url_override,
|
||||
scopes,
|
||||
redirect_uri,
|
||||
frontend_callback_url,
|
||||
attribute_mapping,
|
||||
extra_config,
|
||||
is_enabled,
|
||||
EXTRACT(EPOCH FROM created_at)::bigint AS created_at_unix_secs,
|
||||
EXTRACT(EPOCH FROM updated_at)::bigint AS updated_at_unix_secs
|
||||
"#;
|
||||
|
||||
const DELETE_OAUTH_PROVIDER_CONFIG_SQL: &str = r#"
|
||||
DELETE FROM oauth_providers
|
||||
WHERE provider_type = $1
|
||||
"#;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SqlxOAuthProviderRepository {
|
||||
pool: PgPool,
|
||||
}
|
||||
|
||||
impl SqlxOAuthProviderRepository {
|
||||
pub fn new(pool: PgPool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl OAuthProviderReadRepository for SqlxOAuthProviderRepository {
|
||||
async fn list_oauth_provider_configs(
|
||||
&self,
|
||||
) -> Result<Vec<StoredOAuthProviderConfig>, DataLayerError> {
|
||||
let rows = sqlx::query(LIST_OAUTH_PROVIDER_CONFIGS_SQL)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
rows.iter().map(map_oauth_provider_row).collect()
|
||||
}
|
||||
|
||||
async fn get_oauth_provider_config(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
) -> Result<Option<StoredOAuthProviderConfig>, DataLayerError> {
|
||||
let row = sqlx::query(GET_OAUTH_PROVIDER_CONFIG_SQL)
|
||||
.bind(provider_type)
|
||||
.fetch_optional(&self.pool)
|
||||
.await?;
|
||||
row.as_ref().map(map_oauth_provider_row).transpose()
|
||||
}
|
||||
|
||||
async fn count_locked_users_if_provider_disabled(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
ldap_exclusive: bool,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let locked_count: i64 = sqlx::query_scalar(COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL)
|
||||
.bind(provider_type)
|
||||
.bind(ldap_exclusive)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
usize::try_from(locked_count).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(
|
||||
"oauth_providers.locked_user_count is negative".to_string(),
|
||||
)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl OAuthProviderWriteRepository for SqlxOAuthProviderRepository {
|
||||
async fn upsert_oauth_provider_config(
|
||||
&self,
|
||||
record: &UpsertOAuthProviderConfigRecord,
|
||||
) -> Result<StoredOAuthProviderConfig, DataLayerError> {
|
||||
record.validate()?;
|
||||
let row = sqlx::query(UPSERT_OAUTH_PROVIDER_CONFIG_SQL)
|
||||
.bind(&record.provider_type)
|
||||
.bind(&record.display_name)
|
||||
.bind(&record.client_id)
|
||||
.bind(record.client_secret_encrypted.mode_name())
|
||||
.bind(record.client_secret_encrypted.value())
|
||||
.bind(record.authorization_url_override.as_deref())
|
||||
.bind(record.token_url_override.as_deref())
|
||||
.bind(record.userinfo_url_override.as_deref())
|
||||
.bind(scopes_to_json(record.scopes.as_ref()))
|
||||
.bind(&record.redirect_uri)
|
||||
.bind(&record.frontend_callback_url)
|
||||
.bind(record.attribute_mapping.as_ref())
|
||||
.bind(record.extra_config.as_ref())
|
||||
.bind(record.is_enabled)
|
||||
.fetch_one(&self.pool)
|
||||
.await?;
|
||||
map_oauth_provider_row(&row)
|
||||
}
|
||||
|
||||
async fn delete_oauth_provider_config(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let result = sqlx::query(DELETE_OAUTH_PROVIDER_CONFIG_SQL)
|
||||
.bind(provider_type)
|
||||
.execute(&self.pool)
|
||||
.await?;
|
||||
Ok(result.rows_affected() > 0)
|
||||
}
|
||||
}
|
||||
|
||||
fn optional_unix_secs(value: Option<i64>) -> Option<u64> {
|
||||
value.and_then(|value| u64::try_from(value).ok())
|
||||
}
|
||||
|
||||
fn scopes_to_json(scopes: Option<&Vec<String>>) -> Option<serde_json::Value> {
|
||||
scopes.map(|items| {
|
||||
serde_json::Value::Array(
|
||||
items
|
||||
.iter()
|
||||
.cloned()
|
||||
.map(serde_json::Value::String)
|
||||
.collect(),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn parse_scopes(value: Option<serde_json::Value>) -> Result<Option<Vec<String>>, DataLayerError> {
|
||||
let Some(value) = value else {
|
||||
return Ok(None);
|
||||
};
|
||||
let serde_json::Value::Array(items) = value else {
|
||||
return Err(DataLayerError::UnexpectedValue(
|
||||
"oauth_providers.scopes is not a JSON array".to_string(),
|
||||
));
|
||||
};
|
||||
let mut scopes = Vec::with_capacity(items.len());
|
||||
for item in items {
|
||||
let serde_json::Value::String(scope) = item else {
|
||||
return Err(DataLayerError::UnexpectedValue(
|
||||
"oauth_providers.scopes contains non-string value".to_string(),
|
||||
));
|
||||
};
|
||||
scopes.push(scope);
|
||||
}
|
||||
Ok(Some(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")?,
|
||||
)?
|
||||
.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")?,
|
||||
)
|
||||
.with_timestamps(
|
||||
optional_unix_secs(row.try_get("created_at_unix_secs")?),
|
||||
optional_unix_secs(row.try_get("updated_at_unix_secs")?),
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::SqlxOAuthProviderRepository;
|
||||
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
|
||||
|
||||
#[tokio::test]
|
||||
async fn repository_constructs_from_lazy_pool() {
|
||||
let factory = PostgresPoolFactory::new(PostgresPoolConfig {
|
||||
database_url: "postgres://localhost/aether".to_string(),
|
||||
min_connections: 1,
|
||||
max_connections: 4,
|
||||
acquire_timeout_ms: 1_000,
|
||||
idle_timeout_ms: 5_000,
|
||||
max_lifetime_ms: 30_000,
|
||||
statement_cache_capacity: 64,
|
||||
require_ssl: false,
|
||||
})
|
||||
.expect("factory should build");
|
||||
|
||||
let pool = factory.connect_lazy().expect("pool should build");
|
||||
let _repository = SqlxOAuthProviderRepository::new(pool);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,230 @@
|
||||
use async_trait::async_trait;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredOAuthProviderConfig {
|
||||
pub provider_type: String,
|
||||
pub display_name: String,
|
||||
pub client_id: String,
|
||||
pub client_secret_encrypted: Option<String>,
|
||||
pub authorization_url_override: Option<String>,
|
||||
pub token_url_override: Option<String>,
|
||||
pub userinfo_url_override: Option<String>,
|
||||
pub scopes: Option<Vec<String>>,
|
||||
pub redirect_uri: String,
|
||||
pub frontend_callback_url: String,
|
||||
pub attribute_mapping: Option<serde_json::Value>,
|
||||
pub extra_config: Option<serde_json::Value>,
|
||||
pub is_enabled: bool,
|
||||
pub created_at_unix_secs: Option<u64>,
|
||||
pub updated_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
impl StoredOAuthProviderConfig {
|
||||
pub fn new(
|
||||
provider_type: String,
|
||||
display_name: String,
|
||||
client_id: String,
|
||||
redirect_uri: String,
|
||||
frontend_callback_url: String,
|
||||
) -> Result<Self, crate::DataLayerError> {
|
||||
if provider_type.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"oauth_providers.provider_type is empty".to_string(),
|
||||
));
|
||||
}
|
||||
if display_name.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"oauth_providers.display_name is empty".to_string(),
|
||||
));
|
||||
}
|
||||
if client_id.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"oauth_providers.client_id is empty".to_string(),
|
||||
));
|
||||
}
|
||||
if redirect_uri.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"oauth_providers.redirect_uri is empty".to_string(),
|
||||
));
|
||||
}
|
||||
if frontend_callback_url.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"oauth_providers.frontend_callback_url is empty".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
provider_type,
|
||||
display_name,
|
||||
client_id,
|
||||
client_secret_encrypted: None,
|
||||
authorization_url_override: None,
|
||||
token_url_override: None,
|
||||
userinfo_url_override: None,
|
||||
scopes: None,
|
||||
redirect_uri,
|
||||
frontend_callback_url,
|
||||
attribute_mapping: None,
|
||||
extra_config: None,
|
||||
is_enabled: false,
|
||||
created_at_unix_secs: None,
|
||||
updated_at_unix_secs: None,
|
||||
})
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn with_config_fields(
|
||||
mut self,
|
||||
client_secret_encrypted: Option<String>,
|
||||
authorization_url_override: Option<String>,
|
||||
token_url_override: Option<String>,
|
||||
userinfo_url_override: Option<String>,
|
||||
scopes: Option<Vec<String>>,
|
||||
attribute_mapping: Option<serde_json::Value>,
|
||||
extra_config: Option<serde_json::Value>,
|
||||
is_enabled: bool,
|
||||
) -> Self {
|
||||
self.client_secret_encrypted = client_secret_encrypted;
|
||||
self.authorization_url_override = authorization_url_override;
|
||||
self.token_url_override = token_url_override;
|
||||
self.userinfo_url_override = userinfo_url_override;
|
||||
self.scopes = scopes;
|
||||
self.attribute_mapping = attribute_mapping;
|
||||
self.extra_config = extra_config;
|
||||
self.is_enabled = is_enabled;
|
||||
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, Eq, serde::Serialize, serde::Deserialize, Default)]
|
||||
pub enum EncryptedSecretUpdate {
|
||||
#[default]
|
||||
Preserve,
|
||||
Clear,
|
||||
Set(String),
|
||||
}
|
||||
|
||||
impl EncryptedSecretUpdate {
|
||||
pub fn mode_name(&self) -> &'static str {
|
||||
match self {
|
||||
Self::Preserve => "preserve",
|
||||
Self::Clear => "clear",
|
||||
Self::Set(_) => "set",
|
||||
}
|
||||
}
|
||||
|
||||
pub fn value(&self) -> Option<&str> {
|
||||
match self {
|
||||
Self::Set(value) => Some(value.as_str()),
|
||||
Self::Preserve | Self::Clear => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct UpsertOAuthProviderConfigRecord {
|
||||
pub provider_type: String,
|
||||
pub display_name: String,
|
||||
pub client_id: String,
|
||||
pub client_secret_encrypted: EncryptedSecretUpdate,
|
||||
pub authorization_url_override: Option<String>,
|
||||
pub token_url_override: Option<String>,
|
||||
pub userinfo_url_override: Option<String>,
|
||||
pub scopes: Option<Vec<String>>,
|
||||
pub redirect_uri: String,
|
||||
pub frontend_callback_url: String,
|
||||
pub attribute_mapping: Option<serde_json::Value>,
|
||||
pub extra_config: Option<serde_json::Value>,
|
||||
pub is_enabled: bool,
|
||||
}
|
||||
|
||||
impl UpsertOAuthProviderConfigRecord {
|
||||
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
|
||||
if self.provider_type.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"provider_type is required".to_string(),
|
||||
));
|
||||
}
|
||||
if self.display_name.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"display_name is required".to_string(),
|
||||
));
|
||||
}
|
||||
if self.client_id.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"client_id is required".to_string(),
|
||||
));
|
||||
}
|
||||
if self.redirect_uri.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"redirect_uri is required".to_string(),
|
||||
));
|
||||
}
|
||||
if self.frontend_callback_url.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"frontend_callback_url is required".to_string(),
|
||||
));
|
||||
}
|
||||
if let Some(scopes) = &self.scopes {
|
||||
for scope in scopes {
|
||||
if scope.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"scopes must not contain empty values".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait OAuthProviderReadRepository: Send + Sync {
|
||||
async fn list_oauth_provider_configs(
|
||||
&self,
|
||||
) -> Result<Vec<StoredOAuthProviderConfig>, crate::DataLayerError>;
|
||||
|
||||
async fn get_oauth_provider_config(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
) -> Result<Option<StoredOAuthProviderConfig>, crate::DataLayerError>;
|
||||
|
||||
async fn count_locked_users_if_provider_disabled(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
ldap_exclusive: bool,
|
||||
) -> Result<usize, crate::DataLayerError>;
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait OAuthProviderWriteRepository: Send + Sync {
|
||||
async fn upsert_oauth_provider_config(
|
||||
&self,
|
||||
record: &UpsertOAuthProviderConfigRecord,
|
||||
) -> Result<StoredOAuthProviderConfig, crate::DataLayerError>;
|
||||
|
||||
async fn delete_oauth_provider_config(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
}
|
||||
|
||||
pub trait OAuthProviderRepository:
|
||||
OAuthProviderReadRepository + OAuthProviderWriteRepository + Send + Sync
|
||||
{
|
||||
}
|
||||
|
||||
impl<T> OAuthProviderRepository for T where
|
||||
T: OAuthProviderReadRepository + OAuthProviderWriteRepository + Send + Sync
|
||||
{
|
||||
}
|
||||
Reference in New Issue
Block a user