"""Worker liveness (task 5.6).

Written by the worker, read by the API's health endpoint. The staleness thresholds are
derived from each task's interval rather than hardcoded, so changing the poll interval in
config cannot leave the health check quietly asserting the wrong thing.
"""

from dataclasses import dataclass
from datetime import UTC, datetime, timedelta

from sqlalchemy import select
from sqlalchemy.dialects.postgresql import insert
from sqlalchemy.ext.asyncio import AsyncSession

from app.models import WorkerHeartbeat

#: Task names. These match the APScheduler job ids so a stale row names its own job.
PLANNER = "planner"
SENDER = "sender"

#: A tick is late if it has not landed within this many intervals. Two, not one: a single
#: missed poll is normal jitter, two in a row is a worker that is not running.
GRACE_FACTOR = 3


async def beat(session: AsyncSession, name: str, *, error: str | None = None) -> None:
    """Stamp a task as alive. Upsert so the first tick after a fresh deploy works."""
    now = datetime.now(tz=UTC)
    await session.execute(
        insert(WorkerHeartbeat)
        .values(name=name, beat_at=now, ticks=1, last_error=error)
        .on_conflict_do_update(
            index_elements=[WorkerHeartbeat.name],
            set_={
                "beat_at": now,
                "ticks": WorkerHeartbeat.ticks + 1,
                "last_error": error,
                "updated_at": now,
            },
        )
    )


@dataclass(frozen=True)
class TaskLiveness:
    name: str
    beat_at: datetime | None
    age_seconds: float | None
    expected_interval_seconds: int
    last_error: str | None

    @property
    def is_alive(self) -> bool:
        if self.age_seconds is None:
            return False
        return self.age_seconds <= self.expected_interval_seconds * GRACE_FACTOR


async def read_liveness(
    session: AsyncSession, expected: dict[str, int], *, now: datetime | None = None
) -> list[TaskLiveness]:
    """Report each expected task against its own interval.

    `expected` maps task name to its interval in seconds; a task that has never beaten at
    all still appears, reported dead, rather than being absent from the response.
    """
    moment = now or datetime.now(tz=UTC)
    rows = {r.name: r for r in await session.scalars(select(WorkerHeartbeat))}

    return [
        TaskLiveness(
            name=name,
            beat_at=rows[name].beat_at if name in rows else None,
            age_seconds=((moment - rows[name].beat_at).total_seconds() if name in rows else None),
            expected_interval_seconds=interval,
            last_error=rows[name].last_error if name in rows else None,
        )
        for name, interval in expected.items()
    ]


def stale_after(interval_seconds: int) -> timedelta:
    """The window a caller should allow before treating a task as down."""
    return timedelta(seconds=interval_seconds * GRACE_FACTOR)
