Spaces:
Running
Running
| from __future__ import annotations | |
| from collections.abc import AsyncIterator | |
| from contextlib import asynccontextmanager | |
| from sqlalchemy import event, inspect, text | |
| from sqlalchemy.ext.asyncio import ( | |
| AsyncEngine, | |
| AsyncSession, | |
| async_sessionmaker, | |
| create_async_engine, | |
| ) | |
| # Importing the generation models here guarantees that local SQLite metadata | |
| # and production schema checks see the full security-owned domain even when a | |
| # caller constructs SecurityDatabase without first building the application | |
| # container. | |
| from app.core.database_url import normalize_async_database_url | |
| from app.security.models import Base | |
| REQUIRED_SECURITY_TABLES = frozenset( | |
| { | |
| "api_keys", | |
| "audit_logs", | |
| "rate_limits", | |
| "users", | |
| "workspaces", | |
| "workspace_memberships", | |
| "api_key_principals", | |
| "media_assets", | |
| "media_asset_variants", | |
| "audit_events", | |
| "projects", | |
| "project_generation_jobs", | |
| "generation_requests", | |
| "generation_jobs", | |
| "generation_job_attempts", | |
| "project_editor_states", | |
| "project_render_jobs", | |
| "copilot_runs", | |
| "marketplace_templates", | |
| "marketplace_template_versions", | |
| "marketplace_template_applications", | |
| "brand_kits", | |
| "brand_kit_versions", | |
| } | |
| ) | |
| REQUIRED_POSTGRES_SECURITY_INDEXES = frozenset( | |
| { | |
| "uq_generation_job_provider_external", | |
| "ix_projects_workspace", | |
| "ix_projects_workspace_status", | |
| "ix_projects_workspace_updated", | |
| "ix_projects_created_by", | |
| "ix_media_assets_workspace_project_created", | |
| "ix_project_generation_jobs_workspace_project_created", | |
| "uq_project_generation_job", | |
| "uq_project_editor_state_project", | |
| "ix_project_editor_states_workspace_project", | |
| "uq_project_render_idempotency", | |
| "ix_project_render_jobs_workspace_status", | |
| "ix_project_render_jobs_project_created", | |
| "ix_project_render_jobs_dispatch", | |
| "ix_generation_requests_workspace_project_created", | |
| "ix_generation_requests_workspace_surface_created", | |
| "uq_copilot_run_workspace_idempotency", | |
| "ix_copilot_runs_workspace_created", | |
| "ix_copilot_runs_workspace_status", | |
| "ix_copilot_runs_project_created", | |
| "uq_marketplace_template_workspace_slug", | |
| "ix_marketplace_templates_workspace_status", | |
| "ix_marketplace_templates_discovery", | |
| "uq_marketplace_template_version", | |
| "ix_marketplace_template_versions_template", | |
| "uq_marketplace_template_application_idempotency", | |
| "ix_marketplace_template_applications_project", | |
| "ix_marketplace_template_applications_template", | |
| "ix_brand_kits_workspace_updated", | |
| "ix_brand_kits_workspace_default", | |
| "uq_brand_kits_one_default", | |
| "uq_brand_kit_version", | |
| "ix_brand_versions_kit_created", | |
| } | |
| ) | |
| class SecurityDatabase: | |
| """Owns the authentication database engine and short-lived async sessions.""" | |
| def __init__(self, database_url: str, *, auto_migrate: bool = False) -> None: | |
| database_url = normalize_async_database_url(database_url) | |
| self.database_url = database_url | |
| self.auto_migrate = auto_migrate or database_url.startswith("sqlite") | |
| self.engine: AsyncEngine = create_async_engine( | |
| database_url, | |
| pool_pre_ping=True, | |
| ) | |
| if database_url.startswith("sqlite"): | |
| event.listen(self.engine.sync_engine, "connect", self._configure_sqlite) | |
| self.session_factory = async_sessionmaker( | |
| self.engine, expire_on_commit=False, class_=AsyncSession | |
| ) | |
| def _configure_sqlite(dbapi_connection: object, _record: object) -> None: | |
| cursor = dbapi_connection.cursor() # type: ignore[attr-defined] | |
| cursor.execute("PRAGMA foreign_keys=ON") | |
| cursor.execute("PRAGMA journal_mode=WAL") | |
| cursor.execute("PRAGMA busy_timeout=5000") | |
| cursor.close() | |
| async def initialize(self) -> None: | |
| if not self.auto_migrate: | |
| return | |
| __import__("app.generation.models") | |
| __import__("app.projects.models") | |
| __import__("app.copilot.models") | |
| __import__("app.templates.marketplace_models") | |
| __import__("app.brand.models") | |
| async with self.engine.begin() as connection: | |
| await connection.run_sync(Base.metadata.create_all) | |
| async def close(self) -> None: | |
| await self.engine.dispose() | |
| def is_postgres(self) -> bool: | |
| return self.database_url.startswith(("postgresql", "postgres")) | |
| async def verify_execution_boundary(self, *, expected_role: str, enforce_rls: bool) -> None: | |
| """Ensure tenancy administration never uses the public tenant role.""" | |
| if not self.is_postgres or not enforce_rls: | |
| return | |
| if not expected_role.strip(): | |
| raise RuntimeError( | |
| "SECURITY_DATABASE_ROLE is required when SECURITY_ENFORCE_RLS is enabled." | |
| ) | |
| async with self.engine.connect() as connection: | |
| row = ( | |
| ( | |
| await connection.execute( | |
| text( | |
| "select current_user as role, r.rolbypassrls as bypass_rls, r.rolsuper as superuser " | |
| "from pg_roles r where r.rolname = current_user" | |
| ) | |
| ) | |
| ) | |
| .mappings() | |
| .one_or_none() | |
| ) | |
| if row is None or row["role"] != expected_role.strip(): | |
| raise RuntimeError("DATABASE_URL is not connected as SECURITY_DATABASE_ROLE.") | |
| if not row["bypass_rls"] and not row["superuser"]: | |
| raise RuntimeError( | |
| "SECURITY_DATABASE_ROLE must be backend-only and able to administer forced-RLS tenancy tables." | |
| ) | |
| async def schema_ready(self) -> bool: | |
| return not await self.missing_schema_objects() | |
| async def missing_schema_objects(self) -> list[str]: | |
| """Return missing production schema requirements, including safety indexes.""" | |
| async with self.engine.connect() as connection: | |
| tables = await connection.run_sync(lambda sync: set(inspect(sync).get_table_names())) | |
| if not self.is_postgres: | |
| indexes: set[str] = set() | |
| else: | |
| rows = await connection.execute( | |
| text("select indexname from pg_indexes " "where schemaname = current_schema()") | |
| ) | |
| indexes = set(rows.scalars().all()) | |
| missing = sorted(REQUIRED_SECURITY_TABLES - tables) | |
| if self.is_postgres: | |
| missing.extend( | |
| f"index:{name}" for name in sorted(REQUIRED_POSTGRES_SECURITY_INDEXES - indexes) | |
| ) | |
| return missing | |
| async def missing_tables(self) -> list[str]: | |
| async with self.engine.connect() as connection: | |
| tables = await connection.run_sync(lambda sync: set(inspect(sync).get_table_names())) | |
| return sorted(REQUIRED_SECURITY_TABLES - tables) | |
| async def session(self) -> AsyncIterator[AsyncSession]: | |
| async with self.session_factory() as session: | |
| yield session | |
| async def tenant_session( | |
| self, *, workspace_id: str, user_id: str | |
| ) -> AsyncIterator[AsyncSession]: | |
| """Explicit tenant context for RLS verification and scoped services.""" | |
| if not workspace_id or not user_id: | |
| raise RuntimeError("A security tenant session requires workspace and user IDs.") | |
| async with self.session_factory() as session: | |
| if not self.engine.url.drivername.startswith("sqlite"): | |
| await session.execute( | |
| text("select set_config('app.workspace_id', :workspace_id, true)"), | |
| {"workspace_id": workspace_id}, | |
| ) | |
| await session.execute( | |
| text("select set_config('app.user_id', :user_id, true)"), | |
| {"user_id": user_id}, | |
| ) | |
| yield session | |