#!/usr/bin/env python3
"""Small PostgreSQL repository for the local Xiaohongshu workflow."""

from __future__ import annotations

import json
import os
from pathlib import Path
from typing import Any, Dict, Optional
from decimal import Decimal

import psycopg2
from psycopg2.extras import RealDictCursor


ROOT = Path(__file__).resolve().parent
SCHEMA_PATH = ROOT / "schema.sql"
DEFAULT_DATABASE_URL = "postgresql://postgres:password@127.0.0.1:5432/xhs_automation"


def database_url(value: Optional[str] = None) -> str:
    return value or os.environ.get("XHS_DATABASE_URL", DEFAULT_DATABASE_URL)


def connect(value: Optional[str] = None):
    return psycopg2.connect(database_url(value))


def init_schema(value: Optional[str] = None) -> None:
    sql = SCHEMA_PATH.read_text(encoding="utf-8")
    with connect(value) as conn:
        with conn.cursor() as cur:
            cur.execute(sql)


def save_task(task: Dict[str, Any], value: Optional[str] = None) -> None:
    version = int(str(task.get("version", "v1")).lstrip("v"))
    product_id = task.get("product_id") or task["task_id"]
    with connect(value) as conn:
        with conn.cursor() as cur:
            cur.execute(
                """INSERT INTO products (id, name, price, product_url, selling_points, source_photo)
                   VALUES (%s, %s, %s, %s, %s, %s)
                   ON CONFLICT (id) DO UPDATE SET name=EXCLUDED.name, price=EXCLUDED.price,
                     product_url=EXCLUDED.product_url, selling_points=EXCLUDED.selling_points,
                     source_photo=EXCLUDED.source_photo""",
                (product_id, task["product_name"], task.get("price"), task.get("product_url"), task.get("selling_points"), task.get("source_photo")),
            )
            cur.execute(
                """INSERT INTO content_tasks (task_id, product_id, parent_task_id, version, status, title,
                     cover_copy, body, topics, generation_mode, source_photo, review_url, review_notes)
                   VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
                   ON CONFLICT (task_id) DO UPDATE SET product_id=EXCLUDED.product_id, parent_task_id=EXCLUDED.parent_task_id,
                     version=EXCLUDED.version, status=EXCLUDED.status, title=EXCLUDED.title,
                     cover_copy=EXCLUDED.cover_copy, body=EXCLUDED.body, topics=EXCLUDED.topics,
                     generation_mode=EXCLUDED.generation_mode, source_photo=EXCLUDED.source_photo,
                     review_url=EXCLUDED.review_url, review_notes=EXCLUDED.review_notes, updated_at=now()""",
                (task["task_id"], product_id, task.get("parent_task_id"), version, task.get("status", "待审核"),
                 task["title"], task.get("cover_copy"), task["body"], task.get("topics"), task.get("generation_mode", "demo-assets"),
                 task.get("source_photo"), task.get("review_url"), task.get("review_notes")),
            )
            cur.execute("DELETE FROM content_images WHERE task_id=%s", (task["task_id"],))
            for position, image in enumerate(task.get("images", []), 1):
                cur.execute(
                    "INSERT INTO content_images (task_id, position, url, caption) VALUES (%s, %s, %s, %s)",
                    (task["task_id"], position, image["url"], image.get("caption", "")),
                )


def get_task(task_id: str, value: Optional[str] = None) -> Optional[Dict[str, Any]]:
    with connect(value) as conn:
        with conn.cursor(cursor_factory=RealDictCursor) as cur:
            cur.execute(
                """SELECT t.*, p.name AS product_name, p.price, p.product_url, p.selling_points
                   FROM content_tasks t LEFT JOIN products p ON p.id=t.product_id WHERE t.task_id=%s""",
                (task_id,),
            )
            task = cur.fetchone()
            if not task:
                return None
            cur.execute("SELECT url, caption FROM content_images WHERE task_id=%s ORDER BY position", (task_id,))
            task["images"] = [dict(row) for row in cur.fetchall()]
            if isinstance(task.get("price"), Decimal):
                task["price"] = float(task["price"])
            task["version"] = f"v{task['version']}"
            for key in ("created_at", "updated_at"):
                if task.get(key):
                    task[key] = task[key].isoformat()
            return dict(task)


def record_review(task_id: str, status: str, notes: str = "", value: Optional[str] = None) -> Dict[str, Any]:
    if status not in {"已确认", "驳回"}:
        raise ValueError("status must be 已确认 or 驳回")
    with connect(value) as conn:
        with conn.cursor() as cur:
            cur.execute("UPDATE content_tasks SET status=%s, review_notes=%s, updated_at=now() WHERE task_id=%s", (status, notes, task_id))
            if cur.rowcount != 1:
                raise ValueError(f"task not found: {task_id}")
            cur.execute("INSERT INTO reviews (task_id, status, notes) VALUES (%s, %s, %s)", (task_id, status, notes))
            if status == "已确认":
                cur.execute("INSERT INTO publish_queue (task_id) VALUES (%s)", (task_id,))
    return {"task_id": task_id, "status": status, "notes": notes}


def claim_publish_item(value: Optional[str] = None) -> Optional[Dict[str, Any]]:
    """Atomically claim the oldest approved item for the local publisher worker."""
    with connect(value) as conn:
        with conn.cursor(cursor_factory=RealDictCursor) as cur:
            cur.execute(
                """SELECT id, task_id FROM publish_queue
                   WHERE status = '待发布'
                   ORDER BY created_at, id
                   FOR UPDATE SKIP LOCKED LIMIT 1"""
            )
            item = cur.fetchone()
            if not item:
                return None
            cur.execute(
                """UPDATE publish_queue
                   SET status='发布中', error=NULL, updated_at=now()
                   WHERE id=%s""",
                (item["id"],),
            )
            return dict(item)


def finish_publish_item(queue_id: int, status: str, published_url: str = "", error: str = "", value: Optional[str] = None) -> None:
    if status not in {"已发布", "失败", "暂停"}:
        raise ValueError("invalid publish queue status")
    with connect(value) as conn:
        with conn.cursor() as cur:
            cur.execute(
                """UPDATE publish_queue
                   SET status=%s, published_url=%s, error=%s, updated_at=now()
                   WHERE id=%s""",
                (status, published_url or None, error or None, queue_id),
            )


def mark_task_published(task_id: str, value: Optional[str] = None) -> None:
    with connect(value) as conn:
        with conn.cursor() as cur:
            cur.execute(
                "UPDATE content_tasks SET status='已发布', updated_at=now() WHERE task_id=%s",
                (task_id,),
            )


def export_task(task: Dict[str, Any], path: Path) -> None:
    path.write_text(json.dumps(task, ensure_ascii=False, indent=2), encoding="utf-8")
