"""Runs that survive a restart (Phase 2a).

The in-process store loses everything when the service stops, which is precisely
the situation resumability exists for: a proposal submitted at 17:00 and approved
at 09:00 the next morning spans at least one deploy.

Two rules shape the SQL:

**A save is one transaction.** Run and steps move together or not at all. A run
whose state says `completed` above steps that say `pending` is a corruption no
later read can detect.

**Every save asserts the version it read.** Two processes that both loaded a run
cannot both write it; the second is refused. That is what makes it safe to run
the poller alongside anything else, and it is the same discipline as the CMS's
conditional UPDATE on proposals.
"""

from __future__ import annotations

from datetime import UTC, datetime
from typing import Any

from psycopg.types.json import Jsonb

from app.workflow.engine import ConcurrentUpdateError
from app.workflow.serialisation import row_to_run, run_to_row, step_to_row
from app.workflow.state import RunState, WorkflowRun

#: States a run can still be moved on from. Kept here as values rather than
#: computed, because it is embedded in a partial index — the two must agree, and
#: a mismatch would silently stop the poller from seeing work.
RESUMABLE_STATES = (RunState.PENDING.value, RunState.RUNNING.value, RunState.AWAITING_APPROVAL.value)

_RUN_COLUMNS = "id, workflow, state, error, context, version, created_at, updated_at, archived_at"
_STEP_COLUMNS = (
    "run_id, position, name, tool, state, attempts, output, proposed, amendments, "
    "approved, approved_by, error, confidence, duration_ms, started_at, finished_at"
)

#: Columns holding JSON, which psycopg needs told about explicitly.
_JSON_FIELDS = frozenset({"context", "output", "proposed", "amendments"})


class PostgresRunStore:
    """Durable workflow storage."""

    def __init__(self, connect) -> None:
        # A callable, not a connection: this is constructed at startup, possibly
        # before the database is reachable, and a service that cannot start
        # because it eagerly opened a socket is a service that cannot report why.
        self._connect = connect

    # ── Writing ───────────────────────────────────────────────────────────

    def save(self, run: WorkflowRun) -> None:
        """Insert or update, refusing to overwrite somebody else's work."""
        run.touch()
        expected = run.version
        row = run_to_row(run)
        row["version"] = expected + 1

        with self._connect() as connection, connection.transaction():
            with connection.cursor() as cursor:
                if expected == 0:
                    self._insert_run(cursor, row)
                else:
                    self._update_run(cursor, row, expected)

                # Deleted and rewritten rather than diffed. A step list only ever
                # changes wholesale here, and reconciling row by row would be
                # more code guarding against a case that does not arise.
                cursor.execute("DELETE FROM workflow_steps WHERE run_id = %s", (run.id,))

                for position, step in enumerate(run.steps):
                    self._insert_step(cursor, step_to_row(run.id, position, step))

        run.version = expected + 1

    def _insert_run(self, cursor, row: dict[str, Any]) -> None:
        columns = list(row)
        cursor.execute(
            f"INSERT INTO workflow_runs ({', '.join(columns)}) "
            f"VALUES ({', '.join(['%s'] * len(columns))})",
            tuple(_adapt(c, row[c]) for c in columns),
        )

    def _update_run(self, cursor, row: dict[str, Any], expected: int) -> None:
        updatable = [c for c in row if c != "id"]

        cursor.execute(
            f"UPDATE workflow_runs SET {', '.join(f'{c} = %s' for c in updatable)} "
            "WHERE id = %s AND version = %s",
            (*(_adapt(c, row[c]) for c in updatable), row["id"], expected),
        )

        if cursor.rowcount == 0:
            # Either the row is gone or its version moved. Both mean this run was
            # read before somebody else changed it, and writing now would discard
            # whatever they did.
            raise ConcurrentUpdateError(
                f"Run {row['id']} was version {expected} here but has since changed."
            )

    def _insert_step(self, cursor, row: dict[str, Any]) -> None:
        columns = list(row)
        cursor.execute(
            f"INSERT INTO workflow_steps ({', '.join(columns)}) "
            f"VALUES ({', '.join(['%s'] * len(columns))})",
            tuple(_adapt(c, row[c]) for c in columns),
        )

    # ── Reading ───────────────────────────────────────────────────────────

    def load(self, run_id: str) -> WorkflowRun | None:
        with self._connect() as connection, connection.cursor() as cursor:
            cursor.execute(f"SELECT {_RUN_COLUMNS} FROM workflow_runs WHERE id = %s", (run_id,))
            row = cursor.fetchone()

            if row is None:
                return None

            run_row = dict(zip(_RUN_COLUMNS.split(", "), row, strict=True))

            cursor.execute(
                f"SELECT {_STEP_COLUMNS} FROM workflow_steps WHERE run_id = %s ORDER BY position",
                (run_id,),
            )
            step_rows = [
                dict(zip(_STEP_COLUMNS.split(", "), r, strict=True)) for r in cursor.fetchall()
            ]

        return row_to_run(run_row, step_rows)

    def resumable(self) -> list[WorkflowRun]:
        """Runs that still need something.

        Two queries rather than a join: a join returns one row per step and the
        rebuilding has to group them anyway, while this reads exactly the rows
        needed and lets the partial index do its job.
        """
        with self._connect() as connection, connection.cursor() as cursor:
            cursor.execute(
                f"SELECT {_RUN_COLUMNS} FROM workflow_runs "
                "WHERE archived_at IS NULL AND state = ANY(%s) ORDER BY updated_at",
                (list(RESUMABLE_STATES),),
            )
            run_rows = [
                dict(zip(_RUN_COLUMNS.split(", "), r, strict=True)) for r in cursor.fetchall()
            ]

            if not run_rows:
                return []

            cursor.execute(
                f"SELECT {_STEP_COLUMNS} FROM workflow_steps "
                "WHERE run_id = ANY(%s) ORDER BY run_id, position",
                ([r["id"] for r in run_rows],),
            )
            all_steps = [
                dict(zip(_STEP_COLUMNS.split(", "), r, strict=True)) for r in cursor.fetchall()
            ]

        by_run: dict[str, list[dict[str, Any]]] = {}

        for step in all_steps:
            by_run.setdefault(step["run_id"], []).append(step)

        return [row_to_run(row, by_run.get(row["id"], [])) for row in run_rows]

    # ── Housekeeping ──────────────────────────────────────────────────────

    def archive(self, older_than: datetime) -> int:
        """Move finished runs out of the working set.

        An UPDATE, not a DELETE. The run is the account of what happened to a
        client's document — the reason to stop loading it is volume, not
        irrelevance, and a deletion would destroy the only record of a filing on
        this side of the system.
        """
        with self._connect() as connection, connection.transaction(), connection.cursor() as cursor:
            cursor.execute(
                "UPDATE workflow_runs SET archived_at = %s "
                "WHERE archived_at IS NULL AND updated_at < %s "
                "AND state IN ('completed', 'failed')",
                (datetime.now(UTC), older_than),
            )

            return cursor.rowcount


def _adapt(column: str, value: Any) -> Any:
    """Hand psycopg a JSON value it will store as jsonb rather than a string."""
    return Jsonb(value if value is not None else {}) if column in _JSON_FIELDS else value
