Merge branch 'pr-501'

This commit is contained in:
fawney19
2026-05-20 14:05:00 +08:00
45 changed files with 751 additions and 46 deletions

View File

@@ -2121,6 +2121,7 @@ mod tests {
api_key_allowed_providers: None,
api_key_allowed_api_formats: None,
api_key_allowed_models: None,
api_key_allowed_ips: None,
}
}

View File

@@ -0,0 +1,2 @@
ALTER TABLE api_keys
ADD COLUMN allowed_ips TEXT NULL AFTER allowed_models;

View File

@@ -0,0 +1,5 @@
ALTER TABLE api_keys
ADD COLUMN IF NOT EXISTS allowed_ips jsonb NULL;
ALTER TABLE api_keys
ALTER COLUMN allowed_ips TYPE jsonb USING allowed_ips::jsonb;

View File

@@ -0,0 +1 @@
ALTER TABLE api_keys ADD COLUMN allowed_ips TEXT;

View File

@@ -176,6 +176,7 @@ CREATE TABLE IF NOT EXISTS public.api_keys (
allowed_providers json,
allowed_api_formats json,
allowed_models json,
allowed_ips jsonb,
rate_limit integer DEFAULT 100,
concurrent_limit integer,
force_capabilities json,

View File

@@ -75,6 +75,7 @@ CREATE TABLE IF NOT EXISTS api_keys (
`allowed_models` JSON,
`allowed_providers` JSON,
`allowed_api_formats` JSON,
`allowed_ips` JSON,
`rate_limit` INT DEFAULT 100,
`concurrent_limit` INT,
`force_capabilities` JSON,

View File

@@ -78,6 +78,7 @@ CREATE TABLE IF NOT EXISTS public.api_keys (
allowed_models jsonb,
allowed_providers jsonb,
allowed_api_formats jsonb,
allowed_ips jsonb,
rate_limit integer DEFAULT 100,
concurrent_limit integer,
force_capabilities jsonb,

View File

@@ -73,6 +73,7 @@ CREATE TABLE IF NOT EXISTS api_keys (
allowed_models TEXT,
allowed_providers TEXT,
allowed_api_formats TEXT,
allowed_ips TEXT,
rate_limit INTEGER DEFAULT 100,
concurrent_limit INTEGER,
force_capabilities TEXT,

View File

@@ -333,6 +333,11 @@ name = "allowed_api_formats"
type = "json"
nullable = true
[[table.api_keys.columns]]
name = "allowed_ips"
type = "json"
nullable = true
[[table.api_keys.columns]]
name = "rate_limit"
type = "int32"

View File

@@ -307,6 +307,7 @@ fn empty_database_snapshot_covers_current_cutoff_versions() {
20260515000000,
20260516000000,
20260518000000,
20260518100000,
20260519000000,
20260519120000,
20260519130000,
@@ -469,6 +470,34 @@ fn management_tokens_json_columns_are_normalized_to_jsonb_in_postgres_schema_pat
assert!(generated_identity.contains("permissions jsonb,"));
}
#[test]
fn api_key_allowed_ips_is_jsonb_in_postgres_schema_paths() {
let api_key_allowed_ips_migration = POSTGRES_MIGRATOR
.iter()
.find(|migration| migration.version == 20260518100000)
.expect("api key allowed IPs migration should be embedded");
assert!(api_key_allowed_ips_migration
.sql
.contains("ADD COLUMN IF NOT EXISTS allowed_ips jsonb NULL"));
assert!(api_key_allowed_ips_migration
.sql
.contains("ALTER COLUMN allowed_ips TYPE jsonb USING allowed_ips::jsonb"));
assert!(EMPTY_DATABASE_SNAPSHOT_SQL.contains("allowed_ips jsonb,"));
let bootstrap_schema =
include_str!("../../../schema/bootstrap/postgres/001_types_and_tables.sql");
assert!(bootstrap_schema.contains("allowed_ips jsonb,"));
let driver_schema =
include_str!("../../../schema/drivers/postgres/baseline/001_types_and_tables.sql");
assert!(driver_schema.contains("allowed_ips jsonb,"));
let generated_identity =
include_str!("../../../schema/generated/postgres/baseline/001_identity.sql");
assert!(generated_identity.contains("allowed_ips jsonb,"));
}
#[test]
fn provider_api_keys_api_key_is_nullable() {
let baseline_migration = POSTGRES_MIGRATOR
@@ -602,6 +631,7 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() {
20260512110000,
20260516000000,
20260518000000,
20260518100000,
20260519000000,
20260519120000,
20260519130000,
@@ -623,6 +653,7 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() {
20260512110000,
20260516000000,
20260518000000,
20260518100000,
20260519000000,
20260519120000,
20260519130000,
@@ -1144,6 +1175,7 @@ fn pending_migrations_from_applied_skips_versions_already_applied() {
20260515000000,
20260516000000,
20260518000000,
20260518100000,
20260519000000,
20260519120000,
20260519130000,

View File

@@ -78,6 +78,14 @@ impl InMemoryAuthApiKeySnapshotRepository {
0.0,
snapshot.api_key_is_standalone,
)
.and_then(|record| {
record.with_allowed_ips(
snapshot
.api_key_allowed_ips
.as_ref()
.map(|value| serde_json::json!(value)),
)
})
.expect("derived auth api key export record should build"),
);
if let Some(key_hash) = key_hash {
@@ -496,6 +504,7 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository {
api_key_allowed_providers: record.allowed_providers.clone(),
api_key_allowed_api_formats: record.allowed_api_formats.clone(),
api_key_allowed_models: record.allowed_models.clone(),
api_key_allowed_ips: record.allowed_ips.clone(),
..template
}
} else {
@@ -534,6 +543,12 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository {
.as_ref()
.map(|value| serde_json::json!(value)),
)?
.with_api_key_allowed_ips(
record
.allowed_ips
.as_ref()
.map(|value| serde_json::json!(value)),
)?
};
let now_unix_secs = current_unix_secs() as i64;
@@ -566,6 +581,12 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository {
record.total_cost_usd,
false,
)?
.with_allowed_ips(
record
.allowed_ips
.as_ref()
.map(|value| serde_json::json!(value)),
)?
.with_activity_timestamps(None, Some(now_unix_secs), Some(now_unix_secs))?;
index
@@ -619,6 +640,7 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository {
api_key_allowed_providers: record.allowed_providers.clone(),
api_key_allowed_api_formats: record.allowed_api_formats.clone(),
api_key_allowed_models: record.allowed_models.clone(),
api_key_allowed_ips: record.allowed_ips.clone(),
..template
}
} else {
@@ -657,6 +679,12 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository {
.as_ref()
.map(|value| serde_json::json!(value)),
)?
.with_api_key_allowed_ips(
record
.allowed_ips
.as_ref()
.map(|value| serde_json::json!(value)),
)?
};
let now_unix_secs = current_unix_secs() as i64;
@@ -689,6 +717,12 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository {
record.total_cost_usd,
true,
)?
.with_allowed_ips(
record
.allowed_ips
.as_ref()
.map(|value| serde_json::json!(value)),
)?
.with_activity_timestamps(None, Some(now_unix_secs), Some(now_unix_secs))?;
index
@@ -741,6 +775,14 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository {
export.concurrent_limit = Some(concurrent_limit);
}
}
if let Some(allowed_ips) = record.allowed_ips {
if let Some(snapshot) = index.by_api_key_id.get_mut(&record.api_key_id) {
snapshot.api_key_allowed_ips = allowed_ips.clone();
}
if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) {
export.allowed_ips = allowed_ips;
}
}
Ok(index.export_by_api_key_id.get(&record.api_key_id).cloned())
}
@@ -806,6 +848,14 @@ impl AuthApiKeyWriteRepository for InMemoryAuthApiKeySnapshotRepository {
export.allowed_models = allowed_models;
}
}
if let Some(allowed_ips) = record.allowed_ips {
if let Some(snapshot) = index.by_api_key_id.get_mut(&record.api_key_id) {
snapshot.api_key_allowed_ips = allowed_ips.clone();
}
if let Some(export) = index.export_by_api_key_id.get_mut(&record.api_key_id) {
export.allowed_ips = allowed_ips;
}
}
if record.expires_at_present {
if let Some(snapshot) = index.by_api_key_id.get_mut(&record.api_key_id) {
snapshot.api_key_expires_at_unix_secs = record.expires_at_unix_secs;
@@ -1239,6 +1289,7 @@ mod tests {
name: None,
rate_limit: None,
concurrent_limit: Some(11),
allowed_ips: None,
})
.await
.expect("update should succeed")
@@ -1273,6 +1324,7 @@ mod tests {
allowed_providers: None,
allowed_api_formats: None,
allowed_models: None,
allowed_ips: None,
expires_at_present: false,
expires_at_unix_secs: None,
auto_delete_on_expiry_present: false,

View File

@@ -34,7 +34,8 @@ SELECT
api_keys.expires_at 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
"#;
@@ -49,6 +50,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,
@@ -112,12 +114,12 @@ impl MysqlAuthApiKeyReadRepository {
r#"
INSERT INTO api_keys (
id, user_id, key_hash, key_encrypted, name, allowed_providers,
allowed_api_formats, allowed_models, rate_limit, concurrent_limit,
allowed_api_formats, allowed_models, allowed_ips, rate_limit, concurrent_limit,
force_capabilities, feature_settings, is_active, expires_at, auto_delete_on_expiry,
total_requests, total_tokens, total_cost_usd, is_standalone,
created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&record.api_key_id)
@@ -137,6 +139,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?)
record.allowed_models.as_ref(),
"api_keys.allowed_models",
)?)
.bind(json_string_from_string_list(
record.allowed_ips.as_ref(),
"api_keys.allowed_ips",
)?)
.bind(record.rate_limit)
.bind(record.concurrent_limit)
.bind(optional_json_to_string(
@@ -175,6 +181,7 @@ struct CreateApiKeyInsertRecord {
allowed_providers: Option<Vec<String>>,
allowed_api_formats: Option<Vec<String>>,
allowed_models: Option<Vec<String>>,
allowed_ips: Option<Vec<String>>,
rate_limit: Option<i32>,
concurrent_limit: Option<i32>,
force_capabilities: Option<serde_json::Value>,
@@ -417,6 +424,7 @@ WHERE id = ?
allowed_providers: record.allowed_providers,
allowed_api_formats: record.allowed_api_formats,
allowed_models: record.allowed_models,
allowed_ips: record.allowed_ips,
rate_limit: Some(record.rate_limit),
concurrent_limit: record.concurrent_limit,
force_capabilities: record.force_capabilities,
@@ -444,6 +452,7 @@ WHERE id = ?
allowed_providers: record.allowed_providers,
allowed_api_formats: record.allowed_api_formats,
allowed_models: record.allowed_models,
allowed_ips: record.allowed_ips,
rate_limit: record.rate_limit,
concurrent_limit: record.concurrent_limit,
force_capabilities: record.force_capabilities,
@@ -469,6 +478,7 @@ UPDATE api_keys
SET name = COALESCE(?, name),
rate_limit = COALESCE(?, rate_limit),
concurrent_limit = COALESCE(?, concurrent_limit),
allowed_ips = CASE WHEN ? THEN ? ELSE allowed_ips END,
updated_at = ?
WHERE id = ?
AND user_id = ?
@@ -478,6 +488,11 @@ WHERE id = ?
.bind(record.name.as_deref())
.bind(record.rate_limit)
.bind(record.concurrent_limit)
.bind(record.allowed_ips.is_some())
.bind(json_string_from_nested_string_list(
&record.allowed_ips,
"api_keys.allowed_ips",
)?)
.bind(now)
.bind(&record.api_key_id)
.bind(&record.user_id)
@@ -501,6 +516,7 @@ SET name = COALESCE(?, name),
allowed_providers = CASE WHEN ? THEN ? ELSE allowed_providers END,
allowed_api_formats = CASE WHEN ? THEN ? ELSE allowed_api_formats END,
allowed_models = CASE WHEN ? THEN ? ELSE allowed_models END,
allowed_ips = CASE WHEN ? THEN ? ELSE allowed_ips END,
expires_at = CASE WHEN ? THEN ? ELSE expires_at END,
auto_delete_on_expiry = CASE WHEN ? THEN ? ELSE auto_delete_on_expiry END,
updated_at = ?
@@ -528,6 +544,11 @@ WHERE id = ?
&record.allowed_models,
"api_keys.allowed_models",
)?)
.bind(record.allowed_ips.is_some())
.bind(json_string_from_nested_string_list(
&record.allowed_ips,
"api_keys.allowed_ips",
)?)
.bind(record.expires_at_present)
.bind(optional_i64_from_u64(
record.expires_at_unix_secs,
@@ -922,7 +943,11 @@ fn map_auth_api_key_snapshot_row(
row.try_get("api_key_allowed_models").map_sql_err()?,
"api_keys.allowed_models",
)?,
)?;
)?
.with_api_key_allowed_ips(optional_json_from_string(
row.try_get("api_key_allowed_ips").map_sql_err()?,
"api_keys.allowed_ips",
)?)?;
Ok(snapshot.with_user_rate_limit(row.try_get("user_rate_limit").map_sql_err()?))
}
@@ -965,6 +990,12 @@ fn map_auth_api_key_export_row(
row.try_get("total_cost_usd").map_sql_err()?,
row.try_get("is_standalone").map_sql_err()?,
)
.and_then(|record| {
record.with_allowed_ips(optional_json_from_string(
row.try_get("allowed_ips").map_sql_err()?,
"api_keys.allowed_ips",
)?)
})
.map(|record| record.with_feature_settings(feature_settings))
.and_then(|record| {
record.with_activity_timestamps(

View File

@@ -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::jsonb 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::jsonb 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(
@@ -1538,6 +1593,7 @@ mod tests {
use super::{
SqlxAuthApiKeySnapshotReadRepository, CREATE_STANDALONE_API_KEY_SQL,
CREATE_USER_API_KEY_SQL, UPDATE_STANDALONE_API_KEY_BASIC_SQL,
UPDATE_USER_API_KEY_BASIC_SQL,
};
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
@@ -1546,12 +1602,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,15 +1621,23 @@ 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"));
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"
));
}
#[test]
fn update_user_api_key_basic_sql_casts_allowed_ips_as_jsonb() {
assert!(UPDATE_USER_API_KEY_BASIC_SQL
.contains("allowed_ips = CASE WHEN $6 THEN $7::jsonb ELSE allowed_ips END"));
}
#[tokio::test]
async fn repository_constructs_from_lazy_pool() {
let factory = PostgresPoolFactory::new(PostgresPoolConfig {

View File

@@ -34,7 +34,8 @@ SELECT
api_keys.expires_at 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
"#;
@@ -49,6 +50,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,
@@ -112,12 +114,12 @@ impl SqliteAuthApiKeyReadRepository {
r#"
INSERT INTO api_keys (
id, user_id, key_hash, key_encrypted, name, allowed_providers,
allowed_api_formats, allowed_models, rate_limit, concurrent_limit,
allowed_api_formats, allowed_models, allowed_ips, rate_limit, concurrent_limit,
force_capabilities, feature_settings, is_active, expires_at, auto_delete_on_expiry,
total_requests, total_tokens, total_cost_usd, is_standalone,
created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&record.api_key_id)
@@ -137,6 +139,10 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?)
record.allowed_models.as_ref(),
"api_keys.allowed_models",
)?)
.bind(json_string_from_string_list(
record.allowed_ips.as_ref(),
"api_keys.allowed_ips",
)?)
.bind(record.rate_limit)
.bind(record.concurrent_limit)
.bind(optional_json_to_string(
@@ -175,6 +181,7 @@ struct CreateApiKeyInsertRecord {
allowed_providers: Option<Vec<String>>,
allowed_api_formats: Option<Vec<String>>,
allowed_models: Option<Vec<String>>,
allowed_ips: Option<Vec<String>>,
rate_limit: Option<i32>,
concurrent_limit: Option<i32>,
force_capabilities: Option<serde_json::Value>,
@@ -417,6 +424,7 @@ WHERE id = ?
allowed_providers: record.allowed_providers,
allowed_api_formats: record.allowed_api_formats,
allowed_models: record.allowed_models,
allowed_ips: record.allowed_ips,
rate_limit: Some(record.rate_limit),
concurrent_limit: record.concurrent_limit,
force_capabilities: record.force_capabilities,
@@ -444,6 +452,7 @@ WHERE id = ?
allowed_providers: record.allowed_providers,
allowed_api_formats: record.allowed_api_formats,
allowed_models: record.allowed_models,
allowed_ips: record.allowed_ips,
rate_limit: record.rate_limit,
concurrent_limit: record.concurrent_limit,
force_capabilities: record.force_capabilities,
@@ -469,6 +478,7 @@ UPDATE api_keys
SET name = COALESCE(?, name),
rate_limit = COALESCE(?, rate_limit),
concurrent_limit = COALESCE(?, concurrent_limit),
allowed_ips = CASE WHEN ? THEN ? ELSE allowed_ips END,
updated_at = ?
WHERE id = ?
AND user_id = ?
@@ -478,6 +488,11 @@ WHERE id = ?
.bind(record.name.as_deref())
.bind(record.rate_limit)
.bind(record.concurrent_limit)
.bind(record.allowed_ips.is_some())
.bind(json_string_from_nested_string_list(
&record.allowed_ips,
"api_keys.allowed_ips",
)?)
.bind(now)
.bind(&record.api_key_id)
.bind(&record.user_id)
@@ -501,6 +516,7 @@ SET name = COALESCE(?, name),
allowed_providers = CASE WHEN ? THEN ? ELSE allowed_providers END,
allowed_api_formats = CASE WHEN ? THEN ? ELSE allowed_api_formats END,
allowed_models = CASE WHEN ? THEN ? ELSE allowed_models END,
allowed_ips = CASE WHEN ? THEN ? ELSE allowed_ips END,
expires_at = CASE WHEN ? THEN ? ELSE expires_at END,
auto_delete_on_expiry = CASE WHEN ? THEN ? ELSE auto_delete_on_expiry END,
updated_at = ?
@@ -528,6 +544,11 @@ WHERE id = ?
&record.allowed_models,
"api_keys.allowed_models",
)?)
.bind(record.allowed_ips.is_some())
.bind(json_string_from_nested_string_list(
&record.allowed_ips,
"api_keys.allowed_ips",
)?)
.bind(record.expires_at_present)
.bind(optional_i64_from_u64(
record.expires_at_unix_secs,
@@ -922,7 +943,11 @@ fn map_auth_api_key_snapshot_row(
row.try_get("api_key_allowed_models").map_sql_err()?,
"api_keys.allowed_models",
)?,
)?;
)?
.with_api_key_allowed_ips(optional_json_from_string(
row.try_get("api_key_allowed_ips").map_sql_err()?,
"api_keys.allowed_ips",
)?)?;
Ok(snapshot.with_user_rate_limit(row.try_get("user_rate_limit").map_sql_err()?))
}
@@ -965,6 +990,12 @@ fn map_auth_api_key_export_row(
sqlite_real(row, "total_cost_usd")?,
row.try_get("is_standalone").map_sql_err()?,
)
.and_then(|record| {
record.with_allowed_ips(optional_json_from_string(
row.try_get("allowed_ips").map_sql_err()?,
"api_keys.allowed_ips",
)?)
})
.map(|record| record.with_feature_settings(feature_settings))
.and_then(|record| {
record.with_activity_timestamps(
@@ -1081,6 +1112,7 @@ mod tests {
allowed_providers: Some(vec!["openai".to_string()]),
allowed_api_formats: Some(vec!["openai:chat".to_string()]),
allowed_models: Some(vec!["gpt-4.1".to_string()]),
allowed_ips: Some(vec!["203.0.113.10".to_string()]),
rate_limit: 100,
concurrent_limit: Some(5),
force_capabilities: Some(json!({"cache": true})),
@@ -1104,6 +1136,7 @@ mod tests {
name: Some("Updated User".to_string()),
rate_limit: Some(150),
concurrent_limit: Some(6),
allowed_ips: Some(Some(vec!["10.0.0.0/24".to_string()])),
})
.await
.expect("user key should update")
@@ -1179,6 +1212,7 @@ mod tests {
allowed_providers: Some(vec!["openai".to_string()]),
allowed_api_formats: None,
allowed_models: None,
allowed_ips: None,
rate_limit: None,
concurrent_limit: Some(2),
force_capabilities: None,
@@ -1205,6 +1239,7 @@ mod tests {
allowed_providers: Some(None),
allowed_api_formats: Some(Some(vec!["openai:responses".to_string()])),
allowed_models: Some(Some(vec!["gpt-4.1-mini".to_string()])),
allowed_ips: None,
expires_at_present: true,
expires_at_unix_secs: Some(2_100_000_000),
auto_delete_on_expiry_present: true,

View File

@@ -24,6 +24,7 @@ pub struct StoredAuthApiKeySnapshot {
pub api_key_allowed_providers: Option<Vec<String>>,
pub api_key_allowed_api_formats: Option<Vec<String>>,
pub api_key_allowed_models: Option<Vec<String>>,
pub api_key_allowed_ips: Option<Vec<String>>,
}
impl StoredAuthApiKeySnapshot {
@@ -97,9 +98,18 @@ impl StoredAuthApiKeySnapshot {
api_key_allowed_models,
"api_keys.allowed_models",
)?,
api_key_allowed_ips: None,
})
}
pub fn with_api_key_allowed_ips(
mut self,
api_key_allowed_ips: Option<serde_json::Value>,
) -> Result<Self, crate::DataLayerError> {
self.api_key_allowed_ips = parse_string_list(api_key_allowed_ips, "api_keys.allowed_ips")?;
Ok(self)
}
pub fn is_currently_usable(&self, now_unix_secs: u64) -> bool {
if !self.user_is_active || self.user_is_deleted {
return false;
@@ -148,6 +158,7 @@ pub struct ResolvedAuthApiKeySnapshot {
pub api_key_allowed_providers: Option<Vec<String>>,
pub api_key_allowed_api_formats: Option<Vec<String>>,
pub api_key_allowed_models: Option<Vec<String>>,
pub api_key_allowed_ips: Option<Vec<String>>,
pub currently_usable: bool,
}
@@ -177,6 +188,7 @@ impl ResolvedAuthApiKeySnapshot {
api_key_allowed_providers: snapshot.api_key_allowed_providers,
api_key_allowed_api_formats: snapshot.api_key_allowed_api_formats,
api_key_allowed_models: snapshot.api_key_allowed_models,
api_key_allowed_ips: snapshot.api_key_allowed_ips,
currently_usable,
};
resolved.constrain_non_standalone_api_key_policy_to_user_policy();
@@ -332,6 +344,7 @@ pub struct StoredAuthApiKeyExportRecord {
pub allowed_providers: Option<Vec<String>>,
pub allowed_api_formats: Option<Vec<String>>,
pub allowed_models: Option<Vec<String>>,
pub allowed_ips: Option<Vec<String>>,
pub rate_limit: Option<i32>,
pub concurrent_limit: Option<i32>,
pub force_capabilities: Option<serde_json::Value>,
@@ -403,6 +416,7 @@ impl StoredAuthApiKeyExportRecord {
"api_keys.allowed_api_formats",
)?,
allowed_models: parse_string_list(allowed_models, "api_keys.allowed_models")?,
allowed_ips: None,
rate_limit,
concurrent_limit,
force_capabilities,
@@ -444,6 +458,14 @@ impl StoredAuthApiKeyExportRecord {
self.feature_settings = normalize_optional_json(feature_settings);
self
}
pub fn with_allowed_ips(
mut self,
allowed_ips: Option<serde_json::Value>,
) -> Result<Self, crate::DataLayerError> {
self.allowed_ips = parse_string_list(allowed_ips, "api_keys.allowed_ips")?;
Ok(self)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
@@ -469,6 +491,7 @@ pub struct CreateUserApiKeyRecord {
pub allowed_providers: Option<Vec<String>>,
pub allowed_api_formats: Option<Vec<String>>,
pub allowed_models: Option<Vec<String>>,
pub allowed_ips: Option<Vec<String>>,
pub rate_limit: i32,
pub concurrent_limit: Option<i32>,
pub force_capabilities: Option<serde_json::Value>,
@@ -487,6 +510,7 @@ pub struct UpdateUserApiKeyBasicRecord {
pub name: Option<String>,
pub rate_limit: Option<i32>,
pub concurrent_limit: Option<i32>,
pub allowed_ips: Option<Option<Vec<String>>>,
}
#[derive(Debug, Clone, PartialEq)]
@@ -499,6 +523,7 @@ pub struct CreateStandaloneApiKeyRecord {
pub allowed_providers: Option<Vec<String>>,
pub allowed_api_formats: Option<Vec<String>>,
pub allowed_models: Option<Vec<String>>,
pub allowed_ips: Option<Vec<String>>,
pub rate_limit: Option<i32>,
pub concurrent_limit: Option<i32>,
pub force_capabilities: Option<serde_json::Value>,
@@ -521,6 +546,7 @@ pub struct UpdateStandaloneApiKeyBasicRecord {
pub allowed_providers: Option<Option<Vec<String>>>,
pub allowed_api_formats: Option<Option<Vec<String>>>,
pub allowed_models: Option<Option<Vec<String>>>,
pub allowed_ips: Option<Option<Vec<String>>>,
pub expires_at_present: bool,
pub expires_at_unix_secs: Option<u64>,
pub auto_delete_on_expiry_present: bool,