Files
content-factory/apps/backend/src/infrastructure/repositories.py
T

1082 lines
35 KiB
Python

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