mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-11 11:49:50 +08:00
feat: support api key ip restriction rules
This commit is contained in:
@@ -12,7 +12,7 @@ use serde_json::json;
|
||||
use crate::handlers::shared::{
|
||||
api_key_placeholder_display, deserialize_optional_json_patch,
|
||||
deserialize_optional_string_list_patch, generate_gateway_api_key_plaintext,
|
||||
masked_gateway_api_key_display, normalize_feature_settings,
|
||||
masked_gateway_api_key_display, normalize_feature_settings, normalize_ip_rules,
|
||||
normalize_optional_api_key_concurrent_limit,
|
||||
};
|
||||
|
||||
@@ -35,8 +35,8 @@ struct UsersMeCreateApiKeyRequest {
|
||||
concurrent_limit: Option<i32>,
|
||||
#[serde(default)]
|
||||
feature_settings: Option<serde_json::Value>,
|
||||
#[serde(default)]
|
||||
allowed_ips: Option<Vec<String>>,
|
||||
#[serde(default, alias = "allowed_ips")]
|
||||
ip_rules: Option<Vec<String>>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
@@ -49,8 +49,12 @@ struct UsersMeUpdateApiKeyRequest {
|
||||
concurrent_limit: Option<i32>,
|
||||
#[serde(default, deserialize_with = "deserialize_optional_json_patch")]
|
||||
feature_settings: Option<Option<serde_json::Value>>,
|
||||
#[serde(default, deserialize_with = "deserialize_optional_string_list_patch")]
|
||||
allowed_ips: Option<Option<Vec<String>>>,
|
||||
#[serde(
|
||||
default,
|
||||
alias = "allowed_ips",
|
||||
deserialize_with = "deserialize_optional_string_list_patch"
|
||||
)]
|
||||
ip_rules: Option<Option<Vec<String>>>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
@@ -165,7 +169,7 @@ fn build_users_me_api_key_list_payload(
|
||||
"rate_limit": record.rate_limit,
|
||||
"concurrent_limit": record.concurrent_limit,
|
||||
"allowed_providers": record.allowed_providers,
|
||||
"allowed_ips": record.allowed_ips,
|
||||
"ip_rules": record.ip_rules,
|
||||
"force_capabilities": record.force_capabilities,
|
||||
"feature_settings": record.feature_settings,
|
||||
})
|
||||
@@ -183,7 +187,7 @@ fn build_users_me_api_key_detail_payload(
|
||||
"is_active": record.is_active,
|
||||
"is_locked": is_locked,
|
||||
"allowed_providers": record.allowed_providers,
|
||||
"allowed_ips": record.allowed_ips,
|
||||
"ip_rules": record.ip_rules,
|
||||
"force_capabilities": record.force_capabilities,
|
||||
"feature_settings": record.feature_settings,
|
||||
"rate_limit": record.rate_limit,
|
||||
@@ -206,50 +210,8 @@ fn generate_users_me_api_key_plaintext() -> String {
|
||||
generate_gateway_api_key_plaintext()
|
||||
}
|
||||
|
||||
fn users_me_validate_ip_or_cidr(value: &str) -> bool {
|
||||
let value = value.trim();
|
||||
if value.is_empty() {
|
||||
return false;
|
||||
}
|
||||
if value.parse::<std::net::IpAddr>().is_ok() {
|
||||
return true;
|
||||
}
|
||||
let Some((host, prefix)) = value.split_once('/') else {
|
||||
return false;
|
||||
};
|
||||
let Ok(ip) = host.trim().parse::<std::net::IpAddr>() else {
|
||||
return false;
|
||||
};
|
||||
let Ok(prefix) = prefix.trim().parse::<u8>() else {
|
||||
return false;
|
||||
};
|
||||
match ip {
|
||||
std::net::IpAddr::V4(_) => prefix <= 32,
|
||||
std::net::IpAddr::V6(_) => prefix <= 128,
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_users_me_allowed_ips(
|
||||
values: Option<Vec<String>>,
|
||||
) -> Result<Option<Vec<String>>, String> {
|
||||
let Some(values) = values else {
|
||||
return Ok(None);
|
||||
};
|
||||
if values.is_empty() {
|
||||
return Err("IP 白名单不能为空列表,如需取消限制请不提供此字段".to_string());
|
||||
}
|
||||
let mut normalized = Vec::with_capacity(values.len());
|
||||
for (index, raw) in values.into_iter().enumerate() {
|
||||
let trimmed = raw.trim();
|
||||
if trimmed.is_empty() {
|
||||
return Err(format!("IP 白名单第 {} 项为空", index + 1));
|
||||
}
|
||||
if !users_me_validate_ip_or_cidr(trimmed) {
|
||||
return Err(format!("无效的 IP 地址或 CIDR: {raw}"));
|
||||
}
|
||||
normalized.push(trimmed.to_string());
|
||||
}
|
||||
Ok(Some(normalized))
|
||||
fn normalize_users_me_ip_rules(values: Option<Vec<String>>) -> Result<Option<Vec<String>>, String> {
|
||||
normalize_ip_rules(values)
|
||||
}
|
||||
|
||||
fn hash_users_me_api_key(value: &str) -> String {
|
||||
@@ -595,7 +557,7 @@ pub(super) async fn handle_users_me_api_key_create(
|
||||
return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false);
|
||||
}
|
||||
};
|
||||
let allowed_ips = match normalize_users_me_allowed_ips(payload.allowed_ips) {
|
||||
let ip_rules = match normalize_users_me_ip_rules(payload.ip_rules) {
|
||||
Ok(value) => value,
|
||||
Err(detail) => {
|
||||
return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false);
|
||||
@@ -619,7 +581,7 @@ pub(super) async fn handle_users_me_api_key_create(
|
||||
allowed_providers: None,
|
||||
allowed_api_formats: None,
|
||||
allowed_models: None,
|
||||
allowed_ips,
|
||||
ip_rules,
|
||||
rate_limit,
|
||||
concurrent_limit,
|
||||
force_capabilities: None,
|
||||
@@ -674,7 +636,7 @@ pub(super) async fn handle_users_me_api_key_create(
|
||||
"is_locked": false,
|
||||
"rate_limit": created.rate_limit,
|
||||
"concurrent_limit": created.concurrent_limit,
|
||||
"allowed_ips": created.allowed_ips,
|
||||
"ip_rules": created.ip_rules,
|
||||
"feature_settings": created.feature_settings,
|
||||
"last_used_at": format_users_me_optional_unix_secs_iso8601(created.last_used_at_unix_secs),
|
||||
"created_at": format_users_me_optional_unix_secs_iso8601(created.created_at_unix_secs),
|
||||
@@ -756,8 +718,8 @@ pub(super) async fn handle_users_me_api_key_update(
|
||||
},
|
||||
None => None,
|
||||
};
|
||||
let allowed_ips = match payload.allowed_ips {
|
||||
Some(value) => match normalize_users_me_allowed_ips(value) {
|
||||
let ip_rules = match payload.ip_rules {
|
||||
Some(value) => match normalize_users_me_ip_rules(value) {
|
||||
Ok(value) => Some(value),
|
||||
Err(detail) => {
|
||||
return build_auth_error_response(http::StatusCode::BAD_REQUEST, detail, false);
|
||||
@@ -773,7 +735,7 @@ pub(super) async fn handle_users_me_api_key_update(
|
||||
name,
|
||||
rate_limit,
|
||||
concurrent_limit,
|
||||
allowed_ips,
|
||||
ip_rules,
|
||||
})
|
||||
.await
|
||||
{
|
||||
@@ -1131,16 +1093,16 @@ pub(super) async fn handle_users_me_api_key_capabilities_put(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{normalize_users_me_allowed_ips, UsersMeUpdateApiKeyRequest};
|
||||
use super::{normalize_users_me_ip_rules, UsersMeUpdateApiKeyRequest};
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn normalize_allowed_ips_trims_ip_and_cidr_values() {
|
||||
let values = normalize_users_me_allowed_ips(Some(vec![
|
||||
fn normalize_ip_rules_trims_ip_and_cidr_values() {
|
||||
let values = normalize_users_me_ip_rules(Some(vec![
|
||||
" 203.0.113.10 ".to_string(),
|
||||
"10.0.0.0/24".to_string(),
|
||||
]))
|
||||
.expect("valid whitelist should normalize");
|
||||
.expect("valid IP rules should normalize");
|
||||
|
||||
assert_eq!(
|
||||
values,
|
||||
@@ -1149,33 +1111,33 @@ mod tests {
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalize_allowed_ips_rejects_invalid_cidr() {
|
||||
let err = normalize_users_me_allowed_ips(Some(vec!["10.0.0.0/99".to_string()]))
|
||||
fn normalize_ip_rules_rejects_invalid_cidr() {
|
||||
let err = normalize_users_me_ip_rules(Some(vec!["10.0.0.0/99".to_string()]))
|
||||
.expect_err("invalid cidr should fail");
|
||||
|
||||
assert_eq!(err, "无效的 IP 地址或 CIDR: 10.0.0.0/99");
|
||||
assert_eq!(err, "无效的 IP 限制规则: 10.0.0.0/99(第 1 项)");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn update_payload_distinguishes_missing_null_and_present_allowed_ips() {
|
||||
fn update_payload_distinguishes_missing_null_and_present_ip_rules() {
|
||||
let missing = serde_json::from_value::<UsersMeUpdateApiKeyRequest>(json!({
|
||||
"name": "unchanged-whitelist",
|
||||
"name": "unchanged-ip-rules",
|
||||
}))
|
||||
.expect("missing allowed_ips should deserialize");
|
||||
assert_eq!(missing.allowed_ips, None);
|
||||
.expect("missing ip_rules should deserialize");
|
||||
assert_eq!(missing.ip_rules, None);
|
||||
|
||||
let cleared = serde_json::from_value::<UsersMeUpdateApiKeyRequest>(json!({
|
||||
"allowed_ips": null,
|
||||
"ip_rules": null,
|
||||
}))
|
||||
.expect("null allowed_ips should deserialize");
|
||||
assert_eq!(cleared.allowed_ips, Some(None));
|
||||
.expect("null ip_rules should deserialize");
|
||||
assert_eq!(cleared.ip_rules, Some(None));
|
||||
|
||||
let updated = serde_json::from_value::<UsersMeUpdateApiKeyRequest>(json!({
|
||||
"allowed_ips": ["203.0.113.10", "10.0.0.0/24"],
|
||||
"ip_rules": ["203.0.113.10", "10.0.0.0/24"],
|
||||
}))
|
||||
.expect("present allowed_ips should deserialize");
|
||||
.expect("present ip_rules should deserialize");
|
||||
assert_eq!(
|
||||
updated.allowed_ips,
|
||||
updated.ip_rules,
|
||||
Some(Some(vec![
|
||||
"203.0.113.10".to_string(),
|
||||
"10.0.0.0/24".to_string(),
|
||||
|
||||
@@ -20,7 +20,7 @@ use super::{
|
||||
GatewayPublicRequestContext,
|
||||
};
|
||||
use crate::control::normalize_assignable_management_token_permissions;
|
||||
use crate::handlers::shared::generate_gateway_secret_plaintext;
|
||||
use crate::handlers::shared::{generate_gateway_secret_plaintext, parse_json_ip_rules};
|
||||
use crate::LocalMutationOutcome;
|
||||
|
||||
const USERS_ME_MANAGEMENT_TOKEN_PREFIX: &str = "ae";
|
||||
@@ -177,59 +177,10 @@ fn users_me_management_token_skip(query: Option<&str>) -> usize {
|
||||
.unwrap_or(0)
|
||||
}
|
||||
|
||||
fn users_me_validate_ip_or_cidr(value: &str) -> bool {
|
||||
let value = value.trim();
|
||||
if value.is_empty() {
|
||||
return false;
|
||||
}
|
||||
if value.parse::<std::net::IpAddr>().is_ok() {
|
||||
return true;
|
||||
}
|
||||
let Some((host, prefix)) = value.split_once('/') else {
|
||||
return false;
|
||||
};
|
||||
let Ok(ip) = host.trim().parse::<std::net::IpAddr>() else {
|
||||
return false;
|
||||
};
|
||||
let Ok(prefix) = prefix.trim().parse::<u8>() else {
|
||||
return false;
|
||||
};
|
||||
match ip {
|
||||
std::net::IpAddr::V4(_) => prefix <= 32,
|
||||
std::net::IpAddr::V6(_) => prefix <= 128,
|
||||
}
|
||||
}
|
||||
|
||||
fn users_me_parse_management_token_allowed_ips(
|
||||
value: Option<&serde_json::Value>,
|
||||
) -> Result<Option<serde_json::Value>, String> {
|
||||
let Some(value) = value else {
|
||||
return Ok(None);
|
||||
};
|
||||
match value {
|
||||
serde_json::Value::Null => Ok(None),
|
||||
serde_json::Value::Array(items) => {
|
||||
if items.is_empty() {
|
||||
return Err("IP 白名单不能为空列表,如需取消限制请不提供此字段".to_string());
|
||||
}
|
||||
let mut normalized = Vec::with_capacity(items.len());
|
||||
for (index, item) in items.iter().enumerate() {
|
||||
let Some(raw) = item.as_str() else {
|
||||
return Err("IP 白名单必须是字符串数组".to_string());
|
||||
};
|
||||
let trimmed = raw.trim();
|
||||
if trimmed.is_empty() {
|
||||
return Err(format!("IP 白名单第 {} 项为空", index + 1));
|
||||
}
|
||||
if !users_me_validate_ip_or_cidr(trimmed) {
|
||||
return Err(format!("无效的 IP 地址或 CIDR: {raw}"));
|
||||
}
|
||||
normalized.push(trimmed.to_string());
|
||||
}
|
||||
Ok(Some(json!(normalized)))
|
||||
}
|
||||
_ => Err("IP 白名单必须是字符串数组".to_string()),
|
||||
}
|
||||
parse_json_ip_rules(value)
|
||||
}
|
||||
|
||||
fn users_me_parse_management_token_expires_at(
|
||||
|
||||
Reference in New Issue
Block a user