from __future__ import annotations

import asyncio
import time
from datetime import datetime, timedelta, timezone
from typing import Any

import httpx
from fastapi import APIRouter, BackgroundTasks, HTTPException
from pydantic import BaseModel, Field
from slack_sdk.web.async_client import AsyncWebClient
from slack_sdk.errors import SlackApiError

from app.core.config import settings
from app.core.logger import logger

router = APIRouter(prefix="/sync", tags=["Slack Sync"])

_active_runs: dict[str, dict] = {}  # simple in-memory run tracker


class SlackSyncRequest(BaseModel):
    user_token: str = Field(..., description="xoxp-... Slack user token")
    user_id: str = Field(..., description="Slack user ID to sweep (e.g. U01ABC123)")
    repo_id: str = Field(..., description="repo_id to index data under")
    snapshot_id: str = Field(
        default="",
        description="snapshot_id to use (defaults to snap-YYYY-MM-DD)",
    )
    days_back: int = Field(default=30, ge=1, le=365)


class SlackSyncResponse(BaseModel):
    run_id: str
    status: str
    repo_id: str
    snapshot_id: str


async def _resolve_user_names(
    client: AsyncWebClient,
    user_ids: set[str],
) -> dict[str, str]:
    """Return {user_id: display_name} for all requested IDs."""
    names: dict[str, str] = {}
    for uid in user_ids:
        try:
            resp = await client.users_info(user=uid)
            profile = (resp.get("user") or {}).get("profile") or {}
            names[uid] = profile.get("display_name") or profile.get("real_name") or uid
            await asyncio.sleep(0.2)
        except SlackApiError:
            names[uid] = uid
    return names


async def _fetch_channel_messages(
    client: AsyncWebClient,
    channel_id: str,
    oldest_ts: str,
    user_id: str,
) -> list[dict[str, Any]]:
    messages: list[dict] = []
    cursor = None
    while True:
        try:
            resp = await client.conversations_history(
                channel=channel_id,
                oldest=oldest_ts,
                limit=200,
                cursor=cursor,
            )
        except SlackApiError as exc:
            error = exc.response.get("error", "")
            if error == "ratelimited":
                wait = int(exc.response.headers.get("Retry-After", 10))
                await asyncio.sleep(wait)
                continue
            logger.warning("Skipping channel {channel}: {error}", channel=channel_id, error=error)
            return []

        for msg in resp.get("messages", []):
            if msg.get("type") != "message" or msg.get("subtype"):
                continue
            messages.append(msg)
            if msg.get("reply_count") and (
                msg.get("user") == user_id
                or user_id in (msg.get("reply_users") or [])
            ):
                try:
                    thread = await client.conversations_replies(
                        channel=channel_id, ts=msg["ts"], oldest=oldest_ts, limit=200
                    )
                    messages.extend(thread.get("messages", [])[1:])
                    await asyncio.sleep(0.3)
                except SlackApiError:
                    pass

        cursor = (resp.get("response_metadata") or {}).get("next_cursor")
        if not cursor:
            break
        await asyncio.sleep(0.5)

    return messages


def _format_message(
    m: dict,
    names: dict[str, str],
) -> list[str]:
    """Return one or more lines representing a single message."""
    ts = datetime.fromtimestamp(float(m["ts"]), tz=timezone.utc).strftime("%Y-%m-%d %H:%M")
    author = names.get(m.get("user", ""), m.get("user", "unknown"))
    lines: list[str] = []

    text = (m.get("text") or "").strip()
    if text:
        lines.append(f"[{ts}] {author}: {text}")

    # file attachments (requires files:read scope)
    for f in m.get("files") or []:
        name = f.get("name") or f.get("id", "")
        title = f.get("title") or ""
        mimetype = f.get("mimetype") or ""
        url = f.get("permalink") or f.get("url_private") or ""
        lines.append(f"  [file] {title or name} ({mimetype}) {url}")

    # reactions (requires reactions:read scope)
    for r in m.get("reactions") or []:
        emoji = r.get("name", "")
        count = r.get("count", 0)
        lines.append(f"  [reaction] :{emoji}: x{count}")

    return lines


def _build_doc_content(
    channel_name: str,
    messages: list[dict],
    user_id: str,
    today: str,
    names: dict[str, str] | None = None,
) -> str | None:
    relevant = [
        m for m in messages
        if m.get("user") == user_id or user_id in (m.get("text") or "")
    ]
    if not relevant:
        return None

    resolved = names or {}
    lines = [f"# Slack channel: #{channel_name}  ({today})\n"]
    for m in sorted(relevant, key=lambda x: float(x.get("ts", 0))):
        lines.extend(_format_message(m, resolved))

    return "\n".join(lines) if len(lines) > 1 else None


async def _run_sweep(run_id: str, req: SlackSyncRequest, engine_url: str) -> None:
    today = datetime.now(timezone.utc).strftime("%Y-%m-%d")
    snapshot_id = req.snapshot_id or f"snap-{today}"
    oldest_ts = str((datetime.now(timezone.utc) - timedelta(days=req.days_back)).timestamp())

    _active_runs[run_id] = {"status": "running", "ingested": 0, "error": None}
    client = AsyncWebClient(token=req.user_token)

    try:
        resp = await client.conversations_list(
            types="public_channel,private_channel,im,mpim",
            exclude_archived=True,
            limit=200,
        )
        channels = [
            c for c in resp.get("channels", [])
            if c.get("is_member") or c.get("is_im") or c.get("is_mpim")
        ]
        logger.info("Slack sweep run={run_id} channels={n}", run_id=run_id, n=len(channels))

        # resolve all user IDs to display names upfront (requires users:read)
        all_messages_preview: list[dict] = []
        channel_data: list[tuple[str, str, list[dict]]] = []
        for ch in channels:
            cid = ch["id"]
            cname = ch.get("name") or cid
            messages = await _fetch_channel_messages(client, cid, oldest_ts, req.user_id)
            logger.info("Channel {cname} total_messages={n}", cname=cname, n=len(messages))
            if messages:
                channel_data.append((cid, cname, messages))
                all_messages_preview.extend(messages)

        unique_user_ids = {m.get("user") for m in all_messages_preview if m.get("user")}
        names = await _resolve_user_names(client, unique_user_ids)
        logger.info("Resolved {n} user names", n=len(names))

        ingested = 0
        async with httpx.AsyncClient(base_url=engine_url, timeout=30) as http:
            for _cid, cname, messages in channel_data:

                content = _build_doc_content(cname, messages, req.user_id, today, names)
                if not content:
                    senders = list({m.get("user") for m in messages if m.get("user")})
                    logger.info(
                        "No relevant messages in {cname} — senders={senders}",
                        cname=cname, senders=senders,
                    )
                    continue

                payload = {
                    "repo_id": req.repo_id,
                    "doc_path": f"slack/{cname}/{today}.md",
                    "content": content,
                    "commit_hash": snapshot_id,
                }
                ingest_resp = await http.post("/api/v1/ingest/", json=payload)
                ingest_resp.raise_for_status()
                ingested += 1
                logger.info("Ingested #{cname}", cname=cname)
                await asyncio.sleep(0.2)

        _active_runs[run_id]["ingested"] = ingested

        if ingested > 0:
            async with httpx.AsyncClient(base_url=engine_url, timeout=30) as http:
                embed_resp = await http.post("/api/v1/pipeline/snapshots/embed", json={
                    "repo_id": req.repo_id,
                    "snapshot_id": snapshot_id,
                    "commit_sha": today.replace("-", ""),
                })
                embed_resp.raise_for_status()
                embed_run_id = embed_resp.json().get("run_id")
                logger.info("Embedding queued embed_run_id={id}", id=embed_run_id)
                _active_runs[run_id]["embed_run_id"] = embed_run_id

        _active_runs[run_id]["status"] = "completed"

    except Exception as exc:
        logger.error("Slack sweep failed run={run_id} error={error}", run_id=run_id, error=str(exc))
        _active_runs[run_id]["status"] = "failed"
        _active_runs[run_id]["error"] = str(exc)


@router.post("/slack", response_model=SlackSyncResponse, status_code=202)
async def sync_slack(req: SlackSyncRequest, background_tasks: BackgroundTasks) -> SlackSyncResponse:
    """
    Sweep Slack messages for a user and index them into Qdrant.
    Returns immediately; poll GET /api/v1/sync/slack/{run_id} for status.
    """
    import uuid
    run_id = f"slack_sync_{uuid.uuid4().hex[:12]}"
    today = datetime.now(timezone.utc).strftime("%Y-%m-%d")
    snapshot_id = req.snapshot_id or f"snap-{today}"

    # call ourselves on the same host so we use the same ingest/pipeline routes
    engine_url = str(settings.EMBEDDING_ENGINE_SELF_URL)
    background_tasks.add_task(_run_sweep, run_id, req, engine_url)

    return SlackSyncResponse(
        run_id=run_id,
        status="queued",
        repo_id=req.repo_id,
        snapshot_id=snapshot_id,
    )


@router.get("/slack/{run_id}")
async def sync_slack_status(run_id: str) -> dict:
    """Check the status of a Slack sync run."""
    run = _active_runs.get(run_id)
    if not run:
        raise HTTPException(status_code=404, detail="Run not found")
    return {"run_id": run_id, **run}
