Task 003: add postgres schema and seed data
This commit is contained in:
@@ -0,0 +1,566 @@
|
||||
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
|
||||
|
||||
from src.domain.contracts import (
|
||||
PublishingRules,
|
||||
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.script_config_versions = ScriptConfigVersionsRepository(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 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 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 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 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 _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()
|
||||
Reference in New Issue
Block a user