"""Stanford 式记忆检索打分:relevance × recency × importance 加权后排序。

为什么三者都要:只按相关性会忽略"刚发生的小事"和"一生大事";recency 用指数衰减
模拟遗忘,importance 让里程碑事件不被淹没。三个分量先各自 min-max 归一再加权,
与原论文一致。权重/衰减是可调旋钮(PRD §12 TODO 待校准)。
"""

from __future__ import annotations

from dataclasses import dataclass
from typing import Any


@dataclass
class RetrievalWeights:
    relevance: float = 1.0
    recency: float = 1.0
    importance: float = 1.0


@dataclass
class Candidate:
    item: Any  # 通常是 MemoryRecord
    relevance: float  # 语义相关度 [0,1]
    recency: float  # 新近度 [0,1]
    importance: float  # 重要度 [0,1]


@dataclass
class Ranked:
    item: Any
    total: float


def recency_score(elapsed_hours: float, decay: float = 0.99) -> float:
    """指数衰减:距上次访问越久分越低。elapsed<=0 视为最新(1.0)。"""
    if elapsed_hours <= 0:
        return 1.0
    return decay**elapsed_hours


def _minmax(values: list[float]) -> list[float]:
    lo, hi = min(values), max(values)
    if hi - lo < 1e-9:  # 全相等时不歧视任何一条
        return [1.0 for _ in values]
    return [(v - lo) / (hi - lo) for v in values]


def combine(
    relevance: float, recency: float, importance: float, weights: RetrievalWeights
) -> float:
    return (
        weights.relevance * relevance
        + weights.recency * recency
        + weights.importance * importance
    )


def rank(
    candidates: list[Candidate],
    weights: RetrievalWeights | None = None,
    *,
    normalize: bool = True,
) -> list[Ranked]:
    """对候选记忆打分排序(降序)。normalize=False 时跳过 min-max(便于精确测分量)。"""
    if not candidates:
        return []
    weights = weights or RetrievalWeights()
    rels = [c.relevance for c in candidates]
    recs = [c.recency for c in candidates]
    imps = [c.importance for c in candidates]
    if normalize:
        rels, recs, imps = _minmax(rels), _minmax(recs), _minmax(imps)
    ranked = [
        Ranked(c.item, combine(r, rc, im, weights))
        for c, r, rc, im in zip(candidates, rels, recs, imps)
    ]
    ranked.sort(key=lambda x: x.total, reverse=True)
    return ranked
