diff --git a/backend/migrations/20240902000123_create_tenants/down.sql b/backend/migrations/20240902000123_create_tenants/down.sql new file mode 100644 index 0000000..68e11dc --- /dev/null +++ b/backend/migrations/20240902000123_create_tenants/down.sql @@ -0,0 +1 @@ +DROP TABLE IF EXISTS public.tenants; diff --git a/backend/migrations/20240902000123_create_tenants/up.sql b/backend/migrations/20240902000123_create_tenants/up.sql new file mode 100644 index 0000000..d47fb29 --- /dev/null +++ b/backend/migrations/20240902000123_create_tenants/up.sql @@ -0,0 +1,10 @@ +CREATE TABLE public.tenants ( + tenant_id UUID PRIMARY KEY, + slug TEXT NOT NULL UNIQUE, + storage_root TEXT, + quickwit_index TEXT, + status TEXT NOT NULL DEFAULT 'active', + config JSONB NOT NULL DEFAULT '{}'::jsonb, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW() +); diff --git a/backend/migrations/20240902003000_seed_default_tenant/down.sql b/backend/migrations/20240902003000_seed_default_tenant/down.sql new file mode 100644 index 0000000..8ba7b7d --- /dev/null +++ b/backend/migrations/20240902003000_seed_default_tenant/down.sql @@ -0,0 +1,2 @@ +DELETE FROM public.tenants +WHERE slug = 'admin'; diff --git a/backend/migrations/20240902003000_seed_default_tenant/up.sql b/backend/migrations/20240902003000_seed_default_tenant/up.sql new file mode 100644 index 0000000..438d250 --- /dev/null +++ b/backend/migrations/20240902003000_seed_default_tenant/up.sql @@ -0,0 +1,10 @@ +INSERT INTO public.tenants (tenant_id, slug, storage_root, quickwit_index, status, config) +VALUES ( + gen_random_uuid(), + 'admin', + NULL, + NULL, + 'active', + '{}'::jsonb +) +ON CONFLICT (slug) DO NOTHING; diff --git a/backend/migrations/20240902011000_create_user_memberships/down.sql b/backend/migrations/20240902011000_create_user_memberships/down.sql new file mode 100644 index 0000000..a167636 --- /dev/null +++ b/backend/migrations/20240902011000_create_user_memberships/down.sql @@ -0,0 +1 @@ +DROP TABLE IF EXISTS public.user_memberships; diff --git a/backend/migrations/20240902011000_create_user_memberships/up.sql b/backend/migrations/20240902011000_create_user_memberships/up.sql new file mode 100644 index 0000000..697a31d --- /dev/null +++ b/backend/migrations/20240902011000_create_user_memberships/up.sql @@ -0,0 +1,12 @@ +CREATE TABLE public.user_memberships ( + id UUID PRIMARY KEY, + user_id UUID NOT NULL REFERENCES users(id) ON DELETE CASCADE, + tenant_id UUID NOT NULL REFERENCES public.tenants(tenant_id) ON DELETE CASCADE, + role TEXT NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + UNIQUE (user_id, tenant_id) +); + +CREATE INDEX user_memberships_tenant_id_idx ON public.user_memberships (tenant_id); +CREATE INDEX user_memberships_user_id_idx ON public.user_memberships (user_id); diff --git a/backend/migrations/20240902011500_seed_admin_membership/down.sql b/backend/migrations/20240902011500_seed_admin_membership/down.sql new file mode 100644 index 0000000..e352c7f --- /dev/null +++ b/backend/migrations/20240902011500_seed_admin_membership/down.sql @@ -0,0 +1,7 @@ +DELETE FROM public.user_memberships +WHERE user_id IN ( + SELECT id FROM public.users WHERE username = 'admin' +) + AND tenant_id IN ( + SELECT tenant_id FROM public.tenants WHERE slug = 'admin' +); diff --git a/backend/migrations/20240902011500_seed_admin_membership/up.sql b/backend/migrations/20240902011500_seed_admin_membership/up.sql new file mode 100644 index 0000000..eb6de28 --- /dev/null +++ b/backend/migrations/20240902011500_seed_admin_membership/up.sql @@ -0,0 +1,7 @@ +INSERT INTO public.user_memberships (id, user_id, tenant_id, role) +SELECT gen_random_uuid(), u.id, t.tenant_id, 'admin' +FROM public.users AS u +JOIN public.tenants AS t ON t.slug = 'admin' +LEFT JOIN public.user_memberships AS um ON um.user_id = u.id AND um.tenant_id = t.tenant_id +WHERE u.username = 'admin' AND um.id IS NULL +ON CONFLICT (user_id, tenant_id) DO NOTHING; diff --git a/backend/migrations/20240902013000_drop_user_role/down.sql b/backend/migrations/20240902013000_drop_user_role/down.sql new file mode 100644 index 0000000..6d0888b --- /dev/null +++ b/backend/migrations/20240902013000_drop_user_role/down.sql @@ -0,0 +1,2 @@ +ALTER TABLE users + ADD COLUMN role VARCHAR(16) NOT NULL DEFAULT 'user'; diff --git a/backend/migrations/20240902013000_drop_user_role/up.sql b/backend/migrations/20240902013000_drop_user_role/up.sql new file mode 100644 index 0000000..455814f --- /dev/null +++ b/backend/migrations/20240902013000_drop_user_role/up.sql @@ -0,0 +1,2 @@ +ALTER TABLE users + DROP COLUMN role; diff --git a/backend/migrations/20240902014000_add_tenant_id/down.sql b/backend/migrations/20240902014000_add_tenant_id/down.sql new file mode 100644 index 0000000..d5de82d --- /dev/null +++ b/backend/migrations/20240902014000_add_tenant_id/down.sql @@ -0,0 +1,47 @@ +DROP INDEX IF EXISTS jobs_tenant_id_idx; +ALTER TABLE jobs DROP CONSTRAINT IF EXISTS jobs_tenant_id_fkey; +ALTER TABLE jobs DROP COLUMN IF EXISTS tenant_id; + +DROP INDEX IF EXISTS refresh_tokens_tenant_id_idx; +ALTER TABLE refresh_tokens DROP CONSTRAINT IF EXISTS refresh_tokens_tenant_id_fkey; +ALTER TABLE refresh_tokens DROP COLUMN IF EXISTS tenant_id; + +DROP INDEX IF EXISTS users_tenant_id_idx; +ALTER TABLE users DROP CONSTRAINT IF EXISTS users_tenant_id_fkey; +ALTER TABLE users DROP COLUMN IF EXISTS tenant_id; + +DROP INDEX IF EXISTS folders_tenant_id_idx; +ALTER TABLE folders DROP CONSTRAINT IF EXISTS folders_tenant_id_fkey; +ALTER TABLE folders DROP COLUMN IF EXISTS tenant_id; + +DROP INDEX IF EXISTS tags_tenant_id_idx; +ALTER TABLE tags DROP CONSTRAINT IF EXISTS tags_tenant_id_fkey; +ALTER TABLE tags DROP COLUMN IF EXISTS tenant_id; + +DROP INDEX IF EXISTS correspondents_tenant_id_idx; +ALTER TABLE correspondents DROP CONSTRAINT IF EXISTS correspondents_tenant_id_fkey; +ALTER TABLE correspondents DROP COLUMN IF EXISTS tenant_id; + +DROP INDEX IF EXISTS document_correspondents_tenant_id_idx; +ALTER TABLE document_correspondents DROP CONSTRAINT IF EXISTS document_correspondents_tenant_id_fkey; +ALTER TABLE document_correspondents DROP COLUMN IF EXISTS tenant_id; + +DROP INDEX IF EXISTS document_tags_tenant_id_idx; +ALTER TABLE document_tags DROP CONSTRAINT IF EXISTS document_tags_tenant_id_fkey; +ALTER TABLE document_tags DROP COLUMN IF EXISTS tenant_id; + +DROP INDEX IF EXISTS document_asset_objects_tenant_id_idx; +ALTER TABLE document_asset_objects DROP CONSTRAINT IF EXISTS document_asset_objects_tenant_id_fkey; +ALTER TABLE document_asset_objects DROP COLUMN IF EXISTS tenant_id; + +DROP INDEX IF EXISTS document_assets_tenant_id_idx; +ALTER TABLE document_assets DROP CONSTRAINT IF EXISTS document_assets_tenant_id_fkey; +ALTER TABLE document_assets DROP COLUMN IF EXISTS tenant_id; + +DROP INDEX IF EXISTS document_versions_tenant_id_idx; +ALTER TABLE document_versions DROP CONSTRAINT IF EXISTS document_versions_tenant_id_fkey; +ALTER TABLE document_versions DROP COLUMN IF EXISTS tenant_id; + +DROP INDEX IF EXISTS documents_tenant_id_idx; +ALTER TABLE documents DROP CONSTRAINT IF EXISTS documents_tenant_id_fkey; +ALTER TABLE documents DROP COLUMN IF EXISTS tenant_id; diff --git a/backend/migrations/20240902014000_add_tenant_id/up.sql b/backend/migrations/20240902014000_add_tenant_id/up.sql new file mode 100644 index 0000000..b393c92 --- /dev/null +++ b/backend/migrations/20240902014000_add_tenant_id/up.sql @@ -0,0 +1,119 @@ +-- Add tenant_id to documents +ALTER TABLE documents ADD COLUMN tenant_id UUID; +UPDATE documents +SET tenant_id = (SELECT tenant_id FROM public.tenants WHERE slug = 'admin'); +ALTER TABLE documents ALTER COLUMN tenant_id SET NOT NULL; +ALTER TABLE documents + ADD CONSTRAINT documents_tenant_id_fkey FOREIGN KEY (tenant_id) REFERENCES public.tenants(tenant_id); +CREATE INDEX documents_tenant_id_idx ON documents (tenant_id); + +-- Add tenant_id to document_versions +ALTER TABLE document_versions ADD COLUMN tenant_id UUID; +UPDATE document_versions AS dv +SET tenant_id = d.tenant_id +FROM documents AS d +WHERE dv.document_id = d.id; +ALTER TABLE document_versions ALTER COLUMN tenant_id SET NOT NULL; +ALTER TABLE document_versions + ADD CONSTRAINT document_versions_tenant_id_fkey FOREIGN KEY (tenant_id) REFERENCES public.tenants(tenant_id); +CREATE INDEX document_versions_tenant_id_idx ON document_versions (tenant_id); + +-- Add tenant_id to document_assets +ALTER TABLE document_assets ADD COLUMN tenant_id UUID; +UPDATE document_assets AS da +SET tenant_id = dv.tenant_id +FROM document_versions AS dv +WHERE da.document_version_id = dv.id; +ALTER TABLE document_assets ALTER COLUMN tenant_id SET NOT NULL; +ALTER TABLE document_assets + ADD CONSTRAINT document_assets_tenant_id_fkey FOREIGN KEY (tenant_id) REFERENCES public.tenants(tenant_id); +CREATE INDEX document_assets_tenant_id_idx ON document_assets (tenant_id); + +-- Add tenant_id to document_asset_objects +ALTER TABLE document_asset_objects ADD COLUMN tenant_id UUID; +UPDATE document_asset_objects AS dao +SET tenant_id = da.tenant_id +FROM document_assets AS da +WHERE dao.asset_id = da.id; +ALTER TABLE document_asset_objects ALTER COLUMN tenant_id SET NOT NULL; +ALTER TABLE document_asset_objects + ADD CONSTRAINT document_asset_objects_tenant_id_fkey FOREIGN KEY (tenant_id) REFERENCES public.tenants(tenant_id); +CREATE INDEX document_asset_objects_tenant_id_idx ON document_asset_objects (tenant_id); + +-- Add tenant_id to document_tags +ALTER TABLE document_tags ADD COLUMN tenant_id UUID; +UPDATE document_tags AS dt +SET tenant_id = d.tenant_id +FROM documents AS d +WHERE dt.document_id = d.id; +ALTER TABLE document_tags ALTER COLUMN tenant_id SET NOT NULL; +ALTER TABLE document_tags + ADD CONSTRAINT document_tags_tenant_id_fkey FOREIGN KEY (tenant_id) REFERENCES public.tenants(tenant_id); +CREATE INDEX document_tags_tenant_id_idx ON document_tags (tenant_id); + +-- Add tenant_id to document_correspondents +ALTER TABLE document_correspondents ADD COLUMN tenant_id UUID; +UPDATE document_correspondents AS dc +SET tenant_id = d.tenant_id +FROM documents AS d +WHERE dc.document_id = d.id; +ALTER TABLE document_correspondents ALTER COLUMN tenant_id SET NOT NULL; +ALTER TABLE document_correspondents + ADD CONSTRAINT document_correspondents_tenant_id_fkey FOREIGN KEY (tenant_id) REFERENCES public.tenants(tenant_id); +CREATE INDEX document_correspondents_tenant_id_idx ON document_correspondents (tenant_id); + +-- Add tenant_id to correspondents +ALTER TABLE correspondents ADD COLUMN tenant_id UUID; +UPDATE correspondents +SET tenant_id = (SELECT tenant_id FROM public.tenants WHERE slug = 'admin'); +ALTER TABLE correspondents ALTER COLUMN tenant_id SET NOT NULL; +ALTER TABLE correspondents + ADD CONSTRAINT correspondents_tenant_id_fkey FOREIGN KEY (tenant_id) REFERENCES public.tenants(tenant_id); +CREATE INDEX correspondents_tenant_id_idx ON correspondents (tenant_id); + +-- Add tenant_id to tags +ALTER TABLE tags ADD COLUMN tenant_id UUID; +UPDATE tags +SET tenant_id = (SELECT tenant_id FROM public.tenants WHERE slug = 'admin'); +ALTER TABLE tags ALTER COLUMN tenant_id SET NOT NULL; +ALTER TABLE tags + ADD CONSTRAINT tags_tenant_id_fkey FOREIGN KEY (tenant_id) REFERENCES public.tenants(tenant_id); +CREATE INDEX tags_tenant_id_idx ON tags (tenant_id); + +-- Add tenant_id to folders +ALTER TABLE folders ADD COLUMN tenant_id UUID; +UPDATE folders +SET tenant_id = (SELECT tenant_id FROM public.tenants WHERE slug = 'admin'); +ALTER TABLE folders ALTER COLUMN tenant_id SET NOT NULL; +ALTER TABLE folders + ADD CONSTRAINT folders_tenant_id_fkey FOREIGN KEY (tenant_id) REFERENCES public.tenants(tenant_id); +CREATE INDEX folders_tenant_id_idx ON folders (tenant_id); + +-- Add tenant_id to users +ALTER TABLE users ADD COLUMN tenant_id UUID; +UPDATE users +SET tenant_id = (SELECT tenant_id FROM public.tenants WHERE slug = 'admin'); +ALTER TABLE users ALTER COLUMN tenant_id SET NOT NULL; +ALTER TABLE users + ADD CONSTRAINT users_tenant_id_fkey FOREIGN KEY (tenant_id) REFERENCES public.tenants(tenant_id); +CREATE INDEX users_tenant_id_idx ON users (tenant_id); + +-- Add tenant_id to refresh_tokens +ALTER TABLE refresh_tokens ADD COLUMN tenant_id UUID; +UPDATE refresh_tokens AS rt +SET tenant_id = u.tenant_id +FROM users AS u +WHERE rt.user_id = u.id; +ALTER TABLE refresh_tokens ALTER COLUMN tenant_id SET NOT NULL; +ALTER TABLE refresh_tokens + ADD CONSTRAINT refresh_tokens_tenant_id_fkey FOREIGN KEY (tenant_id) REFERENCES public.tenants(tenant_id); +CREATE INDEX refresh_tokens_tenant_id_idx ON refresh_tokens (tenant_id); + +-- Add tenant_id to jobs +ALTER TABLE jobs ADD COLUMN tenant_id UUID; +UPDATE jobs +SET tenant_id = (SELECT tenant_id FROM public.tenants WHERE slug = 'admin'); +ALTER TABLE jobs ALTER COLUMN tenant_id SET NOT NULL; +ALTER TABLE jobs + ADD CONSTRAINT jobs_tenant_id_fkey FOREIGN KEY (tenant_id) REFERENCES public.tenants(tenant_id); +CREATE INDEX jobs_tenant_id_idx ON jobs (tenant_id); diff --git a/backend/src/auth/jwt.rs b/backend/src/auth/jwt.rs index 464b2a9..ae6b18d 100644 --- a/backend/src/auth/jwt.rs +++ b/backend/src/auth/jwt.rs @@ -15,6 +15,8 @@ pub struct JwtService { expiry: Duration, download_audience: String, download_expiry: Duration, + selector_audience: String, + selector_expiry: Duration, } impl JwtService { @@ -27,16 +29,18 @@ impl JwtService { expiry: Duration::minutes(config.jwt_expiry_minutes), download_audience: config.download_token_audience.clone(), download_expiry: Duration::minutes(config.download_token_expiry_minutes), + selector_audience: format!("{}:tenant-selector", config.jwt_audience), + selector_expiry: Duration::minutes(15), }) } - pub fn generate_token(&self, user_id: Uuid, username: &str, role: &str) -> Result { + pub fn generate_token(&self, user_id: Uuid, tenant_id: Uuid, username: &str) -> Result { let now = Utc::now(); let exp = now + self.expiry; let claims = Claims { sub: user_id, + tenant_id, username: username.to_owned(), - role: role.to_owned(), iss: self.issuer.clone(), aud: self.audience.clone(), iat: now.timestamp() as usize, @@ -54,12 +58,18 @@ impl JwtService { Ok(data.claims) } - pub fn generate_download_token(&self, document_id: Uuid, user_id: Uuid) -> Result { + pub fn generate_download_token( + &self, + document_id: Uuid, + user_id: Uuid, + tenant_id: Uuid, + ) -> Result { let now = Utc::now(); let exp = now + self.download_expiry; let claims = DownloadClaims { doc_id: document_id, user_id, + tenant_id, iss: self.issuer.clone(), aud: self.download_audience.clone(), iat: now.timestamp() as usize, @@ -76,13 +86,35 @@ impl JwtService { let data = decode::(token, &self.decoding, &validation)?; Ok(data.claims) } + + pub fn generate_tenant_selector_token(&self, user_id: Uuid) -> Result { + let now = Utc::now(); + let exp = now + self.selector_expiry; + let claims = TenantSelectionClaims { + sub: user_id, + iss: self.issuer.clone(), + aud: self.selector_audience.clone(), + iat: now.timestamp() as usize, + exp: exp.timestamp() as usize, + }; + + Ok(encode(&Header::default(), &claims, &self.encoding)?) + } + + pub fn verify_tenant_selector_token(&self, token: &str) -> Result { + let mut validation = Validation::default(); + validation.set_audience(&[self.selector_audience.clone()]); + validation.set_issuer(&[self.issuer.clone()]); + let data = decode::(token, &self.decoding, &validation)?; + Ok(data.claims) + } } #[derive(Debug, Clone, Serialize, Deserialize)] pub struct Claims { pub sub: Uuid, + pub tenant_id: Uuid, pub username: String, - pub role: String, pub iss: String, pub aud: String, pub iat: usize, @@ -93,6 +125,16 @@ pub struct Claims { pub struct DownloadClaims { pub doc_id: Uuid, pub user_id: Uuid, + pub tenant_id: Uuid, + pub iss: String, + pub aud: String, + pub iat: usize, + pub exp: usize, +} + +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct TenantSelectionClaims { + pub sub: Uuid, pub iss: String, pub aud: String, pub iat: usize, diff --git a/backend/src/auth/mod.rs b/backend/src/auth/mod.rs index eb1652e..bd08cf4 100644 --- a/backend/src/auth/mod.rs +++ b/backend/src/auth/mod.rs @@ -6,13 +6,17 @@ use axum_extra::headers::{authorization::Bearer, Authorization}; use axum_extra::TypedHeader; use serde::{Deserialize, Serialize}; -use crate::{error::AppError, state::AppState}; +use crate::{ + error::AppError, + state::{AppState, PgPooledConnection}, +}; +use uuid::Uuid; #[derive(Debug, Clone, Serialize, Deserialize)] pub struct AuthenticatedUser { pub user_id: uuid::Uuid, pub username: String, - pub role: String, + pub tenant_id: uuid::Uuid, } #[async_trait] @@ -23,6 +27,10 @@ impl FromRequestParts for AuthenticatedUser { parts: &mut Parts, state: &AppState, ) -> Result { + if let Some(user) = parts.extensions.get::() { + return Ok(user.clone()); + } + let TypedHeader(Authorization(bearer)) = TypedHeader::>::from_request_parts(parts, state) .await @@ -33,10 +41,48 @@ impl FromRequestParts for AuthenticatedUser { .verify_token(bearer.token()) .map_err(|_| AppError::unauthorized())?; - Ok(AuthenticatedUser { + let user = AuthenticatedUser { user_id: claims.sub, username: claims.username, - role: claims.role, + tenant_id: claims.tenant_id, + }; + + parts.extensions.insert(user.clone()); + + Ok(user) + } +} + +pub struct TenantScopedConn { + pub conn: PgPooledConnection, + pub tenant_id: Uuid, + pub user_id: Uuid, + pub user: AuthenticatedUser, +} + +impl TenantScopedConn { + pub fn conn(&mut self) -> &mut PgPooledConnection { + &mut self.conn + } +} + +#[async_trait] +impl FromRequestParts for TenantScopedConn { + type Rejection = AppError; + + async fn from_request_parts( + parts: &mut Parts, + state: &AppState, + ) -> Result { + let user = AuthenticatedUser::from_request_parts(parts, state).await?; + let tenant_id = user.tenant_id; + let conn = state.db_for_tenant(tenant_id)?; + + Ok(Self { + conn, + tenant_id, + user_id: user.user_id, + user, }) } } diff --git a/backend/src/config.rs b/backend/src/config.rs index 2981bc3..2340190 100644 --- a/backend/src/config.rs +++ b/backend/src/config.rs @@ -30,6 +30,7 @@ pub struct AppConfig { pub s3_bucket: String, pub quickwit_endpoint: Option, pub quickwit_index: Option, + pub default_tenant_slug: String, } impl AppConfig { @@ -79,6 +80,8 @@ impl AppConfig { let s3_bucket = env::var("S3_BUCKET").context("S3_BUCKET must be set")?; let quickwit_endpoint = env::var("QUICKWIT_ENDPOINT").ok(); let quickwit_index = env::var("QUICKWIT_INDEX").ok(); + let default_tenant_slug = + env::var("DEFAULT_TENANT_SLUG").unwrap_or_else(|_| "admin".to_string()); Ok(Self { database_url, @@ -104,6 +107,7 @@ impl AppConfig { s3_bucket, quickwit_endpoint, quickwit_index, + default_tenant_slug, }) } diff --git a/backend/src/jobs.rs b/backend/src/jobs.rs index 678cded..d569cec 100644 --- a/backend/src/jobs.rs +++ b/backend/src/jobs.rs @@ -30,6 +30,7 @@ pub type JobQueueResult = Result; pub fn enqueue_job( conn: &mut PgConnection, + tenant_id: Uuid, job_type: &str, payload: Value, run_after: Option, @@ -40,6 +41,7 @@ pub fn enqueue_job( payload, status: STATUS_QUEUED.to_string(), run_after: run_after.unwrap_or_else(|| Utc::now().naive_utc()), + tenant_id, }; diesel::insert_into(jobs::table) diff --git a/backend/src/lib.rs b/backend/src/lib.rs index c7dc874..581469a 100644 --- a/backend/src/lib.rs +++ b/backend/src/lib.rs @@ -9,6 +9,7 @@ pub mod s3; pub mod schema; pub mod state; pub mod storage; +pub mod tenants; pub mod utils; pub mod workers; pub use workers::{default_handlers, Worker}; diff --git a/backend/src/models.rs b/backend/src/models.rs index afeedbb..b05f3e0 100644 --- a/backend/src/models.rs +++ b/backend/src/models.rs @@ -4,15 +4,51 @@ use uuid::Uuid; use crate::schema::*; +#[derive(Debug, Clone, Queryable, Identifiable, Associations)] +#[diesel(table_name = user_memberships)] +#[diesel(belongs_to(User, foreign_key = user_id))] +#[diesel(belongs_to(Tenant, foreign_key = tenant_id))] +pub struct UserMembership { + pub id: Uuid, + pub user_id: Uuid, + pub tenant_id: Uuid, + pub role: String, + pub created_at: NaiveDateTime, + pub updated_at: NaiveDateTime, +} + +#[derive(Debug, Insertable)] +#[diesel(table_name = user_memberships)] +pub struct NewUserMembership { + pub id: Uuid, + pub user_id: Uuid, + pub tenant_id: Uuid, + pub role: String, +} + +#[derive(Debug, Clone, Queryable, Identifiable)] +#[diesel(table_name = tenants)] +#[diesel(primary_key(tenant_id))] +pub struct Tenant { + pub tenant_id: Uuid, + pub slug: String, + pub storage_root: Option, + pub quickwit_index: Option, + pub status: String, + pub config: serde_json::Value, + pub created_at: NaiveDateTime, + pub updated_at: NaiveDateTime, +} + #[derive(Debug, Clone, Queryable, Identifiable)] #[diesel(table_name = users)] pub struct User { pub id: Uuid, pub username: String, pub password_hash: String, - pub role: String, pub created_at: NaiveDateTime, pub updated_at: NaiveDateTime, + pub tenant_id: Uuid, } #[derive(Debug, Insertable)] @@ -21,7 +57,7 @@ pub struct NewUser { pub id: Uuid, pub username: String, pub password_hash: String, - pub role: String, + pub tenant_id: Uuid, } #[derive(Debug, Clone, Queryable, Identifiable)] @@ -32,6 +68,7 @@ pub struct Folder { pub parent_id: Option, pub created_at: NaiveDateTime, pub updated_at: NaiveDateTime, + pub tenant_id: Uuid, } #[derive(Debug, Insertable)] @@ -40,6 +77,7 @@ pub struct NewFolder { pub id: Uuid, pub name: String, pub parent_id: Option, + pub tenant_id: Uuid, } #[derive(Debug, Clone, Queryable, Identifiable, Associations)] @@ -58,6 +96,7 @@ pub struct Document { pub issued_at: Option, pub title: String, pub current_version_id: Uuid, + pub tenant_id: Uuid, } #[derive(Debug, Insertable)] @@ -72,6 +111,7 @@ pub struct NewDocument { pub metadata: serde_json::Value, pub issued_at: Option, pub title: String, + pub tenant_id: Uuid, } #[derive(Debug, Clone, Queryable, Identifiable, Associations)] @@ -87,6 +127,7 @@ pub struct DocumentVersion { pub created_at: NaiveDateTime, pub operations_summary: serde_json::Value, pub metadata: serde_json::Value, + pub tenant_id: Uuid, } #[derive(Debug, Insertable)] @@ -100,6 +141,7 @@ pub struct NewDocumentVersion { pub checksum: String, pub operations_summary: serde_json::Value, pub metadata: serde_json::Value, + pub tenant_id: Uuid, } #[derive(Debug, Clone, Queryable, Identifiable, Associations)] @@ -113,6 +155,7 @@ pub struct DocumentAsset { pub metadata: serde_json::Value, pub created_at: NaiveDateTime, pub cardinality: Option, + pub tenant_id: Uuid, } #[derive(Debug, Insertable)] @@ -124,6 +167,7 @@ pub struct NewDocumentAsset { pub mime_type: String, pub metadata: serde_json::Value, pub cardinality: Option, + pub tenant_id: Uuid, } #[derive(Debug, Clone, Queryable, Identifiable, Associations)] @@ -135,6 +179,7 @@ pub struct DocumentAssetObject { pub ordinal: i32, pub s3_key: String, pub metadata: serde_json::Value, + pub tenant_id: Uuid, } #[derive(Debug, Insertable)] @@ -145,6 +190,7 @@ pub struct NewDocumentAssetObject { pub ordinal: i32, pub s3_key: String, pub metadata: serde_json::Value, + pub tenant_id: Uuid, } #[derive(Debug, Clone, Queryable, Identifiable)] @@ -159,6 +205,7 @@ pub struct Job { pub last_error: Option, pub created_at: NaiveDateTime, pub updated_at: NaiveDateTime, + pub tenant_id: Uuid, } #[derive(Debug, Insertable)] @@ -169,6 +216,7 @@ pub struct NewJob { pub payload: serde_json::Value, pub status: String, pub run_after: NaiveDateTime, + pub tenant_id: Uuid, } #[derive(Debug, Clone, Queryable, Identifiable)] @@ -178,6 +226,7 @@ pub struct Tag { pub label: String, pub color: Option, pub created_at: NaiveDateTime, + pub tenant_id: Uuid, } #[derive(Debug, Insertable)] @@ -186,6 +235,7 @@ pub struct NewTag { pub id: Uuid, pub label: String, pub color: Option, + pub tenant_id: Uuid, } #[allow(dead_code)] @@ -199,6 +249,7 @@ pub struct DocumentTag { pub tag_id: Uuid, pub assigned_at: NaiveDateTime, pub assigned_by: Option, + pub tenant_id: Uuid, } #[derive(Debug, Insertable)] @@ -207,6 +258,7 @@ pub struct NewDocumentTag { pub document_id: Uuid, pub tag_id: Uuid, pub assigned_by: Option, + pub tenant_id: Uuid, } #[derive(Debug, Clone, Queryable, Identifiable)] @@ -217,6 +269,7 @@ pub struct Correspondent { pub metadata: serde_json::Value, pub created_at: NaiveDateTime, pub updated_at: NaiveDateTime, + pub tenant_id: Uuid, } #[derive(Debug, Insertable)] @@ -225,6 +278,7 @@ pub struct NewCorrespondent { pub id: Uuid, pub name: String, pub metadata: serde_json::Value, + pub tenant_id: Uuid, } #[derive(Debug, Clone, Queryable, Associations)] @@ -238,6 +292,7 @@ pub struct DocumentCorrespondent { pub role: String, pub assigned_at: NaiveDateTime, pub assigned_by: Option, + pub tenant_id: Uuid, } #[derive(Debug, Insertable)] @@ -247,6 +302,7 @@ pub struct NewDocumentCorrespondent { pub correspondent_id: Uuid, pub role: String, pub assigned_by: Option, + pub tenant_id: Uuid, } #[derive(Debug, Clone, Queryable, Identifiable, Associations)] @@ -261,6 +317,7 @@ pub struct RefreshToken { pub revoked_at: Option, pub created_at: NaiveDateTime, pub updated_at: NaiveDateTime, + pub tenant_id: Uuid, } #[derive(Debug, Insertable)] @@ -271,4 +328,5 @@ pub struct NewRefreshToken { pub token_hash: String, pub issued_at: NaiveDateTime, pub expires_at: NaiveDateTime, + pub tenant_id: Uuid, } diff --git a/backend/src/routes/auth.rs b/backend/src/routes/auth.rs index 04b19cc..83ce3a4 100644 --- a/backend/src/routes/auth.rs +++ b/backend/src/routes/auth.rs @@ -1,11 +1,15 @@ use axum::{ extract::State, http::{header::SET_COOKIE, HeaderMap, HeaderValue, StatusCode}, + response::{IntoResponse, Response}, Json, }; -use axum_extra::{headers::Cookie, typed_header::TypedHeader}; +use axum_extra::{ + headers::{authorization::Bearer, Authorization, Cookie}, + typed_header::TypedHeader, +}; use chrono::{Duration as ChronoDuration, Utc}; -use diesel::prelude::*; +use diesel::{pg::PgConnection, prelude::*}; use rand::{rngs::OsRng, RngCore}; use serde::{Deserialize, Serialize}; use sha2::{Digest, Sha256}; @@ -14,8 +18,11 @@ use uuid::Uuid; use crate::{ auth::{password, AuthenticatedUser}, error::{AppError, AppResult}, - models::{NewRefreshToken, RefreshToken, User}, - schema::{refresh_tokens, users::dsl}, + models::{NewRefreshToken, RefreshToken, Tenant, User, UserMembership}, + schema::{ + refresh_tokens, tenants::dsl as tenant_dsl, user_memberships::dsl as memberships_dsl, + users::dsl, + }, state::AppState, }; @@ -27,6 +34,8 @@ const REFRESH_COOKIE_NAME: &str = "refresh_token"; pub struct LoginRequest { pub username: String, pub password: String, + #[serde(default)] + pub preferred_tenant_slug: Option, } #[derive(Serialize)] @@ -36,11 +45,28 @@ pub struct LoginResponse { pub expires_in: i64, } +#[derive(Serialize)] +pub struct TenantSummary { + pub tenant_id: Uuid, + pub slug: String, +} + +#[derive(Serialize)] +pub struct TenantSelectionResponse { + pub selection_token: String, + pub tenants: Vec, +} + +#[derive(Deserialize)] +pub struct TenantSelectionRequest { + pub tenant_id: Uuid, +} + pub async fn login( State(state): State, Json(payload): Json, -) -> AppResult<(HeaderMap, Json)> { - let mut conn = state.db()?; +) -> AppResult { + let mut conn = state.db_unscoped()?; let user: User = dsl::users .filter(dsl::username.eq(&payload.username)) @@ -53,55 +79,67 @@ pub async fn login( return Err(AppError::unauthorized()); } - let access_token = state + let memberships: Vec<(UserMembership, Tenant)> = memberships_dsl::user_memberships + .inner_join(tenant_dsl::tenants) + .filter(memberships_dsl::user_id.eq(user.id)) + .load(&mut conn)?; + + if memberships.is_empty() { + return Err(AppError::unauthorized()); + } + + let preferred_slug = payload + .preferred_tenant_slug + .as_ref() + .map(|slug| slug.trim().to_string()) + .filter(|slug| !slug.is_empty()); + + if let Some(tenant) = preferred_slug.as_ref().and_then(|slug| { + memberships + .iter() + .find(|(_, tenant)| tenant.slug.eq_ignore_ascii_case(slug)) + }) { + return issue_session(&state, &mut conn, &user, tenant.1.tenant_id); + } + + if memberships.len() == 1 { + let tenant_id = memberships[0].1.tenant_id; + return issue_session(&state, &mut conn, &user, tenant_id); + } + + let selection_token = state .jwt - .generate_token(user.id, &user.username, &user.role) + .generate_tenant_selector_token(user.id) .map_err(AppError::from)?; - let now = Utc::now(); - let refresh_value = generate_refresh_token(); - let refresh_hash = hash_refresh_token(&refresh_value); - let refresh_expires_at = now + ChronoDuration::days(state.config.refresh_token_expiry_days); + let tenants = memberships + .into_iter() + .map(|(_, tenant)| TenantSummary { + tenant_id: tenant.tenant_id, + slug: tenant.slug, + }) + .collect(); - let new_refresh = NewRefreshToken { - id: Uuid::new_v4(), - user_id: user.id, - token_hash: refresh_hash, - issued_at: now.naive_utc(), - expires_at: refresh_expires_at.naive_utc(), - }; + let response = Json(TenantSelectionResponse { + selection_token, + tenants, + }) + .into_response(); - diesel::insert_into(refresh_tokens::table) - .values(&new_refresh) - .execute(&mut conn)?; - - let mut headers = HeaderMap::new(); - headers.insert( - SET_COOKIE, - build_refresh_cookie(&state, &refresh_value, refresh_expires_at), - ); - - Ok(( - headers, - Json(LoginResponse { - access_token, - token_type: "Bearer".to_string(), - expires_in: state.config.jwt_expiry_minutes * 60, - }), - )) + Ok(response) } pub async fn refresh( State(state): State, jar: Option>, -) -> AppResult<(HeaderMap, Json)> { +) -> AppResult { let cookies = jar.ok_or_else(AppError::unauthorized)?; let refresh_value = cookies .get(REFRESH_COOKIE_NAME) .ok_or_else(AppError::unauthorized)?; let hashed = hash_refresh_token(refresh_value); - let mut conn = state.db()?; + let mut conn = state.db_unscoped()?; let now = Utc::now(); let now_naive = now.naive_utc(); @@ -128,41 +166,39 @@ pub async fn refresh( .first(&mut conn) .map_err(AppError::from)?; - let access_token = state + issue_session(&state, &mut conn, &user, token.tenant_id) +} + +pub async fn select_tenant( + State(state): State, + TypedHeader(Authorization(bearer)): TypedHeader>, + Json(payload): Json, +) -> AppResult { + let claims = state .jwt - .generate_token(user.id, &user.username, &user.role) + .verify_tenant_selector_token(bearer.token()) + .map_err(|_| AppError::unauthorized())?; + + let mut conn = state.db_unscoped()?; + + let membership_exists = memberships_dsl::user_memberships + .filter(memberships_dsl::user_id.eq(claims.sub)) + .filter(memberships_dsl::tenant_id.eq(payload.tenant_id)) + .inner_join(tenant_dsl::tenants) + .select(memberships_dsl::id) + .first::(&mut conn) + .optional()?; + + if membership_exists.is_none() { + return Err(AppError::unauthorized()); + } + + let user: User = dsl::users + .find(claims.sub) + .first(&mut conn) .map_err(AppError::from)?; - let new_refresh_value = generate_refresh_token(); - let new_refresh_hash = hash_refresh_token(&new_refresh_value); - let new_refresh_expires = now + ChronoDuration::days(state.config.refresh_token_expiry_days); - - let new_refresh = NewRefreshToken { - id: Uuid::new_v4(), - user_id: user.id, - token_hash: new_refresh_hash, - issued_at: now_naive, - expires_at: new_refresh_expires.naive_utc(), - }; - - diesel::insert_into(refresh_tokens::table) - .values(&new_refresh) - .execute(&mut conn)?; - - let mut headers = HeaderMap::new(); - headers.insert( - SET_COOKIE, - build_refresh_cookie(&state, &new_refresh_value, new_refresh_expires), - ); - - Ok(( - headers, - Json(LoginResponse { - access_token, - token_type: "Bearer".to_string(), - expires_in: state.config.jwt_expiry_minutes * 60, - }), - )) + issue_session(&state, &mut conn, &user, payload.tenant_id) } pub async fn logout( @@ -170,7 +206,7 @@ pub async fn logout( user: AuthenticatedUser, jar: Option>, ) -> AppResult<(HeaderMap, StatusCode)> { - let mut conn = state.db()?; + let mut conn = state.db_unscoped()?; let now = Utc::now().naive_utc(); let mut rows_affected = 0; @@ -214,6 +250,50 @@ pub async fn me(user: AuthenticatedUser) -> Json { Json(user) } +fn issue_session( + state: &AppState, + conn: &mut PgConnection, + user: &User, + tenant_id: Uuid, +) -> AppResult { + let now = Utc::now(); + let access_token = state + .jwt + .generate_token(user.id, tenant_id, &user.username) + .map_err(AppError::from)?; + + let refresh_value = generate_refresh_token(); + let refresh_hash = hash_refresh_token(&refresh_value); + let refresh_expires_at = now + ChronoDuration::days(state.config.refresh_token_expiry_days); + + let new_refresh = NewRefreshToken { + id: Uuid::new_v4(), + user_id: user.id, + token_hash: refresh_hash, + issued_at: now.naive_utc(), + expires_at: refresh_expires_at.naive_utc(), + tenant_id, + }; + + diesel::insert_into(refresh_tokens::table) + .values(&new_refresh) + .execute(conn)?; + + let mut response = Json(LoginResponse { + access_token, + token_type: "Bearer".to_string(), + expires_in: state.config.jwt_expiry_minutes * 60, + }) + .into_response(); + + response.headers_mut().insert( + SET_COOKIE, + build_refresh_cookie(state, &refresh_value, refresh_expires_at), + ); + + Ok(response) +} + fn hash_refresh_token(token: &str) -> String { let mut hasher = Sha256::new(); hasher.update(token.as_bytes()); diff --git a/backend/src/routes/correspondents.rs b/backend/src/routes/correspondents.rs index 9bf3ba7..e0e2357 100644 --- a/backend/src/routes/correspondents.rs +++ b/backend/src/routes/correspondents.rs @@ -1,11 +1,6 @@ use std::collections::{BTreeMap, HashMap}; -use axum::{ - extract::{Path, State}, - http::StatusCode, - response::IntoResponse, - Json, -}; +use axum::{extract::Path, http::StatusCode, response::IntoResponse, Json}; use chrono::Utc; use diesel::{dsl::count_star, prelude::*, result::DatabaseErrorKind, PgConnection}; use serde::{Deserialize, Serialize}; @@ -13,10 +8,10 @@ use serde_json::Value; use uuid::Uuid; use crate::{ + auth::TenantScopedConn, error::{AppError, AppResult}, models::{Correspondent, NewCorrespondent}, schema::{correspondents, document_correspondents}, - state::AppState, }; use super::documents::to_iso; @@ -59,15 +54,19 @@ struct CorrespondentChangeset<'a> { } pub async fn list_correspondents( - State(state): State, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, ) -> AppResult>> { - let mut conn = state.db()?; - let correspondents_list: Vec = correspondents::table + .filter(correspondents::tenant_id.eq(tenant_id)) .order(correspondents::name.asc()) .load(&mut conn)?; let usage_rows: Vec<(Uuid, String, i64)> = document_correspondents::table + .filter(document_correspondents::tenant_id.eq(tenant_id)) .group_by(( document_correspondents::correspondent_id, document_correspondents::role, @@ -97,7 +96,11 @@ pub async fn list_correspondents( } pub async fn create_correspondent( - State(state): State, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, Json(payload): Json, ) -> AppResult> { let name = payload.name.trim(); @@ -111,9 +114,9 @@ pub async fn create_correspondent( id: new_id, name: name.to_string(), metadata: metadata_value, + tenant_id, }; - let mut conn = state.db()?; match diesel::insert_into(correspondents::table) .values(&new_correspondent) .execute(&mut conn) @@ -130,13 +133,17 @@ pub async fn create_correspondent( } pub async fn update_correspondent( - State(state): State, Path(correspondent_id): Path, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, Json(payload): Json, ) -> AppResult> { - let mut conn = state.db()?; let existing: Correspondent = correspondents::table .find(correspondent_id) + .filter(correspondents::tenant_id.eq(tenant_id)) .first(&mut conn)?; let mut new_name: Option = None; @@ -149,6 +156,7 @@ pub async fn update_correspondent( let duplicate = correspondents::table .filter(correspondents::name.eq(trimmed)) .filter(correspondents::id.ne(correspondent_id)) + .filter(correspondents::tenant_id.eq(tenant_id)) .first::(&mut conn) .optional()?; if duplicate.is_some() { @@ -167,7 +175,7 @@ pub async fn update_correspondent( } if new_name.is_none() && new_metadata.is_none() { - let usage = load_usage_for_correspondent(&mut conn, correspondent_id)?; + let usage = load_usage_for_correspondent(&mut conn, tenant_id, correspondent_id)?; return Ok(Json(build_summary(existing.clone(), usage))); } @@ -180,24 +188,32 @@ pub async fn update_correspondent( } let now = Utc::now().naive_utc(); - diesel::update(correspondents::table.find(correspondent_id)) - .set((&changeset, correspondents::updated_at.eq(now))) - .execute(&mut conn)?; + diesel::update( + correspondents::table + .find(correspondent_id) + .filter(correspondents::tenant_id.eq(tenant_id)), + ) + .set((&changeset, correspondents::updated_at.eq(now))) + .execute(&mut conn)?; let updated: Correspondent = correspondents::table .find(correspondent_id) + .filter(correspondents::tenant_id.eq(tenant_id)) .first(&mut conn)?; - let usage = load_usage_for_correspondent(&mut conn, correspondent_id)?; + let usage = load_usage_for_correspondent(&mut conn, tenant_id, correspondent_id)?; Ok(Json(build_summary(updated, usage))) } pub async fn delete_correspondent( - State(state): State, Path(correspondent_id): Path, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, ) -> AppResult { - let mut conn = state.db()?; - let usage: i64 = document_correspondents::table + .filter(document_correspondents::tenant_id.eq(tenant_id)) .filter(document_correspondents::correspondent_id.eq(correspondent_id)) .select(count_star()) .first(&mut conn)?; @@ -208,8 +224,12 @@ pub async fn delete_correspondent( )); } - let deleted = - diesel::delete(correspondents::table.find(correspondent_id)).execute(&mut conn)?; + let deleted = diesel::delete( + correspondents::table + .filter(correspondents::id.eq(correspondent_id)) + .filter(correspondents::tenant_id.eq(tenant_id)), + ) + .execute(&mut conn)?; if deleted == 0 { return Err(AppError::not_found()); } @@ -243,10 +263,12 @@ fn normalize_metadata(input: Option) -> Value { fn load_usage_for_correspondent( conn: &mut PgConnection, + tenant_id: Uuid, correspondent_id: Uuid, ) -> AppResult> { let rows: Vec<(String, i64)> = document_correspondents::table .filter(document_correspondents::correspondent_id.eq(correspondent_id)) + .filter(document_correspondents::tenant_id.eq(tenant_id)) .group_by(document_correspondents::role) .select((document_correspondents::role, count_star())) .load(conn)?; diff --git a/backend/src/routes/documents.rs b/backend/src/routes/documents.rs index 46c353c..2607e25 100644 --- a/backend/src/routes/documents.rs +++ b/backend/src/routes/documents.rs @@ -18,7 +18,7 @@ use tracing::{debug, error, info, warn}; use uuid::Uuid; use super::folders::gather_descendant_folder_ids; -use crate::auth::AuthenticatedUser; +use crate::auth::TenantScopedConn; use crate::error::{AppError, AppResult}; use crate::jobs::{enqueue_job, JOB_ANALYZE_DOCUMENT}; use crate::models::{ @@ -372,10 +372,13 @@ pub struct AssetObjectsQuery { pub async fn list_documents( State(state): State, Query(params): Query, - user: AuthenticatedUser, + TenantScopedConn { + mut conn, + tenant_id, + user_id, + .. + }: TenantScopedConn, ) -> AppResult>> { - let mut conn = state.db()?; - let DocumentListQuery { folder_id, include_deleted, @@ -385,7 +388,9 @@ pub async fn list_documents( correspondents, } = params; - let mut docs_query = documents::table.into_boxed(); + let mut docs_query = documents::table + .filter(documents::tenant_id.eq(tenant_id)) + .into_boxed(); if !include_deleted { docs_query = docs_query.filter(documents::deleted_at.is_null()); @@ -414,7 +419,7 @@ pub async fn list_documents( match (folder_id, include_descendants) { (Some(folder_id), true) => { - let descendant_ids = gather_descendant_folder_ids(&mut conn, folder_id)?; + let descendant_ids = gather_descendant_folder_ids(&mut conn, tenant_id, folder_id)?; docs_query = docs_query.filter(documents::folder_id.eq_any(descendant_ids)); } (Some(folder_id), false) => { @@ -598,7 +603,7 @@ pub async fn list_documents( let mut correspondents_map = load_correspondents_for_documents(&mut conn, &doc_ids)?; drop(conn); - let primary_versions = load_primary_assets(&state, &docs).await?; + let primary_versions = load_primary_assets(&state, tenant_id, &docs).await?; let mut response = Vec::with_capacity(doc_ids.len()); for doc in docs { let tags = tags_map.get(&doc.id).cloned(); @@ -606,7 +611,7 @@ pub async fn list_documents( let current_version = primary_versions.get(&doc.id).cloned(); response.push(to_document_response( &state, - user.user_id, + user_id, doc, tags, correspondents, @@ -620,11 +625,17 @@ pub async fn list_documents( pub async fn get_document( State(state): State, Path(document_id): Path, - user: AuthenticatedUser, + TenantScopedConn { + mut conn, + tenant_id, + user_id, + .. + }: TenantScopedConn, ) -> AppResult> { - let mut conn = state.db()?; - - let doc: Document = documents::table.find(document_id).first(&mut conn)?; + let doc: Document = documents::table + .find(document_id) + .filter(documents::tenant_id.eq(tenant_id)) + .first(&mut conn)?; if doc.deleted_at.is_some() { return Err(AppError::not_found()); } @@ -638,13 +649,13 @@ pub async fn get_document( let version_id = current_version.id; drop(conn); - let assets = load_asset_responses(&state, version_id).await?; + let assets = load_asset_responses(&state, tenant_id, version_id).await?; let version_response = to_version_response(current_version, true); Ok(Json(DocumentDetailResponse { document: to_document_response( &state, - user.user_id, + user_id, doc, tags_map.get(&document_id).cloned(), correspondents_map.remove(&document_id).unwrap_or_default(), @@ -655,7 +666,9 @@ pub async fn get_document( pub async fn upload_document( State(state): State, - user: AuthenticatedUser, + TenantScopedConn { + tenant_id, user_id, .. + }: TenantScopedConn, mut multipart: Multipart, ) -> AppResult<(StatusCode, Json)> { let mut file_bytes: Option> = None; @@ -733,7 +746,7 @@ pub async fn upload_document( metadata, }; - let outcome = match process_upload(&state, request, user.user_id).await { + let outcome = match process_upload(&state, request, tenant_id, user_id).await { Ok(outcome) => { info!( document_id = %outcome.detail.document.id, @@ -759,18 +772,25 @@ pub async fn upload_document( } pub async fn request_document_assets( - State(state): State, Path(document_id): Path, Query(query): Query, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, ) -> AppResult { - let mut conn = state.db()?; - let document: Document = documents::table.find(document_id).first(&mut conn)?; + let document: Document = documents::table + .find(document_id) + .filter(documents::tenant_id.eq(tenant_id)) + .first(&mut conn)?; if document.deleted_at.is_some() { return Err(AppError::not_found()); } enqueue_job( &mut conn, + tenant_id, JOB_ANALYZE_DOCUMENT, json!({ "document_id": document_id, @@ -785,13 +805,15 @@ pub async fn request_document_assets( } pub async fn reanalyze_all_documents( - State(state): State, - _user: AuthenticatedUser, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, ) -> AppResult<(StatusCode, Json)> { - let mut conn = state.db()?; - let targets: Vec<(Uuid, Uuid)> = documents::table .filter(documents::deleted_at.is_null()) + .filter(documents::tenant_id.eq(tenant_id)) .select((documents::id, documents::current_version_id)) .load(&mut conn)?; @@ -799,6 +821,7 @@ pub async fn reanalyze_all_documents( for (document_id, version_id) in targets { enqueue_job( &mut conn, + tenant_id, JOB_ANALYZE_DOCUMENT, json!({ "document_id": document_id, @@ -815,7 +838,11 @@ pub async fn reanalyze_all_documents( } pub async fn reanalyze_selected_documents( - State(state): State, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, Json(payload): Json, ) -> AppResult<(StatusCode, Json)> { let BulkReanalyzeSelectionRequest { @@ -830,11 +857,10 @@ pub async fn reanalyze_selected_documents( document_ids.sort(); document_ids.dedup(); - let mut conn = state.db()?; - let targets: Vec<(Uuid, Uuid)> = documents::table .filter(documents::id.eq_any(&document_ids)) .filter(documents::deleted_at.is_null()) + .filter(documents::tenant_id.eq(tenant_id)) .select((documents::id, documents::current_version_id)) .load(&mut conn)?; @@ -848,6 +874,7 @@ pub async fn reanalyze_selected_documents( for (document_id, version_id) in targets { enqueue_job( &mut conn, + tenant_id, JOB_ANALYZE_DOCUMENT, json!({ "document_id": document_id, @@ -866,9 +893,16 @@ pub async fn reanalyze_selected_documents( pub async fn list_document_assets( State(state): State, Path(document_id): Path, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, ) -> AppResult>> { - let mut conn = state.db()?; - let document: Document = documents::table.find(document_id).first(&mut conn)?; + let document: Document = documents::table + .find(document_id) + .filter(documents::tenant_id.eq(tenant_id)) + .first(&mut conn)?; if document.deleted_at.is_some() { return Err(AppError::not_found()); } @@ -876,7 +910,7 @@ pub async fn list_document_assets( let version_id = document.current_version_id; drop(conn); - let assets = load_asset_responses(&state, version_id).await?; + let assets = load_asset_responses(&state, tenant_id, version_id).await?; Ok(Json(assets)) } @@ -884,11 +918,15 @@ pub async fn get_document_asset( State(state): State, Path(asset_id): Path, Query(query): Query, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, ) -> AppResult> { - let mut conn = state.db()?; - let asset: DocumentAsset = match document_assets::table .find(asset_id) + .filter(document_assets::tenant_id.eq(tenant_id)) .first(&mut conn) .optional()? { @@ -911,6 +949,7 @@ pub async fn get_document_asset( let objects: Vec = document_asset_objects::table .filter(document_asset_objects::asset_id.eq(asset_id)) + .filter(document_asset_objects::tenant_id.eq(tenant_id)) .filter(document_asset_objects::ordinal.ge(start)) .filter(document_asset_objects::ordinal.le(end)) .order(document_asset_objects::ordinal.asc()) @@ -951,9 +990,16 @@ pub async fn get_document_asset( pub async fn download_document( State(state): State, Path(document_id): Path, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, ) -> AppResult> { - let mut conn = state.db()?; - let doc: Document = documents::table.find(document_id).first(&mut conn)?; + let doc: Document = documents::table + .find(document_id) + .filter(documents::tenant_id.eq(tenant_id)) + .first(&mut conn)?; if doc.deleted_at.is_some() { return Err(AppError::not_found()); } @@ -989,9 +1035,12 @@ pub async fn download_with_token( .verify_download_token(&token) .map_err(|_| AppError::unauthorized())?; - let mut conn = state.db()?; + let mut conn = state.db_for_tenant(claims.tenant_id)?; - let doc: Document = documents::table.find(claims.doc_id).first(&mut conn)?; + let doc: Document = documents::table + .find(claims.doc_id) + .filter(documents::tenant_id.eq(claims.tenant_id)) + .first(&mut conn)?; if doc.deleted_at.is_some() { return Err(AppError::not_found()); } @@ -1004,6 +1053,7 @@ pub async fn download_with_token( let has_active_refresh: bool = select(exists( refresh_dsl::refresh_tokens .filter(refresh_dsl::user_id.eq(claims.user_id)) + .filter(refresh_dsl::tenant_id.eq(claims.tenant_id)) .filter(refresh_dsl::revoked_at.is_null()) .filter(refresh_dsl::expires_at.gt(now)), )) @@ -1028,29 +1078,42 @@ pub async fn download_with_token( } pub async fn delete_document( - State(state): State, Path(document_id): Path, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, ) -> AppResult { - let mut conn = state.db()?; let now = Utc::now().naive_utc(); - diesel::update(documents::table.find(document_id)) - .set(( - documents::deleted_at.eq(Some(now)), - documents::updated_at.eq(now), - )) - .execute(&mut conn)?; + diesel::update( + documents::table + .find(document_id) + .filter(documents::tenant_id.eq(tenant_id)), + ) + .set(( + documents::deleted_at.eq(Some(now)), + documents::updated_at.eq(now), + )) + .execute(&mut conn)?; Ok(StatusCode::NO_CONTENT) } pub async fn update_document( State(state): State, Path(document_id): Path, - user: AuthenticatedUser, + TenantScopedConn { + mut conn, + tenant_id, + user_id, + .. + }: TenantScopedConn, Json(payload): Json, ) -> AppResult> { - let mut conn = state.db()?; - - let mut document: Document = documents::table.find(document_id).first(&mut conn)?; + let mut document: Document = documents::table + .find(document_id) + .filter(documents::tenant_id.eq(tenant_id)) + .first(&mut conn)?; if document.deleted_at.is_some() { return Err(AppError::not_found()); } @@ -1074,7 +1137,11 @@ pub async fn update_document( let now = Utc::now().naive_utc(); let new_filename = filename_with_retained_extension(&title, &document.filename); - let update_result = diesel::update(documents::table.find(document_id)).set(( + let target = documents::table + .find(document_id) + .filter(documents::tenant_id.eq(tenant_id)); + + let update_result = diesel::update(target).set(( documents::title.eq(&title), documents::filename.eq(&new_filename), documents::updated_at.eq(now), @@ -1090,7 +1157,10 @@ pub async fn update_document( Err(err) => return Err(AppError::from(err)), } - document = documents::table.find(document_id).first(&mut conn)?; + document = documents::table + .find(document_id) + .filter(documents::tenant_id.eq(tenant_id)) + .first(&mut conn)?; } let current_version: DocumentVersion = document_versions::table @@ -1102,13 +1172,13 @@ pub async fn update_document( let version_id = current_version.id; drop(conn); - let assets = load_asset_responses(&state, version_id).await?; + let assets = load_asset_responses(&state, tenant_id, version_id).await?; let version_response = to_version_response(current_version, true); Ok(Json(DocumentDetailResponse { document: to_document_response( &state, - user.user_id, + user_id, document, tags_map.get(&document_id).cloned(), correspondents_map.remove(&document_id).unwrap_or_default(), @@ -1120,26 +1190,39 @@ pub async fn update_document( pub async fn move_document( State(state): State, Path(document_id): Path, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, Json(payload): Json, ) -> AppResult { if let Some(folder_id) = payload.folder_id { - ensure_folder_exists(&state, folder_id)?; + ensure_folder_exists(&state, tenant_id, folder_id)?; } - let mut conn = state.db()?; let now = Utc::now().naive_utc(); - diesel::update(documents::table.find(document_id)) - .set(( - documents::folder_id.eq(payload.folder_id), - documents::updated_at.eq(now), - )) - .execute(&mut conn)?; + diesel::update( + documents::table + .find(document_id) + .filter(documents::tenant_id.eq(tenant_id)), + ) + .set(( + documents::folder_id.eq(payload.folder_id), + documents::updated_at.eq(now), + )) + .execute(&mut conn)?; Ok(StatusCode::NO_CONTENT) } pub async fn bulk_move_documents( State(state): State, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, Json(payload): Json, ) -> AppResult<(StatusCode, Json)> { let BulkMoveRequest { @@ -1155,13 +1238,12 @@ pub async fn bulk_move_documents( document_ids.dedup(); if let Some(target_folder) = folder_id { - ensure_folder_exists(&state, target_folder)?; + ensure_folder_exists(&state, tenant_id, target_folder)?; } - let mut conn = state.db()?; - let existing: Vec<(Uuid, Option)> = documents::table .filter(documents::id.eq_any(&document_ids)) + .filter(documents::tenant_id.eq(tenant_id)) .select((documents::id, documents::deleted_at)) .load(&mut conn)?; @@ -1176,20 +1258,28 @@ pub async fn bulk_move_documents( } let now = Utc::now().naive_utc(); - let updated = diesel::update(documents::table.filter(documents::id.eq_any(&document_ids))) - .set(( - documents::folder_id.eq(folder_id), - documents::updated_at.eq(now), - )) - .execute(&mut conn)?; + let updated = diesel::update( + documents::table + .filter(documents::id.eq_any(&document_ids)) + .filter(documents::tenant_id.eq(tenant_id)), + ) + .set(( + documents::folder_id.eq(folder_id), + documents::updated_at.eq(now), + )) + .execute(&mut conn)?; Ok((StatusCode::OK, Json(BulkMoveResponse { updated }))) } pub async fn assign_correspondents( - State(state): State, Path(document_id): Path, - user: AuthenticatedUser, + TenantScopedConn { + mut conn, + tenant_id, + user_id, + .. + }: TenantScopedConn, Json(payload): Json, ) -> AppResult { if payload.assignments.is_empty() { @@ -1199,11 +1289,13 @@ pub async fn assign_correspondents( let (normalized_pairs, correspondents_vec, roles_vec) = normalize_correspondent_assignments(&payload.assignments)?; let replace = payload.replace; - let user_id = user.user_id; + let user_id_val = user_id; - let mut conn = state.db()?; conn.transaction::<(), AppError, _>(|conn| { - let document: Document = documents::table.find(document_id).first(conn)?; + let document: Document = documents::table + .find(document_id) + .filter(documents::tenant_id.eq(tenant_id)) + .first(conn)?; if document.deleted_at.is_some() { return Err(AppError::not_found()); } @@ -1211,6 +1303,7 @@ pub async fn assign_correspondents( if !correspondents_vec.is_empty() { let existing: Vec = correspondents::table .filter(correspondents::id.eq_any(&correspondents_vec)) + .filter(correspondents::tenant_id.eq(tenant_id)) .load(conn)?; if existing.len() != correspondents_vec.len() { return Err(AppError::bad_request( @@ -1224,6 +1317,7 @@ pub async fn assign_correspondents( let deleted = diesel::delete( document_correspondents::table .filter(document_correspondents::document_id.eq(document_id)) + .filter(document_correspondents::tenant_id.eq(tenant_id)) .filter(document_correspondents::role.eq_any(&roles_vec)), ) .execute(conn)?; @@ -1238,7 +1332,8 @@ pub async fn assign_correspondents( document_id, correspondent_id: *correspondent_id, role: role.clone(), - assigned_by: Some(user_id), + assigned_by: Some(user_id_val), + tenant_id, }) .collect(); @@ -1253,9 +1348,13 @@ pub async fn assign_correspondents( } if changed { - diesel::update(documents::table.find(document_id)) - .set(documents::updated_at.eq(Utc::now().naive_utc())) - .execute(conn)?; + diesel::update( + documents::table + .find(document_id) + .filter(documents::tenant_id.eq(tenant_id)), + ) + .set(documents::updated_at.eq(Utc::now().naive_utc())) + .execute(conn)?; } Ok(()) @@ -1265,8 +1364,12 @@ pub async fn assign_correspondents( } pub async fn bulk_assign_correspondents( - State(state): State, - user: AuthenticatedUser, + TenantScopedConn { + mut conn, + tenant_id, + user_id, + .. + }: TenantScopedConn, Json(payload): Json, ) -> AppResult<(StatusCode, Json)> { if payload.document_ids.is_empty() { @@ -1283,12 +1386,11 @@ pub async fn bulk_assign_correspondents( let (normalized_pairs, correspondents_vec, roles_vec) = normalize_correspondent_assignments(&payload.assignments)?; let action = payload.action; - let user_id = user.user_id; - - let mut conn = state.db()?; + let user_id_val = user_id; let (assigned, removed) = conn.transaction::<(usize, usize), AppError, _>(|conn| { let docs: Vec<(Uuid, Option)> = documents::table .filter(documents::id.eq_any(&document_ids)) + .filter(documents::tenant_id.eq(tenant_id)) .select((documents::id, documents::deleted_at)) .load(conn)?; @@ -1307,6 +1409,7 @@ pub async fn bulk_assign_correspondents( if !correspondents_vec.is_empty() { let existing: Vec = correspondents::table .filter(correspondents::id.eq_any(&correspondents_vec)) + .filter(correspondents::tenant_id.eq(tenant_id)) .load(conn)?; if existing.len() != correspondents_vec.len() { return Err(AppError::bad_request( @@ -1322,6 +1425,7 @@ pub async fn bulk_assign_correspondents( removed = diesel::delete( document_correspondents::table .filter(document_correspondents::document_id.eq_any(&document_ids)) + .filter(document_correspondents::tenant_id.eq(tenant_id)) .filter(document_correspondents::role.eq_any(&roles_vec)), ) .execute(conn)?; @@ -1334,7 +1438,8 @@ pub async fn bulk_assign_correspondents( document_id: *doc_id, correspondent_id: *correspondent_id, role: role.clone(), - assigned_by: Some(user_id), + assigned_by: Some(user_id_val), + tenant_id, }); } } @@ -1349,9 +1454,13 @@ pub async fn bulk_assign_correspondents( }; if assigned > 0 || removed > 0 { - diesel::update(documents::table.filter(documents::id.eq_any(&document_ids))) - .set(documents::updated_at.eq(Utc::now().naive_utc())) - .execute(conn)?; + diesel::update( + documents::table + .filter(documents::id.eq_any(&document_ids)) + .filter(documents::tenant_id.eq(tenant_id)), + ) + .set(documents::updated_at.eq(Utc::now().naive_utc())) + .execute(conn)?; } Ok((assigned, removed)) @@ -1371,6 +1480,7 @@ pub async fn bulk_assign_correspondents( removed += diesel::delete( document_correspondents::table .filter(document_correspondents::document_id.eq_any(&document_ids)) + .filter(document_correspondents::tenant_id.eq(tenant_id)) .filter(document_correspondents::role.eq(role.as_str())) .filter(document_correspondents::correspondent_id.eq_any(&ids)), ) @@ -1379,9 +1489,13 @@ pub async fn bulk_assign_correspondents( } if removed > 0 { - diesel::update(documents::table.filter(documents::id.eq_any(&document_ids))) - .set(documents::updated_at.eq(Utc::now().naive_utc())) - .execute(conn)?; + diesel::update( + documents::table + .filter(documents::id.eq_any(&document_ids)) + .filter(documents::tenant_id.eq(tenant_id)), + ) + .set(documents::updated_at.eq(Utc::now().naive_utc())) + .execute(conn)?; } Ok((0, removed)) @@ -1396,9 +1510,13 @@ pub async fn bulk_assign_correspondents( } pub async fn remove_correspondent( - State(state): State, Path((document_id, correspondent_id)): Path<(Uuid, Uuid)>, Query(query): Query, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, ) -> AppResult { let role = normalize_role(&query.role); if role.is_empty() { @@ -1411,8 +1529,10 @@ pub async fn remove_correspondent( ))); } - let mut conn = state.db()?; - let document: Document = documents::table.find(document_id).first(&mut conn)?; + let document: Document = documents::table + .find(document_id) + .filter(documents::tenant_id.eq(tenant_id)) + .first(&mut conn)?; if document.deleted_at.is_some() { return Err(AppError::not_found()); } @@ -1420,6 +1540,7 @@ pub async fn remove_correspondent( let deleted = diesel::delete( document_correspondents::table .filter(document_correspondents::document_id.eq(document_id)) + .filter(document_correspondents::tenant_id.eq(tenant_id)) .filter(document_correspondents::correspondent_id.eq(correspondent_id)) .filter(document_correspondents::role.eq(&role)), ) @@ -1429,33 +1550,41 @@ pub async fn remove_correspondent( return Err(AppError::not_found()); } - diesel::update(documents::table.find(document_id)) - .set(documents::updated_at.eq(Utc::now().naive_utc())) - .execute(&mut conn)?; + diesel::update( + documents::table + .find(document_id) + .filter(documents::tenant_id.eq(tenant_id)), + ) + .set(documents::updated_at.eq(Utc::now().naive_utc())) + .execute(&mut conn)?; Ok(StatusCode::NO_CONTENT) } pub async fn assign_tags( - State(state): State, Path(document_id): Path, - user: AuthenticatedUser, + TenantScopedConn { + mut conn, + tenant_id, + user_id, + .. + }: TenantScopedConn, Json(payload): Json, ) -> AppResult { if payload.tag_ids.is_empty() { return Err(AppError::bad_request("tag_ids must not be empty")); } - let mut conn = state.db()?; - // Ensure document exists documents::table .find(document_id) + .filter(documents::tenant_id.eq(tenant_id)) .first::(&mut conn)?; // Ensure tags exist let existing_tags: Vec = tags::table .filter(tags::id.eq_any(&payload.tag_ids)) + .filter(tags::tenant_id.eq(tenant_id)) .load(&mut conn)?; if existing_tags.len() != payload.tag_ids.len() { return Err(AppError::bad_request("one or more tags do not exist")); @@ -1467,7 +1596,8 @@ pub async fn assign_tags( .map(|tag_id| NewDocumentTag { document_id, tag_id: *tag_id, - assigned_by: Some(user.user_id), + assigned_by: Some(user_id), + tenant_id, }) .collect(); @@ -1480,8 +1610,12 @@ pub async fn assign_tags( } pub async fn bulk_update_tags( - State(state): State, - user: AuthenticatedUser, + TenantScopedConn { + mut conn, + tenant_id, + user_id, + .. + }: TenantScopedConn, Json(payload): Json, ) -> AppResult<(StatusCode, Json)> { let BulkTagRequest { @@ -1502,10 +1636,9 @@ pub async fn bulk_update_tags( tag_ids.sort(); tag_ids.dedup(); - let mut conn = state.db()?; - let docs: Vec<(Uuid, Option)> = documents::table .filter(documents::id.eq_any(&document_ids)) + .filter(documents::tenant_id.eq(tenant_id)) .select((documents::id, documents::deleted_at)) .load(&mut conn)?; @@ -1523,6 +1656,7 @@ pub async fn bulk_update_tags( let existing_tags: Vec = tags::table .filter(tags::id.eq_any(&tag_ids)) + .filter(tags::tenant_id.eq(tenant_id)) .load(&mut conn)?; if existing_tags.len() != tag_ids.len() { return Err(AppError::bad_request("one or more tags do not exist")); @@ -1536,7 +1670,8 @@ pub async fn bulk_update_tags( inserts.push(NewDocumentTag { document_id: *doc_id, tag_id: *tag_id, - assigned_by: Some(user.user_id), + assigned_by: Some(user_id), + tenant_id, }); } } @@ -1556,6 +1691,7 @@ pub async fn bulk_update_tags( let removed = diesel::delete( document_tags::table .filter(document_tags::document_id.eq_any(&document_ids)) + .filter(document_tags::tenant_id.eq(tenant_id)) .filter(document_tags::tag_id.eq_any(&tag_ids)), ) .execute(&mut conn)?; @@ -1568,13 +1704,17 @@ pub async fn bulk_update_tags( } pub async fn remove_tag( - State(state): State, Path((document_id, tag_id)): Path<(Uuid, Uuid)>, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, ) -> AppResult { - let mut conn = state.db()?; diesel::delete( document_tags::table .filter(document_tags::document_id.eq(document_id)) + .filter(document_tags::tenant_id.eq(tenant_id)) .filter(document_tags::tag_id.eq(tag_id)), ) .execute(&mut conn)?; @@ -1585,6 +1725,7 @@ pub async fn remove_tag( async fn process_upload( state: &AppState, request: UploadRequest, + tenant_id: Uuid, user_id: Uuid, ) -> AppResult { let UploadRequest { @@ -1596,7 +1737,7 @@ async fn process_upload( } = request; if let Some(folder) = folder_id { - ensure_folder_exists(state, folder)?; + ensure_folder_exists(state, tenant_id, folder)?; } let doc_id = Uuid::new_v4(); @@ -1610,13 +1751,14 @@ async fn process_upload( let s3_key = document_version_object_key(doc_id, version_number, version_id); { - let mut conn = state.db()?; + let mut conn = state.db_for_tenant(tenant_id)?; let existing = documents::table .inner_join( document_versions::table .on(document_versions::id.eq(documents::current_version_id)), ) + .filter(documents::tenant_id.eq(tenant_id)) .filter(document_versions::checksum.eq(&checksum_hex)) .select((documents::all_columns, document_versions::all_columns)) .first::<(Document, DocumentVersion)>(&mut conn) @@ -1641,7 +1783,7 @@ async fn process_upload( let tags = tags_map.get(&document.id).cloned(); let correspondents = correspondents_map.remove(&document.id).unwrap_or_default(); drop(conn); - let assets = load_asset_responses(state, version.id).await?; + let assets = load_asset_responses(state, tenant_id, version.id).await?; let version_response = to_version_response(version.clone(), true); info!( @@ -1689,7 +1831,7 @@ async fn process_upload( }; let (document, version) = { - let mut conn = state.db()?; + let mut conn = state.db_for_tenant(tenant_id)?; conn.transaction(|conn| { let new_document = NewDocument { id: doc_id, @@ -1701,6 +1843,7 @@ async fn process_upload( issued_at: None, title: derive_document_title(&original_name), metadata: metadata_value.clone(), + tenant_id, }; diesel::insert_into(documents::table) .values(&new_document) @@ -1715,6 +1858,7 @@ async fn process_upload( checksum: checksum_hex.clone(), metadata: Value::Object(Default::default()), operations_summary: Value::Object(Default::default()), + tenant_id, }; diesel::insert_into(document_versions::table) @@ -1739,9 +1883,10 @@ async fn process_upload( )?, }; - if let Ok(mut conn) = state.db() { + if let Ok(mut conn) = state.db_for_tenant(tenant_id) { if let Err(err) = enqueue_job( &mut conn, + tenant_id, JOB_ANALYZE_DOCUMENT, json!({ "document_id": doc_id, @@ -1762,10 +1907,14 @@ async fn process_upload( }) } -fn ensure_folder_exists(state: &AppState, folder_id: Uuid) -> AppResult<()> { - let mut conn = state.db()?; - let exists: bool = diesel::select(exists(folders::table.filter(folders::id.eq(folder_id)))) - .get_result(&mut conn)?; +fn ensure_folder_exists(state: &AppState, tenant_id: Uuid, folder_id: Uuid) -> AppResult<()> { + let mut conn = state.db_for_tenant(tenant_id)?; + let exists: bool = diesel::select(exists( + folders::table + .filter(folders::id.eq(folder_id)) + .filter(folders::tenant_id.eq(tenant_id)), + )) + .get_result(&mut conn)?; if !exists { return Err(AppError::bad_request("folder does not exist")); } @@ -1829,6 +1978,7 @@ pub(crate) fn load_correspondents_for_documents( pub(crate) async fn load_primary_assets( state: &AppState, + tenant_id: Uuid, documents: &[Document], ) -> AppResult)>> { if documents.is_empty() { @@ -1845,7 +1995,7 @@ pub(crate) async fn load_primary_assets( version_ids.sort(); version_ids.dedup(); - let mut conn = state.db()?; + let mut conn = state.db_for_tenant(tenant_id)?; let versions: Vec = document_versions::table .filter(document_versions::id.eq_any(&version_ids)) .load(&mut conn)?; @@ -1905,7 +2055,7 @@ pub(crate) fn to_document_response( current_version: Option<(DocumentVersionResponse, Vec)>, ) -> AppResult { let current_version = if let Some((version, assets)) = current_version { - let download_path = build_download_path(state, doc.id, user_id)?; + let download_path = build_download_path(state, &doc, user_id)?; Some(DocumentCurrentVersionResponse { version, assets, @@ -1937,10 +2087,10 @@ pub(crate) fn to_document_response( }) } -fn build_download_path(state: &AppState, document_id: Uuid, user_id: Uuid) -> AppResult { +fn build_download_path(state: &AppState, document: &Document, user_id: Uuid) -> AppResult { state .jwt - .generate_download_token(document_id, user_id) + .generate_download_token(document.id, user_id, document.tenant_id) .map(|token| format!("/download/{token}")) .map_err(|err| AppError::internal(format!("failed to generate download token: {err}"))) } @@ -2042,9 +2192,10 @@ fn filename_with_retained_extension(title: &str, current_filename: &str) -> Stri async fn load_asset_responses( state: &AppState, + tenant_id: Uuid, version_id: Uuid, ) -> AppResult> { - let mut conn = state.db()?; + let mut conn = state.db_for_tenant(tenant_id)?; let assets: Vec<(DocumentAsset, Option)> = document_assets::table .left_outer_join( document_asset_objects::table.on(document_asset_objects::asset_id @@ -2052,6 +2203,7 @@ async fn load_asset_responses( .and(document_asset_objects::ordinal.eq(1))), ) .filter(document_assets::document_version_id.eq(version_id)) + .filter(document_assets::tenant_id.eq(tenant_id)) .order(document_assets::created_at.asc()) .select(( document_assets::all_columns, diff --git a/backend/src/routes/folders.rs b/backend/src/routes/folders.rs index d907d0b..ea05d52 100644 --- a/backend/src/routes/folders.rs +++ b/backend/src/routes/folders.rs @@ -10,7 +10,7 @@ use crate::models::{Document, Folder, NewFolder}; use crate::schema::{documents, folders}; use crate::state::AppState; use crate::{ - auth::AuthenticatedUser, + auth::TenantScopedConn, error::{AppError, AppResult}, }; @@ -70,15 +70,17 @@ pub struct FolderInfo { } pub async fn ensure_folder_path( - State(state): State, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, Json(payload): Json, ) -> AppResult> { if payload.segments.is_empty() { return Err(AppError::bad_request("segments must not be empty")); } - let mut conn = state.db()?; - let target_folder = conn.transaction::(|conn| { let mut current_parent = payload.parent_id; let mut last_folder: Option = None; @@ -93,12 +95,14 @@ pub async fn ensure_folder_path( folders::table .filter(folders::parent_id.eq(Some(parent_id))) .filter(folders::name.eq(name)) + .filter(folders::tenant_id.eq(tenant_id)) .first(conn) .optional()? } else { folders::table .filter(folders::parent_id.is_null()) .filter(folders::name.eq(name)) + .filter(folders::tenant_id.eq(tenant_id)) .first(conn) .optional()? }; @@ -110,6 +114,7 @@ pub async fn ensure_folder_path( id: Uuid::new_v4(), name: name.to_string(), parent_id: current_parent, + tenant_id, }; diesel::insert_into(folders::table) @@ -132,19 +137,22 @@ pub async fn ensure_folder_path( } pub async fn create_folder( - State(state): State, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, Json(payload): Json, ) -> AppResult> { if payload.name.trim().is_empty() { return Err(AppError::bad_request("name must not be empty")); } - let mut conn = state.db()?; - let new_folder = NewFolder { id: Uuid::new_v4(), name: payload.name.trim().to_string(), parent_id: payload.parent_id, + tenant_id, }; diesel::insert_into(folders::table) @@ -161,10 +169,13 @@ pub async fn list_folder_contents( State(state): State, Path(folder_identifier): Path, Query(query): Query, - user: AuthenticatedUser, + TenantScopedConn { + mut conn, + tenant_id, + user_id, + .. + }: TenantScopedConn, ) -> AppResult> { - let mut conn = state.db()?; - let folder_id = if folder_identifier.eq_ignore_ascii_case("root") { None } else { @@ -176,7 +187,10 @@ pub async fn list_folder_contents( let folder = match folder_id { Some(id) => Some(folder_to_info( - folders::table.find(id).first::(&mut conn)?, + folders::table + .find(id) + .filter(folders::tenant_id.eq(tenant_id)) + .first::(&mut conn)?, )), None => None, }; @@ -184,11 +198,13 @@ pub async fn list_folder_contents( let child_folders: Vec = if let Some(parent_id) = folder_id { folders::table .filter(folders::parent_id.eq(parent_id)) + .filter(folders::tenant_id.eq(tenant_id)) .order(folders::name.asc()) .load(&mut conn)? } else { folders::table .filter(folders::parent_id.is_null()) + .filter(folders::tenant_id.eq(tenant_id)) .order(folders::name.asc()) .load(&mut conn)? }; @@ -197,6 +213,7 @@ pub async fn list_folder_contents( let documents = if query.include_documents { let docs_query = documents::table .filter(documents::deleted_at.is_null()) + .filter(documents::tenant_id.eq(tenant_id)) .order(documents::uploaded_at.desc()); let docs: Vec = if let Some(current_folder) = folder_id { @@ -214,7 +231,7 @@ pub async fn list_folder_contents( let mut correspondents_map = load_correspondents_for_documents(&mut conn, &doc_ids)?; drop(conn); - let primary_versions = load_primary_assets(&state, &docs).await?; + let primary_versions = load_primary_assets(&state, tenant_id, &docs).await?; let mut documents = Vec::with_capacity(doc_ids.len()); for doc in docs { @@ -223,7 +240,7 @@ pub async fn list_folder_contents( let current_version = primary_versions.get(&doc.id).cloned(); documents.push(to_document_response( &state, - user.user_id, + user_id, doc, tags, correspondents, @@ -244,16 +261,23 @@ pub async fn list_folder_contents( } pub async fn delete_folder( - State(state): State, Path(folder_id): Path, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, ) -> AppResult { - let mut conn = state.db()?; - conn.transaction::<_, AppError, _>(|conn| { - folders::table.find(folder_id).first::(conn)?; + folders::table + .find(folder_id) + .filter(folders::tenant_id.eq(tenant_id)) + .first::(conn)?; let has_child_folders: bool = diesel::select(exists( - folders::table.filter(folders::parent_id.eq(Some(folder_id))), + folders::table + .filter(folders::parent_id.eq(Some(folder_id))) + .filter(folders::tenant_id.eq(tenant_id)), )) .get_result(conn)?; @@ -266,6 +290,7 @@ pub async fn delete_folder( let has_documents: bool = diesel::select(exists( documents::table .filter(documents::folder_id.eq(Some(folder_id))) + .filter(documents::tenant_id.eq(tenant_id)) .filter(documents::deleted_at.is_null()), )) .get_result(conn)?; @@ -276,7 +301,12 @@ pub async fn delete_folder( )); } - diesel::delete(folders::table.find(folder_id)).execute(conn)?; + diesel::delete( + folders::table + .filter(folders::id.eq(folder_id)) + .filter(folders::tenant_id.eq(tenant_id)), + ) + .execute(conn)?; Ok(()) })?; @@ -285,14 +315,19 @@ pub async fn delete_folder( } pub async fn update_folder( - State(state): State, Path(folder_id): Path, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, Json(payload): Json, ) -> AppResult { - let mut conn = state.db()?; - conn.transaction::<(), AppError, _>(|conn| { - let folder: Folder = folders::table.find(folder_id).first(conn)?; + let folder: Folder = folders::table + .find(folder_id) + .filter(folders::tenant_id.eq(tenant_id)) + .first(conn)?; let mut next_parent = folder.parent_id; let mut parent_changed = false; @@ -303,9 +338,12 @@ pub async fn update_folder( } if let Some(parent_id) = parent_request { - let _parent: Folder = folders::table.find(parent_id).first(conn)?; + let _parent: Folder = folders::table + .find(parent_id) + .filter(folders::tenant_id.eq(tenant_id)) + .first(conn)?; - let descendant_ids = gather_descendant_folder_ids(conn, folder_id)?; + let descendant_ids = gather_descendant_folder_ids(conn, tenant_id, folder_id)?; if descendant_ids.contains(&parent_id) { return Err(AppError::bad_request( "cannot move folder into itself or a descendant", @@ -341,6 +379,7 @@ pub async fn update_folder( .filter(folders::parent_id.eq(Some(parent_id))) .filter(folders::name.eq(&new_name)) .filter(folders::id.ne(folder_id)) + .filter(folders::tenant_id.eq(tenant_id)) .first::(conn) .optional()? } else { @@ -348,6 +387,7 @@ pub async fn update_folder( .filter(folders::parent_id.is_null()) .filter(folders::name.eq(&new_name)) .filter(folders::id.ne(folder_id)) + .filter(folders::tenant_id.eq(tenant_id)) .first::(conn) .optional()? }; @@ -358,12 +398,16 @@ pub async fn update_folder( )); } - diesel::update(folders::table.find(folder_id)) - .set(( - folders::parent_id.eq(next_parent), - folders::name.eq(&new_name), - )) - .execute(conn)?; + diesel::update( + folders::table + .find(folder_id) + .filter(folders::tenant_id.eq(tenant_id)), + ) + .set(( + folders::parent_id.eq(next_parent), + folders::name.eq(&new_name), + )) + .execute(conn)?; Ok(()) })?; @@ -383,6 +427,7 @@ fn folder_to_info(folder: Folder) -> FolderInfo { pub(super) fn gather_descendant_folder_ids( conn: &mut PgConnection, + tenant_id: Uuid, folder_id: Uuid, ) -> AppResult> { let mut ids = vec![folder_id]; @@ -391,6 +436,7 @@ pub(super) fn gather_descendant_folder_ids( while let Some(current) = queue.pop() { let child_ids: Vec = folders::table .filter(folders::parent_id.eq(Some(current))) + .filter(folders::tenant_id.eq(tenant_id)) .select(folders::id) .load(conn)?; queue.extend(child_ids.iter().copied()); diff --git a/backend/src/routes/mod.rs b/backend/src/routes/mod.rs index d1bcf57..f794f5d 100644 --- a/backend/src/routes/mod.rs +++ b/backend/src/routes/mod.rs @@ -50,6 +50,7 @@ pub fn create_router(state: AppState) -> Router<()> { .route("/login", post(auth::login)) .route("/refresh", post(auth::refresh)) .route("/logout", post(auth::logout)) + .route("/select-tenant", post(auth::select_tenant)) .route("/me", get(auth::me)); let documents_routes = Router::new() diff --git a/backend/src/routes/tags.rs b/backend/src/routes/tags.rs index 2476e19..1e7ed9c 100644 --- a/backend/src/routes/tags.rs +++ b/backend/src/routes/tags.rs @@ -1,19 +1,15 @@ use crate::utils::json::{classify_nullable, NullableValue}; -use axum::{ - extract::{Path, State}, - http::StatusCode, - Json, -}; +use axum::{extract::Path, http::StatusCode, Json}; use diesel::{dsl::count_star, prelude::*}; use serde::{Deserialize, Serialize}; use serde_json::Value; use std::collections::HashMap; use uuid::Uuid; +use crate::auth::TenantScopedConn; use crate::error::{AppError, AppResult}; use crate::models::{NewTag, Tag}; use crate::schema::{document_tags, tags}; -use crate::state::AppState; #[derive(Deserialize)] pub struct CreateTagRequest { @@ -36,12 +32,17 @@ pub struct TagCatalogEntry { pub usage_count: i64, } -pub async fn list_tags(State(state): State) -> AppResult>> { - let mut conn = state.db()?; - +pub async fn list_tags( + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, +) -> AppResult>> { let tag_list: Vec = tags::table.order(tags::label.asc()).load(&mut conn)?; let usage_rows: Vec<(Uuid, i64)> = document_tags::table + .filter(document_tags::tenant_id.eq(tenant_id)) .group_by(document_tags::tag_id) .select((document_tags::tag_id, count_star())) .load(&mut conn)?; @@ -62,18 +63,22 @@ pub async fn list_tags(State(state): State) -> AppResult, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, Json(payload): Json, ) -> AppResult> { if payload.label.trim().is_empty() { return Err(AppError::bad_request("label must not be empty")); } - let mut conn = state.db()?; let new_tag = NewTag { id: Uuid::new_v4(), label: payload.label.trim().to_string(), color: payload.color, + tenant_id, }; match diesel::insert_into(tags::table) @@ -90,7 +95,10 @@ pub async fn create_tag( Err(err) => return Err(AppError::from(err)), } - let tag: Tag = tags::table.find(new_tag.id).first(&mut conn)?; + let tag: Tag = tags::table + .find(new_tag.id) + .filter(tags::tenant_id.eq(tenant_id)) + .first(&mut conn)?; Ok(Json(TagCatalogEntry { id: tag.id, label: tag.label, @@ -100,12 +108,18 @@ pub async fn create_tag( } pub async fn update_tag( - State(state): State, Path(tag_id): Path, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, Json(body): Json, ) -> AppResult> { - let mut conn = state.db()?; - let existing: Tag = tags::table.find(tag_id).first(&mut conn)?; + let existing: Tag = tags::table + .find(tag_id) + .filter(tags::tenant_id.eq(tenant_id)) + .first(&mut conn)?; let label_class = classify_nullable(body.get("label")).map_err(AppError::bad_request)?; let color_class = classify_nullable(body.get("color")).map_err(AppError::bad_request)?; @@ -140,6 +154,7 @@ pub async fn update_tag( let duplicate = tags::table .filter(tags::label.eq(trimmed)) .filter(tags::id.ne(tag_id)) + .filter(tags::tenant_id.eq(tenant_id)) .first::(&mut conn) .optional()?; if duplicate.is_some() { @@ -174,6 +189,7 @@ pub async fn update_tag( if !label_changed && !color_changed { let usage_count: i64 = document_tags::table .filter(document_tags::tag_id.eq(tag_id)) + .filter(document_tags::tenant_id.eq(tenant_id)) .select(count_star()) .first(&mut conn)?; return Ok(Json(TagCatalogEntry { @@ -191,13 +207,21 @@ pub async fn update_tag( .map(|opt| opt.as_ref().map(|value| value.as_str())), }; - diesel::update(tags::table.find(tag_id)) - .set(&changeset) - .execute(&mut conn)?; + diesel::update( + tags::table + .find(tag_id) + .filter(tags::tenant_id.eq(tenant_id)), + ) + .set(&changeset) + .execute(&mut conn)?; - let updated: Tag = tags::table.find(tag_id).first(&mut conn)?; + let updated: Tag = tags::table + .find(tag_id) + .filter(tags::tenant_id.eq(tenant_id)) + .first(&mut conn)?; let usage_count: i64 = document_tags::table .filter(document_tags::tag_id.eq(tag_id)) + .filter(document_tags::tenant_id.eq(tenant_id)) .select(count_star()) .first(&mut conn)?; @@ -210,13 +234,16 @@ pub async fn update_tag( } pub async fn delete_tag( - State(state): State, Path(tag_id): Path, + TenantScopedConn { + mut conn, + tenant_id, + .. + }: TenantScopedConn, ) -> AppResult { - let mut conn = state.db()?; - let usage: i64 = document_tags::table .filter(document_tags::tag_id.eq(tag_id)) + .filter(document_tags::tenant_id.eq(tenant_id)) .select(count_star()) .first(&mut conn)?; @@ -226,7 +253,12 @@ pub async fn delete_tag( )); } - let deleted = diesel::delete(tags::table.find(tag_id)).execute(&mut conn)?; + let deleted = diesel::delete( + tags::table + .find(tag_id) + .filter(tags::tenant_id.eq(tenant_id)), + ) + .execute(&mut conn)?; if deleted == 0 { return Err(AppError::not_found()); } diff --git a/backend/src/routes/webdav/mod.rs b/backend/src/routes/webdav/mod.rs index fb128e0..d521b2e 100644 --- a/backend/src/routes/webdav/mod.rs +++ b/backend/src/routes/webdav/mod.rs @@ -221,7 +221,7 @@ fn fetch_folder_contents( state: &AppState, folder_id: Option, ) -> AppResult { - let mut conn = state.db()?; + let mut conn = state.db_unscoped()?; let folder = match folder_id { Some(id) => Some(folders_dsl::folders.find(id).first::(&mut conn)?), @@ -405,7 +405,7 @@ fn authenticate(state: &AppState, headers: &HeaderMap) -> Result AppResult = None; let mut chain: Vec = Vec::new(); let mut current_folder: Option = None; diff --git a/backend/src/schema.rs b/backend/src/schema.rs index 85591f6..b56d4ea 100644 --- a/backend/src/schema.rs +++ b/backend/src/schema.rs @@ -8,6 +8,7 @@ diesel::table! { metadata -> Jsonb, created_at -> Timestamptz, updated_at -> Timestamptz, + tenant_id -> Uuid, } } @@ -18,6 +19,7 @@ diesel::table! { ordinal -> Int4, s3_key -> Text, metadata -> Jsonb, + tenant_id -> Uuid, } } @@ -30,6 +32,7 @@ diesel::table! { metadata -> Jsonb, created_at -> Timestamptz, cardinality -> Nullable, + tenant_id -> Uuid, } } @@ -41,6 +44,7 @@ diesel::table! { role -> Varchar, assigned_at -> Timestamptz, assigned_by -> Nullable, + tenant_id -> Uuid, } } @@ -50,6 +54,7 @@ diesel::table! { tag_id -> Uuid, assigned_at -> Timestamptz, assigned_by -> Nullable, + tenant_id -> Uuid, } } @@ -66,6 +71,7 @@ diesel::table! { created_at -> Timestamptz, operations_summary -> Jsonb, metadata -> Jsonb, + tenant_id -> Uuid, } } @@ -87,6 +93,7 @@ diesel::table! { #[max_length = 255] title -> Varchar, current_version_id -> Uuid, + tenant_id -> Uuid, } } @@ -98,6 +105,7 @@ diesel::table! { parent_id -> Nullable, created_at -> Timestamptz, updated_at -> Timestamptz, + tenant_id -> Uuid, } } @@ -112,6 +120,7 @@ diesel::table! { last_error -> Nullable, created_at -> Timestamptz, updated_at -> Timestamptz, + tenant_id -> Uuid, } } @@ -125,6 +134,7 @@ diesel::table! { revoked_at -> Nullable, created_at -> Timestamptz, updated_at -> Timestamptz, + tenant_id -> Uuid, } } @@ -136,6 +146,31 @@ diesel::table! { #[max_length = 7] color -> Nullable, created_at -> Timestamptz, + tenant_id -> Uuid, + } +} + +diesel::table! { + tenants (tenant_id) { + tenant_id -> Uuid, + slug -> Text, + storage_root -> Nullable, + quickwit_index -> Nullable, + status -> Text, + config -> Jsonb, + created_at -> Timestamptz, + updated_at -> Timestamptz, + } +} + +diesel::table! { + user_memberships (id) { + id -> Uuid, + user_id -> Uuid, + tenant_id -> Uuid, + role -> Text, + created_at -> Timestamptz, + updated_at -> Timestamptz, } } @@ -146,23 +181,36 @@ diesel::table! { username -> Varchar, #[max_length = 255] password_hash -> Varchar, - #[max_length = 16] - role -> Varchar, created_at -> Timestamptz, updated_at -> Timestamptz, + tenant_id -> Uuid, } } +diesel::joinable!(correspondents -> tenants (tenant_id)); diesel::joinable!(document_asset_objects -> document_assets (asset_id)); +diesel::joinable!(document_asset_objects -> tenants (tenant_id)); diesel::joinable!(document_assets -> document_versions (document_version_id)); +diesel::joinable!(document_assets -> tenants (tenant_id)); diesel::joinable!(document_correspondents -> correspondents (correspondent_id)); diesel::joinable!(document_correspondents -> documents (document_id)); +diesel::joinable!(document_correspondents -> tenants (tenant_id)); diesel::joinable!(document_correspondents -> users (assigned_by)); diesel::joinable!(document_tags -> documents (document_id)); diesel::joinable!(document_tags -> tags (tag_id)); +diesel::joinable!(document_tags -> tenants (tenant_id)); diesel::joinable!(document_tags -> users (assigned_by)); +diesel::joinable!(document_versions -> tenants (tenant_id)); diesel::joinable!(documents -> folders (folder_id)); +diesel::joinable!(documents -> tenants (tenant_id)); +diesel::joinable!(folders -> tenants (tenant_id)); +diesel::joinable!(jobs -> tenants (tenant_id)); +diesel::joinable!(refresh_tokens -> tenants (tenant_id)); diesel::joinable!(refresh_tokens -> users (user_id)); +diesel::joinable!(tags -> tenants (tenant_id)); +diesel::joinable!(user_memberships -> tenants (tenant_id)); +diesel::joinable!(user_memberships -> users (user_id)); +diesel::joinable!(users -> tenants (tenant_id)); diesel::allow_tables_to_appear_in_same_query!( correspondents, @@ -176,5 +224,7 @@ diesel::allow_tables_to_appear_in_same_query!( jobs, refresh_tokens, tags, + tenants, + user_memberships, users, ); diff --git a/backend/src/state.rs b/backend/src/state.rs index d8ac37c..37c6e78 100644 --- a/backend/src/state.rs +++ b/backend/src/state.rs @@ -4,6 +4,7 @@ use diesel::{ pg::PgConnection, r2d2::{ConnectionManager, PooledConnection}, }; +use uuid::Uuid; use crate::{ auth::jwt::JwtService, @@ -11,9 +12,10 @@ use crate::{ db::PgPool, error::{AppError, AppResult}, storage::ObjectStorage, + tenants::{apply_tenant_guc, TenantService}, }; -type PgPooledConnection = PooledConnection>; +pub type PgPooledConnection = PooledConnection>; #[derive(Clone)] pub struct AppState { @@ -21,6 +23,7 @@ pub struct AppState { pub config: Arc, pub storage: Arc, pub jwt: JwtService, + pub tenants: TenantService, } impl AppState { @@ -30,15 +33,25 @@ impl AppState { storage: Arc, jwt: JwtService, ) -> Self { + let config = Arc::new(config); + let tenants = TenantService::new(pool.clone()); + Self { pool, - config: Arc::new(config), + config, storage, jwt, + tenants, } } - pub fn db(&self) -> AppResult { + pub fn db_for_tenant(&self, tenant_id: Uuid) -> AppResult { + let mut conn = self.db_unscoped()?; + apply_tenant_guc(&mut conn, tenant_id)?; + Ok(conn) + } + + pub fn db_unscoped(&self) -> AppResult { self.pool .get() .map_err(|err| AppError::internal(format!("database pool error: {err}"))) diff --git a/backend/src/tenants.rs b/backend/src/tenants.rs new file mode 100644 index 0000000..683ee01 --- /dev/null +++ b/backend/src/tenants.rs @@ -0,0 +1,123 @@ +use std::{ + collections::HashMap, + sync::{Arc, RwLock}, +}; + +use axum::{async_trait, extract::FromRequestParts, http::request::Parts}; +use diesel::{pg::PgConnection, prelude::*, sql_types::Text}; +use uuid::Uuid; + +use crate::{ + db::PgPool, + error::{AppError, AppResult}, + models::Tenant, + schema::tenants::dsl, + state::AppState, +}; + +pub struct TenantRepository; + +impl TenantRepository { + pub fn get_by_id(conn: &mut PgConnection, tenant_id: Uuid) -> AppResult { + dsl::tenants.find(tenant_id).first(conn).map_err(Into::into) + } + + pub fn get_by_slug(conn: &mut PgConnection, slug: &str) -> AppResult { + dsl::tenants + .filter(dsl::slug.eq(slug)) + .first(conn) + .map_err(Into::into) + } +} + +#[derive(Clone)] +pub struct TenantService { + pool: PgPool, + cache_by_id: Arc>>, + cache_by_slug: Arc>>, +} + +impl TenantService { + pub fn new(pool: PgPool) -> Self { + Self { + pool, + cache_by_id: Arc::new(RwLock::new(HashMap::new())), + cache_by_slug: Arc::new(RwLock::new(HashMap::new())), + } + } + + pub fn get_by_id(&self, tenant_id: Uuid) -> AppResult { + if let Some(tenant) = self.cache_by_id.read().unwrap().get(&tenant_id) { + return Ok(tenant.clone()); + } + + let tenant = self.load(|conn| TenantRepository::get_by_id(conn, tenant_id))?; + self.store(&tenant); + Ok(tenant) + } + + pub fn get_by_slug(&self, slug: &str) -> AppResult { + if let Some(tenant) = self.cache_by_slug.read().unwrap().get(slug) { + return Ok(tenant.clone()); + } + + let slug_owned = slug.to_owned(); + let tenant = self.load(|conn| TenantRepository::get_by_slug(conn, &slug_owned))?; + self.store(&tenant); + Ok(tenant) + } + + pub fn tenant_id_for_slug(&self, slug: &str) -> AppResult { + Ok(self.get_by_slug(slug)?.tenant_id) + } + + fn load(&self, loader: F) -> AppResult + where + F: FnOnce(&mut PgConnection) -> AppResult, + { + let mut conn = self + .pool + .get() + .map_err(|err| AppError::internal(format!("database pool error: {err}")))?; + let tenant = loader(&mut conn)?; + Ok(tenant) + } + + fn store(&self, tenant: &Tenant) { + { + let mut by_id = self.cache_by_id.write().unwrap(); + by_id.insert(tenant.tenant_id, tenant.clone()); + } + { + let mut by_slug = self.cache_by_slug.write().unwrap(); + by_slug.insert(tenant.slug.clone(), tenant.clone()); + } + } +} + +pub fn apply_tenant_guc(conn: &mut PgConnection, tenant_id: Uuid) -> AppResult<()> { + diesel::sql_query("SELECT set_config('papercrate.tenant_id', $1, true)") + .bind::(tenant_id.to_string()) + .execute(conn) + .map(|_| ()) + .map_err(AppError::from) +} + +pub struct TenantContext { + pub tenant: Tenant, +} + +#[async_trait] +impl FromRequestParts for TenantContext { + type Rejection = AppError; + + async fn from_request_parts( + _parts: &mut Parts, + state: &AppState, + ) -> Result { + let tenant = state + .tenants + .get_by_slug(&state.config.default_tenant_slug)?; + Ok(Self { tenant }) + } +} diff --git a/backend/src/workers/analyze.rs b/backend/src/workers/analyze.rs index 9e8186d..df8270c 100644 --- a/backend/src/workers/analyze.rs +++ b/backend/src/workers/analyze.rs @@ -51,7 +51,9 @@ impl JobHandler for AnalyzeDocumentJob { }; let state_clone = state.clone(); - match task::spawn_blocking(move || analyze_document(state_clone, payload)).await { + let tenant_id = job.tenant_id; + match task::spawn_blocking(move || analyze_document(state_clone, tenant_id, payload)).await + { Ok(Ok(execution)) => execution, Ok(Err(err)) => { warn!(job_id = %job.id, error = %err, "analyze job will retry"); @@ -71,8 +73,14 @@ impl JobHandler for AnalyzeDocumentJob { } } -fn analyze_document(state: Arc, payload: AnalyzePayload) -> Result { - let mut conn = state.db().map_err(|err| format!("{err:?}"))?; +fn analyze_document( + state: Arc, + tenant_id: Uuid, + payload: AnalyzePayload, +) -> Result { + let mut conn = state + .db_for_tenant(tenant_id) + .map_err(|err| format!("{err:?}"))?; let version: DocumentVersion = document_versions::table .find(payload.document_version_id) @@ -88,12 +96,15 @@ fn analyze_document(state: Arc, payload: AnalyzePayload) -> Result = document_assets::table .filter(document_assets::document_version_id.eq(payload.document_version_id)) .filter(document_assets::asset_type.eq(OCR_TEXT_ASSET_TYPE)) + .filter(document_assets::tenant_id.eq(tenant_id)) .first(&mut conn) .optional() .map_err(|err| format!("{err:?}"))?; @@ -129,6 +140,7 @@ fn analyze_document(state: Arc, payload: AnalyzePayload) -> Result, payload: AnalyzePayload) -> Result, payload: &IndexPayload) -> Result { - let mut conn = state.db().map_err(|err| format!("{err:?}"))?; + let mut base_conn = state.db_unscoped().map_err(|err| format!("{err:?}"))?; let version: DocumentVersion = document_versions::table .find(payload.document_version_id) - .first(&mut conn) + .first(&mut base_conn) .map_err(|err| format!("{err:?}"))?; if version.document_id != payload.document_id { @@ -186,7 +186,14 @@ fn load_context(state: Arc, payload: &IndexPayload) -> Result = document_asset_objects::table @@ -195,7 +202,9 @@ fn load_context(state: Arc, payload: &IndexPayload) -> Result conn, Err(err) => { error!(?err, "failed to obtain database connection in worker"); @@ -87,7 +87,7 @@ impl Worker { let result = handler.handle(self.state.clone(), job.clone()).await; match result { JobExecution::Success => { - if let Ok(mut conn) = self.state.db() { + if let Ok(mut conn) = self.state.db_unscoped() { mark_job_succeeded(&mut conn, job.id)?; info!(job_id = %job.id, job_type = %job.job_type, "job completed successfully"); } else { @@ -96,7 +96,7 @@ impl Worker { } JobExecution::Retry { delay, error } => { warn!(job_id = %job.id, job_type = %job.job_type, %error, "job will retry"); - if let Ok(mut conn) = self.state.db() { + if let Ok(mut conn) = self.state.db_unscoped() { retry_job_after(&mut conn, job.id, delay, &error)?; } else { error!("failed to requeue job for retry due to pool error"); @@ -104,7 +104,7 @@ impl Worker { } JobExecution::Failed { error } => { error!(job_id = %job.id, job_type = %job.job_type, %error, "job failed"); - if let Ok(mut conn) = self.state.db() { + if let Ok(mut conn) = self.state.db_unscoped() { mark_job_failed(&mut conn, job.id, &error)?; } else { error!("failed to mark job failed due to pool error"); @@ -113,7 +113,7 @@ impl Worker { } } else { error!(job_type = %job.job_type, "no handler registered for job type"); - if let Ok(mut conn) = self.state.db() { + if let Ok(mut conn) = self.state.db_unscoped() { mark_job_failed(&mut conn, job.id, "no handler registered")?; } else { error!("failed to mark job failed for missing handler due to pool error"); diff --git a/backend/src/workers/ocr.rs b/backend/src/workers/ocr.rs index d79a04f..8450157 100644 --- a/backend/src/workers/ocr.rs +++ b/backend/src/workers/ocr.rs @@ -213,11 +213,11 @@ struct OcrGeneration { } fn load_ocr_context(state: Arc, payload: &OcrPayload) -> Result { - let mut conn = state.db().map_err(|err| format!("{err:?}"))?; + let mut base_conn = state.db_unscoped().map_err(|err| format!("{err:?}"))?; let version: DocumentVersion = document_versions::table .find(payload.document_version_id) - .first(&mut conn) + .first(&mut base_conn) .map_err(|err| format!("{err:?}"))?; if version.document_id != payload.document_id { @@ -226,12 +226,20 @@ fn load_ocr_context(state: Arc, payload: &OcrPayload) -> Result = document_assets::table .filter(document_assets::document_version_id.eq(payload.document_version_id)) .filter(document_assets::asset_type.eq(OCR_TEXT_ASSET_TYPE)) + .filter(document_assets::tenant_id.eq(tenant_id)) .first(&mut conn) .optional() .map_err(|err| format!("{err:?}"))?; @@ -239,6 +247,7 @@ fn load_ocr_context(state: Arc, payload: &OcrPayload) -> Result = if let Some(asset) = &existing_asset { document_asset_objects::table .filter(document_asset_objects::asset_id.eq(asset.id)) + .filter(document_asset_objects::tenant_id.eq(tenant_id)) .order(document_asset_objects::ordinal.asc()) .load(&mut conn) .map_err(|err| format!("{err:?}"))? @@ -392,7 +401,10 @@ fn persist_ocr_metadata( s3_key: &str, source: &'static str, ) -> Result<(), String> { - let mut conn = state.db().map_err(|err| format!("{err:?}"))?; + let tenant_id = context.document.tenant_id; + let mut conn = state + .db_for_tenant(tenant_id) + .map_err(|err| format!("{err:?}"))?; if let Some(existing_asset) = &context.existing_asset { diesel::delete(document_assets::table.filter(document_assets::id.eq(existing_asset.id))) @@ -410,6 +422,7 @@ fn persist_ocr_metadata( "source": source, }), cardinality: Some(1), + tenant_id, }; diesel::insert_into(document_assets::table) @@ -430,6 +443,7 @@ fn persist_ocr_metadata( let existing_object_id: Option = document_asset_objects::table .filter(document_asset_objects::asset_id.eq(asset_id)) .filter(document_asset_objects::ordinal.eq(1)) + .filter(document_asset_objects::tenant_id.eq(tenant_id)) .select(document_asset_objects::id) .first(&mut conn) .optional() @@ -443,6 +457,7 @@ fn persist_ocr_metadata( ordinal: 1, s3_key: s3_key.to_string(), metadata: json!({}), + tenant_id, }; diesel::insert_into(document_asset_objects::table) @@ -463,9 +478,20 @@ fn persist_ocr_metadata( } fn enqueue_index_job(state: &AppState, payload: &OcrPayload) -> Result<(), String> { - let mut conn = state.db().map_err(|err| format!("{err:?}"))?; + let mut base_conn = state.db_unscoped().map_err(|err| format!("{err:?}"))?; + let tenant_id: Uuid = documents::table + .find(payload.document_id) + .select(documents::tenant_id) + .first(&mut base_conn) + .map_err(|err| format!("{err:?}"))?; + drop(base_conn); + + let mut conn = state + .db_for_tenant(tenant_id) + .map_err(|err| format!("{err:?}"))?; enqueue_job( &mut conn, + tenant_id, JOB_INDEX_DOCUMENT_TEXT, json!({ "document_id": payload.document_id, diff --git a/backend/src/workers/thumbnails.rs b/backend/src/workers/thumbnails.rs index 81788c2..cf4b7f8 100644 --- a/backend/src/workers/thumbnails.rs +++ b/backend/src/workers/thumbnails.rs @@ -343,7 +343,7 @@ fn load_thumbnail_context( state: Arc, payload: &ThumbnailPayload, ) -> Result { - let mut conn = state.db().map_err(|err| format!("{err:?}"))?; + let mut conn = state.db_unscoped().map_err(|err| format!("{err:?}"))?; let version: DocumentVersion = document_versions::table .find(payload.document_version_id) @@ -359,12 +359,15 @@ fn load_thumbnail_context( .first(&mut conn) .map_err(|err| format!("{err:?}"))?; + let tenant_id = document.tenant_id; + let existing_assets: Vec = document_assets::table .filter(document_assets::document_version_id.eq(payload.document_version_id)) .filter(document_assets::asset_type.eq_any(vec![ THUMBNAIL_ASSET_TYPE.to_string(), PREVIEW_ASSET_TYPE.to_string(), ])) + .filter(document_assets::tenant_id.eq(tenant_id)) .load(&mut conn) .map_err(|err| format!("{err:?}"))?; @@ -377,6 +380,7 @@ fn load_thumbnail_context( THUMBNAIL_ASSET_TYPE => { existing_thumbnail_objects = document_asset_objects::table .filter(document_asset_objects::asset_id.eq(asset.id)) + .filter(document_asset_objects::tenant_id.eq(tenant_id)) .order(document_asset_objects::ordinal.asc()) .load(&mut conn) .map_err(|err| format!("{err:?}"))?; @@ -385,6 +389,7 @@ fn load_thumbnail_context( PREVIEW_ASSET_TYPE => { existing_preview_objects = document_asset_objects::table .filter(document_asset_objects::asset_id.eq(asset.id)) + .filter(document_asset_objects::tenant_id.eq(tenant_id)) .order(document_asset_objects::ordinal.asc()) .load(&mut conn) .map_err(|err| format!("{err:?}"))?; @@ -568,7 +573,10 @@ fn persist_assets_metadata( context: &ThumbnailContext, assets: &[AssetPersistence], ) -> Result<(), String> { - let mut conn = state.db().map_err(|err| format!("{err:?}"))?; + let tenant_id = context.document.tenant_id; + let mut conn = state + .db_for_tenant(tenant_id) + .map_err(|err| format!("{err:?}"))?; if let Some(existing_preview) = &context.existing_preview { diesel::delete(document_assets::table.filter(document_assets::id.eq(existing_preview.id))) @@ -607,6 +615,7 @@ fn persist_assets_metadata( "generated_at": Utc::now().to_rfc3339(), }), cardinality: Some(object_count), + tenant_id, }; diesel::insert_into(document_assets::table) @@ -626,7 +635,8 @@ fn persist_assets_metadata( diesel::delete( document_asset_objects::table - .filter(document_asset_objects::asset_id.eq(asset.asset_id)), + .filter(document_asset_objects::asset_id.eq(asset.asset_id)) + .filter(document_asset_objects::tenant_id.eq(tenant_id)), ) .execute(&mut conn) .map_err(|err| format!("{err:?}"))?; @@ -648,6 +658,7 @@ fn persist_assets_metadata( ordinal: object.ordinal, s3_key: object.s3_key.clone(), metadata: object_metadata, + tenant_id, }; diesel::insert_into(document_asset_objects::table) @@ -666,11 +677,17 @@ fn persist_document_page_count( document_version_id: Uuid, page_count: u32, ) -> Result<(), String> { - let mut conn = state.db().map_err(|err| format!("{err:?}"))?; + let mut conn = state.db_unscoped().map_err(|err| format!("{err:?}"))?; + let tenant_id: Uuid = documents::table + .find(document_id) + .select(documents::tenant_id) + .first(&mut conn) + .map_err(|err| format!("{err:?}"))?; let existing_metadata: Value = document_versions::table .filter(document_versions::id.eq(document_version_id)) .filter(document_versions::document_id.eq(document_id)) + .filter(document_versions::tenant_id.eq(tenant_id)) .select(document_versions::metadata) .first(&mut conn) .map_err(|err| format!("{err:?}"))?; @@ -690,7 +707,8 @@ fn persist_document_page_count( diesel::update( document_versions::table .filter(document_versions::id.eq(document_version_id)) - .filter(document_versions::document_id.eq(document_id)), + .filter(document_versions::document_id.eq(document_id)) + .filter(document_versions::tenant_id.eq(tenant_id)), ) .set(document_versions::metadata.eq(updated)) .execute(&mut conn) diff --git a/backend/tests/auth_flow.rs b/backend/tests/auth_flow.rs index ef5f127..973a763 100644 --- a/backend/tests/auth_flow.rs +++ b/backend/tests/auth_flow.rs @@ -8,7 +8,6 @@ use serde::Deserialize; #[derive(Deserialize)] struct AuthenticatedUser { username: String, - role: String, } #[tokio::test] @@ -27,7 +26,6 @@ async fn login_and_me_roundtrip() -> Result<()> { let user: AuthenticatedUser = serde_json::from_slice(&body)?; assert_eq!(user.username, "alice"); - assert_eq!(user.role, "admin"); app.cleanup().await?; Ok(()) diff --git a/backend/tests/common/mod.rs b/backend/tests/common/mod.rs index 14fef39..0a47dc6 100644 --- a/backend/tests/common/mod.rs +++ b/backend/tests/common/mod.rs @@ -11,7 +11,7 @@ use axum::Router; use backend::auth::jwt::JwtService; use backend::config::AppConfig; use backend::db::{self, PgPool}; -use backend::models::{Job, NewUser}; +use backend::models::{Job, NewUser, NewUserMembership}; use backend::routes; use backend::state::AppState; use backend::storage::ObjectStorage; @@ -138,6 +138,7 @@ impl TestApp { s3_bucket: "test-bucket".to_string(), quickwit_endpoint: None, quickwit_index: None, + default_tenant_slug: "admin".to_string(), }; let pool = db::init_pool_with_size(&config.database_url, config.database_max_pool_size)?; @@ -178,18 +179,35 @@ impl TestApp { let username = username.to_string(); let password = password.to_string(); let role = role.to_string(); + let tenant_id = self + .state + .tenants + .tenant_id_for_slug(&self.state.config.default_tenant_slug) + .context("default tenant not found")?; self.with_conn(move |conn| { let password_hash = hash_password(&password)?; let user = NewUser { id: Uuid::new_v4(), username, password_hash, - role, + tenant_id, }; diesel::insert_into(backend::schema::users::table) .values(&user) .execute(conn) .context("failed to insert user")?; + + let membership = NewUserMembership { + id: Uuid::new_v4(), + user_id: user.id, + tenant_id, + role, + }; + + diesel::insert_into(backend::schema::user_memberships::table) + .values(&membership) + .execute(conn) + .context("failed to insert user membership")?; Ok(user.id) }) .await