"""
memory/manager.py - 记忆管理器
================================
基于 SQLite 的双表记忆系统:
  - session_messages 表: 聊天记录（用户消息、助手回复、工具调用等）
  - memories 表: 记忆（全局记忆、工作记忆、用户偏好等）

表分离设计:
  - session_messages: 仅存储对话相关数据，按 session_id 隔离
  - memories: 存储跨会话的持久知识，可跨 session 检索
"""
from __future__ import annotations

import json
import math
import os
import re
import sqlite3
import threading
from collections import Counter
from dataclasses import dataclass, field, asdict
from datetime import datetime, timezone
from typing import Any, Dict, List, Optional, Tuple

from core.logger import get_logger
from core.utils import generate_id, timestamp, truncate_str, get_config_tz

logger = get_logger("myagent.memory")


# ==============================================================================
# 数据模型
# ==============================================================================

@dataclass
class SessionMessageEntry:
    """会话消息条目（聊天记录）"""
    id: str = field(default_factory=lambda: generate_id("msg"))
    session_id: str = ""
    agent_id: int = 1                  # 所属agent ID
    role: str = ""                     # 对话角色: user | assistant | system | tool
    content: str = ""                  # 消息内容
    key: str = ""                      # 特殊标记: tool_call, reasoning, llm_output 等
    metadata: Dict[str, Any] = field(default_factory=dict)
    created_at: str = field(default_factory=timestamp)
    updated_at: str = field(default_factory=timestamp)

    def to_dict(self) -> dict:
        d = asdict(self)
        d["metadata"] = json.dumps(self.metadata, ensure_ascii=False)
        return d

    @classmethod
    def from_row(cls, row: sqlite3.Row) -> "SessionMessageEntry":
        meta = row["metadata"]
        if isinstance(meta, str):
            try:
                meta = json.loads(meta)
            except json.JSONDecodeError:
                meta = {}
        return cls(
            id=row["id"],
            session_id=row["session_id"],
            agent_id=int(dict(row).get("agent_id", 1) or 1),
            role=row["role"],
            content=row["content"],
            key=row["key"],
            metadata=meta,
            created_at=row["created_at"],
            updated_at=row["updated_at"],
        )


@dataclass
class MemoryEntry:
    """记忆条目（全局记忆/长期记忆）"""
    id: str = field(default_factory=lambda: generate_id("mem"))
    session_id: str = ""              # 关联的 session_id（global 表示全局）
    agent_id: int = 1                  # 所属agent ID
    key: str = ""                      # 检索键/标签
    content: str = ""                  # 记忆内容
    summary: str = ""                  # 摘要
    metadata: Dict[str, Any] = field(default_factory=dict)
    importance: float = 0.5           # 重要性 0~1
    access_count: int = 0             # 访问次数
    created_at: str = field(default_factory=timestamp)
    updated_at: str = field(default_factory=timestamp)
    expires_at: str = ""              # 过期时间(空=永不过期)

    def to_dict(self) -> dict:
        d = asdict(self)
        d["metadata"] = json.dumps(self.metadata, ensure_ascii=False)
        return d

    @classmethod
    def from_row(cls, row: sqlite3.Row) -> "MemoryEntry":
        meta = row["metadata"]
        if isinstance(meta, str):
            try:
                meta = json.loads(meta)
            except json.JSONDecodeError:
                meta = {}
        return cls(
            id=row["id"],
            session_id=row["session_id"],
            agent_id=int(dict(row).get("agent_id", 1) or 1),
            key=row["key"],
            content=row["content"],
            summary=row["summary"],
            metadata=meta,
            importance=row["importance"],
            access_count=row["access_count"],
            created_at=row["created_at"],
            updated_at=row["updated_at"],
            expires_at=dict(row).get("expires_at", ""),
        )


# ==============================================================================
# 记忆管理器
# ==============================================================================

class MemoryManager:
    """
    双层记忆管理器。

    记忆层级:
      - session (会话记忆): 对话上下文 + 任务进度，按 session_id 隔离
      - global  (全局记忆): 跨会话持久知识，所有 session 可检索

    使用示例:
        mm = MemoryManager(db_path="~/.myagent/data/memory.db")
        mm.initialize()

        mm.add_session("s1", role="user", content="帮我创建一个Python项目")
        mm.add_global(content="偏好使用Python, TypeScript")

        # 查询对话历史
        history = mm.get_conversation("s1")

        # 语义搜索
        results = mm.search("Python项目", category="global")
    """

    def __init__(self, db_path: str = ""):
        self.db_path = db_path
        self._local = threading.local()
        self._initialized = False

    def _get_conn(self) -> sqlite3.Connection:
        """获取线程本地的数据库连接"""
        if not hasattr(self._local, "conn") or self._local.conn is None:
            # 确保数据库目录存在
            db_parent = os.path.dirname(self.db_path)
            if db_parent:
                os.makedirs(db_parent, exist_ok=True)
            self._local.conn = sqlite3.connect(
                self.db_path,
                check_same_thread=False,
                timeout=10,
            )
            self._local.conn.row_factory = sqlite3.Row
            self._local.conn.execute("PRAGMA journal_mode=WAL")
            self._local.conn.execute("PRAGMA synchronous=NORMAL")
        return self._local.conn

    def initialize(self):
        """初始化数据库表结构（全新版本，无迁移逻辑）"""
        conn = self._get_conn()

        # 创建 agents 表（agent 注册表）
        conn.executescript("""
            CREATE TABLE IF NOT EXISTS agents (
                id INTEGER PRIMARY KEY AUTOINCREMENT,
                name TEXT UNIQUE NOT NULL,
                created_at TEXT NOT NULL
            );
            CREATE INDEX IF NOT EXISTS idx_agent_name ON agents(name);
        """)

        # 确保 default agent 存在
        from datetime import datetime as _dt
        from core.utils import get_config_tz
        now = _dt.now(get_config_tz()).strftime("%Y-%m-%d %H:%M:%S")
        conn.execute("INSERT OR IGNORE INTO agents (id, name, created_at) VALUES (1, '1', ?)", (now,))

        # 创建 session_messages 表（聊天记录）
        conn.executescript("""
            CREATE TABLE IF NOT EXISTS session_messages (
                id TEXT PRIMARY KEY,
                session_id TEXT NOT NULL,
                agent_id INTEGER NOT NULL DEFAULT 1,
                role TEXT DEFAULT '',
                content TEXT DEFAULT '',
                key TEXT DEFAULT '',
                metadata TEXT DEFAULT '{}',
                created_at TEXT NOT NULL,
                updated_at TEXT NOT NULL
            );
            CREATE INDEX IF NOT EXISTS idx_session_msg ON session_messages(session_id);
            CREATE INDEX IF NOT EXISTS idx_agent_msg ON session_messages(agent_id);
            CREATE INDEX IF NOT EXISTS idx_session_created ON session_messages(session_id, created_at);
        """)

        # 创建 memories 表（全局记忆/长期记忆）
        conn.executescript("""
            CREATE TABLE IF NOT EXISTS memories (
                id TEXT PRIMARY KEY,
                session_id TEXT NOT NULL,
                agent_id INTEGER NOT NULL DEFAULT 1,
                key TEXT DEFAULT '',
                content TEXT DEFAULT '',
                summary TEXT DEFAULT '',
                metadata TEXT DEFAULT '{}',
                importance REAL DEFAULT 0.5,
                access_count INTEGER DEFAULT 0,
                created_at TEXT NOT NULL,
                updated_at TEXT NOT NULL,
                expires_at TEXT DEFAULT ''
            );
            CREATE INDEX IF NOT EXISTS idx_memory_session ON memories(session_id);
            CREATE INDEX IF NOT EXISTS idx_memory_key ON memories(key);
            CREATE INDEX IF NOT EXISTS idx_memory_importance ON memories(importance DESC);
            CREATE INDEX IF NOT EXISTS idx_memory_created ON memories(created_at);
        """)

        # 创建 session_names 表
        conn.executescript("""
            CREATE TABLE IF NOT EXISTS session_names (
                session_id TEXT PRIMARY KEY,
                display_name TEXT NOT NULL DEFAULT '',
                updated_at TEXT NOT NULL
            );
        """)

        conn.commit()
        self._initialized = True
        logger.info(f"记忆系统已初始化 (db={self.db_path})")
    def get_agent_id(self, agent_name: str) -> int:
        """获取 agent 的数字 ID，如果不存在则自动注册"""
        conn = self._get_conn()
        try:
            cursor = conn.execute("SELECT id FROM agents WHERE name = ?", (agent_name,))
            row = cursor.fetchone()
            if row:
                return row["id"]
            else:
                # 自动注册新 agent
                from datetime import datetime as _dt
                from core.utils import get_config_tz
                now = _dt.now(get_config_tz()).strftime("%Y-%m-%d %H:%M:%S")
                cursor = conn.execute("INSERT INTO agents (name, created_at) VALUES (?, ?)", (agent_name, now))
                conn.commit()
                return cursor.lastrowid
        except Exception as e:
            logger.error(f"获取 agent ID 失败: {e}")
            # 返回 1 作为兜底（全权Agent），避免 agent_id=0 导致数据归属不明
            return 1

    def register_agent(self, agent_name: str) -> int:
        """注册新 agent，返回分配的 ID"""
        return self.get_agent_id(agent_name)

    def close(self):
        """关闭数据库连接"""
        if hasattr(self._local, "conn") and self._local.conn:
            self._local.conn.close()
            self._local.conn = None

    # ==========================================================================
    # 基础 CRUD
    # ==========================================================================

    def _insert_session_message(self, entry: SessionMessageEntry) -> str:
        """插入会话消息"""
        conn = self._get_conn()
        entry.updated_at = timestamp()
        try:
            conn.execute(
                """INSERT OR REPLACE INTO session_messages
                   (id, session_id, agent_id, role, content, key, metadata, created_at, updated_at)
                   VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)""",
                (
                    entry.id, entry.session_id, int(entry.agent_id), entry.role,
                    entry.content, entry.key,
                    json.dumps(entry.metadata, ensure_ascii=False),
                    entry.created_at, entry.updated_at,
                ),
            )
            conn.commit()
            return entry.id
        except Exception as e:
            logger.error(f"会话消息写入失败: {e}")
            raise

    def _insert_memory(self, entry: MemoryEntry) -> str:
        """插入记忆条目"""
        conn = self._get_conn()
        entry.updated_at = timestamp()
        try:
            conn.execute(
                """INSERT OR REPLACE INTO memories
                   (id, session_id, agent_id, key, content, summary, metadata, importance, access_count, created_at, updated_at, expires_at)
                   VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
                (
                    entry.id, entry.session_id, int(entry.agent_id), entry.key,
                    entry.content, entry.summary,
                    json.dumps(entry.metadata, ensure_ascii=False),
                    entry.importance, entry.access_count,
                    entry.created_at, entry.updated_at, entry.expires_at,
                ),
            )
            conn.commit()
            return entry.id
        except Exception as e:
            logger.error(f"记忆写入失败: {e}")
            raise

    def _query_session_messages(
        self,
        session_id: str = "",
        role: str = "",
        key: str = "",
        limit: int = 100,
        order_by: str = "created_at ASC",
    ) -> List[SessionMessageEntry]:
        """查询会话消息"""
        conn = self._get_conn()
        conditions = []
        params: list = []

        if session_id:
            conditions.append("session_id = ?")
            params.append(session_id)
        if role:
            conditions.append("role = ?")
            params.append(role)
        if key:
            conditions.append("key = ?")
            params.append(key)

        where = " AND ".join(conditions) if conditions else "1=1"
        # 安全白名单
        _allowed_orders = {
            "created_at ASC", "created_at DESC",
            "updated_at ASC", "updated_at DESC",
        }
        _safe_order = order_by if order_by in _allowed_orders else "created_at ASC"
        sql = f"SELECT * FROM session_messages WHERE {where} ORDER BY {_safe_order} LIMIT ?"
        params.append(limit)

        rows = conn.execute(sql, params).fetchall()
        return [SessionMessageEntry.from_row(row) for row in rows]

    def _query_memories(
        self,
        session_id: str = "",
        key: str = "",
        limit: int = 100,
        order_by: str = "created_at DESC",
    ) -> List[MemoryEntry]:
        """查询记忆条目"""
        conn = self._get_conn()
        conditions = []
        params: list = []

        if session_id:
            conditions.append("session_id = ?")
            params.append(session_id)
        if key:
            conditions.append("key = ?")
            params.append(key)

        # 过滤已过期的
        conditions.append("(expires_at = '' OR expires_at > ?)")
        params.append(timestamp())

        where = " AND ".join(conditions) if conditions else "1=1"
        # 安全白名单
        _allowed_orders = {
            "created_at ASC", "created_at DESC",
            "importance ASC", "importance DESC",
            "importance DESC, created_at DESC",
            "access_count ASC", "access_count DESC",
            "updated_at ASC", "updated_at DESC",
        }
        _safe_order = order_by if order_by in _allowed_orders else "created_at DESC"
        sql = f"SELECT * FROM memories WHERE {where} ORDER BY {_safe_order} LIMIT ?"
        params.append(limit)

        rows = conn.execute(sql, params).fetchall()
        return [MemoryEntry.from_row(row) for row in rows]

    def _delete_session_messages(self, session_id: str, older_than: str = "") -> int:
        """删除会话消息，返回删除数量"""
        conn = self._get_conn()
        conditions = ["session_id = ?"]
        params: list = [session_id]

        if older_than:
            conditions.append("created_at < ?")
            params.append(older_than)

        where = " AND ".join(conditions)
        cursor = conn.execute(f"DELETE FROM session_messages WHERE {where}", params)
        conn.commit()
        return cursor.rowcount

    def _delete_memories(self, session_id: str, older_than: str = "") -> int:
        """删除记忆条目，返回删除数量"""
        conn = self._get_conn()
        conditions = ["session_id = ?"]
        params: list = [session_id]

        if older_than:
            conditions.append("created_at < ?")
            params.append(older_than)

        where = " AND ".join(conditions)
        cursor = conn.execute(f"DELETE FROM memories WHERE {where}", params)
        conn.commit()
        return cursor.rowcount

    def _update_session_message(self, message_id: str, content: str, **updates):
        """更新会话消息内容"""
        conn = self._get_conn()
        sets = ["content = ?", "updated_at = ?"]
        params: list = [content, timestamp()]

        for key, val in updates.items():
            if key in ("key", "metadata"):
                sets.append(f"{key} = ?")
                if key == "metadata":
                    params.append(json.dumps(val, ensure_ascii=False))
                else:
                    params.append(val)

        params.append(message_id)
        conn.execute(f"UPDATE session_messages SET {', '.join(sets)} WHERE id = ?", params)
        conn.commit()

    def _update_memory(self, memory_id: str, content: str, **updates):
        """更新记忆内容"""
        conn = self._get_conn()
        sets = ["content = ?", "updated_at = ?"]
        params: list = [content, timestamp()]

        for key, val in updates.items():
            if key in ("summary", "key", "importance", "metadata", "access_count", "expires_at"):
                sets.append(f"{key} = ?")
                if key == "metadata":
                    params.append(json.dumps(val, ensure_ascii=False))
                else:
                    params.append(val)

        params.append(memory_id)
        conn.execute(f"UPDATE memories SET {', '.join(sets)} WHERE id = ?", params)
        conn.commit()

    # ==========================================================================
    # 会话记忆 (session) — 对话上下文 + 任务进度
    # ==========================================================================

    def add_session(self, session_id, agent_id=1, role="", content="", key="", importance=0.5, metadata=None) -> str:
        """添加会话消息（聊天记录）。内容不包含时间前缀，时间仅存于 created_at 和 metadata。"""
        from datetime import datetime as _dt
        from core.utils import get_config_tz
        _now_str = _dt.now(get_config_tz()).strftime("%Y-%m-%d %H:%M:%S")
        # 直接存储原始内容，不再注入时间前缀
        entry = SessionMessageEntry(
            session_id=session_id, agent_id=agent_id, role=role,
            content=truncate_str(content, 10000), key=key,
            metadata={"timestamp": _now_str, **(metadata or {})},
        )
        return self._insert_session_message(entry)

    def get_conversation(self, session_id, limit=500, include_roles=None, agent_id=None) -> List[SessionMessageEntry]:
        """获取对话历史（role 非空的条目），按时间正序排列。

        只排除纯内部审计条目（llm_output 原始LLM输出、conversation_insight 记忆提炼）。
        tool_call 和 tool_result 会返回给前端以展示完整的工具调用过程。

        [FIX] 新增 agent_id 参数：当指定时，仅返回该 agent 的消息，
        防止跨 Agent 消息泄漏（如平台 bot 共享 session_id 的场景）。
        """
        conn = self._get_conn()
        # 只排除纯内部审计条目，保留 tool_call/tool_result 供前端展示
        # [v1.20.2] 使用 rowid 作为二级排序，防止同一秒内的消息顺序不确定
        if agent_id is not None:
            sql = """SELECT * FROM session_messages
                     WHERE session_id = ? AND role != ''
                     AND key NOT IN ('llm_output', 'llm_input', 'tool_result_raw', 'conversation_insight')
                     AND agent_id = ?
                     ORDER BY created_at ASC, rowid ASC LIMIT ?"""
            rows = conn.execute(sql, (session_id, int(agent_id), limit)).fetchall()
        else:
            sql = """SELECT * FROM session_messages
                     WHERE session_id = ? AND role != ''
                     AND key NOT IN ('llm_output', 'llm_input', 'tool_result_raw', 'conversation_insight')
                     ORDER BY created_at ASC, rowid ASC LIMIT ?"""
            rows = conn.execute(sql, (session_id, limit)).fetchall()
        entries = [SessionMessageEntry.from_row(row) for row in rows]
        if include_roles:
            entries = [e for e in entries if e.role in include_roles]
        return entries

    def get_conversation_all(self, session_id, limit=5000) -> List[SessionMessageEntry]:
        """获取全量对话历史（包含所有内部条目），用于完整回溯。limit=0 表示无限制。"""
        conn = self._get_conn()
        if limit and limit > 0:
            sql = """SELECT * FROM session_messages
                     WHERE session_id = ? AND role != ''
                     ORDER BY created_at ASC, rowid ASC LIMIT ?"""
            rows = conn.execute(sql, (session_id, limit)).fetchall()
        else:
            sql = """SELECT * FROM session_messages
                     WHERE session_id = ? AND role != ''
                     ORDER BY created_at ASC, rowid ASC"""
            rows = conn.execute(sql, (session_id,)).fetchall()
        return [SessionMessageEntry.from_row(row) for row in rows]

    def search_by_time_range(
        self,
        session_id: str = "",
        start_time: str = "",
        end_time: str = "",
        keyword: str = "",
        limit: int = 10,
    ) -> List[MemoryEntry]:
        """
        按时间范围 + 关键词搜索记忆（仅搜索 memories 表）。

        Args:
            session_id: 会话 ID（空=跨会话）
            start_time: 起始时间 ISO 格式（如 "2025-01-01 00:00:00"），空=不限
            end_time: 截止时间 ISO 格式，空=不限
            keyword: 关键词过滤（LIKE 模糊匹配），空=不限
            limit: 返回数量

        Returns:
            匹配的记忆条目列表（按时间倒序）
        """
        conn = self._get_conn()
        conditions = ["(expires_at = '' OR expires_at > ?)"]
        params: list = [timestamp()]

        if session_id:
            conditions.append("session_id = ?")
            params.append(session_id)
        if start_time:
            conditions.append("created_at >= ?")
            params.append(start_time)
        if end_time:
            conditions.append("created_at <= ?")
            params.append(end_time)
        if keyword:
            conditions.append("(content LIKE ? OR summary LIKE ? OR key LIKE ?)")
            like_pattern = f"%{keyword}%"
            params.extend([like_pattern, like_pattern, like_pattern])

        where = " AND ".join(conditions)
        sql = f"SELECT * FROM memories WHERE {where} ORDER BY created_at DESC LIMIT ?"
        params.append(limit)
        rows = conn.execute(sql, params).fetchall()
        return [MemoryEntry.from_row(row) for row in rows]

    def get_conversation_text(
        self,
        session_id: str,
        limit: int = 50,
    ) -> str:
        """获取对话历史文本(供 LLM 使用)，临时合并时间信息"""
        entries = self.get_conversation(session_id, limit)
        lines = []
        for e in entries:
            label = e.role.upper()
            if e.role == "user":
                label = "用户"
            elif e.role == "assistant":
                label = "助手"
            elif e.role == "system":
                label = "系统"
            elif e.role == "tool":
                label = "工具"
            # 从 created_at 提取时间，临时合并到内容中给 LLM
            time_str = e.created_at[:19] if e.created_at and len(e.created_at) >= 19 else ""
            if time_str:
                lines.append(f"[{label}] [{time_str}] {e.content}")
            else:
                lines.append(f"[{label}] {e.content}")
        return "\n".join(lines)

    def clear_conversation(self, session_id) -> int:
        """清空会话对话历史"""
        return self._delete_session_messages(session_id)

    def delete_session(self, session_id: str) -> int:
        """删除会话的所有记忆（对话、工作记忆、长期记忆等），彻底删除该会话"""
        # 删除会话消息
        msg_count = self._delete_session_messages(session_id)
        # 删除记忆
        mem_count = self._delete_memories(session_id)
        count = msg_count + mem_count
        # 同时删除 session_names 中的记录
        try:
            conn = self._get_conn()
            conn.execute("DELETE FROM session_names WHERE session_id = ?", (session_id,))
            conn.commit()
        except Exception:
            pass  # session_names 表可能不存在（首次使用）
        if count > 0:
            logger.info(f"删除会话: {session_id} ({count} 条记录)")
        return count

    def rename_session(self, session_id: str, new_name: str) -> bool:
        """重命名会话（设置显示名称别名）"""
        conn = self._get_conn()
        # 确保 session_names 表存在
        conn.execute("""
            CREATE TABLE IF NOT EXISTS session_names (
                session_id TEXT PRIMARY KEY,
                display_name TEXT NOT NULL,
                updated_at TEXT NOT NULL
            )
        """)
        conn.execute(
            "INSERT OR REPLACE INTO session_names (session_id, display_name, updated_at) VALUES (?, ?, ?)",
            (session_id, new_name, timestamp()),
        )
        conn.commit()
        logger.info(f"重命名会话: {session_id} -> {new_name}")
        return True

    def get_session_name(self, session_id: str) -> str:
        """获取会话的自定义显示名称，如果没有则返回空字符串"""
        try:
            conn = self._get_conn()
            row = conn.execute(
                "SELECT display_name FROM session_names WHERE session_id = ?",
                (session_id,),
            ).fetchone()
            return row["display_name"] if row else ""
        except Exception:
            return ""

    def list_session_names(self, session_ids: list) -> dict:
        """批量获取会话的自定义名称映射 {session_id: display_name}"""
        if not session_ids:
            return {}
        try:
            conn = self._get_conn()
            placeholders = ",".join("?" * len(session_ids))
            rows = conn.execute(
                f"SELECT session_id, display_name FROM session_names WHERE session_id IN ({placeholders})",
                session_ids,
            ).fetchall()
            return {r["session_id"]: r["display_name"] for r in rows}
        except Exception:
            return {}

    def prune_conversation(self, session_id: str, max_messages: int = 50) -> int:
        """修剪对话历史，保留最近 N 条"""
        entries = self.get_conversation(session_id, limit=1000)
        if len(entries) <= max_messages:
            return 0
        # 删除最旧的
        to_remove = entries[:-max_messages]
        conn = self._get_conn()
        ids = [e.id for e in to_remove]
        placeholders = ",".join("?" * len(ids))
        cursor = conn.execute(f"DELETE FROM session_messages WHERE id IN ({placeholders})", ids)
        conn.commit()
        logger.debug(f"修剪对话历史: 删除 {len(ids)} 条 (session={session_id})")
        return len(ids)

    # ==========================================================================
    # 全局记忆 (global) — 跨会话持久知识
    # ==========================================================================

    def add_global(self, session_id="global", key="", content="", summary="", importance=0.7, metadata=None) -> str:
        """添加全局记忆（跨会话可检索，保存到 memories 表）"""
        from datetime import datetime
        from core.utils import get_config_tz
        now_str = datetime.now(get_config_tz()).strftime("%Y-%m-%d %H:%M:%S")
        ts_summary = summary or truncate_str(content, 200)
        entry = MemoryEntry(
            session_id=session_id, key=key,
            content=truncate_str(content, 50000), summary=f"[{now_str}] {ts_summary}",
            importance=importance, metadata={"timestamp": now_str, **(metadata or {})},
        )
        return self._insert_memory(entry)

    def add_working_memory(self, session_id: str, key: str, content: str, importance: float = 0.6, metadata=None) -> str:
        """添加工作记忆（任务进度，保存到 memories 表）"""
        from datetime import datetime
        from core.utils import get_config_tz
        now_str = datetime.now(get_config_tz()).strftime("%Y-%m-%d %H:%M:%S")
        entry = MemoryEntry(
            session_id=session_id, key=key,
            content=truncate_str(content, 50000),
            summary=truncate_str(content, 200),
            importance=importance,
            metadata={"timestamp": now_str, **(metadata or {})},
        )
        return self._insert_memory(entry)

    def find_duplicate_memory(
        self,
        content: str,
        session_id: str = "",
        key: str = "",
        similarity_threshold: float = 0.85,
        category: str = "global",
    ) -> Optional[MemoryEntry]:
        """
        查找与给定内容高度相似的已有记忆（查重）。

        使用 TF-IDF 余弦相似度判断新内容是否与已有记忆高度重复。
        相似度超过阈值则返回匹配的旧记忆，否则返回 None。

        Args:
            content: 待检查的新记忆内容
            session_id: 会话 ID（为空则跨会话查重）
            key: 记忆分类键（为空则不限制分类）
            similarity_threshold: 相似度阈值，超过此值视为重复（默认 0.85）
            category: 记忆类别（默认 "global"，可设为 "session" 用于会话记忆查重，但已废弃）

        Returns:
            匹配的旧 MemoryEntry，或 None（无重复）
        """
        if not content or not content.strip():
            return None  # 空内容不查重

        # 注意：find_duplicate_memory 仅用于全局记忆查重，现在 memories 表只存储全局记忆
        # 如果 category 传入 "session"，则改为在 memories 表中按 session_id 查重
        conn = self._get_conn()
        conditions = []
        params: list = []

        if session_id:
            conditions.append("session_id = ?")
            params.append(session_id)
        if key:
            conditions.append("key = ?")
            params.append(key)

        # 过滤过期
        conditions.append("(expires_at = '' OR expires_at > ?)")
        params.append(timestamp())

        where = " AND ".join(conditions) if conditions else "1=1"
        # 取最近的记忆候选进行比对
        sql = f"SELECT * FROM memories WHERE {where} ORDER BY created_at DESC LIMIT 50"
        rows = conn.execute(sql, params).fetchall()

        if not rows:
            return None

        # 使用 TF-IDF 计算新内容与每条已有记忆的相似度
        documents = [(row["id"], row["content"]) for row in rows]
        scores = self._compute_tfidf(content, documents)

        if not scores:
            return None

        max_id = max(scores, key=scores.get)
        max_score = scores[max_id]
        if max_score >= similarity_threshold:
            logger.debug(
                f"记忆查重: 发现相似内容 (相似度={max_score:.2f}, "
                f"阈值={similarity_threshold}, 匹配ID={max_id})"
            )
            # 返回匹配的旧记忆
            for row in rows:
                if row["id"] == max_id:
                    return MemoryEntry.from_row(row)

        return None

    def is_duplicate_memory(
        self,
        content: str,
        session_id: str = "",
        key: str = "",
        similarity_threshold: float = 0.85,
    ) -> bool:
        """
        检查是否已存在高度相似的记忆（查重，仅检查 memories 表）。

        Returns:
            True 表示存在重复记忆，False 表示可以存储
        """
        return self.find_duplicate_memory(
            content, session_id, key, similarity_threshold
        ) is not None

    def update_memory(self, memory_id: str, content: str, summary: str = "", **updates) -> bool:
        """
        更新记忆内容（保留原 ID 和创建时间）。

        Args:
            memory_id: 记忆 ID
            content: 新内容
            summary: 新摘要（为空则自动截取）

        Returns:
            True 更新成功
        """
        if not summary:
            summary = truncate_str(content, 200)
        self._update_memory(memory_id, content, summary=summary, **updates)
        logger.info(f"记忆已更新: {memory_id}")
        return True

    def delete_memory(self, memory_id: str) -> bool:
        """
        删除单条记忆（从 memories 表）。

        Args:
            memory_id: 记忆 ID

        Returns:
            True 删除成功
        """
        conn = self._get_conn()
        cursor = conn.execute("DELETE FROM memories WHERE id = ?", (memory_id,))
        conn.commit()
        if cursor.rowcount > 0:
            logger.info(f"记忆已删除: {memory_id}")
            return True
        return False

    def get_global_memories(self, session_id="global", key="", limit=50) -> List[MemoryEntry]:
        """获取全局记忆（从 memories 表）"""
        return self._query_memories(session_id=session_id, key=key, limit=limit, order_by="importance DESC, created_at DESC")

    def get_preferences(self, session_id: str = "global") -> List[MemoryEntry]:
        """获取用户偏好（从 memories 表）"""
        return self.get_global_memories(session_id, key="user_pref")

    def get_experience(self, session_id: str = "global") -> List[MemoryEntry]:
        """获取技能经验（从 memories 表）"""
        return self.get_global_memories(session_id, key="skill_experience")

    def get_task_summaries(self, session_id: str = "global") -> List[MemoryEntry]:
        """获取历史任务总结（从 memories 表）"""
        return self.get_global_memories(session_id, key="task_summary")

    # ==========================================================================
    # 记忆搜索
    # ==========================================================================

    # ==========================================================================
    # 时间衰减权重 (遗忘曲线)
    # ==========================================================================

    @staticmethod
    def _compute_time_weight(created_at: str, half_life_days: float = 30.0) -> float:
        """
        计算记忆的时间衰减权重（模拟遗忘曲线）。

        使用指数衰减模型:  weight = e^(-ln(2) * age_days / half_life_days)
        - 刚创建的记忆权重为 1.0
        - 经过 half_life_days 后权重降至 0.5
        - 经过 2 * half_life_days 后权重降至 0.25
        - 权重永远不会降到 0（长期记忆仍可被召回）

        Args:
            created_at: 记忆创建时间（ISO 8601 格式，如 "2025-01-15T08:30:00+00:00"）
            half_life_days: 半衰期天数（默认 30 天）
                - 7 天: 适合短期会话记忆，快速遗忘
                - 30 天: 默认值，平衡近期与长期记忆
                - 90 天: 适合知识型场景，长期保留

        Returns:
            时间权重值 (0.0 ~ 1.0)

        Examples:
            >>> _compute_time_weight("2025-04-12T00:00:00+00:00", half_life_days=30)
            1.0  # 今天的记忆
            >>> _compute_time_weight("2025-03-13T00:00:00+00:00", half_life_days=30)
            0.5  # 30天前的记忆
            >>> _compute_time_weight("2025-02-11T00:00:00+00:00", half_life_days=30)
            0.25  # 60天前的记忆
        """
        if not created_at:
            return 0.5  # 无时间信息的记忆给中等权重

        try:
            # 解析 ISO 8601 时间戳
            created_dt = datetime.fromisoformat(created_at.replace("Z", "+00:00"))
            now_dt = datetime.now(get_config_tz())

            age_seconds = (now_dt - created_dt).total_seconds()
            if age_seconds <= 0:
                return 1.0  # 未来的时间戳或零延迟

            age_days = age_seconds / 86400.0

            # 指数衰减: e^(-ln(2) * age / half_life)
            decay = math.exp(-0.693147 * age_days / half_life_days)

            # 限制最低权重为 0.05，确保极老的记忆仍有一丝被召回的可能
            return max(decay, 0.05)
        except (ValueError, TypeError, OSError) as e:
            logger.debug(f"时间衰减计算失败 ({created_at}): {e}")
            return 0.5  # 解析失败时给中等权重

    # ==========================================================================
    # TF-IDF 语义搜索 (无外部依赖)
    # ==========================================================================

    @staticmethod
    def _tokenize(text: str) -> List[str]:
        """
        中文分词（简单实现：单字+双字组合）+ 英文词提取。

        对中文文本提取单字和相邻双字作为 token，
        对英文文本按空格和标点分词并转小写。
        """
        if not text:
            return []

        tokens: List[str] = []
        # 提取中文连续片段
        chinese_segments = re.findall(r'[\u4e00-\u9fff]+', text)
        for seg in chinese_segments:
            # 单字
            tokens.extend(list(seg))
            # 双字组合
            for i in range(len(seg) - 1):
                tokens.append(seg[i:i + 2])

        # 提取英文/数字单词
        english_words = re.findall(r'[a-zA-Z0-9]+', text.lower())
        tokens.extend(english_words)

        return tokens

    @staticmethod
    def _compute_tf(tokens: List[str]) -> Counter:
        """计算词频 (Term Frequency)"""
        return Counter(tokens)

    @classmethod
    def _compute_tfidf(
        cls,
        query: str,
        documents: List[Tuple[str, str]],  # [(id, text), ...]
    ) -> Dict[str, float]:
        """
        计算 TF-IDF 相似度得分。

        Args:
            query: 查询文本
            documents: 文档列表 [(doc_id, text), ...]

        Returns:
            {doc_id: tfidf_score} 按得分降序排列
        """
        if not documents or not query:
            return {}

        query_tokens = cls._tokenize(query)
        if not query_tokens:
            return {}

        # 文档数量
        n_docs = len(documents)

        # 计算每个文档的 TF
        doc_tfs: Dict[str, Counter] = {}
        for doc_id, text in documents:
            doc_tfs[doc_id] = cls._compute_tf(cls._tokenize(text))

        # 计算 IDF: log(N / (1 + df))  df=包含该词的文档数
        doc_freq: Counter = Counter()
        for doc_id, tf in doc_tfs.items():
            for token in set(tf.keys()):
                doc_freq[token] += 1

        idf: Dict[str, float] = {}
        for token, df in doc_freq.items():
            idf[token] = math.log((n_docs + 1) / (1 + df)) + 1

        # 查询的 TF
        query_tf = cls._compute_tf(query_tokens)

        # 计算每个文档与查询的余弦相似度
        scores: Dict[str, float] = {}
        all_tokens = set(query_tf.keys())
        for doc_id, tf in doc_tfs.items():
            all_tokens.update(tf.keys())

        # 计算查询向量的模
        query_norm = math.sqrt(
            sum((query_tf.get(t, 0) * idf.get(t, 0)) ** 2 for t in query_tf)
        )
        if query_norm == 0:
            return {}

        # 计算每个文档与查询的余弦相似度
        for doc_id in doc_tfs:
            doc_tf = doc_tfs[doc_id]
            dot_product = 0.0
            for token in query_tf:
                if token in doc_tf:
                    dot_product += (query_tf[token] * idf.get(token, 0)) * \
                                   (doc_tf[token] * idf.get(token, 0))

            doc_norm = math.sqrt(
                sum((doc_tf.get(t, 0) * idf.get(t, 0)) ** 2 for t in doc_tf)
            ) if doc_tf else 0

            if doc_norm > 0:
                scores[doc_id] = dot_product / (query_norm * doc_norm)
            else:
                scores[doc_id] = 0.0

        return dict(sorted(scores.items(), key=lambda x: x[1], reverse=True))

    def search(
        self,
        query: str,
        session_id: str = "",
        category: str = "",
        limit: int = 10,
        mode: str = "hybrid",
        time_decay: bool = True,
        half_life_days: float = 30.0,
    ) -> List[MemoryEntry]:
        """
        搜索记忆（内置遗忘曲线时间衰减，仅搜索 memories 表）。

        支持三种搜索模式:
          - "keyword": 传统 LIKE 关键词匹配（快速）
          - "semantic": TF-IDF 语义搜索（理解语义相似性）
          - "hybrid": 混合搜索（默认）= 0.4 * keyword_score + 0.6 * tfidf_score

        所有模式均支持时间衰减权重（遗忘曲线），越久远的记忆权重越低。
        最终得分 = 相关性得分 × 时间权重。

        Args:
            query: 搜索查询
            session_id: 会话 ID（空=所有会话）
            category: 记忆类别（已废弃， memories 表只存全局记忆）
            limit: 返回数量
            mode: 搜索模式 "keyword" | "semantic" | "hybrid"
            time_decay: 是否启用时间衰减（默认 True）
            half_life_days: 遗忘曲线半衰期天数（默认 30 天）
                - 7 天: 短期记忆场景，快速遗忘
                - 30 天: 默认，平衡近期与长期
                - 90 天: 知识型场景，长期保留
        """
        conn = self._get_conn()
        conditions = ["1=1"]
        params: list = []

        if session_id:
            conditions.append("session_id = ?")
            params.append(session_id)
        # category 参数已废弃， memories 表只存全局记忆，无需按 category 过滤

        # 过滤过期
        conditions.append("(expires_at = '' OR expires_at > ?)")
        params.append(timestamp())

        where = " AND ".join(conditions)

        if mode == "keyword":
            return self._search_keyword(conn, query, where, params, limit,
                                         time_decay=time_decay, half_life_days=half_life_days)
        elif mode == "semantic":
            return self._search_semantic(conn, query, where, params, limit,
                                         time_decay=time_decay, half_life_days=half_life_days)
        else:
            # 混合模式：取两种搜索结果的加权和 + 时间衰减
            # 注意：传入 update_access=False 避免子搜索重复更新 access_count
            keyword_results = self._search_keyword(conn, query, where, params, limit * 2,
                                                   time_decay=False, half_life_days=half_life_days,
                                                   update_access=False)
            semantic_results = self._search_semantic(conn, query, where, params, limit * 2,
                                                     time_decay=False, half_life_days=half_life_days,
                                                     update_access=False)

            # 合并评分（相关性 + 时间衰减）
            combined: Dict[str, Tuple[MemoryEntry, float]] = {}
            for i, entry in enumerate(keyword_results):
                relevance = 1.0 - (i / max(len(keyword_results), 1))
                score = relevance * 0.4
                combined[entry.id] = (entry, score)

            for i, entry in enumerate(semantic_results):
                relevance = 1.0 - (i / max(len(semantic_results), 1))
                sem_score = relevance * 0.6
                if entry.id in combined:
                    combined[entry.id] = (entry, combined[entry.id][1] + sem_score)
                else:
                    combined[entry.id] = (entry, sem_score)

            # 应用时间衰减权重: final_score = relevance_score × time_weight
            if time_decay:
                for mem_id, (entry, relevance_score) in combined.items():
                    tw = self._compute_time_weight(entry.created_at, half_life_days)
                    combined[mem_id] = (entry, relevance_score * tw)

            # 按综合得分排序
            sorted_results = sorted(
                combined.values(),
                key=lambda x: x[1],
                reverse=True,
            )

            # 更新访问计数
            for entry, _ in sorted_results[:limit]:
                conn.execute(
                    "UPDATE memories SET access_count = access_count + 1 WHERE id = ?",
                    (entry.id,),
                )
            conn.commit()

            return [entry for entry, _ in sorted_results[:limit]]

    def _search_keyword(
        self,
        conn: sqlite3.Connection,
        query: str,
        where: str,
        params: list,
        limit: int,
        time_decay: bool = False,
        half_life_days: float = 30.0,
        update_access: bool = True,
    ) -> List[MemoryEntry]:
        """关键词 LIKE 搜索（支持时间衰减重排序）"""
        like_pattern = f"%{query}%"
        conditions = f"{where} AND (content LIKE ? OR summary LIKE ? OR key LIKE ?)"
        search_params = params + [like_pattern, like_pattern, like_pattern]

        # 取更多候选（后续按时间衰减重排序）
        fetch_limit = limit * 3 if time_decay else limit
        sql = f"""
            SELECT * FROM memories WHERE {conditions}
            ORDER BY importance DESC, access_count DESC
            LIMIT ?
        """
        search_params.append(fetch_limit)
        rows = conn.execute(sql, search_params).fetchall()

        entries = [MemoryEntry.from_row(row) for row in rows]

        # 应用时间衰减重排序
        if time_decay and entries:
            scored = []
            for entry in entries:
                base_score = entry.importance + entry.access_count * 0.01
                tw = self._compute_time_weight(entry.created_at, half_life_days)
                scored.append((entry, base_score * tw))
            scored.sort(key=lambda x: x[1], reverse=True)
            entries = [e for e, _ in scored[:limit]]

        # 更新访问计数
        if update_access:
            for entry in entries:
                conn.execute(
                    "UPDATE memories SET access_count = access_count + 1 WHERE id = ?",
                    (entry.id,),
                )
            conn.commit()

        return entries

    def _search_semantic(
        self,
        conn: sqlite3.Connection,
        query: str,
        where: str,
        params: list,
        limit: int,
        time_decay: bool = False,
        half_life_days: float = 30.0,
        update_access: bool = True,
    ) -> List[MemoryEntry]:
        """TF-IDF 语义搜索（支持时间衰减重排序）"""
        # 先取一批候选文档
        candidate_sql = f"SELECT * FROM memories WHERE {where} ORDER BY created_at DESC LIMIT 200"
        rows = conn.execute(candidate_sql, params).fetchall()

        if not rows:
            return []

        # 构建文档列表（content + summary + key 混合文本）
        documents = []
        row_map: Dict[str, sqlite3.Row] = {}
        for row in rows:
            doc_id = row["id"]
            text = f"{row['content']} {row['summary']} {row['key']}"
            documents.append((doc_id, text))
            row_map[doc_id] = row

        # 计算 TF-IDF 得分
        scores = self._compute_tfidf(query, documents)

        # 应用时间衰减重排序
        if time_decay:
            for doc_id, tfidf_score in scores.items():
                row = row_map.get(doc_id)
                if row:
                    tw = self._compute_time_weight(row["created_at"], half_life_days)
                    scores[doc_id] = tfidf_score * tw

        # 按得分排序，取 top N
        top_ids = list(scores.keys())[:limit]
        result = [MemoryEntry.from_row(row_map[doc_id]) for doc_id in top_ids if doc_id in row_map]

        # 更新访问计数
        if update_access:
            for entry in result:
                conn.execute(
                    "UPDATE memories SET access_count = access_count + 1 WHERE id = ?",
                    (entry.id,),
                )
            conn.commit()

        return result

    def search_across_sessions(
        self,
        query: str,
        category: str = "",
        limit: int = 20,
        mode: str = "hybrid",
        time_decay: bool = True,
        half_life_days: float = 30.0,
    ) -> List[MemoryEntry]:
        """跨会话搜索"""
        return self.search(query, session_id="", category=category, limit=limit, mode=mode,
                          time_decay=time_decay, half_life_days=half_life_days)

    # ==========================================================================
    # 记忆总结与维护
    # ==========================================================================

    def get_recent_for_summary(self, session_id: str, count: int = 20) -> List[SessionMessageEntry]:
        """获取最近的对话用于总结（返回 SessionMessageEntry）"""
        return self.get_conversation(session_id, limit=count)

    def save_summary(
        self,
        session_id: str,
        summary: str,
        original_count: int = 0,
    ) -> str:
        """保存对话总结为全局记忆（保存到 memories 表）"""
        return self.add_global(
            session_id=session_id,
            key="conversation_summary",
            content=summary,
            summary=summary[:500],
            importance=0.6,
            metadata={"original_message_count": original_count},
        )

    def get_error_patterns(self, session_id: str = "global", limit: int = 20) -> List[MemoryEntry]:
        """获取历史错误模式(用于避免重复犯错，从 memories 表)"""
        return self._query_memories(
            session_id=session_id,
            key="error_pattern",
            limit=limit,
            order_by="importance DESC",
        )

    def record_error_pattern(
        self,
        error: str,
        fix: str = "",
        session_id: str = "global",
    ):
        """记录错误模式及修复方案（保存到 memories 表）"""
        self.add_global(
            session_id=session_id,
            key="error_pattern",
            content=f"错误: {error}\n修复: {fix}",
            summary=f"错误: {truncate_str(error, 200)} | 修复: {truncate_str(fix, 200)}",
            importance=0.8,
        )

    # ==========================================================================
    # 统计与管理
    # ==========================================================================

    def get_stats(self) -> Dict[str, Any]:
        """获取记忆系统统计"""
        conn = self._get_conn()
        stats = {}
        # 统计 session_messages 表
        row = conn.execute("SELECT COUNT(*) as cnt FROM session_messages").fetchone()
        stats["session_messages_count"] = row["cnt"]
        # 统计 memories 表
        row = conn.execute("SELECT COUNT(*) as cnt FROM memories").fetchone()
        stats["memories_count"] = row["cnt"]
        stats["total_count"] = stats["session_messages_count"] + stats["memories_count"]
        # 统计 distinct sessions
        sessions = conn.execute(
            "SELECT COUNT(DISTINCT session_id) as cnt FROM session_messages"
        ).fetchone()
        stats["distinct_sessions"] = sessions["cnt"]
        return stats

    def cleanup_expired(self) -> int:
        """清理过期记忆（仅 memories 表支持过期）"""
        conn = self._get_conn()
        now = timestamp()
        cursor = conn.execute(
            "DELETE FROM memories WHERE expires_at != '' AND expires_at < ?",
            (now,),
        )
        conn.commit()
        count = cursor.rowcount
        if count > 0:
            logger.info(f"清理过期记忆: {count} 条")
        return count

    def export_session(self, session_id: str) -> Dict[str, Any]:
        """导出会话所有数据（包括消息和记忆）"""
        messages = self.get_conversation_all(session_id, limit=0)
        memories = self._query_memories(session_id=session_id, limit=10000)
        return {
            "session_id": session_id,
            "exported_at": timestamp(),
            "messages": [asdict(m) for m in messages],
            "memories": [asdict(m) for m in memories],
        }
