"""Database operations for the extraction_queue table.

Provides claim/update/query operations that the worker loop uses to:
1. Claim PENDING tasks (SELECT ... FOR UPDATE SKIP LOCKED)
2. Mark tasks as PROCESSING, DONE, or FAILED
3. Query queue statistics for health monitoring
"""

import json
import logging
from dataclasses import dataclass
from typing import Any
from uuid import UUID

import psycopg
from psycopg.rows import dict_row

from .config import ExtractorConfig

logger = logging.getLogger(__name__)


@dataclass
class ExtractionTask:
    """A task claimed from the extraction_queue."""

    id: UUID
    source_url: str
    source_format: str
    source_name: str | None
    source_tier: int | None
    request_payload: dict[str, Any]
    page_range: str | None
    import_batch_id: UUID | None
    import_task_run_id: UUID | None
    attempts: int


@dataclass
class QueueStats:
    """Queue depth statistics."""

    pending: int = 0
    processing: int = 0
    done: int = 0
    failed: int = 0
    cancelled: int = 0


class ExtractionQueueDB:
    """Database operations for the extraction_queue table."""

    def __init__(self, config: ExtractorConfig) -> None:
        self._config = config
        self._conn: psycopg.Connection[dict[str, Any]] | None = None

    def connect(self) -> None:
        """Establish database connection."""
        self._conn = psycopg.connect(
            self._config.database_url,
            row_factory=dict_row,
            autocommit=False,
        )
        logger.info("Connected to database")

    def close(self) -> None:
        """Close database connection."""
        if self._conn:
            self._conn.close()
            logger.info("Database connection closed")

    @property
    def conn(self) -> psycopg.Connection[dict[str, Any]]:
        if self._conn is None:
            raise RuntimeError("Database connection not established. Call ensure_connected() first.")
        return self._conn

    def ensure_connected(self) -> None:
        """Ensure database connection is open, reconnecting if necessary."""
        if self._conn is None or self._conn.closed:
            logger.info("Reconnecting to database...")
            self.connect()

    @property
    def is_connected(self) -> bool:
        """Whether the database connection is currently open."""
        return self._conn is not None and not self._conn.closed

    def claim_pending_task(self) -> ExtractionTask | None:
        """Claim one PENDING task using SELECT ... FOR UPDATE SKIP LOCKED.

        Returns the claimed task or None if the queue is empty.
        The task is atomically transitioned to PROCESSING status.
        """
        self.ensure_connected()
        with self.conn.cursor() as cur:
            cur.execute(
                """
                SELECT id, source_url, source_format, source_name, source_tier,
                       request_payload, page_range, import_batch_id,
                       import_task_run_id, attempts
                FROM extraction_queue
                WHERE status = 'PENDING' AND attempts < max_attempts
                ORDER BY created_at
                FOR UPDATE SKIP LOCKED
                LIMIT 1
                """,
            )
            row = cur.fetchone()
            if row is None:
                self.conn.rollback()
                return None

            task_id = row["id"]

            # Atomically transition to PROCESSING
            cur.execute(
                """
                UPDATE extraction_queue
                SET status = 'PROCESSING',
                    claimed_at = NOW(),
                    claimed_by = %(hostname)s,
                    attempts = attempts + 1,
                    updated_at = NOW()
                WHERE id = %(id)s
                """,
                {"id": task_id, "hostname": self._config.hostname},
            )
            self.conn.commit()

            return ExtractionTask(
                id=row["id"],
                source_url=row["source_url"],
                source_format=row["source_format"],
                source_name=row["source_name"],
                source_tier=row["source_tier"],
                request_payload=(
                    row["request_payload"]
                    if isinstance(row["request_payload"], dict)
                    else json.loads(row["request_payload"])
                ),
                page_range=row["page_range"],
                import_batch_id=row["import_batch_id"],
                import_task_run_id=row["import_task_run_id"],
                # The row was read before the claim UPDATE incremented it;
                # the task object must carry the attempt now in progress.
                attempts=row["attempts"] + 1,
            )

    def mark_done(
        self,
        task_id: UUID,
        result_payload: list[dict[str, Any]],
        extraction_method: str,
        confidence: float,
        duration_ms: int,
        source_checksum: str | None = None,
    ) -> None:
        """Mark a task as successfully completed with results."""
        self.ensure_connected()
        with self.conn.cursor() as cur:
            cur.execute(
                """
                UPDATE extraction_queue
                SET status = 'DONE',
                    result_payload = %(result)s::jsonb,
                    result_count = %(count)s,
                    extraction_method = %(method)s,
                    extractor_version = %(version)s,
                    confidence = %(confidence)s,
                    duration_ms = %(duration)s,
                    source_checksum = %(checksum)s,
                    updated_at = NOW()
                WHERE id = %(id)s
                """,
                {
                    "id": task_id,
                    "result": json.dumps(result_payload),
                    "count": len(result_payload),
                    "method": extraction_method,
                    "version": self._config.extractor_version,
                    "confidence": confidence,
                    "duration": duration_ms,
                    "checksum": source_checksum,
                },
            )
            self.conn.commit()
        logger.info(
            "Task completed",
            extra={"task_id": str(task_id), "result_count": len(result_payload), "duration_ms": duration_ms},
        )

    def mark_failed(self, task_id: UUID, error_detail: dict[str, Any]) -> None:
        """Mark a task as failed with error details."""
        self.ensure_connected()
        with self.conn.cursor() as cur:
            cur.execute(
                """
                UPDATE extraction_queue
                SET status = 'FAILED',
                    error_detail = %(error)s::jsonb,
                    last_error_at = NOW(),
                    updated_at = NOW()
                WHERE id = %(id)s
                """,
                {
                    "id": task_id,
                    "error": json.dumps(error_detail),
                },
            )
            self.conn.commit()
        logger.warning(
            "Task %s failed: %s (%s)",
            task_id,
            error_detail.get("message", "unknown"),
            error_detail.get("type", "unknown"),
        )

    def get_queue_stats(self) -> QueueStats:
        """Get current queue depth by status."""
        self.ensure_connected()
        with self.conn.cursor() as cur:
            cur.execute(
                """
                SELECT status::text, COUNT(*) as cnt
                FROM extraction_queue
                GROUP BY status
                """
            )
            stats = QueueStats()
            for row in cur.fetchall():
                status = row["status"].lower()
                setattr(stats, status, row["cnt"])
            return stats

    def list_composer_reference_rows(self) -> list[dict[str, Any]]:
        """Load composer reference rows (with aliases where available)."""
        self.ensure_connected()
        with self.conn.cursor() as cur:
            try:
                cur.execute(
                    """
                    SELECT
                        c.id::text AS entity_id,
                        c.name AS name,
                        COALESCE(
                            array_agg(ca.alias_normalized) FILTER (WHERE ca.alias_normalized IS NOT NULL),
                            '{}'
                        ) AS aliases
                    FROM composers c
                    LEFT JOIN composer_aliases ca
                      ON ca.composer_id = c.id
                    GROUP BY c.id, c.name
                    """
                )
            except Exception:
                self.conn.rollback()
                cur.execute(
                    """
                    SELECT
                        c.id::text AS entity_id,
                        c.name AS name,
                        '{}'::text[] AS aliases
                    FROM composers c
                    """
                )
            rows = cur.fetchall()
            self.conn.rollback()
            return rows

    def list_raga_reference_rows(self) -> list[dict[str, Any]]:
        """Load raga reference rows used by identity candidate discovery."""
        self.ensure_connected()
        with self.conn.cursor() as cur:
            cur.execute(
                """
                SELECT
                    r.id::text AS entity_id,
                    r.name AS name
                FROM ragas r
                """
            )
            rows = cur.fetchall()
            self.conn.rollback()
            return rows

    def health_check(self) -> bool:
        """Verify database connectivity."""
        self.ensure_connected()
        try:
            with self.conn.cursor() as cur:
                cur.execute("SELECT 1")
                return True
        except Exception:
            return False
