File size: 3,639 Bytes
6b62834
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
"""Memory Manager β€” coordinates short-term, long-term, and working memory."""

from typing import Optional

from agentic_rag.data.models import Message


class MemoryManager:
    """Coordinates the three-tier memory system.

    - ShortTermMemory: Sliding window of recent messages
    - LongTermMemory: Vector-indexed semantic memories (Milvus-backed, future)
    - WorkingMemory: Ephemeral scratchpad per turn
    """

    def __init__(self, max_short_term_tokens: int = 8000):
        self.max_short_term_tokens = max_short_term_tokens
        self._short_term: dict[str, list[Message]] = {}  # session_id -> messages
        self._working: dict[str, dict] = {}               # session_id -> scratchpad

    # ── Short-term Memory ──────────────────────────

    def add_message(self, session_id: str, message: Message) -> None:
        """Add a message to short-term memory."""
        if session_id not in self._short_term:
            self._short_term[session_id] = []
        self._short_term[session_id].append(message)
        self._trim(session_id)

    def get_messages(self, session_id: str, limit: int = 50) -> list[Message]:
        """Get recent messages for a session."""
        messages = self._short_term.get(session_id, [])
        return messages[-limit:] if len(messages) > limit else messages

    def clear_short_term(self, session_id: str) -> None:
        """Clear short-term memory for a session."""
        self._short_term.pop(session_id, None)

    def _trim(self, session_id: str) -> None:
        """Trim messages to fit within token budget (approximate)."""
        messages = self._short_term.get(session_id, [])
        total_tokens = sum(len(str(m.content)) // 4 for m in messages)
        while total_tokens > self.max_short_term_tokens and len(messages) > 2:
            removed = messages.pop(0)
            total_tokens -= len(str(removed.content)) // 4

    # ── Working Memory ─────────────────────────────

    def set_working(self, session_id: str, key: str, value) -> None:
        """Set a working memory value for the current turn."""
        if session_id not in self._working:
            self._working[session_id] = {}
        self._working[session_id][key] = value

    def get_working(self, session_id: str, key: str, default=None):
        """Get a working memory value."""
        return self._working.get(session_id, {}).get(key, default)

    def clear_working(self, session_id: str) -> None:
        """Clear working memory at end of turn."""
        self._working.pop(session_id, None)

    # ── Long-term Memory (future) ──────────────────

    async def search_long_term(self, session_id: str, query: str, top_k: int = 10) -> list[dict]:
        """Search long-term memory for relevant past interactions."""
        # Placeholder β€” will be implemented with Milvus in Phase 3
        return []

    async def store_long_term(self, session_id: str, fact: str, metadata: dict | None = None) -> None:
        """Store a fact in long-term memory."""
        # Placeholder β€” will be implemented with Milvus in Phase 3
        pass


# Global instance
_memory_manager: Optional[MemoryManager] = None


def get_memory_manager() -> MemoryManager:
    global _memory_manager
    if _memory_manager is None:
        from agentic_rag.config.settings import get_settings
        _memory_manager = MemoryManager(
            max_short_term_tokens=get_settings().memory.short_term_max_tokens,
        )
    return _memory_manager