mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
Persist allowed IPs in auth repositories
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"
|
||||
));
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user