diff --git a/headroom/integrations/litellm_callback.py b/headroom/integrations/litellm_callback.py index 1cf152c59..7256d85b8 100644 --- a/headroom/integrations/litellm_callback.py +++ b/headroom/integrations/litellm_callback.py @@ -104,6 +104,7 @@ class HeadroomCallback(_CustomLogger): data, call_type = cache, data if data is None: return None + if call_type not in ("completion", "acompletion"): return data diff --git a/headroom/proxy/memory_rank_policy.py b/headroom/proxy/memory_rank_policy.py new file mode 100644 index 000000000..5ab6799b3 --- /dev/null +++ b/headroom/proxy/memory_rank_policy.py @@ -0,0 +1,66 @@ +"""Pure memory ranking policy helpers. + +This module owns timestamp parsing and recency score math for proxy memory +ranking. It deliberately avoids backend objects and ranker classes so the +formula can be tested, ported, and reused independently of retrieval adapters. +""" + +from __future__ import annotations + +import math +from datetime import datetime, timezone + +UTC = timezone.utc + + +def parse_memory_created_at(value: object) -> datetime | None: + """Best-effort parse of a memory timestamp into a UTC-aware datetime.""" + if value is None: + return None + if isinstance(value, datetime): + return value if value.tzinfo is not None else value.replace(tzinfo=UTC) + if isinstance(value, str): + try: + return datetime.fromisoformat(value.replace("Z", "+00:00")) + except ValueError: + return None + return None + + +def memory_recency_factor( + *, + now: datetime, + created_at: datetime | None, + decay_days: float, +) -> float: + """Compute the recency multiplier for one memory candidate. + + Missing timestamps and future timestamps are neutral. For normal historical + timestamps the multiplier is ``exp(-age_days / decay_days)``. + """ + if created_at is None: + return 1.0 + if created_at.tzinfo is None: + created_at = created_at.replace(tzinfo=UTC) + if now.tzinfo is None: + now = now.replace(tzinfo=UTC) + + age_days = (now - created_at).total_seconds() / 86400.0 + if age_days <= 0: + return 1.0 + return math.exp(-age_days / decay_days) + + +def boost_memory_score( + *, + score: float, + now: datetime, + created_at: datetime | None, + decay_days: float, +) -> float: + """Apply the recency multiplier to a backend similarity score.""" + return score * memory_recency_factor( + now=now, + created_at=created_at, + decay_days=decay_days, + ) diff --git a/headroom/proxy/memory_ranker.py b/headroom/proxy/memory_ranker.py index 09e37b9a7..7ee3798fb 100644 --- a/headroom/proxy/memory_ranker.py +++ b/headroom/proxy/memory_ranker.py @@ -21,11 +21,16 @@ no network, no disk. from __future__ import annotations -import math from dataclasses import dataclass, field from datetime import datetime, timezone from typing import Protocol +from headroom.proxy.memory_rank_policy import ( + boost_memory_score, + memory_recency_factor, + parse_memory_created_at, +) + # Use ``timezone.utc`` (always available) instead of ``datetime.UTC`` # (Python 3.11+) so this module imports cleanly on older interpreters. _UTC = timezone.utc @@ -73,7 +78,7 @@ class MemoryCandidate: content = str(getattr(memory, "content", "")) if memory is not None else "" memory_id = str(getattr(memory, "id", "") or "") if memory is not None else "" raw_dt = getattr(memory, "created_at", None) if memory is not None else None - created_at = _parse_created_at(raw_dt) + created_at = parse_memory_created_at(raw_dt) raw_related = getattr(result, "related_entities", None) or () related = tuple(str(x) for x in raw_related) source_meta = getattr(memory, "metadata", None) or {} @@ -95,16 +100,7 @@ def _parse_created_at(value: object) -> datetime | None: string (with or without trailing ``Z``). Anything else → ``None`` so the ranker treats the candidate as recency-neutral. """ - if value is None: - return None - if isinstance(value, datetime): - return value if value.tzinfo is not None else value.replace(tzinfo=_UTC) - if isinstance(value, str): - try: - return datetime.fromisoformat(value.replace("Z", "+00:00")) - except ValueError: - return None - return None + return parse_memory_created_at(value) class MemoryRanker(Protocol): @@ -165,8 +161,13 @@ class RecencyBoostRanker: now = datetime.now(_UTC) boosted: list[tuple[int, MemoryCandidate, float]] = [] for idx, c in enumerate(candidates): - factor = self._recency_factor(now, c.created_at) - boosted.append((idx, c, c.score * factor)) + new_score = boost_memory_score( + score=c.score, + now=now, + created_at=c.created_at, + decay_days=self.decay_days, + ) + boosted.append((idx, c, new_score)) # Sort descending by boosted score; stable on ties via the # captured idx — same input order preserved on ties for @@ -182,6 +183,7 @@ class RecencyBoostRanker: created_at=c.created_at, source=c.source, related_entities=c.related_entities, + id=c.id, ) for _, c, new_score in boosted ] @@ -193,14 +195,8 @@ class RecencyBoostRanker: Future timestamps → 1.0 (clock-skew defence). Otherwise: ``exp(-age_days / decay_days)``. """ - if created_at is None: - return 1.0 - # Normalize to UTC-aware for safe subtraction. - if created_at.tzinfo is None: - created_at = created_at.replace(tzinfo=_UTC) - delta = now - created_at - age_days = delta.total_seconds() / 86400.0 - if age_days <= 0: - # Future-dated row (clock skew). Clamp to neutral. - return 1.0 - return math.exp(-age_days / self.decay_days) + return memory_recency_factor( + now=now, + created_at=created_at, + decay_days=self.decay_days, + ) diff --git a/tests/test_memory_rank_policy.py b/tests/test_memory_rank_policy.py new file mode 100644 index 000000000..06f24a616 --- /dev/null +++ b/tests/test_memory_rank_policy.py @@ -0,0 +1,56 @@ +"""Tests for pure memory rank policy formulas.""" + +from __future__ import annotations + +import math +from datetime import datetime, timedelta, timezone + +from headroom.proxy.memory_rank_policy import ( + boost_memory_score, + memory_recency_factor, + parse_memory_created_at, +) + +_UTC = timezone.utc + + +def test_parse_memory_created_at_accepts_zulu_iso_string() -> None: + parsed = parse_memory_created_at("2026-05-19T12:00:00Z") + assert parsed == datetime(2026, 5, 19, 12, 0, tzinfo=_UTC) + + +def test_parse_memory_created_at_normalizes_naive_datetime_to_utc() -> None: + parsed = parse_memory_created_at(datetime(2026, 5, 19, 12, 0)) + assert parsed == datetime(2026, 5, 19, 12, 0, tzinfo=_UTC) + + +def test_parse_memory_created_at_invalid_values_are_neutral() -> None: + assert parse_memory_created_at("not-a-date") is None + assert parse_memory_created_at(123) is None + assert parse_memory_created_at(None) is None + + +def test_memory_recency_factor_uses_exponential_decay() -> None: + now = datetime(2026, 5, 31, tzinfo=_UTC) + created_at = now - timedelta(days=30) + factor = memory_recency_factor(now=now, created_at=created_at, decay_days=30.0) + assert math.isclose(factor, math.exp(-1), rel_tol=1e-12) + + +def test_memory_recency_factor_treats_missing_and_future_dates_as_neutral() -> None: + now = datetime(2026, 5, 31, tzinfo=_UTC) + future = now + timedelta(days=3) + assert memory_recency_factor(now=now, created_at=None, decay_days=30.0) == 1.0 + assert memory_recency_factor(now=now, created_at=future, decay_days=30.0) == 1.0 + + +def test_boost_memory_score_applies_recency_factor() -> None: + now = datetime(2026, 5, 31, tzinfo=_UTC) + created_at = now - timedelta(days=60) + boosted = boost_memory_score( + score=0.9, + now=now, + created_at=created_at, + decay_days=30.0, + ) + assert math.isclose(boosted, 0.9 * math.exp(-2), rel_tol=1e-12) diff --git a/tests/test_memory_ranker.py b/tests/test_memory_ranker.py index 58e2d5d9d..1a14eb8b5 100644 --- a/tests/test_memory_ranker.py +++ b/tests/test_memory_ranker.py @@ -84,6 +84,13 @@ def test_from_backend_result_preserves_memory_id() -> None: assert cand.score == 0.91 +def test_rank_preserves_memory_id() -> None: + """The ranker must not drop the backend ID when rebuilding candidates.""" + cand = MemoryCandidate(content="User prefers Python.", score=0.91, id="mem_abc_123") + out = RecencyBoostRanker().rank([cand]) + assert out[0].id == "mem_abc_123" + + def test_from_backend_result_handles_missing_id() -> None: """Defensive: legacy backend rows without an ID become ``id=""``; the auto-tail formatter renders ``[?]`` for those rows, no crash."""