feat: support api key ip restriction rules

This commit is contained in:
fawney19
2026-05-20 16:11:49 +08:00
parent b6bdc08267
commit f76bbaab52
55 changed files with 698 additions and 551 deletions
@@ -37,7 +37,7 @@ SELECT
api_keys.allowed_providers AS api_key_allowed_providers,
api_keys.allowed_api_formats AS api_key_allowed_api_formats,
api_keys.allowed_models AS api_key_allowed_models,
api_keys.allowed_ips AS api_key_allowed_ips
api_keys.ip_rules AS api_key_ip_rules
FROM api_keys
JOIN users ON users.id = api_keys.user_id
WHERE api_keys.key_hash = $1
@@ -68,7 +68,7 @@ SELECT
api_keys.allowed_providers AS api_key_allowed_providers,
api_keys.allowed_api_formats AS api_key_allowed_api_formats,
api_keys.allowed_models AS api_key_allowed_models,
api_keys.allowed_ips AS api_key_allowed_ips
api_keys.ip_rules AS api_key_ip_rules
FROM api_keys
JOIN users ON users.id = api_keys.user_id
WHERE api_keys.id = $1
@@ -99,7 +99,7 @@ SELECT
api_keys.allowed_providers AS api_key_allowed_providers,
api_keys.allowed_api_formats AS api_key_allowed_api_formats,
api_keys.allowed_models AS api_key_allowed_models,
api_keys.allowed_ips AS api_key_allowed_ips
api_keys.ip_rules AS api_key_ip_rules
FROM api_keys
JOIN users ON users.id = api_keys.user_id
WHERE api_keys.id = $1 AND users.id = $2
@@ -130,7 +130,7 @@ SELECT
api_keys.allowed_providers AS api_key_allowed_providers,
api_keys.allowed_api_formats AS api_key_allowed_api_formats,
api_keys.allowed_models AS api_key_allowed_models,
api_keys.allowed_ips AS api_key_allowed_ips
api_keys.ip_rules AS api_key_ip_rules
FROM api_keys
JOIN users ON users.id = api_keys.user_id
WHERE api_keys.id = ANY($1::TEXT[])
@@ -147,7 +147,7 @@ SELECT
api_keys.allowed_providers,
api_keys.allowed_api_formats,
api_keys.allowed_models,
api_keys.allowed_ips,
api_keys.ip_rules,
api_keys.rate_limit,
api_keys.concurrent_limit,
api_keys.force_capabilities,
@@ -178,7 +178,7 @@ SELECT
api_keys.allowed_providers,
api_keys.allowed_api_formats,
api_keys.allowed_models,
api_keys.allowed_ips,
api_keys.ip_rules,
api_keys.rate_limit,
api_keys.concurrent_limit,
api_keys.force_capabilities,
@@ -208,7 +208,7 @@ SELECT
api_keys.allowed_providers,
api_keys.allowed_api_formats,
api_keys.allowed_models,
api_keys.allowed_ips,
api_keys.ip_rules,
api_keys.rate_limit,
api_keys.concurrent_limit,
api_keys.force_capabilities,
@@ -238,7 +238,7 @@ SELECT
api_keys.allowed_providers,
api_keys.allowed_api_formats,
api_keys.allowed_models,
api_keys.allowed_ips,
api_keys.ip_rules,
api_keys.rate_limit,
api_keys.concurrent_limit,
api_keys.force_capabilities,
@@ -268,7 +268,7 @@ SELECT
api_keys.allowed_providers,
api_keys.allowed_api_formats,
api_keys.allowed_models,
api_keys.allowed_ips,
api_keys.ip_rules,
api_keys.rate_limit,
api_keys.concurrent_limit,
api_keys.force_capabilities,
@@ -342,7 +342,7 @@ SELECT
api_keys.allowed_providers,
api_keys.allowed_api_formats,
api_keys.allowed_models,
api_keys.allowed_ips,
api_keys.ip_rules,
api_keys.rate_limit,
api_keys.concurrent_limit,
api_keys.force_capabilities,
@@ -379,7 +379,7 @@ INSERT INTO api_keys (
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
ip_rules,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -429,7 +429,7 @@ RETURNING
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
ip_rules,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -456,7 +456,7 @@ INSERT INTO api_keys (
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
ip_rules,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -506,7 +506,7 @@ RETURNING
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
ip_rules,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -529,7 +529,7 @@ SET
name = COALESCE($3, name),
rate_limit = COALESCE($4, rate_limit),
concurrent_limit = COALESCE($5, concurrent_limit),
allowed_ips = CASE WHEN $6 THEN $7::jsonb ELSE allowed_ips END,
ip_rules = CASE WHEN $6 THEN $7::jsonb ELSE ip_rules END,
updated_at = NOW()
WHERE user_id = $1
AND id = $2
@@ -543,7 +543,7 @@ RETURNING
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
ip_rules,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -569,7 +569,7 @@ SET
allowed_providers = CASE WHEN $7 THEN $8::json ELSE allowed_providers END,
allowed_api_formats = CASE WHEN $9 THEN $10::json ELSE allowed_api_formats END,
allowed_models = CASE WHEN $11 THEN $12::json ELSE allowed_models END,
allowed_ips = CASE WHEN $13 THEN $14::jsonb ELSE allowed_ips END,
ip_rules = CASE WHEN $13 THEN $14::jsonb ELSE ip_rules END,
expires_at = CASE WHEN $15 THEN $16::timestamptz ELSE expires_at END,
auto_delete_on_expiry = CASE WHEN $17 THEN $18 ELSE auto_delete_on_expiry END,
updated_at = NOW()
@@ -584,7 +584,7 @@ RETURNING
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
ip_rules,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -618,7 +618,7 @@ RETURNING
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
ip_rules,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -651,7 +651,7 @@ RETURNING
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
ip_rules,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -695,7 +695,7 @@ RETURNING
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
ip_rules,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -1145,8 +1145,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
.map(serde_json::to_value)
.transpose()
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
let allowed_ips = record
.allowed_ips
let ip_rules = record
.ip_rules
.map(serde_json::to_value)
.transpose()
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
@@ -1167,7 +1167,7 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
.bind(allowed_providers)
.bind(allowed_api_formats)
.bind(allowed_models)
.bind(allowed_ips)
.bind(ip_rules)
.bind(record.rate_limit)
.bind(record.concurrent_limit)
.bind(record.force_capabilities)
@@ -1202,8 +1202,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
.map(serde_json::to_value)
.transpose()
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
let allowed_ips = record
.allowed_ips
let ip_rules = record
.ip_rules
.map(serde_json::to_value)
.transpose()
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
@@ -1224,7 +1224,7 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
.bind(allowed_providers)
.bind(allowed_api_formats)
.bind(allowed_models)
.bind(allowed_ips)
.bind(ip_rules)
.bind(record.rate_limit)
.bind(record.concurrent_limit)
.bind(record.force_capabilities)
@@ -1244,8 +1244,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
&self,
record: UpdateUserApiKeyBasicRecord,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
let allowed_ips = record
.allowed_ips
let ip_rules = record
.ip_rules
.clone()
.flatten()
.map(serde_json::to_value)
@@ -1257,8 +1257,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
.bind(record.name)
.bind(record.rate_limit)
.bind(record.concurrent_limit)
.bind(record.allowed_ips.is_some())
.bind(allowed_ips)
.bind(record.ip_rules.is_some())
.bind(ip_rules)
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
@@ -1290,8 +1290,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
.map(serde_json::to_value)
.transpose()
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
let allowed_ips = record
.allowed_ips
let ip_rules = record
.ip_rules
.clone()
.flatten()
.map(serde_json::to_value)
@@ -1318,8 +1318,8 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
.bind(allowed_api_formats)
.bind(record.allowed_models.is_some())
.bind(allowed_models)
.bind(record.allowed_ips.is_some())
.bind(allowed_ips)
.bind(record.ip_rules.is_some())
.bind(ip_rules)
.bind(record.expires_at_present)
.bind(expires_at)
.bind(record.auto_delete_on_expiry_present)
@@ -1549,7 +1549,7 @@ fn map_auth_api_key_snapshot_row(
row_get(row, "api_key_allowed_api_formats")?,
row_get(row, "api_key_allowed_models")?,
)?
.with_api_key_allowed_ips(row_get(row, "api_key_allowed_ips")?)?;
.with_api_key_ip_rules(row_get(row, "api_key_ip_rules")?)?;
Ok(snapshot.with_user_rate_limit(row_get(row, "user_rate_limit")?))
}
@@ -1577,7 +1577,7 @@ fn map_auth_api_key_export_row(
row_get(row, "total_cost_usd")?,
row_get(row, "is_standalone")?,
)
.and_then(|record| record.with_allowed_ips(row_get(row, "allowed_ips")?))
.and_then(|record| record.with_ip_rules(row_get(row, "ip_rules")?))
.map(|record| record.with_feature_settings(feature_settings))
.and_then(|record| {
record.with_activity_timestamps(
@@ -1622,7 +1622,7 @@ mod tests {
assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL
.contains("allowed_models = CASE WHEN $11 THEN $12::json ELSE allowed_models END"));
assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL
.contains("allowed_ips = CASE WHEN $13 THEN $14::jsonb ELSE allowed_ips END"));
.contains("ip_rules = CASE WHEN $13 THEN $14::jsonb ELSE ip_rules END"));
assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL
.contains("rate_limit = CASE WHEN $3 THEN $4 ELSE rate_limit END"));
assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL
@@ -1633,9 +1633,9 @@ mod tests {
}
#[test]
fn update_user_api_key_basic_sql_casts_allowed_ips_as_jsonb() {
fn update_user_api_key_basic_sql_casts_ip_rules_as_jsonb() {
assert!(UPDATE_USER_API_KEY_BASIC_SQL
.contains("allowed_ips = CASE WHEN $6 THEN $7::jsonb ELSE allowed_ips END"));
.contains("ip_rules = CASE WHEN $6 THEN $7::jsonb ELSE ip_rules END"));
}
#[tokio::test]