mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-07 01:47:47 +08:00
feat: add user groups and inherited access policies
This commit is contained in:
@@ -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,275 @@ 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 priority DESC, 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 priority DESC, 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 priority DESC, 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.priority DESC, 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 +965,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 +1166,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 +1314,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 +1366,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 +1480,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 +1511,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 +1552,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 +1573,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 +1654,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 +1694,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 +1738,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 +2022,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 +2041,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 +2108,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 +2135,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 +2263,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 +2555,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,
|
||||
|
||||
Reference in New Issue
Block a user