"""What every run store must do, whichever one it is.

One suite, run against each implementation. The in-process store runs always;
PostgresRunStore runs when TAXPILOT_TEST_DSN points at a database:

    TAXPILOT_TEST_DSN=postgresql://postgres:postgres@127.0.0.1:5433/taxpilot_test       python -m pytest

Both halves pass against PostgreSQL 16.10 — see docs/storage.md for starting a
local one, including the clean-PATH requirement that is not obvious.

That is the whole point of writing it this way. A durable store tested only by
its own bespoke tests, against semantics nobody wrote down, diverges from the
one everything else was developed against — and the divergence surfaces in
production, after a restart, which is exactly when nobody wants to find it.
"""

from __future__ import annotations

import os
from datetime import UTC, datetime, timedelta

import pytest

from app.workflow.engine import ConcurrentUpdateError, InMemoryRunStore
from app.workflow.state import RunState, StepRecord, StepState, WorkflowRun

DSN = os.environ.get("TAXPILOT_TEST_DSN")


class StoreUnderTest:
    """The store, plus the one thing a test needs that production must not have.

    ``age`` backdates a stored run. Archival keys off ``updated_at``, and every
    save legitimately stamps that to now — so there is no way through the real
    interface to produce a run that was written ninety days ago, and no reason
    production should have one.

    Everything else delegates, so the tests below exercise the actual store.
    """

    def __init__(self, store, age) -> None:
        self._store = store
        self.age = age

    def __getattr__(self, name):
        return getattr(self._store, name)


def _in_memory() -> StoreUnderTest:
    store = InMemoryRunStore()

    def age(run_id: str, when: datetime) -> None:
        store._runs[run_id].updated_at = when  # noqa: SLF001

    return StoreUnderTest(store, age)


def _postgres() -> StoreUnderTest:
    """A live store, or a skip explaining why not."""
    psycopg = pytest.importorskip("psycopg", reason="psycopg is not installed")

    from app.database.migrator import Migrator
    from app.workflow.postgres_store import PostgresRunStore

    def connect():
        return psycopg.connect(DSN)

    Migrator(connect).apply()

    with connect() as connection, connection.cursor() as cursor:
        cursor.execute("TRUNCATE workflow_runs CASCADE")
        connection.commit()

    def age(run_id: str, when: datetime) -> None:
        with connect() as connection, connection.cursor() as cursor:
            cursor.execute(
                "UPDATE workflow_runs SET updated_at = %s WHERE id = %s", (when, run_id)
            )
            connection.commit()

    return StoreUnderTest(PostgresRunStore(connect), age)


@pytest.fixture(
    params=[
        pytest.param("memory", id="in-memory"),
        pytest.param(
            "postgres",
            id="postgres",
            marks=pytest.mark.skipif(not DSN, reason="TAXPILOT_TEST_DSN is not set"),
        ),
    ]
)
def store(request):
    return _in_memory() if request.param == "memory" else _postgres()


def make_run(**overrides) -> WorkflowRun:
    run = WorkflowRun(
        workflow="document_intake",
        context={"path": "/tmp/cnic.jpg", "sender": "923001234567"},
        steps=[
            StepRecord(name="read", tool="ocr_document"),
            StepRecord(name="file", tool="submit_filing_proposal"),
        ],
    )

    for key, value in overrides.items():
        setattr(run, key, value)

    return run


class TestRoundTrip:
    def test_a_saved_run_comes_back(self, store):
        run = make_run()
        store.save(run)

        loaded = store.load(run.id)

        assert loaded is not None
        assert loaded.workflow == "document_intake"
        assert loaded.context["sender"] == "923001234567"

    def test_steps_come_back_in_order(self, store):
        run = make_run()
        store.save(run)

        # "Resume from the first step not done" depends on this. Out of order,
        # a run restarts halfway through and re-files a document.
        assert [s.name for s in store.load(run.id).steps] == ["read", "file"]

    def test_everything_on_a_step_survives(self, store):
        run = make_run()
        step = run.steps[0]
        step.state = StepState.SUCCEEDED
        step.attempts = 2
        step.output = {"text": "NATIONAL IDENTITY CARD", "readable": True}
        step.proposed = {"client_id": 42}
        step.amendments = {"client_id": 77}
        step.approved = True
        step.approved_by = "hina@firm.test"
        step.confidence = 0.94
        step.duration_ms = 812
        step.started_at = datetime.now(UTC)
        step.finished_at = datetime.now(UTC)

        store.save(run)
        loaded = store.load(run.id).steps[0]

        # A field added to the dataclass and forgotten in the mapping loses data
        # silently, and only after a restart.
        assert loaded.state is StepState.SUCCEEDED
        assert loaded.attempts == 2
        assert loaded.output["readable"] is True
        assert loaded.proposed == {"client_id": 42}
        assert loaded.amendments == {"client_id": 77}
        assert loaded.approved is True
        assert loaded.approved_by == "hina@firm.test"
        assert loaded.confidence == pytest.approx(0.94)
        assert loaded.duration_ms == 812
        assert loaded.started_at is not None

    def test_an_unknown_run_is_absent_rather_than_an_error(self, store):
        assert store.load("does-not-exist") is None

    def test_nested_context_survives(self, store):
        run = make_run()
        run.context["identify"] = {"matches": [{"id": 42, "name": "Muhammad Ali"}]}
        store.save(run)

        assert store.load(run.id).context["identify"]["matches"][0]["id"] == 42


class TestOptimisticLocking:
    def test_saving_advances_the_version(self, store):
        run = make_run()
        store.save(run)

        assert run.version == 1

        store.save(run)

        assert run.version == 2

    def test_a_stale_run_cannot_overwrite_a_newer_one(self, store):
        run = make_run()
        store.save(run)

        # Two processes both loaded it. One writes.
        first = store.load(run.id)
        second = store.load(run.id)

        first.state = RunState.COMPLETED
        store.save(first)

        # The other must not be able to discard that.
        second.state = RunState.FAILED

        with pytest.raises(ConcurrentUpdateError):
            store.save(second)

    def test_the_winning_write_is_the_one_that_stands(self, store):
        run = make_run()
        store.save(run)

        first, second = store.load(run.id), store.load(run.id)
        first.error = "the real outcome"
        store.save(first)

        second.error = "should never be stored"

        with pytest.raises(ConcurrentUpdateError):
            store.save(second)

        assert store.load(run.id).error == "the real outcome"

    def test_a_reloaded_run_can_be_saved_again(self, store):
        run = make_run()
        store.save(run)
        store.save(store.load(run.id))

        # Losing a race is recoverable: read again and retry. If it were not,
        # the only remedy would be to abandon the run.
        fresh = store.load(run.id)
        fresh.state = RunState.COMPLETED
        store.save(fresh)

        assert store.load(run.id).state is RunState.COMPLETED


class TestResumable:
    @pytest.mark.parametrize(
        "state", [RunState.PENDING, RunState.RUNNING, RunState.AWAITING_APPROVAL]
    )
    def test_unfinished_runs_are_returned(self, store, state):
        run = make_run(state=state)
        store.save(run)

        assert [r.id for r in store.resumable()] == [run.id]

    @pytest.mark.parametrize("state", [RunState.COMPLETED, RunState.FAILED])
    def test_finished_runs_are_not(self, store, state):
        store.save(make_run(state=state))

        assert store.resumable() == []

    def test_a_resumable_run_brings_its_steps(self, store):
        run = make_run(state=RunState.AWAITING_APPROVAL)
        run.steps[1].state = StepState.AWAITING_APPROVAL
        run.steps[1].output = {"proposal_id": 901}
        store.save(run)

        # The poller matches decisions to runs through this. Without the steps
        # it cannot tell which proposal a run is waiting on.
        assert store.resumable()[0].steps[1].output["proposal_id"] == 901

    def test_an_archived_run_is_not_resumable(self, store):
        run = make_run(state=RunState.COMPLETED)
        store.save(run)
        store.age(run.id, datetime.now(UTC) - timedelta(days=90))
        store.archive(older_than=datetime.now(UTC) - timedelta(days=30))

        assert store.resumable() == []


class TestArchival:
    def test_old_finished_runs_are_archived(self, store):
        old = make_run(state=RunState.COMPLETED)
        store.save(old)
        store.age(old.id, datetime.now(UTC) - timedelta(days=90))

        assert store.archive(older_than=datetime.now(UTC) - timedelta(days=30)) == 1

    def test_recent_runs_are_left_alone(self, store):
        store.save(make_run(state=RunState.COMPLETED))

        assert store.archive(older_than=datetime.now(UTC) - timedelta(days=30)) == 0

    def test_unfinished_runs_are_never_archived(self, store):
        run = make_run(state=RunState.AWAITING_APPROVAL)
        store.save(run)
        store.age(run.id, datetime.now(UTC) - timedelta(days=90))

        # A proposal waiting a long time is the case the overdue check exists
        # for. Archiving it would hide work nobody has done.
        assert store.archive(older_than=datetime.now(UTC) - timedelta(days=30)) == 0
        assert len(store.resumable()) == 1

    def test_archiving_keeps_the_run(self, store):
        run = make_run(state=RunState.COMPLETED)
        store.save(run)
        store.age(run.id, datetime.now(UTC) - timedelta(days=90))
        store.archive(older_than=datetime.now(UTC) - timedelta(days=30))

        # The account of what happened to a client's document. The reason to
        # stop loading it is volume, not irrelevance.
        assert store.load(run.id) is not None


class TestIsolation:
    """The store must behave like a database, not like a shared object."""

    def test_mutating_a_run_after_saving_does_not_change_the_store(self, store):
        run = make_run()
        store.save(run)

        run.context["injected"] = "after the save"

        # In-process storage that hands back the same object makes this pass by
        # accident, and the behaviour changes the day it becomes Postgres.
        assert "injected" not in store.load(run.id).context

    def test_two_loads_are_independent(self, store):
        run = make_run()
        store.save(run)

        first, second = store.load(run.id), store.load(run.id)
        first.context["only_mine"] = True

        assert "only_mine" not in second.context
