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

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::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"
));
}

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,