from __future__ import annotations import json import sqlite3 from collections.abc import Iterator from contextlib import contextmanager from datetime import datetime from typing import Any from uuid import UUID, uuid4 from src.domain.contracts import ( ArticleSummary, PublishingRules, PublishingStatus, ArticleWorkflowStatus, WorkflowEventSummary, Role, ScriptConfigVersionStatus, TargetSiteConfig, UserSummary, ) from src.infrastructure.schema import setup_database JsonObject = dict[str, Any] def open_backend_repository(dsn: str) -> "BackendRepository": return BackendRepository(dsn) class BackendRepository: def __init__(self, dsn: str) -> None: self.dsn = dsn self.dialect = "sqlite" if dsn.startswith("sqlite:///") else "postgres" self.users = UsersRepository(self) self.target_sites = TargetSitesRepository(self) self.articles = ArticlesRepository(self) self.script_config_versions = ScriptConfigVersionsRepository(self) self.script_config_version_events = ScriptConfigVersionAuditEventsRepository(self) self.schema = SchemaRepository(self) def setup(self) -> None: setup_database(self.dsn) @contextmanager def connection(self) -> Iterator[Any]: connection = self._connect() try: with connection: yield connection finally: connection.close() def _connect(self) -> Any: if self.dialect == "sqlite": sqlite_path = self.dsn.removeprefix("sqlite:///") connection = sqlite3.connect(sqlite_path) connection.row_factory = sqlite3.Row return connection import psycopg from psycopg.rows import dict_row return psycopg.connect(self.dsn, row_factory=dict_row) def placeholder(self) -> str: if self.dialect == "sqlite": return "?" return "%s" def json_cast(self) -> str: if self.dialect == "sqlite": return "" return "::jsonb" class UsersRepository: def __init__(self, repository: BackendRepository) -> None: self._repository = repository def upsert( self, *, user_id: UUID, email: str, display_name: str, role: Role, created_at: datetime, updated_at: datetime, ) -> UserSummary: placeholder = self._repository.placeholder() sql = f""" INSERT INTO users ( id, email, display_name, role, created_at, updated_at ) VALUES ( {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder} ) ON CONFLICT (email) DO UPDATE SET display_name = excluded.display_name, role = excluded.role, updated_at = excluded.updated_at """ params = ( str(user_id), email, display_name, role.value, _datetime_value(created_at), _datetime_value(updated_at), ) with self._repository.connection() as connection: connection.execute(sql, params) return self.get_by_email(email) def get_by_email(self, email: str) -> UserSummary: placeholder = self._repository.placeholder() with self._repository.connection() as connection: row = connection.execute( f""" SELECT id, display_name, role FROM users WHERE email = {placeholder} """, (email,), ).fetchone() if row is None: raise LookupError(f"User not found: {email}") return _user_from_row(row) def list(self) -> list[UserSummary]: with self._repository.connection() as connection: rows = connection.execute( """ SELECT id, display_name, role FROM users ORDER BY email """ ).fetchall() return [_user_from_row(row) for row in rows] def count_by_email(self, email: str) -> int: placeholder = self._repository.placeholder() with self._repository.connection() as connection: row = connection.execute( f"SELECT COUNT(*) AS count FROM users WHERE email = {placeholder}", (email,), ).fetchone() return int(_row_value(row, "count")) class TargetSitesRepository: def __init__(self, repository: BackendRepository) -> None: self._repository = repository def upsert( self, *, site_id: UUID, name: str, slug: str, publishing_type: str, default_language: str, brand_voice: str, audience: str, seo_rules: JsonObject, visual_rules: JsonObject, source_rules: JsonObject, publishing_rules: PublishingRules, active_script_config_version_id: UUID | None, created_at: datetime, updated_at: datetime, ) -> TargetSiteConfig: placeholder = self._repository.placeholder() json_cast = self._repository.json_cast() sql = f""" INSERT INTO target_sites ( id, name, slug, publishing_type, default_language, brand_voice, audience, seo_rules, visual_rules, source_rules, publishing_rules, active_script_config_version_id, created_at, updated_at ) VALUES ( {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}{json_cast}, {placeholder}{json_cast}, {placeholder}{json_cast}, {placeholder}{json_cast}, {placeholder}, {placeholder}, {placeholder} ) ON CONFLICT (slug) DO UPDATE SET name = excluded.name, publishing_type = excluded.publishing_type, default_language = excluded.default_language, brand_voice = excluded.brand_voice, audience = excluded.audience, seo_rules = excluded.seo_rules, visual_rules = excluded.visual_rules, source_rules = excluded.source_rules, publishing_rules = excluded.publishing_rules, active_script_config_version_id = excluded.active_script_config_version_id, updated_at = excluded.updated_at """ params = ( str(site_id), name, slug, publishing_type, default_language, brand_voice, audience, _json_value(seo_rules), _json_value(visual_rules), _json_value(source_rules), _json_value(publishing_rules.model_dump(mode="json")), _uuid_value(active_script_config_version_id), _datetime_value(created_at), _datetime_value(updated_at), ) with self._repository.connection() as connection: connection.execute(sql, params) return self.get_by_slug(slug) def get_by_slug(self, slug: str) -> TargetSiteConfig: placeholder = self._repository.placeholder() with self._repository.connection() as connection: row = connection.execute( f""" SELECT id, name, slug, publishing_type, default_language, brand_voice, audience, seo_rules, visual_rules, source_rules, publishing_rules, active_script_config_version_id, created_at, updated_at FROM target_sites WHERE slug = {placeholder} """, (slug,), ).fetchone() if row is None: raise LookupError(f"Target site not found: {slug}") return _target_site_from_row(row) def get_by_id(self, site_id: UUID) -> TargetSiteConfig: placeholder = self._repository.placeholder() with self._repository.connection() as connection: row = connection.execute( f""" SELECT id, name, slug, publishing_type, default_language, brand_voice, audience, seo_rules, visual_rules, source_rules, publishing_rules, active_script_config_version_id, created_at, updated_at FROM target_sites WHERE id = {placeholder} """, (str(site_id),), ).fetchone() if row is None: raise LookupError(f"Target site not found: {site_id}") return _target_site_from_row(row) def update( self, *, site_id: UUID, name: str, slug: str, publishing_type: str, default_language: str, brand_voice: str, audience: str, seo_rules: JsonObject, visual_rules: JsonObject, source_rules: JsonObject, publishing_rules: PublishingRules, active_script_config_version_id: UUID | None, updated_at: datetime, ) -> TargetSiteConfig: placeholder = self._repository.placeholder() json_cast = self._repository.json_cast() with self._repository.connection() as connection: connection.execute( f""" UPDATE target_sites SET name = {placeholder}, slug = {placeholder}, publishing_type = {placeholder}, default_language = {placeholder}, brand_voice = {placeholder}, audience = {placeholder}, seo_rules = {placeholder}{json_cast}, visual_rules = {placeholder}{json_cast}, source_rules = {placeholder}{json_cast}, publishing_rules = {placeholder}{json_cast}, active_script_config_version_id = {placeholder}, updated_at = {placeholder} WHERE id = {placeholder} """, ( name, slug, publishing_type, default_language, brand_voice, audience, _json_value(seo_rules), _json_value(visual_rules), _json_value(source_rules), _json_value(publishing_rules.model_dump(mode="json")), _uuid_value(active_script_config_version_id), _datetime_value(updated_at), str(site_id), ), ) return self.get_by_id(site_id) def list(self) -> list[TargetSiteConfig]: with self._repository.connection() as connection: rows = connection.execute( """ SELECT id, name, slug, publishing_type, default_language, brand_voice, audience, seo_rules, visual_rules, source_rules, publishing_rules, active_script_config_version_id, created_at, updated_at FROM target_sites ORDER BY slug """ ).fetchall() return [_target_site_from_row(row) for row in rows] class ArticlesRepository: def __init__(self, repository: BackendRepository) -> None: self._repository = repository def create( self, *, target_site_id: UUID, status: ArticleWorkflowStatus, publishing_status: PublishingStatus, brief_description: str, working_title: str | None, language: str, content_type: str, primary_keyword: str | None, assigned_editor_id: UUID | None, created_at: datetime, updated_at: datetime, ) -> ArticleSummary: article_id = uuid4() placeholder = self._repository.placeholder() sql = f""" INSERT INTO articles ( id, target_site_id, status, publishing_status, brief_description, working_title, language, content_type, primary_keyword, assigned_editor_id, created_at, updated_at ) VALUES ( {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder} ) """ with self._repository.connection() as connection: connection.execute( sql, ( str(article_id), str(target_site_id), status.value, publishing_status.value, brief_description, working_title, language, content_type, primary_keyword, _uuid_value(assigned_editor_id), _datetime_value(created_at), _datetime_value(updated_at), ), ) return self.get(article_id) def list(self) -> list[ArticleSummary]: with self._repository.connection() as connection: rows = connection.execute( """ SELECT id, target_site_id, status, publishing_status, brief_description, working_title, language, content_type, primary_keyword, assigned_editor_id, created_at, updated_at FROM articles ORDER BY updated_at DESC """ ).fetchall() return [_article_summary_from_row(row) for row in rows] def get(self, article_id: UUID) -> ArticleSummary: placeholder = self._repository.placeholder() with self._repository.connection() as connection: row = connection.execute( f""" SELECT id, target_site_id, status, publishing_status, brief_description, working_title, language, content_type, primary_keyword, assigned_editor_id, created_at, updated_at FROM articles WHERE id = {placeholder} """, (str(article_id),), ).fetchone() if row is None: raise LookupError(f"Article not found: {article_id}") return _article_summary_from_row(row) def create_workflow_event( self, *, article_id: UUID, event_type: str, from_status: ArticleWorkflowStatus | None, to_status: ArticleWorkflowStatus | None, actor_user_id: UUID | None, payload: dict[str, Any], created_at: datetime, ) -> None: placeholder = self._repository.placeholder() json_cast = self._repository.json_cast() sql = f""" INSERT INTO workflow_events ( id, article_id, event_type, from_status, to_status, actor_user_id, payload, created_at ) VALUES ( {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}{json_cast}, {placeholder} ) """ with self._repository.connection() as connection: connection.execute( sql, ( str(uuid4()), str(article_id), event_type, _article_status_value(from_status), _article_status_value(to_status), _uuid_value(actor_user_id), _json_value(payload), _datetime_value(created_at), ), ) def list_workflow_events(self, article_id: UUID) -> list[WorkflowEventSummary]: placeholder = self._repository.placeholder() with self._repository.connection() as connection: rows = connection.execute( f""" SELECT id, article_id, event_type, from_status, to_status, actor_user_id, payload, created_at FROM workflow_events WHERE article_id = {placeholder} ORDER BY created_at """, (str(article_id),), ).fetchall() return [_workflow_event_from_row(row) for row in rows] class ScriptConfigVersionsRepository: def __init__(self, repository: BackendRepository) -> None: self._repository = repository def upsert( self, *, version_id: UUID, target_site_id: UUID, version: int, status: ScriptConfigVersionStatus, created_by: UUID, created_at: datetime, updated_at: datetime, diff: JsonObject, rollback_target_version_id: UUID | None, activated_at: datetime | None, publishing_yaml: str, publishing_yaml_hash: str, transform_script: str, transform_script_hash: str, ) -> dict[str, Any]: placeholder = self._repository.placeholder() json_cast = self._repository.json_cast() sql = f""" INSERT INTO script_config_versions ( id, target_site_id, version, status, created_by, created_at, updated_at, diff, rollback_target_version_id, activated_at, publishing_yaml, publishing_yaml_hash, transform_script, transform_script_hash ) VALUES ( {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}{json_cast}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder} ) ON CONFLICT (target_site_id, version) DO UPDATE SET status = excluded.status, created_by = excluded.created_by, updated_at = excluded.updated_at, diff = excluded.diff, rollback_target_version_id = excluded.rollback_target_version_id, activated_at = excluded.activated_at, publishing_yaml = excluded.publishing_yaml, publishing_yaml_hash = excluded.publishing_yaml_hash, transform_script = excluded.transform_script, transform_script_hash = excluded.transform_script_hash """ params = ( str(version_id), str(target_site_id), version, status.value, str(created_by), _datetime_value(created_at), _datetime_value(updated_at), _json_value(diff), _uuid_value(rollback_target_version_id), _datetime_value(activated_at), publishing_yaml, publishing_yaml_hash, transform_script, transform_script_hash, ) with self._repository.connection() as connection: connection.execute(sql, params) return self.get_by_site_and_version(target_site_id=target_site_id, version=version) def get_by_site_and_version( self, *, target_site_id: UUID, version: int ) -> dict[str, Any]: placeholder = self._repository.placeholder() with self._repository.connection() as connection: row = connection.execute( f""" SELECT * FROM script_config_versions WHERE target_site_id = {placeholder} AND version = {placeholder} """, (str(target_site_id), version), ).fetchone() if row is None: raise LookupError( f"Script config version not found: {target_site_id} v{version}" ) return _plain_row(row) def get_by_id(self, version_id: UUID) -> dict[str, Any]: placeholder = self._repository.placeholder() with self._repository.connection() as connection: row = connection.execute( f""" SELECT * FROM script_config_versions WHERE id = {placeholder} """, (str(version_id),), ).fetchone() if row is None: raise LookupError(f"Script config version not found: {version_id}") return _plain_row(row) def activate( self, *, target_site_id: UUID, version_id: UUID, activated_at: datetime, rollback_target_version_id: UUID | None = None, ) -> dict[str, Any]: placeholder = self._repository.placeholder() with self._repository.connection() as connection: existing = connection.execute( f""" SELECT id FROM script_config_versions WHERE id = {placeholder} AND target_site_id = {placeholder} """, (str(version_id), str(target_site_id)), ).fetchone() if existing is None: raise LookupError( f"Script config version not found: {target_site_id} {version_id}" ) connection.execute( f""" UPDATE script_config_versions SET status = {placeholder}, updated_at = {placeholder} WHERE target_site_id = {placeholder} AND id <> {placeholder} """, ( ScriptConfigVersionStatus.DEPRECATED.value, _datetime_value(activated_at), str(target_site_id), str(version_id), ), ) connection.execute( f""" UPDATE script_config_versions SET status = {placeholder}, activated_at = {placeholder}, updated_at = {placeholder}, rollback_target_version_id = COALESCE({placeholder}, rollback_target_version_id) WHERE id = {placeholder} """, ( ScriptConfigVersionStatus.ACTIVE.value, _datetime_value(activated_at), _datetime_value(activated_at), _uuid_value(rollback_target_version_id), str(version_id), ), ) connection.execute( f""" UPDATE target_sites SET active_script_config_version_id = {placeholder}, updated_at = {placeholder} WHERE id = {placeholder} """, (str(version_id), _datetime_value(activated_at), str(target_site_id)), ) return self.get_by_id(version_id) def list_for_site(self, target_site_id: UUID) -> list[dict[str, Any]]: placeholder = self._repository.placeholder() with self._repository.connection() as connection: rows = connection.execute( f""" SELECT * FROM script_config_versions WHERE target_site_id = {placeholder} ORDER BY version """, (str(target_site_id),), ).fetchall() return [_plain_row(row) for row in rows] class ScriptConfigVersionAuditEventsRepository: def __init__(self, repository: BackendRepository) -> None: self._repository = repository def create( self, *, target_site_id: UUID, version_id: UUID, event_type: str, actor_user_id: UUID | None, payload: JsonObject, created_at: datetime, ) -> dict[str, Any]: placeholder = self._repository.placeholder() sql = f""" INSERT INTO script_config_version_events ( id, target_site_id, version_id, event_type, actor_user_id, payload, created_at ) VALUES ( {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder}, {placeholder} ) """ params = ( str(uuid4()), str(target_site_id), str(version_id), event_type, _uuid_value(actor_user_id), _json_value(payload), _datetime_value(created_at), ) with self._repository.connection() as connection: connection.execute(sql, params) return self.get(event_type=event_type, target_site_id=target_site_id, version_id=version_id) def list_for_site(self, target_site_id: UUID) -> list[dict[str, Any]]: placeholder = self._repository.placeholder() with self._repository.connection() as connection: rows = connection.execute( f""" SELECT id, target_site_id, version_id, event_type, actor_user_id, payload, created_at FROM script_config_version_events WHERE target_site_id = {placeholder} ORDER BY created_at """, (str(target_site_id),), ).fetchall() return [_script_config_version_event_from_row(row) for row in rows] def get( self, *, target_site_id: UUID, version_id: UUID, event_type: str, ) -> dict[str, Any]: placeholder = self._repository.placeholder() with self._repository.connection() as connection: row = connection.execute( f""" SELECT id, target_site_id, version_id, event_type, actor_user_id, payload, created_at FROM script_config_version_events WHERE target_site_id = {placeholder} AND version_id = {placeholder} AND event_type = {placeholder} """, (str(target_site_id), str(version_id), event_type), ).fetchone() if row is None: raise LookupError( f"Script config version event not found: {target_site_id} {version_id} {event_type}" ) return _script_config_version_event_from_row(row) class SchemaRepository: def __init__(self, repository: BackendRepository) -> None: self._repository = repository def list_tables(self) -> set[str]: with self._repository.connection() as connection: if self._repository.dialect == "sqlite": rows = connection.execute( """ SELECT name FROM sqlite_master WHERE type = 'table' AND name NOT LIKE 'sqlite_%' """ ).fetchall() else: rows = connection.execute( """ SELECT table_name AS name FROM information_schema.tables WHERE table_schema = 'public' AND table_type = 'BASE TABLE' """ ).fetchall() return {_row_value(row, "name") for row in rows} def list_columns(self, table_name: str) -> set[str]: with self._repository.connection() as connection: if self._repository.dialect == "sqlite": rows = connection.execute(f"PRAGMA table_info({table_name})").fetchall() return {_row_value(row, "name") for row in rows} rows = connection.execute( """ SELECT column_name FROM information_schema.columns WHERE table_schema = 'public' AND table_name = %s """, (table_name,), ).fetchall() return {_row_value(row, "column_name") for row in rows} def list_indexes(self, table_name: str) -> set[str]: with self._repository.connection() as connection: if self._repository.dialect == "sqlite": rows = connection.execute(f"PRAGMA index_list({table_name})").fetchall() return {_row_value(row, "name") for row in rows} rows = connection.execute( """ SELECT indexname FROM pg_indexes WHERE schemaname = 'public' AND tablename = %s """, (table_name,), ).fetchall() return {_row_value(row, "indexname") for row in rows} def count_rows(self, table_name: str) -> int: with self._repository.connection() as connection: row = connection.execute(f"SELECT COUNT(*) AS count FROM {table_name}").fetchone() return int(_row_value(row, "count")) def _user_from_row(row: Any) -> UserSummary: return UserSummary( id=_row_value(row, "id"), display_name=_row_value(row, "display_name"), role=_row_value(row, "role"), ) def _target_site_from_row(row: Any) -> TargetSiteConfig: return TargetSiteConfig( id=_row_value(row, "id"), name=_row_value(row, "name"), slug=_row_value(row, "slug"), publishing_type=_row_value(row, "publishing_type"), default_language=_row_value(row, "default_language"), brand_voice=_row_value(row, "brand_voice"), audience=_row_value(row, "audience"), seo_rules=_json_from_row(row, "seo_rules"), visual_rules=_json_from_row(row, "visual_rules"), source_rules=_json_from_row(row, "source_rules"), publishing_rules=_json_from_row(row, "publishing_rules"), active_script_config_version_id=_row_value( row, "active_script_config_version_id" ), created_at=_row_value(row, "created_at"), updated_at=_row_value(row, "updated_at"), ) def _article_summary_from_row(row: Any) -> ArticleSummary: return ArticleSummary( id=_row_value(row, "id"), target_site_id=_row_value(row, "target_site_id"), status=_row_value(row, "status"), publishing_status=_row_value(row, "publishing_status"), brief_description=_row_value(row, "brief_description"), working_title=_row_value(row, "working_title"), language=_row_value(row, "language"), content_type=_row_value(row, "content_type"), primary_keyword=_row_value(row, "primary_keyword"), assigned_editor_id=_row_value(row, "assigned_editor_id"), created_at=_row_value(row, "created_at"), updated_at=_row_value(row, "updated_at"), ) def _workflow_event_from_row(row: Any) -> WorkflowEventSummary: return WorkflowEventSummary( id=_row_value(row, "id"), article_id=_row_value(row, "article_id"), event_type=_row_value(row, "event_type"), from_status=_row_value(row, "from_status"), to_status=_row_value(row, "to_status"), actor_user_id=_row_value(row, "actor_user_id"), payload=_json_from_row(row, "payload"), created_at=_row_value(row, "created_at"), ) def _script_config_version_event_from_row(row: Any) -> dict[str, Any]: return { "id": _row_value(row, "id"), "target_site_id": _row_value(row, "target_site_id"), "version_id": _row_value(row, "version_id"), "event_type": _row_value(row, "event_type"), "actor_user_id": _row_value(row, "actor_user_id"), "payload": _json_from_row(row, "payload"), "created_at": _row_value(row, "created_at"), } def _plain_row(row: Any) -> dict[str, Any]: if isinstance(row, sqlite3.Row): result = dict(row) else: result = dict(row) for key in ( "diff", "seo_rules", "visual_rules", "source_rules", "publishing_rules", ): if key in result: result[key] = _json_decode(result[key]) return result def _row_value(row: Any, key: str) -> Any: return row[key] def _json_from_row(row: Any, key: str) -> JsonObject: return _json_decode(_row_value(row, key)) def _json_decode(value: Any) -> Any: if isinstance(value, str): return json.loads(value) return value def _json_value(value: Any) -> str: return json.dumps(value, sort_keys=True, separators=(",", ":")) def _uuid_value(value: UUID | None) -> str | None: if value is None: return None return str(value) def _datetime_value(value: datetime | None) -> str | None: if value is None: return None return value.isoformat() def _article_status_value(value: ArticleWorkflowStatus | None) -> str | None: if value is None: return None return value.value