mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
fix: align user group access controls
This commit is contained in:
@@ -51,3 +51,52 @@ CREATE TABLE IF NOT EXISTS user_group_members (
|
||||
CONSTRAINT user_group_members_user_id_fk
|
||||
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
INSERT IGNORE INTO user_groups (
|
||||
id,
|
||||
name,
|
||||
normalized_name,
|
||||
description,
|
||||
priority,
|
||||
allowed_providers_mode,
|
||||
allowed_api_formats_mode,
|
||||
allowed_models_mode,
|
||||
rate_limit_mode,
|
||||
created_at,
|
||||
updated_at
|
||||
)
|
||||
VALUES (
|
||||
'00000000-0000-0000-0000-000000000001',
|
||||
'Default',
|
||||
'default',
|
||||
'Default unrestricted group for all users',
|
||||
0,
|
||||
'unrestricted',
|
||||
'unrestricted',
|
||||
'unrestricted',
|
||||
'system',
|
||||
UNIX_TIMESTAMP(),
|
||||
UNIX_TIMESTAMP()
|
||||
);
|
||||
|
||||
INSERT IGNORE INTO system_configs (
|
||||
id,
|
||||
`key`,
|
||||
value,
|
||||
description,
|
||||
created_at,
|
||||
updated_at
|
||||
)
|
||||
VALUES (
|
||||
'00000000-0000-0000-0000-000000000002',
|
||||
'default_user_group_id',
|
||||
'"00000000-0000-0000-0000-000000000001"',
|
||||
'Default unrestricted user group',
|
||||
UNIX_TIMESTAMP(),
|
||||
UNIX_TIMESTAMP()
|
||||
);
|
||||
|
||||
INSERT IGNORE INTO user_group_members (group_id, user_id, created_at)
|
||||
SELECT '00000000-0000-0000-0000-000000000001', id, UNIX_TIMESTAMP()
|
||||
FROM users
|
||||
WHERE is_deleted = 0;
|
||||
|
||||
@@ -58,3 +58,42 @@ CREATE INDEX IF NOT EXISTS user_group_members_user_id_idx
|
||||
|
||||
CREATE INDEX IF NOT EXISTS user_groups_priority_name_idx
|
||||
ON public.user_groups (priority DESC, name ASC, id ASC);
|
||||
|
||||
INSERT INTO public.user_groups (
|
||||
id,
|
||||
name,
|
||||
normalized_name,
|
||||
description,
|
||||
priority,
|
||||
allowed_providers_mode,
|
||||
allowed_api_formats_mode,
|
||||
allowed_models_mode,
|
||||
rate_limit_mode
|
||||
)
|
||||
VALUES (
|
||||
'00000000-0000-0000-0000-000000000001',
|
||||
'Default',
|
||||
'default',
|
||||
'Default unrestricted group for all users',
|
||||
0,
|
||||
'unrestricted',
|
||||
'unrestricted',
|
||||
'unrestricted',
|
||||
'system'
|
||||
)
|
||||
ON CONFLICT (id) DO NOTHING;
|
||||
|
||||
INSERT INTO public.system_configs (id, key, value, description)
|
||||
VALUES (
|
||||
'00000000-0000-0000-0000-000000000002',
|
||||
'default_user_group_id',
|
||||
'"00000000-0000-0000-0000-000000000001"'::json,
|
||||
'Default unrestricted user group'
|
||||
)
|
||||
ON CONFLICT (key) DO NOTHING;
|
||||
|
||||
INSERT INTO public.user_group_members (group_id, user_id)
|
||||
SELECT '00000000-0000-0000-0000-000000000001', id
|
||||
FROM public.users
|
||||
WHERE is_deleted IS FALSE
|
||||
ON CONFLICT (group_id, user_id) DO NOTHING;
|
||||
|
||||
@@ -49,3 +49,52 @@ CREATE INDEX IF NOT EXISTS user_group_members_user_id_idx
|
||||
|
||||
CREATE INDEX IF NOT EXISTS user_groups_priority_name_idx
|
||||
ON user_groups (priority DESC, name ASC, id ASC);
|
||||
|
||||
INSERT OR IGNORE INTO user_groups (
|
||||
id,
|
||||
name,
|
||||
normalized_name,
|
||||
description,
|
||||
priority,
|
||||
allowed_providers_mode,
|
||||
allowed_api_formats_mode,
|
||||
allowed_models_mode,
|
||||
rate_limit_mode,
|
||||
created_at,
|
||||
updated_at
|
||||
)
|
||||
VALUES (
|
||||
'00000000-0000-0000-0000-000000000001',
|
||||
'Default',
|
||||
'default',
|
||||
'Default unrestricted group for all users',
|
||||
0,
|
||||
'unrestricted',
|
||||
'unrestricted',
|
||||
'unrestricted',
|
||||
'system',
|
||||
CAST(strftime('%s', 'now') AS INTEGER),
|
||||
CAST(strftime('%s', 'now') AS INTEGER)
|
||||
);
|
||||
|
||||
INSERT OR IGNORE INTO system_configs (
|
||||
id,
|
||||
key,
|
||||
value,
|
||||
description,
|
||||
created_at,
|
||||
updated_at
|
||||
)
|
||||
VALUES (
|
||||
'00000000-0000-0000-0000-000000000002',
|
||||
'default_user_group_id',
|
||||
'"00000000-0000-0000-0000-000000000001"',
|
||||
'Default unrestricted user group',
|
||||
CAST(strftime('%s', 'now') AS INTEGER),
|
||||
CAST(strftime('%s', 'now') AS INTEGER)
|
||||
);
|
||||
|
||||
INSERT OR IGNORE INTO user_group_members (group_id, user_id, created_at)
|
||||
SELECT '00000000-0000-0000-0000-000000000001', id, CAST(strftime('%s', 'now') AS INTEGER)
|
||||
FROM users
|
||||
WHERE is_deleted = 0;
|
||||
|
||||
@@ -1,5 +1,44 @@
|
||||
-- Restore a normal lookup path before sqlx records this migration in the
|
||||
-- same transaction. sqlx inserts into `_sqlx_migrations` unqualified.
|
||||
INSERT INTO public.user_groups (
|
||||
id,
|
||||
name,
|
||||
normalized_name,
|
||||
description,
|
||||
priority,
|
||||
allowed_providers_mode,
|
||||
allowed_api_formats_mode,
|
||||
allowed_models_mode,
|
||||
rate_limit_mode
|
||||
)
|
||||
VALUES (
|
||||
'00000000-0000-0000-0000-000000000001',
|
||||
'Default',
|
||||
'default',
|
||||
'Default unrestricted group for all users',
|
||||
0,
|
||||
'unrestricted',
|
||||
'unrestricted',
|
||||
'unrestricted',
|
||||
'system'
|
||||
)
|
||||
ON CONFLICT (id) DO NOTHING;
|
||||
|
||||
INSERT INTO public.system_configs (id, key, value, description)
|
||||
VALUES (
|
||||
'00000000-0000-0000-0000-000000000002',
|
||||
'default_user_group_id',
|
||||
'"00000000-0000-0000-0000-000000000001"'::json,
|
||||
'Default unrestricted user group'
|
||||
)
|
||||
ON CONFLICT (key) DO NOTHING;
|
||||
|
||||
INSERT INTO public.user_group_members (group_id, user_id)
|
||||
SELECT '00000000-0000-0000-0000-000000000001', id
|
||||
FROM public.users
|
||||
WHERE is_deleted IS FALSE
|
||||
ON CONFLICT (group_id, user_id) DO NOTHING;
|
||||
|
||||
SELECT pg_catalog.set_config('search_path', 'public', true);
|
||||
|
||||
|
||||
|
||||
@@ -293,7 +293,7 @@ mod tests {
|
||||
.await
|
||||
.expect("system config should list")
|
||||
.len(),
|
||||
1
|
||||
2
|
||||
);
|
||||
assert!(backend
|
||||
.delete_system_config_value("feature.local")
|
||||
|
||||
@@ -23,6 +23,8 @@ pub enum ExportDomain {
|
||||
AuthModules,
|
||||
OAuthProviders,
|
||||
UserOAuthLinks,
|
||||
UserGroups,
|
||||
UserGroupMembers,
|
||||
ProxyNodes,
|
||||
SystemConfigs,
|
||||
Wallets,
|
||||
@@ -43,6 +45,8 @@ impl ExportDomain {
|
||||
Self::AuthModules => "auth_modules",
|
||||
Self::OAuthProviders => "oauth_providers",
|
||||
Self::UserOAuthLinks => "user_oauth_links",
|
||||
Self::UserGroups => "user_groups",
|
||||
Self::UserGroupMembers => "user_group_members",
|
||||
Self::ProxyNodes => "proxy_nodes",
|
||||
Self::SystemConfigs => "system_configs",
|
||||
Self::Wallets => "wallets",
|
||||
@@ -265,6 +269,8 @@ pub fn sqlite_core_export_domains() -> Vec<ExportDomain> {
|
||||
ExportDomain::AuthModules,
|
||||
ExportDomain::OAuthProviders,
|
||||
ExportDomain::UserOAuthLinks,
|
||||
ExportDomain::UserGroups,
|
||||
ExportDomain::UserGroupMembers,
|
||||
ExportDomain::ProxyNodes,
|
||||
ExportDomain::SystemConfigs,
|
||||
ExportDomain::Wallets,
|
||||
@@ -367,19 +373,11 @@ pub async fn export_sqlite_jsonl(
|
||||
continue;
|
||||
}
|
||||
let (table_name, id_column) = sqlite_domain_table(domain)?;
|
||||
let sql = format!("SELECT * FROM {table_name} ORDER BY {id_column} ASC");
|
||||
let order_by = export_order_by(domain, id_column);
|
||||
let sql = format!("SELECT * FROM {table_name} ORDER BY {order_by}");
|
||||
let rows = sqlx::query(&sql).fetch_all(pool).await.map_sql_err()?;
|
||||
for row in rows {
|
||||
let id = row
|
||||
.try_get::<Option<String>, _>(id_column)
|
||||
.map_sql_err()?
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"{} export row has null id column '{}'",
|
||||
domain.as_str(),
|
||||
id_column
|
||||
))
|
||||
})?;
|
||||
let id = sqlite_export_row_id(domain, &row, id_column)?;
|
||||
records.push(DataExportRecord::row(domain, id, sqlite_row_payload(&row)?));
|
||||
}
|
||||
}
|
||||
@@ -456,8 +454,10 @@ pub async fn export_postgres_jsonl(
|
||||
continue;
|
||||
}
|
||||
let (table_name, id_column) = postgres_domain_table(domain)?;
|
||||
let export_id_sql = postgres_export_id_sql(domain, id_column);
|
||||
let order_by = export_order_by(domain, id_column);
|
||||
let sql = format!(
|
||||
"SELECT {id_column}::text AS export_id, to_jsonb(t) AS payload FROM {table_name} AS t ORDER BY {id_column} ASC"
|
||||
"SELECT {export_id_sql} AS export_id, to_jsonb(t) AS payload FROM {table_name} AS t ORDER BY {order_by}"
|
||||
);
|
||||
let rows = sqlx::query(&sql).fetch_all(pool).await.map_sql_err()?;
|
||||
for row in rows {
|
||||
@@ -500,6 +500,7 @@ pub async fn import_postgres_plan(
|
||||
continue;
|
||||
}
|
||||
let (table_name, id_column) = postgres_domain_table(*domain)?;
|
||||
let conflict_columns = postgres_conflict_columns(*domain, id_column);
|
||||
let rows = plan.rows(*domain);
|
||||
if rows.is_empty() {
|
||||
continue;
|
||||
@@ -507,7 +508,15 @@ pub async fn import_postgres_plan(
|
||||
let target_columns =
|
||||
postgres_import_columns_cached(pool, &mut column_cache, table_name).await?;
|
||||
for row in rows {
|
||||
import_postgres_row(pool, table_name, id_column, *domain, row, &target_columns).await?;
|
||||
import_postgres_row(
|
||||
pool,
|
||||
table_name,
|
||||
&conflict_columns,
|
||||
*domain,
|
||||
row,
|
||||
&target_columns,
|
||||
)
|
||||
.await?;
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
}
|
||||
@@ -543,19 +552,11 @@ pub async fn export_mysql_jsonl(
|
||||
continue;
|
||||
}
|
||||
let (table_name, id_column) = mysql_domain_table(domain)?;
|
||||
let sql = format!("SELECT * FROM {table_name} ORDER BY {id_column} ASC");
|
||||
let order_by = export_order_by(domain, id_column);
|
||||
let sql = format!("SELECT * FROM {table_name} ORDER BY {order_by}");
|
||||
let rows = sqlx::query(&sql).fetch_all(pool).await.map_sql_err()?;
|
||||
for row in rows {
|
||||
let id = row
|
||||
.try_get::<Option<String>, _>(id_column)
|
||||
.map_sql_err()?
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"{} export row has null id column '{}'",
|
||||
domain.as_str(),
|
||||
id_column
|
||||
))
|
||||
})?;
|
||||
let id = mysql_export_row_id(domain, &row, id_column)?;
|
||||
records.push(DataExportRecord::row(domain, id, mysql_row_payload(&row)?));
|
||||
}
|
||||
}
|
||||
@@ -617,6 +618,8 @@ fn sqlite_domain_table(
|
||||
ExportDomain::AuthModules => Ok(("auth_modules", "id")),
|
||||
ExportDomain::OAuthProviders => Ok(("oauth_providers", "provider_type")),
|
||||
ExportDomain::UserOAuthLinks => Ok(("user_oauth_links", "id")),
|
||||
ExportDomain::UserGroups => Ok(("user_groups", "id")),
|
||||
ExportDomain::UserGroupMembers => Ok(("user_group_members", "group_id")),
|
||||
ExportDomain::ProxyNodes => Ok(("proxy_nodes", "id")),
|
||||
ExportDomain::SystemConfigs => Ok(("system_configs", "id")),
|
||||
ExportDomain::Wallets => Err(DataLayerError::InvalidInput(
|
||||
@@ -630,6 +633,43 @@ fn sqlite_domain_table(
|
||||
}
|
||||
}
|
||||
|
||||
fn export_order_by(domain: ExportDomain, id_column: &str) -> String {
|
||||
if domain == ExportDomain::UserGroupMembers {
|
||||
"group_id ASC, user_id ASC".to_string()
|
||||
} else {
|
||||
format!("{id_column} ASC")
|
||||
}
|
||||
}
|
||||
|
||||
fn sqlite_export_row_id(
|
||||
domain: ExportDomain,
|
||||
row: &sqlx::sqlite::SqliteRow,
|
||||
id_column: &str,
|
||||
) -> Result<String, DataLayerError> {
|
||||
if domain == ExportDomain::UserGroupMembers {
|
||||
let group_id = sqlite_required_export_text(row, "group_id", domain)?;
|
||||
let user_id = sqlite_required_export_text(row, "user_id", domain)?;
|
||||
return Ok(format!("{group_id}:{user_id}"));
|
||||
}
|
||||
sqlite_required_export_text(row, id_column, domain)
|
||||
}
|
||||
|
||||
fn sqlite_required_export_text(
|
||||
row: &sqlx::sqlite::SqliteRow,
|
||||
column: &str,
|
||||
domain: ExportDomain,
|
||||
) -> Result<String, DataLayerError> {
|
||||
row.try_get::<Option<String>, _>(column)
|
||||
.map_sql_err()?
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"{} export row has null id column '{}'",
|
||||
domain.as_str(),
|
||||
column
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
async fn export_sqlite_billing_records(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
records: &mut Vec<DataExportRecord>,
|
||||
@@ -898,6 +938,8 @@ fn postgres_domain_table(
|
||||
ExportDomain::AuthModules => Ok(("public.auth_modules", "id")),
|
||||
ExportDomain::OAuthProviders => Ok(("public.oauth_providers", "provider_type")),
|
||||
ExportDomain::UserOAuthLinks => Ok(("public.user_oauth_links", "id")),
|
||||
ExportDomain::UserGroups => Ok(("public.user_groups", "id")),
|
||||
ExportDomain::UserGroupMembers => Ok(("public.user_group_members", "group_id")),
|
||||
ExportDomain::ProxyNodes => Ok(("public.proxy_nodes", "id")),
|
||||
ExportDomain::SystemConfigs => Ok(("public.system_configs", "id")),
|
||||
ExportDomain::Wallets => Err(DataLayerError::InvalidInput(
|
||||
@@ -912,6 +954,22 @@ fn postgres_domain_table(
|
||||
}
|
||||
}
|
||||
|
||||
fn postgres_export_id_sql(domain: ExportDomain, id_column: &str) -> String {
|
||||
if domain == ExportDomain::UserGroupMembers {
|
||||
"group_id::text || ':' || user_id::text".to_string()
|
||||
} else {
|
||||
format!("{id_column}::text")
|
||||
}
|
||||
}
|
||||
|
||||
fn postgres_conflict_columns(domain: ExportDomain, id_column: &str) -> Vec<&str> {
|
||||
if domain == ExportDomain::UserGroupMembers {
|
||||
vec!["group_id", "user_id"]
|
||||
} else {
|
||||
vec![id_column]
|
||||
}
|
||||
}
|
||||
|
||||
async fn postgres_import_columns_cached(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
cache: &mut BTreeMap<String, PostgresImportColumns>,
|
||||
@@ -1045,7 +1103,7 @@ async fn export_postgres_wallet_records(
|
||||
async fn import_postgres_row(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
table_name: &str,
|
||||
id_column: &str,
|
||||
conflict_columns: &[&str],
|
||||
domain: ExportDomain,
|
||||
row: &ExportRow,
|
||||
target_columns: &PostgresImportColumns,
|
||||
@@ -1060,23 +1118,22 @@ async fn import_postgres_row(
|
||||
.join(", ");
|
||||
let update_sql = columns
|
||||
.iter()
|
||||
.filter(|column| **column != id_column)
|
||||
.filter(|column| !conflict_columns.contains(column))
|
||||
.map(|column| {
|
||||
let quoted = postgres_quote_identifier(column)?;
|
||||
Ok(format!("{quoted} = EXCLUDED.{quoted}"))
|
||||
})
|
||||
.collect::<Result<Vec<_>, DataLayerError>>()?
|
||||
.join(", ");
|
||||
let conflict_target_sql = conflict_columns
|
||||
.iter()
|
||||
.map(|column| postgres_quote_identifier(column))
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
.join(", ");
|
||||
let conflict_sql = if update_sql.is_empty() {
|
||||
format!(
|
||||
"ON CONFLICT ({}) DO NOTHING",
|
||||
postgres_quote_identifier(id_column)?
|
||||
)
|
||||
format!("ON CONFLICT ({conflict_target_sql}) DO NOTHING")
|
||||
} else {
|
||||
format!(
|
||||
"ON CONFLICT ({}) DO UPDATE SET {update_sql}",
|
||||
postgres_quote_identifier(id_column)?
|
||||
)
|
||||
format!("ON CONFLICT ({conflict_target_sql}) DO UPDATE SET {update_sql}")
|
||||
};
|
||||
let sql = format!(
|
||||
"INSERT INTO {table_name} ({column_sql}) SELECT {column_sql} FROM jsonb_populate_record(NULL::{table_name}, $1::jsonb) {conflict_sql}"
|
||||
@@ -1276,7 +1333,7 @@ async fn import_postgres_billing_row(
|
||||
import_postgres_row(
|
||||
pool,
|
||||
table_name,
|
||||
"id",
|
||||
&["id"],
|
||||
ExportDomain::Billing,
|
||||
&ExportRow {
|
||||
id: row.id.clone(),
|
||||
@@ -1309,7 +1366,7 @@ async fn import_postgres_wallet_row(
|
||||
import_postgres_row(
|
||||
pool,
|
||||
table_name,
|
||||
id_column,
|
||||
&[id_column],
|
||||
ExportDomain::Wallets,
|
||||
&ExportRow {
|
||||
id: row.id.clone(),
|
||||
@@ -1382,6 +1439,8 @@ fn mysql_domain_table(
|
||||
ExportDomain::AuthModules => Ok(("auth_modules", "id")),
|
||||
ExportDomain::OAuthProviders => Ok(("oauth_providers", "provider_type")),
|
||||
ExportDomain::UserOAuthLinks => Ok(("user_oauth_links", "id")),
|
||||
ExportDomain::UserGroups => Ok(("user_groups", "id")),
|
||||
ExportDomain::UserGroupMembers => Ok(("user_group_members", "group_id")),
|
||||
ExportDomain::ProxyNodes => Ok(("proxy_nodes", "id")),
|
||||
ExportDomain::SystemConfigs => Ok(("system_configs", "id")),
|
||||
ExportDomain::Wallets => Err(DataLayerError::InvalidInput(
|
||||
@@ -1394,6 +1453,35 @@ fn mysql_domain_table(
|
||||
}
|
||||
}
|
||||
|
||||
fn mysql_export_row_id(
|
||||
domain: ExportDomain,
|
||||
row: &sqlx::mysql::MySqlRow,
|
||||
id_column: &str,
|
||||
) -> Result<String, DataLayerError> {
|
||||
if domain == ExportDomain::UserGroupMembers {
|
||||
let group_id = mysql_required_export_text(row, "group_id", domain)?;
|
||||
let user_id = mysql_required_export_text(row, "user_id", domain)?;
|
||||
return Ok(format!("{group_id}:{user_id}"));
|
||||
}
|
||||
mysql_required_export_text(row, id_column, domain)
|
||||
}
|
||||
|
||||
fn mysql_required_export_text(
|
||||
row: &sqlx::mysql::MySqlRow,
|
||||
column: &str,
|
||||
domain: ExportDomain,
|
||||
) -> Result<String, DataLayerError> {
|
||||
row.try_get::<Option<String>, _>(column)
|
||||
.map_sql_err()?
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"{} export row has null id column '{}'",
|
||||
domain.as_str(),
|
||||
column
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
async fn export_mysql_billing_records(
|
||||
pool: &crate::driver::mysql::MysqlPool,
|
||||
records: &mut Vec<DataExportRecord>,
|
||||
@@ -2094,6 +2182,10 @@ not-json"#,
|
||||
r#"
|
||||
INSERT INTO users (id, email, username, auth_source, created_at, updated_at)
|
||||
VALUES ('user-1', 'owner@example.com', 'owner', 'local', 1, 2);
|
||||
INSERT INTO user_groups (id, name, normalized_name, description, priority, allowed_models, allowed_models_mode, created_at, updated_at)
|
||||
VALUES ('group-1', 'Export Group', 'export group', 'Exported group', 10, '["gpt-test"]', 'specific', 1, 2);
|
||||
INSERT INTO user_group_members (group_id, user_id, created_at)
|
||||
VALUES ('group-1', 'user-1', 1);
|
||||
INSERT INTO api_keys (id, user_id, key_hash, key_encrypted, name, created_at, updated_at)
|
||||
VALUES ('api-key-1', 'user-1', 'hash-1', 'ciphertext-1', 'Default', 1, 2);
|
||||
INSERT INTO providers (id, name, provider_type, created_at, updated_at)
|
||||
@@ -2136,6 +2228,16 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
import_plan.rows(ExportDomain::Users)[0].payload["email"],
|
||||
"owner@example.com"
|
||||
);
|
||||
assert!(import_plan
|
||||
.rows(ExportDomain::UserGroups)
|
||||
.iter()
|
||||
.any(|row| row.id == "group-1" && row.payload["name"] == "Export Group"));
|
||||
assert!(import_plan
|
||||
.rows(ExportDomain::UserGroupMembers)
|
||||
.iter()
|
||||
.any(|row| row.id == "group-1:user-1"
|
||||
&& row.payload["group_id"] == "group-1"
|
||||
&& row.payload["user_id"] == "user-1"));
|
||||
assert_eq!(
|
||||
import_plan.rows(ExportDomain::ApiKeys)[0].payload["key_encrypted"],
|
||||
"ciphertext-1"
|
||||
@@ -2166,7 +2268,7 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
let imported = import_sqlite_jsonl(&target_pool, &encoded)
|
||||
.await
|
||||
.expect("sqlite import should load exported rows");
|
||||
assert_eq!(imported, 12);
|
||||
assert_eq!(imported, 16);
|
||||
|
||||
let imported_api_key = sqlx::query_as::<_, (String,)>(
|
||||
"SELECT key_encrypted FROM api_keys WHERE id = 'api-key-1'",
|
||||
@@ -2184,6 +2286,15 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
.expect("imported usage should load");
|
||||
assert_eq!(imported_usage.0, "request-1");
|
||||
|
||||
let imported_group_member = sqlx::query_as::<_, (String, String)>(
|
||||
"SELECT group_id, user_id FROM user_group_members WHERE group_id = 'group-1' AND user_id = 'user-1'",
|
||||
)
|
||||
.fetch_one(&target_pool)
|
||||
.await
|
||||
.expect("imported user group member should load");
|
||||
assert_eq!(imported_group_member.0, "group-1");
|
||||
assert_eq!(imported_group_member.1, "user-1");
|
||||
|
||||
let imported_billing_rule = sqlx::query_as::<_, (String,)>(
|
||||
"SELECT expression FROM billing_rules WHERE id = 'billing-rule-1'",
|
||||
)
|
||||
@@ -2217,7 +2328,7 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
let imported = import_postgres_jsonl(&postgres_pool, &encoded)
|
||||
.await
|
||||
.expect("postgres import should load exported rows");
|
||||
assert_eq!(imported, 12);
|
||||
assert_eq!(imported, 16);
|
||||
|
||||
let imported_api_key = sqlx::query_as::<_, (String,)>(
|
||||
"SELECT key_encrypted FROM api_keys WHERE id = 'api-key-1'",
|
||||
@@ -2273,6 +2384,7 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
let config_key = format!("export.config.{suffix}");
|
||||
let wallet_id = format!("export-wallet-{suffix}");
|
||||
let request_id = format!("export-request-{suffix}");
|
||||
let group_id = format!("export-group-{suffix}");
|
||||
|
||||
sqlx::query(
|
||||
"INSERT INTO users (id, email, username, auth_source, email_verified, created_at, updated_at) VALUES ($1, $2, $3, 'local', TRUE, to_timestamp(1), to_timestamp(2))",
|
||||
@@ -2283,6 +2395,23 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("user should seed");
|
||||
sqlx::query(
|
||||
"INSERT INTO user_groups (id, name, normalized_name, priority, allowed_models, allowed_models_mode, created_at, updated_at) VALUES ($1, $2, $3, 10, '[\"provider-model\"]', 'specific', to_timestamp(1), to_timestamp(2))",
|
||||
)
|
||||
.bind(&group_id)
|
||||
.bind(format!("Export Group {suffix}"))
|
||||
.bind(format!("export group {suffix}"))
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("user group should seed");
|
||||
sqlx::query(
|
||||
"INSERT INTO user_group_members (group_id, user_id, created_at) VALUES ($1, $2, to_timestamp(1))",
|
||||
)
|
||||
.bind(&group_id)
|
||||
.bind(&user_id)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("user group member should seed");
|
||||
sqlx::query(
|
||||
"INSERT INTO api_keys (id, user_id, key_hash, key_encrypted, name, created_at, updated_at) VALUES ($1, $2, $3, 'ciphertext-1', 'Default', to_timestamp(1), to_timestamp(2))",
|
||||
)
|
||||
@@ -2389,6 +2518,14 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
.rows(ExportDomain::Users)
|
||||
.iter()
|
||||
.any(|row| row.id == user_id));
|
||||
assert!(import_plan
|
||||
.rows(ExportDomain::UserGroups)
|
||||
.iter()
|
||||
.any(|row| row.id == group_id));
|
||||
assert!(import_plan
|
||||
.rows(ExportDomain::UserGroupMembers)
|
||||
.iter()
|
||||
.any(|row| row.id == format!("{group_id}:{user_id}")));
|
||||
assert!(import_plan
|
||||
.rows(ExportDomain::ApiKeys)
|
||||
.iter()
|
||||
@@ -2432,6 +2569,16 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
.await
|
||||
.expect("imported sqlite api key should load");
|
||||
assert_eq!(imported_api_key.0, "ciphertext-1");
|
||||
let imported_group_member = sqlx::query_as::<_, (String, String)>(
|
||||
"SELECT group_id, user_id FROM user_group_members WHERE group_id = ? AND user_id = ?",
|
||||
)
|
||||
.bind(&group_id)
|
||||
.bind(&user_id)
|
||||
.fetch_one(&target_pool)
|
||||
.await
|
||||
.expect("imported sqlite user group member should load");
|
||||
assert_eq!(imported_group_member.0, group_id);
|
||||
assert_eq!(imported_group_member.1, user_id);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
@@ -2466,6 +2613,7 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
let config_id = format!("export-config-{suffix}");
|
||||
let wallet_id = format!("export-wallet-{suffix}");
|
||||
let request_id = format!("export-request-{suffix}");
|
||||
let group_id = format!("export-group-{suffix}");
|
||||
|
||||
sqlx::query(
|
||||
"INSERT INTO users (id, email, username, auth_source, created_at, updated_at) VALUES (?, ?, ?, 'local', 1, 2)",
|
||||
@@ -2476,6 +2624,23 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("user should seed");
|
||||
sqlx::query(
|
||||
"INSERT INTO user_groups (id, name, normalized_name, priority, allowed_models, allowed_models_mode, created_at, updated_at) VALUES (?, ?, ?, 10, '[\"provider-model\"]', 'specific', 1, 2)",
|
||||
)
|
||||
.bind(&group_id)
|
||||
.bind(format!("Export Group {suffix}"))
|
||||
.bind(format!("export group {suffix}"))
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("user group should seed");
|
||||
sqlx::query(
|
||||
"INSERT INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, 1)",
|
||||
)
|
||||
.bind(&group_id)
|
||||
.bind(&user_id)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("user group member should seed");
|
||||
sqlx::query(
|
||||
"INSERT INTO api_keys (id, user_id, key_hash, key_encrypted, name, created_at, updated_at) VALUES (?, ?, ?, 'ciphertext-1', 'Default', 1, 2)",
|
||||
)
|
||||
@@ -2566,6 +2731,14 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
.rows(ExportDomain::Users)
|
||||
.iter()
|
||||
.any(|row| row.id == user_id));
|
||||
assert!(import_plan
|
||||
.rows(ExportDomain::UserGroups)
|
||||
.iter()
|
||||
.any(|row| row.id == group_id));
|
||||
assert!(import_plan
|
||||
.rows(ExportDomain::UserGroupMembers)
|
||||
.iter()
|
||||
.any(|row| row.id == format!("{group_id}:{user_id}")));
|
||||
assert!(import_plan
|
||||
.rows(ExportDomain::ApiKeys)
|
||||
.iter()
|
||||
@@ -2585,6 +2758,8 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
&pool,
|
||||
vec![
|
||||
ExportDomain::Users,
|
||||
ExportDomain::UserGroups,
|
||||
ExportDomain::UserGroupMembers,
|
||||
ExportDomain::ApiKeys,
|
||||
ExportDomain::ProviderKeys,
|
||||
ExportDomain::Usage,
|
||||
@@ -2596,7 +2771,7 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
let imported = import_mysql_jsonl(&pool, &selected_export)
|
||||
.await
|
||||
.expect("mysql import should be idempotent");
|
||||
assert!(imported >= 4);
|
||||
assert!(imported >= 6);
|
||||
|
||||
let imported_api_key =
|
||||
sqlx::query_as::<_, (String,)>("SELECT key_encrypted FROM api_keys WHERE id = ?")
|
||||
|
||||
@@ -356,6 +356,41 @@ fn memory_group_members(
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn memory_export_row_from_auth_user(
|
||||
repository: &InMemoryUserReadRepository,
|
||||
user: &StoredUserAuthRecord,
|
||||
) -> Result<StoredUserExportRow, DataLayerError> {
|
||||
let model_capability_settings = repository
|
||||
.model_settings_by_user_id
|
||||
.read()
|
||||
.expect("user repository lock")
|
||||
.get(&user.id)
|
||||
.cloned();
|
||||
StoredUserExportRow::new(
|
||||
user.id.clone(),
|
||||
user.email.clone(),
|
||||
user.email_verified,
|
||||
user.username.clone(),
|
||||
user.password_hash.clone(),
|
||||
user.role.clone(),
|
||||
user.auth_source.clone(),
|
||||
user.allowed_providers.clone().map(serde_json::Value::from),
|
||||
user.allowed_api_formats
|
||||
.clone()
|
||||
.map(serde_json::Value::from),
|
||||
user.allowed_models.clone().map(serde_json::Value::from),
|
||||
None,
|
||||
model_capability_settings,
|
||||
user.is_active,
|
||||
)?
|
||||
.with_policy_modes(
|
||||
user.allowed_providers_mode.clone(),
|
||||
user.allowed_api_formats_mode.clone(),
|
||||
user.allowed_models_mode.clone(),
|
||||
"system".to_string(),
|
||||
)
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl UserReadRepository for InMemoryUserReadRepository {
|
||||
async fn list_users_by_ids(
|
||||
@@ -395,22 +430,36 @@ impl UserReadRepository for InMemoryUserReadRepository {
|
||||
async fn list_non_admin_export_users(
|
||||
&self,
|
||||
) -> Result<Vec<StoredUserExportRow>, DataLayerError> {
|
||||
let rows = self.export_rows.read().expect("user repository lock");
|
||||
if !rows.is_empty() {
|
||||
return Ok(rows
|
||||
.iter()
|
||||
.filter(|row| !row.role.eq_ignore_ascii_case("admin"))
|
||||
.cloned()
|
||||
.collect());
|
||||
}
|
||||
Ok(self
|
||||
.export_rows
|
||||
.auth_by_id
|
||||
.read()
|
||||
.expect("user repository lock")
|
||||
.iter()
|
||||
.filter(|row| !row.role.eq_ignore_ascii_case("admin"))
|
||||
.cloned()
|
||||
.collect())
|
||||
.filter(|(_, user)| !user.role.eq_ignore_ascii_case("admin"))
|
||||
.map(|(_, user)| memory_export_row_from_auth_user(self, user))
|
||||
.collect::<Result<Vec<_>, _>>()?)
|
||||
}
|
||||
|
||||
async fn list_export_users(&self) -> Result<Vec<StoredUserExportRow>, DataLayerError> {
|
||||
let rows = self.export_rows.read().expect("user repository lock");
|
||||
if !rows.is_empty() {
|
||||
return Ok(rows.clone());
|
||||
}
|
||||
Ok(self
|
||||
.export_rows
|
||||
.auth_by_id
|
||||
.read()
|
||||
.expect("user repository lock")
|
||||
.clone())
|
||||
.values()
|
||||
.map(|user| memory_export_row_from_auth_user(self, user))
|
||||
.collect::<Result<Vec<_>, _>>()?)
|
||||
}
|
||||
|
||||
async fn list_export_users_page(
|
||||
@@ -500,10 +549,8 @@ impl UserReadRepository for InMemoryUserReadRepository {
|
||||
.cloned()
|
||||
.collect::<Vec<_>>();
|
||||
groups.sort_by(|left, right| {
|
||||
right
|
||||
.priority
|
||||
.cmp(&left.priority)
|
||||
.then_with(|| left.name.cmp(&right.name))
|
||||
left.name
|
||||
.cmp(&right.name)
|
||||
.then_with(|| left.id.cmp(&right.id))
|
||||
});
|
||||
Ok(groups)
|
||||
@@ -687,7 +734,6 @@ impl UserReadRepository for InMemoryUserReadRepository {
|
||||
memberships.sort_by(|left, right| {
|
||||
left.user_id
|
||||
.cmp(&right.user_id)
|
||||
.then_with(|| right.group_priority.cmp(&left.group_priority))
|
||||
.then_with(|| left.group_name.cmp(&right.group_name))
|
||||
.then_with(|| left.group_id.cmp(&right.group_id))
|
||||
});
|
||||
|
||||
@@ -370,7 +370,7 @@ WHERE is_deleted = 0
|
||||
|
||||
async fn list_user_groups(&self) -> Result<Vec<StoredUserGroup>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<MySql>::new(USER_GROUP_COLUMNS);
|
||||
builder.push(" ORDER BY priority DESC, name ASC, id ASC");
|
||||
builder.push(" ORDER BY name ASC, id ASC");
|
||||
self.fetch_group_rows(builder).await
|
||||
}
|
||||
|
||||
@@ -401,7 +401,7 @@ WHERE is_deleted = 0
|
||||
separated.push_bind(group_id);
|
||||
}
|
||||
}
|
||||
builder.push(") ORDER BY priority DESC, name ASC, id ASC");
|
||||
builder.push(") ORDER BY name ASC, id ASC");
|
||||
self.fetch_group_rows(builder).await
|
||||
}
|
||||
|
||||
@@ -568,7 +568,7 @@ WHERE id = ?
|
||||
builder
|
||||
.push(" WHERE id IN (SELECT group_id FROM user_group_members WHERE user_id = ")
|
||||
.push_bind(user_id)
|
||||
.push(") ORDER BY priority DESC, name ASC, id ASC");
|
||||
.push(") ORDER BY name ASC, id ASC");
|
||||
self.fetch_group_rows(builder).await
|
||||
}
|
||||
|
||||
@@ -598,7 +598,9 @@ WHERE user_group_members.user_id IN (
|
||||
separated.push_bind(user_id);
|
||||
}
|
||||
}
|
||||
builder.push(") ORDER BY user_group_members.user_id ASC, user_groups.priority DESC, user_groups.name ASC, user_groups.id ASC");
|
||||
builder.push(
|
||||
") ORDER BY user_group_members.user_id ASC, user_groups.name ASC, user_groups.id ASC",
|
||||
);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_user_group_membership_row).collect()
|
||||
}
|
||||
|
||||
@@ -684,7 +684,7 @@ impl SqlxUserReadRepository {
|
||||
|
||||
pub async fn list_user_groups(&self) -> Result<Vec<StoredUserGroup>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<Postgres>::new(USER_GROUP_COLUMNS);
|
||||
builder.push(" ORDER BY priority DESC, name ASC, id ASC");
|
||||
builder.push(" ORDER BY name ASC, id ASC");
|
||||
collect_query_rows(builder.build().fetch(&self.pool), map_user_group_row).await
|
||||
}
|
||||
|
||||
@@ -720,7 +720,7 @@ impl SqlxUserReadRepository {
|
||||
separated.push_bind(group_id);
|
||||
}
|
||||
}
|
||||
builder.push(") ORDER BY priority DESC, name ASC, id ASC");
|
||||
builder.push(") ORDER BY name ASC, id ASC");
|
||||
collect_query_rows(builder.build().fetch(&self.pool), map_user_group_row).await
|
||||
}
|
||||
|
||||
@@ -872,7 +872,7 @@ WHERE id = $1
|
||||
builder
|
||||
.push(" WHERE id IN (SELECT group_id FROM user_group_members WHERE user_id = ")
|
||||
.push_bind(user_id)
|
||||
.push(") ORDER BY priority DESC, name ASC, id ASC");
|
||||
.push(") ORDER BY name ASC, id ASC");
|
||||
collect_query_rows(builder.build().fetch(&self.pool), map_user_group_row).await
|
||||
}
|
||||
|
||||
@@ -902,7 +902,9 @@ WHERE user_group_members.user_id IN (
|
||||
separated.push_bind(user_id);
|
||||
}
|
||||
}
|
||||
builder.push(") ORDER BY user_group_members.user_id ASC, user_groups.priority DESC, user_groups.name ASC, user_groups.id ASC");
|
||||
builder.push(
|
||||
") ORDER BY user_group_members.user_id ASC, user_groups.name ASC, user_groups.id ASC",
|
||||
);
|
||||
collect_query_rows(
|
||||
builder.build().fetch(&self.pool),
|
||||
map_user_group_membership_row,
|
||||
|
||||
@@ -370,7 +370,7 @@ WHERE is_deleted = 0
|
||||
|
||||
async fn list_user_groups(&self) -> Result<Vec<StoredUserGroup>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(USER_GROUP_COLUMNS);
|
||||
builder.push(" ORDER BY priority DESC, name ASC, id ASC");
|
||||
builder.push(" ORDER BY name ASC, id ASC");
|
||||
self.fetch_group_rows(builder).await
|
||||
}
|
||||
|
||||
@@ -401,7 +401,7 @@ WHERE is_deleted = 0
|
||||
separated.push_bind(group_id);
|
||||
}
|
||||
}
|
||||
builder.push(") ORDER BY priority DESC, name ASC, id ASC");
|
||||
builder.push(") ORDER BY name ASC, id ASC");
|
||||
self.fetch_group_rows(builder).await
|
||||
}
|
||||
|
||||
@@ -568,7 +568,7 @@ WHERE id = ?
|
||||
builder
|
||||
.push(" WHERE id IN (SELECT group_id FROM user_group_members WHERE user_id = ")
|
||||
.push_bind(user_id)
|
||||
.push(") ORDER BY priority DESC, name ASC, id ASC");
|
||||
.push(") ORDER BY name ASC, id ASC");
|
||||
self.fetch_group_rows(builder).await
|
||||
}
|
||||
|
||||
@@ -598,7 +598,9 @@ WHERE user_group_members.user_id IN (
|
||||
separated.push_bind(user_id);
|
||||
}
|
||||
}
|
||||
builder.push(") ORDER BY user_group_members.user_id ASC, user_groups.priority DESC, user_groups.name ASC, user_groups.id ASC");
|
||||
builder.push(
|
||||
") ORDER BY user_group_members.user_id ASC, user_groups.name ASC, user_groups.id ASC",
|
||||
);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_user_group_membership_row).collect()
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user