MediaRouter / app /social /database.py
basyx's picture
Upload 340 files
3493993 verified
Raw
History Blame Contribute Delete
13.3 kB
from __future__ import annotations
import contextvars
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,
)
from app.analytics import models as analytics_models # noqa: F401
from app.core.config import Settings
from app.core.database_url import normalize_async_database_url
from app.social.models import SocialBase
_ANALYTICS_MODELS_REGISTERED = analytics_models
REQUIRED_SOCIAL_TABLES = frozenset(
{
"social_accounts",
"social_account_tokens",
"social_account_capabilities",
"media_variants",
"social_media_assets",
"social_campaigns",
"social_posts",
"social_post_targets",
"social_post_media",
"social_schedules",
"social_jobs",
"social_job_attempts",
"oauth_states",
"social_webhook_events",
"social_post_metrics",
"social_audit_events",
"social_publishing_batches",
"social_publishing_batch_items",
"analytics_sync_runs",
"analytics_metric_snapshots",
"analytics_post_metrics",
"analytics_platform_metrics",
}
)
REQUIRED_SOCIAL_COLUMNS: dict[str, frozenset[str]] = {
"social_media_assets": frozenset({"canonical_asset_id"}),
"social_posts": frozenset(
{"project_id", "canonical_caption", "canonical_hashtags", "revision"}
),
"social_schedules": frozenset({"revision"}),
"social_post_targets": frozenset({"cancellation_requested_at"}),
"social_jobs": frozenset({"provider_state_encrypted", "cancellation_requested_at"}),
"oauth_states": frozenset({"requested_account_type", "requested_scopes"}),
"analytics_sync_runs": frozenset(
{"workspace_id", "idempotency_key", "status", "date_from", "date_to"}
),
}
_trusted_worker_context: contextvars.ContextVar[bool] = contextvars.ContextVar(
"trusted_social_worker_context", default=False
)
class SocialDatabase:
"""Social persistence with migration-only production schema changes."""
def __init__(self, settings: Settings) -> None:
self.settings = settings
self.database_url = normalize_async_database_url(settings.resolved_social_database_url)
self.engine: AsyncEngine = create_async_engine(self.database_url, pool_pre_ping=True)
if self.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
)
self.worker_database_url = normalize_async_database_url(settings.social_worker_database_url)
self.worker_engine: AsyncEngine | None = None
self.worker_session_factory: async_sessionmaker[AsyncSession] | None = None
if self.worker_database_url:
self.worker_engine = create_async_engine(self.worker_database_url, pool_pre_ping=True)
if self.worker_database_url.startswith("sqlite"):
event.listen(self.worker_engine.sync_engine, "connect", self._configure_sqlite)
self.worker_session_factory = async_sessionmaker(
self.worker_engine, expire_on_commit=False, class_=AsyncSession
)
@staticmethod
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 busy_timeout=5000")
cursor.close()
async def initialize(self) -> None:
if self.settings.social_auto_migrate:
async with self.engine.begin() as connection:
await connection.run_sync(SocialBase.metadata.create_all)
@property
def is_postgres(self) -> bool:
return self.database_url.startswith(("postgresql", "postgres"))
async def verify_execution_boundaries(self) -> None:
"""Fail closed when a Postgres deployment cannot enforce RLS safely."""
if not self.is_postgres or not self.settings.social_enforce_rls:
return
tenant_role = self.settings.social_tenant_database_role.strip()
if not tenant_role:
raise RuntimeError(
"SOCIAL_TENANT_DATABASE_ROLE is required when SOCIAL_ENFORCE_RLS is enabled."
)
tenant = await self._role_attributes(self.engine)
if tenant["role"] != tenant_role:
raise RuntimeError(
"SOCIAL_DATABASE_URL is not connected as SOCIAL_TENANT_DATABASE_ROLE."
)
if tenant["bypass_rls"] or tenant["superuser"]:
raise RuntimeError(
"SOCIAL_DATABASE_URL must use a non-privileged tenant role, never a service role."
)
# Startup also adopts historic API-key tenant rows through this
# boundary, so every enabled PostgreSQL social deployment needs it,
# even if the scheduler is temporarily disabled.
if not self.worker_engine:
raise RuntimeError(
"SOCIAL_WORKER_DATABASE_URL is required for PostgreSQL social access."
)
worker_role = self.settings.social_worker_database_role.strip()
if not worker_role:
raise RuntimeError("SOCIAL_WORKER_DATABASE_ROLE is required for trusted worker access.")
worker = await self._role_attributes(self.worker_engine)
if worker["role"] != worker_role or not worker["bypass_rls"]:
raise RuntimeError(
"SOCIAL_WORKER_DATABASE_URL must use the configured BYPASSRLS worker role."
)
@staticmethod
async def _role_attributes(engine: AsyncEngine) -> dict[str, object]:
async with 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:
raise RuntimeError("Unable to verify the active PostgreSQL database role.")
return dict(row)
async def schema_ready(self) -> bool:
"""Check the complete Phase 1 schema without changing the database."""
async with self.engine.connect() as connection:
tables, columns = await connection.run_sync(self._schema_snapshot)
return REQUIRED_SOCIAL_TABLES.issubset(tables) and all(
required.issubset(columns.get(table, set()))
for table, required in REQUIRED_SOCIAL_COLUMNS.items()
)
async def missing_tables(self) -> list[str]:
"""Return absent required tables for an actionable startup warning."""
async with self.engine.connect() as connection:
tables, columns = await connection.run_sync(self._schema_snapshot)
missing = list(REQUIRED_SOCIAL_TABLES - tables)
for table, required in REQUIRED_SOCIAL_COLUMNS.items():
missing.extend(f"{table}.{column}" for column in required - columns.get(table, set()))
return sorted(missing)
async def adopt_legacy_workspace(
self, *, legacy_workspace_id: str, workspace_id: str, user_id: str
) -> int:
"""Move historic API-key tenant rows to the authoritative workspace."""
if legacy_workspace_id == workspace_id:
raise RuntimeError("Legacy and authoritative workspace IDs must differ.")
workspace_tables = (
"social_accounts",
"media_variants",
"social_media_assets",
"social_campaigns",
"social_posts",
"social_jobs",
"social_webhook_events",
"social_audit_events",
"oauth_states",
"social_publishing_batches",
)
changed = 0
async with self.worker_session() as session:
for table in workspace_tables:
result = await session.execute(
text(
f"update {table} set workspace_id = :workspace_id "
"where workspace_id = :legacy_workspace_id"
),
{"workspace_id": workspace_id, "legacy_workspace_id": legacy_workspace_id},
)
changed += max(0, int(result.rowcount or 0))
for table, column in (("social_posts", "created_by"), ("oauth_states", "user_id")):
result = await session.execute(
text(
f"update {table} set {column} = :user_id "
f"where {column} = :legacy_workspace_id"
),
{"user_id": user_id, "legacy_workspace_id": legacy_workspace_id},
)
changed += max(0, int(result.rowcount or 0))
await session.commit()
return changed
@staticmethod
def _schema_snapshot(connection: object) -> tuple[set[str], dict[str, set[str]]]:
inspector = inspect(connection)
tables = set(inspector.get_table_names())
columns = {
table: {column["name"] for column in inspector.get_columns(table)}
for table in REQUIRED_SOCIAL_COLUMNS
if table in tables
}
return tables, columns
async def close(self) -> None:
await self.engine.dispose()
if self.worker_engine is not None:
await self.worker_engine.dispose()
@asynccontextmanager
async def session(self, workspace_id: str | None = None) -> AsyncIterator[AsyncSession]:
"""Open a tenant session; no-context access is OAuth-state compatibility only."""
if workspace_id is None:
async with self.oauth_session() as session:
yield session
return
context = (
self.worker_tenant_session(workspace_id)
if _trusted_worker_context.get()
else self.tenant_session(workspace_id)
)
async with context as session:
yield session
@asynccontextmanager
async def tenant_session(self, workspace_id: str) -> AsyncIterator[AsyncSession]:
if not workspace_id:
raise RuntimeError("A tenant session requires an authoritative workspace ID.")
async with self.session_factory() as session:
if self.is_postgres:
# RLS policies read this transaction-local tenant identity.
await session.execute(
text("select set_config('app.workspace_id', :workspace_id, true)"),
{"workspace_id": workspace_id},
)
yield session
@asynccontextmanager
async def oauth_session(self) -> AsyncIterator[AsyncSession]:
"""The sole non-tenant API session, for random single-use OAuth state."""
async with self.session_factory() as session:
yield session
@asynccontextmanager
async def worker_session(self) -> AsyncIterator[AsyncSession]:
"""Backend-only cross-workspace session for scheduler, jobs, and Vault."""
if self.is_postgres:
if self.worker_session_factory is None:
raise RuntimeError("Trusted worker database access is not configured.")
async with self.worker_session_factory() as session:
yield session
return
# SQLite has no RLS. It remains supported for local/unit-test use only.
async with self.session_factory() as session:
yield session
@asynccontextmanager
async def worker_tenant_session(self, workspace_id: str) -> AsyncIterator[AsyncSession]:
"""Trusted worker session annotated with the job's tenant for auditability."""
if not workspace_id:
raise RuntimeError("A worker tenant session requires a workspace ID.")
if not self.is_postgres:
async with self.tenant_session(workspace_id) as session:
yield session
return
if self.worker_session_factory is None:
raise RuntimeError("Trusted worker database access is not configured.")
async with self.worker_session_factory() as session:
await session.execute(
text("select set_config('app.workspace_id', :workspace_id, true)"),
{"workspace_id": workspace_id},
)
yield session
@asynccontextmanager
async def worker_boundary(self) -> AsyncIterator[None]:
"""Mark a scheduler/publisher call tree as trusted worker execution."""
if self.is_postgres and self.worker_session_factory is None:
raise RuntimeError("Trusted worker database access is not configured.")
token = _trusted_worker_context.set(True)
try:
yield
finally:
_trusted_worker_context.reset(token)