"""Browse imported / published forecasts and find analogs for a draft."""

from __future__ import annotations

from datetime import date

from fastapi import APIRouter, Depends, HTTPException
from psycopg import Connection
from pydantic import BaseModel

from .. import engine
from ..auth import User, current_user, require_center
from ..configs import active_config
from ..db import fetchall, fetchone, get_conn
from ..retrieval import find_analogs
from .runs import _load_run, validate_inputs

router = APIRouter(tags=["forecasts"])


class AnalogIn(BaseModel):
    center_id: int
    zone_id: int | None = None
    valid_date: date
    inputs: dict
    k: int = 8


@router.get("/forecasts")
def list_forecasts(center_id: int, season: str | None = None, q: str | None = None, problem_type: int | None = None,
                   limit: int = 50, offset: int = 0, conn: Connection = Depends(get_conn), user: User = Depends(current_user)):
    require_center(user, center_id)
    where = ["f.center_id = %(c)s"]
    params: dict = {"c": center_id, "limit": min(max(limit, 1), 200), "offset": max(offset, 0)}
    if season:
        where.append("f.season = %(season)s")
        params["season"] = season
    if q:
        where.append("f.tsv @@ websearch_to_tsquery('english', %(q)s)")
        params["q"] = q
    if problem_type:
        where.append("EXISTS (SELECT 1 FROM forecast_problems p WHERE p.forecast_id = f.id AND p.problem_type = %(pt)s)")
        params["pt"] = problem_type
    clause = " AND ".join(where)
    total = fetchone(conn, f"SELECT count(*)::int AS n FROM forecasts f WHERE {clause}", params)["n"]
    items = fetchall(
        conn,
        f"""SELECT f.id, f.valid_date, f.season, f.source, f.guidance_era, f.title, f.danger_upper, f.danger_middle,
                   f.danger_lower, f.normal_caution, left(f.bottom_line, 300) AS bottom_line,
                   (SELECT array_agg(p.problem_type ORDER BY p.rank) FROM forecast_problems p WHERE p.forecast_id = f.id) AS problem_types
            FROM forecasts f WHERE {clause} ORDER BY f.valid_date DESC LIMIT %(limit)s OFFSET %(offset)s""",
        params,
    )
    return {"total": total, "items": items}


@router.get("/forecasts/seasons")
def seasons(center_id: int, conn: Connection = Depends(get_conn), user: User = Depends(current_user)):
    require_center(user, center_id)
    return fetchall(
        conn,
        "SELECT season, count(*)::int AS n, min(valid_date) AS first, max(valid_date) AS last FROM forecasts WHERE center_id = %s GROUP BY season ORDER BY season",
        (center_id,),
    )


@router.get("/forecasts/{forecast_id}")
def get_forecast(forecast_id: int, conn: Connection = Depends(get_conn), user: User = Depends(current_user)):
    f = fetchone(conn, "SELECT f.*, z.name AS zone_name FROM forecasts f LEFT JOIN zones z ON z.id = f.zone_id WHERE f.id = %s", (forecast_id,))
    if not f:
        raise HTTPException(404, "Forecast not found")
    require_center(user, f["center_id"])
    f.pop("tsv", None)
    f.pop("raw", None)
    f["problems"] = fetchall(conn, "SELECT * FROM forecast_problems WHERE forecast_id = %s ORDER BY rank", (forecast_id,))
    return f


def _query_from_inputs(inputs: dict, config: dict) -> list[dict]:
    computed = engine.compute_run(inputs, config)
    return [
        {"type": p["type"], "locations": p["locations"], "likelihood": computed["problems"][i]["likelihood"], "size_max": p["size"]["max"]}
        for i, p in enumerate(inputs["problems"])
    ]


@router.post("/analogs")
def analogs(body: AnalogIn, conn: Connection = Depends(get_conn), user: User = Depends(current_user)):
    require_center(user, body.center_id)
    config = active_config(conn, body.center_id)["config"]
    inputs = validate_inputs(body.inputs, config)
    return find_analogs(conn, body.center_id, _query_from_inputs(inputs, config), body.valid_date, body.zone_id, min(max(body.k, 1), 25))


@router.get("/runs/{run_id}/analogs")
def run_analogs(run_id: int, k: int = 8, conn: Connection = Depends(get_conn), user: User = Depends(current_user)):
    run = _load_run(conn, run_id, user)
    config = fetchone(conn, "SELECT config FROM center_configs WHERE id = %s", (run["config_id"],))["config"]
    exclude = {run["forecast_id"]} if run["forecast_id"] else set()
    return find_analogs(conn, run["center_id"], _query_from_inputs(run["inputs"], config), run["valid_date"], run["zone_id"], min(max(k, 1), 25), exclude)
