"""
訊息 API(階段 3a)。

  - GET    /api/channels/{cid}/messages           ─ 列訊息(預設最新 50 則)
  - POST   /api/channels/{cid}/messages           ─ 發訊息
  - PATCH  /api/messages/{id}                     ─ 編輯訊息(只有作者)
  - DELETE /api/messages/{id}                     ─ 刪訊息(作者或 owner)
  - POST   /api/messages/{id}/reactions           ─ 加表情回應
  - DELETE /api/messages/{id}/reactions/{emoji}   ─ 收回自己的表情回應

階段 3a 簡化:不檢查使用者是否「在頻道裡」(public 頻道任何人都能讀寫)。
階段 3 後續會補上「private 頻道才需要 ChannelMember 才能讀寫」。
"""

from collections import defaultdict
from datetime import datetime, timezone
from uuid import UUID

from fastapi import APIRouter, HTTPException, status
from sqlalchemy import select
from sqlalchemy.orm import selectinload

from app.core.deps import CurrentUser, DBSession
from app.core.ws_manager import manager as ws_manager
from app.models.channel import Channel, ChannelMember, Message, MessageReaction
from app.models.user import User
from app.services.coba_chat import trigger_coba_if_mentioned
from app.services.coba_memory import trigger_index
from app.services.coba_tagger import trigger_tagging
from app.schemas.channel import (
    MessageCreateRequest,
    MessagePublic,
    MessageUpdateRequest,
    ReactionRequest,
)


def _ws_safe(payload: dict) -> dict:
    """把 datetime / UUID 轉成可序列化的字串(WebSocket 用 JSON 走)。"""
    import json
    from datetime import datetime
    from uuid import UUID as _UUID

    def default(o):
        if isinstance(o, datetime):
            return o.isoformat()
        if isinstance(o, _UUID):
            return str(o)
        raise TypeError(f"unserialisable {type(o)}")

    return json.loads(json.dumps(payload, default=default))


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


def _serialize_message(msg: Message, read_by: list | None = None) -> dict:
    """把 SQLAlchemy Message 轉成 MessagePublic 格式(含 reactions 彙整 + 已讀名單)。

    Args:
        msg: 訊息物件
        read_by: 已讀此訊息的使用者清單(不含作者),由 caller 計算後傳入
    """
    grouped = defaultdict(list)
    for r in msg.reactions:
        grouped[r.emoji].append(r.user_id)
    reactions = [
        {"emoji": emoji, "count": len(uids), "user_ids": uids}
        for emoji, uids in grouped.items()
    ]
    return {
        "id": msg.id,
        "channel_id": msg.channel_id,
        "user": {
            "id": msg.user.id,
            "display_name": msg.user.display_name,
            "avatar_url": msg.user.avatar_url,
        } if msg.user else None,
        "content": msg.content,
        "message_type": msg.message_type,
        "reply_to_id": msg.reply_to_id,
        "is_pinned": msg.is_pinned,
        "is_edited": msg.is_edited,
        "created_at": msg.created_at,
        "updated_at": msg.updated_at,
        "reactions": reactions,
        "read_by": read_by or [],
        "tags": msg.tags or [],
    }


async def _get_channel_read_states(channel_id: UUID, db) -> list[dict]:
    """取得頻道內所有成員的「最後讀取時間」+ 個人簡介。

    回傳格式:[{user_id, display_name, avatar_url, last_read_at}, ...]
    給 caller 用來計算每則訊息的 read_by。
    """
    stmt = (
        select(ChannelMember, User)
        .join(User, ChannelMember.user_id == User.id)
        .where(
            ChannelMember.channel_id == channel_id,
            ChannelMember.last_read_at.is_not(None),
            User.deleted_at.is_(None),
        )
    )
    rows = (await db.execute(stmt)).all()
    return [
        {
            "user_id": user.id,
            "display_name": user.display_name,
            "avatar_url": user.avatar_url,
            "last_read_at": member.last_read_at,
        }
        for member, user in rows
    ]


def _normalize_dt(dt):
    """把 datetime 統一成 timezone-aware UTC,避免 naive vs aware 比較炸鍋。

    背景:SQLite 用 ADD COLUMN 加進去的 TIMESTAMP 欄位讀出來是 naive(沒 tz),
    但 server_default=func.now() 建的欄位讀出來是 aware。直接比較會 TypeError。
    這個 helper 把兩種都當 UTC 處理(SQLite 內部本來就是 UTC 存的)。
    """
    if dt is None:
        return None
    if dt.tzinfo is None:
        return dt.replace(tzinfo=timezone.utc)
    return dt


def _build_read_by(msg: Message, member_reads: list[dict]) -> list[dict]:
    """根據訊息建立時間 + 成員 last_read_at,算出哪些人已讀此訊息(不含作者)。"""
    msg_created = _normalize_dt(msg.created_at)
    return [
        {
            "id": m["user_id"],
            "display_name": m["display_name"],
            "avatar_url": m["avatar_url"],
        }
        for m in member_reads
        if _normalize_dt(m["last_read_at"]) >= msg_created and m["user_id"] != msg.user_id
    ]


@router.get(
    "/api/channels/{channel_id}/messages",
    response_model=list[MessagePublic],
    summary="列出頻道訊息",
)
async def list_messages(
    channel_id: UUID,
    current_user: CurrentUser,
    db: DBSession,
    limit: int = 50,
    before: UUID | None = None,
):
    """列訊息,預設最新 50 則。

    傳 ?before=<message_id> 可以分頁載入更舊的(infinite scroll up)。
    """
    # 確認頻道存在
    ch = (await db.execute(
        select(Channel).where(Channel.id == channel_id, Channel.deleted_at.is_(None))
    )).scalar_one_or_none()
    if ch is None:
        raise HTTPException(status_code=404, detail="找不到此頻道")

    stmt = (
        select(Message)
        .where(Message.channel_id == channel_id, Message.deleted_at.is_(None))
        .order_by(Message.created_at.desc())
        .limit(min(limit, 200))
    )
    if before is not None:
        # 找到 before 那則的 created_at,只取更舊的
        cursor_msg = (await db.execute(
            select(Message).where(Message.id == before)
        )).scalar_one_or_none()
        if cursor_msg is not None:
            stmt = stmt.where(Message.created_at < cursor_msg.created_at)

    messages = list((await db.execute(stmt)).scalars().all())
    # 反轉:讓前端拿到時是「舊的在前、新的在後」(LINE 風格)
    messages.reverse()

    # 計算每則訊息的「已讀名單」
    member_reads = await _get_channel_read_states(channel_id, db)
    return [_serialize_message(m, _build_read_by(m, member_reads)) for m in messages]


@router.post(
    "/api/channels/{channel_id}/messages",
    response_model=MessagePublic,
    status_code=status.HTTP_201_CREATED,
    summary="發訊息",
)
async def create_message(
    channel_id: UUID,
    payload: MessageCreateRequest,
    current_user: CurrentUser,
    db: DBSession,
):
    # 確認頻道存在
    ch = (await db.execute(
        select(Channel).where(Channel.id == channel_id, Channel.deleted_at.is_(None))
    )).scalar_one_or_none()
    if ch is None:
        raise HTTPException(status_code=404, detail="找不到此頻道")

    msg = Message(
        channel_id=channel_id,
        user_id=current_user.id,
        content=payload.content,
        message_type=payload.message_type,
        reply_to_id=payload.reply_to_id,
    )
    db.add(msg)

    # 發訊息的人「順便」標記自己已讀(避免在自己訊息下顯示 unread)
    my_member = (await db.execute(
        select(ChannelMember).where(
            ChannelMember.channel_id == channel_id,
            ChannelMember.user_id == current_user.id,
        )
    )).scalar_one_or_none()
    if my_member is None:
        my_member = ChannelMember(channel_id=channel_id, user_id=current_user.id)
        db.add(my_member)
    my_member.last_read_at = datetime.now(timezone.utc)

    await db.commit()
    # 重抓含 user / reactions 的完整版本
    full = (await db.execute(
        select(Message)
        .options(selectinload(Message.reactions))
        .where(Message.id == msg.id)
    )).scalar_one()
    member_reads = await _get_channel_read_states(channel_id, db)
    serialized = _serialize_message(full, _build_read_by(full, member_reads))

    # 廣播給所有連線(包含發送者自己,前端會處理重複)
    await ws_manager.broadcast(_ws_safe({
        "type": "message_new",
        "channel_id": str(channel_id),
        "message": serialized,
    }))

    # 階段 3d / 7:偵測 @庫柏 + engagement window → 觸發 AI 回應
    trigger_coba_if_mentioned(channel_id, msg.id, payload.content, user_id=current_user.id)
    # 階段 3e:每則新訊息自動分類標籤(背景任務,Haiku)
    trigger_tagging(msg.id)
    # 階段 3f:索引到 Qdrant 給 RAG 用(背景任務)
    trigger_index(
        message_id=msg.id,
        channel_id=msg.channel_id,
        user_id=msg.user_id,
        content=msg.content,
        created_at_iso=msg.created_at.isoformat() if msg.created_at else "",
        tags=msg.tags,
    )

    return serialized


async def _get_message_or_404(message_id: UUID, db: DBSession) -> Message:
    stmt = (
        select(Message)
        .options(selectinload(Message.reactions))
        .where(Message.id == message_id, Message.deleted_at.is_(None))
    )
    msg = (await db.execute(stmt)).scalar_one_or_none()
    if msg is None:
        raise HTTPException(status_code=404, detail="找不到此訊息")
    return msg


@router.patch("/api/messages/{message_id}", response_model=MessagePublic, summary="編輯訊息")
async def update_message(
    message_id: UUID,
    payload: MessageUpdateRequest,
    current_user: CurrentUser,
    db: DBSession,
):
    msg = await _get_message_or_404(message_id, db)
    if msg.user_id != current_user.id:
        raise HTTPException(status_code=403, detail="只有作者能編輯自己的訊息")
    msg.content = payload.content
    msg.is_edited = True
    await db.commit()
    await db.refresh(msg)
    member_reads = await _get_channel_read_states(msg.channel_id, db)
    serialized = _serialize_message(msg, _build_read_by(msg, member_reads))
    await ws_manager.broadcast(_ws_safe({
        "type": "message_updated",
        "channel_id": str(msg.channel_id),
        "message": serialized,
    }))
    return serialized


@router.delete("/api/messages/{message_id}", status_code=204, summary="刪除訊息")
async def delete_message(
    message_id: UUID, current_user: CurrentUser, db: DBSession
) -> None:
    msg = await _get_message_or_404(message_id, db)
    # 作者或公司 owner 能刪
    if msg.user_id != current_user.id and current_user.role != "owner":
        raise HTTPException(status_code=403, detail="只有作者或 owner 能刪除訊息")
    msg.deleted_at = datetime.now(timezone.utc)
    await db.commit()
    await ws_manager.broadcast(_ws_safe({
        "type": "message_deleted",
        "channel_id": str(msg.channel_id),
        "message_id": str(msg.id),
    }))


# ---------------------------------------------------------------------------
# Reactions(表情回應)
# ---------------------------------------------------------------------------
@router.post(
    "/api/messages/{message_id}/reactions",
    response_model=MessagePublic,
    summary="對訊息加表情回應",
)
async def add_reaction(
    message_id: UUID,
    payload: ReactionRequest,
    current_user: CurrentUser,
    db: DBSession,
):
    msg = await _get_message_or_404(message_id, db)

    # 已存在就 idempotent(不重複加,直接回完整訊息)
    existing = (await db.execute(
        select(MessageReaction).where(
            MessageReaction.message_id == message_id,
            MessageReaction.user_id == current_user.id,
            MessageReaction.emoji == payload.emoji,
        )
    )).scalar_one_or_none()
    if existing is None:
        db.add(MessageReaction(
            message_id=message_id,
            user_id=current_user.id,
            emoji=payload.emoji,
        ))
        await db.commit()
        # 重抓完整訊息含新 reaction
        msg = await _get_message_or_404(message_id, db)

    member_reads = await _get_channel_read_states(msg.channel_id, db)
    serialized = _serialize_message(msg, _build_read_by(msg, member_reads))
    await ws_manager.broadcast(_ws_safe({
        "type": "reaction_changed",
        "channel_id": str(msg.channel_id),
        "message": serialized,
    }))
    return serialized


@router.delete(
    "/api/messages/{message_id}/reactions/{emoji}",
    response_model=MessagePublic,
    summary="收回自己的表情回應",
)
async def remove_reaction(
    message_id: UUID,
    emoji: str,
    current_user: CurrentUser,
    db: DBSession,
):
    existing = (await db.execute(
        select(MessageReaction).where(
            MessageReaction.message_id == message_id,
            MessageReaction.user_id == current_user.id,
            MessageReaction.emoji == emoji,
        )
    )).scalar_one_or_none()
    if existing:
        await db.delete(existing)
        await db.commit()

    msg = await _get_message_or_404(message_id, db)
    member_reads = await _get_channel_read_states(msg.channel_id, db)
    serialized = _serialize_message(msg, _build_read_by(msg, member_reads))
    await ws_manager.broadcast(_ws_safe({
        "type": "reaction_changed",
        "channel_id": str(msg.channel_id),
        "message": serialized,
    }))
    return serialized
