feat: add user groups and inherited access policies

This commit is contained in:
Entropy.Xu
2026-05-09 21:47:33 +08:00
parent 4a64d078f3
commit 3a814f3d1f
49 changed files with 6381 additions and 250 deletions

View File

@@ -0,0 +1,53 @@
ALTER TABLE users
ADD COLUMN allowed_providers_mode VARCHAR(32) NOT NULL DEFAULT 'unrestricted',
ADD COLUMN allowed_api_formats_mode VARCHAR(32) NOT NULL DEFAULT 'unrestricted',
ADD COLUMN allowed_models_mode VARCHAR(32) NOT NULL DEFAULT 'unrestricted',
ADD COLUMN rate_limit_mode VARCHAR(32) NOT NULL DEFAULT 'system';
UPDATE users
SET allowed_providers_mode = CASE WHEN allowed_providers IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_providers_mode = 'unrestricted';
UPDATE users
SET allowed_api_formats_mode = CASE WHEN allowed_api_formats IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_api_formats_mode = 'unrestricted';
UPDATE users
SET allowed_models_mode = CASE WHEN allowed_models IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_models_mode = 'unrestricted';
UPDATE users
SET rate_limit_mode = CASE WHEN rate_limit IS NULL THEN 'system' ELSE 'custom' END
WHERE rate_limit_mode = 'system';
CREATE TABLE IF NOT EXISTS user_groups (
id VARCHAR(64) PRIMARY KEY,
name VARCHAR(100) NOT NULL,
normalized_name VARCHAR(100) NOT NULL,
description TEXT,
priority INT NOT NULL DEFAULT 0,
allowed_providers TEXT,
allowed_providers_mode VARCHAR(32) NOT NULL DEFAULT 'inherit',
allowed_api_formats TEXT,
allowed_api_formats_mode VARCHAR(32) NOT NULL DEFAULT 'inherit',
allowed_models TEXT,
allowed_models_mode VARCHAR(32) NOT NULL DEFAULT 'inherit',
rate_limit INT,
rate_limit_mode VARCHAR(32) NOT NULL DEFAULT 'inherit',
created_at BIGINT NOT NULL,
updated_at BIGINT NOT NULL,
UNIQUE KEY user_groups_normalized_name_key (normalized_name),
KEY user_groups_priority_name_idx (priority, name, id)
);
CREATE TABLE IF NOT EXISTS user_group_members (
group_id VARCHAR(64) NOT NULL,
user_id VARCHAR(64) NOT NULL,
created_at BIGINT NOT NULL,
PRIMARY KEY (group_id, user_id),
KEY user_group_members_user_id_idx (user_id),
CONSTRAINT user_group_members_group_id_fk
FOREIGN KEY (group_id) REFERENCES user_groups(id) ON DELETE CASCADE,
CONSTRAINT user_group_members_user_id_fk
FOREIGN KEY (user_id) REFERENCES users(id) ON DELETE CASCADE
);

View File

@@ -0,0 +1,60 @@
ALTER TABLE public.users
ADD COLUMN IF NOT EXISTS allowed_providers_mode text DEFAULT 'unrestricted' NOT NULL,
ADD COLUMN IF NOT EXISTS allowed_api_formats_mode text DEFAULT 'unrestricted' NOT NULL,
ADD COLUMN IF NOT EXISTS allowed_models_mode text DEFAULT 'unrestricted' NOT NULL,
ADD COLUMN IF NOT EXISTS rate_limit_mode text DEFAULT 'system' NOT NULL;
UPDATE public.users
SET allowed_providers_mode = CASE WHEN allowed_providers IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_providers_mode = 'unrestricted';
UPDATE public.users
SET allowed_api_formats_mode = CASE WHEN allowed_api_formats IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_api_formats_mode = 'unrestricted';
UPDATE public.users
SET allowed_models_mode = CASE WHEN allowed_models IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_models_mode = 'unrestricted';
UPDATE public.users
SET rate_limit_mode = CASE WHEN rate_limit IS NULL THEN 'system' ELSE 'custom' END
WHERE rate_limit_mode = 'system';
CREATE TABLE IF NOT EXISTS public.user_groups (
id character varying(36) PRIMARY KEY,
name character varying(100) NOT NULL,
normalized_name character varying(100) NOT NULL UNIQUE,
description text,
priority integer DEFAULT 0 NOT NULL,
allowed_providers json,
allowed_providers_mode text DEFAULT 'inherit' NOT NULL,
allowed_api_formats json,
allowed_api_formats_mode text DEFAULT 'inherit' NOT NULL,
allowed_models json,
allowed_models_mode text DEFAULT 'inherit' NOT NULL,
rate_limit integer,
rate_limit_mode text DEFAULT 'inherit' NOT NULL,
created_at timestamp with time zone DEFAULT now() NOT NULL,
updated_at timestamp with time zone DEFAULT now() NOT NULL,
CONSTRAINT user_groups_allowed_providers_mode_check
CHECK (allowed_providers_mode IN ('inherit', 'unrestricted', 'specific', 'deny_all')),
CONSTRAINT user_groups_allowed_api_formats_mode_check
CHECK (allowed_api_formats_mode IN ('inherit', 'unrestricted', 'specific', 'deny_all')),
CONSTRAINT user_groups_allowed_models_mode_check
CHECK (allowed_models_mode IN ('inherit', 'unrestricted', 'specific', 'deny_all')),
CONSTRAINT user_groups_rate_limit_mode_check
CHECK (rate_limit_mode IN ('inherit', 'system', 'custom'))
);
CREATE TABLE IF NOT EXISTS public.user_group_members (
group_id character varying(36) NOT NULL REFERENCES public.user_groups(id) ON DELETE CASCADE,
user_id character varying(36) NOT NULL REFERENCES public.users(id) ON DELETE CASCADE,
created_at timestamp with time zone DEFAULT now() NOT NULL,
PRIMARY KEY (group_id, user_id)
);
CREATE INDEX IF NOT EXISTS user_group_members_user_id_idx
ON public.user_group_members (user_id);
CREATE INDEX IF NOT EXISTS user_groups_priority_name_idx
ON public.user_groups (priority DESC, name ASC, id ASC);

View File

@@ -0,0 +1,51 @@
ALTER TABLE users ADD COLUMN allowed_providers_mode TEXT NOT NULL DEFAULT 'unrestricted';
ALTER TABLE users ADD COLUMN allowed_api_formats_mode TEXT NOT NULL DEFAULT 'unrestricted';
ALTER TABLE users ADD COLUMN allowed_models_mode TEXT NOT NULL DEFAULT 'unrestricted';
ALTER TABLE users ADD COLUMN rate_limit_mode TEXT NOT NULL DEFAULT 'system';
UPDATE users
SET allowed_providers_mode = CASE WHEN allowed_providers IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_providers_mode = 'unrestricted';
UPDATE users
SET allowed_api_formats_mode = CASE WHEN allowed_api_formats IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_api_formats_mode = 'unrestricted';
UPDATE users
SET allowed_models_mode = CASE WHEN allowed_models IS NULL THEN 'unrestricted' ELSE 'specific' END
WHERE allowed_models_mode = 'unrestricted';
UPDATE users
SET rate_limit_mode = CASE WHEN rate_limit IS NULL THEN 'system' ELSE 'custom' END
WHERE rate_limit_mode = 'system';
CREATE TABLE IF NOT EXISTS user_groups (
id TEXT PRIMARY KEY,
name TEXT NOT NULL,
normalized_name TEXT NOT NULL UNIQUE,
description TEXT,
priority INTEGER NOT NULL DEFAULT 0,
allowed_providers TEXT,
allowed_providers_mode TEXT NOT NULL DEFAULT 'inherit',
allowed_api_formats TEXT,
allowed_api_formats_mode TEXT NOT NULL DEFAULT 'inherit',
allowed_models TEXT,
allowed_models_mode TEXT NOT NULL DEFAULT 'inherit',
rate_limit INTEGER,
rate_limit_mode TEXT NOT NULL DEFAULT 'inherit',
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL
);
CREATE TABLE IF NOT EXISTS user_group_members (
group_id TEXT NOT NULL REFERENCES user_groups(id) ON DELETE CASCADE,
user_id TEXT NOT NULL REFERENCES users(id) ON DELETE CASCADE,
created_at INTEGER NOT NULL,
PRIMARY KEY (group_id, user_id)
);
CREATE INDEX IF NOT EXISTS user_group_members_user_id_idx
ON user_group_members (user_id);
CREATE INDEX IF NOT EXISTS user_groups_priority_name_idx
ON user_groups (priority DESC, name ASC, id ASC);

View File

@@ -1238,8 +1238,11 @@ CREATE TABLE IF NOT EXISTS public.users (
password_hash character varying(255),
role public.userrole DEFAULT 'user'::public.userrole NOT NULL,
allowed_providers json,
allowed_providers_mode text DEFAULT 'unrestricted'::text NOT NULL,
allowed_api_formats json,
allowed_api_formats_mode text DEFAULT 'unrestricted'::text NOT NULL,
allowed_models json,
allowed_models_mode text DEFAULT 'unrestricted'::text NOT NULL,
model_capability_settings json,
is_active boolean DEFAULT true NOT NULL,
is_deleted boolean DEFAULT false NOT NULL,
@@ -1251,11 +1254,48 @@ CREATE TABLE IF NOT EXISTS public.users (
ldap_username character varying(255),
email_verified boolean NOT NULL,
rate_limit integer,
rate_limit_mode text DEFAULT 'system'::text NOT NULL,
metadata json
);
--
-- Name: user_groups; Type: TABLE; Schema: public; Owner: -
--
CREATE TABLE IF NOT EXISTS public.user_groups (
id character varying(36) NOT NULL,
name character varying(100) NOT NULL,
normalized_name character varying(100) NOT NULL,
description text,
priority integer DEFAULT 0 NOT NULL,
allowed_providers json,
allowed_providers_mode text DEFAULT 'inherit'::text NOT NULL,
allowed_api_formats json,
allowed_api_formats_mode text DEFAULT 'inherit'::text NOT NULL,
allowed_models json,
allowed_models_mode text DEFAULT 'inherit'::text NOT NULL,
rate_limit integer,
rate_limit_mode text DEFAULT 'inherit'::text NOT NULL,
created_at timestamp with time zone DEFAULT now() NOT NULL,
updated_at timestamp with time zone DEFAULT now() NOT NULL
);
--
-- Name: user_group_members; Type: TABLE; Schema: public; Owner: -
--
CREATE TABLE IF NOT EXISTS public.user_group_members (
group_id character varying(36) NOT NULL,
user_id character varying(36) NOT NULL,
created_at timestamp with time zone DEFAULT now() NOT NULL
);
--
-- Name: video_tasks; Type: TABLE; Schema: public; Owner: -
--

View File

@@ -972,6 +972,111 @@ END $mig$;
--
-- Name: user_group_members user_group_members_pkey; Type: CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_group_members
ADD CONSTRAINT user_group_members_pkey PRIMARY KEY (group_id, user_id);
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_groups user_groups_allowed_api_formats_mode_check; Type: CHECK CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_groups
ADD CONSTRAINT user_groups_allowed_api_formats_mode_check CHECK (allowed_api_formats_mode IN ('inherit', 'unrestricted', 'specific', 'deny_all'));
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_groups user_groups_allowed_models_mode_check; Type: CHECK CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_groups
ADD CONSTRAINT user_groups_allowed_models_mode_check CHECK (allowed_models_mode IN ('inherit', 'unrestricted', 'specific', 'deny_all'));
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_groups user_groups_allowed_providers_mode_check; Type: CHECK CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_groups
ADD CONSTRAINT user_groups_allowed_providers_mode_check CHECK (allowed_providers_mode IN ('inherit', 'unrestricted', 'specific', 'deny_all'));
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_groups user_groups_normalized_name_key; Type: CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_groups
ADD CONSTRAINT user_groups_normalized_name_key UNIQUE (normalized_name);
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_groups user_groups_pkey; Type: CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_groups
ADD CONSTRAINT user_groups_pkey PRIMARY KEY (id);
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_groups user_groups_rate_limit_mode_check; Type: CHECK CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_groups
ADD CONSTRAINT user_groups_rate_limit_mode_check CHECK (rate_limit_mode IN ('inherit', 'system', 'custom'));
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_oauth_links user_oauth_links_pkey; Type: CONSTRAINT; Schema: public; Owner: -
--

View File

@@ -1021,6 +1021,22 @@ CREATE INDEX IF NOT EXISTS ix_system_configs_id ON public.system_configs USING b
--
-- Name: user_group_members_user_id_idx; Type: INDEX; Schema: public; Owner: -
--
CREATE INDEX IF NOT EXISTS user_group_members_user_id_idx ON public.user_group_members USING btree (user_id);
--
-- Name: user_groups_priority_name_idx; Type: INDEX; Schema: public; Owner: -
--
CREATE INDEX IF NOT EXISTS user_groups_priority_name_idx ON public.user_groups USING btree (priority DESC, name, id);
--
-- Name: ix_usage_created_at; Type: INDEX; Schema: public; Owner: -
--

View File

@@ -612,6 +612,36 @@ END $mig$;
--
-- Name: user_group_members user_group_members_group_id_fk; Type: FK CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_group_members
ADD CONSTRAINT user_group_members_group_id_fk FOREIGN KEY (group_id) REFERENCES public.user_groups(id) ON DELETE CASCADE;
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_group_members user_group_members_user_id_fk; Type: FK CONSTRAINT; Schema: public; Owner: -
--
DO $mig$ BEGIN
ALTER TABLE ONLY public.user_group_members
ADD CONSTRAINT user_group_members_user_id_fk FOREIGN KEY (user_id) REFERENCES public.users(id) ON DELETE CASCADE;
EXCEPTION
WHEN duplicate_object THEN NULL;
WHEN duplicate_table THEN NULL;
WHEN invalid_table_definition THEN NULL;
END $mig$;
--
-- Name: user_oauth_links user_oauth_links_provider_type_fkey; Type: FK CONSTRAINT; Schema: public; Owner: -
--

View File

@@ -13,10 +13,14 @@ CREATE TABLE IF NOT EXISTS users (
`is_active` TINYINT(1) NOT NULL DEFAULT 1,
`is_deleted` TINYINT(1) NOT NULL DEFAULT 0,
`allowed_models` JSON,
`allowed_models_mode` VARCHAR(32) NOT NULL DEFAULT 'unrestricted',
`allowed_providers` JSON,
`allowed_providers_mode` VARCHAR(32) NOT NULL DEFAULT 'unrestricted',
`allowed_api_formats` JSON,
`allowed_api_formats_mode` VARCHAR(32) NOT NULL DEFAULT 'unrestricted',
`model_capability_settings` JSON,
`rate_limit` INT,
`rate_limit_mode` VARCHAR(32) NOT NULL DEFAULT 'system',
`metadata` JSON,
`created_at` BIGINT NOT NULL,
`updated_at` BIGINT NOT NULL,
@@ -28,6 +32,35 @@ CREATE TABLE IF NOT EXISTS users (
UNIQUE KEY users_username_key (`username`)
);
CREATE TABLE IF NOT EXISTS user_groups (
`id` VARCHAR(64) NOT NULL,
`name` VARCHAR(100) NOT NULL,
`normalized_name` VARCHAR(100) NOT NULL,
`description` LONGTEXT,
`priority` INT NOT NULL DEFAULT 0,
`allowed_providers` JSON,
`allowed_providers_mode` VARCHAR(32) NOT NULL DEFAULT 'inherit',
`allowed_api_formats` JSON,
`allowed_api_formats_mode` VARCHAR(32) NOT NULL DEFAULT 'inherit',
`allowed_models` JSON,
`allowed_models_mode` VARCHAR(32) NOT NULL DEFAULT 'inherit',
`rate_limit` INT,
`rate_limit_mode` VARCHAR(32) NOT NULL DEFAULT 'inherit',
`created_at` BIGINT NOT NULL,
`updated_at` BIGINT NOT NULL,
PRIMARY KEY (`id`),
UNIQUE KEY user_groups_normalized_name_key (`normalized_name`),
KEY user_groups_priority_name_idx (`priority`, `name`, `id`)
);
CREATE TABLE IF NOT EXISTS user_group_members (
`group_id` VARCHAR(64) NOT NULL,
`user_id` VARCHAR(64) NOT NULL,
`created_at` BIGINT NOT NULL,
PRIMARY KEY (`group_id`, `user_id`),
KEY user_group_members_user_id_idx (`user_id`)
);
CREATE TABLE IF NOT EXISTS api_keys (
`id` VARCHAR(64) NOT NULL,
`user_id` VARCHAR(64) NOT NULL,

View File

@@ -13,10 +13,14 @@ CREATE TABLE IF NOT EXISTS public.users (
is_active boolean DEFAULT true NOT NULL,
is_deleted boolean DEFAULT false NOT NULL,
allowed_models jsonb,
allowed_models_mode character varying(32) DEFAULT 'unrestricted' NOT NULL,
allowed_providers jsonb,
allowed_providers_mode character varying(32) DEFAULT 'unrestricted' NOT NULL,
allowed_api_formats jsonb,
allowed_api_formats_mode character varying(32) DEFAULT 'unrestricted' NOT NULL,
model_capability_settings jsonb,
rate_limit integer,
rate_limit_mode character varying(32) DEFAULT 'system' NOT NULL,
metadata jsonb,
created_at bigint NOT NULL,
updated_at bigint NOT NULL,
@@ -29,6 +33,37 @@ ALTER TABLE ONLY public.users ADD CONSTRAINT users_pkey PRIMARY KEY (id);
ALTER TABLE ONLY public.users ADD CONSTRAINT users_email_key UNIQUE (email);
ALTER TABLE ONLY public.users ADD CONSTRAINT users_username_key UNIQUE (username);
CREATE TABLE IF NOT EXISTS public.user_groups (
id character varying(64) NOT NULL,
name character varying(100) NOT NULL,
normalized_name character varying(100) NOT NULL,
description text,
priority integer DEFAULT 0 NOT NULL,
allowed_providers jsonb,
allowed_providers_mode character varying(32) DEFAULT 'inherit' NOT NULL,
allowed_api_formats jsonb,
allowed_api_formats_mode character varying(32) DEFAULT 'inherit' NOT NULL,
allowed_models jsonb,
allowed_models_mode character varying(32) DEFAULT 'inherit' NOT NULL,
rate_limit integer,
rate_limit_mode character varying(32) DEFAULT 'inherit' NOT NULL,
created_at bigint NOT NULL,
updated_at bigint NOT NULL
);
ALTER TABLE ONLY public.user_groups ADD CONSTRAINT user_groups_pkey PRIMARY KEY (id);
ALTER TABLE ONLY public.user_groups ADD CONSTRAINT user_groups_normalized_name_key UNIQUE (normalized_name);
CREATE INDEX IF NOT EXISTS user_groups_priority_name_idx ON public.user_groups USING btree (priority, name, id);
CREATE TABLE IF NOT EXISTS public.user_group_members (
group_id character varying(64) NOT NULL,
user_id character varying(64) NOT NULL,
created_at bigint NOT NULL
);
ALTER TABLE ONLY public.user_group_members ADD CONSTRAINT user_group_members_pkey PRIMARY KEY (group_id, user_id);
CREATE INDEX IF NOT EXISTS user_group_members_user_id_idx ON public.user_group_members USING btree (user_id);
CREATE TABLE IF NOT EXISTS public.api_keys (
id character varying(64) NOT NULL,
user_id character varying(64) NOT NULL,

View File

@@ -13,10 +13,14 @@ CREATE TABLE IF NOT EXISTS users (
is_active INTEGER NOT NULL DEFAULT 1,
is_deleted INTEGER NOT NULL DEFAULT 0,
allowed_models TEXT,
allowed_models_mode TEXT NOT NULL DEFAULT 'unrestricted',
allowed_providers TEXT,
allowed_providers_mode TEXT NOT NULL DEFAULT 'unrestricted',
allowed_api_formats TEXT,
allowed_api_formats_mode TEXT NOT NULL DEFAULT 'unrestricted',
model_capability_settings TEXT,
rate_limit INTEGER,
rate_limit_mode TEXT NOT NULL DEFAULT 'system',
metadata TEXT,
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
@@ -27,6 +31,34 @@ CREATE TABLE IF NOT EXISTS users (
UNIQUE (username)
);
CREATE TABLE IF NOT EXISTS user_groups (
id TEXT PRIMARY KEY NOT NULL,
name TEXT NOT NULL,
normalized_name TEXT NOT NULL,
description TEXT,
priority INTEGER NOT NULL DEFAULT 0,
allowed_providers TEXT,
allowed_providers_mode TEXT NOT NULL DEFAULT 'inherit',
allowed_api_formats TEXT,
allowed_api_formats_mode TEXT NOT NULL DEFAULT 'inherit',
allowed_models TEXT,
allowed_models_mode TEXT NOT NULL DEFAULT 'inherit',
rate_limit INTEGER,
rate_limit_mode TEXT NOT NULL DEFAULT 'inherit',
created_at INTEGER NOT NULL,
updated_at INTEGER NOT NULL,
UNIQUE (normalized_name)
);
CREATE INDEX IF NOT EXISTS user_groups_priority_name_idx ON user_groups (priority, name, id);
CREATE TABLE IF NOT EXISTS user_group_members (
group_id TEXT NOT NULL,
user_id TEXT NOT NULL,
created_at INTEGER NOT NULL,
PRIMARY KEY (group_id, user_id)
);
CREATE INDEX IF NOT EXISTS user_group_members_user_id_idx ON user_group_members (user_id);
CREATE TABLE IF NOT EXISTS api_keys (
id TEXT PRIMARY KEY NOT NULL,
user_id TEXT NOT NULL,

View File

@@ -64,16 +64,34 @@ name = "allowed_models"
type = "json"
nullable = true
[[table.users.columns]]
name = "allowed_models_mode"
type = "text"
length = 32
default = "unrestricted"
[[table.users.columns]]
name = "allowed_providers"
type = "json"
nullable = true
[[table.users.columns]]
name = "allowed_providers_mode"
type = "text"
length = 32
default = "unrestricted"
[[table.users.columns]]
name = "allowed_api_formats"
type = "json"
nullable = true
[[table.users.columns]]
name = "allowed_api_formats_mode"
type = "text"
length = 32
default = "unrestricted"
[[table.users.columns]]
name = "model_capability_settings"
type = "json"
@@ -84,6 +102,12 @@ name = "rate_limit"
type = "int32"
nullable = true
[[table.users.columns]]
name = "rate_limit_mode"
type = "text"
length = 32
default = "system"
[[table.users.columns]]
name = "metadata"
type = "json"
@@ -122,6 +146,119 @@ columns = ["email"]
name = "users_username_key"
columns = ["username"]
[table.user_groups]
domain = "identity"
order = 15
primary_key = ["id"]
[[table.user_groups.columns]]
name = "id"
type = "text_id"
length = 64
[[table.user_groups.columns]]
name = "name"
type = "text"
length = 100
[[table.user_groups.columns]]
name = "normalized_name"
type = "text"
length = 100
[[table.user_groups.columns]]
name = "description"
type = "long_text"
nullable = true
[[table.user_groups.columns]]
name = "priority"
type = "int32"
default = 0
[[table.user_groups.columns]]
name = "allowed_providers"
type = "json"
nullable = true
[[table.user_groups.columns]]
name = "allowed_providers_mode"
type = "text"
length = 32
default = "inherit"
[[table.user_groups.columns]]
name = "allowed_api_formats"
type = "json"
nullable = true
[[table.user_groups.columns]]
name = "allowed_api_formats_mode"
type = "text"
length = 32
default = "inherit"
[[table.user_groups.columns]]
name = "allowed_models"
type = "json"
nullable = true
[[table.user_groups.columns]]
name = "allowed_models_mode"
type = "text"
length = 32
default = "inherit"
[[table.user_groups.columns]]
name = "rate_limit"
type = "int32"
nullable = true
[[table.user_groups.columns]]
name = "rate_limit_mode"
type = "text"
length = 32
default = "inherit"
[[table.user_groups.columns]]
name = "created_at"
type = "unix_seconds"
[[table.user_groups.columns]]
name = "updated_at"
type = "unix_seconds"
[[table.user_groups.uniques]]
name = "user_groups_normalized_name_key"
columns = ["normalized_name"]
[[table.user_groups.indexes]]
name = "user_groups_priority_name_idx"
columns = ["priority", "name", "id"]
[table.user_group_members]
domain = "identity"
order = 16
primary_key = ["group_id", "user_id"]
[[table.user_group_members.columns]]
name = "group_id"
type = "text_id"
length = 64
[[table.user_group_members.columns]]
name = "user_id"
type = "text_id"
length = 64
[[table.user_group_members.columns]]
name = "created_at"
type = "unix_seconds"
[[table.user_group_members.indexes]]
name = "user_group_members_user_id_idx"
columns = ["user_id"]
[table.api_keys]
domain = "identity"
order = 20

View File

@@ -7,7 +7,7 @@ use tracing::info;
// Generated by build.rs from schema/bootstrap/postgres.
pub(crate) static EMPTY_DATABASE_SNAPSHOT_SQL: &str =
include_str!(concat!(env!("OUT_DIR"), "/empty_database_snapshot.sql"));
pub(crate) const EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION: i64 = 20260508000000;
pub(crate) const EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION: i64 = 20260509000000;
const PUBLIC_BASE_TABLE_COUNT_SQL: &str = r#"
SELECT COUNT(*)::BIGINT
@@ -27,11 +27,12 @@ WHERE table_schema = 'public'
'gemini_file_mappings',
'global_models',
'oauth_providers',
'provider_api_keys',
'proxy_nodes',
'usage_routing_snapshots',
'usage_settlement_snapshots'
)
'provider_api_keys',
'proxy_nodes',
'user_groups',
'usage_routing_snapshots',
'usage_settlement_snapshots'
)
"#;
const INSERT_APPLIED_MIGRATION_SQL: &str = r#"
INSERT INTO _sqlx_migrations (

View File

@@ -294,6 +294,7 @@ fn empty_database_snapshot_covers_current_cutoff_versions() {
20260507000000,
20260507120000,
20260508000000,
20260509000000,
]
);
}
@@ -513,11 +514,21 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() {
assert_eq!(
mysql_versions,
vec![20260403000000, 20260507120000, 20260508000000]
vec![
20260403000000,
20260507120000,
20260508000000,
20260509000000
]
);
assert_eq!(
sqlite_versions,
vec![20260403000000, 20260507120000, 20260508000000]
vec![
20260403000000,
20260507120000,
20260508000000,
20260509000000
]
);
}
@@ -1022,6 +1033,7 @@ fn pending_migrations_from_applied_skips_versions_already_applied() {
20260507000000,
20260507120000,
20260508000000,
20260509000000,
]
);
}

View File

@@ -187,7 +187,7 @@ impl ResolvedAuthApiKeySnapshot {
}
non_empty_allowed_list(self.api_key_allowed_providers.as_deref())
.or_else(|| non_empty_allowed_list(self.user_allowed_providers.as_deref()))
.or(self.user_allowed_providers.as_deref())
}
pub fn effective_allowed_api_formats(&self) -> Option<&[String]> {
@@ -196,7 +196,7 @@ impl ResolvedAuthApiKeySnapshot {
}
non_empty_allowed_list(self.api_key_allowed_api_formats.as_deref())
.or_else(|| non_empty_allowed_list(self.user_allowed_api_formats.as_deref()))
.or(self.user_allowed_api_formats.as_deref())
}
pub fn effective_allowed_models(&self) -> Option<&[String]> {
@@ -205,7 +205,20 @@ impl ResolvedAuthApiKeySnapshot {
}
non_empty_allowed_list(self.api_key_allowed_models.as_deref())
.or_else(|| non_empty_allowed_list(self.user_allowed_models.as_deref()))
.or(self.user_allowed_models.as_deref())
}
pub fn apply_user_policy(
&mut self,
allowed_providers: Option<Vec<String>>,
allowed_api_formats: Option<Vec<String>>,
allowed_models: Option<Vec<String>>,
rate_limit: Option<i32>,
) {
self.user_allowed_providers = allowed_providers;
self.user_allowed_api_formats = allowed_api_formats;
self.user_allowed_models = allowed_models;
self.user_rate_limit = rate_limit;
}
}

View File

@@ -4,9 +4,11 @@ use std::sync::RwLock;
use async_trait::async_trait;
use super::types::{
LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, StoredUserExportRow,
normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord,
StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership,
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
StoredUserSummary, UserExportListQuery, UserExportSummary, UserReadRepository,
StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSummary,
UserReadRepository,
};
use crate::DataLayerError;
@@ -34,6 +36,8 @@ pub struct InMemoryUserReadRepository {
preferences_by_user_id: RwLock<BTreeMap<String, StoredUserPreferenceRecord>>,
sessions_by_id: RwLock<BTreeMap<String, StoredUserSessionRecord>>,
model_settings_by_user_id: RwLock<BTreeMap<String, serde_json::Value>>,
groups_by_id: RwLock<BTreeMap<String, StoredUserGroup>>,
group_members: RwLock<BTreeMap<(String, String), chrono::DateTime<chrono::Utc>>>,
export_rows: RwLock<Vec<StoredUserExportRow>>,
read_only: bool,
}
@@ -57,6 +61,8 @@ impl InMemoryUserReadRepository {
preferences_by_user_id: RwLock::new(BTreeMap::new()),
sessions_by_id: RwLock::new(BTreeMap::new()),
model_settings_by_user_id: RwLock::new(BTreeMap::new()),
groups_by_id: RwLock::new(BTreeMap::new()),
group_members: RwLock::new(BTreeMap::new()),
export_rows: RwLock::new(Vec::new()),
read_only: false,
}
@@ -90,6 +96,8 @@ impl InMemoryUserReadRepository {
preferences_by_user_id: RwLock::new(BTreeMap::new()),
sessions_by_id: RwLock::new(BTreeMap::new()),
model_settings_by_user_id: RwLock::new(BTreeMap::new()),
groups_by_id: RwLock::new(BTreeMap::new()),
group_members: RwLock::new(BTreeMap::new()),
export_rows: RwLock::new(Vec::new()),
read_only: false,
}
@@ -109,6 +117,8 @@ impl InMemoryUserReadRepository {
preferences_by_user_id: RwLock::new(BTreeMap::new()),
sessions_by_id: RwLock::new(BTreeMap::new()),
model_settings_by_user_id: RwLock::new(BTreeMap::new()),
groups_by_id: RwLock::new(BTreeMap::new()),
group_members: RwLock::new(BTreeMap::new()),
export_rows: RwLock::new(items.into_iter().collect()),
read_only: false,
}
@@ -257,6 +267,95 @@ fn upsert_memory_ldap_identifiers(
}
}
fn memory_group_from_record(
record: UpsertUserGroupRecord,
) -> Result<StoredUserGroup, DataLayerError> {
let now = chrono::Utc::now();
let name = normalize_user_group_name(&record.name);
StoredUserGroup::new(
uuid::Uuid::new_v4().to_string(),
name.clone(),
name.to_ascii_lowercase(),
record.description,
record.priority,
record.allowed_providers.map(serde_json::Value::from),
record.allowed_providers_mode,
record.allowed_api_formats.map(serde_json::Value::from),
record.allowed_api_formats_mode,
record.allowed_models.map(serde_json::Value::from),
record.allowed_models_mode,
record.rate_limit,
record.rate_limit_mode,
Some(now),
Some(now),
)
}
fn memory_update_group_from_record(
mut group: StoredUserGroup,
record: UpsertUserGroupRecord,
) -> Result<StoredUserGroup, DataLayerError> {
let name = normalize_user_group_name(&record.name);
group.name = name.clone();
group.normalized_name = name.to_ascii_lowercase();
group.description = record.description;
group.priority = record.priority;
group.allowed_providers = record.allowed_providers;
group.allowed_providers_mode = record.allowed_providers_mode;
group.allowed_api_formats = record.allowed_api_formats;
group.allowed_api_formats_mode = record.allowed_api_formats_mode;
group.allowed_models = record.allowed_models;
group.allowed_models_mode = record.allowed_models_mode;
group.rate_limit = record.rate_limit;
group.rate_limit_mode = record.rate_limit_mode;
group.updated_at = Some(chrono::Utc::now());
StoredUserGroup::new(
group.id,
group.name,
group.normalized_name,
group.description,
group.priority,
group.allowed_providers.map(serde_json::Value::from),
group.allowed_providers_mode,
group.allowed_api_formats.map(serde_json::Value::from),
group.allowed_api_formats_mode,
group.allowed_models.map(serde_json::Value::from),
group.allowed_models_mode,
group.rate_limit,
group.rate_limit_mode,
group.created_at,
group.updated_at,
)
}
fn memory_group_members(
repository: &InMemoryUserReadRepository,
group_id: &str,
) -> Vec<StoredUserGroupMember> {
let members = repository
.group_members
.read()
.expect("user repository lock")
.clone();
let users = repository.auth_by_id.read().expect("user repository lock");
members
.into_iter()
.filter(|((candidate_group_id, _), _)| candidate_group_id == group_id)
.filter_map(|((candidate_group_id, user_id), created_at)| {
users.get(&user_id).map(|user| StoredUserGroupMember {
group_id: candidate_group_id,
user_id: user.id.clone(),
username: user.username.clone(),
email: user.email.clone(),
role: user.role.clone(),
is_active: user.is_active,
is_deleted: user.is_deleted,
created_at: Some(created_at),
})
})
.collect()
}
#[async_trait]
impl UserReadRepository for InMemoryUserReadRepository {
async fn list_users_by_ids(
@@ -329,6 +428,23 @@ impl UserReadRepository for InMemoryUserReadRepository {
if let Some(is_active) = query.is_active {
rows.retain(|row| row.is_active == is_active);
}
if let Some(group_id) = query
.group_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
let member_ids = self
.group_members
.read()
.expect("user repository lock")
.keys()
.filter_map(|(candidate_group_id, user_id)| {
(candidate_group_id == group_id).then(|| user_id.clone())
})
.collect::<std::collections::BTreeSet<_>>();
rows.retain(|row| member_ids.contains(&row.id));
}
if let Some(search) = query
.search
.as_deref()
@@ -375,6 +491,273 @@ impl UserReadRepository for InMemoryUserReadRepository {
.cloned())
}
async fn list_user_groups(&self) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut groups = self
.groups_by_id
.read()
.expect("user repository lock")
.values()
.cloned()
.collect::<Vec<_>>();
groups.sort_by(|left, right| {
right
.priority
.cmp(&left.priority)
.then_with(|| left.name.cmp(&right.name))
.then_with(|| left.id.cmp(&right.id))
});
Ok(groups)
}
async fn find_user_group_by_id(
&self,
group_id: &str,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
Ok(self
.groups_by_id
.read()
.expect("user repository lock")
.get(group_id)
.cloned())
}
async fn list_user_groups_by_ids(
&self,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let groups = self.groups_by_id.read().expect("user repository lock");
Ok(group_ids
.iter()
.filter_map(|group_id| groups.get(group_id).cloned())
.collect())
}
async fn create_user_group(
&self,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
if self.read_only {
return Ok(None);
}
let group = memory_group_from_record(record)?;
let mut groups = self.groups_by_id.write().expect("user repository lock");
if groups
.values()
.any(|existing| existing.normalized_name == group.normalized_name)
{
return Err(DataLayerError::InvalidInput(format!(
"duplicate user group name: {}",
group.name
)));
}
groups.insert(group.id.clone(), group.clone());
Ok(Some(group))
}
async fn update_user_group(
&self,
group_id: &str,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
if self.read_only {
return Ok(None);
}
let mut groups = self.groups_by_id.write().expect("user repository lock");
let Some(existing) = groups.get(group_id).cloned() else {
return Ok(None);
};
let group = memory_update_group_from_record(existing, record)?;
if groups.values().any(|existing| {
existing.id != group.id && existing.normalized_name == group.normalized_name
}) {
return Err(DataLayerError::InvalidInput(format!(
"duplicate user group name: {}",
group.name
)));
}
groups.insert(group.id.clone(), group.clone());
Ok(Some(group))
}
async fn delete_user_group(&self, group_id: &str) -> Result<bool, DataLayerError> {
if self.read_only {
return Ok(false);
}
let removed = self
.groups_by_id
.write()
.expect("user repository lock")
.remove(group_id)
.is_some();
if removed {
self.group_members
.write()
.expect("user repository lock")
.retain(|key, _| key.0 != group_id);
}
Ok(removed)
}
async fn list_user_group_members(
&self,
group_id: &str,
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
Ok(memory_group_members(self, group_id))
}
async fn replace_user_group_members(
&self,
group_id: &str,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
if self.read_only {
return Ok(Vec::new());
}
if !self
.groups_by_id
.read()
.expect("user repository lock")
.contains_key(group_id)
{
return Ok(Vec::new());
}
let valid_user_ids = {
let users = self.auth_by_id.read().expect("user repository lock");
user_ids
.iter()
.map(|value| value.trim())
.filter(|value| !value.is_empty())
.filter(|user_id| users.contains_key(*user_id))
.map(ToOwned::to_owned)
.collect::<std::collections::BTreeSet<_>>()
};
let now = chrono::Utc::now();
let mut members = self.group_members.write().expect("user repository lock");
members.retain(|key, _| key.0 != group_id);
for user_id in valid_user_ids {
members.insert((group_id.to_string(), user_id), now);
}
drop(members);
Ok(memory_group_members(self, group_id))
}
async fn list_user_groups_for_user(
&self,
user_id: &str,
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let group_ids = self
.group_members
.read()
.expect("user repository lock")
.keys()
.filter_map(|(group_id, candidate_user_id)| {
(candidate_user_id == user_id).then(|| group_id.clone())
})
.collect::<Vec<_>>();
self.list_user_groups_by_ids(&group_ids).await
}
async fn list_user_group_memberships_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMembership>, DataLayerError> {
let requested = user_ids
.iter()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.collect::<std::collections::BTreeSet<_>>();
if requested.is_empty() {
return Ok(Vec::new());
}
let groups = self.groups_by_id.read().expect("user repository lock");
let members = self.group_members.read().expect("user repository lock");
let mut memberships = members
.iter()
.filter(|((_, user_id), _)| requested.contains(user_id))
.filter_map(|((group_id, user_id), created_at)| {
groups.get(group_id).map(|group| StoredUserGroupMembership {
user_id: user_id.clone(),
group_id: group.id.clone(),
group_name: group.name.clone(),
group_priority: group.priority,
created_at: Some(*created_at),
})
})
.collect::<Vec<_>>();
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))
});
Ok(memberships)
}
async fn replace_user_groups_for_user(
&self,
user_id: &str,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
if self.read_only {
return Ok(Vec::new());
}
let existing_group_ids = {
let groups = self.groups_by_id.read().expect("user repository lock");
group_ids
.iter()
.map(|value| value.trim())
.filter(|value| !value.is_empty())
.filter(|group_id| groups.contains_key(*group_id))
.map(ToOwned::to_owned)
.collect::<std::collections::BTreeSet<_>>()
};
{
let now = chrono::Utc::now();
let mut members = self.group_members.write().expect("user repository lock");
members.retain(|key, _| key.1 != user_id);
for group_id in &existing_group_ids {
members.insert((group_id.clone(), user_id.to_string()), now);
}
}
self.list_user_groups_by_ids(&existing_group_ids.into_iter().collect::<Vec<_>>())
.await
}
async fn add_user_to_group(
&self,
group_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
if self.read_only {
return Ok(false);
}
if !self
.groups_by_id
.read()
.expect("user repository lock")
.contains_key(group_id)
{
return Ok(false);
}
if !self
.auth_by_id
.read()
.expect("user repository lock")
.contains_key(user_id)
{
return Ok(false);
}
self.group_members
.write()
.expect("user repository lock")
.insert(
(group_id.to_string(), user_id.to_string()),
chrono::Utc::now(),
);
Ok(true)
}
async fn find_user_auth_by_id(
&self,
user_id: &str,
@@ -581,6 +964,11 @@ impl UserReadRepository for InMemoryUserReadRepository {
false,
Some(created_at),
Some(created_at),
)?
.with_policy_modes(
"inherit".to_string(),
"inherit".to_string(),
"inherit".to_string(),
)?;
self.insert_auth_user(user).map(Some)
}
@@ -917,12 +1305,27 @@ impl UserReadRepository for InMemoryUserReadRepository {
}
if allowed_providers_present {
user.allowed_providers = allowed_providers;
user.allowed_providers_mode = if user.allowed_providers.is_some() {
"specific".to_string()
} else {
"unrestricted".to_string()
};
}
if allowed_api_formats_present {
user.allowed_api_formats = allowed_api_formats;
user.allowed_api_formats_mode = if user.allowed_api_formats.is_some() {
"specific".to_string()
} else {
"unrestricted".to_string()
};
}
if allowed_models_present {
user.allowed_models = allowed_models;
user.allowed_models_mode = if user.allowed_models.is_some() {
"specific".to_string()
} else {
"unrestricted".to_string()
};
}
if let Some(is_active) = is_active {
user.is_active = is_active;
@@ -948,16 +1351,75 @@ impl UserReadRepository for InMemoryUserReadRepository {
{
row.role = updated.role.clone();
row.allowed_providers = updated.allowed_providers.clone();
row.allowed_providers_mode = updated.allowed_providers_mode.clone();
row.allowed_api_formats = updated.allowed_api_formats.clone();
row.allowed_api_formats_mode = updated.allowed_api_formats_mode.clone();
row.allowed_models = updated.allowed_models.clone();
row.allowed_models_mode = updated.allowed_models_mode.clone();
if rate_limit_present {
row.rate_limit = rate_limit;
row.rate_limit_mode = if row.rate_limit.is_some() {
"custom".to_string()
} else {
"system".to_string()
};
}
row.is_active = updated.is_active;
}
Ok(Some(updated))
}
async fn update_local_auth_user_policy_modes(
&self,
user_id: &str,
allowed_providers_mode: Option<String>,
allowed_api_formats_mode: Option<String>,
allowed_models_mode: Option<String>,
rate_limit_mode: Option<String>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
if self.read_only {
return Ok(None);
}
let mut auth_by_id = self.auth_by_id.write().expect("user repository lock");
let Some(user) = auth_by_id.get_mut(user_id) else {
return Ok(None);
};
if let Some(mode) = allowed_providers_mode.clone() {
user.allowed_providers_mode = mode;
}
if let Some(mode) = allowed_api_formats_mode.clone() {
user.allowed_api_formats_mode = mode;
}
if let Some(mode) = allowed_models_mode.clone() {
user.allowed_models_mode = mode;
}
let updated = user.clone();
drop(auth_by_id);
if let Some(row) = self
.export_rows
.write()
.expect("user repository lock")
.iter_mut()
.find(|row| row.id == user_id)
{
if let Some(mode) = allowed_providers_mode {
row.allowed_providers_mode = mode;
}
if let Some(mode) = allowed_api_formats_mode {
row.allowed_api_formats_mode = mode;
}
if let Some(mode) = allowed_models_mode {
row.allowed_models_mode = mode;
}
if let Some(mode) = rate_limit_mode {
row.rate_limit_mode = mode;
}
}
Ok(Some(updated))
}
async fn update_user_model_capability_settings(
&self,
user_id: &str,
@@ -1021,18 +1483,29 @@ impl UserReadRepository for InMemoryUserReadRepository {
return Ok(None);
}
self.create_local_auth_user_with_settings(
let now = chrono::Utc::now();
let user = StoredUserAuthRecord::new(
uuid::Uuid::new_v4().to_string(),
email,
email_verified,
username,
password_hash,
Some(password_hash),
"user".to_string(),
"local".to_string(),
None,
None,
None,
true,
false,
Some(now),
None,
)
.await
)?
.with_policy_modes(
"inherit".to_string(),
"inherit".to_string(),
"inherit".to_string(),
)?;
self.insert_auth_user(user).map(Some)
}
async fn create_local_auth_user_with_settings(
@@ -1092,6 +1565,10 @@ impl UserReadRepository for InMemoryUserReadRepository {
.write()
.expect("user repository lock")
.retain(|_, link| link.user_id != user_id);
self.group_members
.write()
.expect("user repository lock")
.retain(|key, _| key.1 != user_id);
let mut identifiers = self
.auth_by_identifier
@@ -2073,6 +2550,7 @@ mod tests {
role: Some("user".to_string()),
is_active: Some(true),
search: None,
group_id: None,
})
.await
.expect("paged export should succeed");

View File

@@ -9,7 +9,8 @@ pub use mysql::MysqlUserReadRepository;
pub use postgres::SqlxUserReadRepository;
pub use sqlite::SqliteUserReadRepository;
pub use types::{
StoredUserAuthRecord, StoredUserExportRow, StoredUserOAuthLinkSummary,
StoredUserPreferenceRecord, StoredUserSessionRecord, StoredUserSummary, UserExportListQuery,
UserExportSummary, UserReadRepository,
normalize_user_group_name, StoredUserAuthRecord, StoredUserExportRow, StoredUserGroup,
StoredUserGroupMember, StoredUserGroupMembership, StoredUserOAuthLinkSummary,
StoredUserPreferenceRecord, StoredUserSessionRecord, StoredUserSummary, UpsertUserGroupRecord,
UserExportListQuery, UserExportSummary, UserReadRepository,
};

View File

@@ -3,9 +3,11 @@ use chrono::{DateTime, TimeZone, Utc};
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use super::types::{
LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, StoredUserExportRow,
normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord,
StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership,
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
StoredUserSummary, UserExportListQuery, UserExportSummary, UserReadRepository,
StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSummary,
UserReadRepository,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
@@ -32,9 +34,13 @@ SELECT
role,
auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
model_capability_settings,
is_active
FROM users
@@ -50,8 +56,11 @@ SELECT
role,
auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -69,8 +78,11 @@ SELECT
users.role AS role,
users.auth_source AS auth_source,
users.allowed_providers AS allowed_providers,
users.allowed_providers_mode AS allowed_providers_mode,
users.allowed_api_formats AS allowed_api_formats,
users.allowed_api_formats_mode AS allowed_api_formats_mode,
users.allowed_models AS allowed_models,
users.allowed_models_mode AS allowed_models_mode,
users.is_active AS is_active,
users.is_deleted AS is_deleted,
users.created_at AS created_at,
@@ -130,6 +142,40 @@ SELECT
FROM user_sessions
"#;
const USER_GROUP_COLUMNS: &str = r#"
SELECT
id,
name,
normalized_name,
description,
priority,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
created_at,
updated_at
FROM user_groups
"#;
const USER_GROUP_MEMBER_COLUMNS: &str = r#"
SELECT
user_group_members.group_id,
users.id AS user_id,
users.username,
users.email,
users.role,
users.is_active,
users.is_deleted,
user_group_members.created_at
FROM user_group_members
JOIN users ON users.id = user_group_members.user_id
"#;
#[derive(Debug, Clone)]
pub struct MysqlUserReadRepository {
pool: MysqlPool,
@@ -163,6 +209,22 @@ impl MysqlUserReadRepository {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_auth_row).collect()
}
async fn fetch_group_rows(
&self,
mut builder: QueryBuilder<'_, MySql>,
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_group_row).collect()
}
async fn fetch_group_member_rows(
&self,
mut builder: QueryBuilder<'_, MySql>,
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_group_member_row).collect()
}
}
#[async_trait]
@@ -224,6 +286,16 @@ impl UserReadRepository for MysqlUserReadRepository {
if let Some(is_active) = query.is_active {
builder.push(" AND is_active = ").push_bind(is_active);
}
if let Some(group_id) = query
.group_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
builder.push(" AND id IN (SELECT user_id FROM user_group_members WHERE group_id = ");
builder.push_bind(group_id);
builder.push(")");
}
if let Some(search) = query
.search
.as_deref()
@@ -296,6 +368,285 @@ WHERE is_deleted = 0
self.fetch_export_rows(builder).await
}
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");
self.fetch_group_rows(builder).await
}
async fn find_user_group_by_id(
&self,
group_id: &str,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(USER_GROUP_COLUMNS);
builder
.push(" WHERE id = ")
.push_bind(group_id)
.push(" LIMIT 1");
Ok(self.fetch_group_rows(builder).await?.into_iter().next())
}
async fn list_user_groups_by_ids(
&self,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
if group_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<MySql>::new(USER_GROUP_COLUMNS);
builder.push(" WHERE id IN (");
{
let mut separated = builder.separated(", ");
for group_id in group_ids {
separated.push_bind(group_id);
}
}
builder.push(") ORDER BY priority DESC, name ASC, id ASC");
self.fetch_group_rows(builder).await
}
async fn create_user_group(
&self,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let now = current_unix_secs();
let id = uuid::Uuid::new_v4().to_string();
let name = normalize_user_group_name(&record.name);
let normalized_name = name.to_ascii_lowercase();
let result = sqlx::query(
r#"
INSERT INTO user_groups (
id, name, normalized_name, description, priority,
allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode,
allowed_models, allowed_models_mode,
rate_limit, rate_limit_mode, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&id)
.bind(name)
.bind(normalized_name)
.bind(record.description)
.bind(record.priority)
.bind(json_string_from_option_vec(
record.allowed_providers.as_ref(),
))
.bind(record.allowed_providers_mode)
.bind(json_string_from_option_vec(
record.allowed_api_formats.as_ref(),
))
.bind(record.allowed_api_formats_mode)
.bind(json_string_from_option_vec(record.allowed_models.as_ref()))
.bind(record.allowed_models_mode)
.bind(record.rate_limit)
.bind(record.rate_limit_mode)
.bind(now)
.bind(now)
.execute(&self.pool)
.await;
match result {
Ok(_) => self.find_user_group_by_id(&id).await,
Err(sqlx::Error::Database(err)) if err.is_unique_violation() => Err(
DataLayerError::InvalidInput("duplicate user group name".to_string()),
),
Err(err) => Err(err).map_sql_err(),
}
}
async fn update_user_group(
&self,
group_id: &str,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let now = current_unix_secs();
let name = normalize_user_group_name(&record.name);
let normalized_name = name.to_ascii_lowercase();
let result = sqlx::query(
r#"
UPDATE user_groups
SET name = ?,
normalized_name = ?,
description = ?,
priority = ?,
allowed_providers = ?,
allowed_providers_mode = ?,
allowed_api_formats = ?,
allowed_api_formats_mode = ?,
allowed_models = ?,
allowed_models_mode = ?,
rate_limit = ?,
rate_limit_mode = ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(name)
.bind(normalized_name)
.bind(record.description)
.bind(record.priority)
.bind(json_string_from_option_vec(
record.allowed_providers.as_ref(),
))
.bind(record.allowed_providers_mode)
.bind(json_string_from_option_vec(
record.allowed_api_formats.as_ref(),
))
.bind(record.allowed_api_formats_mode)
.bind(json_string_from_option_vec(record.allowed_models.as_ref()))
.bind(record.allowed_models_mode)
.bind(record.rate_limit)
.bind(record.rate_limit_mode)
.bind(now)
.bind(group_id)
.execute(&self.pool)
.await;
match result {
Ok(result) if result.rows_affected() == 0 => Ok(None),
Ok(_) => self.find_user_group_by_id(group_id).await,
Err(sqlx::Error::Database(err)) if err.is_unique_violation() => Err(
DataLayerError::InvalidInput("duplicate user group name".to_string()),
),
Err(err) => Err(err).map_sql_err(),
}
}
async fn delete_user_group(&self, group_id: &str) -> Result<bool, DataLayerError> {
let result = sqlx::query("DELETE FROM user_groups WHERE id = ?")
.bind(group_id)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
async fn list_user_group_members(
&self,
group_id: &str,
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(USER_GROUP_MEMBER_COLUMNS);
builder
.push(" WHERE user_group_members.group_id = ")
.push_bind(group_id)
.push(" ORDER BY users.username ASC, users.id ASC");
self.fetch_group_member_rows(builder).await
}
async fn replace_user_group_members(
&self,
group_id: &str,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("DELETE FROM user_group_members WHERE group_id = ?")
.bind(group_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
let now = current_unix_secs();
for user_id in normalized_ids(user_ids) {
sqlx::query(
"INSERT IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)",
)
.bind(group_id)
.bind(user_id)
.bind(now)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
tx.commit().await.map_sql_err()?;
self.list_user_group_members(group_id).await
}
async fn list_user_groups_for_user(
&self,
user_id: &str,
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(USER_GROUP_COLUMNS);
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");
self.fetch_group_rows(builder).await
}
async fn list_user_group_memberships_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMembership>, DataLayerError> {
if user_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<MySql>::new(
r#"
SELECT
user_group_members.user_id,
user_groups.id AS group_id,
user_groups.name AS group_name,
user_groups.priority AS group_priority,
user_group_members.created_at
FROM user_group_members
JOIN user_groups ON user_groups.id = user_group_members.group_id
WHERE user_group_members.user_id IN (
"#,
);
{
let mut separated = builder.separated(", ");
for user_id in user_ids {
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");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_group_membership_row).collect()
}
async fn replace_user_groups_for_user(
&self,
user_id: &str,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("DELETE FROM user_group_members WHERE user_id = ?")
.bind(user_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
let now = current_unix_secs();
for group_id in normalized_ids(group_ids) {
sqlx::query(
"INSERT IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)",
)
.bind(group_id)
.bind(user_id)
.bind(now)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
tx.commit().await.map_sql_err()?;
self.list_user_groups_for_user(user_id).await
}
async fn add_user_to_group(
&self,
group_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
let result = sqlx::query(
"INSERT IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)",
)
.bind(group_id)
.bind(user_id)
.bind(current_unix_secs())
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
async fn find_user_auth_by_id(
&self,
user_id: &str,
@@ -453,9 +804,10 @@ WHERE provider_type = ?
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
is_active, is_deleted, created_at, updated_at, last_login_at
)
VALUES (?, ?, 1, ?, NULL, 'user', 'oauth', 1, 0, ?, ?, ?)
VALUES (?, ?, 1, ?, NULL, 'user', 'oauth', 'inherit', 'inherit', 'inherit', 'inherit', 1, 0, ?, ?, ?)
"#,
)
.bind(&user_id)
@@ -675,14 +1027,38 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
rate_limit: Option<i32>,
is_active: Option<bool>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let allowed_providers_mode = if allowed_providers.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_api_formats_mode = if allowed_api_formats.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_models_mode = if allowed_models.is_some() {
"specific"
} else {
"unrestricted"
};
let rate_limit_mode = if rate_limit.is_some() {
"custom"
} else {
"system"
};
let result = sqlx::query(
r#"
UPDATE users
SET role = CASE WHEN ? THEN COALESCE(?, role) ELSE role END,
allowed_providers = CASE WHEN ? THEN ? ELSE allowed_providers END,
allowed_providers_mode = CASE WHEN ? THEN ? ELSE allowed_providers_mode END,
allowed_api_formats = CASE WHEN ? THEN ? ELSE allowed_api_formats END,
allowed_api_formats_mode = CASE WHEN ? THEN ? ELSE allowed_api_formats_mode END,
allowed_models = CASE WHEN ? THEN ? ELSE allowed_models END,
allowed_models_mode = CASE WHEN ? THEN ? ELSE allowed_models_mode END,
rate_limit = CASE WHEN ? THEN ? ELSE rate_limit END,
rate_limit_mode = CASE WHEN ? THEN ? ELSE rate_limit_mode END,
is_active = CASE WHEN ? THEN ? ELSE is_active END,
updated_at = ?
WHERE id = ?
@@ -695,18 +1071,26 @@ WHERE id = ?
allowed_providers,
"users.allowed_providers",
)?)
.bind(allowed_providers_present)
.bind(allowed_providers_mode)
.bind(allowed_api_formats_present)
.bind(optional_string_list_json(
allowed_api_formats,
"users.allowed_api_formats",
)?)
.bind(allowed_api_formats_present)
.bind(allowed_api_formats_mode)
.bind(allowed_models_present)
.bind(optional_string_list_json(
allowed_models,
"users.allowed_models",
)?)
.bind(allowed_models_present)
.bind(allowed_models_mode)
.bind(rate_limit_present)
.bind(rate_limit)
.bind(rate_limit_present)
.bind(rate_limit_mode)
.bind(is_active.is_some())
.bind(is_active)
.bind(chrono::Utc::now().timestamp())
@@ -720,6 +1104,44 @@ WHERE id = ?
self.find_user_auth_by_id(user_id).await
}
async fn update_local_auth_user_policy_modes(
&self,
user_id: &str,
allowed_providers_mode: Option<String>,
allowed_api_formats_mode: Option<String>,
allowed_models_mode: Option<String>,
rate_limit_mode: Option<String>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let result = sqlx::query(
r#"
UPDATE users
SET allowed_providers_mode = CASE WHEN ? THEN ? ELSE allowed_providers_mode END,
allowed_api_formats_mode = CASE WHEN ? THEN ? ELSE allowed_api_formats_mode END,
allowed_models_mode = CASE WHEN ? THEN ? ELSE allowed_models_mode END,
rate_limit_mode = CASE WHEN ? THEN ? ELSE rate_limit_mode END,
updated_at = ?
WHERE id = ?
"#,
)
.bind(allowed_providers_mode.is_some())
.bind(allowed_providers_mode)
.bind(allowed_api_formats_mode.is_some())
.bind(allowed_api_formats_mode)
.bind(allowed_models_mode.is_some())
.bind(allowed_models_mode)
.bind(rate_limit_mode.is_some())
.bind(rate_limit_mode)
.bind(chrono::Utc::now().timestamp())
.bind(user_id)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.find_user_auth_by_id(user_id).await
}
async fn update_user_model_capability_settings(
&self,
user_id: &str,
@@ -751,18 +1173,29 @@ WHERE id = ?
username: String,
password_hash: String,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
self.create_local_auth_user_with_settings(
email,
email_verified,
username,
password_hash,
"user".to_string(),
None,
None,
None,
None,
let user_id = uuid::Uuid::new_v4().to_string();
let now = chrono::Utc::now().timestamp();
sqlx::query(
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
is_active, is_deleted, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, 'user', 'local', 'inherit', 'inherit', 'inherit', 'inherit', 1, 0, ?, ?)
"#,
)
.bind(&user_id)
.bind(email)
.bind(email_verified)
.bind(username)
.bind(password_hash)
.bind(now)
.bind(now)
.execute(&self.pool)
.await
.map_sql_err()?;
self.find_user_auth_by_id(&user_id).await
}
async fn create_local_auth_user_with_settings(
@@ -779,14 +1212,37 @@ WHERE id = ?
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let user_id = uuid::Uuid::new_v4().to_string();
let now = chrono::Utc::now().timestamp();
let allowed_providers_mode = if allowed_providers.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_api_formats_mode = if allowed_api_formats.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_models_mode = if allowed_models.is_some() {
"specific"
} else {
"unrestricted"
};
let rate_limit_mode = if rate_limit.is_some() {
"custom"
} else {
"system"
};
sqlx::query(
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers, allowed_api_formats, allowed_models, rate_limit,
allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode,
allowed_models, allowed_models_mode,
rate_limit, rate_limit_mode,
is_active, is_deleted, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, 1, 0, ?, ?)
VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, ?, ?, ?, ?, 1, 0, ?, ?)
"#,
)
.bind(&user_id)
@@ -799,15 +1255,19 @@ VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, 1, 0, ?, ?)
allowed_providers,
"users.allowed_providers",
)?)
.bind(allowed_providers_mode)
.bind(optional_string_list_json(
allowed_api_formats,
"users.allowed_api_formats",
)?)
.bind(allowed_api_formats_mode)
.bind(optional_string_list_json(
allowed_models,
"users.allowed_models",
)?)
.bind(allowed_models_mode)
.bind(rate_limit)
.bind(rate_limit_mode)
.bind(now)
.bind(now)
.execute(&self.pool)
@@ -1162,6 +1622,24 @@ fn optional_string_list_json(
.transpose()
}
fn json_string_from_option_vec(value: Option<&Vec<String>>) -> Option<String> {
value.and_then(|items| serde_json::to_string(items).ok())
}
fn normalized_ids(values: &[String]) -> Vec<String> {
values
.iter()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.collect::<std::collections::BTreeSet<_>>()
.into_iter()
.collect()
}
fn current_unix_secs() -> i64 {
chrono::Utc::now().timestamp()
}
fn optional_json_string(
value: Option<serde_json::Value>,
field_name: &str,
@@ -1378,6 +1856,14 @@ fn map_user_export_row(row: &MySqlRow) -> Result<StoredUserExportRow, DataLayerE
)?,
row.try_get("is_active").map_sql_err()?,
)
.and_then(|record| {
record.with_policy_modes(
row.try_get("allowed_providers_mode").map_sql_err()?,
row.try_get("allowed_api_formats_mode").map_sql_err()?,
row.try_get("allowed_models_mode").map_sql_err()?,
row.try_get("rate_limit_mode").map_sql_err()?,
)
})
}
fn map_user_auth_row(row: &MySqlRow) -> Result<StoredUserAuthRecord, DataLayerError> {
@@ -1406,6 +1892,67 @@ fn map_user_auth_row(row: &MySqlRow) -> Result<StoredUserAuthRecord, DataLayerEr
optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
optional_datetime_from_unix_secs(row.try_get("last_login_at").map_sql_err()?),
)
.and_then(|record| {
record.with_policy_modes(
row.try_get("allowed_providers_mode").map_sql_err()?,
row.try_get("allowed_api_formats_mode").map_sql_err()?,
row.try_get("allowed_models_mode").map_sql_err()?,
)
})
}
fn map_user_group_row(row: &MySqlRow) -> Result<StoredUserGroup, DataLayerError> {
StoredUserGroup::new(
row.try_get("id").map_sql_err()?,
row.try_get("name").map_sql_err()?,
row.try_get("normalized_name").map_sql_err()?,
row.try_get("description").map_sql_err()?,
row.try_get("priority").map_sql_err()?,
optional_json_from_string(
row.try_get("allowed_providers").map_sql_err()?,
"user_groups.allowed_providers",
)?,
row.try_get("allowed_providers_mode").map_sql_err()?,
optional_json_from_string(
row.try_get("allowed_api_formats").map_sql_err()?,
"user_groups.allowed_api_formats",
)?,
row.try_get("allowed_api_formats_mode").map_sql_err()?,
optional_json_from_string(
row.try_get("allowed_models").map_sql_err()?,
"user_groups.allowed_models",
)?,
row.try_get("allowed_models_mode").map_sql_err()?,
row.try_get("rate_limit").map_sql_err()?,
row.try_get("rate_limit_mode").map_sql_err()?,
optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
optional_datetime_from_unix_secs(row.try_get("updated_at").map_sql_err()?),
)
}
fn map_user_group_member_row(row: &MySqlRow) -> Result<StoredUserGroupMember, DataLayerError> {
Ok(StoredUserGroupMember {
group_id: row.try_get("group_id").map_sql_err()?,
user_id: row.try_get("user_id").map_sql_err()?,
username: row.try_get("username").map_sql_err()?,
email: row.try_get("email").map_sql_err()?,
role: row.try_get("role").map_sql_err()?,
is_active: row.try_get("is_active").map_sql_err()?,
is_deleted: row.try_get("is_deleted").map_sql_err()?,
created_at: optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
})
}
fn map_user_group_membership_row(
row: &MySqlRow,
) -> Result<StoredUserGroupMembership, DataLayerError> {
Ok(StoredUserGroupMembership {
user_id: row.try_get("user_id").map_sql_err()?,
group_id: row.try_get("group_id").map_sql_err()?,
group_name: row.try_get("group_name").map_sql_err()?,
group_priority: row.try_get("group_priority").map_sql_err()?,
created_at: optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
})
}
fn map_oauth_link_summary_row(

View File

@@ -3,9 +3,11 @@ use futures_util::TryStreamExt;
use sqlx::{PgPool, Postgres, QueryBuilder, Row};
use super::types::{
LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, StoredUserExportRow,
normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord,
StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership,
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
StoredUserSummary, UserExportListQuery, UserExportSummary, UserReadRepository,
StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSummary,
UserReadRepository,
};
use crate::{error::SqlxResultExt, DataLayerError};
@@ -46,9 +48,13 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
model_capability_settings,
is_active
FROM users
@@ -67,9 +73,13 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
model_capability_settings,
is_active
FROM users
@@ -87,9 +97,13 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
model_capability_settings,
is_active
FROM users
@@ -132,9 +146,13 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
model_capability_settings,
is_active
FROM users
@@ -153,8 +171,11 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -174,8 +195,11 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -195,8 +219,11 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -216,8 +243,11 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -237,8 +267,11 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -259,8 +292,11 @@ SELECT
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -550,6 +586,40 @@ SET revoked_at = $2, revoke_reason = $3, updated_at = $2
WHERE user_id = $1 AND revoked_at IS NULL
"#;
const USER_GROUP_COLUMNS: &str = r#"
SELECT
id,
name,
normalized_name,
description,
priority,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
created_at,
updated_at
FROM user_groups
"#;
const USER_GROUP_MEMBER_COLUMNS: &str = r#"
SELECT
user_group_members.group_id,
users.id AS user_id,
users.username,
users.email,
users.role::text AS role,
users.is_active,
users.is_deleted,
user_group_members.created_at
FROM user_group_members
JOIN users ON users.id = user_group_members.user_id
"#;
#[derive(Debug, Clone)]
pub struct SqlxUserReadRepository {
pool: PgPool,
@@ -612,6 +682,275 @@ impl SqlxUserReadRepository {
.await
}
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");
collect_query_rows(builder.build().fetch(&self.pool), map_user_group_row).await
}
pub async fn find_user_group_by_id(
&self,
group_id: &str,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let mut builder = QueryBuilder::<Postgres>::new(USER_GROUP_COLUMNS);
builder
.push(" WHERE id = ")
.push_bind(group_id)
.push(" LIMIT 1");
let row = builder
.build()
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
row.as_ref().map(map_user_group_row).transpose()
}
pub async fn list_user_groups_by_ids(
&self,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
if group_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Postgres>::new(USER_GROUP_COLUMNS);
builder.push(" WHERE id IN (");
{
let mut separated = builder.separated(", ");
for group_id in group_ids {
separated.push_bind(group_id);
}
}
builder.push(") ORDER BY priority DESC, name ASC, id ASC");
collect_query_rows(builder.build().fetch(&self.pool), map_user_group_row).await
}
pub async fn create_user_group(
&self,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let id = uuid::Uuid::new_v4().to_string();
let name = normalize_user_group_name(&record.name);
let normalized_name = name.to_ascii_lowercase();
let result = sqlx::query(
r#"
INSERT INTO user_groups (
id, name, normalized_name, description, priority,
allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode,
allowed_models, allowed_models_mode,
rate_limit, rate_limit_mode
)
VALUES ($1, $2, $3, $4, $5, $6::json, $7, $8::json, $9, $10::json, $11, $12, $13)
"#,
)
.bind(&id)
.bind(name)
.bind(normalized_name)
.bind(record.description)
.bind(record.priority)
.bind(record.allowed_providers.map(serde_json::Value::from))
.bind(record.allowed_providers_mode)
.bind(record.allowed_api_formats.map(serde_json::Value::from))
.bind(record.allowed_api_formats_mode)
.bind(record.allowed_models.map(serde_json::Value::from))
.bind(record.allowed_models_mode)
.bind(record.rate_limit)
.bind(record.rate_limit_mode)
.execute(&self.pool)
.await;
match result {
Ok(_) => self.find_user_group_by_id(&id).await,
Err(sqlx::Error::Database(err)) if err.is_unique_violation() => Err(
DataLayerError::InvalidInput("duplicate user group name".to_string()),
),
Err(err) => Err(err).map_postgres_err(),
}
}
pub async fn update_user_group(
&self,
group_id: &str,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let name = normalize_user_group_name(&record.name);
let normalized_name = name.to_ascii_lowercase();
let result = sqlx::query(
r#"
UPDATE user_groups
SET name = $2,
normalized_name = $3,
description = $4,
priority = $5,
allowed_providers = $6::json,
allowed_providers_mode = $7,
allowed_api_formats = $8::json,
allowed_api_formats_mode = $9,
allowed_models = $10::json,
allowed_models_mode = $11,
rate_limit = $12,
rate_limit_mode = $13,
updated_at = now()
WHERE id = $1
"#,
)
.bind(group_id)
.bind(name)
.bind(normalized_name)
.bind(record.description)
.bind(record.priority)
.bind(record.allowed_providers.map(serde_json::Value::from))
.bind(record.allowed_providers_mode)
.bind(record.allowed_api_formats.map(serde_json::Value::from))
.bind(record.allowed_api_formats_mode)
.bind(record.allowed_models.map(serde_json::Value::from))
.bind(record.allowed_models_mode)
.bind(record.rate_limit)
.bind(record.rate_limit_mode)
.execute(&self.pool)
.await;
match result {
Ok(result) if result.rows_affected() == 0 => Ok(None),
Ok(_) => self.find_user_group_by_id(group_id).await,
Err(sqlx::Error::Database(err)) if err.is_unique_violation() => Err(
DataLayerError::InvalidInput("duplicate user group name".to_string()),
),
Err(err) => Err(err).map_postgres_err(),
}
}
pub async fn delete_user_group(&self, group_id: &str) -> Result<bool, DataLayerError> {
let result = sqlx::query("DELETE FROM user_groups WHERE id = $1")
.bind(group_id)
.execute(&self.pool)
.await
.map_postgres_err()?;
Ok(result.rows_affected() > 0)
}
pub async fn list_user_group_members(
&self,
group_id: &str,
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let mut builder = QueryBuilder::<Postgres>::new(USER_GROUP_MEMBER_COLUMNS);
builder
.push(" WHERE user_group_members.group_id = ")
.push_bind(group_id)
.push(" ORDER BY users.username ASC, users.id ASC");
collect_query_rows(builder.build().fetch(&self.pool), map_user_group_member_row).await
}
pub async fn replace_user_group_members(
&self,
group_id: &str,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let mut tx = self.pool.begin().await.map_postgres_err()?;
sqlx::query("DELETE FROM user_group_members WHERE group_id = $1")
.bind(group_id)
.execute(&mut *tx)
.await
.map_postgres_err()?;
for user_id in normalized_ids(user_ids) {
sqlx::query(
"INSERT INTO user_group_members (group_id, user_id) VALUES ($1, $2) ON CONFLICT (group_id, user_id) DO NOTHING",
)
.bind(group_id)
.bind(user_id)
.execute(&mut *tx)
.await
.map_postgres_err()?;
}
tx.commit().await.map_postgres_err()?;
self.list_user_group_members(group_id).await
}
pub async fn list_user_groups_for_user(
&self,
user_id: &str,
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut builder = QueryBuilder::<Postgres>::new(USER_GROUP_COLUMNS);
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");
collect_query_rows(builder.build().fetch(&self.pool), map_user_group_row).await
}
pub async fn list_user_group_memberships_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMembership>, DataLayerError> {
if user_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Postgres>::new(
r#"
SELECT
user_group_members.user_id,
user_groups.id AS group_id,
user_groups.name AS group_name,
user_groups.priority AS group_priority,
user_group_members.created_at
FROM user_group_members
JOIN user_groups ON user_groups.id = user_group_members.group_id
WHERE user_group_members.user_id IN (
"#,
);
{
let mut separated = builder.separated(", ");
for user_id in user_ids {
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");
collect_query_rows(
builder.build().fetch(&self.pool),
map_user_group_membership_row,
)
.await
}
pub async fn replace_user_groups_for_user(
&self,
user_id: &str,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut tx = self.pool.begin().await.map_postgres_err()?;
sqlx::query("DELETE FROM user_group_members WHERE user_id = $1")
.bind(user_id)
.execute(&mut *tx)
.await
.map_postgres_err()?;
for group_id in normalized_ids(group_ids) {
sqlx::query(
"INSERT INTO user_group_members (group_id, user_id) VALUES ($1, $2) ON CONFLICT (group_id, user_id) DO NOTHING",
)
.bind(group_id)
.bind(user_id)
.execute(&mut *tx)
.await
.map_postgres_err()?;
}
tx.commit().await.map_postgres_err()?;
self.list_user_groups_for_user(user_id).await
}
pub async fn add_user_to_group(
&self,
group_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
let result = sqlx::query(
"INSERT INTO user_group_members (group_id, user_id) VALUES ($1, $2) ON CONFLICT (group_id, user_id) DO NOTHING",
)
.bind(group_id)
.bind(user_id)
.execute(&self.pool)
.await
.map_postgres_err()?;
Ok(result.rows_affected() > 0)
}
pub async fn list_export_users_page(
&self,
query: &UserExportListQuery,
@@ -626,6 +965,16 @@ impl SqlxUserReadRepository {
if let Some(is_active) = query.is_active {
builder.push(" AND is_active = ").push_bind(is_active);
}
if let Some(group_id) = query
.group_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
builder.push(" AND id IN (SELECT user_id FROM user_group_members WHERE group_id = ");
builder.push_bind(group_id);
builder.push(")");
}
if let Some(search) = query
.search
.as_deref()
@@ -817,10 +1166,12 @@ impl SqlxUserReadRepository {
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
is_active, is_deleted, created_at, updated_at, last_login_at
)
VALUES (
$1, $2, TRUE, $3, NULL, 'user'::userrole, 'oauth'::authsource,
'inherit', 'inherit', 'inherit', 'inherit',
TRUE, FALSE, $4, $4, $4
)
"#,
@@ -963,8 +1314,9 @@ SET email = $2,
WHERE id = $1
RETURNING
id, email, email_verified, username, password_hash, role::text AS role,
auth_source::text AS auth_source, allowed_providers, allowed_api_formats,
allowed_models, is_active, is_deleted, created_at, last_login_at
auth_source::text AS auth_source, allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode, allowed_models, allowed_models_mode,
is_active, is_deleted, created_at, last_login_at
"#,
)
.bind(&existing.id)
@@ -1014,8 +1366,9 @@ INSERT INTO users (
VALUES ($1, $2, TRUE, $3, NULL, 'user'::userrole, 'ldap'::authsource, $4, $5, TRUE, FALSE, $6, $6, $6)
RETURNING
id, email, email_verified, username, password_hash, role::text AS role,
auth_source::text AS auth_source, allowed_providers, allowed_api_formats,
allowed_models, is_active, is_deleted, created_at, last_login_at
auth_source::text AS auth_source, allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode, allowed_models, allowed_models_mode,
is_active, is_deleted, created_at, last_login_at
"#,
)
.bind(uuid::Uuid::new_v4().to_string())
@@ -1127,6 +1480,26 @@ WHERE id = $1
rate_limit: Option<i32>,
is_active: Option<bool>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let allowed_providers_mode = if allowed_providers.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_api_formats_mode = if allowed_api_formats.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_models_mode = if allowed_models.is_some() {
"specific"
} else {
"unrestricted"
};
let rate_limit_mode = if rate_limit.is_some() {
"custom"
} else {
"system"
};
let result = sqlx::query(
r#"
UPDATE users
@@ -1138,20 +1511,36 @@ SET role = CASE
WHEN $4::BOOLEAN THEN $5::json
ELSE allowed_providers
END,
allowed_providers_mode = CASE
WHEN $4::BOOLEAN THEN $6
ELSE allowed_providers_mode
END,
allowed_api_formats = CASE
WHEN $6::BOOLEAN THEN $7::json
WHEN $7::BOOLEAN THEN $8::json
ELSE allowed_api_formats
END,
allowed_api_formats_mode = CASE
WHEN $7::BOOLEAN THEN $9
ELSE allowed_api_formats_mode
END,
allowed_models = CASE
WHEN $8::BOOLEAN THEN $9::json
WHEN $10::BOOLEAN THEN $11::json
ELSE allowed_models
END,
allowed_models_mode = CASE
WHEN $10::BOOLEAN THEN $12
ELSE allowed_models_mode
END,
rate_limit = CASE
WHEN $10::BOOLEAN THEN $11
WHEN $13::BOOLEAN THEN $14
ELSE rate_limit
END,
rate_limit_mode = CASE
WHEN $13::BOOLEAN THEN $15
ELSE rate_limit_mode
END,
is_active = CASE
WHEN $12::BOOLEAN AND $13 IS NOT NULL THEN $13
WHEN $16::BOOLEAN AND $17 IS NOT NULL THEN $17
ELSE is_active
END,
updated_at = NOW()
@@ -1163,12 +1552,16 @@ WHERE id = $1
.bind(role)
.bind(allowed_providers_present)
.bind(allowed_providers.map(serde_json::Value::from))
.bind(allowed_providers_mode)
.bind(allowed_api_formats_present)
.bind(allowed_api_formats.map(serde_json::Value::from))
.bind(allowed_api_formats_mode)
.bind(allowed_models_present)
.bind(allowed_models.map(serde_json::Value::from))
.bind(allowed_models_mode)
.bind(rate_limit_present)
.bind(rate_limit)
.bind(rate_limit_mode)
.bind(is_active.is_some())
.bind(is_active)
.execute(&self.pool)
@@ -1180,6 +1573,55 @@ WHERE id = $1
self.find_user_auth_by_id(user_id).await
}
pub async fn update_local_auth_user_policy_modes(
&self,
user_id: &str,
allowed_providers_mode: Option<String>,
allowed_api_formats_mode: Option<String>,
allowed_models_mode: Option<String>,
rate_limit_mode: Option<String>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let result = sqlx::query(
r#"
UPDATE users
SET allowed_providers_mode = CASE
WHEN $2::BOOLEAN THEN $3
ELSE allowed_providers_mode
END,
allowed_api_formats_mode = CASE
WHEN $4::BOOLEAN THEN $5
ELSE allowed_api_formats_mode
END,
allowed_models_mode = CASE
WHEN $6::BOOLEAN THEN $7
ELSE allowed_models_mode
END,
rate_limit_mode = CASE
WHEN $8::BOOLEAN THEN $9
ELSE rate_limit_mode
END,
updated_at = NOW()
WHERE id = $1
"#,
)
.bind(user_id)
.bind(allowed_providers_mode.is_some())
.bind(allowed_providers_mode)
.bind(allowed_api_formats_mode.is_some())
.bind(allowed_api_formats_mode)
.bind(allowed_models_mode.is_some())
.bind(allowed_models_mode)
.bind(rate_limit_mode.is_some())
.bind(rate_limit_mode)
.execute(&self.pool)
.await
.map_postgres_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.find_user_auth_by_id(user_id).await
}
pub async fn update_user_model_capability_settings(
&self,
user_id: &str,
@@ -1212,18 +1654,30 @@ WHERE id = $1
username: String,
password_hash: String,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
self.create_local_auth_user_with_settings(
email,
email_verified,
username,
password_hash,
"user".to_string(),
None,
None,
None,
None,
let user_id = uuid::Uuid::new_v4().to_string();
sqlx::query(
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
is_active, is_deleted, created_at, updated_at
)
VALUES (
$1, $2, $3, $4, $5, 'user'::userrole, 'local'::authsource,
'inherit', 'inherit', 'inherit', 'inherit',
TRUE, FALSE, NOW(), NOW()
)
"#,
)
.bind(&user_id)
.bind(email)
.bind(email_verified)
.bind(username)
.bind(password_hash)
.execute(&self.pool)
.await
.map_postgres_err()?;
self.find_user_auth_by_id(&user_id).await
}
#[allow(clippy::too_many_arguments)]
@@ -1240,16 +1694,39 @@ WHERE id = $1
rate_limit: Option<i32>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let user_id = uuid::Uuid::new_v4().to_string();
let allowed_providers_mode = if allowed_providers.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_api_formats_mode = if allowed_api_formats.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_models_mode = if allowed_models.is_some() {
"specific"
} else {
"unrestricted"
};
let rate_limit_mode = if rate_limit.is_some() {
"custom"
} else {
"system"
};
sqlx::query(
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers, allowed_api_formats, allowed_models, rate_limit,
allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode,
allowed_models, allowed_models_mode,
rate_limit, rate_limit_mode,
is_active, is_deleted, created_at, updated_at
)
VALUES (
$1, $2, $3, $4, $5, $6::userrole, 'local'::authsource,
$7::json, $8::json, $9::json, $10,
$7::json, $8, $9::json, $10, $11::json, $12, $13, $14,
TRUE, FALSE, NOW(), NOW()
)
"#,
@@ -1261,9 +1738,13 @@ VALUES (
.bind(password_hash)
.bind(role)
.bind(allowed_providers.map(serde_json::Value::from))
.bind(allowed_providers_mode)
.bind(allowed_api_formats.map(serde_json::Value::from))
.bind(allowed_api_formats_mode)
.bind(allowed_models.map(serde_json::Value::from))
.bind(allowed_models_mode)
.bind(rate_limit)
.bind(rate_limit_mode)
.execute(&self.pool)
.await
.map_postgres_err()?;
@@ -1541,6 +2022,16 @@ fn normalize_optional_json_value(value: Option<serde_json::Value>) -> Option<ser
}
}
fn normalized_ids(values: &[String]) -> Vec<String> {
values
.iter()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.collect::<std::collections::BTreeSet<_>>()
.into_iter()
.collect()
}
async fn find_postgres_ldap_user_for_update(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
ldap_dn: Option<&str>,
@@ -1550,8 +2041,9 @@ async fn find_postgres_ldap_user_for_update(
let select_columns = r#"
SELECT
id, email, email_verified, username, password_hash, role::text AS role,
auth_source::text AS auth_source, allowed_providers, allowed_api_formats,
allowed_models, is_active, is_deleted, created_at, last_login_at
auth_source::text AS auth_source, allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode, allowed_models, allowed_models_mode,
is_active, is_deleted, created_at, last_login_at
FROM users
"#;
if let Some(ldap_dn) = ldap_dn.filter(|value| !value.trim().is_empty()) {
@@ -1616,6 +2108,14 @@ fn map_user_export_row(row: &sqlx::postgres::PgRow) -> Result<StoredUserExportRo
.map_postgres_err()?,
row.try_get("is_active").map_postgres_err()?,
)
.and_then(|record| {
record.with_policy_modes(
row.try_get("allowed_providers_mode").map_postgres_err()?,
row.try_get("allowed_api_formats_mode").map_postgres_err()?,
row.try_get("allowed_models_mode").map_postgres_err()?,
row.try_get("rate_limit_mode").map_postgres_err()?,
)
})
}
fn map_user_auth_row(row: &sqlx::postgres::PgRow) -> Result<StoredUserAuthRecord, DataLayerError> {
@@ -1635,6 +2135,60 @@ fn map_user_auth_row(row: &sqlx::postgres::PgRow) -> Result<StoredUserAuthRecord
row.try_get("created_at").map_postgres_err()?,
row.try_get("last_login_at").map_postgres_err()?,
)
.and_then(|record| {
record.with_policy_modes(
row.try_get("allowed_providers_mode").map_postgres_err()?,
row.try_get("allowed_api_formats_mode").map_postgres_err()?,
row.try_get("allowed_models_mode").map_postgres_err()?,
)
})
}
fn map_user_group_row(row: &sqlx::postgres::PgRow) -> Result<StoredUserGroup, DataLayerError> {
StoredUserGroup::new(
row.try_get("id").map_postgres_err()?,
row.try_get("name").map_postgres_err()?,
row.try_get("normalized_name").map_postgres_err()?,
row.try_get("description").map_postgres_err()?,
row.try_get("priority").map_postgres_err()?,
row.try_get("allowed_providers").map_postgres_err()?,
row.try_get("allowed_providers_mode").map_postgres_err()?,
row.try_get("allowed_api_formats").map_postgres_err()?,
row.try_get("allowed_api_formats_mode").map_postgres_err()?,
row.try_get("allowed_models").map_postgres_err()?,
row.try_get("allowed_models_mode").map_postgres_err()?,
row.try_get("rate_limit").map_postgres_err()?,
row.try_get("rate_limit_mode").map_postgres_err()?,
row.try_get("created_at").map_postgres_err()?,
row.try_get("updated_at").map_postgres_err()?,
)
}
fn map_user_group_member_row(
row: &sqlx::postgres::PgRow,
) -> Result<StoredUserGroupMember, DataLayerError> {
Ok(StoredUserGroupMember {
group_id: row.try_get("group_id").map_postgres_err()?,
user_id: row.try_get("user_id").map_postgres_err()?,
username: row.try_get("username").map_postgres_err()?,
email: row.try_get("email").map_postgres_err()?,
role: row.try_get("role").map_postgres_err()?,
is_active: row.try_get("is_active").map_postgres_err()?,
is_deleted: row.try_get("is_deleted").map_postgres_err()?,
created_at: row.try_get("created_at").map_postgres_err()?,
})
}
fn map_user_group_membership_row(
row: &sqlx::postgres::PgRow,
) -> Result<StoredUserGroupMembership, DataLayerError> {
Ok(StoredUserGroupMembership {
user_id: row.try_get("user_id").map_postgres_err()?,
group_id: row.try_get("group_id").map_postgres_err()?,
group_name: row.try_get("group_name").map_postgres_err()?,
group_priority: row.try_get("group_priority").map_postgres_err()?,
created_at: row.try_get("created_at").map_postgres_err()?,
})
}
fn map_oauth_link_summary_row(
@@ -1709,6 +2263,88 @@ impl UserReadRepository for SqlxUserReadRepository {
self.find_export_user_by_id(user_id).await
}
async fn list_user_groups(&self) -> Result<Vec<StoredUserGroup>, DataLayerError> {
self.list_user_groups().await
}
async fn find_user_group_by_id(
&self,
group_id: &str,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
self.find_user_group_by_id(group_id).await
}
async fn list_user_groups_by_ids(
&self,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
self.list_user_groups_by_ids(group_ids).await
}
async fn create_user_group(
&self,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
self.create_user_group(record).await
}
async fn update_user_group(
&self,
group_id: &str,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
self.update_user_group(group_id, record).await
}
async fn delete_user_group(&self, group_id: &str) -> Result<bool, DataLayerError> {
self.delete_user_group(group_id).await
}
async fn list_user_group_members(
&self,
group_id: &str,
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
self.list_user_group_members(group_id).await
}
async fn replace_user_group_members(
&self,
group_id: &str,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
self.replace_user_group_members(group_id, user_ids).await
}
async fn list_user_groups_for_user(
&self,
user_id: &str,
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
self.list_user_groups_for_user(user_id).await
}
async fn list_user_group_memberships_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMembership>, DataLayerError> {
self.list_user_group_memberships_by_user_ids(user_ids).await
}
async fn replace_user_groups_for_user(
&self,
user_id: &str,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
self.replace_user_groups_for_user(user_id, group_ids).await
}
async fn add_user_to_group(
&self,
group_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
self.add_user_to_group(group_id, user_id).await
}
async fn find_user_auth_by_id(
&self,
user_id: &str,
@@ -1919,6 +2555,24 @@ impl UserReadRepository for SqlxUserReadRepository {
.await
}
async fn update_local_auth_user_policy_modes(
&self,
user_id: &str,
allowed_providers_mode: Option<String>,
allowed_api_formats_mode: Option<String>,
allowed_models_mode: Option<String>,
rate_limit_mode: Option<String>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
self.update_local_auth_user_policy_modes(
user_id,
allowed_providers_mode,
allowed_api_formats_mode,
allowed_models_mode,
rate_limit_mode,
)
.await
}
async fn update_user_model_capability_settings(
&self,
user_id: &str,

View File

@@ -3,9 +3,11 @@ use chrono::{DateTime, TimeZone, Utc};
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use super::types::{
LdapAuthUserProvisioningOutcome, StoredUserAuthRecord, StoredUserExportRow,
normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord,
StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership,
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
StoredUserSummary, UserExportListQuery, UserExportSummary, UserReadRepository,
StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSummary,
UserReadRepository,
};
use crate::driver::sqlite::SqlitePool;
use crate::error::SqlResultExt;
@@ -32,9 +34,13 @@ SELECT
role,
auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
model_capability_settings,
is_active
FROM users
@@ -50,8 +56,11 @@ SELECT
role,
auth_source,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
is_active,
is_deleted,
created_at,
@@ -69,8 +78,11 @@ SELECT
users.role AS role,
users.auth_source AS auth_source,
users.allowed_providers AS allowed_providers,
users.allowed_providers_mode AS allowed_providers_mode,
users.allowed_api_formats AS allowed_api_formats,
users.allowed_api_formats_mode AS allowed_api_formats_mode,
users.allowed_models AS allowed_models,
users.allowed_models_mode AS allowed_models_mode,
users.is_active AS is_active,
users.is_deleted AS is_deleted,
users.created_at AS created_at,
@@ -130,6 +142,40 @@ SELECT
FROM user_sessions
"#;
const USER_GROUP_COLUMNS: &str = r#"
SELECT
id,
name,
normalized_name,
description,
priority,
allowed_providers,
allowed_providers_mode,
allowed_api_formats,
allowed_api_formats_mode,
allowed_models,
allowed_models_mode,
rate_limit,
rate_limit_mode,
created_at,
updated_at
FROM user_groups
"#;
const USER_GROUP_MEMBER_COLUMNS: &str = r#"
SELECT
user_group_members.group_id,
users.id AS user_id,
users.username,
users.email,
users.role,
users.is_active,
users.is_deleted,
user_group_members.created_at
FROM user_group_members
JOIN users ON users.id = user_group_members.user_id
"#;
#[derive(Debug, Clone)]
pub struct SqliteUserReadRepository {
pool: SqlitePool,
@@ -163,6 +209,22 @@ impl SqliteUserReadRepository {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_auth_row).collect()
}
async fn fetch_group_rows(
&self,
mut builder: QueryBuilder<'_, Sqlite>,
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_group_row).collect()
}
async fn fetch_group_member_rows(
&self,
mut builder: QueryBuilder<'_, Sqlite>,
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_group_member_row).collect()
}
}
#[async_trait]
@@ -224,6 +286,16 @@ impl UserReadRepository for SqliteUserReadRepository {
if let Some(is_active) = query.is_active {
builder.push(" AND is_active = ").push_bind(is_active);
}
if let Some(group_id) = query
.group_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
builder.push(" AND id IN (SELECT user_id FROM user_group_members WHERE group_id = ");
builder.push_bind(group_id);
builder.push(")");
}
if let Some(search) = query
.search
.as_deref()
@@ -296,6 +368,285 @@ WHERE is_deleted = 0
self.fetch_export_rows(builder).await
}
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");
self.fetch_group_rows(builder).await
}
async fn find_user_group_by_id(
&self,
group_id: &str,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(USER_GROUP_COLUMNS);
builder
.push(" WHERE id = ")
.push_bind(group_id)
.push(" LIMIT 1");
Ok(self.fetch_group_rows(builder).await?.into_iter().next())
}
async fn list_user_groups_by_ids(
&self,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
if group_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Sqlite>::new(USER_GROUP_COLUMNS);
builder.push(" WHERE id IN (");
{
let mut separated = builder.separated(", ");
for group_id in group_ids {
separated.push_bind(group_id);
}
}
builder.push(") ORDER BY priority DESC, name ASC, id ASC");
self.fetch_group_rows(builder).await
}
async fn create_user_group(
&self,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let now = current_unix_secs();
let id = uuid::Uuid::new_v4().to_string();
let name = normalize_user_group_name(&record.name);
let normalized_name = name.to_ascii_lowercase();
let result = sqlx::query(
r#"
INSERT INTO user_groups (
id, name, normalized_name, description, priority,
allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode,
allowed_models, allowed_models_mode,
rate_limit, rate_limit_mode, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&id)
.bind(name)
.bind(normalized_name)
.bind(record.description)
.bind(record.priority)
.bind(json_string_from_option_vec(
record.allowed_providers.as_ref(),
))
.bind(record.allowed_providers_mode)
.bind(json_string_from_option_vec(
record.allowed_api_formats.as_ref(),
))
.bind(record.allowed_api_formats_mode)
.bind(json_string_from_option_vec(record.allowed_models.as_ref()))
.bind(record.allowed_models_mode)
.bind(record.rate_limit)
.bind(record.rate_limit_mode)
.bind(now)
.bind(now)
.execute(&self.pool)
.await;
match result {
Ok(_) => self.find_user_group_by_id(&id).await,
Err(sqlx::Error::Database(err)) if err.is_unique_violation() => Err(
DataLayerError::InvalidInput("duplicate user group name".to_string()),
),
Err(err) => Err(err).map_sql_err(),
}
}
async fn update_user_group(
&self,
group_id: &str,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, DataLayerError> {
let now = current_unix_secs();
let name = normalize_user_group_name(&record.name);
let normalized_name = name.to_ascii_lowercase();
let result = sqlx::query(
r#"
UPDATE user_groups
SET name = ?,
normalized_name = ?,
description = ?,
priority = ?,
allowed_providers = ?,
allowed_providers_mode = ?,
allowed_api_formats = ?,
allowed_api_formats_mode = ?,
allowed_models = ?,
allowed_models_mode = ?,
rate_limit = ?,
rate_limit_mode = ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(name)
.bind(normalized_name)
.bind(record.description)
.bind(record.priority)
.bind(json_string_from_option_vec(
record.allowed_providers.as_ref(),
))
.bind(record.allowed_providers_mode)
.bind(json_string_from_option_vec(
record.allowed_api_formats.as_ref(),
))
.bind(record.allowed_api_formats_mode)
.bind(json_string_from_option_vec(record.allowed_models.as_ref()))
.bind(record.allowed_models_mode)
.bind(record.rate_limit)
.bind(record.rate_limit_mode)
.bind(now)
.bind(group_id)
.execute(&self.pool)
.await;
match result {
Ok(result) if result.rows_affected() == 0 => Ok(None),
Ok(_) => self.find_user_group_by_id(group_id).await,
Err(sqlx::Error::Database(err)) if err.is_unique_violation() => Err(
DataLayerError::InvalidInput("duplicate user group name".to_string()),
),
Err(err) => Err(err).map_sql_err(),
}
}
async fn delete_user_group(&self, group_id: &str) -> Result<bool, DataLayerError> {
let result = sqlx::query("DELETE FROM user_groups WHERE id = ?")
.bind(group_id)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
async fn list_user_group_members(
&self,
group_id: &str,
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(USER_GROUP_MEMBER_COLUMNS);
builder
.push(" WHERE user_group_members.group_id = ")
.push_bind(group_id)
.push(" ORDER BY users.username ASC, users.id ASC");
self.fetch_group_member_rows(builder).await
}
async fn replace_user_group_members(
&self,
group_id: &str,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMember>, DataLayerError> {
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("DELETE FROM user_group_members WHERE group_id = ?")
.bind(group_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
let now = current_unix_secs();
for user_id in normalized_ids(user_ids) {
sqlx::query(
"INSERT OR IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)",
)
.bind(group_id)
.bind(user_id)
.bind(now)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
tx.commit().await.map_sql_err()?;
self.list_user_group_members(group_id).await
}
async fn list_user_groups_for_user(
&self,
user_id: &str,
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new(USER_GROUP_COLUMNS);
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");
self.fetch_group_rows(builder).await
}
async fn list_user_group_memberships_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMembership>, DataLayerError> {
if user_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Sqlite>::new(
r#"
SELECT
user_group_members.user_id,
user_groups.id AS group_id,
user_groups.name AS group_name,
user_groups.priority AS group_priority,
user_group_members.created_at
FROM user_group_members
JOIN user_groups ON user_groups.id = user_group_members.group_id
WHERE user_group_members.user_id IN (
"#,
);
{
let mut separated = builder.separated(", ");
for user_id in user_ids {
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");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_user_group_membership_row).collect()
}
async fn replace_user_groups_for_user(
&self,
user_id: &str,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, DataLayerError> {
let mut tx = self.pool.begin().await.map_sql_err()?;
sqlx::query("DELETE FROM user_group_members WHERE user_id = ?")
.bind(user_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
let now = current_unix_secs();
for group_id in normalized_ids(group_ids) {
sqlx::query(
"INSERT OR IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)",
)
.bind(group_id)
.bind(user_id)
.bind(now)
.execute(&mut *tx)
.await
.map_sql_err()?;
}
tx.commit().await.map_sql_err()?;
self.list_user_groups_for_user(user_id).await
}
async fn add_user_to_group(
&self,
group_id: &str,
user_id: &str,
) -> Result<bool, DataLayerError> {
let result = sqlx::query(
"INSERT OR IGNORE INTO user_group_members (group_id, user_id, created_at) VALUES (?, ?, ?)",
)
.bind(group_id)
.bind(user_id)
.bind(current_unix_secs())
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
async fn find_user_auth_by_id(
&self,
user_id: &str,
@@ -453,9 +804,10 @@ WHERE provider_type = ?
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
is_active, is_deleted, created_at, updated_at, last_login_at
)
VALUES (?, ?, 1, ?, NULL, 'user', 'oauth', 1, 0, ?, ?, ?)
VALUES (?, ?, 1, ?, NULL, 'user', 'oauth', 'inherit', 'inherit', 'inherit', 'inherit', 1, 0, ?, ?, ?)
"#,
)
.bind(&user_id)
@@ -675,14 +1027,38 @@ VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
rate_limit: Option<i32>,
is_active: Option<bool>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let allowed_providers_mode = if allowed_providers.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_api_formats_mode = if allowed_api_formats.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_models_mode = if allowed_models.is_some() {
"specific"
} else {
"unrestricted"
};
let rate_limit_mode = if rate_limit.is_some() {
"custom"
} else {
"system"
};
let result = sqlx::query(
r#"
UPDATE users
SET role = CASE WHEN ? THEN COALESCE(?, role) ELSE role END,
allowed_providers = CASE WHEN ? THEN ? ELSE allowed_providers END,
allowed_providers_mode = CASE WHEN ? THEN ? ELSE allowed_providers_mode END,
allowed_api_formats = CASE WHEN ? THEN ? ELSE allowed_api_formats END,
allowed_api_formats_mode = CASE WHEN ? THEN ? ELSE allowed_api_formats_mode END,
allowed_models = CASE WHEN ? THEN ? ELSE allowed_models END,
allowed_models_mode = CASE WHEN ? THEN ? ELSE allowed_models_mode END,
rate_limit = CASE WHEN ? THEN ? ELSE rate_limit END,
rate_limit_mode = CASE WHEN ? THEN ? ELSE rate_limit_mode END,
is_active = CASE WHEN ? THEN ? ELSE is_active END,
updated_at = ?
WHERE id = ?
@@ -695,18 +1071,26 @@ WHERE id = ?
allowed_providers,
"users.allowed_providers",
)?)
.bind(allowed_providers_present)
.bind(allowed_providers_mode)
.bind(allowed_api_formats_present)
.bind(optional_string_list_json(
allowed_api_formats,
"users.allowed_api_formats",
)?)
.bind(allowed_api_formats_present)
.bind(allowed_api_formats_mode)
.bind(allowed_models_present)
.bind(optional_string_list_json(
allowed_models,
"users.allowed_models",
)?)
.bind(allowed_models_present)
.bind(allowed_models_mode)
.bind(rate_limit_present)
.bind(rate_limit)
.bind(rate_limit_present)
.bind(rate_limit_mode)
.bind(is_active.is_some())
.bind(is_active)
.bind(chrono::Utc::now().timestamp())
@@ -720,6 +1104,44 @@ WHERE id = ?
self.find_user_auth_by_id(user_id).await
}
async fn update_local_auth_user_policy_modes(
&self,
user_id: &str,
allowed_providers_mode: Option<String>,
allowed_api_formats_mode: Option<String>,
allowed_models_mode: Option<String>,
rate_limit_mode: Option<String>,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let result = sqlx::query(
r#"
UPDATE users
SET allowed_providers_mode = CASE WHEN ? THEN ? ELSE allowed_providers_mode END,
allowed_api_formats_mode = CASE WHEN ? THEN ? ELSE allowed_api_formats_mode END,
allowed_models_mode = CASE WHEN ? THEN ? ELSE allowed_models_mode END,
rate_limit_mode = CASE WHEN ? THEN ? ELSE rate_limit_mode END,
updated_at = ?
WHERE id = ?
"#,
)
.bind(allowed_providers_mode.is_some())
.bind(allowed_providers_mode)
.bind(allowed_api_formats_mode.is_some())
.bind(allowed_api_formats_mode)
.bind(allowed_models_mode.is_some())
.bind(allowed_models_mode)
.bind(rate_limit_mode.is_some())
.bind(rate_limit_mode)
.bind(chrono::Utc::now().timestamp())
.bind(user_id)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.find_user_auth_by_id(user_id).await
}
async fn update_user_model_capability_settings(
&self,
user_id: &str,
@@ -751,18 +1173,29 @@ WHERE id = ?
username: String,
password_hash: String,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
self.create_local_auth_user_with_settings(
email,
email_verified,
username,
password_hash,
"user".to_string(),
None,
None,
None,
None,
let user_id = uuid::Uuid::new_v4().to_string();
let now = chrono::Utc::now().timestamp();
sqlx::query(
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
is_active, is_deleted, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, 'user', 'local', 'inherit', 'inherit', 'inherit', 'inherit', 1, 0, ?, ?)
"#,
)
.bind(&user_id)
.bind(email)
.bind(email_verified)
.bind(username)
.bind(password_hash)
.bind(now)
.bind(now)
.execute(&self.pool)
.await
.map_sql_err()?;
self.find_user_auth_by_id(&user_id).await
}
async fn create_local_auth_user_with_settings(
@@ -779,14 +1212,37 @@ WHERE id = ?
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let user_id = uuid::Uuid::new_v4().to_string();
let now = chrono::Utc::now().timestamp();
let allowed_providers_mode = if allowed_providers.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_api_formats_mode = if allowed_api_formats.is_some() {
"specific"
} else {
"unrestricted"
};
let allowed_models_mode = if allowed_models.is_some() {
"specific"
} else {
"unrestricted"
};
let rate_limit_mode = if rate_limit.is_some() {
"custom"
} else {
"system"
};
sqlx::query(
r#"
INSERT INTO users (
id, email, email_verified, username, password_hash, role, auth_source,
allowed_providers, allowed_api_formats, allowed_models, rate_limit,
allowed_providers, allowed_providers_mode,
allowed_api_formats, allowed_api_formats_mode,
allowed_models, allowed_models_mode,
rate_limit, rate_limit_mode,
is_active, is_deleted, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, 1, 0, ?, ?)
VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, ?, ?, ?, ?, 1, 0, ?, ?)
"#,
)
.bind(&user_id)
@@ -799,15 +1255,19 @@ VALUES (?, ?, ?, ?, ?, ?, 'local', ?, ?, ?, ?, 1, 0, ?, ?)
allowed_providers,
"users.allowed_providers",
)?)
.bind(allowed_providers_mode)
.bind(optional_string_list_json(
allowed_api_formats,
"users.allowed_api_formats",
)?)
.bind(allowed_api_formats_mode)
.bind(optional_string_list_json(
allowed_models,
"users.allowed_models",
)?)
.bind(allowed_models_mode)
.bind(rate_limit)
.bind(rate_limit_mode)
.bind(now)
.bind(now)
.execute(&self.pool)
@@ -1166,6 +1626,24 @@ fn optional_string_list_json(
.transpose()
}
fn json_string_from_option_vec(value: Option<&Vec<String>>) -> Option<String> {
value.and_then(|items| serde_json::to_string(items).ok())
}
fn normalized_ids(values: &[String]) -> Vec<String> {
values
.iter()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.collect::<std::collections::BTreeSet<_>>()
.into_iter()
.collect()
}
fn current_unix_secs() -> i64 {
chrono::Utc::now().timestamp()
}
fn optional_json_string(
value: Option<serde_json::Value>,
field_name: &str,
@@ -1382,6 +1860,14 @@ fn map_user_export_row(row: &SqliteRow) -> Result<StoredUserExportRow, DataLayer
)?,
row.try_get("is_active").map_sql_err()?,
)
.and_then(|record| {
record.with_policy_modes(
row.try_get("allowed_providers_mode").map_sql_err()?,
row.try_get("allowed_api_formats_mode").map_sql_err()?,
row.try_get("allowed_models_mode").map_sql_err()?,
row.try_get("rate_limit_mode").map_sql_err()?,
)
})
}
fn map_user_auth_row(row: &SqliteRow) -> Result<StoredUserAuthRecord, DataLayerError> {
@@ -1410,6 +1896,67 @@ fn map_user_auth_row(row: &SqliteRow) -> Result<StoredUserAuthRecord, DataLayerE
optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
optional_datetime_from_unix_secs(row.try_get("last_login_at").map_sql_err()?),
)
.and_then(|record| {
record.with_policy_modes(
row.try_get("allowed_providers_mode").map_sql_err()?,
row.try_get("allowed_api_formats_mode").map_sql_err()?,
row.try_get("allowed_models_mode").map_sql_err()?,
)
})
}
fn map_user_group_row(row: &SqliteRow) -> Result<StoredUserGroup, DataLayerError> {
StoredUserGroup::new(
row.try_get("id").map_sql_err()?,
row.try_get("name").map_sql_err()?,
row.try_get("normalized_name").map_sql_err()?,
row.try_get("description").map_sql_err()?,
row.try_get("priority").map_sql_err()?,
optional_json_from_string(
row.try_get("allowed_providers").map_sql_err()?,
"user_groups.allowed_providers",
)?,
row.try_get("allowed_providers_mode").map_sql_err()?,
optional_json_from_string(
row.try_get("allowed_api_formats").map_sql_err()?,
"user_groups.allowed_api_formats",
)?,
row.try_get("allowed_api_formats_mode").map_sql_err()?,
optional_json_from_string(
row.try_get("allowed_models").map_sql_err()?,
"user_groups.allowed_models",
)?,
row.try_get("allowed_models_mode").map_sql_err()?,
row.try_get("rate_limit").map_sql_err()?,
row.try_get("rate_limit_mode").map_sql_err()?,
optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
optional_datetime_from_unix_secs(row.try_get("updated_at").map_sql_err()?),
)
}
fn map_user_group_member_row(row: &SqliteRow) -> Result<StoredUserGroupMember, DataLayerError> {
Ok(StoredUserGroupMember {
group_id: row.try_get("group_id").map_sql_err()?,
user_id: row.try_get("user_id").map_sql_err()?,
username: row.try_get("username").map_sql_err()?,
email: row.try_get("email").map_sql_err()?,
role: row.try_get("role").map_sql_err()?,
is_active: row.try_get("is_active").map_sql_err()?,
is_deleted: row.try_get("is_deleted").map_sql_err()?,
created_at: optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
})
}
fn map_user_group_membership_row(
row: &SqliteRow,
) -> Result<StoredUserGroupMembership, DataLayerError> {
Ok(StoredUserGroupMembership {
user_id: row.try_get("user_id").map_sql_err()?,
group_id: row.try_get("group_id").map_sql_err()?,
group_name: row.try_get("group_name").map_sql_err()?,
group_priority: row.try_get("group_priority").map_sql_err()?,
created_at: optional_datetime_from_unix_secs(row.try_get("created_at").map_sql_err()?),
})
}
fn map_oauth_link_summary_row(
@@ -1560,6 +2107,7 @@ INSERT INTO users (
role: Some("user".to_string()),
is_active: Some(true),
search: None,
group_id: None,
})
.await
.expect("export page should load");

View File

@@ -57,8 +57,11 @@ pub struct StoredUserAuthRecord {
pub role: String,
pub auth_source: String,
pub allowed_providers: Option<Vec<String>>,
pub allowed_providers_mode: String,
pub allowed_api_formats: Option<Vec<String>>,
pub allowed_api_formats_mode: String,
pub allowed_models: Option<Vec<String>>,
pub allowed_models_mode: String,
pub is_active: bool,
pub is_deleted: bool,
pub created_at: Option<DateTime<Utc>>,
@@ -113,16 +116,44 @@ impl StoredUserAuthRecord {
role,
auth_source,
allowed_providers: parse_string_list(allowed_providers, "users.allowed_providers")?,
allowed_providers_mode: "unrestricted".to_string(),
allowed_api_formats: parse_string_list(
allowed_api_formats,
"users.allowed_api_formats",
)?,
allowed_api_formats_mode: "unrestricted".to_string(),
allowed_models: parse_string_list(allowed_models, "users.allowed_models")?,
allowed_models_mode: "unrestricted".to_string(),
is_active,
is_deleted,
created_at,
last_login_at,
})
.map(|record| record.with_legacy_policy_modes())
}
pub fn with_policy_modes(
mut self,
allowed_providers_mode: String,
allowed_api_formats_mode: String,
allowed_models_mode: String,
) -> Result<Self, crate::DataLayerError> {
self.allowed_providers_mode =
normalize_list_policy_mode(&allowed_providers_mode, "users.allowed_providers_mode")?;
self.allowed_api_formats_mode = normalize_list_policy_mode(
&allowed_api_formats_mode,
"users.allowed_api_formats_mode",
)?;
self.allowed_models_mode =
normalize_list_policy_mode(&allowed_models_mode, "users.allowed_models_mode")?;
Ok(self)
}
fn with_legacy_policy_modes(mut self) -> Self {
self.allowed_providers_mode = legacy_list_policy_mode(&self.allowed_providers);
self.allowed_api_formats_mode = legacy_list_policy_mode(&self.allowed_api_formats);
self.allowed_models_mode = legacy_list_policy_mode(&self.allowed_models);
self
}
pub fn to_summary(&self) -> Result<StoredUserSummary, crate::DataLayerError> {
@@ -197,9 +228,13 @@ pub struct StoredUserExportRow {
pub role: String,
pub auth_source: String,
pub allowed_providers: Option<Vec<String>>,
pub allowed_providers_mode: String,
pub allowed_api_formats: Option<Vec<String>>,
pub allowed_api_formats_mode: String,
pub allowed_models: Option<Vec<String>>,
pub allowed_models_mode: String,
pub rate_limit: Option<i32>,
pub rate_limit_mode: String,
pub model_capability_settings: Option<Value>,
pub is_active: bool,
}
@@ -251,15 +286,52 @@ impl StoredUserExportRow {
role,
auth_source,
allowed_providers: parse_string_list(allowed_providers, "users.allowed_providers")?,
allowed_providers_mode: "unrestricted".to_string(),
allowed_api_formats: parse_string_list(
allowed_api_formats,
"users.allowed_api_formats",
)?,
allowed_api_formats_mode: "unrestricted".to_string(),
allowed_models: parse_string_list(allowed_models, "users.allowed_models")?,
allowed_models_mode: "unrestricted".to_string(),
rate_limit,
rate_limit_mode: "system".to_string(),
model_capability_settings: normalize_optional_json(model_capability_settings),
is_active,
})
.map(|record| record.with_legacy_policy_modes())
}
pub fn with_policy_modes(
mut self,
allowed_providers_mode: String,
allowed_api_formats_mode: String,
allowed_models_mode: String,
rate_limit_mode: String,
) -> Result<Self, crate::DataLayerError> {
self.allowed_providers_mode =
normalize_list_policy_mode(&allowed_providers_mode, "users.allowed_providers_mode")?;
self.allowed_api_formats_mode = normalize_list_policy_mode(
&allowed_api_formats_mode,
"users.allowed_api_formats_mode",
)?;
self.allowed_models_mode =
normalize_list_policy_mode(&allowed_models_mode, "users.allowed_models_mode")?;
self.rate_limit_mode =
normalize_rate_limit_policy_mode(&rate_limit_mode, "users.rate_limit_mode")?;
Ok(self)
}
fn with_legacy_policy_modes(mut self) -> Self {
self.allowed_providers_mode = legacy_list_policy_mode(&self.allowed_providers);
self.allowed_api_formats_mode = legacy_list_policy_mode(&self.allowed_api_formats);
self.allowed_models_mode = legacy_list_policy_mode(&self.allowed_models);
self.rate_limit_mode = if self.rate_limit.is_some() {
"custom".to_string()
} else {
"system".to_string()
};
self
}
}
@@ -404,6 +476,139 @@ pub struct StoredUserPreferenceRecord {
pub announcement_notifications: bool,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredUserGroup {
pub id: String,
pub name: String,
pub normalized_name: String,
pub description: Option<String>,
pub priority: i32,
pub allowed_providers: Option<Vec<String>>,
pub allowed_providers_mode: String,
pub allowed_api_formats: Option<Vec<String>>,
pub allowed_api_formats_mode: String,
pub allowed_models: Option<Vec<String>>,
pub allowed_models_mode: String,
pub rate_limit: Option<i32>,
pub rate_limit_mode: String,
pub created_at: Option<DateTime<Utc>>,
pub updated_at: Option<DateTime<Utc>>,
}
impl StoredUserGroup {
#[allow(clippy::too_many_arguments)]
pub fn new(
id: String,
name: String,
normalized_name: String,
description: Option<String>,
priority: i32,
allowed_providers: Option<Value>,
allowed_providers_mode: String,
allowed_api_formats: Option<Value>,
allowed_api_formats_mode: String,
allowed_models: Option<Value>,
allowed_models_mode: String,
rate_limit: Option<i32>,
rate_limit_mode: String,
created_at: Option<DateTime<Utc>>,
updated_at: Option<DateTime<Utc>>,
) -> Result<Self, crate::DataLayerError> {
if id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"user_groups.id is empty".to_string(),
));
}
if name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"user_groups.name is empty".to_string(),
));
}
if normalized_name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"user_groups.normalized_name is empty".to_string(),
));
}
Ok(Self {
id,
name,
normalized_name,
description,
priority,
allowed_providers: parse_string_list(
allowed_providers,
"user_groups.allowed_providers",
)?,
allowed_providers_mode: normalize_list_policy_mode(
&allowed_providers_mode,
"user_groups.allowed_providers_mode",
)?,
allowed_api_formats: parse_string_list(
allowed_api_formats,
"user_groups.allowed_api_formats",
)?,
allowed_api_formats_mode: normalize_list_policy_mode(
&allowed_api_formats_mode,
"user_groups.allowed_api_formats_mode",
)?,
allowed_models: parse_string_list(allowed_models, "user_groups.allowed_models")?,
allowed_models_mode: normalize_list_policy_mode(
&allowed_models_mode,
"user_groups.allowed_models_mode",
)?,
rate_limit,
rate_limit_mode: normalize_rate_limit_policy_mode(
&rate_limit_mode,
"user_groups.rate_limit_mode",
)?,
created_at,
updated_at,
})
}
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredUserGroupMember {
pub group_id: String,
pub user_id: String,
pub username: String,
pub email: Option<String>,
pub role: String,
pub is_active: bool,
pub is_deleted: bool,
pub created_at: Option<DateTime<Utc>>,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredUserGroupMembership {
pub user_id: String,
pub group_id: String,
pub group_name: String,
pub group_priority: i32,
pub created_at: Option<DateTime<Utc>>,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct UpsertUserGroupRecord {
pub name: String,
pub description: Option<String>,
pub priority: i32,
pub allowed_providers: Option<Vec<String>>,
pub allowed_providers_mode: String,
pub allowed_api_formats: Option<Vec<String>>,
pub allowed_api_formats_mode: String,
pub allowed_models: Option<Vec<String>>,
pub allowed_models_mode: String,
pub rate_limit: Option<i32>,
pub rate_limit_mode: String,
}
impl UpsertUserGroupRecord {
pub fn normalized_name(&self) -> String {
normalize_user_group_name(&self.name).to_ascii_lowercase()
}
}
impl StoredUserPreferenceRecord {
pub fn default_for_user(user_id: impl Into<String>) -> Self {
Self {
@@ -429,6 +634,7 @@ pub struct UserExportListQuery {
pub role: Option<String>,
pub is_active: Option<bool>,
pub search: Option<String>,
pub group_id: Option<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
@@ -463,6 +669,64 @@ pub trait UserReadRepository: Send + Sync {
user_id: &str,
) -> Result<Option<StoredUserExportRow>, crate::DataLayerError>;
async fn list_user_groups(&self) -> Result<Vec<StoredUserGroup>, crate::DataLayerError>;
async fn find_user_group_by_id(
&self,
group_id: &str,
) -> Result<Option<StoredUserGroup>, crate::DataLayerError>;
async fn list_user_groups_by_ids(
&self,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, crate::DataLayerError>;
async fn create_user_group(
&self,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, crate::DataLayerError>;
async fn update_user_group(
&self,
group_id: &str,
record: UpsertUserGroupRecord,
) -> Result<Option<StoredUserGroup>, crate::DataLayerError>;
async fn delete_user_group(&self, group_id: &str) -> Result<bool, crate::DataLayerError>;
async fn list_user_group_members(
&self,
group_id: &str,
) -> Result<Vec<StoredUserGroupMember>, crate::DataLayerError>;
async fn replace_user_group_members(
&self,
group_id: &str,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMember>, crate::DataLayerError>;
async fn list_user_groups_for_user(
&self,
user_id: &str,
) -> Result<Vec<StoredUserGroup>, crate::DataLayerError>;
async fn list_user_group_memberships_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserGroupMembership>, crate::DataLayerError>;
async fn replace_user_groups_for_user(
&self,
user_id: &str,
group_ids: &[String],
) -> Result<Vec<StoredUserGroup>, crate::DataLayerError>;
async fn add_user_to_group(
&self,
group_id: &str,
user_id: &str,
) -> Result<bool, crate::DataLayerError>;
async fn list_non_admin_export_users(
&self,
) -> Result<Vec<StoredUserExportRow>, crate::DataLayerError>;
@@ -602,6 +866,15 @@ pub trait UserReadRepository: Send + Sync {
is_active: Option<bool>,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
async fn update_local_auth_user_policy_modes(
&self,
user_id: &str,
allowed_providers_mode: Option<String>,
allowed_api_formats_mode: Option<String>,
allowed_models_mode: Option<String>,
rate_limit_mode: Option<String>,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
async fn update_user_model_capability_settings(
&self,
user_id: &str,
@@ -717,6 +990,47 @@ fn normalize_optional_json(value: Option<Value>) -> Option<Value> {
}
}
pub fn normalize_user_group_name(value: &str) -> String {
value.split_whitespace().collect::<Vec<_>>().join(" ")
}
pub fn normalize_list_policy_mode(
value: &str,
field_name: &str,
) -> Result<String, crate::DataLayerError> {
match value.trim().to_ascii_lowercase().as_str() {
"inherit" => Ok("inherit".to_string()),
"unrestricted" => Ok("unrestricted".to_string()),
"specific" => Ok("specific".to_string()),
"deny_all" => Ok("deny_all".to_string()),
_ => Err(crate::DataLayerError::UnexpectedValue(format!(
"{field_name} is not a valid list policy mode"
))),
}
}
pub fn normalize_rate_limit_policy_mode(
value: &str,
field_name: &str,
) -> Result<String, crate::DataLayerError> {
match value.trim().to_ascii_lowercase().as_str() {
"inherit" => Ok("inherit".to_string()),
"system" => Ok("system".to_string()),
"custom" => Ok("custom".to_string()),
_ => Err(crate::DataLayerError::UnexpectedValue(format!(
"{field_name} is not a valid rate limit policy mode"
))),
}
}
fn legacy_list_policy_mode(values: &Option<Vec<String>>) -> String {
if values.is_some() {
"specific".to_string()
} else {
"unrestricted".to_string()
}
}
fn parse_string_list(
value: Option<Value>,
field_name: &str,