File size: 12,079 Bytes
0e3d4b8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
"""Persistent Memory — long-term memory that survives restarts.

Three layers:
- Working memory: current conversation context (in-memory, per session)
- Episodic memory: past conversations and events (SQLite, persistent)
- Semantic memory: extracted facts and knowledge (SQLite, persistent)

All layers are persisted to SQLite so the LLM remembers everything
across restarts. Semantic memory is auto-extracted from conversations.

Linked to the recursive link graph for knowledge graph traversal.
"""

from __future__ import annotations

import hashlib
import json
import logging
import os
import sqlite3
import time
from dataclasses import dataclass, field
from typing import Any

logger = logging.getLogger(__name__)


@dataclass
class EpisodicMemory:
    """A single episodic memory (event/conversation)."""
    id: str
    session_id: str
    role: str  # "user", "assistant", "system", "event"
    content: str
    channel: str = "cli"
    timestamp: float = field(default_factory=time.time)
    importance: float = 0.5  # 0-1, higher = more important
    tags: list[str] = field(default_factory=list)


@dataclass
class SemanticMemory:
    """A extracted fact or piece of knowledge."""
    id: str
    fact: str
    source: str = ""  # what conversation/event it came from
    confidence: float = 0.5
    timestamp: float = field(default_factory=time.time)
    access_count: int = 0
    last_accessed: float = field(default_factory=time.time)
    tags: list[str] = field(default_factory=list)


class PersistentMemory:
    """Persistent long-term memory backed by SQLite.

    Survives restarts. Three layers:
    - Working: current session context (in-memory)
    - Episodic: past conversations (SQLite)
    - Semantic: extracted facts (SQLite)

    Auto-extracts semantic memories from conversations.
    Provides context injection for inference.
    """

    IMPORTANCE_DECAY = 0.0001  # per second
    MAX_WORKING_MEMORY = 20  # max items in working memory

    def __init__(self, db_path: str) -> None:
        self.db_path = db_path
        os.makedirs(os.path.dirname(db_path) or ".", exist_ok=True)

        self._working: list[dict[str, Any]] = []
        self._session_id: str = ""
        self._stats = {
            "episodic_stored": 0,
            "semantic_extracted": 0,
            "memories_recalled": 0,
            "context_injections": 0,
        }

        self._init_db()
        self._load_stats()

    def _init_db(self) -> None:
        """Initialize SQLite tables."""
        with sqlite3.connect(self.db_path) as conn:
            conn.executescript("""
                CREATE TABLE IF NOT EXISTS episodic (
                    id TEXT PRIMARY KEY,
                    session_id TEXT,
                    role TEXT,
                    content TEXT,
                    channel TEXT,
                    timestamp REAL,
                    importance REAL,
                    tags TEXT
                );
                CREATE INDEX IF NOT EXISTS idx_episodic_session ON episodic(session_id);
                CREATE INDEX IF NOT EXISTS idx_episodic_importance ON episodic(importance);
                CREATE INDEX IF NOT EXISTS idx_episodic_timestamp ON episodic(timestamp);

                CREATE TABLE IF NOT EXISTS semantic (
                    id TEXT PRIMARY KEY,
                    fact TEXT,
                    source TEXT,
                    confidence REAL,
                    timestamp REAL,
                    access_count INTEGER DEFAULT 0,
                    last_accessed REAL,
                    tags TEXT
                );
                CREATE INDEX IF NOT EXISTS idx_semantic_confidence ON semantic(confidence);
                CREATE INDEX IF NOT EXISTS idx_semantic_tags ON semantic(tags);
            """)

    def _load_stats(self) -> None:
        """Load counts from DB."""
        with sqlite3.connect(self.db_path) as conn:
            self._stats["episodic_stored"] = conn.execute("SELECT COUNT(*) FROM episodic").fetchone()[0]
            self._stats["semantic_extracted"] = conn.execute("SELECT COUNT(*) FROM semantic").fetchone()[0]

    def set_session(self, session_id: str) -> None:
        """Set the current session ID."""
        self._session_id = session_id
        self._working.clear()

    def add_episodic(self, role: str, content: str, channel: str = "cli",
                     importance: float = 0.5, tags: list[str] | None = None) -> str:
        """Store an episodic memory (conversation turn or event)."""
        mem_id = hashlib.sha256(f"{role}:{content}:{time.time()}".encode()).hexdigest()[:16]
        mem = EpisodicMemory(
            id=mem_id, session_id=self._session_id, role=role,
            content=content, channel=channel, importance=importance,
            tags=tags or [],
        )

        with sqlite3.connect(self.db_path) as conn:
            conn.execute(
                "INSERT OR REPLACE INTO episodic VALUES (?,?,?,?,?,?,?,?)",
                (mem.id, mem.session_id, mem.role, mem.content, mem.channel,
                 mem.timestamp, mem.importance, json.dumps(mem.tags))
            )

        # Also add to working memory
        self._working.append({"role": role, "content": content, "timestamp": time.time()})
        if len(self._working) > self.MAX_WORKING_MEMORY:
            self._working = self._working[-self.MAX_WORKING_MEMORY:]

        self._stats["episodic_stored"] += 1
        return mem_id

    def add_semantic(self, fact: str, source: str = "", confidence: float = 0.5,
                     tags: list[str] | None = None) -> str:
        """Store a semantic memory (extracted fact)."""
        fact_id = hashlib.sha256(f"{fact}:{time.time()}".encode()).hexdigest()[:16]
        mem = SemanticMemory(
            id=fact_id, fact=fact, source=source, confidence=confidence,
            tags=tags or [],
        )

        with sqlite3.connect(self.db_path) as conn:
            conn.execute(
                "INSERT OR REPLACE INTO semantic VALUES (?,?,?,?,?,?,?,?)",
                (mem.id, mem.fact, mem.source, mem.confidence, mem.timestamp,
                 mem.access_count, mem.last_accessed, json.dumps(mem.tags))
            )

        self._stats["semantic_extracted"] += 1
        return fact_id

    def recall_episodic(self, query: str, max_results: int = 5) -> list[EpisodicMemory]:
        """Recall episodic memories related to a query."""
        # Simple keyword search
        keywords = query.lower().split()
        with sqlite3.connect(self.db_path) as conn:
            rows = conn.execute(
                "SELECT * FROM episodic ORDER BY importance DESC, timestamp DESC LIMIT ?",
                (max_results * 3,)
            ).fetchall()

        results = []
        for row in rows:
            mem = self._row_to_episodic(row)
            # Score by keyword overlap
            content_lower = mem.content.lower()
            score = sum(1 for kw in keywords if kw in content_lower)
            if score > 0:
                # Apply time decay
                age = time.time() - mem.timestamp
                mem.importance = max(0.01, mem.importance - age * self.IMPORTANCE_DECAY)
                results.append((score + mem.importance, mem))

        results.sort(key=lambda x: -x[0])
        self._stats["memories_recalled"] += len(results[:max_results])
        return [mem for _, mem in results[:max_results]]

    def recall_semantic(self, query: str, max_results: int = 5) -> list[SemanticMemory]:
        """Recall semantic memories (facts) related to a query."""
        keywords = query.lower().split()
        with sqlite3.connect(self.db_path) as conn:
            rows = conn.execute(
                "SELECT * FROM semantic ORDER BY confidence DESC, last_accessed DESC LIMIT ?",
                (max_results * 3,)
            ).fetchall()

        results = []
        for row in rows:
            mem = self._row_to_semantic(row)
            fact_lower = mem.fact.lower()
            score = sum(1 for kw in keywords if kw in fact_lower)
            if score > 0:
                results.append((score + mem.confidence, mem))
                # Update access count
                with sqlite3.connect(self.db_path) as conn:
                    conn.execute(
                        "UPDATE semantic SET access_count = access_count + 1, last_accessed = ? WHERE id = ?",
                        (time.time(), mem.id)
                    )

        results.sort(key=lambda x: -x[0])
        return [mem for _, mem in results[:max_results]]

    def extract_semantic(self, user_message: str, assistant_response: str) -> list[str]:
        """Auto-extract semantic memories (facts) from a conversation.

        Simple extraction: look for statements that contain facts.
        """
        facts: list[str] = []
        # Simple heuristics for fact extraction
        sentences = assistant_response.replace("!", ".").replace("?", ".").split(".")
        for sentence in sentences:
            s = sentence.strip()
            if len(s) < 10 or len(s) > 200:
                continue
            # Skip questions and commands
            if s.endswith("?") or s.startswith("You ") or s.startswith("I "):
                continue
            # Look for factual statements (contains "is", "are", "was", "has", etc.)
            fact_indicators = [" is ", " are ", " was ", " has ", " have ", " can ", " cannot ",
                             " means ", " refers to ", " defined as ", " consists of "]
            if any(ind in s.lower() for ind in fact_indicators):
                fact_id = self.add_semantic(s, source=self._session_id, confidence=0.6)
                facts.append(fact_id)

        return facts

    def get_context(self, query: str, max_episodic: int = 3, max_semantic: int = 3) -> str:
        """Get memory context to inject into the prompt for inference."""
        parts: list[str] = []

        # Working memory (current session)
        if self._working:
            recent = self._working[-5:]
            working_text = " | ".join(f"{m['role']}: {m['content'][:80]}" for m in recent)
            parts.append(f"Recent: {working_text}")

        # Episodic memory
        episodic = self.recall_episodic(query, max_results=max_episodic)
        if episodic:
            ep_text = " | ".join(f"{m.role}: {m.content[:80]}" for m in episodic)
            parts.append(f"Past: {ep_text}")

        # Semantic memory
        semantic = self.recall_semantic(query, max_results=max_semantic)
        if semantic:
            sem_text = " | ".join(m.fact[:80] for m in semantic)
            parts.append(f"Facts: {sem_text}")

        if parts:
            self._stats["context_injections"] += 1

        return " | ".join(parts)

    def get_working_memory(self) -> list[dict[str, Any]]:
        """Get current working memory (this session)."""
        return self._working.copy()

    def clear_working(self) -> None:
        """Clear working memory."""
        self._working.clear()

    def _row_to_episodic(self, row: tuple) -> EpisodicMemory:
        return EpisodicMemory(
            id=row[0], session_id=row[1], role=row[2], content=row[3],
            channel=row[4], timestamp=row[5], importance=row[6],
            tags=json.loads(row[7]) if row[7] else [],
        )

    def _row_to_semantic(self, row: tuple) -> SemanticMemory:
        return SemanticMemory(
            id=row[0], fact=row[1], source=row[2], confidence=row[3],
            timestamp=row[4], access_count=row[5], last_accessed=row[6],
            tags=json.loads(row[7]) if row[7] else [],
        )

    def get_stats(self) -> dict[str, Any]:
        with sqlite3.connect(self.db_path) as conn:
            ep_count = conn.execute("SELECT COUNT(*) FROM episodic").fetchone()[0]
            sem_count = conn.execute("SELECT COUNT(*) FROM semantic").fetchone()[0]
        return {
            **self._stats,
            "episodic_total": ep_count,
            "semantic_total": sem_count,
            "working_size": len(self._working),
            "session_id": self._session_id,
        }