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

@@ -271,6 +271,8 @@ pub(crate) fn build_local_auth_rejection_response(
control_decision: Option<&GatewayControlDecision>, control_decision: Option<&GatewayControlDecision>,
rejection: &GatewayLocalAuthRejection, rejection: &GatewayLocalAuthRejection,
) -> Result<Response<Body>, GatewayError> { ) -> Result<Response<Body>, GatewayError> {
const ACCESS_POLICY_SUBJECT: &str = "当前用户、用户组或密钥的访问控制策略";
match rejection { match rejection {
GatewayLocalAuthRejection::InvalidApiKey => build_local_http_error_response( GatewayLocalAuthRejection::InvalidApiKey => build_local_http_error_response(
trace_id, trace_id,
@@ -298,7 +300,7 @@ pub(crate) fn build_local_auth_rejection_response(
trace_id, trace_id,
control_decision, control_decision,
StatusCode::FORBIDDEN, StatusCode::FORBIDDEN,
&format!("当前密钥不允许访问 {provider} 提供商"), &format!("{ACCESS_POLICY_SUBJECT}不允许访问 {provider} 提供商"),
) )
} }
GatewayLocalAuthRejection::ApiFormatNotAllowed { api_format } => { GatewayLocalAuthRejection::ApiFormatNotAllowed { api_format } => {
@@ -306,14 +308,14 @@ pub(crate) fn build_local_auth_rejection_response(
trace_id, trace_id,
control_decision, control_decision,
StatusCode::FORBIDDEN, StatusCode::FORBIDDEN,
&format!("当前密钥不允许访问 {api_format} 格式"), &format!("{ACCESS_POLICY_SUBJECT}不允许访问 {api_format} 格式"),
) )
} }
GatewayLocalAuthRejection::ModelNotAllowed { model } => build_local_http_error_response( GatewayLocalAuthRejection::ModelNotAllowed { model } => build_local_http_error_response(
trace_id, trace_id,
control_decision, control_decision,
StatusCode::FORBIDDEN, StatusCode::FORBIDDEN,
&format!("当前密钥不允许访问模型 {model}"), &format!("{ACCESS_POLICY_SUBJECT}不允许访问模型 {model}"),
), ),
} }
} }

View File

@@ -575,6 +575,95 @@ pub(super) fn classify_admin_operations_family_route(
"admin:wallets", "admin:wallets",
false, false,
)) ))
} else if method == http::Method::GET
&& matches!(
normalized_path,
"/api/admin/user-groups" | "/api/admin/user-groups/"
)
{
Some(classified(
"admin_proxy",
"users_manage",
"list_user_groups",
"admin:users",
false,
))
} else if method == http::Method::POST
&& matches!(
normalized_path,
"/api/admin/user-groups" | "/api/admin/user-groups/"
)
{
Some(classified(
"admin_proxy",
"users_manage",
"create_user_group",
"admin:users",
false,
))
} else if method == http::Method::PUT
&& matches!(
normalized_path,
"/api/admin/user-groups/default" | "/api/admin/user-groups/default/"
)
{
Some(classified(
"admin_proxy",
"users_manage",
"set_default_user_group",
"admin:users",
false,
))
} else if method == http::Method::GET
&& normalized_path.starts_with("/api/admin/user-groups/")
&& normalized_path.ends_with("/members")
&& normalized_path.matches('/').count() == 5
{
Some(classified(
"admin_proxy",
"users_manage",
"list_user_group_members",
"admin:users",
false,
))
} else if method == http::Method::PUT
&& normalized_path.starts_with("/api/admin/user-groups/")
&& normalized_path.ends_with("/members")
&& normalized_path.matches('/').count() == 5
{
Some(classified(
"admin_proxy",
"users_manage",
"replace_user_group_members",
"admin:users",
false,
))
} else if method == http::Method::PUT
&& normalized_path.starts_with("/api/admin/user-groups/")
&& normalized_path.matches('/').count() == 4
&& !normalized_path.ends_with("/default")
&& !normalized_path.ends_with("/members")
{
Some(classified(
"admin_proxy",
"users_manage",
"update_user_group",
"admin:users",
false,
))
} else if method == http::Method::DELETE
&& normalized_path.starts_with("/api/admin/user-groups/")
&& normalized_path.matches('/').count() == 4
&& !normalized_path.ends_with("/default")
&& !normalized_path.ends_with("/members")
{
Some(classified(
"admin_proxy",
"users_manage",
"delete_user_group",
"admin:users",
false,
))
} else if method == http::Method::GET } else if method == http::Method::GET
&& matches!(normalized_path, "/api/admin/users" | "/api/admin/users/") && matches!(normalized_path, "/api/admin/users" | "/api/admin/users/")
{ {

View File

@@ -1,6 +1,8 @@
use http::Uri; use http::Uri;
use super::{classify_control_route, headers}; use crate::handlers::shared::local_proxy_route_requires_buffered_body;
use super::{classify_control_route, headers, GatewayPublicRequestContext};
#[test] #[test]
fn classifies_admin_users_list_as_admin_proxy_route() { fn classifies_admin_users_list_as_admin_proxy_route() {
@@ -68,6 +70,86 @@ fn classifies_admin_user_batch_routes_as_admin_proxy_route() {
); );
} }
#[test]
fn classifies_admin_user_group_routes_as_admin_proxy_route() {
let headers = headers(&[]);
let list_uri: Uri = "/api/admin/user-groups".parse().expect("uri should parse");
let list = classify_control_route(&http::Method::GET, &list_uri, &headers)
.expect("route should classify");
assert_eq!(list.route_family.as_deref(), Some("users_manage"));
assert_eq!(list.route_kind.as_deref(), Some("list_user_groups"));
let create_uri: Uri = "/api/admin/user-groups".parse().expect("uri should parse");
let create = classify_control_route(&http::Method::POST, &create_uri, &headers)
.expect("route should classify");
assert_eq!(create.route_family.as_deref(), Some("users_manage"));
assert_eq!(create.route_kind.as_deref(), Some("create_user_group"));
let update_uri: Uri = "/api/admin/user-groups/group-1"
.parse()
.expect("uri should parse");
let update = classify_control_route(&http::Method::PUT, &update_uri, &headers)
.expect("route should classify");
assert_eq!(update.route_family.as_deref(), Some("users_manage"));
assert_eq!(update.route_kind.as_deref(), Some("update_user_group"));
let members_uri: Uri = "/api/admin/user-groups/group-1/members"
.parse()
.expect("uri should parse");
let members = classify_control_route(&http::Method::PUT, &members_uri, &headers)
.expect("route should classify");
assert_eq!(members.route_family.as_deref(), Some("users_manage"));
assert_eq!(
members.route_kind.as_deref(),
Some("replace_user_group_members")
);
let default_uri: Uri = "/api/admin/user-groups/default"
.parse()
.expect("uri should parse");
let default = classify_control_route(&http::Method::PUT, &default_uri, &headers)
.expect("route should classify");
assert_eq!(default.route_family.as_deref(), Some("users_manage"));
assert_eq!(
default.route_kind.as_deref(),
Some("set_default_user_group")
);
assert_eq!(
default.auth_endpoint_signature.as_deref(),
Some("admin:users")
);
}
#[test]
fn admin_user_group_write_routes_buffer_request_body() {
let headers = headers(&[]);
let routes = [
(http::Method::POST, "/api/admin/user-groups"),
(http::Method::PUT, "/api/admin/user-groups/group-1"),
(http::Method::PUT, "/api/admin/user-groups/group-1/members"),
(http::Method::PUT, "/api/admin/user-groups/default"),
];
for (method, path) in routes {
let uri: Uri = path.parse().expect("uri should parse");
let decision =
classify_control_route(&method, &uri, &headers).expect("route should classify");
let context = GatewayPublicRequestContext::from_request_parts(
"trace-user-group-write",
&method,
&uri,
&headers,
Some(decision),
);
assert!(
local_proxy_route_requires_buffered_body(&context),
"{method} {path} should buffer request body"
);
}
}
#[test] #[test]
fn classifies_admin_user_detail_routes_as_admin_proxy_route() { fn classifies_admin_user_detail_routes_as_admin_proxy_route() {
let headers = headers(&[]); let headers = headers(&[]);

View File

@@ -86,6 +86,139 @@ impl GatewayDataState {
} }
} }
pub(crate) async fn list_user_groups(
&self,
) -> Result<Vec<aether_data::repository::users::StoredUserGroup>, DataLayerError> {
match &self.user_reader {
Some(repository) => repository.list_user_groups().await,
None => Ok(Vec::new()),
}
}
pub(crate) async fn find_user_group_by_id(
&self,
group_id: &str,
) -> Result<Option<aether_data::repository::users::StoredUserGroup>, DataLayerError> {
match &self.user_reader {
Some(repository) => repository.find_user_group_by_id(group_id).await,
None => Ok(None),
}
}
pub(crate) async fn list_user_groups_by_ids(
&self,
group_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroup>, DataLayerError> {
match &self.user_reader {
Some(repository) => repository.list_user_groups_by_ids(group_ids).await,
None => Ok(Vec::new()),
}
}
pub(crate) async fn create_user_group(
&self,
record: aether_data::repository::users::UpsertUserGroupRecord,
) -> Result<Option<aether_data::repository::users::StoredUserGroup>, DataLayerError> {
match &self.user_reader {
Some(repository) => repository.create_user_group(record).await,
None => Ok(None),
}
}
pub(crate) async fn update_user_group(
&self,
group_id: &str,
record: aether_data::repository::users::UpsertUserGroupRecord,
) -> Result<Option<aether_data::repository::users::StoredUserGroup>, DataLayerError> {
match &self.user_reader {
Some(repository) => repository.update_user_group(group_id, record).await,
None => Ok(None),
}
}
pub(crate) async fn delete_user_group(&self, group_id: &str) -> Result<bool, DataLayerError> {
match &self.user_reader {
Some(repository) => repository.delete_user_group(group_id).await,
None => Ok(false),
}
}
pub(crate) async fn list_user_group_members(
&self,
group_id: &str,
) -> Result<Vec<aether_data::repository::users::StoredUserGroupMember>, DataLayerError> {
match &self.user_reader {
Some(repository) => repository.list_user_group_members(group_id).await,
None => Ok(Vec::new()),
}
}
pub(crate) async fn replace_user_group_members(
&self,
group_id: &str,
user_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroupMember>, DataLayerError> {
match &self.user_reader {
Some(repository) => {
repository
.replace_user_group_members(group_id, user_ids)
.await
}
None => Ok(Vec::new()),
}
}
pub(crate) async fn list_user_groups_for_user(
&self,
user_id: &str,
) -> Result<Vec<aether_data::repository::users::StoredUserGroup>, DataLayerError> {
match &self.user_reader {
Some(repository) => repository.list_user_groups_for_user(user_id).await,
None => Ok(Vec::new()),
}
}
pub(crate) async fn list_user_group_memberships_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroupMembership>, DataLayerError>
{
match &self.user_reader {
Some(repository) => {
repository
.list_user_group_memberships_by_user_ids(user_ids)
.await
}
None => Ok(Vec::new()),
}
}
pub(crate) async fn replace_user_groups_for_user(
&self,
user_id: &str,
group_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroup>, DataLayerError> {
match &self.user_reader {
Some(repository) => {
repository
.replace_user_groups_for_user(user_id, group_ids)
.await
}
None => Ok(Vec::new()),
}
}
pub(crate) async fn add_user_to_group(
&self,
group_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
match &self.user_reader {
Some(repository) => repository.add_user_to_group(group_id, user_id).await,
None => Ok(false),
}
}
pub(crate) async fn list_user_oauth_links( pub(crate) async fn list_user_oauth_links(
&self, &self,
user_id: &str, user_id: &str,
@@ -424,6 +557,28 @@ impl GatewayDataState {
.await .await
} }
pub(crate) 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 Some(repository) = self.user_reader.as_ref() else {
return Ok(None);
};
repository
.update_local_auth_user_policy_modes(
user_id,
allowed_providers_mode,
allowed_api_formats_mode,
allowed_models_mode,
rate_limit_mode,
)
.await
}
pub(crate) async fn touch_auth_user_last_login( pub(crate) async fn touch_auth_user_last_login(
&self, &self,
user_id: &str, user_id: &str,
@@ -1470,13 +1625,14 @@ impl GatewayDataState {
api_key_id: &str, api_key_id: &str,
now_unix_secs: u64, now_unix_secs: u64,
) -> Result<Option<GatewayAuthApiKeySnapshot>, DataLayerError> { ) -> Result<Option<GatewayAuthApiKeySnapshot>, DataLayerError> {
read_resolved_auth_api_key_snapshot_by_user_api_key_ids( let snapshot = read_resolved_auth_api_key_snapshot_by_user_api_key_ids(
self, self,
user_id, user_id,
api_key_id, api_key_id,
now_unix_secs, now_unix_secs,
) )
.await .await?;
self.apply_user_group_effective_policies(snapshot).await
} }
pub(crate) async fn read_auth_api_key_snapshot_by_key_hash( pub(crate) async fn read_auth_api_key_snapshot_by_key_hash(
@@ -1484,7 +1640,148 @@ impl GatewayDataState {
key_hash: &str, key_hash: &str,
now_unix_secs: u64, now_unix_secs: u64,
) -> Result<Option<GatewayAuthApiKeySnapshot>, DataLayerError> { ) -> Result<Option<GatewayAuthApiKeySnapshot>, DataLayerError> {
read_resolved_auth_api_key_snapshot_by_key_hash(self, key_hash, now_unix_secs).await let snapshot =
read_resolved_auth_api_key_snapshot_by_key_hash(self, key_hash, now_unix_secs).await?;
self.apply_user_group_effective_policies(snapshot).await
}
async fn apply_user_group_effective_policies(
&self,
snapshot: Option<GatewayAuthApiKeySnapshot>,
) -> Result<Option<GatewayAuthApiKeySnapshot>, DataLayerError> {
let Some(mut snapshot) = snapshot else {
return Ok(None);
};
let Some(repository) = self.user_reader.as_ref() else {
return Ok(Some(snapshot));
};
let Some(user) = repository.find_user_auth_by_id(&snapshot.user_id).await? else {
return Ok(Some(snapshot));
};
let export_row = repository.find_export_user_by_id(&snapshot.user_id).await?;
let mut groups = repository
.list_user_groups_for_user(&snapshot.user_id)
.await?;
groups.sort_by(|left, right| {
right
.priority
.cmp(&left.priority)
.then_with(|| left.name.cmp(&right.name))
.then_with(|| left.id.cmp(&right.id))
});
let allowed_providers = resolve_effective_list_policy(
user.allowed_providers,
&user.allowed_providers_mode,
&groups,
|group| {
(
&group.allowed_providers_mode,
group.allowed_providers.clone(),
)
},
);
let allowed_api_formats = resolve_effective_list_policy(
user.allowed_api_formats,
&user.allowed_api_formats_mode,
&groups,
|group| {
(
&group.allowed_api_formats_mode,
group.allowed_api_formats.clone(),
)
},
);
let allowed_models = resolve_effective_list_policy(
user.allowed_models,
&user.allowed_models_mode,
&groups,
|group| (&group.allowed_models_mode, group.allowed_models.clone()),
);
let snapshot_user_rate_limit = snapshot.user_rate_limit;
let export_user_rate_limit = export_row.as_ref().and_then(|row| row.rate_limit);
let user_rate_limit_mode = match export_row.as_ref() {
Some(row)
if row.rate_limit.is_none()
&& row.rate_limit_mode == "system"
&& snapshot_user_rate_limit.is_some() =>
{
"custom"
}
Some(row) => row.rate_limit_mode.as_str(),
None if snapshot_user_rate_limit.is_some() => "custom",
None => "system",
};
let user_rate_limit = resolve_effective_rate_limit_policy(
export_user_rate_limit.or(snapshot_user_rate_limit),
user_rate_limit_mode,
&groups,
);
snapshot.apply_user_policy(
allowed_providers,
allowed_api_formats,
allowed_models,
user_rate_limit,
);
Ok(Some(snapshot))
}
}
fn resolve_effective_list_policy(
user_values: Option<Vec<String>>,
user_mode: &str,
groups: &[aether_data::repository::users::StoredUserGroup],
group_field: impl Fn(
&aether_data::repository::users::StoredUserGroup,
) -> (&str, Option<Vec<String>>),
) -> Option<Vec<String>> {
match user_mode {
"unrestricted" => None,
"specific" => Some(user_values.unwrap_or_default()),
"deny_all" => Some(Vec::new()),
"inherit" => groups
.iter()
.find_map(|group| {
let (mode, values) = group_field(group);
match mode {
"unrestricted" => Some(None),
"specific" => Some(Some(values.unwrap_or_default())),
"deny_all" => Some(Some(Vec::new())),
_ => None,
}
})
.flatten_or_unrestricted(),
_ => None,
}
}
trait FlattenPolicyOption<T> {
fn flatten_or_unrestricted(self) -> Option<T>;
}
impl<T> FlattenPolicyOption<T> for Option<Option<T>> {
fn flatten_or_unrestricted(self) -> Option<T> {
self.unwrap_or_default()
}
}
fn resolve_effective_rate_limit_policy(
user_rate_limit: Option<i32>,
user_mode: &str,
groups: &[aether_data::repository::users::StoredUserGroup],
) -> Option<i32> {
match user_mode {
"custom" => Some(user_rate_limit.unwrap_or(0)),
"system" => None,
"inherit" => groups
.iter()
.find_map(|group| match group.rate_limit_mode.as_str() {
"custom" => Some(Some(group.rate_limit.unwrap_or(0))),
"system" => Some(None),
_ => None,
})
.flatten_or_unrestricted(),
_ => None,
} }
} }

View File

@@ -63,6 +63,88 @@ impl<'a> AdminAppState<'a> {
self.app.find_user_auth_by_identifier(identifier).await self.app.find_user_auth_by_identifier(identifier).await
} }
pub(crate) async fn list_user_groups(
&self,
) -> Result<Vec<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.app.list_user_groups().await
}
pub(crate) async fn find_user_group_by_id(
&self,
group_id: &str,
) -> Result<Option<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.app.find_user_group_by_id(group_id).await
}
pub(crate) async fn list_user_groups_by_ids(
&self,
group_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.app.list_user_groups_by_ids(group_ids).await
}
pub(crate) async fn create_user_group(
&self,
record: aether_data::repository::users::UpsertUserGroupRecord,
) -> Result<Option<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.app.create_user_group(record).await
}
pub(crate) async fn update_user_group(
&self,
group_id: &str,
record: aether_data::repository::users::UpsertUserGroupRecord,
) -> Result<Option<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.app.update_user_group(group_id, record).await
}
pub(crate) async fn delete_user_group(&self, group_id: &str) -> Result<bool, GatewayError> {
self.app.delete_user_group(group_id).await
}
pub(crate) async fn list_user_group_members(
&self,
group_id: &str,
) -> Result<Vec<aether_data::repository::users::StoredUserGroupMember>, GatewayError> {
self.app.list_user_group_members(group_id).await
}
pub(crate) async fn replace_user_group_members(
&self,
group_id: &str,
user_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroupMember>, GatewayError> {
self.app
.replace_user_group_members(group_id, user_ids)
.await
}
pub(crate) async fn list_user_groups_for_user(
&self,
user_id: &str,
) -> Result<Vec<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.app.list_user_groups_for_user(user_id).await
}
pub(crate) async fn list_user_group_memberships_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroupMembership>, GatewayError> {
self.app
.list_user_group_memberships_by_user_ids(user_ids)
.await
}
pub(crate) async fn replace_user_groups_for_user(
&self,
user_id: &str,
group_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.app
.replace_user_groups_for_user(user_id, group_ids)
.await
}
pub(crate) async fn is_other_user_auth_email_taken( pub(crate) async fn is_other_user_auth_email_taken(
&self, &self,
email: &str, email: &str,
@@ -187,6 +269,25 @@ impl<'a> AdminAppState<'a> {
.await .await
} }
pub(crate) 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<aether_data::repository::users::StoredUserAuthRecord>, GatewayError> {
self.app
.update_local_auth_user_policy_modes(
user_id,
allowed_providers_mode,
allowed_api_formats_mode,
allowed_models_mode,
rate_limit_mode,
)
.await
}
pub(crate) async fn update_auth_user_wallet_limit_mode( pub(crate) async fn update_auth_user_wallet_limit_mode(
&self, &self,
user_id: &str, user_id: &str,

View File

@@ -22,11 +22,14 @@ struct AdminUserSelectionFilters {
role: Option<String>, role: Option<String>,
#[serde(default)] #[serde(default)]
is_active: Option<bool>, is_active: Option<bool>,
#[serde(default)]
group_id: Option<String>,
} }
#[derive(Debug, Clone, Default)] #[derive(Debug, Clone, Default)]
struct AdminUserSelectionRequest { struct AdminUserSelectionRequest {
user_ids: Vec<String>, user_ids: Vec<String>,
group_ids: Vec<String>,
filters: Option<AdminUserSelectionFilters>, filters: Option<AdminUserSelectionFilters>,
filters_scope_present: bool, filters_scope_present: bool,
} }
@@ -51,6 +54,7 @@ struct NormalizedAdminUserSelectionFilters {
search: Option<String>, search: Option<String>,
role: Option<String>, role: Option<String>,
is_active: Option<bool>, is_active: Option<bool>,
group_id: Option<String>,
} }
#[derive(Debug, Clone, serde::Serialize)] #[derive(Debug, Clone, serde::Serialize)]
@@ -60,12 +64,22 @@ struct AdminUserSelectionItem {
email: Option<String>, email: Option<String>,
role: String, role: String,
is_active: bool, is_active: bool,
matched_by: Vec<String>,
}
#[derive(Debug, Clone, serde::Serialize)]
struct AdminUserSelectionWarning {
#[serde(rename = "type")]
warning_type: String,
group_id: Option<String>,
message: String,
} }
#[derive(Debug, Clone, Default)] #[derive(Debug, Clone, Default)]
struct ResolvedAdminUserSelection { struct ResolvedAdminUserSelection {
items: Vec<AdminUserSelectionItem>, items: Vec<AdminUserSelectionItem>,
missing_user_ids: Vec<String>, missing_user_ids: Vec<String>,
warnings: Vec<AdminUserSelectionWarning>,
} }
#[derive(Debug, Clone, Default)] #[derive(Debug, Clone, Default)]
@@ -112,6 +126,7 @@ pub(in super::super) async fn build_admin_resolve_user_selection_response(
Ok(Json(json!({ Ok(Json(json!({
"total": resolved.items.len(), "total": resolved.items.len(),
"items": resolved.items, "items": resolved.items,
"warnings": resolved.warnings,
})) }))
.into_response()) .into_response())
} }
@@ -229,6 +244,7 @@ pub(in super::super) async fn build_admin_user_batch_action_response(
"success": success, "success": success,
"failed": failed, "failed": failed,
"failures": failures, "failures": failures,
"warnings": resolved.warnings,
"action": request.action.trim().to_ascii_lowercase(), "action": request.action.trim().to_ascii_lowercase(),
"modified_fields": mutation.modified_fields, "modified_fields": mutation.modified_fields,
})) }))
@@ -284,6 +300,11 @@ fn parse_selection_request_value(value: Value) -> Result<AdminUserSelectionReque
Some(value) => serde_json::from_value::<Vec<String>>(value.clone()) Some(value) => serde_json::from_value::<Vec<String>>(value.clone())
.map_err(|_| "user_ids 必须是字符串数组".to_string())?, .map_err(|_| "user_ids 必须是字符串数组".to_string())?,
}; };
let group_ids = match map.get("group_ids") {
None | Some(Value::Null) => Vec::new(),
Some(value) => serde_json::from_value::<Vec<String>>(value.clone())
.map_err(|_| "group_ids 必须是字符串数组".to_string())?,
};
let (filters_scope_present, filters) = match map.get("filters") { let (filters_scope_present, filters) = match map.get("filters") {
Some(Value::Object(_)) => { Some(Value::Object(_)) => {
@@ -299,6 +320,7 @@ fn parse_selection_request_value(value: Value) -> Result<AdminUserSelectionReque
Ok(AdminUserSelectionRequest { Ok(AdminUserSelectionRequest {
user_ids, user_ids,
group_ids,
filters, filters,
filters_scope_present, filters_scope_present,
}) })
@@ -310,12 +332,36 @@ async fn resolve_admin_user_selection(
) -> Result<ResolvedAdminUserSelection, String> { ) -> Result<ResolvedAdminUserSelection, String> {
let filters = normalize_selection_filters(selection.filters)?; let filters = normalize_selection_filters(selection.filters)?;
let explicit_user_ids = normalize_user_ids(selection.user_ids); let explicit_user_ids = normalize_user_ids(selection.user_ids);
if explicit_user_ids.is_empty() && !selection.filters_scope_present { let explicit_group_ids = normalize_user_ids(selection.group_ids);
return Err("至少需要选择一个用户或明确提供筛选条件".to_string()); if explicit_user_ids.is_empty()
&& explicit_group_ids.is_empty()
&& !selection.filters_scope_present
{
return Err("至少需要选择一个用户、用户组或明确提供筛选条件".to_string());
} }
let should_resolve_filters = selection.filters_scope_present; let should_resolve_filters = selection.filters_scope_present;
let mut items_by_id = BTreeMap::new(); let mut items_by_id = BTreeMap::new();
let mut missing_user_ids = Vec::new(); let mut missing_user_ids = Vec::new();
let mut warnings = Vec::new();
if !explicit_group_ids.is_empty() {
let groups = state
.list_user_groups_by_ids(&explicit_group_ids)
.await
.map_err(|_| "用户分组数据不可用".to_string())?;
let found_group_ids = groups
.iter()
.map(|group| group.id.clone())
.collect::<BTreeSet<_>>();
let missing_group_ids = explicit_group_ids
.iter()
.filter(|group_id| !found_group_ids.contains(*group_id))
.cloned()
.collect::<Vec<_>>();
if !missing_group_ids.is_empty() {
return Err(format!("用户分组不存在: {}", missing_group_ids.join(", ")));
}
}
if !explicit_user_ids.is_empty() { if !explicit_user_ids.is_empty() {
let users = state let users = state
@@ -325,15 +371,14 @@ async fn resolve_admin_user_selection(
for user_id in explicit_user_ids { for user_id in explicit_user_ids {
match users.get(&user_id).filter(|user| !user.is_deleted) { match users.get(&user_id).filter(|user| !user.is_deleted) {
Some(user) => { Some(user) => {
items_by_id.insert( insert_or_update_selection_item(
&mut items_by_id,
user.id.clone(), user.id.clone(),
AdminUserSelectionItem { user.username.clone(),
user_id: user.id.clone(), user.email.clone(),
username: user.username.clone(), user.role.clone(),
email: user.email.clone(), user.is_active,
role: user.role.clone(), "direct".to_string(),
is_active: user.is_active,
},
); );
} }
None => missing_user_ids.push(user_id), None => missing_user_ids.push(user_id),
@@ -341,24 +386,71 @@ async fn resolve_admin_user_selection(
} }
} }
if should_resolve_filters { for group_id in &explicit_group_ids {
let users = state let members = state
.list_export_users() .list_user_group_members(group_id)
.await .await
.map_err(|_| "用户数据不可用".to_string())?; .map_err(|_| "用户分组成员数据不可用".to_string())?;
let mut matched_count = 0usize;
for member in members.into_iter().filter(|member| !member.is_deleted) {
matched_count += 1;
insert_or_update_selection_item(
&mut items_by_id,
member.user_id,
member.username,
member.email,
member.role,
member.is_active,
format!("group:{group_id}"),
);
}
if matched_count == 0 {
warnings.push(AdminUserSelectionWarning {
warning_type: "empty_group".to_string(),
group_id: Some(group_id.clone()),
message: "分组内没有可操作用户".to_string(),
});
}
}
if should_resolve_filters {
let users = if filters.as_ref().is_some_and(|filters| {
filters.search.is_some()
|| filters.role.is_some()
|| filters.is_active.is_some()
|| filters.group_id.is_some()
}) {
state
.list_export_users_page(&aether_data::repository::users::UserExportListQuery {
skip: 0,
limit: 100_000,
role: filters.as_ref().and_then(|filters| filters.role.clone()),
is_active: filters.as_ref().and_then(|filters| filters.is_active),
search: filters.as_ref().and_then(|filters| filters.search.clone()),
group_id: filters
.as_ref()
.and_then(|filters| filters.group_id.clone()),
})
.await
.map_err(|_| "用户数据不可用".to_string())?
} else {
state
.list_export_users()
.await
.map_err(|_| "用户数据不可用".to_string())?
};
for user in users for user in users
.into_iter() .into_iter()
.filter(|user| admin_user_matches_filters(user, filters.as_ref())) .filter(|user| admin_user_matches_filters(user, filters.as_ref()))
{ {
items_by_id.insert( insert_or_update_selection_item(
user.id.clone(), &mut items_by_id,
AdminUserSelectionItem { user.id,
user_id: user.id, user.username,
username: user.username, user.email,
email: user.email, user.role,
role: user.role, user.is_active,
is_active: user.is_active, "filter".to_string(),
},
); );
} }
} }
@@ -374,9 +466,41 @@ async fn resolve_admin_user_selection(
Ok(ResolvedAdminUserSelection { Ok(ResolvedAdminUserSelection {
items, items,
missing_user_ids, missing_user_ids,
warnings,
}) })
} }
fn insert_or_update_selection_item(
items_by_id: &mut BTreeMap<String, AdminUserSelectionItem>,
user_id: String,
username: String,
email: Option<String>,
role: String,
is_active: bool,
matched_by: String,
) {
match items_by_id.get_mut(&user_id) {
Some(item) => {
if !item.matched_by.iter().any(|value| value == &matched_by) {
item.matched_by.push(matched_by);
}
}
None => {
items_by_id.insert(
user_id.clone(),
AdminUserSelectionItem {
user_id,
username,
email,
role,
is_active,
matched_by: vec![matched_by],
},
);
}
}
}
fn normalize_selection_filters( fn normalize_selection_filters(
filters: Option<AdminUserSelectionFilters>, filters: Option<AdminUserSelectionFilters>,
) -> Result<Option<NormalizedAdminUserSelectionFilters>, String> { ) -> Result<Option<NormalizedAdminUserSelectionFilters>, String> {
@@ -401,6 +525,10 @@ fn normalize_selection_filters(
search, search,
role, role,
is_active: filters.is_active, is_active: filters.is_active,
group_id: filters
.group_id
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty()),
})) }))
} }

View File

@@ -0,0 +1,457 @@
use super::{
build_admin_users_bad_request_response, build_admin_users_read_only_response,
format_optional_datetime_iso8601, normalize_admin_user_api_formats,
normalize_admin_user_string_list,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::GatewayError;
use axum::{
body::Body,
http,
response::{IntoResponse, Response},
Json,
};
use serde_json::json;
const DEFAULT_USER_GROUP_CONFIG_KEY: &str = "default_user_group_id";
#[derive(Debug, serde::Deserialize)]
struct AdminUserGroupPayload {
name: String,
#[serde(default)]
description: Option<String>,
#[serde(default)]
priority: Option<i32>,
#[serde(default)]
allowed_providers: Option<Vec<String>>,
#[serde(default = "default_list_mode")]
allowed_providers_mode: String,
#[serde(default)]
allowed_api_formats: Option<Vec<String>>,
#[serde(default = "default_list_mode")]
allowed_api_formats_mode: String,
#[serde(default)]
allowed_models: Option<Vec<String>>,
#[serde(default = "default_list_mode")]
allowed_models_mode: String,
#[serde(default)]
rate_limit: Option<i32>,
#[serde(default = "default_rate_limit_mode")]
rate_limit_mode: String,
}
#[derive(Debug, serde::Deserialize)]
struct AdminUserGroupMembersPayload {
user_ids: Vec<String>,
}
#[derive(Debug, serde::Deserialize)]
struct AdminDefaultUserGroupPayload {
#[serde(default)]
group_id: Option<String>,
}
pub(in super::super) async fn build_admin_list_user_groups_response(
state: &AdminAppState<'_>,
) -> Result<Response<Body>, GatewayError> {
let default_group_id = read_default_user_group_id(state).await?;
let items = state
.list_user_groups()
.await?
.into_iter()
.map(|group| user_group_payload(group, default_group_id.as_deref()))
.collect::<Vec<_>>();
Ok(Json(json!({
"items": items,
"default_group_id": default_group_id,
}))
.into_response())
}
pub(in super::super) async fn build_admin_create_user_group_response(
state: &AdminAppState<'_>,
request_body: Option<&axum::body::Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_auth_user_write_capability() {
return Ok(build_admin_users_read_only_response(
"当前为只读模式,无法创建用户分组",
));
}
let record = match parse_group_record(request_body) {
Ok(value) => value,
Err(detail) => return Ok(bad_request_owned(detail)),
};
let group = match state.create_user_group(record).await {
Ok(Some(group)) => group,
Ok(None) => {
return Ok(build_admin_users_read_only_response(
"当前为只读模式,无法创建用户分组",
))
}
Err(err) if is_duplicate_group_name_error(&err) => {
return Ok(bad_request_owned("用户分组名称已存在".to_string()))
}
Err(err) => return Err(err),
};
let default_group_id = read_default_user_group_id(state).await?;
Ok(attach_admin_audit_response(
Json(user_group_payload(group, default_group_id.as_deref())).into_response(),
"admin_user_group_created",
"create_user_group",
"user_group",
"user_groups",
))
}
pub(in super::super) async fn build_admin_update_user_group_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&axum::body::Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_auth_user_write_capability() {
return Ok(build_admin_users_read_only_response(
"当前为只读模式,无法更新用户分组",
));
}
let Some(group_id) = user_group_id_from_path(request_context.path()) else {
return Ok(build_admin_users_bad_request_response("缺少 group_id"));
};
let record = match parse_group_record(request_body) {
Ok(value) => value,
Err(detail) => return Ok(bad_request_owned(detail)),
};
let group = match state.update_user_group(&group_id, record).await {
Ok(Some(group)) => group,
Ok(None) => return Ok(not_found("用户分组不存在")),
Err(err) if is_duplicate_group_name_error(&err) => {
return Ok(bad_request_owned("用户分组名称已存在".to_string()))
}
Err(err) => return Err(err),
};
let default_group_id = read_default_user_group_id(state).await?;
Ok(attach_admin_audit_response(
Json(user_group_payload(group, default_group_id.as_deref())).into_response(),
"admin_user_group_updated",
"update_user_group",
"user_group",
&group_id,
))
}
pub(in super::super) async fn build_admin_delete_user_group_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_auth_user_write_capability() {
return Ok(build_admin_users_read_only_response(
"当前为只读模式,无法删除用户分组",
));
}
let Some(group_id) = user_group_id_from_path(request_context.path()) else {
return Ok(build_admin_users_bad_request_response("缺少 group_id"));
};
if read_default_user_group_id(state).await?.as_deref() == Some(group_id.as_str()) {
state
.delete_system_config_value(DEFAULT_USER_GROUP_CONFIG_KEY)
.await?;
}
if !state.delete_user_group(&group_id).await? {
return Ok(not_found("用户分组不存在"));
}
Ok(attach_admin_audit_response(
Json(json!({ "deleted": true })).into_response(),
"admin_user_group_deleted",
"delete_user_group",
"user_group",
&group_id,
))
}
pub(in super::super) async fn build_admin_list_user_group_members_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
) -> Result<Response<Body>, GatewayError> {
let Some(group_id) = user_group_member_group_id_from_path(request_context.path()) else {
return Ok(build_admin_users_bad_request_response("缺少 group_id"));
};
if state.find_user_group_by_id(&group_id).await?.is_none() {
return Ok(not_found("用户分组不存在"));
}
let items = state
.list_user_group_members(&group_id)
.await?
.into_iter()
.map(|member| {
json!({
"group_id": member.group_id,
"user_id": member.user_id,
"username": member.username,
"email": member.email,
"role": member.role,
"is_active": member.is_active,
"is_deleted": member.is_deleted,
"created_at": format_optional_datetime_iso8601(member.created_at),
})
})
.collect::<Vec<_>>();
Ok(Json(json!({ "items": items })).into_response())
}
pub(in super::super) async fn build_admin_replace_user_group_members_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,
request_body: Option<&axum::body::Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_auth_user_write_capability() {
return Ok(build_admin_users_read_only_response(
"当前为只读模式,无法更新分组成员",
));
}
let Some(group_id) = user_group_member_group_id_from_path(request_context.path()) else {
return Ok(build_admin_users_bad_request_response("缺少 group_id"));
};
if state.find_user_group_by_id(&group_id).await?.is_none() {
return Ok(not_found("用户分组不存在"));
}
let payload = match parse_members_payload(request_body) {
Ok(value) => value,
Err(detail) => return Ok(bad_request_owned(detail)),
};
let user_ids = normalize_ids(payload.user_ids);
let known_users = state.resolve_auth_user_summaries_by_ids(&user_ids).await?;
if known_users.len() != user_ids.len() {
return Ok(bad_request_owned("成员包含不存在的用户".to_string()));
}
let items = state
.replace_user_group_members(&group_id, &user_ids)
.await?;
Ok(attach_admin_audit_response(
Json(json!({
"items": items.into_iter().map(|member| json!({
"group_id": member.group_id,
"user_id": member.user_id,
"username": member.username,
"email": member.email,
"role": member.role,
"is_active": member.is_active,
"is_deleted": member.is_deleted,
"created_at": format_optional_datetime_iso8601(member.created_at),
})).collect::<Vec<_>>()
}))
.into_response(),
"admin_user_group_members_updated",
"update_user_group_members",
"user_group",
&group_id,
))
}
pub(in super::super) async fn build_admin_set_default_user_group_response(
state: &AdminAppState<'_>,
request_body: Option<&axum::body::Bytes>,
) -> Result<Response<Body>, GatewayError> {
if !state.has_auth_user_write_capability() {
return Ok(build_admin_users_read_only_response(
"当前为只读模式,无法设置默认用户组",
));
}
let payload = match request_body {
Some(body) if !body.is_empty() => {
serde_json::from_slice::<AdminDefaultUserGroupPayload>(body)
.map_err(|_| "请求数据验证失败".to_string())
}
_ => Err("请求数据验证失败".to_string()),
};
let payload = match payload {
Ok(value) => value,
Err(detail) => return Ok(bad_request_owned(detail)),
};
let group_id = payload
.group_id
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty());
if let Some(group_id) = group_id.as_deref() {
if state.find_user_group_by_id(group_id).await?.is_none() {
return Ok(bad_request_owned("默认用户组不存在".to_string()));
}
state
.upsert_system_config_json_value(
DEFAULT_USER_GROUP_CONFIG_KEY,
&json!(group_id),
Some("Default group for self-registered users"),
)
.await?;
} else {
state
.delete_system_config_value(DEFAULT_USER_GROUP_CONFIG_KEY)
.await?;
}
Ok(Json(json!({ "default_group_id": group_id })).into_response())
}
pub(crate) async fn read_default_user_group_id(
state: &AdminAppState<'_>,
) -> Result<Option<String>, GatewayError> {
Ok(state
.read_system_config_json_value(DEFAULT_USER_GROUP_CONFIG_KEY)
.await?
.and_then(|value| value.as_str().map(str::to_string))
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty()))
}
fn parse_group_record(
request_body: Option<&axum::body::Bytes>,
) -> Result<aether_data::repository::users::UpsertUserGroupRecord, String> {
let Some(body) = request_body.filter(|body| !body.is_empty()) else {
return Err("请求数据验证失败".to_string());
};
let payload = serde_json::from_slice::<AdminUserGroupPayload>(body)
.map_err(|_| "请求数据验证失败".to_string())?;
let name = aether_data::repository::users::normalize_user_group_name(&payload.name);
if name.is_empty() {
return Err("分组名称不能为空".to_string());
}
if payload.rate_limit.is_some_and(|value| value < 0) {
return Err("rate_limit 必须大于等于 0".to_string());
}
let allowed_providers =
normalize_admin_user_string_list(payload.allowed_providers, "allowed_providers")?;
let allowed_api_formats = normalize_admin_user_api_formats(payload.allowed_api_formats)?;
let allowed_models =
normalize_admin_user_string_list(payload.allowed_models, "allowed_models")?;
Ok(aether_data::repository::users::UpsertUserGroupRecord {
name,
description: payload
.description
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty()),
priority: payload.priority.unwrap_or_default(),
allowed_providers,
allowed_providers_mode: normalize_list_mode(&payload.allowed_providers_mode)?,
allowed_api_formats,
allowed_api_formats_mode: normalize_list_mode(&payload.allowed_api_formats_mode)?,
allowed_models,
allowed_models_mode: normalize_list_mode(&payload.allowed_models_mode)?,
rate_limit: payload.rate_limit,
rate_limit_mode: normalize_rate_mode(&payload.rate_limit_mode)?,
})
}
fn parse_members_payload(
request_body: Option<&axum::body::Bytes>,
) -> Result<AdminUserGroupMembersPayload, String> {
let Some(body) = request_body.filter(|body| !body.is_empty()) else {
return Err("请求数据验证失败".to_string());
};
serde_json::from_slice::<AdminUserGroupMembersPayload>(body)
.map_err(|_| "请求数据验证失败".to_string())
}
fn user_group_payload(
group: aether_data::repository::users::StoredUserGroup,
default_group_id: Option<&str>,
) -> serde_json::Value {
json!({
"id": group.id,
"name": group.name,
"normalized_name": group.normalized_name,
"description": group.description,
"priority": group.priority,
"allowed_providers": group.allowed_providers,
"allowed_providers_mode": group.allowed_providers_mode,
"allowed_api_formats": group.allowed_api_formats,
"allowed_api_formats_mode": group.allowed_api_formats_mode,
"allowed_models": group.allowed_models,
"allowed_models_mode": group.allowed_models_mode,
"rate_limit": group.rate_limit,
"rate_limit_mode": group.rate_limit_mode,
"is_default": default_group_id == Some(group.id.as_str()),
"created_at": format_optional_datetime_iso8601(group.created_at),
"updated_at": format_optional_datetime_iso8601(group.updated_at),
})
}
fn normalize_list_mode(value: &str) -> Result<String, String> {
match value.trim().to_ascii_lowercase().as_str() {
"inherit" | "unrestricted" | "specific" | "deny_all" => {
Ok(value.trim().to_ascii_lowercase())
}
_ => Err("权限列表模式不合法".to_string()),
}
}
fn normalize_rate_mode(value: &str) -> Result<String, String> {
match value.trim().to_ascii_lowercase().as_str() {
"inherit" | "system" | "custom" => Ok(value.trim().to_ascii_lowercase()),
_ => Err("限速模式不合法".to_string()),
}
}
fn default_list_mode() -> String {
"inherit".to_string()
}
fn default_rate_limit_mode() -> String {
"inherit".to_string()
}
fn normalize_ids(values: Vec<String>) -> Vec<String> {
values
.into_iter()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.collect::<std::collections::BTreeSet<_>>()
.into_iter()
.collect()
}
fn user_group_id_from_path(request_path: &str) -> Option<String> {
let value = request_path
.strip_prefix("/api/admin/user-groups/")?
.trim()
.trim_matches('/')
.to_string();
if value.is_empty() || value.contains('/') || value == "default" {
None
} else {
Some(value)
}
}
fn user_group_member_group_id_from_path(request_path: &str) -> Option<String> {
let value = request_path
.strip_prefix("/api/admin/user-groups/")?
.trim()
.trim_matches('/');
let group_id = value.strip_suffix("/members")?.trim_matches('/');
if group_id.is_empty() || group_id.contains('/') {
None
} else {
Some(group_id.to_string())
}
}
fn bad_request_owned(detail: String) -> Response<Body> {
(
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response()
}
fn not_found(detail: &'static str) -> Response<Body> {
(
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": detail })),
)
.into_response()
}
fn is_duplicate_group_name_error(err: &GatewayError) -> bool {
match err {
GatewayError::Internal(message) => message.contains("duplicate user group name"),
_ => false,
}
}

View File

@@ -1,10 +1,12 @@
use super::super::{ use super::super::{
admin_default_user_initial_gift, build_admin_users_read_only_response, admin_default_user_initial_gift, build_admin_users_read_only_response,
normalize_admin_optional_user_email, normalize_admin_user_api_formats, legacy_admin_list_policy_mode, legacy_admin_rate_limit_policy_mode,
normalize_admin_user_role, normalize_admin_user_string_list, normalize_admin_username, normalize_admin_list_policy_mode, normalize_admin_optional_user_email,
validate_admin_user_password, AdminCreateUserRequest, normalize_admin_rate_limit_policy_mode, normalize_admin_user_api_formats,
normalize_admin_user_group_ids, normalize_admin_user_role, normalize_admin_user_string_list,
normalize_admin_username, validate_admin_user_password, AdminCreateUserRequest,
}; };
use super::support::{admin_user_password_policy, build_admin_user_payload}; use super::support::{admin_user_password_policy, build_admin_user_payload_with_groups};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::attach_admin_audit_response; use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::GatewayError; use crate::GatewayError;
@@ -136,6 +138,72 @@ pub(in super::super) async fn build_admin_create_user_response(
.into_response()) .into_response())
} }
}; };
let allowed_providers_mode = match payload.allowed_providers_mode.as_deref() {
Some(value) => match normalize_admin_list_policy_mode(value) {
Ok(value) => value,
Err(detail) => {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response())
}
},
None => legacy_admin_list_policy_mode(&allowed_providers),
};
let allowed_api_formats_mode = match payload.allowed_api_formats_mode.as_deref() {
Some(value) => match normalize_admin_list_policy_mode(value) {
Ok(value) => value,
Err(detail) => {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response())
}
},
None => legacy_admin_list_policy_mode(&allowed_api_formats),
};
let allowed_models_mode = match payload.allowed_models_mode.as_deref() {
Some(value) => match normalize_admin_list_policy_mode(value) {
Ok(value) => value,
Err(detail) => {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response())
}
},
None => legacy_admin_list_policy_mode(&allowed_models),
};
let rate_limit_mode = match payload.rate_limit_mode.as_deref() {
Some(value) => match normalize_admin_rate_limit_policy_mode(value) {
Ok(value) => value,
Err(detail) => {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response())
}
},
None => legacy_admin_rate_limit_policy_mode(payload.rate_limit),
};
let group_ids = normalize_admin_user_group_ids(payload.group_ids);
let groups = if group_ids.is_empty() {
Vec::new()
} else {
let groups = state.list_user_groups_by_ids(&group_ids).await?;
if groups.len() != group_ids.len() {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "用户分组不存在" })),
)
.into_response());
}
groups
};
if let Some(email) = email.as_deref() { if let Some(email) = email.as_deref() {
if state.find_user_auth_by_identifier(email).await?.is_some() { if state.find_user_auth_by_identifier(email).await?.is_some() {
@@ -209,12 +277,33 @@ pub(in super::super) async fn build_admin_create_user_response(
"当前为只读模式,无法初始化用户钱包", "当前为只读模式,无法初始化用户钱包",
)); ));
} }
let Some(user) = state
.update_local_auth_user_policy_modes(
&user.id,
Some(allowed_providers_mode.clone()),
Some(allowed_api_formats_mode.clone()),
Some(allowed_models_mode.clone()),
Some(rate_limit_mode.clone()),
)
.await?
else {
return Ok(build_admin_users_read_only_response(
"当前为只读模式,无法创建用户",
));
};
if !group_ids.is_empty() {
state
.replace_user_groups_for_user(&user.id, &group_ids)
.await?;
}
Ok(attach_admin_audit_response( Ok(attach_admin_audit_response(
Json(build_admin_user_payload( Json(build_admin_user_payload_with_groups(
&user, &user,
payload.rate_limit, payload.rate_limit,
Some(rate_limit_mode.as_str()),
payload.unlimited, payload.unlimited,
&groups,
)) ))
.into_response(), .into_response(),
"admin_user_created", "admin_user_created",

View File

@@ -1,6 +1,7 @@
use super::super::{build_admin_users_bad_request_response, format_optional_datetime_iso8601}; use super::super::{build_admin_users_bad_request_response, format_optional_datetime_iso8601};
use super::support::{ use super::support::{
admin_user_id_from_detail_path, build_admin_user_payload, find_admin_export_user, admin_user_id_from_detail_path, build_admin_user_export_payload,
build_admin_user_payload_with_groups, find_admin_export_user,
}; };
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{query_param_optional_bool, query_param_value}; use crate::handlers::admin::shared::{query_param_optional_bool, query_param_value};
@@ -32,6 +33,9 @@ pub(in super::super) async fn build_admin_list_users_response(
let search = query_param_value(request_context.query_string(), "search") let search = query_param_value(request_context.query_string(), "search")
.map(|value| value.trim().to_string()) .map(|value| value.trim().to_string())
.filter(|value| !value.is_empty()); .filter(|value| !value.is_empty());
let group_id = query_param_value(request_context.query_string(), "group_id")
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty());
let paged_rows = state let paged_rows = state
.list_export_users_page(&aether_data::repository::users::UserExportListQuery { .list_export_users_page(&aether_data::repository::users::UserExportListQuery {
@@ -40,16 +44,25 @@ pub(in super::super) async fn build_admin_list_users_response(
role: role.clone(), role: role.clone(),
is_active, is_active,
search, search,
group_id,
}) })
.await?; .await?;
let user_ids = paged_rows let user_ids = paged_rows
.iter() .iter()
.map(|row| row.id.clone()) .map(|row| row.id.clone())
.collect::<Vec<_>>(); .collect::<Vec<_>>();
let (auth_rows_result, wallet_rows_result, usage_totals_result) = tokio::join!( let (
auth_rows_result,
wallet_rows_result,
usage_totals_result,
memberships_result,
groups_result,
) = tokio::join!(
state.list_user_auth_by_ids(&user_ids), state.list_user_auth_by_ids(&user_ids),
state.list_wallet_snapshots_by_user_ids(&user_ids), state.list_wallet_snapshots_by_user_ids(&user_ids),
state.summarize_usage_totals_by_user_ids(&user_ids), state.summarize_usage_totals_by_user_ids(&user_ids),
state.list_user_group_memberships_by_user_ids(&user_ids),
state.list_user_groups(),
); );
let auth_by_user_id = auth_rows_result? let auth_by_user_id = auth_rows_result?
.into_iter() .into_iter()
@@ -63,6 +76,17 @@ pub(in super::super) async fn build_admin_list_users_response(
.into_iter() .into_iter()
.map(|item| (item.user_id.clone(), item)) .map(|item| (item.user_id.clone(), item))
.collect::<BTreeMap<_, _>>(); .collect::<BTreeMap<_, _>>();
let groups_by_id = groups_result?
.into_iter()
.map(|group| (group.id.clone(), group))
.collect::<BTreeMap<_, _>>();
let mut group_ids_by_user_id = BTreeMap::<String, Vec<String>>::new();
for membership in memberships_result? {
group_ids_by_user_id
.entry(membership.user_id)
.or_default()
.push(membership.group_id);
}
let mut payload = Vec::with_capacity(paged_rows.len()); let mut payload = Vec::with_capacity(paged_rows.len());
for row in paged_rows { for row in paged_rows {
@@ -71,25 +95,25 @@ pub(in super::super) async fn build_admin_list_users_response(
.get(&row.id) .get(&row.id)
.is_some_and(|wallet| wallet.limit_mode.eq_ignore_ascii_case("unlimited")); .is_some_and(|wallet| wallet.limit_mode.eq_ignore_ascii_case("unlimited"));
let usage_totals = usage_totals_by_user_id.get(&row.id); let usage_totals = usage_totals_by_user_id.get(&row.id);
payload.push(json!({ let groups = group_ids_by_user_id
"id": row.id, .get(&row.id)
"email": row.email, .into_iter()
"username": row.username, .flatten()
"role": row.role, .filter_map(|group_id| groups_by_id.get(group_id).cloned())
"allowed_providers": row.allowed_providers, .collect::<Vec<_>>();
"allowed_api_formats": row.allowed_api_formats, payload.push(build_admin_user_export_payload(
"allowed_models": row.allowed_models, &row,
"rate_limit": row.rate_limit, unlimited,
"unlimited": unlimited, auth.as_ref().and_then(|user| user.created_at),
"is_active": row.is_active, auth.as_ref().and_then(|user| user.last_login_at),
"created_at": format_optional_datetime_iso8601(auth.as_ref().and_then(|user| user.created_at)), usage_totals
"updated_at": serde_json::Value::Null, .map(|item| item.request_count)
"last_login_at": format_optional_datetime_iso8601( .unwrap_or_default(),
auth.as_ref().and_then(|user| user.last_login_at), usage_totals
), .map(|item| item.total_tokens)
"request_count": usage_totals.map(|item| item.request_count).unwrap_or_default(), .unwrap_or_default(),
"total_tokens": usage_totals.map(|item| item.total_tokens).unwrap_or_default(), &groups,
})); ));
} }
Ok(Json(payload).into_response()) Ok(Json(payload).into_response())
@@ -116,13 +140,16 @@ pub(in super::super) async fn build_admin_get_user_response(
)) ))
.await?; .await?;
let export_row = find_admin_export_user(state, &user_id).await?; let export_row = find_admin_export_user(state, &user_id).await?;
let groups = state.list_user_groups_for_user(&user_id).await?;
let unlimited = wallet let unlimited = wallet
.as_ref() .as_ref()
.is_some_and(|wallet| wallet.limit_mode.eq_ignore_ascii_case("unlimited")); .is_some_and(|wallet| wallet.limit_mode.eq_ignore_ascii_case("unlimited"));
Ok(Json(build_admin_user_payload( Ok(Json(build_admin_user_payload_with_groups(
&user, &user,
export_row.as_ref().and_then(|row| row.rate_limit), export_row.as_ref().and_then(|row| row.rate_limit),
export_row.as_ref().map(|row| row.rate_limit_mode.as_str()),
unlimited, unlimited,
&groups,
)) ))
.into_response()) .into_response())
} }

View File

@@ -36,6 +36,16 @@ pub(super) fn build_admin_user_payload(
user: &aether_data::repository::users::StoredUserAuthRecord, user: &aether_data::repository::users::StoredUserAuthRecord,
rate_limit: Option<i32>, rate_limit: Option<i32>,
unlimited: bool, unlimited: bool,
) -> serde_json::Value {
build_admin_user_payload_with_groups(user, rate_limit, None, unlimited, &[])
}
pub(super) fn build_admin_user_payload_with_groups(
user: &aether_data::repository::users::StoredUserAuthRecord,
rate_limit: Option<i32>,
rate_limit_mode: Option<&str>,
unlimited: bool,
groups: &[aether_data::repository::users::StoredUserGroup],
) -> serde_json::Value { ) -> serde_json::Value {
json!({ json!({
"id": user.id, "id": user.id,
@@ -43,14 +53,233 @@ pub(super) fn build_admin_user_payload(
"username": user.username, "username": user.username,
"role": user.role, "role": user.role,
"allowed_providers": user.allowed_providers, "allowed_providers": user.allowed_providers,
"allowed_providers_mode": user.allowed_providers_mode,
"allowed_api_formats": user.allowed_api_formats, "allowed_api_formats": user.allowed_api_formats,
"allowed_api_formats_mode": user.allowed_api_formats_mode,
"allowed_models": user.allowed_models, "allowed_models": user.allowed_models,
"allowed_models_mode": user.allowed_models_mode,
"rate_limit": rate_limit, "rate_limit": rate_limit,
"rate_limit_mode": rate_limit_mode.unwrap_or("system"),
"unlimited": unlimited, "unlimited": unlimited,
"is_active": user.is_active, "is_active": user.is_active,
"created_at": format_optional_datetime_iso8601(user.created_at), "created_at": format_optional_datetime_iso8601(user.created_at),
"updated_at": serde_json::Value::Null, "updated_at": serde_json::Value::Null,
"last_login_at": format_optional_datetime_iso8601(user.last_login_at), "last_login_at": format_optional_datetime_iso8601(user.last_login_at),
"groups": groups.iter().map(user_group_badge_payload).collect::<Vec<_>>(),
"effective_policy": effective_policy_payload(
user.allowed_providers.as_ref(),
&user.allowed_providers_mode,
user.allowed_api_formats.as_ref(),
&user.allowed_api_formats_mode,
user.allowed_models.as_ref(),
&user.allowed_models_mode,
rate_limit,
rate_limit_mode.unwrap_or("system"),
groups,
),
})
}
#[allow(clippy::too_many_arguments)]
pub(super) fn build_admin_user_export_payload(
row: &aether_data::repository::users::StoredUserExportRow,
unlimited: bool,
created_at: Option<chrono::DateTime<chrono::Utc>>,
last_login_at: Option<chrono::DateTime<chrono::Utc>>,
request_count: u64,
total_tokens: u64,
groups: &[aether_data::repository::users::StoredUserGroup],
) -> serde_json::Value {
json!({
"id": row.id,
"email": row.email,
"username": row.username,
"role": row.role,
"allowed_providers": row.allowed_providers,
"allowed_providers_mode": row.allowed_providers_mode,
"allowed_api_formats": row.allowed_api_formats,
"allowed_api_formats_mode": row.allowed_api_formats_mode,
"allowed_models": row.allowed_models,
"allowed_models_mode": row.allowed_models_mode,
"rate_limit": row.rate_limit,
"rate_limit_mode": row.rate_limit_mode,
"unlimited": unlimited,
"is_active": row.is_active,
"created_at": format_optional_datetime_iso8601(created_at),
"updated_at": serde_json::Value::Null,
"last_login_at": format_optional_datetime_iso8601(last_login_at),
"request_count": request_count,
"total_tokens": total_tokens,
"groups": groups.iter().map(user_group_badge_payload).collect::<Vec<_>>(),
"effective_policy": effective_policy_payload(
row.allowed_providers.as_ref(),
&row.allowed_providers_mode,
row.allowed_api_formats.as_ref(),
&row.allowed_api_formats_mode,
row.allowed_models.as_ref(),
&row.allowed_models_mode,
row.rate_limit,
&row.rate_limit_mode,
groups,
),
})
}
pub(super) fn user_group_badge_payload(
group: &aether_data::repository::users::StoredUserGroup,
) -> serde_json::Value {
json!({
"id": group.id,
"name": group.name,
"priority": group.priority,
})
}
#[allow(clippy::too_many_arguments)]
fn effective_policy_payload(
allowed_providers: Option<&Vec<String>>,
allowed_providers_mode: &str,
allowed_api_formats: Option<&Vec<String>>,
allowed_api_formats_mode: &str,
allowed_models: Option<&Vec<String>>,
allowed_models_mode: &str,
rate_limit: Option<i32>,
rate_limit_mode: &str,
groups: &[aether_data::repository::users::StoredUserGroup],
) -> serde_json::Value {
let mut sorted_groups = groups.to_vec();
sorted_groups.sort_by(|left, right| {
right
.priority
.cmp(&left.priority)
.then_with(|| left.name.cmp(&right.name))
.then_with(|| left.id.cmp(&right.id))
});
json!({
"allowed_providers": effective_list_policy_payload(
allowed_providers,
allowed_providers_mode,
&sorted_groups,
|group| (&group.allowed_providers_mode, group.allowed_providers.as_ref()),
),
"allowed_api_formats": effective_list_policy_payload(
allowed_api_formats,
allowed_api_formats_mode,
&sorted_groups,
|group| (&group.allowed_api_formats_mode, group.allowed_api_formats.as_ref()),
),
"allowed_models": effective_list_policy_payload(
allowed_models,
allowed_models_mode,
&sorted_groups,
|group| (&group.allowed_models_mode, group.allowed_models.as_ref()),
),
"rate_limit": effective_rate_limit_policy_payload(rate_limit, rate_limit_mode, &sorted_groups),
})
}
fn effective_list_policy_payload(
user_values: Option<&Vec<String>>,
user_mode: &str,
groups: &[aether_data::repository::users::StoredUserGroup],
group_field: impl Fn(
&aether_data::repository::users::StoredUserGroup,
) -> (&String, Option<&Vec<String>>),
) -> serde_json::Value {
match user_mode {
"unrestricted" => policy_payload("unrestricted", serde_json::Value::Null, "user", None),
"specific" => policy_payload(
"specific",
json!(user_values.cloned().unwrap_or_default()),
"user",
None,
),
"deny_all" => policy_payload("deny_all", json!(Vec::<String>::new()), "user", None),
"inherit" => {
for group in groups {
let (mode, values) = group_field(group);
match mode.as_str() {
"unrestricted" => {
return policy_payload(
"unrestricted",
serde_json::Value::Null,
"group",
Some(group),
)
}
"specific" => {
return policy_payload(
"specific",
json!(values.cloned().unwrap_or_default()),
"group",
Some(group),
)
}
"deny_all" => {
return policy_payload(
"deny_all",
json!(Vec::<String>::new()),
"group",
Some(group),
)
}
_ => {}
}
}
policy_payload("unrestricted", serde_json::Value::Null, "fallback", None)
}
_ => policy_payload("unrestricted", serde_json::Value::Null, "fallback", None),
}
}
fn effective_rate_limit_policy_payload(
user_rate_limit: Option<i32>,
user_mode: &str,
groups: &[aether_data::repository::users::StoredUserGroup],
) -> serde_json::Value {
match user_mode {
"custom" => policy_payload("custom", json!(user_rate_limit.unwrap_or(0)), "user", None),
"system" => policy_payload("system", serde_json::Value::Null, "user", None),
"inherit" => {
for group in groups {
match group.rate_limit_mode.as_str() {
"custom" => {
return policy_payload(
"custom",
json!(group.rate_limit.unwrap_or(0)),
"group",
Some(group),
)
}
"system" => {
return policy_payload(
"system",
serde_json::Value::Null,
"group",
Some(group),
)
}
_ => {}
}
}
policy_payload("system", serde_json::Value::Null, "fallback", None)
}
_ => policy_payload("system", serde_json::Value::Null, "fallback", None),
}
}
fn policy_payload(
mode: &str,
value: serde_json::Value,
source: &str,
group: Option<&aether_data::repository::users::StoredUserGroup>,
) -> serde_json::Value {
json!({
"mode": mode,
"value": value,
"source": source,
"group_id": group.map(|group| group.id.as_str()),
"group_name": group.map(|group| group.name.as_str()),
}) })
} }

View File

@@ -1,12 +1,14 @@
use super::super::{ use super::super::{
build_admin_users_bad_request_response, build_admin_users_data_unavailable_response, build_admin_users_bad_request_response, build_admin_users_data_unavailable_response,
build_admin_users_read_only_response, normalize_admin_optional_user_email, build_admin_users_read_only_response, normalize_admin_list_policy_mode,
normalize_admin_user_api_formats, normalize_admin_user_role, normalize_admin_user_string_list, normalize_admin_optional_user_email, normalize_admin_rate_limit_policy_mode,
normalize_admin_username, validate_admin_user_password, AdminUpdateUserPatch, normalize_admin_user_api_formats, normalize_admin_user_group_ids, normalize_admin_user_role,
normalize_admin_user_string_list, normalize_admin_username, validate_admin_user_password,
AdminUpdateUserPatch,
}; };
use super::support::{ use super::support::{
admin_user_id_from_detail_path, admin_user_password_policy, build_admin_user_payload, admin_user_id_from_detail_path, admin_user_password_policy,
find_admin_export_user, build_admin_user_payload_with_groups, find_admin_export_user,
}; };
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::attach_admin_audit_response; use crate::handlers::admin::shared::attach_admin_audit_response;
@@ -177,15 +179,105 @@ pub(in super::super) async fn build_admin_update_user_response(
} else { } else {
None None
}; };
let allowed_providers_mode = if field_presence.contains("allowed_providers_mode") {
match payload.allowed_providers_mode.as_deref() {
Some(value) => match normalize_admin_list_policy_mode(value) {
Ok(value) => Some(value),
Err(detail) => {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response())
}
},
None => None,
}
} else {
None
};
let allowed_api_formats_mode = if field_presence.contains("allowed_api_formats_mode") {
match payload.allowed_api_formats_mode.as_deref() {
Some(value) => match normalize_admin_list_policy_mode(value) {
Ok(value) => Some(value),
Err(detail) => {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response())
}
},
None => None,
}
} else {
None
};
let allowed_models_mode = if field_presence.contains("allowed_models_mode") {
match payload.allowed_models_mode.as_deref() {
Some(value) => match normalize_admin_list_policy_mode(value) {
Ok(value) => Some(value),
Err(detail) => {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response())
}
},
None => None,
}
} else {
None
};
let rate_limit_mode = if field_presence.contains("rate_limit_mode") {
match payload.rate_limit_mode.as_deref() {
Some(value) => match normalize_admin_rate_limit_policy_mode(value) {
Ok(value) => Some(value),
Err(detail) => {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": detail })),
)
.into_response())
}
},
None => None,
}
} else {
None
};
let group_ids = if field_presence.contains("group_ids") {
Some(normalize_admin_user_group_ids(payload.group_ids))
} else {
None
};
if let Some(group_ids) = group_ids.as_ref() {
if !group_ids.is_empty() {
let groups = state.list_user_groups_by_ids(group_ids).await?;
if groups.len() != group_ids.len() {
return Ok((
http::StatusCode::BAD_REQUEST,
Json(json!({ "detail": "用户分组不存在" })),
)
.into_response());
}
}
}
let needs_auth_user_write = email.is_some() let needs_auth_user_write = email.is_some()
|| username.is_some() || username.is_some()
|| payload.password.is_some() || payload.password.is_some()
|| role.is_some() || role.is_some()
|| field_presence.contains("allowed_providers") || field_presence.contains("allowed_providers")
|| allowed_providers_mode.is_some()
|| field_presence.contains("allowed_api_formats") || field_presence.contains("allowed_api_formats")
|| allowed_api_formats_mode.is_some()
|| field_presence.contains("allowed_models") || field_presence.contains("allowed_models")
|| allowed_models_mode.is_some()
|| field_presence.contains("rate_limit") || field_presence.contains("rate_limit")
|| payload.is_active.is_some(); || rate_limit_mode.is_some()
|| payload.is_active.is_some()
|| group_ids.is_some();
if needs_auth_user_write && !state.has_auth_user_write_capability() { if needs_auth_user_write && !state.has_auth_user_write_capability() {
return Ok(build_admin_users_read_only_response( return Ok(build_admin_users_read_only_response(
"当前为只读模式,无法更新用户", "当前为只读模式,无法更新用户",
@@ -210,6 +302,34 @@ pub(in super::super) async fn build_admin_update_user_response(
.into_response()); .into_response());
} }
} }
if allowed_providers_mode.is_some()
|| allowed_api_formats_mode.is_some()
|| allowed_models_mode.is_some()
|| rate_limit_mode.is_some()
{
if state
.update_local_auth_user_policy_modes(
&user_id,
allowed_providers_mode,
allowed_api_formats_mode,
allowed_models_mode,
rate_limit_mode,
)
.await?
.is_none()
{
return Ok((
http::StatusCode::NOT_FOUND,
Json(json!({ "detail": "用户不存在" })),
)
.into_response());
}
}
if let Some(group_ids) = group_ids.as_ref() {
state
.replace_user_groups_for_user(&user_id, group_ids)
.await?;
}
if let Some(password) = payload.password.as_deref() { if let Some(password) = payload.password.as_deref() {
let password_policy = admin_user_password_policy(state).await?; let password_policy = admin_user_password_policy(state).await?;
@@ -322,13 +442,21 @@ pub(in super::super) async fn build_admin_update_user_response(
.as_ref() .as_ref()
.is_some_and(|wallet| wallet.limit_mode.eq_ignore_ascii_case("unlimited")); .is_some_and(|wallet| wallet.limit_mode.eq_ignore_ascii_case("unlimited"));
let export_row = find_admin_export_user(state, &user_id).await?; let export_row = find_admin_export_user(state, &user_id).await?;
let groups = state.list_user_groups_for_user(&user_id).await?;
let rate_limit = export_row let rate_limit = export_row
.as_ref() .as_ref()
.and_then(|row| row.rate_limit) .and_then(|row| row.rate_limit)
.or(payload.rate_limit); .or(payload.rate_limit);
Ok(attach_admin_audit_response( Ok(attach_admin_audit_response(
Json(build_admin_user_payload(&user, rate_limit, unlimited)).into_response(), Json(build_admin_user_payload_with_groups(
&user,
rate_limit,
export_row.as_ref().map(|row| row.rate_limit_mode.as_str()),
unlimited,
&groups,
))
.into_response(),
"admin_user_updated", "admin_user_updated",
"update_user", "update_user",
"user", "user",

View File

@@ -4,6 +4,7 @@ const ADMIN_USERS_DATA_UNAVAILABLE_DETAIL: &str = "Admin user management data un
mod api_keys; mod api_keys;
mod batch; mod batch;
mod groups;
mod lifecycle; mod lifecycle;
mod route_seam; mod route_seam;
mod routes; mod routes;
@@ -23,6 +24,12 @@ pub(crate) use self::api_keys::{
use self::batch::{ use self::batch::{
build_admin_resolve_user_selection_response, build_admin_user_batch_action_response, build_admin_resolve_user_selection_response, build_admin_user_batch_action_response,
}; };
use self::groups::{
build_admin_create_user_group_response, build_admin_delete_user_group_response,
build_admin_list_user_group_members_response, build_admin_list_user_groups_response,
build_admin_replace_user_group_members_response, build_admin_set_default_user_group_response,
build_admin_update_user_group_response,
};
use self::lifecycle::{ use self::lifecycle::{
build_admin_create_user_response, build_admin_delete_user_response, build_admin_create_user_response, build_admin_delete_user_response,
build_admin_get_user_response, build_admin_list_users_response, build_admin_get_user_response, build_admin_list_users_response,
@@ -36,10 +43,12 @@ use self::shared::AdminUpdateUserPatch;
use self::shared::{ use self::shared::{
admin_default_user_initial_gift, build_admin_users_bad_request_response, admin_default_user_initial_gift, build_admin_users_bad_request_response,
build_admin_users_data_unavailable_response, build_admin_users_read_only_response, build_admin_users_data_unavailable_response, build_admin_users_read_only_response,
format_optional_datetime_iso8601, normalize_admin_optional_user_email, format_optional_datetime_iso8601, legacy_admin_list_policy_mode,
normalize_admin_user_role, normalize_admin_username, validate_admin_user_password, legacy_admin_rate_limit_policy_mode, normalize_admin_list_policy_mode,
AdminCreateUserApiKeyRequest, AdminCreateUserRequest, AdminToggleUserApiKeyLockRequest, normalize_admin_optional_user_email, normalize_admin_rate_limit_policy_mode,
AdminUpdateUserApiKeyRequest, normalize_admin_user_group_ids, normalize_admin_user_role, normalize_admin_username,
validate_admin_user_password, AdminCreateUserApiKeyRequest, AdminCreateUserRequest,
AdminToggleUserApiKeyLockRequest, AdminUpdateUserApiKeyRequest,
}; };
pub(crate) use self::shared::{normalize_admin_user_api_formats, normalize_admin_user_string_list}; pub(crate) use self::shared::{normalize_admin_user_api_formats, normalize_admin_user_string_list};

View File

@@ -1,13 +1,16 @@
use super::{ use super::{
build_admin_create_user_api_key_response, build_admin_create_user_response, build_admin_create_user_api_key_response, build_admin_create_user_group_response,
build_admin_delete_user_api_key_response, build_admin_delete_user_response, build_admin_create_user_response, build_admin_delete_user_api_key_response,
build_admin_delete_user_group_response, build_admin_delete_user_response,
build_admin_delete_user_session_response, build_admin_delete_user_sessions_response, build_admin_delete_user_session_response, build_admin_delete_user_sessions_response,
build_admin_get_user_response, build_admin_list_user_api_keys_response, build_admin_get_user_response, build_admin_list_user_api_keys_response,
build_admin_list_user_group_members_response, build_admin_list_user_groups_response,
build_admin_list_user_sessions_response, build_admin_list_users_response, build_admin_list_user_sessions_response, build_admin_list_users_response,
build_admin_resolve_user_selection_response, build_admin_reveal_user_api_key_response, build_admin_replace_user_group_members_response, build_admin_resolve_user_selection_response,
build_admin_reveal_user_api_key_response, build_admin_set_default_user_group_response,
build_admin_toggle_user_api_key_lock_response, build_admin_update_user_api_key_response, build_admin_toggle_user_api_key_lock_response, build_admin_update_user_api_key_response,
build_admin_update_user_response, build_admin_user_batch_action_response, build_admin_update_user_group_response, build_admin_update_user_response,
build_admin_users_data_unavailable_response, build_admin_user_batch_action_response, build_admin_users_data_unavailable_response,
}; };
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError; use crate::GatewayError;
@@ -15,8 +18,27 @@ use axum::{body::Body, http, response::Response};
fn is_admin_users_route(request_context: &AdminRequestContext<'_>) -> bool { fn is_admin_users_route(request_context: &AdminRequestContext<'_>) -> bool {
let path = request_context.path(); let path = request_context.path();
(request_context.method() == http::Method::GET ((request_context.method() == http::Method::GET
&& matches!(path, "/api/admin/users" | "/api/admin/users/")) || request_context.method() == http::Method::POST)
&& matches!(path, "/api/admin/user-groups" | "/api/admin/user-groups/"))
|| (request_context.method() == http::Method::PUT
&& matches!(
path,
"/api/admin/user-groups/default" | "/api/admin/user-groups/default/"
))
|| ((request_context.method() == http::Method::PUT
|| request_context.method() == http::Method::DELETE)
&& path.starts_with("/api/admin/user-groups/")
&& path.matches('/').count() == 4
&& !path.ends_with("/members")
&& !path.ends_with("/default"))
|| ((request_context.method() == http::Method::GET
|| request_context.method() == http::Method::PUT)
&& path.starts_with("/api/admin/user-groups/")
&& path.ends_with("/members")
&& path.matches('/').count() == 5)
|| (request_context.method() == http::Method::GET
&& matches!(path, "/api/admin/users" | "/api/admin/users/"))
|| (request_context.method() == http::Method::POST || (request_context.method() == http::Method::POST
&& matches!(path, "/api/admin/users" | "/api/admin/users/")) && matches!(path, "/api/admin/users" | "/api/admin/users/"))
|| (request_context.method() == http::Method::POST || (request_context.method() == http::Method::POST
@@ -84,6 +106,26 @@ pub(super) async fn maybe_build_local_admin_users_routes_response(
} }
match decision.route_kind.as_deref() { match decision.route_kind.as_deref() {
Some("list_user_groups") => Ok(Some(build_admin_list_user_groups_response(state).await?)),
Some("create_user_group") => Ok(Some(
build_admin_create_user_group_response(state, request_body).await?,
)),
Some("update_user_group") => Ok(Some(
build_admin_update_user_group_response(state, request_context, request_body).await?,
)),
Some("delete_user_group") => Ok(Some(
build_admin_delete_user_group_response(state, request_context).await?,
)),
Some("list_user_group_members") => Ok(Some(
build_admin_list_user_group_members_response(state, request_context).await?,
)),
Some("replace_user_group_members") => Ok(Some(
build_admin_replace_user_group_members_response(state, request_context, request_body)
.await?,
)),
Some("set_default_user_group") => Ok(Some(
build_admin_set_default_user_group_response(state, request_body).await?,
)),
Some("create_user") => Ok(Some( Some("create_user") => Ok(Some(
build_admin_create_user_response(state, request_context, request_body).await?, build_admin_create_user_response(state, request_context, request_body).await?,
)), )),

View File

@@ -68,11 +68,21 @@ pub(super) struct AdminCreateUserRequest {
#[serde(default)] #[serde(default)]
pub(super) allowed_providers: Option<Vec<String>>, pub(super) allowed_providers: Option<Vec<String>>,
#[serde(default)] #[serde(default)]
pub(super) allowed_providers_mode: Option<String>,
#[serde(default)]
pub(super) allowed_api_formats: Option<Vec<String>>, pub(super) allowed_api_formats: Option<Vec<String>>,
#[serde(default)] #[serde(default)]
pub(super) allowed_api_formats_mode: Option<String>,
#[serde(default)]
pub(super) allowed_models: Option<Vec<String>>, pub(super) allowed_models: Option<Vec<String>>,
#[serde(default)] #[serde(default)]
pub(super) allowed_models_mode: Option<String>,
#[serde(default)]
pub(super) rate_limit: Option<i32>, pub(super) rate_limit: Option<i32>,
#[serde(default)]
pub(super) rate_limit_mode: Option<String>,
#[serde(default)]
pub(super) group_ids: Vec<String>,
} }
#[derive(Debug, serde::Deserialize)] #[derive(Debug, serde::Deserialize)]
@@ -90,12 +100,22 @@ pub(super) struct AdminUpdateUserRequest {
#[serde(default)] #[serde(default)]
pub(super) allowed_providers: Option<Vec<String>>, pub(super) allowed_providers: Option<Vec<String>>,
#[serde(default)] #[serde(default)]
pub(super) allowed_providers_mode: Option<String>,
#[serde(default)]
pub(super) allowed_api_formats: Option<Vec<String>>, pub(super) allowed_api_formats: Option<Vec<String>>,
#[serde(default)] #[serde(default)]
pub(super) allowed_api_formats_mode: Option<String>,
#[serde(default)]
pub(super) allowed_models: Option<Vec<String>>, pub(super) allowed_models: Option<Vec<String>>,
#[serde(default)] #[serde(default)]
pub(super) allowed_models_mode: Option<String>,
#[serde(default)]
pub(super) rate_limit: Option<i32>, pub(super) rate_limit: Option<i32>,
#[serde(default)] #[serde(default)]
pub(super) rate_limit_mode: Option<String>,
#[serde(default)]
pub(super) group_ids: Vec<String>,
#[serde(default)]
pub(super) is_active: Option<bool>, pub(super) is_active: Option<bool>,
} }
@@ -263,6 +283,48 @@ pub(crate) fn normalize_admin_user_api_formats(
Ok(Some(normalized)) Ok(Some(normalized))
} }
pub(super) fn normalize_admin_list_policy_mode(value: &str) -> Result<String, String> {
match value.trim().to_ascii_lowercase().as_str() {
"inherit" | "unrestricted" | "specific" | "deny_all" => {
Ok(value.trim().to_ascii_lowercase())
}
_ => Err("权限列表模式不合法".to_string()),
}
}
pub(super) fn normalize_admin_rate_limit_policy_mode(value: &str) -> Result<String, String> {
match value.trim().to_ascii_lowercase().as_str() {
"inherit" | "system" | "custom" => Ok(value.trim().to_ascii_lowercase()),
_ => Err("限速模式不合法".to_string()),
}
}
pub(super) fn legacy_admin_list_policy_mode(values: &Option<Vec<String>>) -> String {
if values.is_some() {
"specific".to_string()
} else {
"unrestricted".to_string()
}
}
pub(super) fn legacy_admin_rate_limit_policy_mode(value: Option<i32>) -> String {
if value.is_some() {
"custom".to_string()
} else {
"system".to_string()
}
}
pub(super) fn normalize_admin_user_group_ids(values: Vec<String>) -> Vec<String> {
values
.into_iter()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.collect::<std::collections::BTreeSet<_>>()
.into_iter()
.collect()
}
pub(super) fn admin_default_user_initial_gift(value: Option<&serde_json::Value>) -> f64 { pub(super) fn admin_default_user_initial_gift(value: Option<&serde_json::Value>) -> f64 {
match value { match value {
Some(serde_json::Value::Number(number)) => number.as_f64().unwrap_or(10.0), Some(serde_json::Value::Number(number)) => number.as_f64().unwrap_or(10.0),

View File

@@ -513,6 +513,17 @@ pub(super) async fn handle_auth_register(
false, false,
); );
}; };
if let Err(err) = state
.assign_default_group_to_self_registered_user(&user.id)
.await
{
let _ = state.delete_local_auth_user(&user.id).await;
return build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("auth default user group assignment failed: {err:?}"),
false,
);
}
if require_verification { if require_verification {
if let Some(email) = email.as_deref() { if let Some(email) = email.as_deref() {

View File

@@ -325,6 +325,13 @@ pub(crate) fn admin_proxy_local_requires_buffered_body(
| (Some("users_manage"), http::Method::POST, Some("resolve_user_selection")) | (Some("users_manage"), http::Method::POST, Some("resolve_user_selection"))
| (Some("users_manage"), http::Method::POST, Some("batch_action_users")) | (Some("users_manage"), http::Method::POST, Some("batch_action_users"))
| (Some("users_manage"), http::Method::PUT, Some("update_user")) | (Some("users_manage"), http::Method::PUT, Some("update_user"))
| (Some("users_manage"), http::Method::POST, Some("create_user_group"))
| (Some("users_manage"), http::Method::PUT, Some("update_user_group"))
| (
Some("users_manage"),
http::Method::PUT,
Some("replace_user_group_members" | "set_default_user_group"),
)
| (Some("users_manage"), http::Method::POST, Some("create_user_api_key")) | (Some("users_manage"), http::Method::POST, Some("create_user_api_key"))
| (Some("users_manage"), http::Method::PUT, Some("update_user_api_key")) | (Some("users_manage"), http::Method::PUT, Some("update_user_api_key"))
| (Some("users_manage"), http::Method::PATCH, Some("lock_user_api_key")) | (Some("users_manage"), http::Method::PATCH, Some("lock_user_api_key"))

View File

@@ -211,6 +211,13 @@ pub(crate) async fn resolve_identity_oauth_login_user(
return Err(IdentityOAuthAccountError::Storage(format!("{err:?}"))); return Err(IdentityOAuthAccountError::Storage(format!("{err:?}")));
} }
} }
if let Err(err) = state
.assign_default_group_to_self_registered_user(&user.id)
.await
{
let _ = state.delete_local_auth_user(&user.id).await;
return Err(IdentityOAuthAccountError::Storage(format!("{err:?}")));
}
if let Err(err) = upsert_oauth_link(state, &user.id, claims, now).await { if let Err(err) = upsert_oauth_link(state, &user.id, claims, now).await {
let _ = state.delete_local_auth_user(&user.id).await; let _ = state.delete_local_auth_user(&user.id).await;
return Err(err); return Err(err);

View File

@@ -3,6 +3,32 @@ use std::collections::{BTreeMap, BTreeSet};
use crate::{AppState, GatewayError}; use crate::{AppState, GatewayError};
impl AppState { impl AppState {
pub(crate) async fn assign_default_group_to_self_registered_user(
&self,
user_id: &str,
) -> Result<(), GatewayError> {
let group_id = self
.read_system_config_json_value("default_user_group_id")
.await?
.and_then(|value| value.as_str().map(str::to_string))
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty());
let Some(group_id) = group_id else {
return Ok(());
};
if self.find_user_group_by_id(&group_id).await?.is_none() {
return Err(GatewayError::Internal(format!(
"default_user_group_id points to missing group: {group_id}"
)));
}
if !self.add_user_to_group(&group_id, user_id).await? {
return Err(GatewayError::Internal(format!(
"failed to add user {user_id} to default group {group_id}"
)));
}
Ok(())
}
pub(crate) async fn resolve_auth_user_summaries_by_ids( pub(crate) async fn resolve_auth_user_summaries_by_ids(
&self, &self,
user_ids: &[String], user_ids: &[String],
@@ -132,6 +158,126 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string())) .map_err(|err| GatewayError::Internal(err.to_string()))
} }
pub(crate) async fn list_user_groups(
&self,
) -> Result<Vec<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.data
.list_user_groups()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn find_user_group_by_id(
&self,
group_id: &str,
) -> Result<Option<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.data
.find_user_group_by_id(group_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn list_user_groups_by_ids(
&self,
group_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.data
.list_user_groups_by_ids(group_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn create_user_group(
&self,
record: aether_data::repository::users::UpsertUserGroupRecord,
) -> Result<Option<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.data
.create_user_group(record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn update_user_group(
&self,
group_id: &str,
record: aether_data::repository::users::UpsertUserGroupRecord,
) -> Result<Option<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.data
.update_user_group(group_id, record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn delete_user_group(&self, group_id: &str) -> Result<bool, GatewayError> {
self.data
.delete_user_group(group_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn list_user_group_members(
&self,
group_id: &str,
) -> Result<Vec<aether_data::repository::users::StoredUserGroupMember>, GatewayError> {
self.data
.list_user_group_members(group_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn replace_user_group_members(
&self,
group_id: &str,
user_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroupMember>, GatewayError> {
self.data
.replace_user_group_members(group_id, user_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn list_user_groups_for_user(
&self,
user_id: &str,
) -> Result<Vec<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.data
.list_user_groups_for_user(user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn list_user_group_memberships_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroupMembership>, GatewayError> {
self.data
.list_user_group_memberships_by_user_ids(user_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn replace_user_groups_for_user(
&self,
user_id: &str,
group_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.data
.replace_user_groups_for_user(user_id, group_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn add_user_to_group(
&self,
group_id: &str,
user_id: &str,
) -> Result<bool, GatewayError> {
self.data
.add_user_to_group(group_id, user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn is_other_user_auth_email_taken( pub(crate) async fn is_other_user_auth_email_taken(
&self, &self,
email: &str, email: &str,
@@ -289,6 +435,12 @@ impl AppState {
Some(now), Some(now),
None, None,
) )
.map_err(|err| GatewayError::Internal(err.to_string()))?
.with_policy_modes(
"inherit".to_string(),
"inherit".to_string(),
"inherit".to_string(),
)
.map_err(|err| GatewayError::Internal(err.to_string()))?; .map_err(|err| GatewayError::Internal(err.to_string()))?;
store store
.lock() .lock()
@@ -418,6 +570,45 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string())) .map_err(|err| GatewayError::Internal(err.to_string()))
} }
pub(crate) 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<aether_data::repository::users::StoredUserAuthRecord>, GatewayError> {
#[cfg(test)]
if let Some(store) = self.auth_user_store.as_ref() {
let mut guard = store.lock().expect("auth user store should lock");
let Some(user) = guard.get_mut(user_id) else {
return Ok(None);
};
if let Some(mode) = allowed_providers_mode {
user.allowed_providers_mode = mode;
}
if let Some(mode) = allowed_api_formats_mode {
user.allowed_api_formats_mode = mode;
}
if let Some(mode) = allowed_models_mode {
user.allowed_models_mode = mode;
}
let _ = rate_limit_mode;
return Ok(Some(user.clone()));
}
self.data
.update_local_auth_user_policy_modes(
user_id,
allowed_providers_mode,
allowed_api_formats_mode,
allowed_models_mode,
rate_limit_mode,
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn touch_auth_user_last_login( pub(crate) async fn touch_auth_user_last_login(
&self, &self,
user_id: &str, user_id: &str,
@@ -508,6 +699,12 @@ impl AppState {
Some(now), Some(now),
None, None,
) )
.map_err(|err| GatewayError::Internal(err.to_string()))?
.with_policy_modes(
"inherit".to_string(),
"inherit".to_string(),
"inherit".to_string(),
)
.map_err(|err| GatewayError::Internal(err.to_string()))?; .map_err(|err| GatewayError::Internal(err.to_string()))?;
let gift_balance = if unlimited { let gift_balance = if unlimited {
0.0 0.0

View File

@@ -2927,8 +2927,11 @@ async fn gateway_handles_admin_usage_cache_affinity_interval_timeline_with_legac
role: "user".to_string(), role: "user".to_string(),
auth_source: "local".to_string(), auth_source: "local".to_string(),
allowed_providers: None, allowed_providers: None,
allowed_providers_mode: "unrestricted".to_string(),
allowed_api_formats: None, allowed_api_formats: None,
allowed_api_formats_mode: "unrestricted".to_string(),
allowed_models: None, allowed_models: None,
allowed_models_mode: "unrestricted".to_string(),
is_active: true, is_active: true,
is_deleted: false, is_deleted: false,
created_at: None, created_at: None,

View File

@@ -377,7 +377,10 @@ async fn embeddings_route_rejects_chat_only_model() {
Some(EXECUTION_PATH_LOCAL_AUTH_DENIED) Some(EXECUTION_PATH_LOCAL_AUTH_DENIED)
); );
let payload: serde_json::Value = response.json().await.expect("body should parse"); let payload: serde_json::Value = response.json().await.expect("body should parse");
assert_eq!(payload["error"]["message"], "当前密钥不允许访问模型 gpt-5"); assert_eq!(
payload["error"]["message"],
"当前用户、用户组或密钥的访问控制策略不允许访问模型 gpt-5"
);
gateway_handle.abort(); gateway_handle.abort();
} }
@@ -416,7 +419,7 @@ async fn embeddings_route_rejects_chat_only_api_format() {
let payload: serde_json::Value = response.json().await.expect("body should parse"); let payload: serde_json::Value = response.json().await.expect("body should parse");
assert_eq!( assert_eq!(
payload["error"]["message"], payload["error"]["message"],
"当前密钥不允许访问 openai:embedding 格式" "当前用户、用户组或密钥的访问控制策略不允许访问 openai:embedding 格式"
); );
gateway_handle.abort(); gateway_handle.abort();

View File

@@ -501,7 +501,7 @@ async fn gateway_locally_denies_disallowed_claude_api_format_without_hitting_con
assert_eq!(payload["error"]["type"], "http_error"); assert_eq!(payload["error"]["type"], "http_error");
assert_eq!( assert_eq!(
payload["error"]["message"], payload["error"]["message"],
"当前密钥不允许访问 claude:messages 格式" "当前用户、用户组或密钥的访问控制策略不允许访问 claude:messages 格式"
); );
assert_eq!(*auth_context_hits.lock().expect("mutex should lock"), 0); assert_eq!(*auth_context_hits.lock().expect("mutex should lock"), 0);
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
@@ -584,7 +584,7 @@ async fn gateway_locally_denies_disallowed_provider_without_hitting_control_or_u
assert_eq!(payload["error"]["type"], "http_error"); assert_eq!(payload["error"]["type"], "http_error");
assert_eq!( assert_eq!(
payload["error"]["message"], payload["error"]["message"],
"当前密钥不允许访问 claude 提供商" "当前用户、用户组或密钥的访问控制策略不允许访问 claude 提供商"
); );
assert_eq!(*auth_context_hits.lock().expect("mutex should lock"), 0); assert_eq!(*auth_context_hits.lock().expect("mutex should lock"), 0);
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
@@ -656,7 +656,7 @@ async fn gateway_locally_denies_disallowed_gemini_model_without_hitting_control_
assert_eq!(payload["error"]["type"], "http_error"); assert_eq!(payload["error"]["type"], "http_error");
assert_eq!( assert_eq!(
payload["error"]["message"], payload["error"]["message"],
"当前密钥不允许访问模型 gemini-2.5-pro" "当前用户、用户组或密钥的访问控制策略不允许访问模型 gemini-2.5-pro"
); );
assert_eq!(*auth_context_hits.lock().expect("mutex should lock"), 0); assert_eq!(*auth_context_hits.lock().expect("mutex should lock"), 0);
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);
@@ -806,7 +806,10 @@ async fn gateway_locally_denies_disallowed_openai_model_without_hitting_control_
); );
let payload: serde_json::Value = response.json().await.expect("response json should parse"); let payload: serde_json::Value = response.json().await.expect("response json should parse");
assert_eq!(payload["error"]["type"], "http_error"); assert_eq!(payload["error"]["type"], "http_error");
assert_eq!(payload["error"]["message"], "当前密钥不允许访问模型 gpt-5"); assert_eq!(
payload["error"]["message"],
"当前用户、用户组或密钥的访问控制策略不允许访问模型 gpt-5"
);
assert_eq!(*auth_context_hits.lock().expect("mutex should lock"), 0); assert_eq!(*auth_context_hits.lock().expect("mutex should lock"), 0);
assert_eq!(*public_hits.lock().expect("mutex should lock"), 0); assert_eq!(*public_hits.lock().expect("mutex should lock"), 0);

View File

@@ -319,7 +319,7 @@ async fn rerank_route_rejects_chat_only_api_format() {
let payload: serde_json::Value = response.json().await.expect("body should parse"); let payload: serde_json::Value = response.json().await.expect("body should parse");
assert_eq!( assert_eq!(
payload["error"]["message"], payload["error"]["message"],
"当前密钥不允许访问 openai:rerank 格式" "当前用户、用户组或密钥的访问控制策略不允许访问 openai:rerank 格式"
); );
gateway_handle.abort(); gateway_handle.abort();

View File

@@ -0,0 +1,53 @@
ALTER TABLE users
ADD COLUMN allowed_providers_mode VARCHAR(32) NOT NULL DEFAULT 'unrestricted',
ADD COLUMN allowed_api_formats_mode VARCHAR(32) NOT NULL DEFAULT 'unrestricted',
ADD COLUMN allowed_models_mode VARCHAR(32) NOT NULL DEFAULT 'unrestricted',
ADD COLUMN rate_limit_mode VARCHAR(32) NOT NULL DEFAULT 'system';
UPDATE users
SET allowed_providers_mode = CASE WHEN allowed_providers IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_providers_mode = 'unrestricted';
UPDATE users
SET allowed_api_formats_mode = CASE WHEN allowed_api_formats IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_api_formats_mode = 'unrestricted';
UPDATE users
SET allowed_models_mode = CASE WHEN allowed_models IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_models_mode = 'unrestricted';
UPDATE users
SET rate_limit_mode = CASE WHEN rate_limit IS NULL THEN 'system' ELSE 'custom' END
WHERE rate_limit_mode = 'system';
CREATE TABLE IF NOT EXISTS user_groups (
id VARCHAR(64) PRIMARY KEY,
name VARCHAR(100) NOT NULL,
normalized_name VARCHAR(100) NOT NULL,
description TEXT,
priority INT NOT NULL DEFAULT 0,
allowed_providers TEXT,
allowed_providers_mode VARCHAR(32) NOT NULL DEFAULT 'inherit',
allowed_api_formats TEXT,
allowed_api_formats_mode VARCHAR(32) NOT NULL DEFAULT 'inherit',
allowed_models TEXT,
allowed_models_mode VARCHAR(32) NOT NULL DEFAULT 'inherit',
rate_limit INT,
rate_limit_mode VARCHAR(32) NOT NULL DEFAULT 'inherit',
created_at BIGINT NOT NULL,
updated_at BIGINT NOT NULL,
UNIQUE KEY user_groups_normalized_name_key (normalized_name),
KEY user_groups_priority_name_idx (priority, name, id)
);
CREATE TABLE IF NOT EXISTS user_group_members (
group_id VARCHAR(64) NOT NULL,
user_id VARCHAR(64) NOT NULL,
created_at BIGINT NOT NULL,
PRIMARY KEY (group_id, user_id),
KEY user_group_members_user_id_idx (user_id),
CONSTRAINT user_group_members_group_id_fk
FOREIGN KEY (group_id) REFERENCES user_groups(id) ON DELETE CASCADE,
CONSTRAINT user_group_members_user_id_fk
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);

View File

@@ -0,0 +1,60 @@
ALTER TABLE public.users
ADD COLUMN IF NOT EXISTS allowed_providers_mode text DEFAULT 'unrestricted' NOT NULL,
ADD COLUMN IF NOT EXISTS allowed_api_formats_mode text DEFAULT 'unrestricted' NOT NULL,
ADD COLUMN IF NOT EXISTS allowed_models_mode text DEFAULT 'unrestricted' NOT NULL,
ADD COLUMN IF NOT EXISTS rate_limit_mode text DEFAULT 'system' NOT NULL;
UPDATE public.users
SET allowed_providers_mode = CASE WHEN allowed_providers IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_providers_mode = 'unrestricted';
UPDATE public.users
SET allowed_api_formats_mode = CASE WHEN allowed_api_formats IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_api_formats_mode = 'unrestricted';
UPDATE public.users
SET allowed_models_mode = CASE WHEN allowed_models IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_models_mode = 'unrestricted';
UPDATE public.users
SET rate_limit_mode = CASE WHEN rate_limit IS NULL THEN 'system' ELSE 'custom' END
WHERE rate_limit_mode = 'system';
CREATE TABLE IF NOT EXISTS public.user_groups (
id character varying(36) PRIMARY KEY,
name character varying(100) NOT NULL,
normalized_name character varying(100) NOT NULL UNIQUE,
description text,
priority integer DEFAULT 0 NOT NULL,
allowed_providers json,
allowed_providers_mode text DEFAULT 'inherit' NOT NULL,
allowed_api_formats json,
allowed_api_formats_mode text DEFAULT 'inherit' NOT NULL,
allowed_models json,
allowed_models_mode text DEFAULT 'inherit' NOT NULL,
rate_limit integer,
rate_limit_mode text DEFAULT 'inherit' NOT NULL,
created_at timestamp with time zone DEFAULT now() NOT NULL,
updated_at timestamp with time zone DEFAULT now() NOT NULL,
CONSTRAINT user_groups_allowed_providers_mode_check
CHECK (allowed_providers_mode IN ('inherit', 'unrestricted', 'specific', 'deny_all')),
CONSTRAINT user_groups_allowed_api_formats_mode_check
CHECK (allowed_api_formats_mode IN ('inherit', 'unrestricted', 'specific', 'deny_all')),
CONSTRAINT user_groups_allowed_models_mode_check
CHECK (allowed_models_mode IN ('inherit', 'unrestricted', 'specific', 'deny_all')),
CONSTRAINT user_groups_rate_limit_mode_check
CHECK (rate_limit_mode IN ('inherit', 'system', 'custom'))
);
CREATE TABLE IF NOT EXISTS public.user_group_members (
group_id character varying(36) NOT NULL REFERENCES public.user_groups(id) ON DELETE CASCADE,
user_id character varying(36) NOT NULL REFERENCES public.users(id) ON DELETE CASCADE,
created_at timestamp with time zone DEFAULT now() NOT NULL,
PRIMARY KEY (group_id, user_id)
);
CREATE INDEX IF NOT EXISTS user_group_members_user_id_idx
ON public.user_group_members (user_id);
CREATE INDEX IF NOT EXISTS user_groups_priority_name_idx
ON public.user_groups (priority DESC, name ASC, id ASC);

View File

@@ -0,0 +1,51 @@
ALTER TABLE users ADD COLUMN allowed_providers_mode TEXT NOT NULL DEFAULT 'unrestricted';
ALTER TABLE users ADD COLUMN allowed_api_formats_mode TEXT NOT NULL DEFAULT 'unrestricted';
ALTER TABLE users ADD COLUMN allowed_models_mode TEXT NOT NULL DEFAULT 'unrestricted';
ALTER TABLE users ADD COLUMN rate_limit_mode TEXT NOT NULL DEFAULT 'system';
UPDATE users
SET allowed_providers_mode = CASE WHEN allowed_providers IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_providers_mode = 'unrestricted';
UPDATE users
SET allowed_api_formats_mode = CASE WHEN allowed_api_formats IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_api_formats_mode = 'unrestricted';
UPDATE users
SET allowed_models_mode = CASE WHEN allowed_models IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_models_mode = 'unrestricted';
UPDATE users
SET rate_limit_mode = CASE WHEN rate_limit IS NULL THEN 'system' ELSE 'custom' END
WHERE rate_limit_mode = 'system';
CREATE TABLE IF NOT EXISTS user_groups (
id TEXT PRIMARY KEY,
name TEXT NOT NULL,
normalized_name TEXT NOT NULL UNIQUE,
description TEXT,
priority INTEGER NOT NULL DEFAULT 0,
allowed_providers TEXT,
allowed_providers_mode TEXT NOT NULL DEFAULT 'inherit',
allowed_api_formats TEXT,
allowed_api_formats_mode TEXT NOT NULL DEFAULT 'inherit',
allowed_models TEXT,
allowed_models_mode TEXT NOT NULL DEFAULT 'inherit',
rate_limit INTEGER,
rate_limit_mode TEXT NOT NULL DEFAULT 'inherit',
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS user_group_members (
group_id TEXT NOT NULL REFERENCES user_groups(id) ON DELETE CASCADE,
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
created_at INTEGER NOT NULL,
PRIMARY KEY (group_id, user_id)
);
CREATE INDEX IF NOT EXISTS user_group_members_user_id_idx
ON user_group_members (user_id);
CREATE INDEX IF NOT EXISTS user_groups_priority_name_idx
ON user_groups (priority DESC, name ASC, id ASC);

View File

@@ -1238,8 +1238,11 @@ CREATE TABLE IF NOT EXISTS public.users (
password_hash character varying(255), password_hash character varying(255),
role public.userrole DEFAULT 'user'::public.userrole NOT NULL, role public.userrole DEFAULT 'user'::public.userrole NOT NULL,
allowed_providers json, allowed_providers json,
allowed_providers_mode text DEFAULT 'unrestricted'::text NOT NULL,
allowed_api_formats json, allowed_api_formats json,
allowed_api_formats_mode text DEFAULT 'unrestricted'::text NOT NULL,
allowed_models json, allowed_models json,
allowed_models_mode text DEFAULT 'unrestricted'::text NOT NULL,
model_capability_settings json, model_capability_settings json,
is_active boolean DEFAULT true NOT NULL, is_active boolean DEFAULT true NOT NULL,
is_deleted boolean DEFAULT false NOT NULL, is_deleted boolean DEFAULT false NOT NULL,
@@ -1251,11 +1254,48 @@ CREATE TABLE IF NOT EXISTS public.users (
ldap_username character varying(255), ldap_username character varying(255),
email_verified boolean NOT NULL, email_verified boolean NOT NULL,
rate_limit integer, rate_limit integer,
rate_limit_mode text DEFAULT 'system'::text NOT NULL,
metadata json metadata json
); );
--
-- Name: user_groups; Type: TABLE; Schema: public; Owner: -
--
CREATE TABLE IF NOT EXISTS public.user_groups (
id character varying(36) NOT NULL,
name character varying(100) NOT NULL,
normalized_name character varying(100) NOT NULL,
description text,
priority integer DEFAULT 0 NOT NULL,
allowed_providers json,
allowed_providers_mode text DEFAULT 'inherit'::text NOT NULL,
allowed_api_formats json,
allowed_api_formats_mode text DEFAULT 'inherit'::text NOT NULL,
allowed_models json,
allowed_models_mode text DEFAULT 'inherit'::text NOT NULL,
rate_limit integer,
rate_limit_mode text DEFAULT 'inherit'::text NOT NULL,
created_at timestamp with time zone DEFAULT now() NOT NULL,
updated_at timestamp with time zone DEFAULT now() NOT NULL
);
--
-- Name: user_group_members; Type: TABLE; Schema: public; Owner: -
--
CREATE TABLE IF NOT EXISTS public.user_group_members (
group_id character varying(36) NOT NULL,
user_id character varying(36) NOT NULL,
created_at timestamp with time zone DEFAULT now() NOT NULL
);
-- --
-- Name: video_tasks; Type: TABLE; Schema: public; Owner: - -- Name: video_tasks; Type: TABLE; Schema: public; Owner: -
-- --

View File

@@ -972,6 +972,111 @@ END $mig$;
--
-- Name: user_group_members user_group_members_pkey; Type: CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_group_members
ADD CONSTRAINT user_group_members_pkey PRIMARY KEY (group_id, user_id);
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_groups user_groups_allowed_api_formats_mode_check; Type: CHECK CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_groups
ADD CONSTRAINT user_groups_allowed_api_formats_mode_check CHECK (allowed_api_formats_mode IN ('inherit', 'unrestricted', 'specific', 'deny_all'));
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_groups user_groups_allowed_models_mode_check; Type: CHECK CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_groups
ADD CONSTRAINT user_groups_allowed_models_mode_check CHECK (allowed_models_mode IN ('inherit', 'unrestricted', 'specific', 'deny_all'));
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_groups user_groups_allowed_providers_mode_check; Type: CHECK CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_groups
ADD CONSTRAINT user_groups_allowed_providers_mode_check CHECK (allowed_providers_mode IN ('inherit', 'unrestricted', 'specific', 'deny_all'));
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_groups user_groups_normalized_name_key; Type: CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_groups
ADD CONSTRAINT user_groups_normalized_name_key UNIQUE (normalized_name);
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_groups user_groups_pkey; Type: CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_groups
ADD CONSTRAINT user_groups_pkey PRIMARY KEY (id);
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_groups user_groups_rate_limit_mode_check; Type: CHECK CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_groups
ADD CONSTRAINT user_groups_rate_limit_mode_check CHECK (rate_limit_mode IN ('inherit', 'system', 'custom'));
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
-- --
-- Name: user_oauth_links user_oauth_links_pkey; Type: CONSTRAINT; Schema: public; Owner: - -- Name: user_oauth_links user_oauth_links_pkey; Type: CONSTRAINT; Schema: public; Owner: -
-- --

View File

@@ -1021,6 +1021,22 @@ CREATE INDEX IF NOT EXISTS ix_system_configs_id ON public.system_configs USING b
--
-- Name: user_group_members_user_id_idx; Type: INDEX; Schema: public; Owner: -
--
CREATE INDEX IF NOT EXISTS user_group_members_user_id_idx ON public.user_group_members USING btree (user_id);
--
-- Name: user_groups_priority_name_idx; Type: INDEX; Schema: public; Owner: -
--
CREATE INDEX IF NOT EXISTS user_groups_priority_name_idx ON public.user_groups USING btree (priority DESC, name, id);
-- --
-- Name: ix_usage_created_at; Type: INDEX; Schema: public; Owner: - -- Name: ix_usage_created_at; Type: INDEX; Schema: public; Owner: -
-- --

View File

@@ -612,6 +612,36 @@ END $mig$;
--
-- Name: user_group_members user_group_members_group_id_fk; Type: FK CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_group_members
ADD CONSTRAINT user_group_members_group_id_fk FOREIGN KEY (group_id) REFERENCES public.user_groups(id) ON DELETE CASCADE;
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_group_members user_group_members_user_id_fk; Type: FK CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_group_members
ADD CONSTRAINT user_group_members_user_id_fk FOREIGN KEY (user_id) REFERENCES public.users(id) ON DELETE CASCADE;
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
-- --
-- Name: user_oauth_links user_oauth_links_provider_type_fkey; Type: FK CONSTRAINT; Schema: public; Owner: - -- Name: user_oauth_links user_oauth_links_provider_type_fkey; Type: FK CONSTRAINT; Schema: public; Owner: -
-- --

View File

@@ -13,10 +13,14 @@ CREATE TABLE IF NOT EXISTS users (
`is_active` TINYINT(1) NOT NULL DEFAULT 1, `is_active` TINYINT(1) NOT NULL DEFAULT 1,
`is_deleted` TINYINT(1) NOT NULL DEFAULT 0, `is_deleted` TINYINT(1) NOT NULL DEFAULT 0,
`allowed_models` JSON, `allowed_models` JSON,
`allowed_models_mode` VARCHAR(32) NOT NULL DEFAULT 'unrestricted',
`allowed_providers` JSON, `allowed_providers` JSON,
`allowed_providers_mode` VARCHAR(32) NOT NULL DEFAULT 'unrestricted',
`allowed_api_formats` JSON, `allowed_api_formats` JSON,
`allowed_api_formats_mode` VARCHAR(32) NOT NULL DEFAULT 'unrestricted',
`model_capability_settings` JSON, `model_capability_settings` JSON,
`rate_limit` INT, `rate_limit` INT,
`rate_limit_mode` VARCHAR(32) NOT NULL DEFAULT 'system',
`metadata` JSON, `metadata` JSON,
`created_at` BIGINT NOT NULL, `created_at` BIGINT NOT NULL,
`updated_at` BIGINT NOT NULL, `updated_at` BIGINT NOT NULL,
@@ -28,6 +32,35 @@ CREATE TABLE IF NOT EXISTS users (
UNIQUE KEY users_username_key (`username`) UNIQUE KEY users_username_key (`username`)
); );
CREATE TABLE IF NOT EXISTS user_groups (
`id` VARCHAR(64) NOT NULL,
`name` VARCHAR(100) NOT NULL,
`normalized_name` VARCHAR(100) NOT NULL,
`description` LONGTEXT,
`priority` INT NOT NULL DEFAULT 0,
`allowed_providers` JSON,
`allowed_providers_mode` VARCHAR(32) NOT NULL DEFAULT 'inherit',
`allowed_api_formats` JSON,
`allowed_api_formats_mode` VARCHAR(32) NOT NULL DEFAULT 'inherit',
`allowed_models` JSON,
`allowed_models_mode` VARCHAR(32) NOT NULL DEFAULT 'inherit',
`rate_limit` INT,
`rate_limit_mode` VARCHAR(32) NOT NULL DEFAULT 'inherit',
`created_at` BIGINT NOT NULL,
`updated_at` BIGINT NOT NULL,
PRIMARY KEY (`id`),
UNIQUE KEY user_groups_normalized_name_key (`normalized_name`),
KEY user_groups_priority_name_idx (`priority`, `name`, `id`)
);
CREATE TABLE IF NOT EXISTS user_group_members (
`group_id` VARCHAR(64) NOT NULL,
`user_id` VARCHAR(64) NOT NULL,
`created_at` BIGINT NOT NULL,
PRIMARY KEY (`group_id`, `user_id`),
KEY user_group_members_user_id_idx (`user_id`)
);
CREATE TABLE IF NOT EXISTS api_keys ( CREATE TABLE IF NOT EXISTS api_keys (
`id` VARCHAR(64) NOT NULL, `id` VARCHAR(64) NOT NULL,
`user_id` VARCHAR(64) NOT NULL, `user_id` VARCHAR(64) NOT NULL,

View File

@@ -13,10 +13,14 @@ CREATE TABLE IF NOT EXISTS public.users (
is_active boolean DEFAULT true NOT NULL, is_active boolean DEFAULT true NOT NULL,
is_deleted boolean DEFAULT false NOT NULL, is_deleted boolean DEFAULT false NOT NULL,
allowed_models jsonb, allowed_models jsonb,
allowed_models_mode character varying(32) DEFAULT 'unrestricted' NOT NULL,
allowed_providers jsonb, allowed_providers jsonb,
allowed_providers_mode character varying(32) DEFAULT 'unrestricted' NOT NULL,
allowed_api_formats jsonb, allowed_api_formats jsonb,
allowed_api_formats_mode character varying(32) DEFAULT 'unrestricted' NOT NULL,
model_capability_settings jsonb, model_capability_settings jsonb,
rate_limit integer, rate_limit integer,
rate_limit_mode character varying(32) DEFAULT 'system' NOT NULL,
metadata jsonb, metadata jsonb,
created_at bigint NOT NULL, created_at bigint NOT NULL,
updated_at bigint NOT NULL, updated_at bigint NOT NULL,
@@ -29,6 +33,37 @@ ALTER TABLE ONLY public.users ADD CONSTRAINT users_pkey PRIMARY KEY (id);
ALTER TABLE ONLY public.users ADD CONSTRAINT users_email_key UNIQUE (email); ALTER TABLE ONLY public.users ADD CONSTRAINT users_email_key UNIQUE (email);
ALTER TABLE ONLY public.users ADD CONSTRAINT users_username_key UNIQUE (username); ALTER TABLE ONLY public.users ADD CONSTRAINT users_username_key UNIQUE (username);
CREATE TABLE IF NOT EXISTS public.user_groups (
id character varying(64) NOT NULL,
name character varying(100) NOT NULL,
normalized_name character varying(100) NOT NULL,
description text,
priority integer DEFAULT 0 NOT NULL,
allowed_providers jsonb,
allowed_providers_mode character varying(32) DEFAULT 'inherit' NOT NULL,
allowed_api_formats jsonb,
allowed_api_formats_mode character varying(32) DEFAULT 'inherit' NOT NULL,
allowed_models jsonb,
allowed_models_mode character varying(32) DEFAULT 'inherit' NOT NULL,
rate_limit integer,
rate_limit_mode character varying(32) DEFAULT 'inherit' NOT NULL,
created_at bigint NOT NULL,
updated_at bigint NOT NULL
);
ALTER TABLE ONLY public.user_groups ADD CONSTRAINT user_groups_pkey PRIMARY KEY (id);
ALTER TABLE ONLY public.user_groups ADD CONSTRAINT user_groups_normalized_name_key UNIQUE (normalized_name);
CREATE INDEX IF NOT EXISTS user_groups_priority_name_idx ON public.user_groups USING btree (priority, name, id);
CREATE TABLE IF NOT EXISTS public.user_group_members (
group_id character varying(64) NOT NULL,
user_id character varying(64) NOT NULL,
created_at bigint NOT NULL
);
ALTER TABLE ONLY public.user_group_members ADD CONSTRAINT user_group_members_pkey PRIMARY KEY (group_id, user_id);
CREATE INDEX IF NOT EXISTS user_group_members_user_id_idx ON public.user_group_members USING btree (user_id);
CREATE TABLE IF NOT EXISTS public.api_keys ( CREATE TABLE IF NOT EXISTS public.api_keys (
id character varying(64) NOT NULL, id character varying(64) NOT NULL,
user_id character varying(64) NOT NULL, user_id character varying(64) NOT NULL,

View File

@@ -13,10 +13,14 @@ CREATE TABLE IF NOT EXISTS users (
is_active INTEGER NOT NULL DEFAULT 1, is_active INTEGER NOT NULL DEFAULT 1,
is_deleted INTEGER NOT NULL DEFAULT 0, is_deleted INTEGER NOT NULL DEFAULT 0,
allowed_models TEXT, allowed_models TEXT,
allowed_models_mode TEXT NOT NULL DEFAULT 'unrestricted',
allowed_providers TEXT, allowed_providers TEXT,
allowed_providers_mode TEXT NOT NULL DEFAULT 'unrestricted',
allowed_api_formats TEXT, allowed_api_formats TEXT,
allowed_api_formats_mode TEXT NOT NULL DEFAULT 'unrestricted',
model_capability_settings TEXT, model_capability_settings TEXT,
rate_limit INTEGER, rate_limit INTEGER,
rate_limit_mode TEXT NOT NULL DEFAULT 'system',
metadata TEXT, metadata TEXT,
created_at INTEGER NOT NULL, created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL, updated_at INTEGER NOT NULL,
@@ -27,6 +31,34 @@ CREATE TABLE IF NOT EXISTS users (
UNIQUE (username) UNIQUE (username)
); );
CREATE TABLE IF NOT EXISTS user_groups (
id TEXT PRIMARY KEY NOT NULL,
name TEXT NOT NULL,
normalized_name TEXT NOT NULL,
description TEXT,
priority INTEGER NOT NULL DEFAULT 0,
allowed_providers TEXT,
allowed_providers_mode TEXT NOT NULL DEFAULT 'inherit',
allowed_api_formats TEXT,
allowed_api_formats_mode TEXT NOT NULL DEFAULT 'inherit',
allowed_models TEXT,
allowed_models_mode TEXT NOT NULL DEFAULT 'inherit',
rate_limit INTEGER,
rate_limit_mode TEXT NOT NULL DEFAULT 'inherit',
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
UNIQUE (normalized_name)
);
CREATE INDEX IF NOT EXISTS user_groups_priority_name_idx ON user_groups (priority, name, id);
CREATE TABLE IF NOT EXISTS user_group_members (
group_id TEXT NOT NULL,
user_id TEXT NOT NULL,
created_at INTEGER NOT NULL,
PRIMARY KEY (group_id, user_id)
);
CREATE INDEX IF NOT EXISTS user_group_members_user_id_idx ON user_group_members (user_id);
CREATE TABLE IF NOT EXISTS api_keys ( CREATE TABLE IF NOT EXISTS api_keys (
id TEXT PRIMARY KEY NOT NULL, id TEXT PRIMARY KEY NOT NULL,
user_id TEXT NOT NULL, user_id TEXT NOT NULL,

View File

@@ -64,16 +64,34 @@ name = "allowed_models"
type = "json" type = "json"
nullable = true nullable = true
[[table.users.columns]]
name = "allowed_models_mode"
type = "text"
length = 32
default = "unrestricted"
[[table.users.columns]] [[table.users.columns]]
name = "allowed_providers" name = "allowed_providers"
type = "json" type = "json"
nullable = true nullable = true
[[table.users.columns]]
name = "allowed_providers_mode"
type = "text"
length = 32
default = "unrestricted"
[[table.users.columns]] [[table.users.columns]]
name = "allowed_api_formats" name = "allowed_api_formats"
type = "json" type = "json"
nullable = true nullable = true
[[table.users.columns]]
name = "allowed_api_formats_mode"
type = "text"
length = 32
default = "unrestricted"
[[table.users.columns]] [[table.users.columns]]
name = "model_capability_settings" name = "model_capability_settings"
type = "json" type = "json"
@@ -84,6 +102,12 @@ name = "rate_limit"
type = "int32" type = "int32"
nullable = true nullable = true
[[table.users.columns]]
name = "rate_limit_mode"
type = "text"
length = 32
default = "system"
[[table.users.columns]] [[table.users.columns]]
name = "metadata" name = "metadata"
type = "json" type = "json"
@@ -122,6 +146,119 @@ columns = ["email"]
name = "users_username_key" name = "users_username_key"
columns = ["username"] columns = ["username"]
[table.user_groups]
domain = "identity"
order = 15
primary_key = ["id"]
[[table.user_groups.columns]]
name = "id"
type = "text_id"
length = 64
[[table.user_groups.columns]]
name = "name"
type = "text"
length = 100
[[table.user_groups.columns]]
name = "normalized_name"
type = "text"
length = 100
[[table.user_groups.columns]]
name = "description"
type = "long_text"
nullable = true
[[table.user_groups.columns]]
name = "priority"
type = "int32"
default = 0
[[table.user_groups.columns]]
name = "allowed_providers"
type = "json"
nullable = true
[[table.user_groups.columns]]
name = "allowed_providers_mode"
type = "text"
length = 32
default = "inherit"
[[table.user_groups.columns]]
name = "allowed_api_formats"
type = "json"
nullable = true
[[table.user_groups.columns]]
name = "allowed_api_formats_mode"
type = "text"
length = 32
default = "inherit"
[[table.user_groups.columns]]
name = "allowed_models"
type = "json"
nullable = true
[[table.user_groups.columns]]
name = "allowed_models_mode"
type = "text"
length = 32
default = "inherit"
[[table.user_groups.columns]]
name = "rate_limit"
type = "int32"
nullable = true
[[table.user_groups.columns]]
name = "rate_limit_mode"
type = "text"
length = 32
default = "inherit"
[[table.user_groups.columns]]
name = "created_at"
type = "unix_seconds"
[[table.user_groups.columns]]
name = "updated_at"
type = "unix_seconds"
[[table.user_groups.uniques]]
name = "user_groups_normalized_name_key"
columns = ["normalized_name"]
[[table.user_groups.indexes]]
name = "user_groups_priority_name_idx"
columns = ["priority", "name", "id"]
[table.user_group_members]
domain = "identity"
order = 16
primary_key = ["group_id", "user_id"]
[[table.user_group_members.columns]]
name = "group_id"
type = "text_id"
length = 64
[[table.user_group_members.columns]]
name = "user_id"
type = "text_id"
length = 64
[[table.user_group_members.columns]]
name = "created_at"
type = "unix_seconds"
[[table.user_group_members.indexes]]
name = "user_group_members_user_id_idx"
columns = ["user_id"]
[table.api_keys] [table.api_keys]
domain = "identity" domain = "identity"
order = 20 order = 20

View File

@@ -7,7 +7,7 @@ use tracing::info;
// Generated by build.rs from schema/bootstrap/postgres. // Generated by build.rs from schema/bootstrap/postgres.
pub(crate) static EMPTY_DATABASE_SNAPSHOT_SQL: &str = pub(crate) static EMPTY_DATABASE_SNAPSHOT_SQL: &str =
include_str!(concat!(env!("OUT_DIR"), "/empty_database_snapshot.sql")); include_str!(concat!(env!("OUT_DIR"), "/empty_database_snapshot.sql"));
pub(crate) const EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION: i64 = 20260509000000; pub(crate) const EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION: i64 = 20260509120000;
const PUBLIC_BASE_TABLE_COUNT_SQL: &str = r#" const PUBLIC_BASE_TABLE_COUNT_SQL: &str = r#"
SELECT COUNT(*)::BIGINT SELECT COUNT(*)::BIGINT
@@ -27,11 +27,12 @@ WHERE table_schema = 'public'
'gemini_file_mappings', 'gemini_file_mappings',
'global_models', 'global_models',
'oauth_providers', 'oauth_providers',
'provider_api_keys', 'provider_api_keys',
'proxy_nodes', 'proxy_nodes',
'usage_routing_snapshots', 'user_groups',
'usage_settlement_snapshots' 'usage_routing_snapshots',
) 'usage_settlement_snapshots'
)
"#; "#;
const INSERT_APPLIED_MIGRATION_SQL: &str = r#" const INSERT_APPLIED_MIGRATION_SQL: &str = r#"
INSERT INTO _sqlx_migrations ( INSERT INTO _sqlx_migrations (

View File

@@ -295,6 +295,7 @@ fn empty_database_snapshot_covers_current_cutoff_versions() {
20260507120000, 20260507120000,
20260508000000, 20260508000000,
20260509000000, 20260509000000,
20260509120000,
] ]
); );
} }
@@ -518,7 +519,8 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() {
20260403000000, 20260403000000,
20260507120000, 20260507120000,
20260508000000, 20260508000000,
20260509000000 20260509000000,
20260509120000
] ]
); );
assert_eq!( assert_eq!(
@@ -527,7 +529,8 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() {
20260403000000, 20260403000000,
20260507120000, 20260507120000,
20260508000000, 20260508000000,
20260509000000 20260509000000,
20260509120000
] ]
); );
} }
@@ -1034,6 +1037,7 @@ fn pending_migrations_from_applied_skips_versions_already_applied() {
20260507120000, 20260507120000,
20260508000000, 20260508000000,
20260509000000, 20260509000000,
20260509120000,
] ]
); );
} }

View File

@@ -187,7 +187,7 @@ impl ResolvedAuthApiKeySnapshot {
} }
non_empty_allowed_list(self.api_key_allowed_providers.as_deref()) 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]> { 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()) 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]> { 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()) 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 async_trait::async_trait;
use super::types::{ use super::types::{
LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, StoredUserExportRow, normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord,
StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership,
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord, StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
StoredUserSummary, UserExportListQuery, UserExportSummary, UserReadRepository, StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSummary,
UserReadRepository,
}; };
use crate::DataLayerError; use crate::DataLayerError;
@@ -34,6 +36,8 @@ pub struct InMemoryUserReadRepository {
preferences_by_user_id: RwLock<BTreeMap<String, StoredUserPreferenceRecord>>, preferences_by_user_id: RwLock<BTreeMap<String, StoredUserPreferenceRecord>>,
sessions_by_id: RwLock<BTreeMap<String, StoredUserSessionRecord>>, sessions_by_id: RwLock<BTreeMap<String, StoredUserSessionRecord>>,
model_settings_by_user_id: RwLock<BTreeMap<String, serde_json::Value>>, 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>>, export_rows: RwLock<Vec<StoredUserExportRow>>,
read_only: bool, read_only: bool,
} }
@@ -57,6 +61,8 @@ impl InMemoryUserReadRepository {
preferences_by_user_id: RwLock::new(BTreeMap::new()), preferences_by_user_id: RwLock::new(BTreeMap::new()),
sessions_by_id: RwLock::new(BTreeMap::new()), sessions_by_id: RwLock::new(BTreeMap::new()),
model_settings_by_user_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()), export_rows: RwLock::new(Vec::new()),
read_only: false, read_only: false,
} }
@@ -90,6 +96,8 @@ impl InMemoryUserReadRepository {
preferences_by_user_id: RwLock::new(BTreeMap::new()), preferences_by_user_id: RwLock::new(BTreeMap::new()),
sessions_by_id: RwLock::new(BTreeMap::new()), sessions_by_id: RwLock::new(BTreeMap::new()),
model_settings_by_user_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()), export_rows: RwLock::new(Vec::new()),
read_only: false, read_only: false,
} }
@@ -109,6 +117,8 @@ impl InMemoryUserReadRepository {
preferences_by_user_id: RwLock::new(BTreeMap::new()), preferences_by_user_id: RwLock::new(BTreeMap::new()),
sessions_by_id: RwLock::new(BTreeMap::new()), sessions_by_id: RwLock::new(BTreeMap::new()),
model_settings_by_user_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()), export_rows: RwLock::new(items.into_iter().collect()),
read_only: false, 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] #[async_trait]
impl UserReadRepository for InMemoryUserReadRepository { impl UserReadRepository for InMemoryUserReadRepository {
async fn list_users_by_ids( async fn list_users_by_ids(
@@ -329,6 +428,23 @@ impl UserReadRepository for InMemoryUserReadRepository {
if let Some(is_active) = query.is_active { if let Some(is_active) = query.is_active {
rows.retain(|row| row.is_active == 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 if let Some(search) = query
.search .search
.as_deref() .as_deref()
@@ -375,6 +491,273 @@ impl UserReadRepository for InMemoryUserReadRepository {
.cloned()) .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( async fn find_user_auth_by_id(
&self, &self,
user_id: &str, user_id: &str,
@@ -581,6 +964,11 @@ impl UserReadRepository for InMemoryUserReadRepository {
false, false,
Some(created_at), Some(created_at),
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) self.insert_auth_user(user).map(Some)
} }
@@ -917,12 +1305,27 @@ impl UserReadRepository for InMemoryUserReadRepository {
} }
if allowed_providers_present { if allowed_providers_present {
user.allowed_providers = allowed_providers; 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 { if allowed_api_formats_present {
user.allowed_api_formats = allowed_api_formats; 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 { if allowed_models_present {
user.allowed_models = allowed_models; 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 { if let Some(is_active) = is_active {
user.is_active = is_active; user.is_active = is_active;
@@ -948,16 +1351,75 @@ impl UserReadRepository for InMemoryUserReadRepository {
{ {
row.role = updated.role.clone(); row.role = updated.role.clone();
row.allowed_providers = updated.allowed_providers.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 = 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 = updated.allowed_models.clone();
row.allowed_models_mode = updated.allowed_models_mode.clone();
if rate_limit_present { if rate_limit_present {
row.rate_limit = rate_limit; 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; row.is_active = updated.is_active;
} }
Ok(Some(updated)) 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( async fn update_user_model_capability_settings(
&self, &self,
user_id: &str, user_id: &str,
@@ -1021,18 +1483,29 @@ impl UserReadRepository for InMemoryUserReadRepository {
return Ok(None); 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,
email_verified, email_verified,
username, username,
password_hash, Some(password_hash),
"user".to_string(), "user".to_string(),
"local".to_string(),
None, None,
None, None,
None, None,
true,
false,
Some(now),
None, 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( async fn create_local_auth_user_with_settings(
@@ -1092,6 +1565,10 @@ impl UserReadRepository for InMemoryUserReadRepository {
.write() .write()
.expect("user repository lock") .expect("user repository lock")
.retain(|_, link| link.user_id != user_id); .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 let mut identifiers = self
.auth_by_identifier .auth_by_identifier
@@ -2073,6 +2550,7 @@ mod tests {
role: Some("user".to_string()), role: Some("user".to_string()),
is_active: Some(true), is_active: Some(true),
search: None, search: None,
group_id: None,
}) })
.await .await
.expect("paged export should succeed"); .expect("paged export should succeed");

View File

@@ -9,7 +9,8 @@ pub use mysql::MysqlUserReadRepository;
pub use postgres::SqlxUserReadRepository; pub use postgres::SqlxUserReadRepository;
pub use sqlite::SqliteUserReadRepository; pub use sqlite::SqliteUserReadRepository;
pub use types::{ pub use types::{
StoredUserAuthRecord, StoredUserExportRow, StoredUserOAuthLinkSummary, normalize_user_group_name, StoredUserAuthRecord, StoredUserExportRow, StoredUserGroup,
StoredUserPreferenceRecord, StoredUserSessionRecord, StoredUserSummary, UserExportListQuery, StoredUserGroupMember, StoredUserGroupMembership, StoredUserOAuthLinkSummary,
UserExportSummary, UserReadRepository, 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 sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use super::types::{ use super::types::{
LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, StoredUserExportRow, normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord,
StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership,
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord, StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
StoredUserSummary, UserExportListQuery, UserExportSummary, UserReadRepository, StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSummary,
UserReadRepository,
}; };
use crate::driver::mysql::MysqlPool; use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt; use crate::error::SqlResultExt;
@@ -32,9 +34,13 @@ SELECT
role, role,
auth_source, auth_source,
allowed_providers, allowed_providers,
allowed_providers_mode,
allowed_api_formats, allowed_api_formats,
allowed_api_formats_mode,
allowed_models, allowed_models,
allowed_models_mode,
rate_limit, rate_limit,
rate_limit_mode,
model_capability_settings, model_capability_settings,
is_active is_active
FROM users FROM users
@@ -50,8 +56,11 @@ SELECT
role, role,
auth_source, auth_source,
allowed_providers, allowed_providers,
allowed_providers_mode,
allowed_api_formats, allowed_api_formats,
allowed_api_formats_mode,
allowed_models, allowed_models,
allowed_models_mode,
is_active, is_active,
is_deleted, is_deleted,
created_at, created_at,
@@ -69,8 +78,11 @@ SELECT
users.role AS role, users.role AS role,
users.auth_source AS auth_source, users.auth_source AS auth_source,
users.allowed_providers AS allowed_providers, 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 AS allowed_api_formats,
users.allowed_api_formats_mode AS allowed_api_formats_mode,
users.allowed_models AS allowed_models, users.allowed_models AS allowed_models,
users.allowed_models_mode AS allowed_models_mode,
users.is_active AS is_active, users.is_active AS is_active,
users.is_deleted AS is_deleted, users.is_deleted AS is_deleted,
users.created_at AS created_at, users.created_at AS created_at,
@@ -130,6 +142,40 @@ SELECT
FROM user_sessions 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)] #[derive(Debug, Clone)]
pub struct MysqlUserReadRepository { pub struct MysqlUserReadRepository {
pool: MysqlPool, pool: MysqlPool,
@@ -163,6 +209,22 @@ impl MysqlUserReadRepository {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_auth_row).collect() 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] #[async_trait]
@@ -224,6 +286,16 @@ impl UserReadRepository for MysqlUserReadRepository {
if let Some(is_active) = query.is_active { if let Some(is_active) = query.is_active {
builder.push(" AND is_active = ").push_bind(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 if let Some(search) = query
.search .search
.as_deref() .as_deref()
@@ -296,6 +368,285 @@ WHERE is_deleted = 0
self.fetch_export_rows(builder).await 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( async fn find_user_auth_by_id(
&self, &self,
user_id: &str, user_id: &str,
@@ -453,9 +804,10 @@ WHERE provider_type = ?
r#" r#"
INSERT INTO users ( INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source, 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 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) .bind(&user_id)
@@ -675,14 +1027,38 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
rate_limit: Option<i32>, rate_limit: Option<i32>,
is_active: Option<bool>, is_active: Option<bool>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> { ) -> 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( let result = sqlx::query(
r#" r#"
UPDATE users UPDATE users
SET role = CASE WHEN ? THEN COALESCE(?, role) ELSE role END, SET role = CASE WHEN ? THEN COALESCE(?, role) ELSE role END,
allowed_providers = CASE WHEN ? THEN ? ELSE allowed_providers 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 = 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 = 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 = 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, is_active = CASE WHEN ? THEN ? ELSE is_active END,
updated_at = ? updated_at = ?
WHERE id = ? WHERE id = ?
@@ -695,18 +1071,26 @@ WHERE id = ?
allowed_providers, allowed_providers,
"users.allowed_providers", "users.allowed_providers",
)?) )?)
.bind(allowed_providers_present)
.bind(allowed_providers_mode)
.bind(allowed_api_formats_present) .bind(allowed_api_formats_present)
.bind(optional_string_list_json( .bind(optional_string_list_json(
allowed_api_formats, allowed_api_formats,
"users.allowed_api_formats", "users.allowed_api_formats",
)?) )?)
.bind(allowed_api_formats_present)
.bind(allowed_api_formats_mode)
.bind(allowed_models_present) .bind(allowed_models_present)
.bind(optional_string_list_json( .bind(optional_string_list_json(
allowed_models, allowed_models,
"users.allowed_models", "users.allowed_models",
)?) )?)
.bind(allowed_models_present)
.bind(allowed_models_mode)
.bind(rate_limit_present) .bind(rate_limit_present)
.bind(rate_limit) .bind(rate_limit)
.bind(rate_limit_present)
.bind(rate_limit_mode)
.bind(is_active.is_some()) .bind(is_active.is_some())
.bind(is_active) .bind(is_active)
.bind(chrono::Utc::now().timestamp()) .bind(chrono::Utc::now().timestamp())
@@ -720,6 +1104,44 @@ WHERE id = ?
self.find_user_auth_by_id(user_id).await 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( async fn update_user_model_capability_settings(
&self, &self,
user_id: &str, user_id: &str,
@@ -751,18 +1173,29 @@ WHERE id = ?
username: String, username: String,
password_hash: String, password_hash: String,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> { ) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
self.create_local_auth_user_with_settings( let user_id = uuid::Uuid::new_v4().to_string();
email, let now = chrono::Utc::now().timestamp();
email_verified, sqlx::query(
username, r#"
password_hash, INSERT INTO users (
"user".to_string(), id, email, email_verified, username, password_hash, role, auth_source,
None, allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
None, is_active, is_deleted, created_at, updated_at
None, )
None, 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 .await
.map_sql_err()?;
self.find_user_auth_by_id(&user_id).await
} }
async fn create_local_auth_user_with_settings( async fn create_local_auth_user_with_settings(
@@ -779,14 +1212,37 @@ WHERE id = ?
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> { ) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let user_id = uuid::Uuid::new_v4().to_string(); let user_id = uuid::Uuid::new_v4().to_string();
let now = chrono::Utc::now().timestamp(); 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( sqlx::query(
r#" r#"
INSERT INTO users ( INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source, 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 is_active, is_deleted, created_at, updated_at
) )
VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, 1, 0, ?, ?) VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, ?, ?, ?, ?, 1, 0, ?, ?)
"#, "#,
) )
.bind(&user_id) .bind(&user_id)
@@ -799,15 +1255,19 @@ VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, 1, 0, ?, ?)
allowed_providers, allowed_providers,
"users.allowed_providers", "users.allowed_providers",
)?) )?)
.bind(allowed_providers_mode)
.bind(optional_string_list_json( .bind(optional_string_list_json(
allowed_api_formats, allowed_api_formats,
"users.allowed_api_formats", "users.allowed_api_formats",
)?) )?)
.bind(allowed_api_formats_mode)
.bind(optional_string_list_json( .bind(optional_string_list_json(
allowed_models, allowed_models,
"users.allowed_models", "users.allowed_models",
)?) )?)
.bind(allowed_models_mode)
.bind(rate_limit) .bind(rate_limit)
.bind(rate_limit_mode)
.bind(now) .bind(now)
.bind(now) .bind(now)
.execute(&self.pool) .execute(&self.pool)
@@ -1162,6 +1622,24 @@ fn optional_string_list_json(
.transpose() .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( fn optional_json_string(
value: Option<serde_json::Value>, value: Option<serde_json::Value>,
field_name: &str, 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()?, 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> { 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("created_at").map_sql_err()?),
optional_datetime_from_unix_secs(row.try_get("last_login_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( fn map_oauth_link_summary_row(

View File

@@ -3,9 +3,11 @@ use futures_util::TryStreamExt;
use sqlx::{PgPool, Postgres, QueryBuilder, Row}; use sqlx::{PgPool, Postgres, QueryBuilder, Row};
use super::types::{ use super::types::{
LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, StoredUserExportRow, normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord,
StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership,
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord, StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
StoredUserSummary, UserExportListQuery, UserExportSummary, UserReadRepository, StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSummary,
UserReadRepository,
}; };
use crate::{error::SqlxResultExt, DataLayerError}; use crate::{error::SqlxResultExt, DataLayerError};
@@ -46,9 +48,13 @@ SELECT
role::text AS role, role::text AS role,
auth_source::text AS auth_source, auth_source::text AS auth_source,
allowed_providers, allowed_providers,
allowed_providers_mode,
allowed_api_formats, allowed_api_formats,
allowed_api_formats_mode,
allowed_models, allowed_models,
allowed_models_mode,
rate_limit, rate_limit,
rate_limit_mode,
model_capability_settings, model_capability_settings,
is_active is_active
FROM users FROM users
@@ -67,9 +73,13 @@ SELECT
role::text AS role, role::text AS role,
auth_source::text AS auth_source, auth_source::text AS auth_source,
allowed_providers, allowed_providers,
allowed_providers_mode,
allowed_api_formats, allowed_api_formats,
allowed_api_formats_mode,
allowed_models, allowed_models,
allowed_models_mode,
rate_limit, rate_limit,
rate_limit_mode,
model_capability_settings, model_capability_settings,
is_active is_active
FROM users FROM users
@@ -87,9 +97,13 @@ SELECT
role::text AS role, role::text AS role,
auth_source::text AS auth_source, auth_source::text AS auth_source,
allowed_providers, allowed_providers,
allowed_providers_mode,
allowed_api_formats, allowed_api_formats,
allowed_api_formats_mode,
allowed_models, allowed_models,
allowed_models_mode,
rate_limit, rate_limit,
rate_limit_mode,
model_capability_settings, model_capability_settings,
is_active is_active
FROM users FROM users
@@ -132,9 +146,13 @@ SELECT
role::text AS role, role::text AS role,
auth_source::text AS auth_source, auth_source::text AS auth_source,
allowed_providers, allowed_providers,
allowed_providers_mode,
allowed_api_formats, allowed_api_formats,
allowed_api_formats_mode,
allowed_models, allowed_models,
allowed_models_mode,
rate_limit, rate_limit,
rate_limit_mode,
model_capability_settings, model_capability_settings,
is_active is_active
FROM users FROM users
@@ -153,8 +171,11 @@ SELECT
role::text AS role, role::text AS role,
auth_source::text AS auth_source, auth_source::text AS auth_source,
allowed_providers, allowed_providers,
allowed_providers_mode,
allowed_api_formats, allowed_api_formats,
allowed_api_formats_mode,
allowed_models, allowed_models,
allowed_models_mode,
is_active, is_active,
is_deleted, is_deleted,
created_at, created_at,
@@ -174,8 +195,11 @@ SELECT
role::text AS role, role::text AS role,
auth_source::text AS auth_source, auth_source::text AS auth_source,
allowed_providers, allowed_providers,
allowed_providers_mode,
allowed_api_formats, allowed_api_formats,
allowed_api_formats_mode,
allowed_models, allowed_models,
allowed_models_mode,
is_active, is_active,
is_deleted, is_deleted,
created_at, created_at,
@@ -195,8 +219,11 @@ SELECT
role::text AS role, role::text AS role,
auth_source::text AS auth_source, auth_source::text AS auth_source,
allowed_providers, allowed_providers,
allowed_providers_mode,
allowed_api_formats, allowed_api_formats,
allowed_api_formats_mode,
allowed_models, allowed_models,
allowed_models_mode,
is_active, is_active,
is_deleted, is_deleted,
created_at, created_at,
@@ -216,8 +243,11 @@ SELECT
role::text AS role, role::text AS role,
auth_source::text AS auth_source, auth_source::text AS auth_source,
allowed_providers, allowed_providers,
allowed_providers_mode,
allowed_api_formats, allowed_api_formats,
allowed_api_formats_mode,
allowed_models, allowed_models,
allowed_models_mode,
is_active, is_active,
is_deleted, is_deleted,
created_at, created_at,
@@ -237,8 +267,11 @@ SELECT
role::text AS role, role::text AS role,
auth_source::text AS auth_source, auth_source::text AS auth_source,
allowed_providers, allowed_providers,
allowed_providers_mode,
allowed_api_formats, allowed_api_formats,
allowed_api_formats_mode,
allowed_models, allowed_models,
allowed_models_mode,
is_active, is_active,
is_deleted, is_deleted,
created_at, created_at,
@@ -259,8 +292,11 @@ SELECT
role::text AS role, role::text AS role,
auth_source::text AS auth_source, auth_source::text AS auth_source,
allowed_providers, allowed_providers,
allowed_providers_mode,
allowed_api_formats, allowed_api_formats,
allowed_api_formats_mode,
allowed_models, allowed_models,
allowed_models_mode,
is_active, is_active,
is_deleted, is_deleted,
created_at, 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 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)] #[derive(Debug, Clone)]
pub struct SqlxUserReadRepository { pub struct SqlxUserReadRepository {
pool: PgPool, pool: PgPool,
@@ -612,6 +682,275 @@ impl SqlxUserReadRepository {
.await .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( pub async fn list_export_users_page(
&self, &self,
query: &UserExportListQuery, query: &UserExportListQuery,
@@ -626,6 +965,16 @@ impl SqlxUserReadRepository {
if let Some(is_active) = query.is_active { if let Some(is_active) = query.is_active {
builder.push(" AND is_active = ").push_bind(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 if let Some(search) = query
.search .search
.as_deref() .as_deref()
@@ -817,10 +1166,12 @@ impl SqlxUserReadRepository {
r#" r#"
INSERT INTO users ( INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source, 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 is_active, is_deleted, created_at, updated_at, last_login_at
) )
VALUES ( VALUES (
$1, $2, TRUE, $3, NULL, 'user'::userrole, 'oauth'::authsource, $1, $2, TRUE, $3, NULL, 'user'::userrole, 'oauth'::authsource,
'inherit', 'inherit', 'inherit', 'inherit',
TRUE, FALSE, $4, $4, $4 TRUE, FALSE, $4, $4, $4
) )
"#, "#,
@@ -963,8 +1314,9 @@ SET email = $2,
WHERE id = $1 WHERE id = $1
RETURNING RETURNING
id, email, email_verified, username, password_hash, role::text AS role, id, email, email_verified, username, password_hash, role::text AS role,
auth_source::text AS auth_source, allowed_providers, allowed_api_formats, auth_source::text AS auth_source, allowed_providers, allowed_providers_mode,
allowed_models, is_active, is_deleted, created_at, last_login_at allowed_api_formats, allowed_api_formats_mode, allowed_models, allowed_models_mode,
is_active, is_deleted, created_at, last_login_at
"#, "#,
) )
.bind(&existing.id) .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) VALUES ($1, $2, TRUE, $3, NULL, 'user'::userrole, 'ldap'::authsource, $4, $5, TRUE, FALSE, $6, $6, $6)
RETURNING RETURNING
id, email, email_verified, username, password_hash, role::text AS role, id, email, email_verified, username, password_hash, role::text AS role,
auth_source::text AS auth_source, allowed_providers, allowed_api_formats, auth_source::text AS auth_source, allowed_providers, allowed_providers_mode,
allowed_models, is_active, is_deleted, created_at, last_login_at 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()) .bind(uuid::Uuid::new_v4().to_string())
@@ -1127,6 +1480,26 @@ WHERE id = $1
rate_limit: Option<i32>, rate_limit: Option<i32>,
is_active: Option<bool>, is_active: Option<bool>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> { ) -> 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( let result = sqlx::query(
r#" r#"
UPDATE users UPDATE users
@@ -1138,20 +1511,36 @@ SET role = CASE
WHEN $4::BOOLEAN THEN $5::json WHEN $4::BOOLEAN THEN $5::json
ELSE allowed_providers ELSE allowed_providers
END, END,
allowed_providers_mode = CASE
WHEN $4::BOOLEAN THEN $6
ELSE allowed_providers_mode
END,
allowed_api_formats = CASE allowed_api_formats = CASE
WHEN $6::BOOLEAN THEN $7::json WHEN $7::BOOLEAN THEN $8::json
ELSE allowed_api_formats ELSE allowed_api_formats
END, END,
allowed_api_formats_mode = CASE
WHEN $7::BOOLEAN THEN $9
ELSE allowed_api_formats_mode
END,
allowed_models = CASE allowed_models = CASE
WHEN $8::BOOLEAN THEN $9::json WHEN $10::BOOLEAN THEN $11::json
ELSE allowed_models ELSE allowed_models
END, END,
allowed_models_mode = CASE
WHEN $10::BOOLEAN THEN $12
ELSE allowed_models_mode
END,
rate_limit = CASE rate_limit = CASE
WHEN $10::BOOLEAN THEN $11 WHEN $13::BOOLEAN THEN $14
ELSE rate_limit ELSE rate_limit
END, END,
rate_limit_mode = CASE
WHEN $13::BOOLEAN THEN $15
ELSE rate_limit_mode
END,
is_active = CASE 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 ELSE is_active
END, END,
updated_at = NOW() updated_at = NOW()
@@ -1163,12 +1552,16 @@ WHERE id = $1
.bind(role) .bind(role)
.bind(allowed_providers_present) .bind(allowed_providers_present)
.bind(allowed_providers.map(serde_json::Value::from)) .bind(allowed_providers.map(serde_json::Value::from))
.bind(allowed_providers_mode)
.bind(allowed_api_formats_present) .bind(allowed_api_formats_present)
.bind(allowed_api_formats.map(serde_json::Value::from)) .bind(allowed_api_formats.map(serde_json::Value::from))
.bind(allowed_api_formats_mode)
.bind(allowed_models_present) .bind(allowed_models_present)
.bind(allowed_models.map(serde_json::Value::from)) .bind(allowed_models.map(serde_json::Value::from))
.bind(allowed_models_mode)
.bind(rate_limit_present) .bind(rate_limit_present)
.bind(rate_limit) .bind(rate_limit)
.bind(rate_limit_mode)
.bind(is_active.is_some()) .bind(is_active.is_some())
.bind(is_active) .bind(is_active)
.execute(&self.pool) .execute(&self.pool)
@@ -1180,6 +1573,55 @@ WHERE id = $1
self.find_user_auth_by_id(user_id).await 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( pub async fn update_user_model_capability_settings(
&self, &self,
user_id: &str, user_id: &str,
@@ -1212,18 +1654,30 @@ WHERE id = $1
username: String, username: String,
password_hash: String, password_hash: String,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> { ) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
self.create_local_auth_user_with_settings( let user_id = uuid::Uuid::new_v4().to_string();
email, sqlx::query(
email_verified, r#"
username, INSERT INTO users (
password_hash, id, email, email_verified, username, password_hash, role, auth_source,
"user".to_string(), allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
None, is_active, is_deleted, created_at, updated_at
None, )
None, VALUES (
None, $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 .await
.map_postgres_err()?;
self.find_user_auth_by_id(&user_id).await
} }
#[allow(clippy::too_many_arguments)] #[allow(clippy::too_many_arguments)]
@@ -1240,16 +1694,39 @@ WHERE id = $1
rate_limit: Option<i32>, rate_limit: Option<i32>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> { ) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let user_id = uuid::Uuid::new_v4().to_string(); 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( sqlx::query(
r#" r#"
INSERT INTO users ( INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source, 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 is_active, is_deleted, created_at, updated_at
) )
VALUES ( VALUES (
$1, $2, $3, $4, $5, $6::userrole, 'local'::authsource, $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() TRUE, FALSE, NOW(), NOW()
) )
"#, "#,
@@ -1261,9 +1738,13 @@ VALUES (
.bind(password_hash) .bind(password_hash)
.bind(role) .bind(role)
.bind(allowed_providers.map(serde_json::Value::from)) .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.map(serde_json::Value::from))
.bind(allowed_api_formats_mode)
.bind(allowed_models.map(serde_json::Value::from)) .bind(allowed_models.map(serde_json::Value::from))
.bind(allowed_models_mode)
.bind(rate_limit) .bind(rate_limit)
.bind(rate_limit_mode)
.execute(&self.pool) .execute(&self.pool)
.await .await
.map_postgres_err()?; .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( async fn find_postgres_ldap_user_for_update(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>, tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
ldap_dn: Option<&str>, ldap_dn: Option<&str>,
@@ -1550,8 +2041,9 @@ async fn find_postgres_ldap_user_for_update(
let select_columns = r#" let select_columns = r#"
SELECT SELECT
id, email, email_verified, username, password_hash, role::text AS role, id, email, email_verified, username, password_hash, role::text AS role,
auth_source::text AS auth_source, allowed_providers, allowed_api_formats, auth_source::text AS auth_source, allowed_providers, allowed_providers_mode,
allowed_models, is_active, is_deleted, created_at, last_login_at allowed_api_formats, allowed_api_formats_mode, allowed_models, allowed_models_mode,
is_active, is_deleted, created_at, last_login_at
FROM users FROM users
"#; "#;
if let Some(ldap_dn) = ldap_dn.filter(|value| !value.trim().is_empty()) { 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()?, .map_postgres_err()?,
row.try_get("is_active").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> { 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("created_at").map_postgres_err()?,
row.try_get("last_login_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( fn map_oauth_link_summary_row(
@@ -1709,6 +2263,88 @@ impl UserReadRepository for SqlxUserReadRepository {
self.find_export_user_by_id(user_id).await 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( async fn find_user_auth_by_id(
&self, &self,
user_id: &str, user_id: &str,
@@ -1919,6 +2555,24 @@ impl UserReadRepository for SqlxUserReadRepository {
.await .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( async fn update_user_model_capability_settings(
&self, &self,
user_id: &str, user_id: &str,

View File

@@ -3,9 +3,11 @@ use chrono::{DateTime, TimeZone, Utc};
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite}; use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use super::types::{ use super::types::{
LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, StoredUserExportRow, normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord,
StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership,
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord, StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
StoredUserSummary, UserExportListQuery, UserExportSummary, UserReadRepository, StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSummary,
UserReadRepository,
}; };
use crate::driver::sqlite::SqlitePool; use crate::driver::sqlite::SqlitePool;
use crate::error::SqlResultExt; use crate::error::SqlResultExt;
@@ -32,9 +34,13 @@ SELECT
role, role,
auth_source, auth_source,
allowed_providers, allowed_providers,
allowed_providers_mode,
allowed_api_formats, allowed_api_formats,
allowed_api_formats_mode,
allowed_models, allowed_models,
allowed_models_mode,
rate_limit, rate_limit,
rate_limit_mode,
model_capability_settings, model_capability_settings,
is_active is_active
FROM users FROM users
@@ -50,8 +56,11 @@ SELECT
role, role,
auth_source, auth_source,
allowed_providers, allowed_providers,
allowed_providers_mode,
allowed_api_formats, allowed_api_formats,
allowed_api_formats_mode,
allowed_models, allowed_models,
allowed_models_mode,
is_active, is_active,
is_deleted, is_deleted,
created_at, created_at,
@@ -69,8 +78,11 @@ SELECT
users.role AS role, users.role AS role,
users.auth_source AS auth_source, users.auth_source AS auth_source,
users.allowed_providers AS allowed_providers, 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 AS allowed_api_formats,
users.allowed_api_formats_mode AS allowed_api_formats_mode,
users.allowed_models AS allowed_models, users.allowed_models AS allowed_models,
users.allowed_models_mode AS allowed_models_mode,
users.is_active AS is_active, users.is_active AS is_active,
users.is_deleted AS is_deleted, users.is_deleted AS is_deleted,
users.created_at AS created_at, users.created_at AS created_at,
@@ -130,6 +142,40 @@ SELECT
FROM user_sessions 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)] #[derive(Debug, Clone)]
pub struct SqliteUserReadRepository { pub struct SqliteUserReadRepository {
pool: SqlitePool, pool: SqlitePool,
@@ -163,6 +209,22 @@ impl SqliteUserReadRepository {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_auth_row).collect() 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] #[async_trait]
@@ -224,6 +286,16 @@ impl UserReadRepository for SqliteUserReadRepository {
if let Some(is_active) = query.is_active { if let Some(is_active) = query.is_active {
builder.push(" AND is_active = ").push_bind(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 if let Some(search) = query
.search .search
.as_deref() .as_deref()
@@ -296,6 +368,285 @@ WHERE is_deleted = 0
self.fetch_export_rows(builder).await 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( async fn find_user_auth_by_id(
&self, &self,
user_id: &str, user_id: &str,
@@ -453,9 +804,10 @@ WHERE provider_type = ?
r#" r#"
INSERT INTO users ( INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source, 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 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) .bind(&user_id)
@@ -675,14 +1027,38 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
rate_limit: Option<i32>, rate_limit: Option<i32>,
is_active: Option<bool>, is_active: Option<bool>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> { ) -> 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( let result = sqlx::query(
r#" r#"
UPDATE users UPDATE users
SET role = CASE WHEN ? THEN COALESCE(?, role) ELSE role END, SET role = CASE WHEN ? THEN COALESCE(?, role) ELSE role END,
allowed_providers = CASE WHEN ? THEN ? ELSE allowed_providers 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 = 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 = 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 = 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, is_active = CASE WHEN ? THEN ? ELSE is_active END,
updated_at = ? updated_at = ?
WHERE id = ? WHERE id = ?
@@ -695,18 +1071,26 @@ WHERE id = ?
allowed_providers, allowed_providers,
"users.allowed_providers", "users.allowed_providers",
)?) )?)
.bind(allowed_providers_present)
.bind(allowed_providers_mode)
.bind(allowed_api_formats_present) .bind(allowed_api_formats_present)
.bind(optional_string_list_json( .bind(optional_string_list_json(
allowed_api_formats, allowed_api_formats,
"users.allowed_api_formats", "users.allowed_api_formats",
)?) )?)
.bind(allowed_api_formats_present)
.bind(allowed_api_formats_mode)
.bind(allowed_models_present) .bind(allowed_models_present)
.bind(optional_string_list_json( .bind(optional_string_list_json(
allowed_models, allowed_models,
"users.allowed_models", "users.allowed_models",
)?) )?)
.bind(allowed_models_present)
.bind(allowed_models_mode)
.bind(rate_limit_present) .bind(rate_limit_present)
.bind(rate_limit) .bind(rate_limit)
.bind(rate_limit_present)
.bind(rate_limit_mode)
.bind(is_active.is_some()) .bind(is_active.is_some())
.bind(is_active) .bind(is_active)
.bind(chrono::Utc::now().timestamp()) .bind(chrono::Utc::now().timestamp())
@@ -720,6 +1104,44 @@ WHERE id = ?
self.find_user_auth_by_id(user_id).await 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( async fn update_user_model_capability_settings(
&self, &self,
user_id: &str, user_id: &str,
@@ -751,18 +1173,29 @@ WHERE id = ?
username: String, username: String,
password_hash: String, password_hash: String,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> { ) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
self.create_local_auth_user_with_settings( let user_id = uuid::Uuid::new_v4().to_string();
email, let now = chrono::Utc::now().timestamp();
email_verified, sqlx::query(
username, r#"
password_hash, INSERT INTO users (
"user".to_string(), id, email, email_verified, username, password_hash, role, auth_source,
None, allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
None, is_active, is_deleted, created_at, updated_at
None, )
None, 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 .await
.map_sql_err()?;
self.find_user_auth_by_id(&user_id).await
} }
async fn create_local_auth_user_with_settings( async fn create_local_auth_user_with_settings(
@@ -779,14 +1212,37 @@ WHERE id = ?
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> { ) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let user_id = uuid::Uuid::new_v4().to_string(); let user_id = uuid::Uuid::new_v4().to_string();
let now = chrono::Utc::now().timestamp(); 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( sqlx::query(
r#" r#"
INSERT INTO users ( INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source, 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 is_active, is_deleted, created_at, updated_at
) )
VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, 1, 0, ?, ?) VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, ?, ?, ?, ?, 1, 0, ?, ?)
"#, "#,
) )
.bind(&user_id) .bind(&user_id)
@@ -799,15 +1255,19 @@ VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, 1, 0, ?, ?)
allowed_providers, allowed_providers,
"users.allowed_providers", "users.allowed_providers",
)?) )?)
.bind(allowed_providers_mode)
.bind(optional_string_list_json( .bind(optional_string_list_json(
allowed_api_formats, allowed_api_formats,
"users.allowed_api_formats", "users.allowed_api_formats",
)?) )?)
.bind(allowed_api_formats_mode)
.bind(optional_string_list_json( .bind(optional_string_list_json(
allowed_models, allowed_models,
"users.allowed_models", "users.allowed_models",
)?) )?)
.bind(allowed_models_mode)
.bind(rate_limit) .bind(rate_limit)
.bind(rate_limit_mode)
.bind(now) .bind(now)
.bind(now) .bind(now)
.execute(&self.pool) .execute(&self.pool)
@@ -1166,6 +1626,24 @@ fn optional_string_list_json(
.transpose() .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( fn optional_json_string(
value: Option<serde_json::Value>, value: Option<serde_json::Value>,
field_name: &str, 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()?, 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> { 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("created_at").map_sql_err()?),
optional_datetime_from_unix_secs(row.try_get("last_login_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( fn map_oauth_link_summary_row(
@@ -1560,6 +2107,7 @@ INSERT INTO users (
role: Some("user".to_string()), role: Some("user".to_string()),
is_active: Some(true), is_active: Some(true),
search: None, search: None,
group_id: None,
}) })
.await .await
.expect("export page should load"); .expect("export page should load");

View File

@@ -57,8 +57,11 @@ pub struct StoredUserAuthRecord {
pub role: String, pub role: String,
pub auth_source: String, pub auth_source: String,
pub allowed_providers: Option<Vec<String>>, pub allowed_providers: Option<Vec<String>>,
pub allowed_providers_mode: String,
pub allowed_api_formats: Option<Vec<String>>, pub allowed_api_formats: Option<Vec<String>>,
pub allowed_api_formats_mode: String,
pub allowed_models: Option<Vec<String>>, pub allowed_models: Option<Vec<String>>,
pub allowed_models_mode: String,
pub is_active: bool, pub is_active: bool,
pub is_deleted: bool, pub is_deleted: bool,
pub created_at: Option<DateTime<Utc>>, pub created_at: Option<DateTime<Utc>>,
@@ -113,16 +116,44 @@ impl StoredUserAuthRecord {
role, role,
auth_source, auth_source,
allowed_providers: parse_string_list(allowed_providers, "users.allowed_providers")?, 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: parse_string_list(
allowed_api_formats, allowed_api_formats,
"users.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: parse_string_list(allowed_models, "users.allowed_models")?,
allowed_models_mode: "unrestricted".to_string(),
is_active, is_active,
is_deleted, is_deleted,
created_at, created_at,
last_login_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> { pub fn to_summary(&self) -> Result<StoredUserSummary, crate::DataLayerError> {
@@ -197,9 +228,13 @@ pub struct StoredUserExportRow {
pub role: String, pub role: String,
pub auth_source: String, pub auth_source: String,
pub allowed_providers: Option<Vec<String>>, pub allowed_providers: Option<Vec<String>>,
pub allowed_providers_mode: String,
pub allowed_api_formats: Option<Vec<String>>, pub allowed_api_formats: Option<Vec<String>>,
pub allowed_api_formats_mode: String,
pub allowed_models: Option<Vec<String>>, pub allowed_models: Option<Vec<String>>,
pub allowed_models_mode: String,
pub rate_limit: Option<i32>, pub rate_limit: Option<i32>,
pub rate_limit_mode: String,
pub model_capability_settings: Option<Value>, pub model_capability_settings: Option<Value>,
pub is_active: bool, pub is_active: bool,
} }
@@ -251,15 +286,52 @@ impl StoredUserExportRow {
role, role,
auth_source, auth_source,
allowed_providers: parse_string_list(allowed_providers, "users.allowed_providers")?, 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: parse_string_list(
allowed_api_formats, allowed_api_formats,
"users.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: parse_string_list(allowed_models, "users.allowed_models")?,
allowed_models_mode: "unrestricted".to_string(),
rate_limit, rate_limit,
rate_limit_mode: "system".to_string(),
model_capability_settings: normalize_optional_json(model_capability_settings), model_capability_settings: normalize_optional_json(model_capability_settings),
is_active, 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, 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 { impl StoredUserPreferenceRecord {
pub fn default_for_user(user_id: impl Into<String>) -> Self { pub fn default_for_user(user_id: impl Into<String>) -> Self {
Self { Self {
@@ -429,6 +634,7 @@ pub struct UserExportListQuery {
pub role: Option<String>, pub role: Option<String>,
pub is_active: Option<bool>, pub is_active: Option<bool>,
pub search: Option<String>, pub search: Option<String>,
pub group_id: Option<String>,
} }
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)] #[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
@@ -463,6 +669,64 @@ pub trait UserReadRepository: Send + Sync {
user_id: &str, user_id: &str,
) -> Result<Option<StoredUserExportRow>, crate::DataLayerError>; ) -> 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( async fn list_non_admin_export_users(
&self, &self,
) -> Result<Vec<StoredUserExportRow>, crate::DataLayerError>; ) -> Result<Vec<StoredUserExportRow>, crate::DataLayerError>;
@@ -602,6 +866,15 @@ pub trait UserReadRepository: Send + Sync {
is_active: Option<bool>, is_active: Option<bool>,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>; ) -> 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( async fn update_user_model_capability_settings(
&self, &self,
user_id: &str, 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( fn parse_string_list(
value: Option<Value>, value: Option<Value>,
field_name: &str, field_name: &str,

View File

@@ -3,6 +3,29 @@ import { cachedRequest } from '@/utils/cache'
import type { UserSession as SessionRecord } from '@/types/session' import type { UserSession as SessionRecord } from '@/types/session'
export type UserRole = 'admin' | 'user' export type UserRole = 'admin' | 'user'
export type ListPolicyMode = 'inherit' | 'unrestricted' | 'specific' | 'deny_all'
export type RateLimitPolicyMode = 'inherit' | 'system' | 'custom'
export interface UserGroupSummary {
id: string
name: string
priority: number
}
export interface EffectivePolicyField<T> {
mode: string
value: T | null
source: 'user' | 'group' | 'fallback' | string
group_id?: string | null
group_name?: string | null
}
export interface UserEffectivePolicy {
allowed_providers?: EffectivePolicyField<string[]>
allowed_api_formats?: EffectivePolicyField<string[]>
allowed_models?: EffectivePolicyField<string[]>
rate_limit?: EffectivePolicyField<number>
}
export interface User { export interface User {
id: string // UUID id: string // UUID
@@ -12,9 +35,15 @@ export interface User {
is_active: boolean is_active: boolean
unlimited: boolean unlimited: boolean
allowed_providers: string[] | null // 允许使用的提供商 ID 列表 allowed_providers: string[] | null // 允许使用的提供商 ID 列表
allowed_providers_mode?: ListPolicyMode
allowed_api_formats: string[] | null // 允许使用的 API 格式列表 allowed_api_formats: string[] | null // 允许使用的 API 格式列表
allowed_api_formats_mode?: ListPolicyMode
allowed_models: string[] | null // 允许使用的模型名称列表 allowed_models: string[] | null // 允许使用的模型名称列表
allowed_models_mode?: ListPolicyMode
rate_limit?: number | null // null = 跟随系统默认0 = 不限制 rate_limit?: number | null // null = 跟随系统默认0 = 不限制
rate_limit_mode?: RateLimitPolicyMode
groups?: UserGroupSummary[]
effective_policy?: UserEffectivePolicy
created_at: string created_at: string
updated_at?: string updated_at?: string
last_login_at?: string | null last_login_at?: string | null
@@ -30,9 +59,14 @@ export interface CreateUserRequest {
initial_gift_usd?: number | null initial_gift_usd?: number | null
unlimited?: boolean unlimited?: boolean
allowed_providers?: string[] | null allowed_providers?: string[] | null
allowed_providers_mode?: ListPolicyMode
allowed_api_formats?: string[] | null allowed_api_formats?: string[] | null
allowed_api_formats_mode?: ListPolicyMode
allowed_models?: string[] | null allowed_models?: string[] | null
allowed_models_mode?: ListPolicyMode
rate_limit?: number | null rate_limit?: number | null
rate_limit_mode?: RateLimitPolicyMode
group_ids?: string[]
} }
export interface UpdateUserRequest { export interface UpdateUserRequest {
@@ -42,19 +76,26 @@ export interface UpdateUserRequest {
unlimited?: boolean unlimited?: boolean
password?: string password?: string
allowed_providers?: string[] | null allowed_providers?: string[] | null
allowed_providers_mode?: ListPolicyMode
allowed_api_formats?: string[] | null allowed_api_formats?: string[] | null
allowed_api_formats_mode?: ListPolicyMode
allowed_models?: string[] | null allowed_models?: string[] | null
allowed_models_mode?: ListPolicyMode
rate_limit?: number | null rate_limit?: number | null
rate_limit_mode?: RateLimitPolicyMode
group_ids?: string[]
} }
export interface UserBatchSelectionFilters { export interface UserBatchSelectionFilters {
search?: string search?: string
role?: UserRole role?: UserRole
is_active?: boolean is_active?: boolean
group_id?: string
} }
export interface UserBatchSelection { export interface UserBatchSelection {
user_ids?: string[] user_ids?: string[]
group_ids?: string[]
filters?: UserBatchSelectionFilters | null filters?: UserBatchSelectionFilters | null
} }
@@ -64,11 +105,19 @@ export interface UserBatchSelectionItem {
email?: string | null email?: string | null
role: UserRole role: UserRole
is_active: boolean is_active: boolean
matched_by?: string[]
}
export interface UserBatchSelectionWarning {
type: string
group_id?: string | null
message: string
} }
export interface ResolveUserBatchSelectionResponse { export interface ResolveUserBatchSelectionResponse {
total: number total: number
items: UserBatchSelectionItem[] items: UserBatchSelectionItem[]
warnings?: UserBatchSelectionWarning[]
} }
export interface UserBatchAccessControlPayload { export interface UserBatchAccessControlPayload {
@@ -120,10 +169,60 @@ export interface UserBatchActionResponse {
success: number success: number
failed: number failed: number
failures: UserBatchActionFailure[] failures: UserBatchActionFailure[]
warnings?: UserBatchSelectionWarning[]
action?: string action?: string
modified_fields?: string[] modified_fields?: string[]
} }
export interface UserGroup {
id: string
name: string
normalized_name?: string
description?: string | null
priority: number
allowed_providers?: string[] | null
allowed_providers_mode: ListPolicyMode
allowed_api_formats?: string[] | null
allowed_api_formats_mode: ListPolicyMode
allowed_models?: string[] | null
allowed_models_mode: ListPolicyMode
rate_limit?: number | null
rate_limit_mode: RateLimitPolicyMode
is_default?: boolean
created_at?: string | null
updated_at?: string | null
}
export interface UpsertUserGroupRequest {
name: string
description?: string | null
priority?: number
allowed_providers?: string[] | null
allowed_providers_mode?: ListPolicyMode
allowed_api_formats?: string[] | null
allowed_api_formats_mode?: ListPolicyMode
allowed_models?: string[] | null
allowed_models_mode?: ListPolicyMode
rate_limit?: number | null
rate_limit_mode?: RateLimitPolicyMode
}
export interface UserGroupMember {
group_id: string
user_id: string
username: string
email?: string | null
role: UserRole
is_active: boolean
is_deleted: boolean
created_at?: string | null
}
export interface ListUserGroupsResponse {
items: UserGroup[]
default_group_id?: string | null
}
export interface ApiKey { export interface ApiKey {
id: string // UUID id: string // UUID
key?: string // 完整的 key只在创建时返回 key?: string // 完整的 key只在创建时返回
@@ -151,6 +250,9 @@ export type UserSession = SessionRecord
export interface GetAllUsersOptions { export interface GetAllUsersOptions {
search?: string search?: string
role?: UserRole
is_active?: boolean
group_id?: string
skip?: number skip?: number
limit?: number limit?: number
cacheTtlMs?: number cacheTtlMs?: number
@@ -163,6 +265,9 @@ export const usersApi = {
const search = options.search?.trim() const search = options.search?.trim()
if (search) params.search = search if (search) params.search = search
if (options.role) params.role = options.role
if (options.is_active !== undefined) params.is_active = options.is_active ? 'true' : 'false'
if (options.group_id) params.group_id = options.group_id
if (options.skip !== undefined) params.skip = options.skip if (options.skip !== undefined) params.skip = options.skip
if (options.limit !== undefined) params.limit = options.limit if (options.limit !== undefined) params.limit = options.limit
@@ -171,6 +276,9 @@ export const usersApi = {
: [ : [
'admin:users:list', 'admin:users:list',
search ?? '', search ?? '',
options.role ?? '',
options.is_active ?? '',
options.group_id ?? '',
options.skip ?? '', options.skip ?? '',
options.limit ?? '', options.limit ?? '',
].join(':') ].join(':')
@@ -220,6 +328,46 @@ export const usersApi = {
return response.data return response.data
}, },
async listUserGroups(): Promise<ListUserGroupsResponse> {
const response = await apiClient.get<ListUserGroupsResponse>('/api/admin/user-groups')
return response.data
},
async createUserGroup(payload: UpsertUserGroupRequest): Promise<UserGroup> {
const response = await apiClient.post<UserGroup>('/api/admin/user-groups', payload)
return response.data
},
async updateUserGroup(groupId: string, payload: UpsertUserGroupRequest): Promise<UserGroup> {
const response = await apiClient.put<UserGroup>(`/api/admin/user-groups/${groupId}`, payload)
return response.data
},
async deleteUserGroup(groupId: string): Promise<void> {
await apiClient.delete(`/api/admin/user-groups/${groupId}`)
},
async listUserGroupMembers(groupId: string): Promise<UserGroupMember[]> {
const response = await apiClient.get<{ items: UserGroupMember[] }>(`/api/admin/user-groups/${groupId}/members`)
return response.data.items
},
async replaceUserGroupMembers(groupId: string, userIds: string[]): Promise<UserGroupMember[]> {
const response = await apiClient.put<{ items: UserGroupMember[] }>(
`/api/admin/user-groups/${groupId}/members`,
{ user_ids: userIds },
)
return response.data.items
},
async setDefaultUserGroup(groupId: string | null): Promise<{ default_group_id?: string | null }> {
const response = await apiClient.put<{ default_group_id?: string | null }>(
'/api/admin/user-groups/default',
{ group_id: groupId },
)
return response.data
},
async deleteUser(userId: string): Promise<void> { async deleteUser(userId: string): Promise<void> {
await apiClient.delete(`/api/admin/users/${userId}`) await apiClient.delete(`/api/admin/users/${userId}`)
}, },

View File

@@ -54,7 +54,7 @@
</div> </div>
</div> </div>
<div class="max-h-48 overflow-y-auto"> <div class="max-h-64 overflow-y-auto">
<div <div
v-if="hasOptions" v-if="hasOptions"
class="sticky top-0 z-10 flex cursor-pointer items-center gap-2 border-b bg-popover/95 px-3 py-2 backdrop-blur hover:bg-muted/50 supports-[backdrop-filter]:bg-popover/85" class="sticky top-0 z-10 flex cursor-pointer items-center gap-2 border-b bg-popover/95 px-3 py-2 backdrop-blur hover:bg-muted/50 supports-[backdrop-filter]:bg-popover/85"

View File

@@ -52,6 +52,20 @@
</div> </div>
<div class="space-y-2.5"> <div class="space-y-2.5">
<div class="grid gap-2 rounded-xl border border-border/70 bg-muted/20 p-3 sm:grid-cols-[9rem_minmax(0,1fr)] sm:items-start">
<div>
<Label class="text-sm font-medium">按分组选择</Label>
<p class="mt-1 text-[11px] text-muted-foreground">可与直接用户或筛选条件混合</p>
</div>
<MultiSelect
v-model="selectedGroupIds"
:options="groupOptions"
:search-threshold="0"
placeholder="选择一个或多个分组"
empty-text="暂无用户分组"
/>
</div>
<div class="flex items-center justify-between gap-3"> <div class="flex items-center justify-between gap-3">
<Label class="text-sm font-medium">选择批量动作</Label> <Label class="text-sm font-medium">选择批量动作</Label>
<span class="text-[11px] text-muted-foreground">只会提交当前动作对应的字段</span> <span class="text-[11px] text-muted-foreground">只会提交当前动作对应的字段</span>
@@ -339,6 +353,7 @@ import type {
UserBatchSelectionFilters, UserBatchSelectionFilters,
UserBatchSelectionItem, UserBatchSelectionItem,
UserRole, UserRole,
UserGroup,
} from '@/api/users' } from '@/api/users'
type AccessFieldMode = 'skip' | 'unrestricted' | 'specific' type AccessFieldMode = 'skip' | 'unrestricted' | 'specific'
@@ -358,6 +373,7 @@ const props = defineProps<{
selectAllFiltered: boolean selectAllFiltered: boolean
selectedCount: number selectedCount: number
filters: UserBatchSelectionFilters filters: UserBatchSelectionFilters
groups: UserGroup[]
}>() }>()
const emit = defineEmits<{ const emit = defineEmits<{
@@ -408,6 +424,7 @@ const apiFormatMode = ref<AccessFieldMode>('skip')
const modelMode = ref<AccessFieldMode>('skip') const modelMode = ref<AccessFieldMode>('skip')
const rateLimitMode = ref<RateLimitMode>('skip') const rateLimitMode = ref<RateLimitMode>('skip')
const quotaMode = ref<QuotaMode>('skip') const quotaMode = ref<QuotaMode>('skip')
const selectedGroupIds = ref<string[]>([])
const allowedProviders = ref<string[]>([]) const allowedProviders = ref<string[]>([])
const allowedApiFormats = ref<string[]>([]) const allowedApiFormats = ref<string[]>([])
const allowedModels = ref<string[]>([]) const allowedModels = ref<string[]>([])
@@ -418,8 +435,13 @@ const resolvedTotal = ref<number | null>(null)
const executing = ref(false) const executing = ref(false)
const lastResult = ref<UserBatchActionResponse | null>(null) const lastResult = ref<UserBatchActionResponse | null>(null)
const groupOptions = computed(() => props.groups.map((group) => ({
label: `${group.name}${group.is_default ? '(默认)' : ''}`,
value: group.id,
})))
const hasAnyTarget = computed(() => props.selectedCount > 0 || selectedGroupIds.value.length > 0)
const impactCount = computed(() => resolvedTotal.value ?? props.selectedCount) const impactCount = computed(() => resolvedTotal.value ?? props.selectedCount)
const canExecute = computed(() => props.selectedCount > 0 && !previewLoading.value && !executing.value) const canExecute = computed(() => hasAnyTarget.value && !previewLoading.value && !executing.value)
const selectedActionLabel = computed(() => ( const selectedActionLabel = computed(() => (
actionOptions.find((action) => action.value === selectedAction.value)?.label ?? '批量操作' actionOptions.find((action) => action.value === selectedAction.value)?.label ?? '批量操作'
)) ))
@@ -456,6 +478,7 @@ function resetLocalState(): void {
modelMode.value = 'skip' modelMode.value = 'skip'
rateLimitMode.value = 'skip' rateLimitMode.value = 'skip'
quotaMode.value = 'skip' quotaMode.value = 'skip'
selectedGroupIds.value = []
allowedProviders.value = [] allowedProviders.value = []
allowedApiFormats.value = [] allowedApiFormats.value = []
allowedModels.value = [] allowedModels.value = []
@@ -482,14 +505,15 @@ function actionIconClass(action: UserBatchAction): string {
} }
function buildSelection(): UserBatchSelection { function buildSelection(): UserBatchSelection {
const group_ids = selectedGroupIds.value.length > 0 ? [...selectedGroupIds.value] : undefined
if (props.selectAllFiltered) { if (props.selectAllFiltered) {
return { filters: props.filters } return { filters: props.filters, group_ids }
} }
return { user_ids: [...props.selectedIds] } return { user_ids: [...props.selectedIds], group_ids }
} }
async function resolvePreview(): Promise<void> { async function resolvePreview(): Promise<void> {
if (props.selectedCount === 0) { if (!hasAnyTarget.value) {
resolvedTotal.value = 0 resolvedTotal.value = 0
previewItems.value = [] previewItems.value = []
return return
@@ -508,6 +532,10 @@ async function resolvePreview(): Promise<void> {
} }
} }
watch(selectedGroupIds, () => {
if (props.open) void resolvePreview()
})
function buildAccessControlPayload(): UserBatchAccessControlPayload | null { function buildAccessControlPayload(): UserBatchAccessControlPayload | null {
const payload: UserBatchAccessControlPayload = {} const payload: UserBatchAccessControlPayload = {}
if (providerMode.value === 'unrestricted') payload.allowed_providers = null if (providerMode.value === 'unrestricted') payload.allowed_providers = null

View File

@@ -180,6 +180,18 @@
</Select> </Select>
</div> </div>
</div> </div>
<div class="space-y-2">
<Label class="text-sm font-medium">所属分组</Label>
<MultiSelect
v-model="form.group_ids"
:options="groupOptions"
:search-threshold="0"
placeholder="可选择多个分组"
empty-text="暂无分组"
no-results-text="未找到匹配的分组"
/>
</div>
</div> </div>
<!-- 右侧:访问限制 --> <!-- 右侧:访问限制 -->
@@ -191,22 +203,27 @@
<!-- 提供商 --> <!-- 提供商 -->
<div class="space-y-2"> <div class="space-y-2">
<Label class="text-sm font-medium">允许的提供商</Label> <Label class="text-sm font-medium">允许的提供商</Label>
<div class="flex items-center gap-3"> <div class="grid gap-2 sm:grid-cols-[7rem_minmax(0,1fr)]">
<div class="flex-1 min-w-0"> <Select v-model="form.allowed_providers_mode">
<MultiSelect <SelectTrigger class="h-10">
v-model="form.allowed_providers" <SelectValue />
:options="providerOptions" </SelectTrigger>
:search-threshold="0" <SelectContent>
:disabled="form.provider_unrestricted" <SelectItem value="inherit">继承</SelectItem>
:placeholder="form.provider_unrestricted ? '不限制' : '未选择(全部禁用)'" <SelectItem value="unrestricted">不限制</SelectItem>
empty-text="暂无可用提供商" <SelectItem value="specific">指定列表</SelectItem>
no-results-text="未找到匹配的提供商" <SelectItem value="deny_all">全部禁用</SelectItem>
search-placeholder="搜索提供商名称..." </SelectContent>
/> </Select>
</div> <MultiSelect
<Switch v-model="form.allowed_providers"
v-model="form.provider_unrestricted" :options="providerOptions"
class="shrink-0" :search-threshold="0"
:disabled="form.allowed_providers_mode !== 'specific'"
placeholder="未选择时表示全部禁用"
empty-text="暂无可用提供商"
no-results-text="未找到匹配的提供商"
search-placeholder="搜索提供商名称..."
/> />
</div> </div>
</div> </div>
@@ -214,22 +231,27 @@
<!-- 端点 --> <!-- 端点 -->
<div class="space-y-2"> <div class="space-y-2">
<Label class="text-sm font-medium">允许的端点</Label> <Label class="text-sm font-medium">允许的端点</Label>
<div class="flex items-center gap-3"> <div class="grid gap-2 sm:grid-cols-[7rem_minmax(0,1fr)]">
<div class="flex-1 min-w-0"> <Select v-model="form.allowed_api_formats_mode">
<MultiSelect <SelectTrigger class="h-10">
v-model="form.allowed_api_formats" <SelectValue />
:options="apiFormatOptions" </SelectTrigger>
:search-threshold="0" <SelectContent>
:disabled="form.api_format_unrestricted" <SelectItem value="inherit">继承</SelectItem>
:placeholder="form.api_format_unrestricted ? '不限制' : '未选择(全部禁用)'" <SelectItem value="unrestricted">不限制</SelectItem>
empty-text="暂无可用端点" <SelectItem value="specific">指定列表</SelectItem>
no-results-text="未找到匹配的端点" <SelectItem value="deny_all">全部禁用</SelectItem>
search-placeholder="搜索端点..." </SelectContent>
/> </Select>
</div> <MultiSelect
<Switch v-model="form.allowed_api_formats"
v-model="form.api_format_unrestricted" :options="apiFormatOptions"
class="shrink-0" :search-threshold="0"
:disabled="form.allowed_api_formats_mode !== 'specific'"
placeholder="未选择时表示全部禁用"
empty-text="暂无可用端点"
no-results-text="未找到匹配的端点"
search-placeholder="搜索端点..."
/> />
</div> </div>
</div> </div>
@@ -237,22 +259,27 @@
<!-- 模型 --> <!-- 模型 -->
<div class="space-y-2"> <div class="space-y-2">
<Label class="text-sm font-medium">允许的模型</Label> <Label class="text-sm font-medium">允许的模型</Label>
<div class="flex items-center gap-3"> <div class="grid gap-2 sm:grid-cols-[7rem_minmax(0,1fr)]">
<div class="flex-1 min-w-0"> <Select v-model="form.allowed_models_mode">
<MultiSelect <SelectTrigger class="h-10">
v-model="form.allowed_models" <SelectValue />
:options="modelOptions" </SelectTrigger>
:search-threshold="0" <SelectContent>
:disabled="form.model_unrestricted" <SelectItem value="inherit">继承</SelectItem>
:placeholder="form.model_unrestricted ? '不限制' : '未选择(全部禁用)'" <SelectItem value="unrestricted">不限制</SelectItem>
empty-text="暂无可用模型" <SelectItem value="specific">指定列表</SelectItem>
no-results-text="未找到匹配的模型" <SelectItem value="deny_all">全部禁用</SelectItem>
search-placeholder="输入模型名搜索..." </SelectContent>
/> </Select>
</div> <MultiSelect
<Switch v-model="form.allowed_models"
v-model="form.model_unrestricted" :options="modelOptions"
class="shrink-0" :search-threshold="0"
:disabled="form.allowed_models_mode !== 'specific'"
placeholder="未选择时表示全部禁用"
empty-text="暂无可用模型"
no-results-text="未找到匹配的模型"
search-placeholder="输入模型名搜索..."
/> />
</div> </div>
</div> </div>
@@ -263,9 +290,18 @@
class="text-sm font-medium" class="text-sm font-medium"
>速率限制 (请求/分钟)</Label> >速率限制 (请求/分钟)</Label>
<div class="flex items-center gap-3"> <div class="flex items-center gap-3">
<Select v-model="form.rate_limit_mode">
<SelectTrigger class="h-10 w-28">
<SelectValue />
</SelectTrigger>
<SelectContent>
<SelectItem value="inherit">继承</SelectItem>
<SelectItem value="system">系统默认</SelectItem>
<SelectItem value="custom">指定数值</SelectItem>
</SelectContent>
</Select>
<div class="flex-1 min-w-0"> <div class="flex-1 min-w-0">
<Input <Input
v-if="!form.rate_limit_inherited"
id="form-rate-limit" id="form-rate-limit"
:model-value="form.rate_limit ?? ''" :model-value="form.rate_limit ?? ''"
type="number" type="number"
@@ -273,17 +309,10 @@
max="10000" max="10000"
placeholder="0 = 不限速" placeholder="0 = 不限速"
class="h-10" class="h-10"
:disabled="form.rate_limit_mode !== 'custom'"
@update:model-value="(v) => form.rate_limit = parseNumberInput(v, { min: 0, max: 10000 })" @update:model-value="(v) => form.rate_limit = parseNumberInput(v, { min: 0, max: 10000 })"
/> />
<span
v-else
class="flex h-10 w-full items-center rounded-lg border bg-background px-3 text-sm text-muted-foreground opacity-60"
>跟随系统默认</span>
</div> </div>
<Switch
v-model="form.rate_limit_inherited"
class="shrink-0"
/>
</div> </div>
</div> </div>
@@ -366,6 +395,7 @@ import {
validatePasswordByPolicy, validatePasswordByPolicy,
type PasswordPolicyLevel, type PasswordPolicyLevel,
} from '@/utils/passwordPolicy' } from '@/utils/passwordPolicy'
import type { ListPolicyMode, RateLimitPolicyMode, UserGroup } from '@/api/users'
export interface UserFormData { export interface UserFormData {
id?: string id?: string
@@ -376,14 +406,20 @@ export interface UserFormData {
role: 'admin' | 'user' role: 'admin' | 'user'
is_active?: boolean is_active?: boolean
allowed_providers?: string[] | null allowed_providers?: string[] | null
allowed_providers_mode?: ListPolicyMode
allowed_api_formats?: string[] | null allowed_api_formats?: string[] | null
allowed_api_formats_mode?: ListPolicyMode
allowed_models?: string[] | null allowed_models?: string[] | null
allowed_models_mode?: ListPolicyMode
rate_limit?: number | null rate_limit?: number | null
rate_limit_mode?: RateLimitPolicyMode
group_ids?: string[]
} }
const props = defineProps<{ const props = defineProps<{
open: boolean open: boolean
user: UserFormData | null user: UserFormData | null
groups?: UserGroup[]
}>() }>()
const emit = defineEmits<{ const emit = defineEmits<{
@@ -413,16 +449,22 @@ const form = ref({
role: 'user' as 'admin' | 'user', role: 'user' as 'admin' | 'user',
unlimited: false, unlimited: false,
is_active: true, is_active: true,
provider_unrestricted: true, allowed_providers_mode: 'unrestricted' as ListPolicyMode,
api_format_unrestricted: true, allowed_api_formats_mode: 'unrestricted' as ListPolicyMode,
model_unrestricted: true, allowed_models_mode: 'unrestricted' as ListPolicyMode,
rate_limit_inherited: true, rate_limit_mode: 'system' as RateLimitPolicyMode,
allowed_providers: [] as string[], allowed_providers: [] as string[],
allowed_api_formats: [] as string[], allowed_api_formats: [] as string[],
allowed_models: [] as string[], allowed_models: [] as string[],
rate_limit: undefined as number | undefined, rate_limit: undefined as number | undefined,
group_ids: [] as string[],
}) })
const groupOptions = computed(() => (props.groups || []).map((group) => ({
label: group.name,
value: group.id,
})))
function createFieldNonce(): string { function createFieldNonce(): string {
return Math.random().toString(36).slice(2, 10) return Math.random().toString(36).slice(2, 10)
} }
@@ -438,14 +480,15 @@ function resetForm() {
role: 'user', role: 'user',
unlimited: false, unlimited: false,
is_active: true, is_active: true,
provider_unrestricted: true, allowed_providers_mode: 'unrestricted',
api_format_unrestricted: true, allowed_api_formats_mode: 'unrestricted',
model_unrestricted: true, allowed_models_mode: 'unrestricted',
rate_limit_inherited: true, rate_limit_mode: 'system',
allowed_providers: [], allowed_providers: [],
allowed_api_formats: [], allowed_api_formats: [],
allowed_models: [], allowed_models: [],
rate_limit: undefined, rate_limit: undefined,
group_ids: [],
} }
} }
@@ -462,14 +505,15 @@ function loadUserData() {
role: props.user.role, role: props.user.role,
unlimited: props.user.unlimited ?? false, unlimited: props.user.unlimited ?? false,
is_active: props.user.is_active ?? true, is_active: props.user.is_active ?? true,
provider_unrestricted: props.user.allowed_providers == null, allowed_providers_mode: props.user.allowed_providers_mode ?? (props.user.allowed_providers == null ? 'unrestricted' : 'specific'),
api_format_unrestricted: props.user.allowed_api_formats == null, allowed_api_formats_mode: props.user.allowed_api_formats_mode ?? (props.user.allowed_api_formats == null ? 'unrestricted' : 'specific'),
model_unrestricted: props.user.allowed_models == null, allowed_models_mode: props.user.allowed_models_mode ?? (props.user.allowed_models == null ? 'unrestricted' : 'specific'),
rate_limit_inherited: props.user.rate_limit == null, rate_limit_mode: props.user.rate_limit_mode ?? (props.user.rate_limit == null ? 'system' : 'custom'),
allowed_providers: props.user.allowed_providers ? [...props.user.allowed_providers] : [], allowed_providers: props.user.allowed_providers ? [...props.user.allowed_providers] : [],
allowed_api_formats: props.user.allowed_api_formats ? [...props.user.allowed_api_formats] : [], allowed_api_formats: props.user.allowed_api_formats ? [...props.user.allowed_api_formats] : [],
allowed_models: props.user.allowed_models ? [...props.user.allowed_models] : [], allowed_models: props.user.allowed_models ? [...props.user.allowed_models] : [],
rate_limit: props.user.rate_limit ?? undefined, rate_limit: props.user.rate_limit ?? undefined,
group_ids: props.user.group_ids ? [...props.user.group_ids] : [],
} }
} }
@@ -545,16 +589,21 @@ async function handleSubmit() {
email: form.value.email.trim() || '', email: form.value.email.trim() || '',
unlimited: form.value.unlimited, unlimited: form.value.unlimited,
role: form.value.role, role: form.value.role,
allowed_providers: form.value.provider_unrestricted allowed_providers: form.value.allowed_providers_mode === 'specific'
? null ? [...form.value.allowed_providers]
: [...form.value.allowed_providers], : null,
allowed_api_formats: form.value.api_format_unrestricted allowed_providers_mode: form.value.allowed_providers_mode,
? null allowed_api_formats: form.value.allowed_api_formats_mode === 'specific'
: [...form.value.allowed_api_formats], ? [...form.value.allowed_api_formats]
allowed_models: form.value.model_unrestricted : null,
? null allowed_api_formats_mode: form.value.allowed_api_formats_mode,
: [...form.value.allowed_models], allowed_models: form.value.allowed_models_mode === 'specific'
rate_limit: form.value.rate_limit_inherited ? null : (form.value.rate_limit ?? 0), ? [...form.value.allowed_models]
: null,
allowed_models_mode: form.value.allowed_models_mode,
rate_limit: form.value.rate_limit_mode === 'custom' ? (form.value.rate_limit ?? 0) : null,
rate_limit_mode: form.value.rate_limit_mode,
group_ids: [...form.value.group_ids],
} }
if (isEditMode.value && props.user?.id) { if (isEditMode.value && props.user?.id) {

View File

@@ -0,0 +1,543 @@
<template>
<Dialog
:model-value="open"
title="用户分组"
description="管理用户组、默认注册组、成员和组级访问控制"
size="6xl"
persistent
@update:model-value="handleDialogUpdate"
>
<div class="grid min-h-[560px] gap-4 lg:grid-cols-[17rem_minmax(0,1fr)]">
<div class="rounded-xl border border-border/70 bg-muted/20 p-3">
<div class="mb-3 flex items-center justify-between gap-2">
<Label class="text-sm font-semibold">分组</Label>
<Button
size="sm"
class="h-8 px-2 text-xs"
@click="startCreate"
>
<Plus class="mr-1.5 h-3.5 w-3.5" />
新建
</Button>
</div>
<div
v-if="loading"
class="rounded-lg border border-dashed border-border/70 px-3 py-8 text-center text-xs text-muted-foreground"
>
正在加载...
</div>
<div
v-else-if="groups.length === 0"
class="rounded-lg border border-dashed border-border/70 px-3 py-8 text-center text-xs text-muted-foreground"
>
暂无分组
</div>
<div
v-else
class="space-y-1.5"
>
<button
v-for="group in groups"
:key="group.id"
type="button"
:class="groupButtonClass(group.id)"
@click="selectGroup(group.id)"
>
<span class="min-w-0 flex-1 text-left">
<span class="flex items-center gap-1.5">
<span class="truncate text-sm font-medium">{{ group.name }}</span>
<Badge
v-if="group.is_default"
variant="secondary"
class="h-5 px-1.5 py-0 text-[10px]"
>
默认
</Badge>
</span>
<span class="mt-0.5 block text-[11px] text-muted-foreground">
优先级 {{ group.priority }}
</span>
</span>
<ChevronRight class="h-4 w-4 shrink-0 text-muted-foreground" />
</button>
</div>
</div>
<div class="min-w-0 rounded-xl border border-border/70 bg-background p-4">
<div class="mb-4 flex flex-wrap items-center justify-between gap-3">
<div class="min-w-0">
<h4 class="truncate text-base font-semibold text-foreground">
{{ editingGroupId ? '编辑分组' : '新建分组' }}
</h4>
<p class="text-xs text-muted-foreground">
{{ selectedGroup?.is_default ? '当前为自助注册默认组' : '默认组只影响本地注册和 OAuth 自动创建用户' }}
</p>
</div>
<div class="flex items-center gap-2">
<Button
v-if="editingGroupId"
variant="outline"
size="sm"
class="h-8 border-rose-200 px-2 text-xs text-rose-600 hover:bg-rose-50 dark:border-rose-900/60 dark:hover:bg-rose-950/40"
:disabled="saving"
@click="deleteSelectedGroup"
>
<Trash2 class="mr-1.5 h-3.5 w-3.5" />
删除
</Button>
</div>
</div>
<div class="grid gap-5 lg:grid-cols-2">
<div class="space-y-4">
<div class="grid gap-3 sm:grid-cols-[minmax(0,1fr)_8rem]">
<div class="space-y-2">
<Label class="text-sm font-medium">名称</Label>
<Input
v-model="form.name"
class="h-10"
placeholder="例如:生产团队"
/>
</div>
<div class="space-y-2">
<Label class="text-sm font-medium">优先级</Label>
<Input
:model-value="form.priority"
type="number"
class="h-10"
@update:model-value="(value) => form.priority = parseNumberInput(value, { min: -10000, max: 10000 }) ?? 0"
/>
</div>
</div>
<div class="flex items-center justify-between gap-3 rounded-lg border border-border/70 bg-muted/20 px-3 py-2">
<div class="min-w-0">
<Label class="text-sm font-medium">默认注册组</Label>
<div class="mt-0.5 text-[11px] text-muted-foreground">
本地注册 / OAuth 自动创建
</div>
</div>
<Switch
v-model="form.is_default"
class="shrink-0"
/>
</div>
<div class="space-y-2">
<Label class="text-sm font-medium">描述</Label>
<Textarea
v-model="form.description"
class="min-h-20"
placeholder="可选"
/>
</div>
<div class="space-y-2">
<Label class="text-sm font-medium">成员</Label>
<MultiSelect
v-model="memberUserIds"
:options="userOptions"
:search-threshold="0"
placeholder="选择用户"
empty-text="暂无用户"
no-results-text="未找到匹配用户"
/>
</div>
</div>
<div class="space-y-4 lg:border-l lg:border-border/60 lg:pl-5">
<div class="flex items-baseline justify-between gap-2 pb-2 border-b border-border/60">
<span class="text-sm font-medium">组权限</span>
<span class="text-[11px] text-muted-foreground">
用户选择继承时按优先级取首个已配置组
</span>
</div>
<PolicyFieldEditor
v-model:mode="form.allowed_providers_mode"
v-model:values="form.allowed_providers"
label="允许的提供商"
:options="providerOptions"
/>
<PolicyFieldEditor
v-model:mode="form.allowed_api_formats_mode"
v-model:values="form.allowed_api_formats"
label="允许的端点"
:options="apiFormatOptions"
/>
<PolicyFieldEditor
v-model:mode="form.allowed_models_mode"
v-model:values="form.allowed_models"
label="允许的模型"
:options="modelOptions"
/>
<div class="space-y-2">
<Label class="text-sm font-medium">速率限制 (请求/分钟)</Label>
<div class="flex items-start gap-2">
<div class="w-28 shrink-0">
<Select v-model="form.rate_limit_mode">
<SelectTrigger class="h-10 w-full">
<SelectValue />
</SelectTrigger>
<SelectContent>
<SelectItem value="inherit">不配置</SelectItem>
<SelectItem value="system">系统默认</SelectItem>
<SelectItem value="custom">指定数值</SelectItem>
</SelectContent>
</Select>
</div>
<div class="min-w-0 flex-1">
<Input
:model-value="form.rate_limit ?? ''"
type="number"
min="0"
max="10000"
class="h-10"
:disabled="form.rate_limit_mode !== 'custom'"
:placeholder="rateLimitPlaceholder"
@update:model-value="(value) => form.rate_limit = parseNumberInput(value, { min: 0, max: 10000 })"
/>
</div>
</div>
</div>
</div>
</div>
</div>
</div>
<template #footer>
<Button
variant="outline"
:disabled="saving"
@click="emit('close')"
>
关闭
</Button>
<Button
:disabled="saving || !form.name.trim()"
@click="saveGroup"
>
{{ saving ? '保存中...' : '保存分组' }}
</Button>
</template>
</Dialog>
</template>
<script setup lang="ts">
import { computed, defineComponent, h, ref, watch } from 'vue'
import { ChevronRight, Plus, Trash2 } from 'lucide-vue-next'
import {
Badge,
Button,
Dialog,
Input,
Label,
Select,
SelectContent,
SelectItem,
SelectTrigger,
SelectValue,
Switch,
Textarea,
} from '@/components/ui'
import { MultiSelect } from '@/components/common'
import { useUsersStore } from '@/stores/users'
import { useToast } from '@/composables/useToast'
import { useConfirm } from '@/composables/useConfirm'
import { parseApiError } from '@/utils/errorParser'
import { parseNumberInput } from '@/utils/form'
import { cn } from '@/lib/utils'
import { useUserAccessControlOptions } from '@/features/users/composables/useUserAccessControlOptions'
import type {
ListPolicyMode,
RateLimitPolicyMode,
UpsertUserGroupRequest,
User,
UserGroup,
} from '@/api/users'
const PolicyFieldEditor = defineComponent({
name: 'PolicyFieldEditor',
props: {
label: { type: String, required: true },
mode: { type: String as () => ListPolicyMode, required: true },
values: { type: Array as () => string[], required: true },
options: { type: Array as () => Array<{ label: string; value: string }>, required: true },
},
emits: ['update:mode', 'update:values'],
setup(props, { emit }) {
return () => h('div', { class: 'space-y-2' }, [
h(Label, { class: 'text-sm font-medium' }, () => props.label),
h('div', { class: 'flex items-start gap-2' }, [
h('div', { class: 'w-28 shrink-0' }, [
h(Select, {
modelValue: props.mode,
'onUpdate:modelValue': (value: string) => emit('update:mode', value),
}, () => [
h(SelectTrigger, { class: 'h-10 w-full' }, () => h(SelectValue)),
h(SelectContent, null, () => [
h(SelectItem, { value: 'inherit' }, () => '不配置'),
h(SelectItem, { value: 'unrestricted' }, () => '不限制'),
h(SelectItem, { value: 'specific' }, () => '指定列表'),
h(SelectItem, { value: 'deny_all' }, () => '全部禁用'),
]),
]),
]),
h('div', { class: 'min-w-0 flex-1' }, [
h(MultiSelect, {
modelValue: props.values,
'onUpdate:modelValue': (value: string[]) => emit('update:values', value),
options: props.options,
disabled: props.mode !== 'specific',
searchThreshold: 0,
placeholder: listPolicyValuePlaceholder(props.mode),
emptyText: '暂无选项',
dropdownMinWidth: '16rem',
}),
]),
]),
])
},
})
function listPolicyValuePlaceholder(mode: ListPolicyMode): string {
switch (mode) {
case 'inherit':
return '该组不配置此项'
case 'unrestricted':
return '不限制所有选项'
case 'deny_all':
return '全部禁用'
case 'specific':
default:
return '未选择时表示全部禁用'
}
}
const props = defineProps<{
open: boolean
users: User[]
}>()
const emit = defineEmits<{
close: []
changed: []
}>()
const usersStore = useUsersStore()
const { success, error } = useToast()
const { confirmDanger } = useConfirm()
const {
providerOptions,
apiFormatOptions,
modelOptions,
loadAccessControlOptions,
} = useUserAccessControlOptions()
const loading = ref(false)
const saving = ref(false)
const groups = ref<UserGroup[]>([])
const defaultGroupId = ref<string | null>(null)
const editingGroupId = ref<string | null>(null)
const memberUserIds = ref<string[]>([])
const form = ref({
name: '',
description: '',
priority: 0,
is_default: false,
allowed_providers_mode: 'inherit' as ListPolicyMode,
allowed_api_formats_mode: 'inherit' as ListPolicyMode,
allowed_models_mode: 'inherit' as ListPolicyMode,
allowed_providers: [] as string[],
allowed_api_formats: [] as string[],
allowed_models: [] as string[],
rate_limit_mode: 'inherit' as RateLimitPolicyMode,
rate_limit: undefined as number | undefined,
})
const selectedGroup = computed(() => groups.value.find((group) => group.id === editingGroupId.value) ?? null)
const rateLimitPlaceholder = computed(() => {
switch (form.value.rate_limit_mode) {
case 'inherit':
return '该组不配置速率'
case 'system':
return '使用系统默认'
case 'custom':
default:
return '0 = 不限速'
}
})
const userOptions = computed(() => props.users.map((user) => ({
label: `${user.username}${user.email ? ` (${user.email})` : ''}`,
value: user.id,
})))
watch(
() => props.open,
(open) => {
if (!open) return
void loadDialogData()
void loadAccessControlOptions().catch((err) => {
error(parseApiError(err, '加载访问控制选项失败'))
})
},
)
function handleDialogUpdate(value: boolean): void {
if (!value) emit('close')
}
async function loadDialogData(): Promise<void> {
loading.value = true
try {
const response = await usersStore.listUserGroups()
groups.value = response.items
defaultGroupId.value = response.default_group_id ?? null
if (editingGroupId.value && !groups.value.some((group) => group.id === editingGroupId.value)) {
editingGroupId.value = null
}
const nextGroup = editingGroupId.value
? groups.value.find((group) => group.id === editingGroupId.value) ?? null
: groups.value[0] ?? null
if (nextGroup) {
await selectGroup(nextGroup.id)
} else {
startCreate()
}
} catch (err) {
error(parseApiError(err, '加载用户分组失败'))
} finally {
loading.value = false
}
}
async function selectGroup(groupId: string): Promise<void> {
const group = groups.value.find((item) => item.id === groupId)
if (!group) return
editingGroupId.value = group.id
form.value = {
name: group.name,
description: group.description ?? '',
priority: group.priority,
is_default: group.is_default === true,
allowed_providers_mode: group.allowed_providers_mode,
allowed_api_formats_mode: group.allowed_api_formats_mode,
allowed_models_mode: group.allowed_models_mode,
allowed_providers: group.allowed_providers ? [...group.allowed_providers] : [],
allowed_api_formats: group.allowed_api_formats ? [...group.allowed_api_formats] : [],
allowed_models: group.allowed_models ? [...group.allowed_models] : [],
rate_limit_mode: group.rate_limit_mode,
rate_limit: group.rate_limit ?? undefined,
}
try {
const members = await usersStore.listUserGroupMembers(group.id)
memberUserIds.value = members.map((member) => member.user_id)
} catch (err) {
memberUserIds.value = []
error(parseApiError(err, '加载分组成员失败'))
}
}
function startCreate(): void {
editingGroupId.value = null
form.value = {
name: '',
description: '',
priority: 0,
is_default: false,
allowed_providers_mode: 'inherit',
allowed_api_formats_mode: 'inherit',
allowed_models_mode: 'inherit',
allowed_providers: [],
allowed_api_formats: [],
allowed_models: [],
rate_limit_mode: 'inherit',
rate_limit: undefined,
}
memberUserIds.value = []
}
function groupButtonClass(groupId: string): string {
return cn(
'flex w-full items-center gap-2 rounded-lg border px-3 py-2 transition-colors',
editingGroupId.value === groupId
? 'border-primary/50 bg-primary/10'
: 'border-transparent hover:border-border hover:bg-background',
)
}
function buildPayload(): UpsertUserGroupRequest {
return {
name: form.value.name.trim(),
description: form.value.description.trim() || null,
priority: form.value.priority,
allowed_providers_mode: form.value.allowed_providers_mode,
allowed_api_formats_mode: form.value.allowed_api_formats_mode,
allowed_models_mode: form.value.allowed_models_mode,
allowed_providers: form.value.allowed_providers_mode === 'specific'
? [...form.value.allowed_providers]
: null,
allowed_api_formats: form.value.allowed_api_formats_mode === 'specific'
? [...form.value.allowed_api_formats]
: null,
allowed_models: form.value.allowed_models_mode === 'specific'
? [...form.value.allowed_models]
: null,
rate_limit_mode: form.value.rate_limit_mode,
rate_limit: form.value.rate_limit_mode === 'custom'
? (form.value.rate_limit ?? 0)
: null,
}
}
async function saveGroup(): Promise<void> {
if (!form.value.name.trim()) return
saving.value = true
try {
const wasDefault = selectedGroup.value?.is_default === true
const wantsDefault = form.value.is_default
const saved = editingGroupId.value
? await usersStore.updateUserGroup(editingGroupId.value, buildPayload())
: await usersStore.createUserGroup(buildPayload())
await usersStore.replaceUserGroupMembers(saved.id, memberUserIds.value)
if (wantsDefault) {
await usersStore.setDefaultUserGroup(saved.id)
} else if (wasDefault) {
await usersStore.setDefaultUserGroup(null)
}
success('用户分组已保存')
emit('changed')
editingGroupId.value = saved.id
await loadDialogData()
} catch (err) {
error(parseApiError(err, '保存用户分组失败'))
} finally {
saving.value = false
}
}
async function deleteSelectedGroup(): Promise<void> {
if (!selectedGroup.value) return
const group = selectedGroup.value
const confirmed = await confirmDanger(
`确定要删除用户分组 ${group.name} 吗?成员关系会一并清理。`,
'删除用户分组',
)
if (!confirmed) return
saving.value = true
try {
await usersStore.deleteUserGroup(group.id)
success('用户分组已删除')
emit('changed')
editingGroupId.value = null
await loadDialogData()
} catch (err) {
error(parseApiError(err, '删除用户分组失败'))
} finally {
saving.value = false
}
}
</script>

View File

@@ -12,6 +12,10 @@ import {
type ResolveUserBatchSelectionResponse, type ResolveUserBatchSelectionResponse,
type UserBatchActionRequest, type UserBatchActionRequest,
type UserBatchActionResponse, type UserBatchActionResponse,
type UserGroup,
type UserGroupMember,
type UpsertUserGroupRequest,
type ListUserGroupsResponse,
} from '@/api/users' } from '@/api/users'
import { parseApiError } from '@/utils/errorParser' import { parseApiError } from '@/utils/errorParser'
@@ -20,7 +24,13 @@ export const useUsersStore = defineStore('users', () => {
const loading = ref(false) const loading = ref(false)
const error = ref<string | null>(null) const error = ref<string | null>(null)
async function fetchUsers(options: { cacheTtlMs?: number } = {}) { async function fetchUsers(options: {
cacheTtlMs?: number
search?: string
role?: 'admin' | 'user'
is_active?: boolean
group_id?: string
} = {}) {
loading.value = true loading.value = true
error.value = null error.value = null
@@ -111,6 +121,75 @@ export const useUsersStore = defineStore('users', () => {
} }
} }
async function listUserGroups(): Promise<ListUserGroupsResponse> {
try {
return await usersApi.listUserGroups()
} catch (err: unknown) {
error.value = parseApiError(err, '获取用户分组失败')
throw err
}
}
async function createUserGroup(payload: UpsertUserGroupRequest): Promise<UserGroup> {
try {
return await usersApi.createUserGroup(payload)
} catch (err: unknown) {
error.value = parseApiError(err, '创建用户分组失败')
throw err
}
}
async function updateUserGroup(
groupId: string,
payload: UpsertUserGroupRequest,
): Promise<UserGroup> {
try {
return await usersApi.updateUserGroup(groupId, payload)
} catch (err: unknown) {
error.value = parseApiError(err, '更新用户分组失败')
throw err
}
}
async function deleteUserGroup(groupId: string): Promise<void> {
try {
await usersApi.deleteUserGroup(groupId)
} catch (err: unknown) {
error.value = parseApiError(err, '删除用户分组失败')
throw err
}
}
async function listUserGroupMembers(groupId: string): Promise<UserGroupMember[]> {
try {
return await usersApi.listUserGroupMembers(groupId)
} catch (err: unknown) {
error.value = parseApiError(err, '获取分组成员失败')
throw err
}
}
async function replaceUserGroupMembers(
groupId: string,
userIds: string[],
): Promise<UserGroupMember[]> {
try {
return await usersApi.replaceUserGroupMembers(groupId, userIds)
} catch (err: unknown) {
error.value = parseApiError(err, '更新分组成员失败')
throw err
}
}
async function setDefaultUserGroup(groupId: string | null): Promise<{ default_group_id?: string | null }> {
try {
return await usersApi.setDefaultUserGroup(groupId)
} catch (err: unknown) {
error.value = parseApiError(err, '设置默认用户组失败')
throw err
}
}
async function getUserApiKeys(userId: string): Promise<ApiKey[]> { async function getUserApiKeys(userId: string): Promise<ApiKey[]> {
try { try {
return await usersApi.getUserApiKeys(userId) return await usersApi.getUserApiKeys(userId)
@@ -199,6 +278,13 @@ export const useUsersStore = defineStore('users', () => {
deleteUser, deleteUser,
resolveBatchSelection, resolveBatchSelection,
batchAction, batchAction,
listUserGroups,
createUserGroup,
updateUserGroup,
deleteUserGroup,
listUserGroupMembers,
replaceUserGroupMembers,
setDefaultUserGroup,
getUserApiKeys, getUserApiKeys,
createApiKey, createApiKey,
updateApiKey, updateApiKey,

View File

@@ -15,6 +15,15 @@
</h3> </h3>
<div class="flex items-center gap-2"> <div class="flex items-center gap-2">
<!-- 新增用户按钮 --> <!-- 新增用户按钮 -->
<Button
variant="ghost"
size="icon"
class="h-8 w-8"
title="分组管理"
@click="showUserGroupsDialog = true"
>
<FolderKanban class="w-3.5 h-3.5" />
</Button>
<Button <Button
variant="ghost" variant="ghost"
size="icon" size="icon"
@@ -61,6 +70,25 @@
</SelectItem> </SelectItem>
</SelectContent> </SelectContent>
</Select> </Select>
<Select
v-model="filterGroup"
>
<SelectTrigger class="w-24 h-8 text-xs border-border/60">
<SelectValue placeholder="分组" />
</SelectTrigger>
<SelectContent>
<SelectItem value="all">
全部
</SelectItem>
<SelectItem
v-for="group in userGroups"
:key="group.id"
:value="group.id"
>
{{ group.name }}
</SelectItem>
</SelectContent>
</Select>
<Select <Select
v-model="filterStatus" v-model="filterStatus"
> >
@@ -149,9 +177,38 @@
</Select> </Select>
</div> </div>
<Select v-model="filterGroup">
<SelectTrigger class="w-32 h-8 text-xs border-border/60">
<SelectValue placeholder="全部分组" />
</SelectTrigger>
<SelectContent>
<SelectItem value="all">
全部分组
</SelectItem>
<SelectItem
v-for="group in userGroups"
:key="group.id"
:value="group.id"
>
{{ group.name }}
</SelectItem>
</SelectContent>
</Select>
<!-- 分隔线 --> <!-- 分隔线 -->
<div class="h-4 w-px bg-border" /> <div class="h-4 w-px bg-border" />
<!-- 新增用户按钮 -->
<Button
variant="ghost"
size="icon"
class="h-8 w-8"
title="分组管理"
@click="showUserGroupsDialog = true"
>
<FolderKanban class="w-3.5 h-3.5" />
</Button>
<!-- 新增用户按钮 --> <!-- 新增用户按钮 -->
<Button <Button
variant="ghost" variant="ghost"
@@ -207,7 +264,7 @@
<Button <Button
size="sm" size="sm"
class="h-7 px-3 text-[11px]" class="h-7 px-3 text-[11px]"
:disabled="selectedCount === 0 || usersStore.loading" :disabled="(selectedCount === 0 && userGroups.length === 0) || usersStore.loading"
@click="openUserBatchDialog" @click="openUserBatchDialog"
> >
批量操作 批量操作
@@ -317,6 +374,19 @@
> >
{{ user.email || '-' }} {{ user.email || '-' }}
</div> </div>
<div
v-if="user.groups?.length"
class="mt-1 flex flex-wrap gap-1"
>
<Badge
v-for="group in user.groups"
:key="group.id"
variant="outline"
class="h-5 px-1.5 py-0 text-[10px]"
>
{{ group.name }}
</Badge>
</div>
</div> </div>
</div> </div>
</TableCell> </TableCell>
@@ -550,9 +620,18 @@
<Badge <Badge
variant="secondary" variant="secondary"
class="h-5 px-1.5 py-0 text-[10px] font-medium" class="h-5 px-1.5 py-0 text-[10px] font-medium"
:title="formatUserEffectiveRateLimitSource(user)"
> >
{{ formatRateLimitInheritable(user.rate_limit) }} {{ formatRateLimitInheritable(user.rate_limit) }}
</Badge> </Badge>
<Badge
v-for="group in user.groups || []"
:key="group.id"
variant="outline"
class="h-5 px-1.5 py-0 text-[10px] font-medium"
>
{{ group.name }}
</Badge>
</div> </div>
<div class="rounded-xl border border-border/60 bg-muted/40 p-3.5"> <div class="rounded-xl border border-border/60 bg-muted/40 p-3.5">
@@ -697,6 +776,7 @@
ref="userFormDialogRef" ref="userFormDialogRef"
:open="showUserFormDialog" :open="showUserFormDialog"
:user="editingUser" :user="editingUser"
:groups="userGroups"
@close="closeUserFormDialog" @close="closeUserFormDialog"
@submit="handleUserFormSubmit" @submit="handleUserFormSubmit"
/> />
@@ -707,10 +787,18 @@
:select-all-filtered="selectAllFiltered" :select-all-filtered="selectAllFiltered"
:selected-count="selectedCount" :selected-count="selectedCount"
:filters="batchSelectionFilters" :filters="batchSelectionFilters"
:groups="userGroups"
@close="showUserBatchDialog = false" @close="showUserBatchDialog = false"
@completed="handleUserBatchCompleted" @completed="handleUserBatchCompleted"
/> />
<UserGroupsDialog
:open="showUserGroupsDialog"
:users="usersStore.users"
@close="showUserGroupsDialog = false"
@changed="handleUserGroupsChanged"
/>
<!-- API Keys 管理对话框 --> <!-- API Keys 管理对话框 -->
<Dialog <Dialog
v-model="showApiKeysDialog" v-model="showApiKeysDialog"
@@ -1134,7 +1222,7 @@
<script setup lang="ts"> <script setup lang="ts">
import { ref, computed, onMounted, watch } from 'vue' import { ref, computed, onMounted, watch } from 'vue'
import { useUsersStore } from '@/stores/users' import { useUsersStore } from '@/stores/users'
import type { User, ApiKey, UserSession, UserBatchActionResponse, UserBatchSelectionFilters } from '@/api/users' import type { User, ApiKey, UserSession, UserBatchActionResponse, UserBatchSelectionFilters, UserGroup } from '@/api/users'
import { formatSessionMeta } from '@/types/session' import { formatSessionMeta } from '@/types/session'
import { adminWalletApi, type AdminWallet } from '@/api/admin-wallets' import { adminWalletApi, type AdminWallet } from '@/api/admin-wallets'
import { useToast } from '@/composables/useToast' import { useToast } from '@/composables/useToast'
@@ -1184,12 +1272,14 @@ import {
CheckCircle, CheckCircle,
Lock, Lock,
LockOpen, LockOpen,
MonitorSmartphone MonitorSmartphone,
FolderKanban,
} from 'lucide-vue-next' } from 'lucide-vue-next'
// 功能组件 // 功能组件
import UserFormDialog, { type UserFormData } from '@/features/users/components/UserFormDialog.vue' import UserFormDialog, { type UserFormData } from '@/features/users/components/UserFormDialog.vue'
import UserBatchActionDialog from '@/features/users/components/UserBatchActionDialog.vue' import UserBatchActionDialog from '@/features/users/components/UserBatchActionDialog.vue'
import UserGroupsDialog from '@/features/users/components/UserGroupsDialog.vue'
import WalletOpsDrawer from '@/features/wallet/components/WalletOpsDrawer.vue' import WalletOpsDrawer from '@/features/wallet/components/WalletOpsDrawer.vue'
import { parseApiError } from '@/utils/errorParser' import { parseApiError } from '@/utils/errorParser'
import { formatTokens, formatRateLimitInheritable, formatRateLimitSimple, isRateLimitInherited, isRateLimitUnlimited } from '@/utils/format' import { formatTokens, formatRateLimitInheritable, formatRateLimitSimple, isRateLimitInherited, isRateLimitUnlimited } from '@/utils/format'
@@ -1233,10 +1323,13 @@ const userWalletMap = ref<Record<string, AdminWallet>>({})
const showWalletActionDialogState = ref(false) const showWalletActionDialogState = ref(false)
const walletActionTarget = ref<{ user: User; wallet: AdminWallet } | null>(null) const walletActionTarget = ref<{ user: User; wallet: AdminWallet } | null>(null)
const showUserBatchDialog = ref(false) const showUserBatchDialog = ref(false)
const showUserGroupsDialog = ref(false)
const searchQuery = ref('') const searchQuery = ref('')
const filterRole = ref('all') const filterRole = ref('all')
const filterStatus = ref('all') const filterStatus = ref('all')
const filterGroup = ref('all')
const userGroups = ref<UserGroup[]>([])
const userRoleFilterOptions = [ const userRoleFilterOptions = [
{ value: 'all', label: '全部角色' }, { value: 'all', label: '全部角色' },
{ value: 'admin', label: '管理员' }, { value: 'admin', label: '管理员' },
@@ -1285,6 +1378,10 @@ const filteredUsers = computed(() => {
) )
} }
if (filterGroup.value !== 'all') {
filtered = filtered.filter(u => (u.groups || []).some(group => group.id === filterGroup.value))
}
return filtered return filtered
}) })
@@ -1322,11 +1419,12 @@ const batchSelectionFilters = computed<UserBatchSelectionFilters>(() => {
if (filterRole.value === 'admin' || filterRole.value === 'user') filters.role = filterRole.value if (filterRole.value === 'admin' || filterRole.value === 'user') filters.role = filterRole.value
if (filterStatus.value === 'active') filters.is_active = true if (filterStatus.value === 'active') filters.is_active = true
if (filterStatus.value === 'inactive') filters.is_active = false if (filterStatus.value === 'inactive') filters.is_active = false
if (filterGroup.value !== 'all') filters.group_id = filterGroup.value
return filters return filters
}) })
// Watch filter changes and reset to first page // Watch filter changes and reset to first page
watch([searchQuery, filterRole, filterStatus], () => { watch([searchQuery, filterRole, filterStatus, filterGroup], () => {
currentPage.value = 1 currentPage.value = 1
resetBatchSelection() resetBatchSelection()
}) })
@@ -1339,14 +1437,33 @@ onMounted(() => {
async function refreshUsers(options: { preferCache?: boolean } = {}) { async function refreshUsers(options: { preferCache?: boolean } = {}) {
const cacheTtlMs = options.preferCache ? USERS_PAGE_CACHE_TTL_MS : 0 const cacheTtlMs = options.preferCache ? USERS_PAGE_CACHE_TTL_MS : 0
await usersStore.fetchUsers({ cacheTtlMs }) await Promise.all([
usersStore.fetchUsers({ cacheTtlMs }),
loadUserGroups(),
])
void loadUserWallets({ void loadUserWallets({
cacheTtlMs: options.preferCache ? USER_WALLETS_CACHE_TTL_MS : 0, cacheTtlMs: options.preferCache ? USER_WALLETS_CACHE_TTL_MS : 0,
}) })
} }
async function loadUserGroups(): Promise<void> {
try {
const response = await usersStore.listUserGroups()
userGroups.value = response.items
if (filterGroup.value !== 'all' && !userGroups.value.some((group) => group.id === filterGroup.value)) {
filterGroup.value = 'all'
}
} catch (err) {
log.error('加载用户分组失败:', err)
}
}
async function handleUserGroupsChanged(): Promise<void> {
await refreshUsers()
}
function openUserBatchDialog(): void { function openUserBatchDialog(): void {
if (selectedCount.value === 0) return if (selectedCount.value === 0 && userGroups.value.length === 0) return
showUserBatchDialog.value = true showUserBatchDialog.value = true
} }
@@ -1429,6 +1546,18 @@ function formatConcurrentLimitSimple(concurrentLimit?: number | null): string {
return `${concurrentLimit} 并发` return `${concurrentLimit} 并发`
} }
function formatUserEffectiveRateLimitSource(user: User): string {
const source = user.effective_policy?.rate_limit
if (!source) return ''
if (source.source === 'group' && source.group_name) {
return `继承自分组:${source.group_name}`
}
if (source.source === 'user') {
return '用户单独配置'
}
return '系统默认'
}
function isNegativeWalletValue(value: number | null): boolean { function isNegativeWalletValue(value: number | null): boolean {
return typeof value === 'number' && value < 0 return typeof value === 'number' && value < 0
} }
@@ -1470,7 +1599,12 @@ function editUser(user: User) {
allowed_providers: user.allowed_providers == null ? null : [...user.allowed_providers], allowed_providers: user.allowed_providers == null ? null : [...user.allowed_providers],
allowed_api_formats: user.allowed_api_formats == null ? null : [...user.allowed_api_formats], allowed_api_formats: user.allowed_api_formats == null ? null : [...user.allowed_api_formats],
allowed_models: user.allowed_models == null ? null : [...user.allowed_models], allowed_models: user.allowed_models == null ? null : [...user.allowed_models],
rate_limit: user.rate_limit ?? null rate_limit: user.rate_limit ?? null,
allowed_providers_mode: user.allowed_providers_mode ?? (user.allowed_providers == null ? 'unrestricted' : 'specific'),
allowed_api_formats_mode: user.allowed_api_formats_mode ?? (user.allowed_api_formats == null ? 'unrestricted' : 'specific'),
allowed_models_mode: user.allowed_models_mode ?? (user.allowed_models == null ? 'unrestricted' : 'specific'),
rate_limit_mode: user.rate_limit_mode ?? (user.rate_limit == null ? 'system' : 'custom'),
group_ids: (user.groups || []).map(group => group.id),
} }
showUserFormDialog.value = true showUserFormDialog.value = true
} }
@@ -1491,9 +1625,14 @@ async function handleUserFormSubmit(data: UserFormData & { password?: string; un
unlimited: data.unlimited, unlimited: data.unlimited,
role: data.role, role: data.role,
allowed_providers: data.allowed_providers, allowed_providers: data.allowed_providers,
allowed_providers_mode: data.allowed_providers_mode,
allowed_api_formats: data.allowed_api_formats, allowed_api_formats: data.allowed_api_formats,
allowed_api_formats_mode: data.allowed_api_formats_mode,
allowed_models: data.allowed_models, allowed_models: data.allowed_models,
rate_limit: data.rate_limit ?? null allowed_models_mode: data.allowed_models_mode,
rate_limit: data.rate_limit ?? null,
rate_limit_mode: data.rate_limit_mode,
group_ids: data.group_ids ?? [],
} }
if (data.password) { if (data.password) {
updateData.password = data.password updateData.password = data.password
@@ -1511,9 +1650,14 @@ async function handleUserFormSubmit(data: UserFormData & { password?: string; un
unlimited: data.unlimited, unlimited: data.unlimited,
role: data.role, role: data.role,
allowed_providers: data.allowed_providers, allowed_providers: data.allowed_providers,
allowed_providers_mode: data.allowed_providers_mode,
allowed_api_formats: data.allowed_api_formats, allowed_api_formats: data.allowed_api_formats,
allowed_api_formats_mode: data.allowed_api_formats_mode,
allowed_models: data.allowed_models, allowed_models: data.allowed_models,
rate_limit: data.rate_limit ?? null allowed_models_mode: data.allowed_models_mode,
rate_limit: data.rate_limit ?? null,
rate_limit_mode: data.rate_limit_mode,
group_ids: data.group_ids ?? [],
}) })
// 如果创建时指定为禁用,则更新状态 // 如果创建时指定为禁用,则更新状态
if (data.is_active === false && newUser) { if (data.is_active === false && newUser) {