mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
fix: align user group access controls
This commit is contained in:
@@ -61,6 +61,8 @@ pub(crate) const TRUSTED_ADMIN_SESSION_ID_HEADER: &str = "x-aether-admin-session
|
||||
pub(crate) const TRUSTED_ADMIN_MANAGEMENT_TOKEN_ID_HEADER: &str =
|
||||
"x-aether-admin-management-token-id";
|
||||
pub(crate) const TRUSTED_RATE_LIMIT_PREFLIGHT_HEADER: &str = "x-aether-rate-limit-preflight";
|
||||
pub(crate) const DEFAULT_USER_GROUP_CONFIG_KEY: &str = "default_user_group_id";
|
||||
pub(crate) const BUILTIN_DEFAULT_USER_GROUP_ID: &str = "00000000-0000-0000-0000-000000000001";
|
||||
|
||||
pub(crate) const FRONTDOOR_REPLACEABLE_ROUTE_GROUPS: &[&str] = &["frontdoor_compat_router"];
|
||||
pub(crate) const FRONTDOOR_REPLACEABLE_MIDDLEWARE_GROUPS: &[&str] = &["cors"];
|
||||
|
||||
@@ -1663,14 +1663,12 @@ impl GatewayDataState {
|
||||
.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))
|
||||
left.name
|
||||
.cmp(&right.name)
|
||||
.then_with(|| left.id.cmp(&right.id))
|
||||
});
|
||||
|
||||
let allowed_providers = resolve_effective_list_policy(
|
||||
let mut allowed_providers = resolve_effective_list_policy(
|
||||
user.allowed_providers,
|
||||
&user.allowed_providers_mode,
|
||||
&groups,
|
||||
@@ -1681,7 +1679,7 @@ impl GatewayDataState {
|
||||
)
|
||||
},
|
||||
);
|
||||
let allowed_api_formats = resolve_effective_list_policy(
|
||||
let mut allowed_api_formats = resolve_effective_list_policy(
|
||||
user.allowed_api_formats,
|
||||
&user.allowed_api_formats_mode,
|
||||
&groups,
|
||||
@@ -1692,7 +1690,7 @@ impl GatewayDataState {
|
||||
)
|
||||
},
|
||||
);
|
||||
let allowed_models = resolve_effective_list_policy(
|
||||
let mut allowed_models = resolve_effective_list_policy(
|
||||
user.allowed_models,
|
||||
&user.allowed_models_mode,
|
||||
&groups,
|
||||
@@ -1717,6 +1715,20 @@ impl GatewayDataState {
|
||||
user_rate_limit_mode,
|
||||
&groups,
|
||||
);
|
||||
if !snapshot.api_key_is_standalone {
|
||||
constrain_api_key_list_policy_to_user_policy(
|
||||
&mut allowed_providers,
|
||||
&mut snapshot.api_key_allowed_providers,
|
||||
);
|
||||
constrain_api_key_list_policy_to_user_policy(
|
||||
&mut allowed_api_formats,
|
||||
&mut snapshot.api_key_allowed_api_formats,
|
||||
);
|
||||
constrain_api_key_list_policy_to_user_policy(
|
||||
&mut allowed_models,
|
||||
&mut snapshot.api_key_allowed_models,
|
||||
);
|
||||
}
|
||||
snapshot.apply_user_policy(
|
||||
allowed_providers,
|
||||
allowed_api_formats,
|
||||
@@ -1735,64 +1747,135 @@ fn resolve_effective_list_policy(
|
||||
&aether_data::repository::users::StoredUserGroup,
|
||||
) -> (&str, Option<Vec<String>>),
|
||||
) -> Option<Vec<String>> {
|
||||
match user_mode {
|
||||
"unrestricted" => None,
|
||||
"specific" => Some(user_values.unwrap_or_default()),
|
||||
let group_policy = groups.iter().fold(None, |effective, group| {
|
||||
let (mode, values) = group_field(group);
|
||||
intersect_list_policies(effective, list_restriction_from_mode(mode, values))
|
||||
});
|
||||
let user_policy = list_restriction_from_mode(user_mode, user_values);
|
||||
intersect_list_policies(group_policy, user_policy)
|
||||
}
|
||||
|
||||
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()),
|
||||
"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(),
|
||||
let group_policy = groups.iter().fold(None, |effective, group| {
|
||||
intersect_rate_limit_policies(
|
||||
effective,
|
||||
rate_limit_restriction_from_mode(&group.rate_limit_mode, group.rate_limit),
|
||||
)
|
||||
});
|
||||
let user_policy = rate_limit_restriction_from_mode(user_mode, user_rate_limit);
|
||||
rate_limit_policy_value(intersect_rate_limit_policies(group_policy, user_policy))
|
||||
}
|
||||
|
||||
#[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 constrain_api_key_list_policy_to_user_policy(
|
||||
user_policy: &mut Option<Vec<String>>,
|
||||
api_key_policy: &mut Option<Vec<String>>,
|
||||
) {
|
||||
let Some(api_key_values) = api_key_policy.as_ref().filter(|values| !values.is_empty()) else {
|
||||
return;
|
||||
};
|
||||
let Some(user_values) = user_policy.clone() else {
|
||||
return;
|
||||
};
|
||||
let effective = intersect_list_policies(Some(api_key_values.to_vec()), Some(user_values))
|
||||
.unwrap_or_default();
|
||||
*user_policy = Some(effective.clone());
|
||||
*api_key_policy = Some(effective);
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use super::*;
|
||||
use aether_data::repository::auth::{
|
||||
InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeyExportRecord,
|
||||
StoredAuthApiKeySnapshot,
|
||||
};
|
||||
use aether_data::repository::users::StoredUserGroup;
|
||||
|
||||
use crate::data::GatewayDataState;
|
||||
|
||||
@@ -1823,6 +1906,125 @@ mod tests {
|
||||
.expect("snapshot should build")
|
||||
}
|
||||
|
||||
fn sample_group(
|
||||
id: &str,
|
||||
priority: i32,
|
||||
allowed_models: Option<Vec<&str>>,
|
||||
allowed_models_mode: &str,
|
||||
rate_limit: Option<i32>,
|
||||
rate_limit_mode: &str,
|
||||
) -> StoredUserGroup {
|
||||
StoredUserGroup {
|
||||
id: id.to_string(),
|
||||
name: id.to_string(),
|
||||
normalized_name: id.to_string(),
|
||||
description: None,
|
||||
priority,
|
||||
allowed_providers: None,
|
||||
allowed_providers_mode: "unrestricted".to_string(),
|
||||
allowed_api_formats: None,
|
||||
allowed_api_formats_mode: "unrestricted".to_string(),
|
||||
allowed_models: allowed_models.map(|values| {
|
||||
values
|
||||
.into_iter()
|
||||
.map(ToOwned::to_owned)
|
||||
.collect::<Vec<_>>()
|
||||
}),
|
||||
allowed_models_mode: allowed_models_mode.to_string(),
|
||||
rate_limit,
|
||||
rate_limit_mode: rate_limit_mode.to_string(),
|
||||
created_at: None,
|
||||
updated_at: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn list_policy_intersects_group_and_user_restrictions() {
|
||||
let groups = vec![
|
||||
sample_group("default", 0, None, "unrestricted", None, "system"),
|
||||
sample_group(
|
||||
"restricted",
|
||||
10,
|
||||
Some(vec!["gpt-5", "gpt-4.1"]),
|
||||
"specific",
|
||||
None,
|
||||
"system",
|
||||
),
|
||||
];
|
||||
|
||||
let policy = resolve_effective_list_policy(
|
||||
Some(vec!["gpt-4.1".to_string(), "gemini-2.5-pro".to_string()]),
|
||||
"specific",
|
||||
&groups,
|
||||
|group| (&group.allowed_models_mode, group.allowed_models.clone()),
|
||||
);
|
||||
|
||||
assert_eq!(policy, Some(vec!["gpt-4.1".to_string()]));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn user_unrestricted_does_not_bypass_group_restrictions() {
|
||||
let groups = vec![sample_group(
|
||||
"restricted",
|
||||
10,
|
||||
Some(vec!["gpt-5"]),
|
||||
"specific",
|
||||
None,
|
||||
"system",
|
||||
)];
|
||||
|
||||
let policy = resolve_effective_list_policy(None, "unrestricted", &groups, |group| {
|
||||
(&group.allowed_models_mode, group.allowed_models.clone())
|
||||
});
|
||||
|
||||
assert_eq!(policy, Some(vec!["gpt-5".to_string()]));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rate_limit_policy_uses_most_restrictive_custom_limit() {
|
||||
let groups = vec![sample_group(
|
||||
"restricted",
|
||||
10,
|
||||
None,
|
||||
"unrestricted",
|
||||
Some(60),
|
||||
"custom",
|
||||
)];
|
||||
|
||||
assert_eq!(
|
||||
resolve_effective_rate_limit_policy(Some(120), "custom", &groups),
|
||||
Some(60)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rate_limit_unlimited_does_not_bypass_limited_group() {
|
||||
let groups = vec![sample_group(
|
||||
"restricted",
|
||||
10,
|
||||
None,
|
||||
"unrestricted",
|
||||
Some(60),
|
||||
"custom",
|
||||
)];
|
||||
|
||||
assert_eq!(
|
||||
resolve_effective_rate_limit_policy(Some(0), "custom", &groups),
|
||||
Some(60)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn api_key_specific_policy_cannot_expand_user_policy() {
|
||||
let mut user_policy = Some(vec!["gpt-5".to_string()]);
|
||||
let mut api_key_policy = Some(vec!["gpt-4.1".to_string()]);
|
||||
|
||||
constrain_api_key_list_policy_to_user_policy(&mut user_policy, &mut api_key_policy);
|
||||
|
||||
assert_eq!(user_policy, Some(Vec::<String>::new()));
|
||||
assert_eq!(api_key_policy, Some(Vec::<String>::new()));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn data_state_lists_auth_api_key_export_records() {
|
||||
let repository = Arc::new(
|
||||
|
||||
@@ -172,6 +172,10 @@ impl<'a> AdminAppState<'a> {
|
||||
let user_api_keys = self
|
||||
.list_auth_api_key_export_records_by_user_ids(&user_ids)
|
||||
.await?;
|
||||
let groups = self.list_user_groups().await?;
|
||||
let memberships = self
|
||||
.list_user_group_memberships_by_user_ids(&user_ids)
|
||||
.await?;
|
||||
let standalone_api_keys = self.list_auth_api_key_export_standalone_records().await?;
|
||||
let standalone_api_key_ids = standalone_api_keys
|
||||
.iter()
|
||||
@@ -205,12 +209,49 @@ impl<'a> AdminAppState<'a> {
|
||||
.or_default()
|
||||
.push(key);
|
||||
}
|
||||
let mut memberships_by_user_id = BTreeMap::<
|
||||
String,
|
||||
Vec<aether_data::repository::users::StoredUserGroupMembership>,
|
||||
>::new();
|
||||
for membership in memberships {
|
||||
memberships_by_user_id
|
||||
.entry(membership.user_id.clone())
|
||||
.or_default()
|
||||
.push(membership);
|
||||
}
|
||||
let user_groups_data = groups
|
||||
.iter()
|
||||
.map(|group| {
|
||||
json!({
|
||||
"id": group.id.clone(),
|
||||
"name": group.name.clone(),
|
||||
"description": group.description.clone(),
|
||||
"allowed_providers": group.allowed_providers.clone(),
|
||||
"allowed_providers_mode": group.allowed_providers_mode.clone(),
|
||||
"allowed_api_formats": group.allowed_api_formats.clone(),
|
||||
"allowed_api_formats_mode": group.allowed_api_formats_mode.clone(),
|
||||
"allowed_models": group.allowed_models.clone(),
|
||||
"allowed_models_mode": group.allowed_models_mode.clone(),
|
||||
"rate_limit": group.rate_limit,
|
||||
"rate_limit_mode": group.rate_limit_mode.clone(),
|
||||
})
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
let users_data = users
|
||||
.iter()
|
||||
.map(|user| {
|
||||
let wallet = wallets_by_user_id.get(&user.id);
|
||||
let wallet_payload = serialize_admin_system_users_export_wallet(wallet);
|
||||
let memberships = memberships_by_user_id.remove(&user.id).unwrap_or_default();
|
||||
let group_ids = memberships
|
||||
.iter()
|
||||
.map(|membership| membership.group_id.clone())
|
||||
.collect::<Vec<_>>();
|
||||
let group_names = memberships
|
||||
.iter()
|
||||
.map(|membership| membership.group_name.clone())
|
||||
.collect::<Vec<_>>();
|
||||
let api_keys = api_keys_by_user_id.remove(&user.id).unwrap_or_default();
|
||||
let api_keys_payload = api_keys
|
||||
.iter()
|
||||
@@ -226,10 +267,16 @@ impl<'a> AdminAppState<'a> {
|
||||
"password_hash": user.password_hash.clone(),
|
||||
"role": user.role.clone(),
|
||||
"allowed_providers": user.allowed_providers.clone(),
|
||||
"allowed_providers_mode": user.allowed_providers_mode.clone(),
|
||||
"allowed_api_formats": user.allowed_api_formats.clone(),
|
||||
"allowed_api_formats_mode": user.allowed_api_formats_mode.clone(),
|
||||
"allowed_models": user.allowed_models.clone(),
|
||||
"allowed_models_mode": user.allowed_models_mode.clone(),
|
||||
"rate_limit": user.rate_limit,
|
||||
"rate_limit_mode": user.rate_limit_mode.clone(),
|
||||
"model_capability_settings": user.model_capability_settings.clone(),
|
||||
"group_ids": group_ids,
|
||||
"group_names": group_names,
|
||||
"unlimited": wallet
|
||||
.map(|entry| entry.limit_mode.eq_ignore_ascii_case("unlimited"))
|
||||
.unwrap_or(false),
|
||||
@@ -254,6 +301,7 @@ impl<'a> AdminAppState<'a> {
|
||||
Ok(json!({
|
||||
"version": ADMIN_SYSTEM_USERS_EXPORT_VERSION,
|
||||
"exported_at": Utc::now().to_rfc3339(),
|
||||
"user_groups": user_groups_data,
|
||||
"users": users_data,
|
||||
"standalone_keys": standalone_keys_data,
|
||||
}))
|
||||
|
||||
@@ -10,7 +10,9 @@ use crate::handlers::admin::shared::{
|
||||
};
|
||||
use crate::handlers::admin::system::shared::configs::apply_admin_system_config_update;
|
||||
use crate::handlers::admin::users::{
|
||||
hash_admin_user_api_key, normalize_admin_user_api_formats, normalize_admin_user_string_list,
|
||||
hash_admin_user_api_key, normalize_admin_list_policy_mode,
|
||||
normalize_admin_rate_limit_policy_mode, normalize_admin_user_api_formats,
|
||||
normalize_admin_user_string_list,
|
||||
};
|
||||
use crate::handlers::public::normalize_admin_base_url;
|
||||
use crate::GatewayError;
|
||||
@@ -397,6 +399,7 @@ fn build_import_provider_model_record(
|
||||
|
||||
#[derive(Debug, Clone, Default, serde::Serialize)]
|
||||
struct AdminSystemUsersImportStats {
|
||||
user_groups: AdminSystemConfigImportCounter,
|
||||
users: AdminSystemConfigImportCounter,
|
||||
api_keys: AdminSystemConfigImportCounter,
|
||||
standalone_keys: AdminSystemConfigImportCounter,
|
||||
@@ -542,6 +545,78 @@ fn imported_optional_value(value: Option<&Value>) -> Option<Value> {
|
||||
value.cloned().filter(|value| !value.is_null())
|
||||
}
|
||||
|
||||
fn imported_optional_list_policy_mode(
|
||||
value: Option<&Value>,
|
||||
field_name: &str,
|
||||
) -> Result<Option<String>, String> {
|
||||
let Some(value) = imported_optional_string(value)? else {
|
||||
return Ok(None);
|
||||
};
|
||||
let value = value.to_ascii_lowercase();
|
||||
normalize_admin_list_policy_mode(&value)
|
||||
.map(Some)
|
||||
.map_err(|_| format!("{field_name} 不合法"))
|
||||
}
|
||||
|
||||
fn imported_optional_rate_limit_policy_mode(
|
||||
value: Option<&Value>,
|
||||
field_name: &str,
|
||||
) -> Result<Option<String>, String> {
|
||||
let Some(value) = imported_optional_string(value)? else {
|
||||
return Ok(None);
|
||||
};
|
||||
let value = value.to_ascii_lowercase();
|
||||
normalize_admin_rate_limit_policy_mode(&value)
|
||||
.map(Some)
|
||||
.map_err(|_| format!("{field_name} 不合法"))
|
||||
}
|
||||
|
||||
fn legacy_imported_list_policy_mode(values: &Option<Vec<String>>) -> String {
|
||||
if values.is_some() {
|
||||
"specific".to_string()
|
||||
} else {
|
||||
"unrestricted".to_string()
|
||||
}
|
||||
}
|
||||
|
||||
fn legacy_imported_rate_limit_policy_mode(value: Option<i32>) -> String {
|
||||
if value.is_some() {
|
||||
"custom".to_string()
|
||||
} else {
|
||||
"system".to_string()
|
||||
}
|
||||
}
|
||||
|
||||
fn imported_user_list_policy_mode(
|
||||
object: &Map<String, Value>,
|
||||
mode_field: &str,
|
||||
value_field: &str,
|
||||
values: &Option<Vec<String>>,
|
||||
) -> Result<Option<String>, String> {
|
||||
imported_optional_list_policy_mode(object.get(mode_field), mode_field).map(|mode| {
|
||||
mode.or_else(|| {
|
||||
object
|
||||
.contains_key(value_field)
|
||||
.then(|| legacy_imported_list_policy_mode(values))
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn imported_user_rate_limit_policy_mode(
|
||||
object: &Map<String, Value>,
|
||||
mode_field: &str,
|
||||
value_field: &str,
|
||||
value: Option<i32>,
|
||||
) -> Result<Option<String>, String> {
|
||||
imported_optional_rate_limit_policy_mode(object.get(mode_field), mode_field).map(|mode| {
|
||||
mode.or_else(|| {
|
||||
object
|
||||
.contains_key(value_field)
|
||||
.then(|| legacy_imported_rate_limit_policy_mode(value))
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn imported_rfc3339_to_unix_secs(
|
||||
value: Option<&Value>,
|
||||
field_name: &str,
|
||||
@@ -601,6 +676,130 @@ fn normalize_imported_user_api_formats(
|
||||
)?)
|
||||
}
|
||||
|
||||
fn build_imported_user_group_record(
|
||||
group: &Map<String, Value>,
|
||||
field_name: &str,
|
||||
) -> Result<
|
||||
(
|
||||
Option<String>,
|
||||
String,
|
||||
aether_data::repository::users::UpsertUserGroupRecord,
|
||||
),
|
||||
String,
|
||||
> {
|
||||
let export_id = imported_optional_string(group.get("id"))?;
|
||||
let name = imported_optional_string(group.get("name"))?
|
||||
.ok_or_else(|| format!("{field_name}.name 不能为空"))?;
|
||||
let name = aether_data::repository::users::normalize_user_group_name(&name);
|
||||
if name.is_empty() {
|
||||
return Err(format!("{field_name}.name 不能为空"));
|
||||
}
|
||||
let description = imported_optional_string(group.get("description"))?;
|
||||
let allowed_providers = normalize_imported_user_string_list(group, "allowed_providers")?;
|
||||
let allowed_api_formats = normalize_imported_user_api_formats(group, "allowed_api_formats")?;
|
||||
let allowed_models = normalize_imported_user_string_list(group, "allowed_models")?;
|
||||
let rate_limit = imported_optional_i32(group.get("rate_limit"), "rate_limit")?;
|
||||
|
||||
let allowed_providers_mode = imported_optional_list_policy_mode(
|
||||
group.get("allowed_providers_mode"),
|
||||
"allowed_providers_mode",
|
||||
)?
|
||||
.unwrap_or_else(|| {
|
||||
if group.contains_key("allowed_providers") {
|
||||
legacy_imported_list_policy_mode(&allowed_providers)
|
||||
} else {
|
||||
"inherit".to_string()
|
||||
}
|
||||
});
|
||||
let allowed_api_formats_mode = imported_optional_list_policy_mode(
|
||||
group.get("allowed_api_formats_mode"),
|
||||
"allowed_api_formats_mode",
|
||||
)?
|
||||
.unwrap_or_else(|| {
|
||||
if group.contains_key("allowed_api_formats") {
|
||||
legacy_imported_list_policy_mode(&allowed_api_formats)
|
||||
} else {
|
||||
"inherit".to_string()
|
||||
}
|
||||
});
|
||||
let allowed_models_mode = imported_optional_list_policy_mode(
|
||||
group.get("allowed_models_mode"),
|
||||
"allowed_models_mode",
|
||||
)?
|
||||
.unwrap_or_else(|| {
|
||||
if group.contains_key("allowed_models") {
|
||||
legacy_imported_list_policy_mode(&allowed_models)
|
||||
} else {
|
||||
"inherit".to_string()
|
||||
}
|
||||
});
|
||||
let rate_limit_mode =
|
||||
imported_optional_rate_limit_policy_mode(group.get("rate_limit_mode"), "rate_limit_mode")?
|
||||
.unwrap_or_else(|| {
|
||||
if group.contains_key("rate_limit") {
|
||||
legacy_imported_rate_limit_policy_mode(rate_limit)
|
||||
} else {
|
||||
"inherit".to_string()
|
||||
}
|
||||
});
|
||||
|
||||
let normalized_name = name.to_ascii_lowercase();
|
||||
|
||||
Ok((
|
||||
export_id,
|
||||
normalized_name,
|
||||
aether_data::repository::users::UpsertUserGroupRecord {
|
||||
name,
|
||||
description,
|
||||
priority: 0,
|
||||
allowed_providers,
|
||||
allowed_providers_mode,
|
||||
allowed_api_formats,
|
||||
allowed_api_formats_mode,
|
||||
allowed_models,
|
||||
allowed_models_mode,
|
||||
rate_limit,
|
||||
rate_limit_mode,
|
||||
},
|
||||
))
|
||||
}
|
||||
|
||||
fn resolve_imported_user_group_ids(
|
||||
user: &Map<String, Value>,
|
||||
imported_group_id_map: &BTreeMap<String, String>,
|
||||
imported_group_name_map: &BTreeMap<String, String>,
|
||||
groups_by_name: &BTreeMap<String, aether_data::repository::users::StoredUserGroup>,
|
||||
) -> Result<Vec<String>, String> {
|
||||
let raw_group_ids =
|
||||
imported_string_list_from_value(user.get("group_ids"), "group_ids")?.unwrap_or_default();
|
||||
let raw_group_names = imported_string_list_from_value(user.get("group_names"), "group_names")?
|
||||
.unwrap_or_default();
|
||||
let mut group_ids = BTreeSet::new();
|
||||
for raw_group_id in raw_group_ids {
|
||||
if let Some(group_id) = imported_group_id_map.get(&raw_group_id) {
|
||||
group_ids.insert(group_id.clone());
|
||||
continue;
|
||||
}
|
||||
group_ids.insert(raw_group_id);
|
||||
}
|
||||
for raw_group_name in raw_group_names {
|
||||
let normalized_name =
|
||||
aether_data::repository::users::normalize_user_group_name(&raw_group_name)
|
||||
.to_ascii_lowercase();
|
||||
if normalized_name.is_empty() {
|
||||
continue;
|
||||
}
|
||||
if let Some(group_id) = imported_group_name_map.get(&normalized_name) {
|
||||
group_ids.insert(group_id.clone());
|
||||
continue;
|
||||
}
|
||||
if let Some(group) = groups_by_name.get(&normalized_name) {
|
||||
group_ids.insert(group.id.clone());
|
||||
}
|
||||
}
|
||||
Ok(group_ids.into_iter().collect())
|
||||
}
|
||||
|
||||
fn normalize_imported_wallet_target(
|
||||
wallet: Option<&Map<String, Value>>,
|
||||
unlimited: bool,
|
||||
@@ -1686,6 +1885,11 @@ impl<'a> AdminAppState<'a> {
|
||||
Some(_) => return Ok(Err(invalid_request("standalone_keys 必须是数组"))),
|
||||
None => &empty,
|
||||
};
|
||||
let imported_user_groups = match root.get("user_groups") {
|
||||
Some(Value::Array(items)) => items,
|
||||
Some(_) => return Ok(Err(invalid_request("user_groups 必须是数组"))),
|
||||
None => &empty,
|
||||
};
|
||||
|
||||
let standalone_owner_id = match operator_id {
|
||||
Some(candidate) => match self.find_user_auth_by_id(candidate).await? {
|
||||
@@ -1709,6 +1913,86 @@ impl<'a> AdminAppState<'a> {
|
||||
));
|
||||
|
||||
let mut stats = AdminSystemUsersImportStats::default();
|
||||
let default_group_id = self.effective_default_user_group_id().await?;
|
||||
let existing_groups = self.list_user_groups().await?;
|
||||
let mut groups_by_name = existing_groups
|
||||
.into_iter()
|
||||
.map(|group| {
|
||||
(
|
||||
aether_data::repository::users::normalize_user_group_name(&group.name)
|
||||
.to_ascii_lowercase(),
|
||||
group,
|
||||
)
|
||||
})
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
let mut imported_group_id_map = BTreeMap::<String, String>::new();
|
||||
let mut imported_group_name_map = BTreeMap::<String, String>::new();
|
||||
|
||||
for (index, raw_group) in imported_user_groups.iter().enumerate() {
|
||||
let group = match imported_object_field(raw_group, &format!("user_groups[{index}]")) {
|
||||
Ok(value) => value,
|
||||
Err(detail) => return Ok(Err(invalid_request(detail))),
|
||||
};
|
||||
let (export_id, normalized_name, record) = invalid_value!(
|
||||
build_imported_user_group_record(group, &format!("user_groups[{index}]"))
|
||||
);
|
||||
if default_group_id
|
||||
.as_deref()
|
||||
.is_some_and(|group_id| export_id.as_deref() == Some(group_id))
|
||||
|| normalized_name == "default"
|
||||
{
|
||||
if let Some(default_group_id) = default_group_id.as_ref() {
|
||||
if let Some(export_id) = export_id {
|
||||
imported_group_id_map.insert(export_id, default_group_id.clone());
|
||||
}
|
||||
imported_group_name_map.insert(normalized_name, default_group_id.clone());
|
||||
}
|
||||
stats.user_groups.skipped += 1;
|
||||
continue;
|
||||
}
|
||||
if let Some(existing) = groups_by_name.get(&normalized_name).cloned() {
|
||||
if let Some(export_id) = export_id {
|
||||
imported_group_id_map.insert(export_id, existing.id.clone());
|
||||
}
|
||||
imported_group_name_map.insert(normalized_name.clone(), existing.id.clone());
|
||||
match merge_mode {
|
||||
AdminImportMergeMode::Skip => {
|
||||
stats.user_groups.skipped += 1;
|
||||
}
|
||||
AdminImportMergeMode::Error => {
|
||||
return Ok(Err(invalid_request(format!(
|
||||
"用户组 '{}' 已存在",
|
||||
existing.name
|
||||
))));
|
||||
}
|
||||
AdminImportMergeMode::Overwrite => {
|
||||
let Some(updated) = self.update_user_group(&existing.id, record).await?
|
||||
else {
|
||||
return Ok(Err((
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
json!({ "detail": "Admin system data unavailable" }),
|
||||
)));
|
||||
};
|
||||
groups_by_name.insert(normalized_name, updated);
|
||||
stats.user_groups.updated += 1;
|
||||
}
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
let Some(created) = self.create_user_group(record).await? else {
|
||||
return Ok(Err((
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
json!({ "detail": "Admin system data unavailable" }),
|
||||
)));
|
||||
};
|
||||
if let Some(export_id) = export_id {
|
||||
imported_group_id_map.insert(export_id, created.id.clone());
|
||||
}
|
||||
imported_group_name_map.insert(normalized_name.clone(), created.id.clone());
|
||||
groups_by_name.insert(normalized_name, created);
|
||||
stats.user_groups.created += 1;
|
||||
}
|
||||
|
||||
for (index, raw_user) in users.iter().enumerate() {
|
||||
let user = match imported_object_field(raw_user, &format!("users[{index}]")) {
|
||||
@@ -1760,6 +2044,53 @@ impl<'a> AdminAppState<'a> {
|
||||
invalid_value!(normalize_imported_user_string_list(user, "allowed_models"));
|
||||
let rate_limit =
|
||||
invalid_value!(imported_optional_i32(user.get("rate_limit"), "rate_limit"));
|
||||
let allowed_providers_mode = invalid_value!(imported_user_list_policy_mode(
|
||||
user,
|
||||
"allowed_providers_mode",
|
||||
"allowed_providers",
|
||||
&allowed_providers,
|
||||
));
|
||||
let allowed_api_formats_mode = invalid_value!(imported_user_list_policy_mode(
|
||||
user,
|
||||
"allowed_api_formats_mode",
|
||||
"allowed_api_formats",
|
||||
&allowed_api_formats,
|
||||
));
|
||||
let allowed_models_mode = invalid_value!(imported_user_list_policy_mode(
|
||||
user,
|
||||
"allowed_models_mode",
|
||||
"allowed_models",
|
||||
&allowed_models,
|
||||
));
|
||||
let rate_limit_mode = invalid_value!(imported_user_rate_limit_policy_mode(
|
||||
user,
|
||||
"rate_limit_mode",
|
||||
"rate_limit",
|
||||
rate_limit,
|
||||
));
|
||||
let imported_user_group_ids = invalid_value!(resolve_imported_user_group_ids(
|
||||
user,
|
||||
&imported_group_id_map,
|
||||
&imported_group_name_map,
|
||||
&groups_by_name,
|
||||
));
|
||||
let group_ids = if user.contains_key("group_ids") || user.contains_key("group_names") {
|
||||
let group_ids = self
|
||||
.include_default_user_group_ids(&imported_user_group_ids)
|
||||
.await?;
|
||||
if !group_ids.is_empty() {
|
||||
let existing_groups = self.list_user_groups_by_ids(&group_ids).await?;
|
||||
if existing_groups.len() != group_ids.len() {
|
||||
return Ok(Err(invalid_request(format!(
|
||||
"用户 '{}' 的用户组不存在",
|
||||
email.clone().unwrap_or(username.clone())
|
||||
))));
|
||||
}
|
||||
}
|
||||
Some(group_ids)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let is_active =
|
||||
invalid_value!(imported_optional_bool(user.get("is_active"))).unwrap_or(true);
|
||||
let model_capability_settings = invalid_value!(imported_optional_json_object(
|
||||
@@ -1883,6 +2214,31 @@ impl<'a> AdminAppState<'a> {
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
if allowed_providers_mode.is_some()
|
||||
|| allowed_api_formats_mode.is_some()
|
||||
|| allowed_models_mode.is_some()
|
||||
|| rate_limit_mode.is_some()
|
||||
{
|
||||
let updated_policy_modes = self
|
||||
.update_local_auth_user_policy_modes(
|
||||
&existing.id,
|
||||
allowed_providers_mode.clone(),
|
||||
allowed_api_formats_mode.clone(),
|
||||
allowed_models_mode.clone(),
|
||||
rate_limit_mode.clone(),
|
||||
)
|
||||
.await?;
|
||||
if updated_policy_modes.is_none() {
|
||||
return Ok(Err((
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
json!({ "detail": "Admin system data unavailable" }),
|
||||
)));
|
||||
}
|
||||
}
|
||||
if let Some(group_ids) = group_ids.as_ref() {
|
||||
self.replace_user_groups_for_user(&existing.id, group_ids)
|
||||
.await?;
|
||||
}
|
||||
self.sync_imported_user_wallet(
|
||||
&existing.id,
|
||||
&wallet_target,
|
||||
@@ -1921,6 +2277,34 @@ impl<'a> AdminAppState<'a> {
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
let created = if allowed_providers_mode.is_some()
|
||||
|| allowed_api_formats_mode.is_some()
|
||||
|| allowed_models_mode.is_some()
|
||||
|| rate_limit_mode.is_some()
|
||||
{
|
||||
let Some(updated_policy_modes) = self
|
||||
.update_local_auth_user_policy_modes(
|
||||
&created.id,
|
||||
allowed_providers_mode.clone(),
|
||||
allowed_api_formats_mode.clone(),
|
||||
allowed_models_mode.clone(),
|
||||
rate_limit_mode.clone(),
|
||||
)
|
||||
.await?
|
||||
else {
|
||||
return Ok(Err((
|
||||
http::StatusCode::SERVICE_UNAVAILABLE,
|
||||
json!({ "detail": "Admin system data unavailable" }),
|
||||
)));
|
||||
};
|
||||
updated_policy_modes
|
||||
} else {
|
||||
created
|
||||
};
|
||||
if let Some(group_ids) = group_ids.as_ref() {
|
||||
self.replace_user_groups_for_user(&created.id, group_ids)
|
||||
.await?;
|
||||
}
|
||||
self.sync_imported_user_wallet(
|
||||
&created.id,
|
||||
&wallet_target,
|
||||
@@ -2454,9 +2838,10 @@ mod tests {
|
||||
#[test]
|
||||
fn users_import_requires_supported_export_version() {
|
||||
assert!(validate_imported_system_users_export_version(Some(&json!("1.3"))).is_ok());
|
||||
assert!(validate_imported_system_users_export_version(Some(&json!("1.4"))).is_ok());
|
||||
assert_eq!(
|
||||
validate_imported_system_users_export_version(Some(&json!("2.2"))).unwrap_err(),
|
||||
"不支持的用户数据版本: 2.2,支持的版本: 1.3"
|
||||
"不支持的用户数据版本: 2.2,支持的版本: 1.3, 1.4"
|
||||
);
|
||||
assert_eq!(
|
||||
validate_imported_system_users_export_version(Some(&json!(null))).unwrap_err(),
|
||||
|
||||
@@ -145,6 +145,31 @@ impl<'a> AdminAppState<'a> {
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn include_default_user_group_ids(
|
||||
&self,
|
||||
group_ids: &[String],
|
||||
) -> Result<Vec<String>, GatewayError> {
|
||||
self.app.include_default_user_group_ids(group_ids).await
|
||||
}
|
||||
|
||||
pub(crate) async fn effective_default_user_group_id(
|
||||
&self,
|
||||
) -> Result<Option<String>, GatewayError> {
|
||||
self.app.effective_default_user_group_id().await
|
||||
}
|
||||
|
||||
pub(crate) async fn add_user_to_group(
|
||||
&self,
|
||||
group_id: &str,
|
||||
user_id: &str,
|
||||
) -> Result<bool, GatewayError> {
|
||||
self.app.add_user_to_group(group_id, user_id).await
|
||||
}
|
||||
|
||||
pub(crate) async fn add_all_users_to_group(&self, group_id: &str) -> Result<(), GatewayError> {
|
||||
self.app.add_all_users_to_group(group_id).await
|
||||
}
|
||||
|
||||
pub(crate) async fn is_other_user_auth_email_taken(
|
||||
&self,
|
||||
email: &str,
|
||||
|
||||
@@ -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()),
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
|
||||
use crate::constants::{BUILTIN_DEFAULT_USER_GROUP_ID, DEFAULT_USER_GROUP_CONFIG_KEY};
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
impl AppState {
|
||||
@@ -7,20 +8,10 @@ impl AppState {
|
||||
&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 group_id = self.effective_default_user_group_id().await?;
|
||||
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}"
|
||||
@@ -29,6 +20,61 @@ impl AppState {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) async fn configured_default_user_group_id(
|
||||
&self,
|
||||
) -> Result<Option<String>, GatewayError> {
|
||||
Ok(self
|
||||
.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()))
|
||||
}
|
||||
|
||||
pub(crate) async fn effective_default_user_group_id(
|
||||
&self,
|
||||
) -> Result<Option<String>, GatewayError> {
|
||||
if let Some(group_id) = self.configured_default_user_group_id().await? {
|
||||
if self.find_user_group_by_id(&group_id).await?.is_none() {
|
||||
return Err(GatewayError::Internal(format!(
|
||||
"{DEFAULT_USER_GROUP_CONFIG_KEY} points to missing group: {group_id}"
|
||||
)));
|
||||
}
|
||||
return Ok(Some(group_id));
|
||||
}
|
||||
if self
|
||||
.find_user_group_by_id(BUILTIN_DEFAULT_USER_GROUP_ID)
|
||||
.await?
|
||||
.is_some()
|
||||
{
|
||||
return Ok(Some(BUILTIN_DEFAULT_USER_GROUP_ID.to_string()));
|
||||
}
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
pub(crate) async fn include_default_user_group_ids(
|
||||
&self,
|
||||
group_ids: &[String],
|
||||
) -> Result<Vec<String>, GatewayError> {
|
||||
let mut group_ids = group_ids
|
||||
.iter()
|
||||
.map(|value| value.trim())
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.collect::<BTreeSet<_>>();
|
||||
if let Some(default_group_id) = self.effective_default_user_group_id().await? {
|
||||
group_ids.insert(default_group_id);
|
||||
}
|
||||
Ok(group_ids.into_iter().collect())
|
||||
}
|
||||
|
||||
pub(crate) async fn add_all_users_to_group(&self, group_id: &str) -> Result<(), GatewayError> {
|
||||
for user in self.list_export_users().await? {
|
||||
self.add_user_to_group(group_id, &user.id).await?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub(crate) async fn resolve_auth_user_summaries_by_ids(
|
||||
&self,
|
||||
user_ids: &[String],
|
||||
|
||||
@@ -13,7 +13,9 @@ use aether_data::repository::global_models::InMemoryGlobalModelReadRepository;
|
||||
use aether_data::repository::oauth_providers::InMemoryOAuthProviderRepository;
|
||||
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
|
||||
use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository;
|
||||
use aether_data::repository::users::{InMemoryUserReadRepository, StoredUserExportRow};
|
||||
use aether_data::repository::users::{
|
||||
InMemoryUserReadRepository, StoredUserAuthRecord, UpsertUserGroupRecord, UserReadRepository,
|
||||
};
|
||||
use aether_data::repository::wallet::{InMemoryWalletRepository, StoredWalletSnapshot};
|
||||
use aether_data_contracts::repository::global_models::StoredPublicGlobalModel;
|
||||
use axum::body::Body;
|
||||
@@ -596,8 +598,8 @@ async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_pr
|
||||
}),
|
||||
);
|
||||
|
||||
let user_repository = Arc::new(InMemoryUserReadRepository::seed_export_users(vec![
|
||||
StoredUserExportRow::new(
|
||||
let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![
|
||||
StoredUserAuthRecord::new(
|
||||
"user-1".to_string(),
|
||||
Some("alice@example.com".to_string()),
|
||||
true,
|
||||
@@ -608,12 +610,40 @@ async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_pr
|
||||
Some(json!(["openai"])),
|
||||
Some(json!(["openai:chat"])),
|
||||
Some(json!(["gpt-5"])),
|
||||
Some(120),
|
||||
Some(json!({"gpt-5": {"cache_1h": true}})),
|
||||
true,
|
||||
false,
|
||||
Some(chrono::Utc::now()),
|
||||
None,
|
||||
)
|
||||
.expect("user export row should build"),
|
||||
.expect("user export row should build")
|
||||
.with_policy_modes(
|
||||
"specific".to_string(),
|
||||
"specific".to_string(),
|
||||
"specific".to_string(),
|
||||
)
|
||||
.expect("user policy modes should build"),
|
||||
]));
|
||||
let user_group = user_repository
|
||||
.create_user_group(UpsertUserGroupRecord {
|
||||
name: "Restricted GPT".to_string(),
|
||||
description: Some("GPT-only users".to_string()),
|
||||
priority: 10,
|
||||
allowed_providers: Some(vec!["openai".to_string()]),
|
||||
allowed_providers_mode: "specific".to_string(),
|
||||
allowed_api_formats: Some(vec!["openai:chat".to_string()]),
|
||||
allowed_api_formats_mode: "specific".to_string(),
|
||||
allowed_models: Some(vec!["gpt-5".to_string()]),
|
||||
allowed_models_mode: "specific".to_string(),
|
||||
rate_limit: Some(60),
|
||||
rate_limit_mode: "custom".to_string(),
|
||||
})
|
||||
.await
|
||||
.expect("user group should create")
|
||||
.expect("user group should exist");
|
||||
user_repository
|
||||
.replace_user_groups_for_user("user-1", std::slice::from_ref(&user_group.id))
|
||||
.await
|
||||
.expect("user group membership should create");
|
||||
let auth_repository = Arc::new(
|
||||
InMemoryAuthApiKeySnapshotRepository::default().with_export_records(vec![
|
||||
StoredAuthApiKeyExportRecord::new(
|
||||
@@ -728,9 +758,24 @@ async fn gateway_handles_admin_system_users_export_locally_with_trusted_admin_pr
|
||||
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
let payload: serde_json::Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(payload["version"], "1.3");
|
||||
assert_eq!(payload["version"], "1.4");
|
||||
assert!(payload["exported_at"].as_str().is_some());
|
||||
assert_eq!(payload["user_groups"][0]["name"], "Restricted GPT");
|
||||
assert!(payload["user_groups"][0].get("priority").is_none());
|
||||
assert_eq!(
|
||||
payload["user_groups"][0]["allowed_models"],
|
||||
json!(["gpt-5"])
|
||||
);
|
||||
assert_eq!(payload["users"][0]["email"], "alice@example.com");
|
||||
assert_eq!(
|
||||
payload["users"][0]["allowed_models_mode"],
|
||||
json!("specific")
|
||||
);
|
||||
assert_eq!(payload["users"][0]["rate_limit_mode"], json!("system"));
|
||||
assert_eq!(
|
||||
payload["users"][0]["group_names"],
|
||||
json!(["Restricted GPT"])
|
||||
);
|
||||
assert_eq!(payload["users"][0]["wallet"]["balance"], json!(12.5));
|
||||
assert_eq!(
|
||||
payload["users"][0]["wallet"]["recharge_balance"],
|
||||
|
||||
@@ -12,7 +12,7 @@ use aether_data::repository::oauth_providers::{
|
||||
InMemoryOAuthProviderRepository, OAuthProviderReadRepository, StoredOAuthProviderConfig,
|
||||
};
|
||||
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
|
||||
use aether_data::repository::users::StoredUserAuthRecord;
|
||||
use aether_data::repository::users::{StoredUserAuthRecord, UserReadRepository};
|
||||
use aether_data::repository::wallet::{StoredWalletSnapshot, WalletLookupKey};
|
||||
use aether_data_contracts::repository::global_models::{
|
||||
AdminGlobalModelListQuery, AdminProviderModelListQuery, GlobalModelReadRepository,
|
||||
@@ -556,11 +556,14 @@ async fn gateway_imports_admin_system_users_locally_and_persists_data() {
|
||||
}));
|
||||
|
||||
let auth_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default());
|
||||
let user_repository =
|
||||
Arc::new(aether_data::repository::users::InMemoryUserReadRepository::default());
|
||||
let (upstream_url, upstream_handle) = start_server(upstream).await;
|
||||
let state = AppState::new()
|
||||
.expect("gateway should build")
|
||||
.with_data_state_for_tests(
|
||||
GatewayDataState::with_auth_api_key_repository_for_tests(Arc::clone(&auth_repository))
|
||||
.with_user_reader(user_repository)
|
||||
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
|
||||
)
|
||||
.with_auth_users_for_tests([sample_import_admin_user("admin-user-123")])
|
||||
@@ -575,8 +578,21 @@ async fn gateway_imports_admin_system_users_locally_and_persists_data() {
|
||||
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
|
||||
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
|
||||
.json(&json!({
|
||||
"version": "1.3",
|
||||
"version": "1.4",
|
||||
"merge_mode": "overwrite",
|
||||
"user_groups": [{
|
||||
"id": "source-group-1",
|
||||
"name": "GPT Import",
|
||||
"description": "Imported group",
|
||||
"allowed_providers": ["openai"],
|
||||
"allowed_providers_mode": "specific",
|
||||
"allowed_api_formats": ["openai:chat"],
|
||||
"allowed_api_formats_mode": "specific",
|
||||
"allowed_models": ["gpt-5"],
|
||||
"allowed_models_mode": "specific",
|
||||
"rate_limit": 44,
|
||||
"rate_limit_mode": "custom"
|
||||
}],
|
||||
"users": [{
|
||||
"email": "alice@example.com",
|
||||
"email_verified": true,
|
||||
@@ -587,6 +603,10 @@ async fn gateway_imports_admin_system_users_locally_and_persists_data() {
|
||||
"allowed_api_formats": ["openai:chat"],
|
||||
"allowed_models": ["gpt-5"],
|
||||
"rate_limit": 77,
|
||||
"allowed_models_mode": "specific",
|
||||
"rate_limit_mode": "custom",
|
||||
"group_ids": ["source-group-1"],
|
||||
"group_names": ["GPT Import"],
|
||||
"is_active": true,
|
||||
"wallet": {
|
||||
"balance": 20.0,
|
||||
@@ -654,6 +674,7 @@ async fn gateway_imports_admin_system_users_locally_and_persists_data() {
|
||||
let payload: Value = response.json().await.expect("json body should parse");
|
||||
assert_eq!(status, StatusCode::OK, "payload={payload}");
|
||||
assert_eq!(payload["message"], "用户数据导入成功");
|
||||
assert_eq!(payload["stats"]["user_groups"]["created"], json!(1));
|
||||
assert_eq!(payload["stats"]["users"]["created"], json!(1));
|
||||
assert_eq!(payload["stats"]["api_keys"]["created"], json!(1));
|
||||
assert_eq!(payload["stats"]["standalone_keys"]["created"], json!(1));
|
||||
@@ -683,8 +704,22 @@ async fn gateway_imports_admin_system_users_locally_and_persists_data() {
|
||||
imported_user.allowed_models,
|
||||
Some(vec!["gpt-5".to_string()])
|
||||
);
|
||||
assert_eq!(imported_user.allowed_models_mode, "specific");
|
||||
assert!(imported_user.is_active);
|
||||
|
||||
let imported_groups = state
|
||||
.list_user_groups_for_user(&imported_user.id)
|
||||
.await
|
||||
.expect("user groups should load");
|
||||
assert_eq!(imported_groups.len(), 1);
|
||||
assert_eq!(imported_groups[0].name, "GPT Import");
|
||||
assert_eq!(imported_groups[0].allowed_models_mode, "specific");
|
||||
assert_eq!(
|
||||
imported_groups[0].allowed_models,
|
||||
Some(vec!["gpt-5".to_string()])
|
||||
);
|
||||
assert_eq!(imported_groups[0].rate_limit, Some(44));
|
||||
|
||||
let user_wallet = state
|
||||
.find_wallet(WalletLookupKey::UserId(&imported_user.id))
|
||||
.await
|
||||
|
||||
Reference in New Issue
Block a user