Persist allowed IPs in auth repositories

This commit is contained in:
RWDai
2026-05-18 20:47:17 +08:00
parent 437024cdb1
commit 276d19b63c
4 changed files with 199 additions and 24 deletions
@@ -36,7 +36,8 @@ SELECT
CAST(EXTRACT(EPOCH FROM api_keys.expires_at) AS BIGINT) AS api_key_expires_at_unix_secs,
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_models AS api_key_allowed_models,
api_keys.allowed_ips AS api_key_allowed_ips
FROM api_keys
JOIN users ON users.id = api_keys.user_id
WHERE api_keys.key_hash = $1
@@ -66,7 +67,8 @@ SELECT
CAST(EXTRACT(EPOCH FROM api_keys.expires_at) AS BIGINT) AS api_key_expires_at_unix_secs,
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_models AS api_key_allowed_models,
api_keys.allowed_ips AS api_key_allowed_ips
FROM api_keys
JOIN users ON users.id = api_keys.user_id
WHERE api_keys.id = $1
@@ -96,7 +98,8 @@ SELECT
CAST(EXTRACT(EPOCH FROM api_keys.expires_at) AS BIGINT) AS api_key_expires_at_unix_secs,
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_models AS api_key_allowed_models,
api_keys.allowed_ips AS api_key_allowed_ips
FROM api_keys
JOIN users ON users.id = api_keys.user_id
WHERE api_keys.id = $1 AND users.id = $2
@@ -126,7 +129,8 @@ SELECT
CAST(EXTRACT(EPOCH FROM api_keys.expires_at) AS BIGINT) AS api_key_expires_at_unix_secs,
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_models AS api_key_allowed_models,
api_keys.allowed_ips AS api_key_allowed_ips
FROM api_keys
JOIN users ON users.id = api_keys.user_id
WHERE api_keys.id = ANY($1::TEXT[])
@@ -143,6 +147,7 @@ SELECT
api_keys.allowed_providers,
api_keys.allowed_api_formats,
api_keys.allowed_models,
api_keys.allowed_ips,
api_keys.rate_limit,
api_keys.concurrent_limit,
api_keys.force_capabilities,
@@ -173,6 +178,7 @@ SELECT
api_keys.allowed_providers,
api_keys.allowed_api_formats,
api_keys.allowed_models,
api_keys.allowed_ips,
api_keys.rate_limit,
api_keys.concurrent_limit,
api_keys.force_capabilities,
@@ -202,6 +208,7 @@ SELECT
api_keys.allowed_providers,
api_keys.allowed_api_formats,
api_keys.allowed_models,
api_keys.allowed_ips,
api_keys.rate_limit,
api_keys.concurrent_limit,
api_keys.force_capabilities,
@@ -231,6 +238,7 @@ SELECT
api_keys.allowed_providers,
api_keys.allowed_api_formats,
api_keys.allowed_models,
api_keys.allowed_ips,
api_keys.rate_limit,
api_keys.concurrent_limit,
api_keys.force_capabilities,
@@ -260,6 +268,7 @@ SELECT
api_keys.allowed_providers,
api_keys.allowed_api_formats,
api_keys.allowed_models,
api_keys.allowed_ips,
api_keys.rate_limit,
api_keys.concurrent_limit,
api_keys.force_capabilities,
@@ -333,6 +342,7 @@ SELECT
api_keys.allowed_providers,
api_keys.allowed_api_formats,
api_keys.allowed_models,
api_keys.allowed_ips,
api_keys.rate_limit,
api_keys.concurrent_limit,
api_keys.force_capabilities,
@@ -369,6 +379,7 @@ INSERT INTO api_keys (
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -396,15 +407,16 @@ VALUES (
$9,
$10,
$11,
NULL,
$12,
NULL,
$13,
$14,
FALSE,
FALSE,
$15,
FALSE,
FALSE,
$16,
$17,
$18,
NOW(),
NOW()
)
@@ -417,6 +429,7 @@ RETURNING
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -443,6 +456,7 @@ INSERT INTO api_keys (
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -470,15 +484,16 @@ VALUES (
$9,
$10,
$11,
NULL,
$12,
NULL,
$13,
$14,
$15,
FALSE,
TRUE,
$15,
$16,
$17,
$18,
NOW(),
NOW()
)
@@ -491,6 +506,7 @@ RETURNING
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -513,6 +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::json ELSE allowed_ips END,
updated_at = NOW()
WHERE user_id = $1
AND id = $2
@@ -526,6 +543,7 @@ RETURNING
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -551,8 +569,9 @@ 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,
expires_at = CASE WHEN $13 THEN $14::timestamptz ELSE expires_at END,
auto_delete_on_expiry = CASE WHEN $15 THEN $16 ELSE auto_delete_on_expiry END,
allowed_ips = CASE WHEN $13 THEN $14::json ELSE allowed_ips 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()
WHERE id = $1
AND is_standalone = TRUE
@@ -565,6 +584,7 @@ RETURNING
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -598,6 +618,7 @@ RETURNING
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -630,6 +651,7 @@ RETURNING
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -673,6 +695,7 @@ RETURNING
allowed_providers,
allowed_api_formats,
allowed_models,
allowed_ips,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -1122,6 +1145,11 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
.map(serde_json::to_value)
.transpose()
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
let allowed_ips = record
.allowed_ips
.map(serde_json::to_value)
.transpose()
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
let expires_at = record
.expires_at_unix_secs
.map(|value| {
@@ -1139,6 +1167,7 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
.bind(allowed_providers)
.bind(allowed_api_formats)
.bind(allowed_models)
.bind(allowed_ips)
.bind(record.rate_limit)
.bind(record.concurrent_limit)
.bind(record.force_capabilities)
@@ -1173,6 +1202,11 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
.map(serde_json::to_value)
.transpose()
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
let allowed_ips = record
.allowed_ips
.map(serde_json::to_value)
.transpose()
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
let expires_at = record
.expires_at_unix_secs
.map(|value| {
@@ -1190,6 +1224,7 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
.bind(allowed_providers)
.bind(allowed_api_formats)
.bind(allowed_models)
.bind(allowed_ips)
.bind(record.rate_limit)
.bind(record.concurrent_limit)
.bind(record.force_capabilities)
@@ -1209,12 +1244,21 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
&self,
record: UpdateUserApiKeyBasicRecord,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
let allowed_ips = record
.allowed_ips
.clone()
.flatten()
.map(serde_json::to_value)
.transpose()
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
let row = sqlx::query(UPDATE_USER_API_KEY_BASIC_SQL)
.bind(record.user_id)
.bind(record.api_key_id)
.bind(record.name)
.bind(record.rate_limit)
.bind(record.concurrent_limit)
.bind(record.allowed_ips.is_some())
.bind(allowed_ips)
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
@@ -1246,6 +1290,13 @@ impl AuthApiKeyWriteRepository for SqlxAuthApiKeySnapshotReadRepository {
.map(serde_json::to_value)
.transpose()
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
let allowed_ips = record
.allowed_ips
.clone()
.flatten()
.map(serde_json::to_value)
.transpose()
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
let expires_at = record
.expires_at_unix_secs
.map(|value| {
@@ -1267,6 +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.expires_at_present)
.bind(expires_at)
.bind(record.auto_delete_on_expiry_present)
@@ -1495,7 +1548,8 @@ fn map_auth_api_key_snapshot_row(
row_get(row, "api_key_allowed_providers")?,
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")?)?;
Ok(snapshot.with_user_rate_limit(row_get(row, "user_rate_limit")?))
}
@@ -1523,6 +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")?))
.map(|record| record.with_feature_settings(feature_settings))
.and_then(|record| {
record.with_activity_timestamps(
@@ -1546,12 +1601,12 @@ mod tests {
assert!(CREATE_USER_API_KEY_SQL
.contains("expires_at,\n auto_delete_on_expiry,\n is_locked,\n is_standalone,"));
assert!(
CREATE_USER_API_KEY_SQL.contains("$12,\n $13,\n $14,\n FALSE,\n FALSE,\n $15,")
CREATE_USER_API_KEY_SQL.contains("$13,\n $14,\n $15,\n FALSE,\n FALSE,\n $16,")
);
assert!(CREATE_STANDALONE_API_KEY_SQL
.contains("expires_at,\n auto_delete_on_expiry,\n is_locked,\n is_standalone,"));
assert!(CREATE_STANDALONE_API_KEY_SQL
.contains("$12,\n $13,\n $14,\n FALSE,\n TRUE,\n $15,"));
.contains("$13,\n $14,\n $15,\n FALSE,\n TRUE,\n $16,"));
}
#[test]
@@ -1565,12 +1620,14 @@ 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::json ELSE allowed_ips 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
.contains("expires_at = CASE WHEN $13 THEN $14::timestamptz ELSE expires_at END"));
.contains("expires_at = CASE WHEN $15 THEN $16::timestamptz ELSE expires_at END"));
assert!(UPDATE_STANDALONE_API_KEY_BASIC_SQL.contains(
"auto_delete_on_expiry = CASE WHEN $15 THEN $16 ELSE auto_delete_on_expiry END"
"auto_delete_on_expiry = CASE WHEN $17 THEN $18 ELSE auto_delete_on_expiry END"
));
}