1"""Database-backed memory backend — persists entries via DatabaseProviderProtocol."""
2
3from __future__ import annotations
4
5from datetime import UTC, datetime
6import math
7from typing import TYPE_CHECKING, Any, cast
8
9from lexigram.contracts.ai.memory import (
10 MemoryEntry,
11 MemoryQuery,
12 MemorySearchResult,
13)
14from lexigram.contracts.core import HealthCheckResult, HealthStatus
15from lexigram.serialization.backends.json import dumps_str, loads
16
17if TYPE_CHECKING:
18 from lexigram.contracts.data import DatabaseProviderProtocol
19
20_SELECT_ALL = (
21 "SELECT id, owner_id, content, role, timestamp, importance, metadata, embedding"
22 " FROM memory_entries WHERE owner_id = $1"
23)
24_INSERT = (
25 "INSERT INTO memory_entries (id, owner_id, content, role, timestamp, importance, metadata, embedding)"
26 " VALUES ($1,$2,$3,$4,$5,$6,$7,$8)"
27 " ON CONFLICT (id) DO UPDATE SET content=EXCLUDED.content, importance=EXCLUDED.importance"
28)
29_DELETE = "DELETE FROM memory_entries WHERE id = $1 AND owner_id = $2"
30_CLEAR = "DELETE FROM memory_entries WHERE owner_id = $1"
31_RECENT = f"{_SELECT_ALL} ORDER BY timestamp DESC LIMIT $2"
32
33
34def _normalise_timestamp(value: Any) -> datetime:
35 if isinstance(value, datetime):
36 return value if value.tzinfo is not None else value.replace(tzinfo=UTC)
37 if isinstance(value, str):
38 parsed = datetime.fromisoformat(value)
39 return parsed if parsed.tzinfo is not None else parsed.replace(tzinfo=UTC)
40 raise ValueError(f"Unsupported timestamp value type: {type(value).__name__}")
41
42
43def _row_to_entry(row: dict[str, Any]) -> MemoryEntry:
44 metadata_raw = row.get("metadata") or {}
45 metadata: dict[str, Any]
46 if isinstance(metadata_raw, str):
47 metadata = cast("dict[str, Any]", loads(metadata_raw))
48 else:
49 metadata = cast("dict[str, Any]", metadata_raw)
50
51 embedding_raw = row.get("embedding")
52 embedding = cast(
53 "list[float] | None", list(embedding_raw) if embedding_raw else None
54 )
55
56 return MemoryEntry(
57 id=str(row["id"]),
58 owner_id=str(row["owner_id"]),
59 content=str(row["content"]),
60 role=str(row["role"]),
61 timestamp=_normalise_timestamp(row["timestamp"]),
62 importance=float(row.get("importance", 0.5)),
63 metadata=metadata,
64 embedding=embedding,
65 )
66
67
68class DatabaseMemoryBackend:
69 """MemoryStoreProtocol backed by an SQL database."""
70
71 def __init__(self, provider: DatabaseProviderProtocol) -> None:
72 self._provider = provider
73
74 async def store(self, entry: MemoryEntry) -> None:
75 async with self._provider.scoped_context():
76 conn = await self._provider.get_scoped_connection()
77 await conn.execute(
78 _INSERT,
79 [
80 entry.id,
81 entry.owner_id,
82 entry.content,
83 entry.role,
84 entry.timestamp,
85 entry.importance,
86 dumps_str(entry.metadata),
87 entry.embedding,
88 ],
89 )
90
91 async def retrieve(self, query: MemoryQuery) -> list[MemorySearchResult]:
92 async with self._provider.scoped_context():
93 conn = await self._provider.get_scoped_connection()
94 query_result = await conn.execute(_SELECT_ALL, [query.owner_id])
95
96 entries = [_row_to_entry(row) for row in query_result.rows]
97 scored: list[tuple[float, MemoryEntry]] = []
98 for entry in entries:
99 age_seconds = (datetime.now(UTC) - entry.timestamp).total_seconds()
100 recency = math.exp(-age_seconds / 86400.0)
101 score = (
102 query.recency_weight * recency
103 + query.importance_weight * entry.importance
104 + query.relevance_weight * 0.5
105 )
106 if score < query.min_relevance:
107 continue
108 if query.time_range:
109 start, end = query.time_range
110 if not (start <= entry.timestamp <= end):
111 continue
112 if query.filters and not all(
113 entry.metadata.get(key) == value for key, value in query.filters.items()
114 ):
115 continue
116 scored.append((score, entry))
117
118 scored.sort(key=lambda item: item[0], reverse=True)
119 return [
120 MemorySearchResult(entry=entry, score=score, source="database")
121 for score, entry in scored[: query.top_k]
122 ]
123
124 async def get_recent(self, n: int, owner_id: str) -> list[MemoryEntry]:
125 async with self._provider.scoped_context():
126 conn = await self._provider.get_scoped_connection()
127 query_result = await conn.execute(_RECENT, [owner_id, n])
128 return [_row_to_entry(row) for row in query_result.rows]
129
130 async def delete(self, entry_id: str, owner_id: str) -> None:
131 async with self._provider.scoped_context():
132 conn = await self._provider.get_scoped_connection()
133 await conn.execute(_DELETE, [entry_id, owner_id])
134
135 async def clear(self, owner_id: str) -> None:
136 async with self._provider.scoped_context():
137 conn = await self._provider.get_scoped_connection()
138 await conn.execute(_CLEAR, [owner_id])
139
140 async def health_check(self, timeout: float = 5.0) -> HealthCheckResult:
141 provider_result = await self._provider.health_check(timeout=timeout)
142 status = (
143 HealthStatus.HEALTHY
144 if provider_result.is_healthy()
145 else HealthStatus.DEGRADED
146 )
147 return HealthCheckResult(
148 component="memory.database",
149 status=status,
150 message=provider_result.message,
151 error=provider_result.error,
152 details={"timeout": timeout},
153 )
154
155
156__all__ = ["DatabaseMemoryBackend"]