"""Mem0 托管云版 MemoryStore。

agent_id → mem0 的 user_id,实现 agent 级记忆隔离。我们的 importance/created_at/kind
塞进 metadata 随记忆一起存,检索时取回用于本地 recency/importance 加权。

注:mem0 默认对写入做 LLM 抽取(infer=True),可能改写 episodic 原文;若需逐字
存记忆流可设 infer=False(原文入库,但失去 mem0 的去重/自编辑)。该取舍待评估
(PRD §12 TODO)。
"""

from __future__ import annotations

from typing import Any

from genesis.memory.record import MemoryKind, MemoryRecord
from genesis.memory.store import MemoryStore, SearchHit


def _extract_id(res: Any) -> str:
    """从 mem0 add 的返回里尽力取第一条记忆 id(返回结构随版本而异)。"""
    items = res.get("results", res) if isinstance(res, dict) else res
    if isinstance(items, list) and items and isinstance(items[0], dict):
        return items[0].get("id", "") or ""
    return ""


def _to_record(item: dict, agent_id: str) -> MemoryRecord:
    md = item.get("metadata") or {}
    kind_val = md.get("kind", "observation")
    kind = MemoryKind(kind_val) if kind_val in MemoryKind._value2member_map_ else MemoryKind.OBSERVATION
    return MemoryRecord(
        agent_id=agent_id,
        content=item.get("memory", "") or "",
        created_at=float(md.get("created_at", 0.0)),
        kind=kind,
        importance=float(md.get("importance", 0.5)),
        id=item.get("id"),
    )


class Mem0Store(MemoryStore):
    def __init__(self, api_key: str, *, infer: bool = True) -> None:
        from mem0 import MemoryClient  # 延迟导入,避免无 key 环境下也强依赖

        self.client = MemoryClient(api_key=api_key)
        self.infer = infer

    def add(self, record: MemoryRecord) -> str:
        res = self.client.add(
            [{"role": "user", "content": record.content}],
            user_id=record.agent_id,
            metadata={
                "kind": record.kind.value,
                "importance": record.importance,
                "created_at": record.created_at,
            },
            infer=self.infer,
        )
        record.id = _extract_id(res)
        return record.id

    def search(self, query: str, agent_id: str, limit: int = 20) -> list[SearchHit]:
        # mem0 v2:实体参数必须走 filters,不能传顶层 user_id
        res = self.client.search(query, filters={"user_id": agent_id}, limit=limit)
        items = res.get("results", res) if isinstance(res, dict) else res
        hits = []
        for it in items or []:
            if not isinstance(it, dict):
                continue
            hits.append(
                SearchHit(record=_to_record(it, agent_id), relevance=float(it.get("score", 0.0)))
            )
        return hits

    def all(self, agent_id: str) -> list[MemoryRecord]:
        res = self.client.get_all(filters={"user_id": agent_id})
        items = res.get("results", res) if isinstance(res, dict) else res
        return [_to_record(it, agent_id) for it in (items or []) if isinstance(it, dict)]

    def clear(self, agent_id: str) -> None:
        """删除某 agent 的全部记忆(测试清理用)。"""
        self.client.delete_all(user_id=agent_id)
