fix: align user group access controls

This commit is contained in:
fawney19
2026-05-10 17:28:23 +08:00
parent a8ac944959
commit d246f6b42c
33 changed files with 1955 additions and 1015 deletions
@@ -3,6 +3,7 @@ use super::{
format_optional_datetime_iso8601, normalize_admin_user_api_formats,
normalize_admin_user_string_list,
};
use crate::constants::DEFAULT_USER_GROUP_CONFIG_KEY;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::attach_admin_audit_response;
use crate::GatewayError;
@@ -14,16 +15,12 @@ use axum::{
};
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,
@@ -121,6 +118,11 @@ pub(in super::super) async fn build_admin_update_user_group_response(
Ok(value) => value,
Err(detail) => return Ok(bad_request_owned(detail)),
};
if read_default_user_group_id(state).await?.as_deref() == Some(group_id.as_str())
&& !is_unrestricted_default_group_record(&record)
{
return Ok(bad_request_owned("默认用户组不能配置访问限制".to_string()));
}
let group = match state.update_user_group(&group_id, record).await {
Ok(Some(group)) => group,
Ok(None) => return Ok(not_found("用户分组不存在")),
@@ -152,9 +154,7 @@ pub(in super::super) async fn build_admin_delete_user_group_response(
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?;
return Ok(bad_request_owned("默认用户组不能删除".to_string()));
}
if !state.delete_user_group(&group_id).await? {
return Ok(not_found("用户分组不存在"));
@@ -214,6 +214,9 @@ pub(in super::super) async fn build_admin_replace_user_group_members_response(
if state.find_user_group_by_id(&group_id).await?.is_none() {
return Ok(not_found("用户分组不存在"));
}
if read_default_user_group_id(state).await?.as_deref() == Some(group_id.as_str()) {
return Ok(bad_request_owned("默认用户组成员由系统维护".to_string()));
}
let payload = match parse_members_payload(request_body) {
Ok(value) => value,
Err(detail) => return Ok(bad_request_owned(detail)),
@@ -272,8 +275,11 @@ pub(in super::super) async fn build_admin_set_default_user_group_response(
.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() {
let Some(group) = state.find_user_group_by_id(group_id).await? else {
return Ok(bad_request_owned("默认用户组不存在".to_string()));
};
if !is_unrestricted_default_group(&group) {
return Ok(bad_request_owned("默认用户组不能配置访问限制".to_string()));
}
state
.upsert_system_config_json_value(
@@ -287,18 +293,23 @@ pub(in super::super) async fn build_admin_set_default_user_group_response(
.delete_system_config_value(DEFAULT_USER_GROUP_CONFIG_KEY)
.await?;
}
Ok(Json(json!({ "default_group_id": group_id })).into_response())
let effective_group_id = read_default_user_group_id(state).await?;
if let Some(group_id) = effective_group_id.as_deref() {
state.add_all_users_to_group(group_id).await?;
}
Ok(attach_admin_audit_response(
Json(json!({ "default_group_id": effective_group_id })).into_response(),
"admin_default_user_group_set",
"set_default_user_group",
"user_group",
group_id.as_deref().unwrap_or("default_user_group"),
))
}
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()))
state.effective_default_user_group_id().await
}
fn parse_group_record(
@@ -327,7 +338,7 @@ fn parse_group_record(
.description
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty()),
priority: payload.priority.unwrap_or_default(),
priority: 0,
allowed_providers,
allowed_providers_mode: normalize_list_mode(&payload.allowed_providers_mode)?,
allowed_api_formats,
@@ -358,7 +369,6 @@ fn user_group_payload(
"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,
@@ -407,6 +417,30 @@ fn normalize_ids(values: Vec<String>) -> Vec<String> {
.collect()
}
fn is_unrestricted_default_group(group: &aether_data::repository::users::StoredUserGroup) -> bool {
list_mode_has_no_restriction(&group.allowed_providers_mode)
&& list_mode_has_no_restriction(&group.allowed_api_formats_mode)
&& list_mode_has_no_restriction(&group.allowed_models_mode)
&& rate_mode_has_no_restriction(&group.rate_limit_mode)
}
fn is_unrestricted_default_group_record(
record: &aether_data::repository::users::UpsertUserGroupRecord,
) -> bool {
list_mode_has_no_restriction(&record.allowed_providers_mode)
&& list_mode_has_no_restriction(&record.allowed_api_formats_mode)
&& list_mode_has_no_restriction(&record.allowed_models_mode)
&& rate_mode_has_no_restriction(&record.rate_limit_mode)
}
fn list_mode_has_no_restriction(mode: &str) -> bool {
matches!(mode, "inherit" | "unrestricted")
}
fn rate_mode_has_no_restriction(mode: &str) -> bool {
matches!(mode, "inherit" | "system")
}
fn user_group_id_from_path(request_path: &str) -> Option<String> {
let value = request_path
.strip_prefix("/api/admin/user-groups/")?
@@ -190,7 +190,10 @@ pub(in super::super) async fn build_admin_create_user_response(
},
None => legacy_admin_rate_limit_policy_mode(payload.rate_limit),
};
let group_ids = normalize_admin_user_group_ids(payload.group_ids);
let requested_group_ids = normalize_admin_user_group_ids(payload.group_ids);
let group_ids = state
.include_default_user_group_ids(&requested_group_ids)
.await?;
let groups = if group_ids.is_empty() {
Vec::new()
} else {
@@ -131,7 +131,6 @@ pub(super) fn user_group_badge_payload(
json!({
"id": group.id,
"name": group.name,
"priority": group.priority,
})
}
@@ -149,10 +148,8 @@ fn effective_policy_payload(
) -> 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))
left.name
.cmp(&right.name)
.then_with(|| left.id.cmp(&right.id))
});
json!({
@@ -186,50 +183,28 @@ fn effective_list_policy_payload(
&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)
let mut effective = None;
let mut group_sources = Vec::new();
for group in groups {
let (mode, values) = group_field(group);
if let Some(restriction) = list_restriction_from_mode(mode, values.cloned()) {
effective = intersect_list_policies(effective, Some(restriction));
group_sources.push(group);
}
_ => policy_payload("unrestricted", serde_json::Value::Null, "fallback", None),
}
let mut has_user_source = false;
if let Some(restriction) = list_restriction_from_mode(user_mode, user_values.cloned()) {
effective = intersect_list_policies(effective, Some(restriction));
has_user_source = true;
}
let (mode, value) = match effective {
Some(values) if values.is_empty() => ("deny_all", json!(Vec::<String>::new())),
Some(values) => ("specific", json!(values)),
None => ("unrestricted", serde_json::Value::Null),
};
let source = combined_policy_source(has_user_source, group_sources.len(), "fallback");
policy_payload(mode, value, source, group_sources.as_slice())
}
fn effective_rate_limit_policy_payload(
@@ -237,34 +212,26 @@ fn effective_rate_limit_policy_payload(
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)
let mut effective = None;
let mut group_sources = Vec::new();
for group in groups {
if let Some(restriction) =
rate_limit_restriction_from_mode(&group.rate_limit_mode, group.rate_limit)
{
effective = intersect_rate_limit_policies(effective, Some(restriction));
group_sources.push(group);
}
_ => policy_payload("system", serde_json::Value::Null, "fallback", None),
}
let mut has_user_source = false;
if let Some(restriction) = rate_limit_restriction_from_mode(user_mode, user_rate_limit) {
effective = intersect_rate_limit_policies(effective, Some(restriction));
has_user_source = true;
}
let source = combined_policy_source(has_user_source, group_sources.len(), "fallback");
match rate_limit_policy_value(effective) {
Some(rate_limit) => policy_payload("custom", json!(rate_limit), source, &group_sources),
None => policy_payload("system", serde_json::Value::Null, source, &group_sources),
}
}
@@ -272,17 +239,113 @@ fn policy_payload(
mode: &str,
value: serde_json::Value,
source: &str,
group: Option<&aether_data::repository::users::StoredUserGroup>,
groups: &[&aether_data::repository::users::StoredUserGroup],
) -> serde_json::Value {
let single_group = groups.first().copied().filter(|_| groups.len() == 1);
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()),
"group_id": single_group.map(|group| group.id.as_str()),
"group_name": single_group.map(|group| group.name.as_str()),
"group_ids": groups.iter().map(|group| group.id.as_str()).collect::<Vec<_>>(),
"group_names": groups.iter().map(|group| group.name.as_str()).collect::<Vec<_>>(),
})
}
fn list_restriction_from_mode(mode: &str, values: Option<Vec<String>>) -> Option<Vec<String>> {
match mode {
"specific" => Some(values.unwrap_or_default()),
"deny_all" => Some(Vec::new()),
_ => None,
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum RateLimitRestriction {
Unlimited,
Limited(i32),
}
fn rate_limit_restriction_from_mode(
mode: &str,
rate_limit: Option<i32>,
) -> Option<RateLimitRestriction> {
match mode {
"custom" => {
let rate_limit = rate_limit.unwrap_or(0).max(0);
if rate_limit == 0 {
Some(RateLimitRestriction::Unlimited)
} else {
Some(RateLimitRestriction::Limited(rate_limit))
}
}
_ => None,
}
}
fn intersect_list_policies(
left: Option<Vec<String>>,
right: Option<Vec<String>>,
) -> Option<Vec<String>> {
match (left, right) {
(None, None) => None,
(Some(values), None) | (None, Some(values)) => Some(values),
(Some(left_values), Some(right_values)) => {
let right_values = right_values
.into_iter()
.collect::<std::collections::BTreeSet<_>>();
Some(
left_values
.into_iter()
.filter(|value| right_values.contains(value))
.collect(),
)
}
}
}
fn intersect_rate_limit_policies(
left: Option<RateLimitRestriction>,
right: Option<RateLimitRestriction>,
) -> Option<RateLimitRestriction> {
match (left, right) {
(None, None) => None,
(Some(value), None) | (None, Some(value)) => Some(value),
(Some(RateLimitRestriction::Unlimited), Some(RateLimitRestriction::Unlimited)) => {
Some(RateLimitRestriction::Unlimited)
}
(Some(RateLimitRestriction::Limited(value)), Some(RateLimitRestriction::Unlimited))
| (Some(RateLimitRestriction::Unlimited), Some(RateLimitRestriction::Limited(value))) => {
Some(RateLimitRestriction::Limited(value))
}
(Some(RateLimitRestriction::Limited(left)), Some(RateLimitRestriction::Limited(right))) => {
Some(RateLimitRestriction::Limited(left.min(right)))
}
}
}
fn rate_limit_policy_value(policy: Option<RateLimitRestriction>) -> Option<i32> {
match policy {
None => None,
Some(RateLimitRestriction::Unlimited) => Some(0),
Some(RateLimitRestriction::Limited(value)) => Some(value),
}
}
fn combined_policy_source(
has_user_source: bool,
group_source_count: usize,
fallback_source: &'static str,
) -> &'static str {
match (has_user_source, group_source_count) {
(true, 0) => "user",
(false, 1) => "group",
(false, 0) => fallback_source,
_ => "combined",
}
}
pub(super) fn admin_user_id_from_detail_path(request_path: &str) -> Option<String> {
let value = request_path
.strip_prefix("/api/admin/users/")?
@@ -248,7 +248,12 @@ pub(in super::super) async fn build_admin_update_user_response(
None
};
let group_ids = if field_presence.contains("group_ids") {
Some(normalize_admin_user_group_ids(payload.group_ids))
let requested_group_ids = normalize_admin_user_group_ids(payload.group_ids);
Some(
state
.include_default_user_group_ids(&requested_group_ids)
.await?,
)
} else {
None
};
@@ -302,29 +307,6 @@ pub(in super::super) async fn build_admin_update_user_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)
@@ -394,6 +376,29 @@ pub(in super::super) async fn build_admin_update_user_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(unlimited) = payload.unlimited {
match state
@@ -44,13 +44,15 @@ use self::shared::{
admin_default_user_initial_gift, build_admin_users_bad_request_response,
build_admin_users_data_unavailable_response, build_admin_users_read_only_response,
format_optional_datetime_iso8601, legacy_admin_list_policy_mode,
legacy_admin_rate_limit_policy_mode, normalize_admin_list_policy_mode,
normalize_admin_optional_user_email, normalize_admin_rate_limit_policy_mode,
legacy_admin_rate_limit_policy_mode, normalize_admin_optional_user_email,
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_list_policy_mode, normalize_admin_rate_limit_policy_mode,
normalize_admin_user_api_formats, normalize_admin_user_string_list,
};
pub(crate) async fn maybe_build_local_admin_users_response(
request: AdminRouteRequest<'_>,
@@ -283,7 +283,7 @@ pub(crate) fn normalize_admin_user_api_formats(
Ok(Some(normalized))
}
pub(super) fn normalize_admin_list_policy_mode(value: &str) -> Result<String, String> {
pub(crate) 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())
@@ -292,7 +292,7 @@ pub(super) fn normalize_admin_list_policy_mode(value: &str) -> Result<String, St
}
}
pub(super) fn normalize_admin_rate_limit_policy_mode(value: &str) -> Result<String, String> {
pub(crate) 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()),