File size: 9,934 Bytes
2eef9ea
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
backend/memory/memory_store.py

Two-tier memory architecture:

  Short-term (Redis):
    - Current session context
    - Recent tool results
    - Working memory for active task
    - TTL: configurable (default 24h)

  Long-term (SQLite/PostgreSQL):
    - Episodic memory: what happened in past tasks
    - Semantic memory: learned facts and patterns
    - Procedural memory: successful task strategies
    - Persists across restarts

Memory retrieval uses simple keyword matching (production: use embeddings + vector DB).
"""
from __future__ import annotations
import hashlib
import json
import time
from datetime import datetime, timezone
from typing import Any

from sqlalchemy import Column, String, Float, Text, Integer, DateTime, create_engine, select
from sqlalchemy.orm import DeclarativeBase, Session
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker

from ..core.config import get_settings
from ..core.logger import get_logger

log = get_logger(__name__)


# ── SQLAlchemy models ─────────────────────────────────────────────────────────

class Base(DeclarativeBase):
    pass


class MemoryRecord(Base):
    __tablename__ = "memories"
    id          = Column(String, primary_key=True)
    task_id     = Column(String, index=True)
    content     = Column(Text, nullable=False)
    memory_type = Column(String, default="episodic")    # episodic|semantic|procedural
    importance  = Column(Float, default=0.5)
    tags        = Column(Text, default="[]")            # JSON list
    access_count = Column(Integer, default=0)
    created_at  = Column(DateTime, default=datetime.utcnow)
    last_accessed = Column(DateTime, default=datetime.utcnow)


class TaskRecord(Base):
    __tablename__ = "tasks"
    task_id     = Column(String, primary_key=True)
    task        = Column(Text, nullable=False)
    status      = Column(String, default="pending")
    final_output = Column(Text)
    quality_score = Column(Float)
    total_tokens = Column(Integer, default=0)
    created_at  = Column(DateTime, default=datetime.utcnow)
    completed_at = Column(DateTime)
    state_json  = Column(Text)    # Full state snapshot


# ── Database setup ─────────────────────────────────────────────────────────────

_engine = None
_session_factory = None


async def init_db():
    global _engine, _session_factory
    settings = get_settings()
    _engine = create_async_engine(settings.database_url, echo=False)
    _session_factory = async_sessionmaker(_engine, expire_on_commit=False)
    async with _engine.begin() as conn:
        await conn.run_sync(Base.metadata.create_all)
    log.info("Database initialized", url=settings.database_url)


async def get_session() -> AsyncSession:
    return _session_factory()


# ── Redis short-term memory ───────────────────────────────────────────────────

_redis_client = None
_redis_checked = False  # prevents re-attempting after a failed connect


def get_redis():
    global _redis_client, _redis_checked
    if _redis_checked:
        return _redis_client
    _redis_checked = True
    import redis
    settings = get_settings()
    try:
        client = redis.from_url(
            settings.redis_url, decode_responses=True,
            socket_connect_timeout=2, socket_timeout=2,
        )
        client.ping()
        _redis_client = client
    except Exception as e:
        log.warning("Redis unavailable β€” short-term memory disabled", error=str(e))
        _redis_client = None
    return _redis_client


class ShortTermMemory:
    """Redis-backed working memory for active tasks."""

    PREFIX = "agent:stm:"

    def set(self, task_id: str, key: str, value: Any, ttl: int | None = None) -> None:
        r = get_redis()
        if r is None:
            return
        full_key = f"{self.PREFIX}{task_id}:{key}"
        r.set(full_key, json.dumps(value), ex=ttl or get_settings().redis_ttl)

    def get(self, task_id: str, key: str) -> Any | None:
        r = get_redis()
        if r is None:
            return None
        raw = r.get(f"{self.PREFIX}{task_id}:{key}")
        return json.loads(raw) if raw else None

    def get_all(self, task_id: str) -> dict[str, Any]:
        r = get_redis()
        if r is None:
            return {}
        pattern = f"{self.PREFIX}{task_id}:*"
        keys = r.keys(pattern)
        result = {}
        for k in keys:
            sub_key = k.replace(f"{self.PREFIX}{task_id}:", "")
            raw = r.get(k)
            if raw:
                result[sub_key] = json.loads(raw)
        return result

    def store_state(self, task_id: str, state: dict) -> None:
        """Cache full workflow state for resumption."""
        self.set(task_id, "state", state, ttl=3600)

    def get_state(self, task_id: str) -> dict | None:
        return self.get(task_id, "state")

    def clear(self, task_id: str) -> None:
        r = get_redis()
        if r is None:
            return
        for k in r.keys(f"{self.PREFIX}{task_id}:*"):
            r.delete(k)


# ── Long-term memory ──────────────────────────────────────────────────────────

class LongTermMemory:
    """SQLite/PostgreSQL-backed episodic + semantic memory."""

    async def store(self, memory: dict, task_id: str = "") -> str:
        mem_id = hashlib.sha256(
            (memory["content"] + str(time.time())).encode()
        ).hexdigest()[:12]

        async with await get_session() as session:
            record = MemoryRecord(
                id=mem_id,
                task_id=task_id,
                content=memory["content"],
                memory_type=memory.get("memory_type", "episodic"),
                importance=memory.get("importance", 0.5),
                tags=json.dumps(memory.get("tags", [])),
            )
            session.add(record)
            await session.commit()

        log.debug("Memory stored", id=mem_id, type=memory.get("memory_type"))
        return mem_id

    async def retrieve(
        self,
        query: str,
        memory_type: str | None = None,
        limit: int = 5,
        min_importance: float = 0.3,
    ) -> list[dict]:
        """
        Retrieve relevant memories using keyword matching.
        Production upgrade: embed query + cosine similarity with pgvector.
        """
        async with await get_session() as session:
            result = await session.execute(
                select(MemoryRecord)
                .where(MemoryRecord.importance >= min_importance)
                .order_by(MemoryRecord.importance.desc())
                .limit(50)
            )
            all_memories = result.scalars().all()

        # Score by keyword overlap
        query_words = set(query.lower().split())
        scored = []
        for m in all_memories:
            if memory_type and m.memory_type != memory_type:
                continue
            content_words = set(m.content.lower().split())
            overlap = len(query_words & content_words)
            if overlap > 0:
                scored.append((overlap, m))

        scored.sort(key=lambda x: x[0], reverse=True)

        results = []
        for _, m in scored[:limit]:
            results.append({
                "memory_id": m.id,
                "content": m.content,
                "memory_type": m.memory_type,
                "importance": m.importance,
                "tags": json.loads(m.tags),
                "created_at": m.created_at.isoformat() if m.created_at else "",
            })

        return results

    async def store_task(self, task_id: str, task: str, state: dict) -> None:
        """Persist task record to DB."""
        async with await get_session() as session:
            existing = await session.get(TaskRecord, task_id)
            if existing:
                existing.status = state.get("status", "unknown")
                existing.final_output = state.get("final_output")
                existing.quality_score = state.get("quality_score")
                existing.total_tokens = state.get("total_tokens", 0)
                existing.state_json = json.dumps(state, default=str)
                if state.get("status") in ("completed", "failed"):
                    existing.completed_at = datetime.utcnow()
            else:
                record = TaskRecord(
                    task_id=task_id,
                    task=task,
                    status=state.get("status", "pending"),
                    state_json=json.dumps(state, default=str),
                )
                session.add(record)
            await session.commit()

    async def get_recent_tasks(self, limit: int = 10) -> list[dict]:
        async with await get_session() as session:
            result = await session.execute(
                select(TaskRecord).order_by(TaskRecord.created_at.desc()).limit(limit)
            )
            tasks = result.scalars().all()
        return [
            {
                "task_id": t.task_id,
                "task": t.task[:100],
                "status": t.status,
                "quality_score": t.quality_score,
                "total_tokens": t.total_tokens,
                "created_at": t.created_at.isoformat() if t.created_at else "",
            }
            for t in tasks
        ]


# ── Unified memory interface ──────────────────────────────────────────────────

short_term = ShortTermMemory()
long_term = LongTermMemory()