"""
Coba 長期記憶(階段 3f)— RAG 用的向量索引。

用 Qdrant 本地檔案模式(in-process,無需 server)+ FastEmbed 多語言 embedding。
階段 5 部署到 NAS 時可換成 Qdrant 獨立 server,介面不變。

職責:
  - 把每則新訊息向量化、索引到 Qdrant
  - 提供「給我跟某問題相關的 N 則過去訊息」的搜尋介面
  - 支援冷啟 backfill(把 SQLite 既有訊息一次補進去)

embedding 模型:
  - sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2
    → 384 維 / 多語言(含中文)/ ONNX 加速
"""

import asyncio
import logging
from pathlib import Path
from uuid import UUID

from qdrant_client import AsyncQdrantClient
from qdrant_client.http.models import (
    Distance,
    FieldCondition,
    Filter,
    MatchValue,
    PointStruct,
    VectorParams,
)
from sqlalchemy import select


logger = logging.getLogger("coba.memory")


# 路徑設定
QDRANT_DIR = Path(__file__).resolve().parent.parent.parent / "qdrant_data"
QDRANT_DIR.mkdir(exist_ok=True, parents=True)

COLLECTION = "messages"
EMBEDDING_MODEL = "sentence-transformers/paraphrase-multilingual-MiniLM-L12-v2"
VECTOR_DIM = 384


# 全域單例(整個 app 共用)
_client: AsyncQdrantClient | None = None
_embedding_model = None
_initialized: bool = False


def _get_embedding_model():
    """延遲載入 fastembed 模型(第一次用會下載 ~80MB)。"""
    global _embedding_model
    if _embedding_model is None:
        from fastembed import TextEmbedding
        logger.info("載入 embedding 模型 %s(第一次會下載)...", EMBEDDING_MODEL)
        _embedding_model = TextEmbedding(model_name=EMBEDDING_MODEL)
        logger.info("embedding 模型 ready")
    return _embedding_model


def _embed_sync(text: str) -> list[float]:
    """把文字轉成 384 維向量(同步,在 thread 跑)。"""
    model = _get_embedding_model()
    # FastEmbed 的 embed() 回傳 generator,取第一個
    return list(next(iter(model.embed([text]))))


async def embed_text(text: str) -> list[float]:
    """async wrapper:embedding 計算放 thread 不卡 event loop。"""
    return await asyncio.to_thread(_embed_sync, text)


async def get_client() -> AsyncQdrantClient:
    """取得 Qdrant client(本地檔案模式,自動建 collection)。"""
    global _client, _initialized
    if _client is None:
        _client = AsyncQdrantClient(path=str(QDRANT_DIR))
    if not _initialized:
        # 確認 collection 存在
        try:
            await _client.get_collection(COLLECTION)
        except Exception:
            await _client.create_collection(
                collection_name=COLLECTION,
                vectors_config=VectorParams(size=VECTOR_DIM, distance=Distance.COSINE),
            )
            logger.info("已建立 Qdrant collection: %s", COLLECTION)
        _initialized = True
    return _client


async def index_message(
    message_id: UUID,
    channel_id: UUID,
    user_id: UUID | None,
    content: str,
    created_at_iso: str,
    tags: list[str] | None = None,
) -> None:
    """把一則訊息加入向量索引。

    - 太短的訊息(< 5 字)跳過,索引也沒意義
    - 失敗靜默忽略(不影響對話主流程)
    """
    if len(content.strip()) < 5:
        return
    try:
        client = await get_client()
        vector = await embed_text(content)

        await client.upsert(
            collection_name=COLLECTION,
            points=[
                PointStruct(
                    id=str(message_id),
                    vector=vector,
                    payload={
                        "message_id": str(message_id),
                        "channel_id": str(channel_id),
                        "user_id": str(user_id) if user_id else None,
                        "content": content,
                        "created_at": created_at_iso,
                        "tags": tags or [],
                    },
                )
            ],
        )
    except Exception as e:
        logger.warning("index_message 失敗 msg=%s: %s", str(message_id)[:8], e)


async def search_relevant(query: str, limit: int = 5, channel_id: UUID | None = None) -> list[dict]:
    """根據 query 找 N 則最相關的過去訊息。

    回傳格式:[{message_id, content, user_id, created_at, score}, ...]
    """
    if len(query.strip()) < 3:
        return []
    try:
        client = await get_client()
        vector = await embed_text(query)

        # 可選:限定特定 channel
        flt = None
        if channel_id is not None:
            flt = Filter(
                must=[FieldCondition(key="channel_id", match=MatchValue(value=str(channel_id)))]
            )

        # Qdrant 1.7+ 換成 query_points,1.17 已移除 search()
        # 兩種 API 都試一遍,讓不同版本都能跑
        try:
            qr = await client.query_points(
                collection_name=COLLECTION,
                query=vector,
                query_filter=flt,
                limit=limit,
            )
            results = qr.points
        except AttributeError:
            results = await client.search(
                collection_name=COLLECTION,
                query_vector=vector,
                query_filter=flt,
                limit=limit,
            )
        return [
            {
                "message_id": p.payload.get("message_id"),
                "content": p.payload.get("content"),
                "user_id": p.payload.get("user_id"),
                "created_at": p.payload.get("created_at"),
                "tags": p.payload.get("tags", []),
                "score": p.score,
            }
            for p in results
        ]
    except Exception as e:
        logger.warning("search_relevant 失敗 query=%r: %s", query[:30], e)
        return []


async def backfill_existing_messages() -> int:
    """把 SQLite 既有訊息全部索引一次(冷啟用)。

    冪等:upsert 自動覆蓋既有 point,跑兩次沒事。
    回傳:成功索引的訊息數。
    """
    from app.database import AsyncSessionLocal
    from app.models.channel import Message

    count = 0
    async with AsyncSessionLocal() as db:
        msgs = list((await db.execute(
            select(Message).where(Message.deleted_at.is_(None))
        )).scalars().all())
        for m in msgs:
            if len(m.content.strip()) < 5:
                continue
            try:
                await index_message(
                    message_id=m.id,
                    channel_id=m.channel_id,
                    user_id=m.user_id,
                    content=m.content,
                    created_at_iso=m.created_at.isoformat() if m.created_at else "",
                    tags=m.tags or [],
                )
                count += 1
            except Exception:
                continue
    logger.info("backfill 完成,共索引 %d 則訊息", count)
    return count


def trigger_index(message_id: UUID, channel_id: UUID, user_id: UUID | None,
                  content: str, created_at_iso: str, tags: list[str] | None = None) -> None:
    """同步入口:背景索引,不 block HTTP 回應。"""
    asyncio.create_task(index_message(
        message_id=message_id,
        channel_id=channel_id,
        user_id=user_id,
        content=content,
        created_at_iso=created_at_iso,
        tags=tags,
    ))
