Merge remote-tracking branch 'upstream/aether-rust-pioneer' into fix-management-token-oauth-jsonb

# Conflicts:
#	crates/aether-data/src/lifecycle/bootstrap/postgres.rs
#	crates/aether-data/src/lifecycle/migrate/tests.rs
This commit is contained in:
Entropy.Xu
2026-05-10 18:36:53 +08:00
109 changed files with 11022 additions and 948 deletions

View File

@@ -293,7 +293,7 @@ mod tests {
.await
.expect("system config should list")
.len(),
1
2
);
assert!(backend
.delete_system_config_value("feature.local")

View File

@@ -29,6 +29,7 @@ WHERE table_schema = 'public'
'oauth_providers',
'provider_api_keys',
'proxy_nodes',
'user_groups',
'usage_routing_snapshots',
'usage_settlement_snapshots'
)

View File

@@ -23,6 +23,8 @@ pub enum ExportDomain {
AuthModules,
OAuthProviders,
UserOAuthLinks,
UserGroups,
UserGroupMembers,
ProxyNodes,
SystemConfigs,
Wallets,
@@ -43,6 +45,8 @@ impl ExportDomain {
Self::AuthModules => "auth_modules",
Self::OAuthProviders => "oauth_providers",
Self::UserOAuthLinks => "user_oauth_links",
Self::UserGroups => "user_groups",
Self::UserGroupMembers => "user_group_members",
Self::ProxyNodes => "proxy_nodes",
Self::SystemConfigs => "system_configs",
Self::Wallets => "wallets",
@@ -265,6 +269,8 @@ pub fn sqlite_core_export_domains() -> Vec<ExportDomain> {
ExportDomain::AuthModules,
ExportDomain::OAuthProviders,
ExportDomain::UserOAuthLinks,
ExportDomain::UserGroups,
ExportDomain::UserGroupMembers,
ExportDomain::ProxyNodes,
ExportDomain::SystemConfigs,
ExportDomain::Wallets,
@@ -367,19 +373,11 @@ pub async fn export_sqlite_jsonl(
continue;
}
let (table_name, id_column) = sqlite_domain_table(domain)?;
let sql = format!("SELECT * FROM {table_name} ORDER BY {id_column} ASC");
let order_by = export_order_by(domain, id_column);
let sql = format!("SELECT * FROM {table_name} ORDER BY {order_by}");
let rows = sqlx::query(&sql).fetch_all(pool).await.map_sql_err()?;
for row in rows {
let id = row
.try_get::<Option<String>, _>(id_column)
.map_sql_err()?
.ok_or_else(|| {
DataLayerError::UnexpectedValue(format!(
"{} export row has null id column '{}'",
domain.as_str(),
id_column
))
})?;
let id = sqlite_export_row_id(domain, &row, id_column)?;
records.push(DataExportRecord::row(domain, id, sqlite_row_payload(&row)?));
}
}
@@ -456,8 +454,10 @@ pub async fn export_postgres_jsonl(
continue;
}
let (table_name, id_column) = postgres_domain_table(domain)?;
let export_id_sql = postgres_export_id_sql(domain, id_column);
let order_by = export_order_by(domain, id_column);
let sql = format!(
"SELECT {id_column}::text AS export_id, to_jsonb(t) AS payload FROM {table_name} AS t ORDER BY {id_column} ASC"
"SELECT {export_id_sql} AS export_id, to_jsonb(t) AS payload FROM {table_name} AS t ORDER BY {order_by}"
);
let rows = sqlx::query(&sql).fetch_all(pool).await.map_sql_err()?;
for row in rows {
@@ -500,6 +500,7 @@ pub async fn import_postgres_plan(
continue;
}
let (table_name, id_column) = postgres_domain_table(*domain)?;
let conflict_columns = postgres_conflict_columns(*domain, id_column);
let rows = plan.rows(*domain);
if rows.is_empty() {
continue;
@@ -507,7 +508,15 @@ pub async fn import_postgres_plan(
let target_columns =
postgres_import_columns_cached(pool, &mut column_cache, table_name).await?;
for row in rows {
import_postgres_row(pool, table_name, id_column, *domain, row, &target_columns).await?;
import_postgres_row(
pool,
table_name,
&conflict_columns,
*domain,
row,
&target_columns,
)
.await?;
imported = imported.saturating_add(1);
}
}
@@ -543,19 +552,11 @@ pub async fn export_mysql_jsonl(
continue;
}
let (table_name, id_column) = mysql_domain_table(domain)?;
let sql = format!("SELECT * FROM {table_name} ORDER BY {id_column} ASC");
let order_by = export_order_by(domain, id_column);
let sql = format!("SELECT * FROM {table_name} ORDER BY {order_by}");
let rows = sqlx::query(&sql).fetch_all(pool).await.map_sql_err()?;
for row in rows {
let id = row
.try_get::<Option<String>, _>(id_column)
.map_sql_err()?
.ok_or_else(|| {
DataLayerError::UnexpectedValue(format!(
"{} export row has null id column '{}'",
domain.as_str(),
id_column
))
})?;
let id = mysql_export_row_id(domain, &row, id_column)?;
records.push(DataExportRecord::row(domain, id, mysql_row_payload(&row)?));
}
}
@@ -617,6 +618,8 @@ fn sqlite_domain_table(
ExportDomain::AuthModules => Ok(("auth_modules", "id")),
ExportDomain::OAuthProviders => Ok(("oauth_providers", "provider_type")),
ExportDomain::UserOAuthLinks => Ok(("user_oauth_links", "id")),
ExportDomain::UserGroups => Ok(("user_groups", "id")),
ExportDomain::UserGroupMembers => Ok(("user_group_members", "group_id")),
ExportDomain::ProxyNodes => Ok(("proxy_nodes", "id")),
ExportDomain::SystemConfigs => Ok(("system_configs", "id")),
ExportDomain::Wallets => Err(DataLayerError::InvalidInput(
@@ -630,6 +633,43 @@ fn sqlite_domain_table(
}
}
fn export_order_by(domain: ExportDomain, id_column: &str) -> String {
if domain == ExportDomain::UserGroupMembers {
"group_id ASC, user_id ASC".to_string()
} else {
format!("{id_column} ASC")
}
}
fn sqlite_export_row_id(
domain: ExportDomain,
row: &sqlx::sqlite::SqliteRow,
id_column: &str,
) -> Result<String, DataLayerError> {
if domain == ExportDomain::UserGroupMembers {
let group_id = sqlite_required_export_text(row, "group_id", domain)?;
let user_id = sqlite_required_export_text(row, "user_id", domain)?;
return Ok(format!("{group_id}:{user_id}"));
}
sqlite_required_export_text(row, id_column, domain)
}
fn sqlite_required_export_text(
row: &sqlx::sqlite::SqliteRow,
column: &str,
domain: ExportDomain,
) -> Result<String, DataLayerError> {
row.try_get::<Option<String>, _>(column)
.map_sql_err()?
.ok_or_else(|| {
DataLayerError::UnexpectedValue(format!(
"{} export row has null id column '{}'",
domain.as_str(),
column
))
})
}
async fn export_sqlite_billing_records(
pool: &crate::driver::sqlite::SqlitePool,
records: &mut Vec<DataExportRecord>,
@@ -898,6 +938,8 @@ fn postgres_domain_table(
ExportDomain::AuthModules => Ok(("public.auth_modules", "id")),
ExportDomain::OAuthProviders => Ok(("public.oauth_providers", "provider_type")),
ExportDomain::UserOAuthLinks => Ok(("public.user_oauth_links", "id")),
ExportDomain::UserGroups => Ok(("public.user_groups", "id")),
ExportDomain::UserGroupMembers => Ok(("public.user_group_members", "group_id")),
ExportDomain::ProxyNodes => Ok(("public.proxy_nodes", "id")),
ExportDomain::SystemConfigs => Ok(("public.system_configs", "id")),
ExportDomain::Wallets => Err(DataLayerError::InvalidInput(
@@ -912,6 +954,22 @@ fn postgres_domain_table(
}
}
fn postgres_export_id_sql(domain: ExportDomain, id_column: &str) -> String {
if domain == ExportDomain::UserGroupMembers {
"group_id::text || ':' || user_id::text".to_string()
} else {
format!("{id_column}::text")
}
}
fn postgres_conflict_columns(domain: ExportDomain, id_column: &str) -> Vec<&str> {
if domain == ExportDomain::UserGroupMembers {
vec!["group_id", "user_id"]
} else {
vec![id_column]
}
}
async fn postgres_import_columns_cached(
pool: &crate::driver::postgres::PostgresPool,
cache: &mut BTreeMap<String, PostgresImportColumns>,
@@ -1045,7 +1103,7 @@ async fn export_postgres_wallet_records(
async fn import_postgres_row(
pool: &crate::driver::postgres::PostgresPool,
table_name: &str,
id_column: &str,
conflict_columns: &[&str],
domain: ExportDomain,
row: &ExportRow,
target_columns: &PostgresImportColumns,
@@ -1060,23 +1118,22 @@ async fn import_postgres_row(
.join(", ");
let update_sql = columns
.iter()
.filter(|column| **column != id_column)
.filter(|column| !conflict_columns.contains(column))
.map(|column| {
let quoted = postgres_quote_identifier(column)?;
Ok(format!("{quoted} = EXCLUDED.{quoted}"))
})
.collect::<Result<Vec<_>, DataLayerError>>()?
.join(", ");
let conflict_target_sql = conflict_columns
.iter()
.map(|column| postgres_quote_identifier(column))
.collect::<Result<Vec<_>, _>>()?
.join(", ");
let conflict_sql = if update_sql.is_empty() {
format!(
"ON CONFLICT ({}) DO NOTHING",
postgres_quote_identifier(id_column)?
)
format!("ON CONFLICT ({conflict_target_sql}) DO NOTHING")
} else {
format!(
"ON CONFLICT ({}) DO UPDATE SET {update_sql}",
postgres_quote_identifier(id_column)?
)
format!("ON CONFLICT ({conflict_target_sql}) DO UPDATE SET {update_sql}")
};
let sql = format!(
"INSERT INTO {table_name} ({column_sql}) SELECT {column_sql} FROM jsonb_populate_record(NULL::{table_name}, $1::jsonb) {conflict_sql}"
@@ -1276,7 +1333,7 @@ async fn import_postgres_billing_row(
import_postgres_row(
pool,
table_name,
"id",
&["id"],
ExportDomain::Billing,
&ExportRow {
id: row.id.clone(),
@@ -1309,7 +1366,7 @@ async fn import_postgres_wallet_row(
import_postgres_row(
pool,
table_name,
id_column,
&[id_column],
ExportDomain::Wallets,
&ExportRow {
id: row.id.clone(),
@@ -1382,6 +1439,8 @@ fn mysql_domain_table(
ExportDomain::AuthModules => Ok(("auth_modules", "id")),
ExportDomain::OAuthProviders => Ok(("oauth_providers", "provider_type")),
ExportDomain::UserOAuthLinks => Ok(("user_oauth_links", "id")),
ExportDomain::UserGroups => Ok(("user_groups", "id")),
ExportDomain::UserGroupMembers => Ok(("user_group_members", "group_id")),
ExportDomain::ProxyNodes => Ok(("proxy_nodes", "id")),
ExportDomain::SystemConfigs => Ok(("system_configs", "id")),
ExportDomain::Wallets => Err(DataLayerError::InvalidInput(
@@ -1394,6 +1453,35 @@ fn mysql_domain_table(
}
}
fn mysql_export_row_id(
domain: ExportDomain,
row: &sqlx::mysql::MySqlRow,
id_column: &str,
) -> Result<String, DataLayerError> {
if domain == ExportDomain::UserGroupMembers {
let group_id = mysql_required_export_text(row, "group_id", domain)?;
let user_id = mysql_required_export_text(row, "user_id", domain)?;
return Ok(format!("{group_id}:{user_id}"));
}
mysql_required_export_text(row, id_column, domain)
}
fn mysql_required_export_text(
row: &sqlx::mysql::MySqlRow,
column: &str,
domain: ExportDomain,
) -> Result<String, DataLayerError> {
row.try_get::<Option<String>, _>(column)
.map_sql_err()?
.ok_or_else(|| {
DataLayerError::UnexpectedValue(format!(
"{} export row has null id column '{}'",
domain.as_str(),
column
))
})
}
async fn export_mysql_billing_records(
pool: &crate::driver::mysql::MysqlPool,
records: &mut Vec<DataExportRecord>,
@@ -2094,6 +2182,10 @@ not-json"#,
r#"
INSERT INTO users (id, email, username, auth_source, created_at, updated_at)
VALUES ('user-1', 'owner@example.com', 'owner', 'local', 1, 2);
INSERT INTO user_groups (id, name, normalized_name, description, priority, allowed_models, allowed_models_mode, created_at, updated_at)
VALUES ('group-1', 'Export Group', 'export group', 'Exported group', 10, '["gpt-test"]', 'specific', 1, 2);
INSERT INTO user_group_members (group_id, user_id, created_at)
VALUES ('group-1', 'user-1', 1);
INSERT INTO api_keys (id, user_id, key_hash, key_encrypted, name, created_at, updated_at)
VALUES ('api-key-1', 'user-1', 'hash-1', 'ciphertext-1', 'Default', 1, 2);
INSERT INTO providers (id, name, provider_type, created_at, updated_at)
@@ -2136,6 +2228,16 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
import_plan.rows(ExportDomain::Users)[0].payload["email"],
"owner@example.com"
);
assert!(import_plan
.rows(ExportDomain::UserGroups)
.iter()
.any(|row| row.id == "group-1" && row.payload["name"] == "Export Group"));
assert!(import_plan
.rows(ExportDomain::UserGroupMembers)
.iter()
.any(|row| row.id == "group-1:user-1"
&& row.payload["group_id"] == "group-1"
&& row.payload["user_id"] == "user-1"));
assert_eq!(
import_plan.rows(ExportDomain::ApiKeys)[0].payload["key_encrypted"],
"ciphertext-1"
@@ -2166,7 +2268,7 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
let imported = import_sqlite_jsonl(&target_pool, &encoded)
.await
.expect("sqlite import should load exported rows");
assert_eq!(imported, 12);
assert_eq!(imported, 16);
let imported_api_key = sqlx::query_as::<_, (String,)>(
"SELECT key_encrypted FROM api_keys WHERE id = 'api-key-1'",
@@ -2184,6 +2286,15 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
.expect("imported usage should load");
assert_eq!(imported_usage.0, "request-1");
let imported_group_member = sqlx::query_as::<_, (String, String)>(
"SELECT group_id, user_id FROM user_group_members WHERE group_id = 'group-1' AND user_id = 'user-1'",
)
.fetch_one(&target_pool)
.await
.expect("imported user group member should load");
assert_eq!(imported_group_member.0, "group-1");
assert_eq!(imported_group_member.1, "user-1");
let imported_billing_rule = sqlx::query_as::<_, (String,)>(
"SELECT expression FROM billing_rules WHERE id = 'billing-rule-1'",
)
@@ -2217,7 +2328,7 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
let imported = import_postgres_jsonl(&postgres_pool, &encoded)
.await
.expect("postgres import should load exported rows");
assert_eq!(imported, 12);
assert_eq!(imported, 16);
let imported_api_key = sqlx::query_as::<_, (String,)>(
"SELECT key_encrypted FROM api_keys WHERE id = 'api-key-1'",
@@ -2273,6 +2384,7 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
let config_key = format!("export.config.{suffix}");
let wallet_id = format!("export-wallet-{suffix}");
let request_id = format!("export-request-{suffix}");
let group_id = format!("export-group-{suffix}");
sqlx::query(
"INSERT INTO users (id, email, username, auth_source, email_verified, created_at, updated_at) VALUES ($1, $2, $3, 'local', TRUE, to_timestamp(1), to_timestamp(2))",
@@ -2283,6 +2395,23 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
.execute(&pool)
.await
.expect("user should seed");
sqlx::query(
"INSERT INTO user_groups (id, name, normalized_name, priority, allowed_models, allowed_models_mode, created_at, updated_at) VALUES ($1, $2, $3, 10, '[\"provider-model\"]', 'specific', to_timestamp(1), to_timestamp(2))",
)
.bind(&group_id)
.bind(format!("Export Group {suffix}"))
.bind(format!("export group {suffix}"))
.execute(&pool)
.await
.expect("user group should seed");
sqlx::query(
"INSERT INTO user_group_members (group_id, user_id, created_at) VALUES ($1, $2, to_timestamp(1))",
)
.bind(&group_id)
.bind(&user_id)
.execute(&pool)
.await
.expect("user group member should seed");
sqlx::query(
"INSERT INTO api_keys (id, user_id, key_hash, key_encrypted, name, created_at, updated_at) VALUES ($1, $2, $3, 'ciphertext-1', 'Default', to_timestamp(1), to_timestamp(2))",
)
@@ -2389,6 +2518,14 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
.rows(ExportDomain::Users)
.iter()
.any(|row| row.id == user_id));
assert!(import_plan
.rows(ExportDomain::UserGroups)
.iter()
.any(|row| row.id == group_id));
assert!(import_plan
.rows(ExportDomain::UserGroupMembers)
.iter()
.any(|row| row.id == format!("{group_id}:{user_id}")));
assert!(import_plan
.rows(ExportDomain::ApiKeys)
.iter()
@@ -2432,6 +2569,16 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
.await
.expect("imported sqlite api key should load");
assert_eq!(imported_api_key.0, "ciphertext-1");
let imported_group_member = sqlx::query_as::<_, (String, String)>(
"SELECT group_id, user_id FROM user_group_members WHERE group_id = ? AND user_id = ?",
)
.bind(&group_id)
.bind(&user_id)
.fetch_one(&target_pool)
.await
.expect("imported sqlite user group member should load");
assert_eq!(imported_group_member.0, group_id);
assert_eq!(imported_group_member.1, user_id);
}
#[tokio::test]
@@ -2466,6 +2613,7 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
let config_id = format!("export-config-{suffix}");
let wallet_id = format!("export-wallet-{suffix}");
let request_id = format!("export-request-{suffix}");
let group_id = format!("export-group-{suffix}");
sqlx::query(
"INSERT INTO users (id, email, username, auth_source, created_at, updated_at) VALUES (?, ?, ?, 'local', 1, 2)",
@@ -2476,6 +2624,23 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
.execute(&pool)
.await
.expect("user should seed");
sqlx::query(
"INSERT INTO user_groups (id, name, normalized_name, priority, allowed_models, allowed_models_mode, created_at, updated_at) VALUES (?, ?, ?, 10, '[\"provider-model\"]', 'specific', 1, 2)",
)
.bind(&group_id)
.bind(format!("Export Group {suffix}"))
.bind(format!("export group {suffix}"))
.execute(&pool)
.await
.expect("user group should seed");
sqlx::query(
"INSERT INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, 1)",
)
.bind(&group_id)
.bind(&user_id)
.execute(&pool)
.await
.expect("user group member should seed");
sqlx::query(
"INSERT INTO api_keys (id, user_id, key_hash, key_encrypted, name, created_at, updated_at) VALUES (?, ?, ?, 'ciphertext-1', 'Default', 1, 2)",
)
@@ -2566,6 +2731,14 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
.rows(ExportDomain::Users)
.iter()
.any(|row| row.id == user_id));
assert!(import_plan
.rows(ExportDomain::UserGroups)
.iter()
.any(|row| row.id == group_id));
assert!(import_plan
.rows(ExportDomain::UserGroupMembers)
.iter()
.any(|row| row.id == format!("{group_id}:{user_id}")));
assert!(import_plan
.rows(ExportDomain::ApiKeys)
.iter()
@@ -2585,6 +2758,8 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
&pool,
vec![
ExportDomain::Users,
ExportDomain::UserGroups,
ExportDomain::UserGroupMembers,
ExportDomain::ApiKeys,
ExportDomain::ProviderKeys,
ExportDomain::Usage,
@@ -2596,7 +2771,7 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
let imported = import_mysql_jsonl(&pool, &selected_export)
.await
.expect("mysql import should be idempotent");
assert!(imported >= 4);
assert!(imported >= 6);
let imported_api_key =
sqlx::query_as::<_, (String,)>("SELECT key_encrypted FROM api_keys WHERE id = ?")

View File

@@ -295,6 +295,7 @@ fn empty_database_snapshot_covers_current_cutoff_versions() {
20260507120000,
20260508000000,
20260509000000,
20260509120000,
20260510000000,
]
);
@@ -557,7 +558,8 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() {
20260403000000,
20260507120000,
20260508000000,
20260509000000
20260509000000,
20260509120000
]
);
assert_eq!(
@@ -566,7 +568,8 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() {
20260403000000,
20260507120000,
20260508000000,
20260509000000
20260509000000,
20260509120000
]
);
}
@@ -1073,6 +1076,7 @@ fn pending_migrations_from_applied_skips_versions_already_applied() {
20260507120000,
20260508000000,
20260509000000,
20260509120000,
20260510000000,
]
);

View File

@@ -187,7 +187,7 @@ impl ResolvedAuthApiKeySnapshot {
}
non_empty_allowed_list(self.api_key_allowed_providers.as_deref())
.or_else(|| non_empty_allowed_list(self.user_allowed_providers.as_deref()))
.or(self.user_allowed_providers.as_deref())
}
pub fn effective_allowed_api_formats(&self) -> Option<&[String]> {
@@ -196,7 +196,7 @@ impl ResolvedAuthApiKeySnapshot {
}
non_empty_allowed_list(self.api_key_allowed_api_formats.as_deref())
.or_else(|| non_empty_allowed_list(self.user_allowed_api_formats.as_deref()))
.or(self.user_allowed_api_formats.as_deref())
}
pub fn effective_allowed_models(&self) -> Option<&[String]> {
@@ -205,7 +205,20 @@ impl ResolvedAuthApiKeySnapshot {
}
non_empty_allowed_list(self.api_key_allowed_models.as_deref())
.or_else(|| non_empty_allowed_list(self.user_allowed_models.as_deref()))
.or(self.user_allowed_models.as_deref())
}
pub fn apply_user_policy(
&mut self,
allowed_providers: Option<Vec<String>>,
allowed_api_formats: Option<Vec<String>>,
allowed_models: Option<Vec<String>>,
rate_limit: Option<i32>,
) {
self.user_allowed_providers = allowed_providers;
self.user_allowed_api_formats = allowed_api_formats;
self.user_allowed_models = allowed_models;
self.user_rate_limit = rate_limit;
}
}

View File

@@ -4,9 +4,11 @@ use std::sync::RwLock;
use async_trait::async_trait;
use super::types::{
LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, StoredUserExportRow,
normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord,
StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership,
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
StoredUserSummary, UserExportListQuery, UserExportSummary, UserReadRepository,
StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSummary,
UserReadRepository,
};
use crate::DataLayerError;
@@ -34,6 +36,8 @@ pub struct InMemoryUserReadRepository {
preferences_by_user_id: RwLock<BTreeMap<String, StoredUserPreferenceRecord>>,
sessions_by_id: RwLock<BTreeMap<String, StoredUserSessionRecord>>,
model_settings_by_user_id: RwLock<BTreeMap<String, serde_json::Value>>,
groups_by_id: RwLock<BTreeMap<String, StoredUserGroup>>,
group_members: RwLock<BTreeMap<(String, String), chrono::DateTime<chrono::Utc>>>,
export_rows: RwLock<Vec<StoredUserExportRow>>,
read_only: bool,
}
@@ -57,6 +61,8 @@ impl InMemoryUserReadRepository {
preferences_by_user_id: RwLock::new(BTreeMap::new()),
sessions_by_id: RwLock::new(BTreeMap::new()),
model_settings_by_user_id: RwLock::new(BTreeMap::new()),
groups_by_id: RwLock::new(BTreeMap::new()),
group_members: RwLock::new(BTreeMap::new()),
export_rows: RwLock::new(Vec::new()),
read_only: false,
}
@@ -90,6 +96,8 @@ impl InMemoryUserReadRepository {
preferences_by_user_id: RwLock::new(BTreeMap::new()),
sessions_by_id: RwLock::new(BTreeMap::new()),
model_settings_by_user_id: RwLock::new(BTreeMap::new()),
groups_by_id: RwLock::new(BTreeMap::new()),
group_members: RwLock::new(BTreeMap::new()),
export_rows: RwLock::new(Vec::new()),
read_only: false,
}
@@ -109,6 +117,8 @@ impl InMemoryUserReadRepository {
preferences_by_user_id: RwLock::new(BTreeMap::new()),
sessions_by_id: RwLock::new(BTreeMap::new()),
model_settings_by_user_id: RwLock::new(BTreeMap::new()),
groups_by_id: RwLock::new(BTreeMap::new()),
group_members: RwLock::new(BTreeMap::new()),
export_rows: RwLock::new(items.into_iter().collect()),
read_only: false,
}
@@ -257,6 +267,130 @@ fn upsert_memory_ldap_identifiers(
}
}
fn memory_group_from_record(
record: UpsertUserGroupRecord,
) -> Result<StoredUserGroup, DataLayerError> {
let now = chrono::Utc::now();
let name = normalize_user_group_name(&record.name);
StoredUserGroup::new(
uuid::Uuid::new_v4().to_string(),
name.clone(),
name.to_ascii_lowercase(),
record.description,
record.priority,
record.allowed_providers.map(serde_json::Value::from),
record.allowed_providers_mode,
record.allowed_api_formats.map(serde_json::Value::from),
record.allowed_api_formats_mode,
record.allowed_models.map(serde_json::Value::from),
record.allowed_models_mode,
record.rate_limit,
record.rate_limit_mode,
Some(now),
Some(now),
)
}
fn memory_update_group_from_record(
mut group: StoredUserGroup,
record: UpsertUserGroupRecord,
) -> Result<StoredUserGroup, DataLayerError> {
let name = normalize_user_group_name(&record.name);
group.name = name.clone();
group.normalized_name = name.to_ascii_lowercase();
group.description = record.description;
group.priority = record.priority;
group.allowed_providers = record.allowed_providers;
group.allowed_providers_mode = record.allowed_providers_mode;
group.allowed_api_formats = record.allowed_api_formats;
group.allowed_api_formats_mode = record.allowed_api_formats_mode;
group.allowed_models = record.allowed_models;
group.allowed_models_mode = record.allowed_models_mode;
group.rate_limit = record.rate_limit;
group.rate_limit_mode = record.rate_limit_mode;
group.updated_at = Some(chrono::Utc::now());
StoredUserGroup::new(
group.id,
group.name,
group.normalized_name,
group.description,
group.priority,
group.allowed_providers.map(serde_json::Value::from),
group.allowed_providers_mode,
group.allowed_api_formats.map(serde_json::Value::from),
group.allowed_api_formats_mode,
group.allowed_models.map(serde_json::Value::from),
group.allowed_models_mode,
group.rate_limit,
group.rate_limit_mode,
group.created_at,
group.updated_at,
)
}
fn memory_group_members(
repository: &InMemoryUserReadRepository,
group_id: &str,
) -> Vec<StoredUserGroupMember> {
let members = repository
.group_members
.read()
.expect("user repository lock")
.clone();
let users = repository.auth_by_id.read().expect("user repository lock");
members
.into_iter()
.filter(|((candidate_group_id, _), _)| candidate_group_id == group_id)
.filter_map(|((candidate_group_id, user_id), created_at)| {
users.get(&user_id).map(|user| StoredUserGroupMember {
group_id: candidate_group_id,
user_id: user.id.clone(),
username: user.username.clone(),
email: user.email.clone(),
role: user.role.clone(),
is_active: user.is_active,
is_deleted: user.is_deleted,
created_at: Some(created_at),
})
})
.collect()
}
fn memory_export_row_from_auth_user(
repository: &InMemoryUserReadRepository,
user: &StoredUserAuthRecord,
) -> Result<StoredUserExportRow, DataLayerError> {
let model_capability_settings = repository
.model_settings_by_user_id
.read()
.expect("user repository lock")
.get(&user.id)
.cloned();
StoredUserExportRow::new(
user.id.clone(),
user.email.clone(),
user.email_verified,
user.username.clone(),
user.password_hash.clone(),
user.role.clone(),
user.auth_source.clone(),
user.allowed_providers.clone().map(serde_json::Value::from),
user.allowed_api_formats
.clone()
.map(serde_json::Value::from),
user.allowed_models.clone().map(serde_json::Value::from),
None,
model_capability_settings,
user.is_active,
)?
.with_policy_modes(
user.allowed_providers_mode.clone(),
user.allowed_api_formats_mode.clone(),
user.allowed_models_mode.clone(),
"system".to_string(),
)
}
#[async_trait]
impl UserReadRepository for InMemoryUserReadRepository {
async fn list_users_by_ids(
@@ -296,22 +430,36 @@ impl UserReadRepository for InMemoryUserReadRepository {
async fn list_non_admin_export_users(
&self,
) -> Result<Vec<StoredUserExportRow>, DataLayerError> {
let rows = self.export_rows.read().expect("user repository lock");
if !rows.is_empty() {
return Ok(rows
.iter()
.filter(|row| !row.role.eq_ignore_ascii_case("admin"))
.cloned()
.collect());
}
Ok(self
.export_rows
.auth_by_id
.read()
.expect("user repository lock")
.iter()
.filter(|row| !row.role.eq_ignore_ascii_case("admin"))
.cloned()
.collect())
.filter(|(_, user)| !user.role.eq_ignore_ascii_case("admin"))
.map(|(_, user)| memory_export_row_from_auth_user(self, user))
.collect::<Result<Vec<_>, _>>()?)
}
async fn list_export_users(&self) -> Result<Vec<StoredUserExportRow>, DataLayerError> {
let rows = self.export_rows.read().expect("user repository lock");
if !rows.is_empty() {
return Ok(rows.clone());
}
Ok(self
.export_rows
.auth_by_id
.read()
.expect("user repository lock")
.clone())
.values()
.map(|user| memory_export_row_from_auth_user(self, user))
.collect::<Result<Vec<_>, _>>()?)
}
async fn list_export_users_page(
@@ -329,6 +477,23 @@ impl UserReadRepository for InMemoryUserReadRepository {
if let Some(is_active) = query.is_active {
rows.retain(|row| row.is_active == is_active);
}
if let Some(group_id) = query
.group_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
let member_ids = self
.group_members
.read()
.expect("user repository lock")
.keys()
.filter_map(|(candidate_group_id, user_id)| {
(candidate_group_id == group_id).then(|| user_id.clone())
})
.collect::<std::collections::BTreeSet<_>>();
rows.retain(|row| member_ids.contains(&row.id));
}
if let Some(search) = query
.search
.as_deref()
@@ -375,6 +540,270 @@ impl UserReadRepository for InMemoryUserReadRepository {
.cloned())
}
async fn list_user_groups(&self) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut groups = self
.groups_by_id
.read()
.expect("user repository lock")
.values()
.cloned()
.collect::<Vec<_>>();
groups.sort_by(|left, right| {
left.name
.cmp(&right.name)
.then_with(|| left.id.cmp(&right.id))
});
Ok(groups)
}
async fn find_user_group_by_id(
&self,
group_id: &str,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
Ok(self
.groups_by_id
.read()
.expect("user repository lock")
.get(group_id)
.cloned())
}
async fn list_user_groups_by_ids(
&self,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let groups = self.groups_by_id.read().expect("user repository lock");
Ok(group_ids
.iter()
.filter_map(|group_id| groups.get(group_id).cloned())
.collect())
}
async fn create_user_group(
&self,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
if self.read_only {
return Ok(None);
}
let group = memory_group_from_record(record)?;
let mut groups = self.groups_by_id.write().expect("user repository lock");
if groups
.values()
.any(|existing| existing.normalized_name == group.normalized_name)
{
return Err(DataLayerError::InvalidInput(format!(
"duplicate user group name: {}",
group.name
)));
}
groups.insert(group.id.clone(), group.clone());
Ok(Some(group))
}
async fn update_user_group(
&self,
group_id: &str,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
if self.read_only {
return Ok(None);
}
let mut groups = self.groups_by_id.write().expect("user repository lock");
let Some(existing) = groups.get(group_id).cloned() else {
return Ok(None);
};
let group = memory_update_group_from_record(existing, record)?;
if groups.values().any(|existing| {
existing.id != group.id && existing.normalized_name == group.normalized_name
}) {
return Err(DataLayerError::InvalidInput(format!(
"duplicate user group name: {}",
group.name
)));
}
groups.insert(group.id.clone(), group.clone());
Ok(Some(group))
}
async fn delete_user_group(&self, group_id: &str) -> Result<bool, DataLayerError> {
if self.read_only {
return Ok(false);
}
let removed = self
.groups_by_id
.write()
.expect("user repository lock")
.remove(group_id)
.is_some();
if removed {
self.group_members
.write()
.expect("user repository lock")
.retain(|key, _| key.0 != group_id);
}
Ok(removed)
}
async fn list_user_group_members(
&self,
group_id: &str,
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
Ok(memory_group_members(self, group_id))
}
async fn replace_user_group_members(
&self,
group_id: &str,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
if self.read_only {
return Ok(Vec::new());
}
if !self
.groups_by_id
.read()
.expect("user repository lock")
.contains_key(group_id)
{
return Ok(Vec::new());
}
let valid_user_ids = {
let users = self.auth_by_id.read().expect("user repository lock");
user_ids
.iter()
.map(|value| value.trim())
.filter(|value| !value.is_empty())
.filter(|user_id| users.contains_key(*user_id))
.map(ToOwned::to_owned)
.collect::<std::collections::BTreeSet<_>>()
};
let now = chrono::Utc::now();
let mut members = self.group_members.write().expect("user repository lock");
members.retain(|key, _| key.0 != group_id);
for user_id in valid_user_ids {
members.insert((group_id.to_string(), user_id), now);
}
drop(members);
Ok(memory_group_members(self, group_id))
}
async fn list_user_groups_for_user(
&self,
user_id: &str,
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let group_ids = self
.group_members
.read()
.expect("user repository lock")
.keys()
.filter_map(|(group_id, candidate_user_id)| {
(candidate_user_id == user_id).then(|| group_id.clone())
})
.collect::<Vec<_>>();
self.list_user_groups_by_ids(&group_ids).await
}
async fn list_user_group_memberships_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMembership>, DataLayerError> {
let requested = user_ids
.iter()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.collect::<std::collections::BTreeSet<_>>();
if requested.is_empty() {
return Ok(Vec::new());
}
let groups = self.groups_by_id.read().expect("user repository lock");
let members = self.group_members.read().expect("user repository lock");
let mut memberships = members
.iter()
.filter(|((_, user_id), _)| requested.contains(user_id))
.filter_map(|((group_id, user_id), created_at)| {
groups.get(group_id).map(|group| StoredUserGroupMembership {
user_id: user_id.clone(),
group_id: group.id.clone(),
group_name: group.name.clone(),
group_priority: group.priority,
created_at: Some(*created_at),
})
})
.collect::<Vec<_>>();
memberships.sort_by(|left, right| {
left.user_id
.cmp(&right.user_id)
.then_with(|| left.group_name.cmp(&right.group_name))
.then_with(|| left.group_id.cmp(&right.group_id))
});
Ok(memberships)
}
async fn replace_user_groups_for_user(
&self,
user_id: &str,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
if self.read_only {
return Ok(Vec::new());
}
let existing_group_ids = {
let groups = self.groups_by_id.read().expect("user repository lock");
group_ids
.iter()
.map(|value| value.trim())
.filter(|value| !value.is_empty())
.filter(|group_id| groups.contains_key(*group_id))
.map(ToOwned::to_owned)
.collect::<std::collections::BTreeSet<_>>()
};
{
let now = chrono::Utc::now();
let mut members = self.group_members.write().expect("user repository lock");
members.retain(|key, _| key.1 != user_id);
for group_id in &existing_group_ids {
members.insert((group_id.clone(), user_id.to_string()), now);
}
}
self.list_user_groups_by_ids(&existing_group_ids.into_iter().collect::<Vec<_>>())
.await
}
async fn add_user_to_group(
&self,
group_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
if self.read_only {
return Ok(false);
}
if !self
.groups_by_id
.read()
.expect("user repository lock")
.contains_key(group_id)
{
return Ok(false);
}
if !self
.auth_by_id
.read()
.expect("user repository lock")
.contains_key(user_id)
{
return Ok(false);
}
self.group_members
.write()
.expect("user repository lock")
.insert(
(group_id.to_string(), user_id.to_string()),
chrono::Utc::now(),
);
Ok(true)
}
async fn find_user_auth_by_id(
&self,
user_id: &str,
@@ -581,6 +1010,11 @@ impl UserReadRepository for InMemoryUserReadRepository {
false,
Some(created_at),
Some(created_at),
)?
.with_policy_modes(
"inherit".to_string(),
"inherit".to_string(),
"inherit".to_string(),
)?;
self.insert_auth_user(user).map(Some)
}
@@ -917,12 +1351,27 @@ impl UserReadRepository for InMemoryUserReadRepository {
}
if allowed_providers_present {
user.allowed_providers = allowed_providers;
user.allowed_providers_mode = if user.allowed_providers.is_some() {
"specific".to_string()
} else {
"unrestricted".to_string()
};
}
if allowed_api_formats_present {
user.allowed_api_formats = allowed_api_formats;
user.allowed_api_formats_mode = if user.allowed_api_formats.is_some() {
"specific".to_string()
} else {
"unrestricted".to_string()
};
}
if allowed_models_present {
user.allowed_models = allowed_models;
user.allowed_models_mode = if user.allowed_models.is_some() {
"specific".to_string()
} else {
"unrestricted".to_string()
};
}
if let Some(is_active) = is_active {
user.is_active = is_active;
@@ -948,16 +1397,75 @@ impl UserReadRepository for InMemoryUserReadRepository {
{
row.role = updated.role.clone();
row.allowed_providers = updated.allowed_providers.clone();
row.allowed_providers_mode = updated.allowed_providers_mode.clone();
row.allowed_api_formats = updated.allowed_api_formats.clone();
row.allowed_api_formats_mode = updated.allowed_api_formats_mode.clone();
row.allowed_models = updated.allowed_models.clone();
row.allowed_models_mode = updated.allowed_models_mode.clone();
if rate_limit_present {
row.rate_limit = rate_limit;
row.rate_limit_mode = if row.rate_limit.is_some() {
"custom".to_string()
} else {
"system".to_string()
};
}
row.is_active = updated.is_active;
}
Ok(Some(updated))
}
async fn update_local_auth_user_policy_modes(
&self,
user_id: &str,
allowed_providers_mode: Option<String>,
allowed_api_formats_mode: Option<String>,
allowed_models_mode: Option<String>,
rate_limit_mode: Option<String>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
if self.read_only {
return Ok(None);
}
let mut auth_by_id = self.auth_by_id.write().expect("user repository lock");
let Some(user) = auth_by_id.get_mut(user_id) else {
return Ok(None);
};
if let Some(mode) = allowed_providers_mode.clone() {
user.allowed_providers_mode = mode;
}
if let Some(mode) = allowed_api_formats_mode.clone() {
user.allowed_api_formats_mode = mode;
}
if let Some(mode) = allowed_models_mode.clone() {
user.allowed_models_mode = mode;
}
let updated = user.clone();
drop(auth_by_id);
if let Some(row) = self
.export_rows
.write()
.expect("user repository lock")
.iter_mut()
.find(|row| row.id == user_id)
{
if let Some(mode) = allowed_providers_mode {
row.allowed_providers_mode = mode;
}
if let Some(mode) = allowed_api_formats_mode {
row.allowed_api_formats_mode = mode;
}
if let Some(mode) = allowed_models_mode {
row.allowed_models_mode = mode;
}
if let Some(mode) = rate_limit_mode {
row.rate_limit_mode = mode;
}
}
Ok(Some(updated))
}
async fn update_user_model_capability_settings(
&self,
user_id: &str,
@@ -1021,18 +1529,29 @@ impl UserReadRepository for InMemoryUserReadRepository {
return Ok(None);
}
self.create_local_auth_user_with_settings(
let now = chrono::Utc::now();
let user = StoredUserAuthRecord::new(
uuid::Uuid::new_v4().to_string(),
email,
email_verified,
username,
password_hash,
Some(password_hash),
"user".to_string(),
"local".to_string(),
None,
None,
None,
true,
false,
Some(now),
None,
)
.await
)?
.with_policy_modes(
"inherit".to_string(),
"inherit".to_string(),
"inherit".to_string(),
)?;
self.insert_auth_user(user).map(Some)
}
async fn create_local_auth_user_with_settings(
@@ -1092,6 +1611,10 @@ impl UserReadRepository for InMemoryUserReadRepository {
.write()
.expect("user repository lock")
.retain(|_, link| link.user_id != user_id);
self.group_members
.write()
.expect("user repository lock")
.retain(|key, _| key.1 != user_id);
let mut identifiers = self
.auth_by_identifier
@@ -2073,6 +2596,7 @@ mod tests {
role: Some("user".to_string()),
is_active: Some(true),
search: None,
group_id: None,
})
.await
.expect("paged export should succeed");

View File

@@ -9,7 +9,8 @@ pub use mysql::MysqlUserReadRepository;
pub use postgres::SqlxUserReadRepository;
pub use sqlite::SqliteUserReadRepository;
pub use types::{
StoredUserAuthRecord, StoredUserExportRow, StoredUserOAuthLinkSummary,
StoredUserPreferenceRecord, StoredUserSessionRecord, StoredUserSummary, UserExportListQuery,
UserExportSummary, UserReadRepository,
normalize_user_group_name, StoredUserAuthRecord, StoredUserExportRow, StoredUserGroup,
StoredUserGroupMember, StoredUserGroupMembership, StoredUserOAuthLinkSummary,
StoredUserPreferenceRecord, StoredUserSessionRecord, StoredUserSummary, UpsertUserGroupRecord,
UserExportListQuery, UserExportSummary, UserReadRepository,
};

View File

@@ -3,9 +3,11 @@ use chrono::{DateTime, TimeZone, Utc};
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use super::types::{
LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, StoredUserExportRow,
normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord,
StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership,
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
StoredUserSummary, UserExportListQuery, UserExportSummary, UserReadRepository,
StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSummary,
UserReadRepository,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
@@ -32,9 +34,13 @@ SELECT
role,
auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
model_capability_settings,
is_active
FROM users
@@ -50,8 +56,11 @@ SELECT
role,
auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -69,8 +78,11 @@ SELECT
users.role AS role,
users.auth_source AS auth_source,
users.allowed_providers AS allowed_providers,
users.allowed_providers_mode AS allowed_providers_mode,
users.allowed_api_formats AS allowed_api_formats,
users.allowed_api_formats_mode AS allowed_api_formats_mode,
users.allowed_models AS allowed_models,
users.allowed_models_mode AS allowed_models_mode,
users.is_active AS is_active,
users.is_deleted AS is_deleted,
users.created_at AS created_at,
@@ -130,6 +142,40 @@ SELECT
FROM user_sessions
"#;
const USER_GROUP_COLUMNS: &str = r#"
SELECT
id,
name,
normalized_name,
description,
priority,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
created_at,
updated_at
FROM user_groups
"#;
const USER_GROUP_MEMBER_COLUMNS: &str = r#"
SELECT
user_group_members.group_id,
users.id AS user_id,
users.username,
users.email,
users.role,
users.is_active,
users.is_deleted,
user_group_members.created_at
FROM user_group_members
JOIN users ON users.id = user_group_members.user_id
"#;
#[derive(Debug, Clone)]
pub struct MysqlUserReadRepository {
pool: MysqlPool,
@@ -163,6 +209,22 @@ impl MysqlUserReadRepository {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_auth_row).collect()
}
async fn fetch_group_rows(
&self,
mut builder: QueryBuilder<'_, MySql>,
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_group_row).collect()
}
async fn fetch_group_member_rows(
&self,
mut builder: QueryBuilder<'_, MySql>,
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_group_member_row).collect()
}
}
#[async_trait]
@@ -224,6 +286,16 @@ impl UserReadRepository for MysqlUserReadRepository {
if let Some(is_active) = query.is_active {
builder.push(" AND is_active = ").push_bind(is_active);
}
if let Some(group_id) = query
.group_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
builder.push(" AND id IN (SELECT user_id FROM user_group_members WHERE group_id = ");
builder.push_bind(group_id);
builder.push(")");
}
if let Some(search) = query
.search
.as_deref()
@@ -296,6 +368,287 @@ WHERE is_deleted = 0
self.fetch_export_rows(builder).await
}
async fn list_user_groups(&self) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(USER_GROUP_COLUMNS);
builder.push(" ORDER BY name ASC, id ASC");
self.fetch_group_rows(builder).await
}
async fn find_user_group_by_id(
&self,
group_id: &str,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(USER_GROUP_COLUMNS);
builder
.push(" WHERE id = ")
.push_bind(group_id)
.push(" LIMIT 1");
Ok(self.fetch_group_rows(builder).await?.into_iter().next())
}
async fn list_user_groups_by_ids(
&self,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
if group_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<MySql>::new(USER_GROUP_COLUMNS);
builder.push(" WHERE id IN (");
{
let mut separated = builder.separated(", ");
for group_id in group_ids {
separated.push_bind(group_id);
}
}
builder.push(") ORDER BY name ASC, id ASC");
self.fetch_group_rows(builder).await
}
async fn create_user_group(
&self,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let now = current_unix_secs();
let id = uuid::Uuid::new_v4().to_string();
let name = normalize_user_group_name(&record.name);
let normalized_name = name.to_ascii_lowercase();
let result = sqlx::query(
r#"
INSERT INTO user_groups (
id, name, normalized_name, description, priority,
allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode,
allowed_models, allowed_models_mode,
rate_limit, rate_limit_mode, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&id)
.bind(name)
.bind(normalized_name)
.bind(record.description)
.bind(record.priority)
.bind(json_string_from_option_vec(
record.allowed_providers.as_ref(),
))
.bind(record.allowed_providers_mode)
.bind(json_string_from_option_vec(
record.allowed_api_formats.as_ref(),
))
.bind(record.allowed_api_formats_mode)
.bind(json_string_from_option_vec(record.allowed_models.as_ref()))
.bind(record.allowed_models_mode)
.bind(record.rate_limit)
.bind(record.rate_limit_mode)
.bind(now)
.bind(now)
.execute(&self.pool)
.await;
match result {
Ok(_) => self.find_user_group_by_id(&id).await,
Err(sqlx::Error::Database(err)) if err.is_unique_violation() => Err(
DataLayerError::InvalidInput("duplicate user group name".to_string()),
),
Err(err) => Err(err).map_sql_err(),
}
}
async fn update_user_group(
&self,
group_id: &str,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let now = current_unix_secs();
let name = normalize_user_group_name(&record.name);
let normalized_name = name.to_ascii_lowercase();
let result = sqlx::query(
r#"
UPDATE user_groups
SET name = ?,
normalized_name = ?,
description = ?,
priority = ?,
allowed_providers = ?,
allowed_providers_mode = ?,
allowed_api_formats = ?,
allowed_api_formats_mode = ?,
allowed_models = ?,
allowed_models_mode = ?,
rate_limit = ?,
rate_limit_mode = ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(name)
.bind(normalized_name)
.bind(record.description)
.bind(record.priority)
.bind(json_string_from_option_vec(
record.allowed_providers.as_ref(),
))
.bind(record.allowed_providers_mode)
.bind(json_string_from_option_vec(
record.allowed_api_formats.as_ref(),
))
.bind(record.allowed_api_formats_mode)
.bind(json_string_from_option_vec(record.allowed_models.as_ref()))
.bind(record.allowed_models_mode)
.bind(record.rate_limit)
.bind(record.rate_limit_mode)
.bind(now)
.bind(group_id)
.execute(&self.pool)
.await;
match result {
Ok(result) if result.rows_affected() == 0 => Ok(None),
Ok(_) => self.find_user_group_by_id(group_id).await,
Err(sqlx::Error::Database(err)) if err.is_unique_violation() => Err(
DataLayerError::InvalidInput("duplicate user group name".to_string()),
),
Err(err) => Err(err).map_sql_err(),
}
}
async fn delete_user_group(&self, group_id: &str) -> Result<bool, DataLayerError> {
let result = sqlx::query("DELETE FROM user_groups WHERE id = ?")
.bind(group_id)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
async fn list_user_group_members(
&self,
group_id: &str,
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(USER_GROUP_MEMBER_COLUMNS);
builder
.push(" WHERE user_group_members.group_id = ")
.push_bind(group_id)
.push(" ORDER BY users.username ASC, users.id ASC");
self.fetch_group_member_rows(builder).await
}
async fn replace_user_group_members(
&self,
group_id: &str,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("DELETE FROM user_group_members WHERE group_id = ?")
.bind(group_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
let now = current_unix_secs();
for user_id in normalized_ids(user_ids) {
sqlx::query(
"INSERT IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)",
)
.bind(group_id)
.bind(user_id)
.bind(now)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
tx.commit().await.map_sql_err()?;
self.list_user_group_members(group_id).await
}
async fn list_user_groups_for_user(
&self,
user_id: &str,
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(USER_GROUP_COLUMNS);
builder
.push(" WHERE id IN (SELECT group_id FROM user_group_members WHERE user_id = ")
.push_bind(user_id)
.push(") ORDER BY name ASC, id ASC");
self.fetch_group_rows(builder).await
}
async fn list_user_group_memberships_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMembership>, DataLayerError> {
if user_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<MySql>::new(
r#"
SELECT
user_group_members.user_id,
user_groups.id AS group_id,
user_groups.name AS group_name,
user_groups.priority AS group_priority,
user_group_members.created_at
FROM user_group_members
JOIN user_groups ON user_groups.id = user_group_members.group_id
WHERE user_group_members.user_id IN (
"#,
);
{
let mut separated = builder.separated(", ");
for user_id in user_ids {
separated.push_bind(user_id);
}
}
builder.push(
") ORDER BY user_group_members.user_id ASC, user_groups.name ASC, user_groups.id ASC",
);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_group_membership_row).collect()
}
async fn replace_user_groups_for_user(
&self,
user_id: &str,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("DELETE FROM user_group_members WHERE user_id = ?")
.bind(user_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
let now = current_unix_secs();
for group_id in normalized_ids(group_ids) {
sqlx::query(
"INSERT IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)",
)
.bind(group_id)
.bind(user_id)
.bind(now)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
tx.commit().await.map_sql_err()?;
self.list_user_groups_for_user(user_id).await
}
async fn add_user_to_group(
&self,
group_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
let result = sqlx::query(
"INSERT IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)",
)
.bind(group_id)
.bind(user_id)
.bind(current_unix_secs())
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
async fn find_user_auth_by_id(
&self,
user_id: &str,
@@ -453,9 +806,10 @@ WHERE provider_type = ?
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
is_active, is_deleted, created_at, updated_at, last_login_at
)
VALUES (?, ?, 1, ?, NULL, 'user', 'oauth', 1, 0, ?, ?, ?)
VALUES (?, ?, 1, ?, NULL, 'user', 'oauth', 'inherit', 'inherit', 'inherit', 'inherit', 1, 0, ?, ?, ?)
"#,
)
.bind(&user_id)
@@ -675,14 +1029,38 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
rate_limit: Option<i32>,
is_active: Option<bool>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let allowed_providers_mode = if allowed_providers.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_api_formats_mode = if allowed_api_formats.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_models_mode = if allowed_models.is_some() {
"specific"
} else {
"unrestricted"
};
let rate_limit_mode = if rate_limit.is_some() {
"custom"
} else {
"system"
};
let result = sqlx::query(
r#"
UPDATE users
SET role = CASE WHEN ? THEN COALESCE(?, role) ELSE role END,
allowed_providers = CASE WHEN ? THEN ? ELSE allowed_providers END,
allowed_providers_mode = CASE WHEN ? THEN ? ELSE allowed_providers_mode END,
allowed_api_formats = CASE WHEN ? THEN ? ELSE allowed_api_formats END,
allowed_api_formats_mode = CASE WHEN ? THEN ? ELSE allowed_api_formats_mode END,
allowed_models = CASE WHEN ? THEN ? ELSE allowed_models END,
allowed_models_mode = CASE WHEN ? THEN ? ELSE allowed_models_mode END,
rate_limit = CASE WHEN ? THEN ? ELSE rate_limit END,
rate_limit_mode = CASE WHEN ? THEN ? ELSE rate_limit_mode END,
is_active = CASE WHEN ? THEN ? ELSE is_active END,
updated_at = ?
WHERE id = ?
@@ -695,18 +1073,26 @@ WHERE id = ?
allowed_providers,
"users.allowed_providers",
)?)
.bind(allowed_providers_present)
.bind(allowed_providers_mode)
.bind(allowed_api_formats_present)
.bind(optional_string_list_json(
allowed_api_formats,
"users.allowed_api_formats",
)?)
.bind(allowed_api_formats_present)
.bind(allowed_api_formats_mode)
.bind(allowed_models_present)
.bind(optional_string_list_json(
allowed_models,
"users.allowed_models",
)?)
.bind(allowed_models_present)
.bind(allowed_models_mode)
.bind(rate_limit_present)
.bind(rate_limit)
.bind(rate_limit_present)
.bind(rate_limit_mode)
.bind(is_active.is_some())
.bind(is_active)
.bind(chrono::Utc::now().timestamp())
@@ -720,6 +1106,44 @@ WHERE id = ?
self.find_user_auth_by_id(user_id).await
}
async fn update_local_auth_user_policy_modes(
&self,
user_id: &str,
allowed_providers_mode: Option<String>,
allowed_api_formats_mode: Option<String>,
allowed_models_mode: Option<String>,
rate_limit_mode: Option<String>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let result = sqlx::query(
r#"
UPDATE users
SET allowed_providers_mode = CASE WHEN ? THEN ? ELSE allowed_providers_mode END,
allowed_api_formats_mode = CASE WHEN ? THEN ? ELSE allowed_api_formats_mode END,
allowed_models_mode = CASE WHEN ? THEN ? ELSE allowed_models_mode END,
rate_limit_mode = CASE WHEN ? THEN ? ELSE rate_limit_mode END,
updated_at = ?
WHERE id = ?
"#,
)
.bind(allowed_providers_mode.is_some())
.bind(allowed_providers_mode)
.bind(allowed_api_formats_mode.is_some())
.bind(allowed_api_formats_mode)
.bind(allowed_models_mode.is_some())
.bind(allowed_models_mode)
.bind(rate_limit_mode.is_some())
.bind(rate_limit_mode)
.bind(chrono::Utc::now().timestamp())
.bind(user_id)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.find_user_auth_by_id(user_id).await
}
async fn update_user_model_capability_settings(
&self,
user_id: &str,
@@ -751,18 +1175,29 @@ WHERE id = ?
username: String,
password_hash: String,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
self.create_local_auth_user_with_settings(
email,
email_verified,
username,
password_hash,
"user".to_string(),
None,
None,
None,
None,
let user_id = uuid::Uuid::new_v4().to_string();
let now = chrono::Utc::now().timestamp();
sqlx::query(
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
is_active, is_deleted, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, 'user', 'local', 'inherit', 'inherit', 'inherit', 'inherit', 1, 0, ?, ?)
"#,
)
.bind(&user_id)
.bind(email)
.bind(email_verified)
.bind(username)
.bind(password_hash)
.bind(now)
.bind(now)
.execute(&self.pool)
.await
.map_sql_err()?;
self.find_user_auth_by_id(&user_id).await
}
async fn create_local_auth_user_with_settings(
@@ -779,14 +1214,37 @@ WHERE id = ?
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let user_id = uuid::Uuid::new_v4().to_string();
let now = chrono::Utc::now().timestamp();
let allowed_providers_mode = if allowed_providers.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_api_formats_mode = if allowed_api_formats.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_models_mode = if allowed_models.is_some() {
"specific"
} else {
"unrestricted"
};
let rate_limit_mode = if rate_limit.is_some() {
"custom"
} else {
"system"
};
sqlx::query(
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers, allowed_api_formats, allowed_models, rate_limit,
allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode,
allowed_models, allowed_models_mode,
rate_limit, rate_limit_mode,
is_active, is_deleted, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, 1, 0, ?, ?)
VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, ?, ?, ?, ?, 1, 0, ?, ?)
"#,
)
.bind(&user_id)
@@ -799,15 +1257,19 @@ VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, 1, 0, ?, ?)
allowed_providers,
"users.allowed_providers",
)?)
.bind(allowed_providers_mode)
.bind(optional_string_list_json(
allowed_api_formats,
"users.allowed_api_formats",
)?)
.bind(allowed_api_formats_mode)
.bind(optional_string_list_json(
allowed_models,
"users.allowed_models",
)?)
.bind(allowed_models_mode)
.bind(rate_limit)
.bind(rate_limit_mode)
.bind(now)
.bind(now)
.execute(&self.pool)
@@ -1162,6 +1624,24 @@ fn optional_string_list_json(
.transpose()
}
fn json_string_from_option_vec(value: Option<&Vec<String>>) -> Option<String> {
value.and_then(|items| serde_json::to_string(items).ok())
}
fn normalized_ids(values: &[String]) -> Vec<String> {
values
.iter()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.collect::<std::collections::BTreeSet<_>>()
.into_iter()
.collect()
}
fn current_unix_secs() -> i64 {
chrono::Utc::now().timestamp()
}
fn optional_json_string(
value: Option<serde_json::Value>,
field_name: &str,
@@ -1378,6 +1858,14 @@ fn map_user_export_row(row: &MySqlRow) -> Result<StoredUserExportRow, DataLayerE
)?,
row.try_get("is_active").map_sql_err()?,
)
.and_then(|record| {
record.with_policy_modes(
row.try_get("allowed_providers_mode").map_sql_err()?,
row.try_get("allowed_api_formats_mode").map_sql_err()?,
row.try_get("allowed_models_mode").map_sql_err()?,
row.try_get("rate_limit_mode").map_sql_err()?,
)
})
}
fn map_user_auth_row(row: &MySqlRow) -> Result<StoredUserAuthRecord, DataLayerError> {
@@ -1406,6 +1894,67 @@ fn map_user_auth_row(row: &MySqlRow) -> Result<StoredUserAuthRecord, DataLayerEr
optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
optional_datetime_from_unix_secs(row.try_get("last_login_at").map_sql_err()?),
)
.and_then(|record| {
record.with_policy_modes(
row.try_get("allowed_providers_mode").map_sql_err()?,
row.try_get("allowed_api_formats_mode").map_sql_err()?,
row.try_get("allowed_models_mode").map_sql_err()?,
)
})
}
fn map_user_group_row(row: &MySqlRow) -> Result<StoredUserGroup, DataLayerError> {
StoredUserGroup::new(
row.try_get("id").map_sql_err()?,
row.try_get("name").map_sql_err()?,
row.try_get("normalized_name").map_sql_err()?,
row.try_get("description").map_sql_err()?,
row.try_get("priority").map_sql_err()?,
optional_json_from_string(
row.try_get("allowed_providers").map_sql_err()?,
"user_groups.allowed_providers",
)?,
row.try_get("allowed_providers_mode").map_sql_err()?,
optional_json_from_string(
row.try_get("allowed_api_formats").map_sql_err()?,
"user_groups.allowed_api_formats",
)?,
row.try_get("allowed_api_formats_mode").map_sql_err()?,
optional_json_from_string(
row.try_get("allowed_models").map_sql_err()?,
"user_groups.allowed_models",
)?,
row.try_get("allowed_models_mode").map_sql_err()?,
row.try_get("rate_limit").map_sql_err()?,
row.try_get("rate_limit_mode").map_sql_err()?,
optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
optional_datetime_from_unix_secs(row.try_get("updated_at").map_sql_err()?),
)
}
fn map_user_group_member_row(row: &MySqlRow) -> Result<StoredUserGroupMember, DataLayerError> {
Ok(StoredUserGroupMember {
group_id: row.try_get("group_id").map_sql_err()?,
user_id: row.try_get("user_id").map_sql_err()?,
username: row.try_get("username").map_sql_err()?,
email: row.try_get("email").map_sql_err()?,
role: row.try_get("role").map_sql_err()?,
is_active: row.try_get("is_active").map_sql_err()?,
is_deleted: row.try_get("is_deleted").map_sql_err()?,
created_at: optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
})
}
fn map_user_group_membership_row(
row: &MySqlRow,
) -> Result<StoredUserGroupMembership, DataLayerError> {
Ok(StoredUserGroupMembership {
user_id: row.try_get("user_id").map_sql_err()?,
group_id: row.try_get("group_id").map_sql_err()?,
group_name: row.try_get("group_name").map_sql_err()?,
group_priority: row.try_get("group_priority").map_sql_err()?,
created_at: optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
})
}
fn map_oauth_link_summary_row(

View File

@@ -3,9 +3,11 @@ use futures_util::TryStreamExt;
use sqlx::{PgPool, Postgres, QueryBuilder, Row};
use super::types::{
LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, StoredUserExportRow,
normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord,
StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership,
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
StoredUserSummary, UserExportListQuery, UserExportSummary, UserReadRepository,
StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSummary,
UserReadRepository,
};
use crate::{error::SqlxResultExt, DataLayerError};
@@ -46,9 +48,13 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
model_capability_settings,
is_active
FROM users
@@ -67,9 +73,13 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
model_capability_settings,
is_active
FROM users
@@ -87,9 +97,13 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
model_capability_settings,
is_active
FROM users
@@ -132,9 +146,13 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
model_capability_settings,
is_active
FROM users
@@ -153,8 +171,11 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -174,8 +195,11 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -195,8 +219,11 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -216,8 +243,11 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -237,8 +267,11 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -259,8 +292,11 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -550,6 +586,40 @@ SET revoked_at = $2, revoke_reason = $3, updated_at = $2
WHERE user_id = $1 AND revoked_at IS NULL
"#;
const USER_GROUP_COLUMNS: &str = r#"
SELECT
id,
name,
normalized_name,
description,
priority,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
created_at,
updated_at
FROM user_groups
"#;
const USER_GROUP_MEMBER_COLUMNS: &str = r#"
SELECT
user_group_members.group_id,
users.id AS user_id,
users.username,
users.email,
users.role::text AS role,
users.is_active,
users.is_deleted,
user_group_members.created_at
FROM user_group_members
JOIN users ON users.id = user_group_members.user_id
"#;
#[derive(Debug, Clone)]
pub struct SqlxUserReadRepository {
pool: PgPool,
@@ -612,6 +682,277 @@ impl SqlxUserReadRepository {
.await
}
pub async fn list_user_groups(&self) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut builder = QueryBuilder::<Postgres>::new(USER_GROUP_COLUMNS);
builder.push(" ORDER BY name ASC, id ASC");
collect_query_rows(builder.build().fetch(&self.pool), map_user_group_row).await
}
pub async fn find_user_group_by_id(
&self,
group_id: &str,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let mut builder = QueryBuilder::<Postgres>::new(USER_GROUP_COLUMNS);
builder
.push(" WHERE id = ")
.push_bind(group_id)
.push(" LIMIT 1");
let row = builder
.build()
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
row.as_ref().map(map_user_group_row).transpose()
}
pub async fn list_user_groups_by_ids(
&self,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
if group_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Postgres>::new(USER_GROUP_COLUMNS);
builder.push(" WHERE id IN (");
{
let mut separated = builder.separated(", ");
for group_id in group_ids {
separated.push_bind(group_id);
}
}
builder.push(") ORDER BY name ASC, id ASC");
collect_query_rows(builder.build().fetch(&self.pool), map_user_group_row).await
}
pub async fn create_user_group(
&self,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let id = uuid::Uuid::new_v4().to_string();
let name = normalize_user_group_name(&record.name);
let normalized_name = name.to_ascii_lowercase();
let result = sqlx::query(
r#"
INSERT INTO user_groups (
id, name, normalized_name, description, priority,
allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode,
allowed_models, allowed_models_mode,
rate_limit, rate_limit_mode
)
VALUES ($1, $2, $3, $4, $5, $6::json, $7, $8::json, $9, $10::json, $11, $12, $13)
"#,
)
.bind(&id)
.bind(name)
.bind(normalized_name)
.bind(record.description)
.bind(record.priority)
.bind(record.allowed_providers.map(serde_json::Value::from))
.bind(record.allowed_providers_mode)
.bind(record.allowed_api_formats.map(serde_json::Value::from))
.bind(record.allowed_api_formats_mode)
.bind(record.allowed_models.map(serde_json::Value::from))
.bind(record.allowed_models_mode)
.bind(record.rate_limit)
.bind(record.rate_limit_mode)
.execute(&self.pool)
.await;
match result {
Ok(_) => self.find_user_group_by_id(&id).await,
Err(sqlx::Error::Database(err)) if err.is_unique_violation() => Err(
DataLayerError::InvalidInput("duplicate user group name".to_string()),
),
Err(err) => Err(err).map_postgres_err(),
}
}
pub async fn update_user_group(
&self,
group_id: &str,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let name = normalize_user_group_name(&record.name);
let normalized_name = name.to_ascii_lowercase();
let result = sqlx::query(
r#"
UPDATE user_groups
SET name = $2,
normalized_name = $3,
description = $4,
priority = $5,
allowed_providers = $6::json,
allowed_providers_mode = $7,
allowed_api_formats = $8::json,
allowed_api_formats_mode = $9,
allowed_models = $10::json,
allowed_models_mode = $11,
rate_limit = $12,
rate_limit_mode = $13,
updated_at = now()
WHERE id = $1
"#,
)
.bind(group_id)
.bind(name)
.bind(normalized_name)
.bind(record.description)
.bind(record.priority)
.bind(record.allowed_providers.map(serde_json::Value::from))
.bind(record.allowed_providers_mode)
.bind(record.allowed_api_formats.map(serde_json::Value::from))
.bind(record.allowed_api_formats_mode)
.bind(record.allowed_models.map(serde_json::Value::from))
.bind(record.allowed_models_mode)
.bind(record.rate_limit)
.bind(record.rate_limit_mode)
.execute(&self.pool)
.await;
match result {
Ok(result) if result.rows_affected() == 0 => Ok(None),
Ok(_) => self.find_user_group_by_id(group_id).await,
Err(sqlx::Error::Database(err)) if err.is_unique_violation() => Err(
DataLayerError::InvalidInput("duplicate user group name".to_string()),
),
Err(err) => Err(err).map_postgres_err(),
}
}
pub async fn delete_user_group(&self, group_id: &str) -> Result<bool, DataLayerError> {
let result = sqlx::query("DELETE FROM user_groups WHERE id = $1")
.bind(group_id)
.execute(&self.pool)
.await
.map_postgres_err()?;
Ok(result.rows_affected() > 0)
}
pub async fn list_user_group_members(
&self,
group_id: &str,
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let mut builder = QueryBuilder::<Postgres>::new(USER_GROUP_MEMBER_COLUMNS);
builder
.push(" WHERE user_group_members.group_id = ")
.push_bind(group_id)
.push(" ORDER BY users.username ASC, users.id ASC");
collect_query_rows(builder.build().fetch(&self.pool), map_user_group_member_row).await
}
pub async fn replace_user_group_members(
&self,
group_id: &str,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let mut tx = self.pool.begin().await.map_postgres_err()?;
sqlx::query("DELETE FROM user_group_members WHERE group_id = $1")
.bind(group_id)
.execute(&mut *tx)
.await
.map_postgres_err()?;
for user_id in normalized_ids(user_ids) {
sqlx::query(
"INSERT INTO user_group_members (group_id, user_id) VALUES ($1, $2) ON CONFLICT (group_id, user_id) DO NOTHING",
)
.bind(group_id)
.bind(user_id)
.execute(&mut *tx)
.await
.map_postgres_err()?;
}
tx.commit().await.map_postgres_err()?;
self.list_user_group_members(group_id).await
}
pub async fn list_user_groups_for_user(
&self,
user_id: &str,
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut builder = QueryBuilder::<Postgres>::new(USER_GROUP_COLUMNS);
builder
.push(" WHERE id IN (SELECT group_id FROM user_group_members WHERE user_id = ")
.push_bind(user_id)
.push(") ORDER BY name ASC, id ASC");
collect_query_rows(builder.build().fetch(&self.pool), map_user_group_row).await
}
pub async fn list_user_group_memberships_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMembership>, DataLayerError> {
if user_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Postgres>::new(
r#"
SELECT
user_group_members.user_id,
user_groups.id AS group_id,
user_groups.name AS group_name,
user_groups.priority AS group_priority,
user_group_members.created_at
FROM user_group_members
JOIN user_groups ON user_groups.id = user_group_members.group_id
WHERE user_group_members.user_id IN (
"#,
);
{
let mut separated = builder.separated(", ");
for user_id in user_ids {
separated.push_bind(user_id);
}
}
builder.push(
") ORDER BY user_group_members.user_id ASC, user_groups.name ASC, user_groups.id ASC",
);
collect_query_rows(
builder.build().fetch(&self.pool),
map_user_group_membership_row,
)
.await
}
pub async fn replace_user_groups_for_user(
&self,
user_id: &str,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut tx = self.pool.begin().await.map_postgres_err()?;
sqlx::query("DELETE FROM user_group_members WHERE user_id = $1")
.bind(user_id)
.execute(&mut *tx)
.await
.map_postgres_err()?;
for group_id in normalized_ids(group_ids) {
sqlx::query(
"INSERT INTO user_group_members (group_id, user_id) VALUES ($1, $2) ON CONFLICT (group_id, user_id) DO NOTHING",
)
.bind(group_id)
.bind(user_id)
.execute(&mut *tx)
.await
.map_postgres_err()?;
}
tx.commit().await.map_postgres_err()?;
self.list_user_groups_for_user(user_id).await
}
pub async fn add_user_to_group(
&self,
group_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
let result = sqlx::query(
"INSERT INTO user_group_members (group_id, user_id) VALUES ($1, $2) ON CONFLICT (group_id, user_id) DO NOTHING",
)
.bind(group_id)
.bind(user_id)
.execute(&self.pool)
.await
.map_postgres_err()?;
Ok(result.rows_affected() > 0)
}
pub async fn list_export_users_page(
&self,
query: &UserExportListQuery,
@@ -626,6 +967,16 @@ impl SqlxUserReadRepository {
if let Some(is_active) = query.is_active {
builder.push(" AND is_active = ").push_bind(is_active);
}
if let Some(group_id) = query
.group_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
builder.push(" AND id IN (SELECT user_id FROM user_group_members WHERE group_id = ");
builder.push_bind(group_id);
builder.push(")");
}
if let Some(search) = query
.search
.as_deref()
@@ -817,10 +1168,12 @@ impl SqlxUserReadRepository {
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
is_active, is_deleted, created_at, updated_at, last_login_at
)
VALUES (
$1, $2, TRUE, $3, NULL, 'user'::userrole, 'oauth'::authsource,
'inherit', 'inherit', 'inherit', 'inherit',
TRUE, FALSE, $4, $4, $4
)
"#,
@@ -963,8 +1316,9 @@ SET email = $2,
WHERE id = $1
RETURNING
id, email, email_verified, username, password_hash, role::text AS role,
auth_source::text AS auth_source, allowed_providers, allowed_api_formats,
allowed_models, is_active, is_deleted, created_at, last_login_at
auth_source::text AS auth_source, allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode, allowed_models, allowed_models_mode,
is_active, is_deleted, created_at, last_login_at
"#,
)
.bind(&existing.id)
@@ -1014,8 +1368,9 @@ INSERT INTO users (
VALUES ($1, $2, TRUE, $3, NULL, 'user'::userrole, 'ldap'::authsource, $4, $5, TRUE, FALSE, $6, $6, $6)
RETURNING
id, email, email_verified, username, password_hash, role::text AS role,
auth_source::text AS auth_source, allowed_providers, allowed_api_formats,
allowed_models, is_active, is_deleted, created_at, last_login_at
auth_source::text AS auth_source, allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode, allowed_models, allowed_models_mode,
is_active, is_deleted, created_at, last_login_at
"#,
)
.bind(uuid::Uuid::new_v4().to_string())
@@ -1127,6 +1482,26 @@ WHERE id = $1
rate_limit: Option<i32>,
is_active: Option<bool>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let allowed_providers_mode = if allowed_providers.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_api_formats_mode = if allowed_api_formats.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_models_mode = if allowed_models.is_some() {
"specific"
} else {
"unrestricted"
};
let rate_limit_mode = if rate_limit.is_some() {
"custom"
} else {
"system"
};
let result = sqlx::query(
r#"
UPDATE users
@@ -1138,20 +1513,36 @@ SET role = CASE
WHEN $4::BOOLEAN THEN $5::json
ELSE allowed_providers
END,
allowed_providers_mode = CASE
WHEN $4::BOOLEAN THEN $6
ELSE allowed_providers_mode
END,
allowed_api_formats = CASE
WHEN $6::BOOLEAN THEN $7::json
WHEN $7::BOOLEAN THEN $8::json
ELSE allowed_api_formats
END,
allowed_api_formats_mode = CASE
WHEN $7::BOOLEAN THEN $9
ELSE allowed_api_formats_mode
END,
allowed_models = CASE
WHEN $8::BOOLEAN THEN $9::json
WHEN $10::BOOLEAN THEN $11::json
ELSE allowed_models
END,
allowed_models_mode = CASE
WHEN $10::BOOLEAN THEN $12
ELSE allowed_models_mode
END,
rate_limit = CASE
WHEN $10::BOOLEAN THEN $11
WHEN $13::BOOLEAN THEN $14
ELSE rate_limit
END,
rate_limit_mode = CASE
WHEN $13::BOOLEAN THEN $15
ELSE rate_limit_mode
END,
is_active = CASE
WHEN $12::BOOLEAN AND $13 IS NOT NULL THEN $13
WHEN $16::BOOLEAN AND $17 IS NOT NULL THEN $17
ELSE is_active
END,
updated_at = NOW()
@@ -1163,12 +1554,16 @@ WHERE id = $1
.bind(role)
.bind(allowed_providers_present)
.bind(allowed_providers.map(serde_json::Value::from))
.bind(allowed_providers_mode)
.bind(allowed_api_formats_present)
.bind(allowed_api_formats.map(serde_json::Value::from))
.bind(allowed_api_formats_mode)
.bind(allowed_models_present)
.bind(allowed_models.map(serde_json::Value::from))
.bind(allowed_models_mode)
.bind(rate_limit_present)
.bind(rate_limit)
.bind(rate_limit_mode)
.bind(is_active.is_some())
.bind(is_active)
.execute(&self.pool)
@@ -1180,6 +1575,55 @@ WHERE id = $1
self.find_user_auth_by_id(user_id).await
}
pub async fn update_local_auth_user_policy_modes(
&self,
user_id: &str,
allowed_providers_mode: Option<String>,
allowed_api_formats_mode: Option<String>,
allowed_models_mode: Option<String>,
rate_limit_mode: Option<String>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let result = sqlx::query(
r#"
UPDATE users
SET allowed_providers_mode = CASE
WHEN $2::BOOLEAN THEN $3
ELSE allowed_providers_mode
END,
allowed_api_formats_mode = CASE
WHEN $4::BOOLEAN THEN $5
ELSE allowed_api_formats_mode
END,
allowed_models_mode = CASE
WHEN $6::BOOLEAN THEN $7
ELSE allowed_models_mode
END,
rate_limit_mode = CASE
WHEN $8::BOOLEAN THEN $9
ELSE rate_limit_mode
END,
updated_at = NOW()
WHERE id = $1
"#,
)
.bind(user_id)
.bind(allowed_providers_mode.is_some())
.bind(allowed_providers_mode)
.bind(allowed_api_formats_mode.is_some())
.bind(allowed_api_formats_mode)
.bind(allowed_models_mode.is_some())
.bind(allowed_models_mode)
.bind(rate_limit_mode.is_some())
.bind(rate_limit_mode)
.execute(&self.pool)
.await
.map_postgres_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.find_user_auth_by_id(user_id).await
}
pub async fn update_user_model_capability_settings(
&self,
user_id: &str,
@@ -1212,18 +1656,30 @@ WHERE id = $1
username: String,
password_hash: String,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
self.create_local_auth_user_with_settings(
email,
email_verified,
username,
password_hash,
"user".to_string(),
None,
None,
None,
None,
let user_id = uuid::Uuid::new_v4().to_string();
sqlx::query(
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
is_active, is_deleted, created_at, updated_at
)
VALUES (
$1, $2, $3, $4, $5, 'user'::userrole, 'local'::authsource,
'inherit', 'inherit', 'inherit', 'inherit',
TRUE, FALSE, NOW(), NOW()
)
"#,
)
.bind(&user_id)
.bind(email)
.bind(email_verified)
.bind(username)
.bind(password_hash)
.execute(&self.pool)
.await
.map_postgres_err()?;
self.find_user_auth_by_id(&user_id).await
}
#[allow(clippy::too_many_arguments)]
@@ -1240,16 +1696,39 @@ WHERE id = $1
rate_limit: Option<i32>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let user_id = uuid::Uuid::new_v4().to_string();
let allowed_providers_mode = if allowed_providers.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_api_formats_mode = if allowed_api_formats.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_models_mode = if allowed_models.is_some() {
"specific"
} else {
"unrestricted"
};
let rate_limit_mode = if rate_limit.is_some() {
"custom"
} else {
"system"
};
sqlx::query(
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers, allowed_api_formats, allowed_models, rate_limit,
allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode,
allowed_models, allowed_models_mode,
rate_limit, rate_limit_mode,
is_active, is_deleted, created_at, updated_at
)
VALUES (
$1, $2, $3, $4, $5, $6::userrole, 'local'::authsource,
$7::json, $8::json, $9::json, $10,
$7::json, $8, $9::json, $10, $11::json, $12, $13, $14,
TRUE, FALSE, NOW(), NOW()
)
"#,
@@ -1261,9 +1740,13 @@ VALUES (
.bind(password_hash)
.bind(role)
.bind(allowed_providers.map(serde_json::Value::from))
.bind(allowed_providers_mode)
.bind(allowed_api_formats.map(serde_json::Value::from))
.bind(allowed_api_formats_mode)
.bind(allowed_models.map(serde_json::Value::from))
.bind(allowed_models_mode)
.bind(rate_limit)
.bind(rate_limit_mode)
.execute(&self.pool)
.await
.map_postgres_err()?;
@@ -1541,6 +2024,16 @@ fn normalize_optional_json_value(value: Option<serde_json::Value>) -> Option<ser
}
}
fn normalized_ids(values: &[String]) -> Vec<String> {
values
.iter()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.collect::<std::collections::BTreeSet<_>>()
.into_iter()
.collect()
}
async fn find_postgres_ldap_user_for_update(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
ldap_dn: Option<&str>,
@@ -1550,8 +2043,9 @@ async fn find_postgres_ldap_user_for_update(
let select_columns = r#"
SELECT
id, email, email_verified, username, password_hash, role::text AS role,
auth_source::text AS auth_source, allowed_providers, allowed_api_formats,
allowed_models, is_active, is_deleted, created_at, last_login_at
auth_source::text AS auth_source, allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode, allowed_models, allowed_models_mode,
is_active, is_deleted, created_at, last_login_at
FROM users
"#;
if let Some(ldap_dn) = ldap_dn.filter(|value| !value.trim().is_empty()) {
@@ -1616,6 +2110,14 @@ fn map_user_export_row(row: &sqlx::postgres::PgRow) -> Result<StoredUserExportRo
.map_postgres_err()?,
row.try_get("is_active").map_postgres_err()?,
)
.and_then(|record| {
record.with_policy_modes(
row.try_get("allowed_providers_mode").map_postgres_err()?,
row.try_get("allowed_api_formats_mode").map_postgres_err()?,
row.try_get("allowed_models_mode").map_postgres_err()?,
row.try_get("rate_limit_mode").map_postgres_err()?,
)
})
}
fn map_user_auth_row(row: &sqlx::postgres::PgRow) -> Result<StoredUserAuthRecord, DataLayerError> {
@@ -1635,6 +2137,60 @@ fn map_user_auth_row(row: &sqlx::postgres::PgRow) -> Result<StoredUserAuthRecord
row.try_get("created_at").map_postgres_err()?,
row.try_get("last_login_at").map_postgres_err()?,
)
.and_then(|record| {
record.with_policy_modes(
row.try_get("allowed_providers_mode").map_postgres_err()?,
row.try_get("allowed_api_formats_mode").map_postgres_err()?,
row.try_get("allowed_models_mode").map_postgres_err()?,
)
})
}
fn map_user_group_row(row: &sqlx::postgres::PgRow) -> Result<StoredUserGroup, DataLayerError> {
StoredUserGroup::new(
row.try_get("id").map_postgres_err()?,
row.try_get("name").map_postgres_err()?,
row.try_get("normalized_name").map_postgres_err()?,
row.try_get("description").map_postgres_err()?,
row.try_get("priority").map_postgres_err()?,
row.try_get("allowed_providers").map_postgres_err()?,
row.try_get("allowed_providers_mode").map_postgres_err()?,
row.try_get("allowed_api_formats").map_postgres_err()?,
row.try_get("allowed_api_formats_mode").map_postgres_err()?,
row.try_get("allowed_models").map_postgres_err()?,
row.try_get("allowed_models_mode").map_postgres_err()?,
row.try_get("rate_limit").map_postgres_err()?,
row.try_get("rate_limit_mode").map_postgres_err()?,
row.try_get("created_at").map_postgres_err()?,
row.try_get("updated_at").map_postgres_err()?,
)
}
fn map_user_group_member_row(
row: &sqlx::postgres::PgRow,
) -> Result<StoredUserGroupMember, DataLayerError> {
Ok(StoredUserGroupMember {
group_id: row.try_get("group_id").map_postgres_err()?,
user_id: row.try_get("user_id").map_postgres_err()?,
username: row.try_get("username").map_postgres_err()?,
email: row.try_get("email").map_postgres_err()?,
role: row.try_get("role").map_postgres_err()?,
is_active: row.try_get("is_active").map_postgres_err()?,
is_deleted: row.try_get("is_deleted").map_postgres_err()?,
created_at: row.try_get("created_at").map_postgres_err()?,
})
}
fn map_user_group_membership_row(
row: &sqlx::postgres::PgRow,
) -> Result<StoredUserGroupMembership, DataLayerError> {
Ok(StoredUserGroupMembership {
user_id: row.try_get("user_id").map_postgres_err()?,
group_id: row.try_get("group_id").map_postgres_err()?,
group_name: row.try_get("group_name").map_postgres_err()?,
group_priority: row.try_get("group_priority").map_postgres_err()?,
created_at: row.try_get("created_at").map_postgres_err()?,
})
}
fn map_oauth_link_summary_row(
@@ -1709,6 +2265,88 @@ impl UserReadRepository for SqlxUserReadRepository {
self.find_export_user_by_id(user_id).await
}
async fn list_user_groups(&self) -> Result<Vec<StoredUserGroup>, DataLayerError> {
self.list_user_groups().await
}
async fn find_user_group_by_id(
&self,
group_id: &str,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
self.find_user_group_by_id(group_id).await
}
async fn list_user_groups_by_ids(
&self,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
self.list_user_groups_by_ids(group_ids).await
}
async fn create_user_group(
&self,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
self.create_user_group(record).await
}
async fn update_user_group(
&self,
group_id: &str,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
self.update_user_group(group_id, record).await
}
async fn delete_user_group(&self, group_id: &str) -> Result<bool, DataLayerError> {
self.delete_user_group(group_id).await
}
async fn list_user_group_members(
&self,
group_id: &str,
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
self.list_user_group_members(group_id).await
}
async fn replace_user_group_members(
&self,
group_id: &str,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
self.replace_user_group_members(group_id, user_ids).await
}
async fn list_user_groups_for_user(
&self,
user_id: &str,
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
self.list_user_groups_for_user(user_id).await
}
async fn list_user_group_memberships_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMembership>, DataLayerError> {
self.list_user_group_memberships_by_user_ids(user_ids).await
}
async fn replace_user_groups_for_user(
&self,
user_id: &str,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
self.replace_user_groups_for_user(user_id, group_ids).await
}
async fn add_user_to_group(
&self,
group_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
self.add_user_to_group(group_id, user_id).await
}
async fn find_user_auth_by_id(
&self,
user_id: &str,
@@ -1919,6 +2557,24 @@ impl UserReadRepository for SqlxUserReadRepository {
.await
}
async fn update_local_auth_user_policy_modes(
&self,
user_id: &str,
allowed_providers_mode: Option<String>,
allowed_api_formats_mode: Option<String>,
allowed_models_mode: Option<String>,
rate_limit_mode: Option<String>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
self.update_local_auth_user_policy_modes(
user_id,
allowed_providers_mode,
allowed_api_formats_mode,
allowed_models_mode,
rate_limit_mode,
)
.await
}
async fn update_user_model_capability_settings(
&self,
user_id: &str,

View File

@@ -3,9 +3,11 @@ use chrono::{DateTime, TimeZone, Utc};
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use super::types::{
LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, StoredUserExportRow,
normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord,
StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership,
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
StoredUserSummary, UserExportListQuery, UserExportSummary, UserReadRepository,
StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSummary,
UserReadRepository,
};
use crate::driver::sqlite::SqlitePool;
use crate::error::SqlResultExt;
@@ -32,9 +34,13 @@ SELECT
role,
auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
model_capability_settings,
is_active
FROM users
@@ -50,8 +56,11 @@ SELECT
role,
auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -69,8 +78,11 @@ SELECT
users.role AS role,
users.auth_source AS auth_source,
users.allowed_providers AS allowed_providers,
users.allowed_providers_mode AS allowed_providers_mode,
users.allowed_api_formats AS allowed_api_formats,
users.allowed_api_formats_mode AS allowed_api_formats_mode,
users.allowed_models AS allowed_models,
users.allowed_models_mode AS allowed_models_mode,
users.is_active AS is_active,
users.is_deleted AS is_deleted,
users.created_at AS created_at,
@@ -130,6 +142,40 @@ SELECT
FROM user_sessions
"#;
const USER_GROUP_COLUMNS: &str = r#"
SELECT
id,
name,
normalized_name,
description,
priority,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
created_at,
updated_at
FROM user_groups
"#;
const USER_GROUP_MEMBER_COLUMNS: &str = r#"
SELECT
user_group_members.group_id,
users.id AS user_id,
users.username,
users.email,
users.role,
users.is_active,
users.is_deleted,
user_group_members.created_at
FROM user_group_members
JOIN users ON users.id = user_group_members.user_id
"#;
#[derive(Debug, Clone)]
pub struct SqliteUserReadRepository {
pool: SqlitePool,
@@ -163,6 +209,22 @@ impl SqliteUserReadRepository {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_auth_row).collect()
}
async fn fetch_group_rows(
&self,
mut builder: QueryBuilder<'_, Sqlite>,
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_group_row).collect()
}
async fn fetch_group_member_rows(
&self,
mut builder: QueryBuilder<'_, Sqlite>,
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_group_member_row).collect()
}
}
#[async_trait]
@@ -224,6 +286,16 @@ impl UserReadRepository for SqliteUserReadRepository {
if let Some(is_active) = query.is_active {
builder.push(" AND is_active = ").push_bind(is_active);
}
if let Some(group_id) = query
.group_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
builder.push(" AND id IN (SELECT user_id FROM user_group_members WHERE group_id = ");
builder.push_bind(group_id);
builder.push(")");
}
if let Some(search) = query
.search
.as_deref()
@@ -296,6 +368,287 @@ WHERE is_deleted = 0
self.fetch_export_rows(builder).await
}
async fn list_user_groups(&self) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(USER_GROUP_COLUMNS);
builder.push(" ORDER BY name ASC, id ASC");
self.fetch_group_rows(builder).await
}
async fn find_user_group_by_id(
&self,
group_id: &str,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(USER_GROUP_COLUMNS);
builder
.push(" WHERE id = ")
.push_bind(group_id)
.push(" LIMIT 1");
Ok(self.fetch_group_rows(builder).await?.into_iter().next())
}
async fn list_user_groups_by_ids(
&self,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
if group_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Sqlite>::new(USER_GROUP_COLUMNS);
builder.push(" WHERE id IN (");
{
let mut separated = builder.separated(", ");
for group_id in group_ids {
separated.push_bind(group_id);
}
}
builder.push(") ORDER BY name ASC, id ASC");
self.fetch_group_rows(builder).await
}
async fn create_user_group(
&self,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let now = current_unix_secs();
let id = uuid::Uuid::new_v4().to_string();
let name = normalize_user_group_name(&record.name);
let normalized_name = name.to_ascii_lowercase();
let result = sqlx::query(
r#"
INSERT INTO user_groups (
id, name, normalized_name, description, priority,
allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode,
allowed_models, allowed_models_mode,
rate_limit, rate_limit_mode, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&id)
.bind(name)
.bind(normalized_name)
.bind(record.description)
.bind(record.priority)
.bind(json_string_from_option_vec(
record.allowed_providers.as_ref(),
))
.bind(record.allowed_providers_mode)
.bind(json_string_from_option_vec(
record.allowed_api_formats.as_ref(),
))
.bind(record.allowed_api_formats_mode)
.bind(json_string_from_option_vec(record.allowed_models.as_ref()))
.bind(record.allowed_models_mode)
.bind(record.rate_limit)
.bind(record.rate_limit_mode)
.bind(now)
.bind(now)
.execute(&self.pool)
.await;
match result {
Ok(_) => self.find_user_group_by_id(&id).await,
Err(sqlx::Error::Database(err)) if err.is_unique_violation() => Err(
DataLayerError::InvalidInput("duplicate user group name".to_string()),
),
Err(err) => Err(err).map_sql_err(),
}
}
async fn update_user_group(
&self,
group_id: &str,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let now = current_unix_secs();
let name = normalize_user_group_name(&record.name);
let normalized_name = name.to_ascii_lowercase();
let result = sqlx::query(
r#"
UPDATE user_groups
SET name = ?,
normalized_name = ?,
description = ?,
priority = ?,
allowed_providers = ?,
allowed_providers_mode = ?,
allowed_api_formats = ?,
allowed_api_formats_mode = ?,
allowed_models = ?,
allowed_models_mode = ?,
rate_limit = ?,
rate_limit_mode = ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(name)
.bind(normalized_name)
.bind(record.description)
.bind(record.priority)
.bind(json_string_from_option_vec(
record.allowed_providers.as_ref(),
))
.bind(record.allowed_providers_mode)
.bind(json_string_from_option_vec(
record.allowed_api_formats.as_ref(),
))
.bind(record.allowed_api_formats_mode)
.bind(json_string_from_option_vec(record.allowed_models.as_ref()))
.bind(record.allowed_models_mode)
.bind(record.rate_limit)
.bind(record.rate_limit_mode)
.bind(now)
.bind(group_id)
.execute(&self.pool)
.await;
match result {
Ok(result) if result.rows_affected() == 0 => Ok(None),
Ok(_) => self.find_user_group_by_id(group_id).await,
Err(sqlx::Error::Database(err)) if err.is_unique_violation() => Err(
DataLayerError::InvalidInput("duplicate user group name".to_string()),
),
Err(err) => Err(err).map_sql_err(),
}
}
async fn delete_user_group(&self, group_id: &str) -> Result<bool, DataLayerError> {
let result = sqlx::query("DELETE FROM user_groups WHERE id = ?")
.bind(group_id)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
async fn list_user_group_members(
&self,
group_id: &str,
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(USER_GROUP_MEMBER_COLUMNS);
builder
.push(" WHERE user_group_members.group_id = ")
.push_bind(group_id)
.push(" ORDER BY users.username ASC, users.id ASC");
self.fetch_group_member_rows(builder).await
}
async fn replace_user_group_members(
&self,
group_id: &str,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("DELETE FROM user_group_members WHERE group_id = ?")
.bind(group_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
let now = current_unix_secs();
for user_id in normalized_ids(user_ids) {
sqlx::query(
"INSERT OR IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)",
)
.bind(group_id)
.bind(user_id)
.bind(now)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
tx.commit().await.map_sql_err()?;
self.list_user_group_members(group_id).await
}
async fn list_user_groups_for_user(
&self,
user_id: &str,
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(USER_GROUP_COLUMNS);
builder
.push(" WHERE id IN (SELECT group_id FROM user_group_members WHERE user_id = ")
.push_bind(user_id)
.push(") ORDER BY name ASC, id ASC");
self.fetch_group_rows(builder).await
}
async fn list_user_group_memberships_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMembership>, DataLayerError> {
if user_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Sqlite>::new(
r#"
SELECT
user_group_members.user_id,
user_groups.id AS group_id,
user_groups.name AS group_name,
user_groups.priority AS group_priority,
user_group_members.created_at
FROM user_group_members
JOIN user_groups ON user_groups.id = user_group_members.group_id
WHERE user_group_members.user_id IN (
"#,
);
{
let mut separated = builder.separated(", ");
for user_id in user_ids {
separated.push_bind(user_id);
}
}
builder.push(
") ORDER BY user_group_members.user_id ASC, user_groups.name ASC, user_groups.id ASC",
);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_group_membership_row).collect()
}
async fn replace_user_groups_for_user(
&self,
user_id: &str,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("DELETE FROM user_group_members WHERE user_id = ?")
.bind(user_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
let now = current_unix_secs();
for group_id in normalized_ids(group_ids) {
sqlx::query(
"INSERT OR IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)",
)
.bind(group_id)
.bind(user_id)
.bind(now)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
tx.commit().await.map_sql_err()?;
self.list_user_groups_for_user(user_id).await
}
async fn add_user_to_group(
&self,
group_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
let result = sqlx::query(
"INSERT OR IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)",
)
.bind(group_id)
.bind(user_id)
.bind(current_unix_secs())
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
async fn find_user_auth_by_id(
&self,
user_id: &str,
@@ -453,9 +806,10 @@ WHERE provider_type = ?
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
is_active, is_deleted, created_at, updated_at, last_login_at
)
VALUES (?, ?, 1, ?, NULL, 'user', 'oauth', 1, 0, ?, ?, ?)
VALUES (?, ?, 1, ?, NULL, 'user', 'oauth', 'inherit', 'inherit', 'inherit', 'inherit', 1, 0, ?, ?, ?)
"#,
)
.bind(&user_id)
@@ -675,14 +1029,38 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
rate_limit: Option<i32>,
is_active: Option<bool>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let allowed_providers_mode = if allowed_providers.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_api_formats_mode = if allowed_api_formats.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_models_mode = if allowed_models.is_some() {
"specific"
} else {
"unrestricted"
};
let rate_limit_mode = if rate_limit.is_some() {
"custom"
} else {
"system"
};
let result = sqlx::query(
r#"
UPDATE users
SET role = CASE WHEN ? THEN COALESCE(?, role) ELSE role END,
allowed_providers = CASE WHEN ? THEN ? ELSE allowed_providers END,
allowed_providers_mode = CASE WHEN ? THEN ? ELSE allowed_providers_mode END,
allowed_api_formats = CASE WHEN ? THEN ? ELSE allowed_api_formats END,
allowed_api_formats_mode = CASE WHEN ? THEN ? ELSE allowed_api_formats_mode END,
allowed_models = CASE WHEN ? THEN ? ELSE allowed_models END,
allowed_models_mode = CASE WHEN ? THEN ? ELSE allowed_models_mode END,
rate_limit = CASE WHEN ? THEN ? ELSE rate_limit END,
rate_limit_mode = CASE WHEN ? THEN ? ELSE rate_limit_mode END,
is_active = CASE WHEN ? THEN ? ELSE is_active END,
updated_at = ?
WHERE id = ?
@@ -695,18 +1073,26 @@ WHERE id = ?
allowed_providers,
"users.allowed_providers",
)?)
.bind(allowed_providers_present)
.bind(allowed_providers_mode)
.bind(allowed_api_formats_present)
.bind(optional_string_list_json(
allowed_api_formats,
"users.allowed_api_formats",
)?)
.bind(allowed_api_formats_present)
.bind(allowed_api_formats_mode)
.bind(allowed_models_present)
.bind(optional_string_list_json(
allowed_models,
"users.allowed_models",
)?)
.bind(allowed_models_present)
.bind(allowed_models_mode)
.bind(rate_limit_present)
.bind(rate_limit)
.bind(rate_limit_present)
.bind(rate_limit_mode)
.bind(is_active.is_some())
.bind(is_active)
.bind(chrono::Utc::now().timestamp())
@@ -720,6 +1106,44 @@ WHERE id = ?
self.find_user_auth_by_id(user_id).await
}
async fn update_local_auth_user_policy_modes(
&self,
user_id: &str,
allowed_providers_mode: Option<String>,
allowed_api_formats_mode: Option<String>,
allowed_models_mode: Option<String>,
rate_limit_mode: Option<String>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let result = sqlx::query(
r#"
UPDATE users
SET allowed_providers_mode = CASE WHEN ? THEN ? ELSE allowed_providers_mode END,
allowed_api_formats_mode = CASE WHEN ? THEN ? ELSE allowed_api_formats_mode END,
allowed_models_mode = CASE WHEN ? THEN ? ELSE allowed_models_mode END,
rate_limit_mode = CASE WHEN ? THEN ? ELSE rate_limit_mode END,
updated_at = ?
WHERE id = ?
"#,
)
.bind(allowed_providers_mode.is_some())
.bind(allowed_providers_mode)
.bind(allowed_api_formats_mode.is_some())
.bind(allowed_api_formats_mode)
.bind(allowed_models_mode.is_some())
.bind(allowed_models_mode)
.bind(rate_limit_mode.is_some())
.bind(rate_limit_mode)
.bind(chrono::Utc::now().timestamp())
.bind(user_id)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.find_user_auth_by_id(user_id).await
}
async fn update_user_model_capability_settings(
&self,
user_id: &str,
@@ -751,18 +1175,29 @@ WHERE id = ?
username: String,
password_hash: String,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
self.create_local_auth_user_with_settings(
email,
email_verified,
username,
password_hash,
"user".to_string(),
None,
None,
None,
None,
let user_id = uuid::Uuid::new_v4().to_string();
let now = chrono::Utc::now().timestamp();
sqlx::query(
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
is_active, is_deleted, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, 'user', 'local', 'inherit', 'inherit', 'inherit', 'inherit', 1, 0, ?, ?)
"#,
)
.bind(&user_id)
.bind(email)
.bind(email_verified)
.bind(username)
.bind(password_hash)
.bind(now)
.bind(now)
.execute(&self.pool)
.await
.map_sql_err()?;
self.find_user_auth_by_id(&user_id).await
}
async fn create_local_auth_user_with_settings(
@@ -779,14 +1214,37 @@ WHERE id = ?
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let user_id = uuid::Uuid::new_v4().to_string();
let now = chrono::Utc::now().timestamp();
let allowed_providers_mode = if allowed_providers.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_api_formats_mode = if allowed_api_formats.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_models_mode = if allowed_models.is_some() {
"specific"
} else {
"unrestricted"
};
let rate_limit_mode = if rate_limit.is_some() {
"custom"
} else {
"system"
};
sqlx::query(
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers, allowed_api_formats, allowed_models, rate_limit,
allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode,
allowed_models, allowed_models_mode,
rate_limit, rate_limit_mode,
is_active, is_deleted, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, 1, 0, ?, ?)
VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, ?, ?, ?, ?, 1, 0, ?, ?)
"#,
)
.bind(&user_id)
@@ -799,15 +1257,19 @@ VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, 1, 0, ?, ?)
allowed_providers,
"users.allowed_providers",
)?)
.bind(allowed_providers_mode)
.bind(optional_string_list_json(
allowed_api_formats,
"users.allowed_api_formats",
)?)
.bind(allowed_api_formats_mode)
.bind(optional_string_list_json(
allowed_models,
"users.allowed_models",
)?)
.bind(allowed_models_mode)
.bind(rate_limit)
.bind(rate_limit_mode)
.bind(now)
.bind(now)
.execute(&self.pool)
@@ -1166,6 +1628,24 @@ fn optional_string_list_json(
.transpose()
}
fn json_string_from_option_vec(value: Option<&Vec<String>>) -> Option<String> {
value.and_then(|items| serde_json::to_string(items).ok())
}
fn normalized_ids(values: &[String]) -> Vec<String> {
values
.iter()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.collect::<std::collections::BTreeSet<_>>()
.into_iter()
.collect()
}
fn current_unix_secs() -> i64 {
chrono::Utc::now().timestamp()
}
fn optional_json_string(
value: Option<serde_json::Value>,
field_name: &str,
@@ -1382,6 +1862,14 @@ fn map_user_export_row(row: &SqliteRow) -> Result<StoredUserExportRow, DataLayer
)?,
row.try_get("is_active").map_sql_err()?,
)
.and_then(|record| {
record.with_policy_modes(
row.try_get("allowed_providers_mode").map_sql_err()?,
row.try_get("allowed_api_formats_mode").map_sql_err()?,
row.try_get("allowed_models_mode").map_sql_err()?,
row.try_get("rate_limit_mode").map_sql_err()?,
)
})
}
fn map_user_auth_row(row: &SqliteRow) -> Result<StoredUserAuthRecord, DataLayerError> {
@@ -1410,6 +1898,67 @@ fn map_user_auth_row(row: &SqliteRow) -> Result<StoredUserAuthRecord, DataLayerE
optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
optional_datetime_from_unix_secs(row.try_get("last_login_at").map_sql_err()?),
)
.and_then(|record| {
record.with_policy_modes(
row.try_get("allowed_providers_mode").map_sql_err()?,
row.try_get("allowed_api_formats_mode").map_sql_err()?,
row.try_get("allowed_models_mode").map_sql_err()?,
)
})
}
fn map_user_group_row(row: &SqliteRow) -> Result<StoredUserGroup, DataLayerError> {
StoredUserGroup::new(
row.try_get("id").map_sql_err()?,
row.try_get("name").map_sql_err()?,
row.try_get("normalized_name").map_sql_err()?,
row.try_get("description").map_sql_err()?,
row.try_get("priority").map_sql_err()?,
optional_json_from_string(
row.try_get("allowed_providers").map_sql_err()?,
"user_groups.allowed_providers",
)?,
row.try_get("allowed_providers_mode").map_sql_err()?,
optional_json_from_string(
row.try_get("allowed_api_formats").map_sql_err()?,
"user_groups.allowed_api_formats",
)?,
row.try_get("allowed_api_formats_mode").map_sql_err()?,
optional_json_from_string(
row.try_get("allowed_models").map_sql_err()?,
"user_groups.allowed_models",
)?,
row.try_get("allowed_models_mode").map_sql_err()?,
row.try_get("rate_limit").map_sql_err()?,
row.try_get("rate_limit_mode").map_sql_err()?,
optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
optional_datetime_from_unix_secs(row.try_get("updated_at").map_sql_err()?),
)
}
fn map_user_group_member_row(row: &SqliteRow) -> Result<StoredUserGroupMember, DataLayerError> {
Ok(StoredUserGroupMember {
group_id: row.try_get("group_id").map_sql_err()?,
user_id: row.try_get("user_id").map_sql_err()?,
username: row.try_get("username").map_sql_err()?,
email: row.try_get("email").map_sql_err()?,
role: row.try_get("role").map_sql_err()?,
is_active: row.try_get("is_active").map_sql_err()?,
is_deleted: row.try_get("is_deleted").map_sql_err()?,
created_at: optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
})
}
fn map_user_group_membership_row(
row: &SqliteRow,
) -> Result<StoredUserGroupMembership, DataLayerError> {
Ok(StoredUserGroupMembership {
user_id: row.try_get("user_id").map_sql_err()?,
group_id: row.try_get("group_id").map_sql_err()?,
group_name: row.try_get("group_name").map_sql_err()?,
group_priority: row.try_get("group_priority").map_sql_err()?,
created_at: optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
})
}
fn map_oauth_link_summary_row(
@@ -1560,6 +2109,7 @@ INSERT INTO users (
role: Some("user".to_string()),
is_active: Some(true),
search: None,
group_id: None,
})
.await
.expect("export page should load");

View File

@@ -57,8 +57,11 @@ pub struct StoredUserAuthRecord {
pub role: String,
pub auth_source: String,
pub allowed_providers: Option<Vec<String>>,
pub allowed_providers_mode: String,
pub allowed_api_formats: Option<Vec<String>>,
pub allowed_api_formats_mode: String,
pub allowed_models: Option<Vec<String>>,
pub allowed_models_mode: String,
pub is_active: bool,
pub is_deleted: bool,
pub created_at: Option<DateTime<Utc>>,
@@ -113,16 +116,44 @@ impl StoredUserAuthRecord {
role,
auth_source,
allowed_providers: parse_string_list(allowed_providers, "users.allowed_providers")?,
allowed_providers_mode: "unrestricted".to_string(),
allowed_api_formats: parse_string_list(
allowed_api_formats,
"users.allowed_api_formats",
)?,
allowed_api_formats_mode: "unrestricted".to_string(),
allowed_models: parse_string_list(allowed_models, "users.allowed_models")?,
allowed_models_mode: "unrestricted".to_string(),
is_active,
is_deleted,
created_at,
last_login_at,
})
.map(|record| record.with_legacy_policy_modes())
}
pub fn with_policy_modes(
mut self,
allowed_providers_mode: String,
allowed_api_formats_mode: String,
allowed_models_mode: String,
) -> Result<Self, crate::DataLayerError> {
self.allowed_providers_mode =
normalize_list_policy_mode(&allowed_providers_mode, "users.allowed_providers_mode")?;
self.allowed_api_formats_mode = normalize_list_policy_mode(
&allowed_api_formats_mode,
"users.allowed_api_formats_mode",
)?;
self.allowed_models_mode =
normalize_list_policy_mode(&allowed_models_mode, "users.allowed_models_mode")?;
Ok(self)
}
fn with_legacy_policy_modes(mut self) -> Self {
self.allowed_providers_mode = legacy_list_policy_mode(&self.allowed_providers);
self.allowed_api_formats_mode = legacy_list_policy_mode(&self.allowed_api_formats);
self.allowed_models_mode = legacy_list_policy_mode(&self.allowed_models);
self
}
pub fn to_summary(&self) -> Result<StoredUserSummary, crate::DataLayerError> {
@@ -197,9 +228,13 @@ pub struct StoredUserExportRow {
pub role: String,
pub auth_source: String,
pub allowed_providers: Option<Vec<String>>,
pub allowed_providers_mode: String,
pub allowed_api_formats: Option<Vec<String>>,
pub allowed_api_formats_mode: String,
pub allowed_models: Option<Vec<String>>,
pub allowed_models_mode: String,
pub rate_limit: Option<i32>,
pub rate_limit_mode: String,
pub model_capability_settings: Option<Value>,
pub is_active: bool,
}
@@ -251,15 +286,52 @@ impl StoredUserExportRow {
role,
auth_source,
allowed_providers: parse_string_list(allowed_providers, "users.allowed_providers")?,
allowed_providers_mode: "unrestricted".to_string(),
allowed_api_formats: parse_string_list(
allowed_api_formats,
"users.allowed_api_formats",
)?,
allowed_api_formats_mode: "unrestricted".to_string(),
allowed_models: parse_string_list(allowed_models, "users.allowed_models")?,
allowed_models_mode: "unrestricted".to_string(),
rate_limit,
rate_limit_mode: "system".to_string(),
model_capability_settings: normalize_optional_json(model_capability_settings),
is_active,
})
.map(|record| record.with_legacy_policy_modes())
}
pub fn with_policy_modes(
mut self,
allowed_providers_mode: String,
allowed_api_formats_mode: String,
allowed_models_mode: String,
rate_limit_mode: String,
) -> Result<Self, crate::DataLayerError> {
self.allowed_providers_mode =
normalize_list_policy_mode(&allowed_providers_mode, "users.allowed_providers_mode")?;
self.allowed_api_formats_mode = normalize_list_policy_mode(
&allowed_api_formats_mode,
"users.allowed_api_formats_mode",
)?;
self.allowed_models_mode =
normalize_list_policy_mode(&allowed_models_mode, "users.allowed_models_mode")?;
self.rate_limit_mode =
normalize_rate_limit_policy_mode(&rate_limit_mode, "users.rate_limit_mode")?;
Ok(self)
}
fn with_legacy_policy_modes(mut self) -> Self {
self.allowed_providers_mode = legacy_list_policy_mode(&self.allowed_providers);
self.allowed_api_formats_mode = legacy_list_policy_mode(&self.allowed_api_formats);
self.allowed_models_mode = legacy_list_policy_mode(&self.allowed_models);
self.rate_limit_mode = if self.rate_limit.is_some() {
"custom".to_string()
} else {
"system".to_string()
};
self
}
}
@@ -404,6 +476,139 @@ pub struct StoredUserPreferenceRecord {
pub announcement_notifications: bool,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredUserGroup {
pub id: String,
pub name: String,
pub normalized_name: String,
pub description: Option<String>,
pub priority: i32,
pub allowed_providers: Option<Vec<String>>,
pub allowed_providers_mode: String,
pub allowed_api_formats: Option<Vec<String>>,
pub allowed_api_formats_mode: String,
pub allowed_models: Option<Vec<String>>,
pub allowed_models_mode: String,
pub rate_limit: Option<i32>,
pub rate_limit_mode: String,
pub created_at: Option<DateTime<Utc>>,
pub updated_at: Option<DateTime<Utc>>,
}
impl StoredUserGroup {
#[allow(clippy::too_many_arguments)]
pub fn new(
id: String,
name: String,
normalized_name: String,
description: Option<String>,
priority: i32,
allowed_providers: Option<Value>,
allowed_providers_mode: String,
allowed_api_formats: Option<Value>,
allowed_api_formats_mode: String,
allowed_models: Option<Value>,
allowed_models_mode: String,
rate_limit: Option<i32>,
rate_limit_mode: String,
created_at: Option<DateTime<Utc>>,
updated_at: Option<DateTime<Utc>>,
) -> Result<Self, crate::DataLayerError> {
if id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"user_groups.id is empty".to_string(),
));
}
if name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"user_groups.name is empty".to_string(),
));
}
if normalized_name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"user_groups.normalized_name is empty".to_string(),
));
}
Ok(Self {
id,
name,
normalized_name,
description,
priority,
allowed_providers: parse_string_list(
allowed_providers,
"user_groups.allowed_providers",
)?,
allowed_providers_mode: normalize_list_policy_mode(
&allowed_providers_mode,
"user_groups.allowed_providers_mode",
)?,
allowed_api_formats: parse_string_list(
allowed_api_formats,
"user_groups.allowed_api_formats",
)?,
allowed_api_formats_mode: normalize_list_policy_mode(
&allowed_api_formats_mode,
"user_groups.allowed_api_formats_mode",
)?,
allowed_models: parse_string_list(allowed_models, "user_groups.allowed_models")?,
allowed_models_mode: normalize_list_policy_mode(
&allowed_models_mode,
"user_groups.allowed_models_mode",
)?,
rate_limit,
rate_limit_mode: normalize_rate_limit_policy_mode(
&rate_limit_mode,
"user_groups.rate_limit_mode",
)?,
created_at,
updated_at,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredUserGroupMember {
pub group_id: String,
pub user_id: String,
pub username: String,
pub email: Option<String>,
pub role: String,
pub is_active: bool,
pub is_deleted: bool,
pub created_at: Option<DateTime<Utc>>,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredUserGroupMembership {
pub user_id: String,
pub group_id: String,
pub group_name: String,
pub group_priority: i32,
pub created_at: Option<DateTime<Utc>>,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct UpsertUserGroupRecord {
pub name: String,
pub description: Option<String>,
pub priority: i32,
pub allowed_providers: Option<Vec<String>>,
pub allowed_providers_mode: String,
pub allowed_api_formats: Option<Vec<String>>,
pub allowed_api_formats_mode: String,
pub allowed_models: Option<Vec<String>>,
pub allowed_models_mode: String,
pub rate_limit: Option<i32>,
pub rate_limit_mode: String,
}
impl UpsertUserGroupRecord {
pub fn normalized_name(&self) -> String {
normalize_user_group_name(&self.name).to_ascii_lowercase()
}
}
impl StoredUserPreferenceRecord {
pub fn default_for_user(user_id: impl Into<String>) -> Self {
Self {
@@ -429,6 +634,7 @@ pub struct UserExportListQuery {
pub role: Option<String>,
pub is_active: Option<bool>,
pub search: Option<String>,
pub group_id: Option<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
@@ -463,6 +669,64 @@ pub trait UserReadRepository: Send + Sync {
user_id: &str,
) -> Result<Option<StoredUserExportRow>, crate::DataLayerError>;
async fn list_user_groups(&self) -> Result<Vec<StoredUserGroup>, crate::DataLayerError>;
async fn find_user_group_by_id(
&self,
group_id: &str,
) -> Result<Option<StoredUserGroup>, crate::DataLayerError>;
async fn list_user_groups_by_ids(
&self,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, crate::DataLayerError>;
async fn create_user_group(
&self,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, crate::DataLayerError>;
async fn update_user_group(
&self,
group_id: &str,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, crate::DataLayerError>;
async fn delete_user_group(&self, group_id: &str) -> Result<bool, crate::DataLayerError>;
async fn list_user_group_members(
&self,
group_id: &str,
) -> Result<Vec<StoredUserGroupMember>, crate::DataLayerError>;
async fn replace_user_group_members(
&self,
group_id: &str,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMember>, crate::DataLayerError>;
async fn list_user_groups_for_user(
&self,
user_id: &str,
) -> Result<Vec<StoredUserGroup>, crate::DataLayerError>;
async fn list_user_group_memberships_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMembership>, crate::DataLayerError>;
async fn replace_user_groups_for_user(
&self,
user_id: &str,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, crate::DataLayerError>;
async fn add_user_to_group(
&self,
group_id: &str,
user_id: &str,
) -> Result<bool, crate::DataLayerError>;
async fn list_non_admin_export_users(
&self,
) -> Result<Vec<StoredUserExportRow>, crate::DataLayerError>;
@@ -602,6 +866,15 @@ pub trait UserReadRepository: Send + Sync {
is_active: Option<bool>,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
async fn update_local_auth_user_policy_modes(
&self,
user_id: &str,
allowed_providers_mode: Option<String>,
allowed_api_formats_mode: Option<String>,
allowed_models_mode: Option<String>,
rate_limit_mode: Option<String>,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
async fn update_user_model_capability_settings(
&self,
user_id: &str,
@@ -717,6 +990,47 @@ fn normalize_optional_json(value: Option<Value>) -> Option<Value> {
}
}
pub fn normalize_user_group_name(value: &str) -> String {
value.split_whitespace().collect::<Vec<_>>().join(" ")
}
pub fn normalize_list_policy_mode(
value: &str,
field_name: &str,
) -> Result<String, crate::DataLayerError> {
match value.trim().to_ascii_lowercase().as_str() {
"inherit" => Ok("inherit".to_string()),
"unrestricted" => Ok("unrestricted".to_string()),
"specific" => Ok("specific".to_string()),
"deny_all" => Ok("deny_all".to_string()),
_ => Err(crate::DataLayerError::UnexpectedValue(format!(
"{field_name} is not a valid list policy mode"
))),
}
}
pub fn normalize_rate_limit_policy_mode(
value: &str,
field_name: &str,
) -> Result<String, crate::DataLayerError> {
match value.trim().to_ascii_lowercase().as_str() {
"inherit" => Ok("inherit".to_string()),
"system" => Ok("system".to_string()),
"custom" => Ok("custom".to_string()),
_ => Err(crate::DataLayerError::UnexpectedValue(format!(
"{field_name} is not a valid rate limit policy mode"
))),
}
}
fn legacy_list_policy_mode(values: &Option<Vec<String>>) -> String {
if values.is_some() {
"specific".to_string()
} else {
"unrestricted".to_string()
}
}
fn parse_string_list(
value: Option<Value>,
field_name: &str,