Merge remote-tracking branch 'entropy-xu/codex/user-groups-default-permissions' into aether-rust-pioneer

This commit is contained in:
fawney19
2026-05-10 11:26:29 +08:00
49 changed files with 6373 additions and 250 deletions

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,95 @@ 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()
}
#[async_trait]
impl UserReadRepository for InMemoryUserReadRepository {
async fn list_users_by_ids(
@@ -329,6 +428,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 +491,273 @@ 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| {
right
.priority
.cmp(&left.priority)
.then_with(|| 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(|| right.group_priority.cmp(&left.group_priority))
.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 +964,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 +1305,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 +1351,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 +1483,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 +1565,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 +2550,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,285 @@ 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 priority DESC, 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 priority DESC, 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 priority DESC, 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.priority DESC, 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 +804,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 +1027,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 +1071,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 +1104,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 +1173,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 +1212,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 +1255,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 +1622,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 +1856,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 +1892,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,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,

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,285 @@ 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 priority DESC, 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 priority DESC, 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 priority DESC, 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.priority DESC, 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 +804,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 +1027,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 +1071,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 +1104,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 +1173,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 +1212,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 +1255,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 +1626,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 +1860,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 +1896,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 +2107,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,