Coverage for /home/admin/Documents/AI/applications/lexigram-dev/lexigram/experimental/ai/lexigram-ai-memory/src/lexigram/ai/memory/backends/database.py: 32%

73 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-25 07:19 +0800

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"]