Coverage for src/dataknobs_data/backends/memory.py: 33%

239 statements  

« prev     ^ index     » next       coverage.py v7.10.3, created at 2025-08-17 19:59 -0500

1"""In-memory database backend implementation.""" 

2 

3import asyncio 

4import threading 

5import time 

6import uuid 

7from collections import OrderedDict 

8from typing import Any, AsyncIterator, Iterator, Optional 

9 

10from dataknobs_config import ConfigurableBase 

11 

12from ..database import AsyncDatabase, SyncDatabase 

13from ..query import Query 

14from ..query_logic import ComplexQuery 

15from ..records import Record 

16from ..streaming import AsyncStreamingMixin, StreamConfig, StreamResult, StreamingMixin 

17 

18 

19class AsyncMemoryDatabase(AsyncDatabase, AsyncStreamingMixin, ConfigurableBase): 

20 """Async in-memory database implementation.""" 

21 

22 def __init__(self, config: dict[str, Any] | None = None): 

23 super().__init__(config) 

24 self._storage: OrderedDict[str, Record] = OrderedDict() 

25 self._lock = asyncio.Lock() 

26 

27 @classmethod 

28 def from_config(cls, config: dict) -> "AsyncMemoryDatabase": 

29 """Create from config dictionary.""" 

30 return cls(config) 

31 

32 async def connect(self) -> None: 

33 """Connect to the database (no-op for memory backend).""" 

34 pass 

35 

36 def _generate_id(self) -> str: 

37 """Generate a unique ID for a record.""" 

38 return str(uuid.uuid4()) 

39 

40 async def create(self, record: Record) -> str: 

41 """Create a new record in memory.""" 

42 async with self._lock: 

43 # Use record's ID if it has one, otherwise generate a new one 

44 id = record.id if record.id else self._generate_id() 

45 self._storage[id] = record.copy(deep=True) 

46 return id 

47 

48 async def read(self, id: str) -> Record | None: 

49 """Read a record from memory.""" 

50 async with self._lock: 

51 record = self._storage.get(id) 

52 return record.copy(deep=True) if record else None 

53 

54 async def update(self, id: str, record: Record) -> bool: 

55 """Update a record in memory.""" 

56 async with self._lock: 

57 if id in self._storage: 

58 self._storage[id] = record.copy(deep=True) 

59 return True 

60 return False 

61 

62 async def delete(self, id: str) -> bool: 

63 """Delete a record from memory.""" 

64 async with self._lock: 

65 if id in self._storage: 

66 del self._storage[id] 

67 return True 

68 return False 

69 

70 async def exists(self, id: str) -> bool: 

71 """Check if a record exists in memory.""" 

72 async with self._lock: 

73 return id in self._storage 

74 

75 async def upsert(self, id: str, record: Record) -> str: 

76 """Update or insert a record with the specified ID.""" 

77 async with self._lock: 

78 self._storage[id] = record.copy(deep=True) 

79 return id 

80 

81 async def search(self, query: Query | ComplexQuery) -> list[Record]: 

82 """Search for records matching the query.""" 

83 # Handle ComplexQuery using base class implementation 

84 if isinstance(query, ComplexQuery): 

85 return await self._search_with_complex_query(query) 

86 

87 async with self._lock: 

88 results = [] 

89 

90 for id, record in self._storage.items(): 

91 # Apply filters 

92 matches = True 

93 for filter in query.filters: 

94 field_value = record.get_value(filter.field) 

95 if not filter.matches(field_value): 

96 matches = False 

97 break 

98 

99 if matches: 

100 results.append((id, record)) 

101 

102 # Apply sorting 

103 if query.sort_specs: 

104 for sort_spec in reversed(query.sort_specs): 

105 reverse = sort_spec.order.value == "desc" 

106 results.sort(key=lambda x: x[1].get_value(sort_spec.field, ""), reverse=reverse) 

107 

108 # Extract records 

109 records = [record for _, record in results] 

110 

111 # Apply offset and limit 

112 if query.offset_value: 

113 records = records[query.offset_value :] 

114 if query.limit_value: 

115 records = records[: query.limit_value] 

116 

117 # Apply field projection 

118 if query.fields: 

119 projected_records = [] 

120 for record in records: 

121 projected_records.append(record.project(query.fields)) 

122 records = projected_records 

123 

124 # Return deep copies 

125 return [record.copy(deep=True) for record in records] 

126 

127 async def _count_all(self) -> int: 

128 """Count all records in memory.""" 

129 async with self._lock: 

130 return len(self._storage) 

131 

132 async def clear(self) -> int: 

133 """Clear all records from memory.""" 

134 async with self._lock: 

135 count = len(self._storage) 

136 self._storage.clear() 

137 return count 

138 

139 async def create_batch(self, records: list[Record]) -> list[str]: 

140 """Create multiple records efficiently.""" 

141 async with self._lock: 

142 ids = [] 

143 for record in records: 

144 # Use record's ID if it has one, otherwise generate a new one 

145 id = record.id if record.id else self._generate_id() 

146 self._storage[id] = record.copy(deep=True) 

147 ids.append(id) 

148 return ids 

149 

150 async def read_batch(self, ids: list[str]) -> list[Record | None]: 

151 """Read multiple records efficiently.""" 

152 async with self._lock: 

153 results = [] 

154 for id in ids: 

155 record = self._storage.get(id) 

156 results.append(record.copy(deep=True) if record else None) 

157 return results 

158 

159 async def delete_batch(self, ids: list[str]) -> list[bool]: 

160 """Delete multiple records efficiently.""" 

161 async with self._lock: 

162 results = [] 

163 for id in ids: 

164 if id in self._storage: 

165 del self._storage[id] 

166 results.append(True) 

167 else: 

168 results.append(False) 

169 return results 

170 

171 async def stream_read( 

172 self, 

173 query: Optional[Query] = None, 

174 config: Optional[StreamConfig] = None 

175 ) -> AsyncIterator[Record]: 

176 """Stream records from memory.""" 

177 config = config or StreamConfig() 

178 

179 # Get all matching records 

180 if query: 

181 records = await self.search(query) 

182 else: 

183 async with self._lock: 

184 records = list(self._storage.values()) 

185 

186 # Yield records in batches 

187 for i in range(0, len(records), config.batch_size): 

188 batch = records[i:i + config.batch_size] 

189 for record in batch: 

190 yield record.copy(deep=True) 

191 # Small yield to prevent blocking 

192 await asyncio.sleep(0) 

193 

194 async def stream_write( 

195 self, 

196 records: AsyncIterator[Record], 

197 config: Optional[StreamConfig] = None 

198 ) -> StreamResult: 

199 """Stream records into memory.""" 

200 # Use the default implementation from mixin 

201 return await self._default_stream_write(records, config) 

202 

203 

204class SyncMemoryDatabase(SyncDatabase, StreamingMixin, ConfigurableBase): 

205 """Synchronous in-memory database implementation.""" 

206 

207 def __init__(self, config: dict[str, Any] | None = None): 

208 super().__init__(config) 

209 self._storage: OrderedDict[str, Record] = OrderedDict() 

210 self._lock = threading.RLock() 

211 

212 @classmethod 

213 def from_config(cls, config: dict) -> "SyncMemoryDatabase": 

214 """Create from config dictionary.""" 

215 return cls(config) 

216 

217 def connect(self) -> None: 

218 """Connect to the database (no-op for memory backend).""" 

219 pass 

220 

221 def _generate_id(self) -> str: 

222 """Generate a unique ID for a record.""" 

223 return str(uuid.uuid4()) 

224 

225 def create(self, record: Record) -> str: 

226 """Create a new record in memory.""" 

227 with self._lock: 

228 # Use record's ID if it has one, otherwise generate a new one 

229 id = record.id if record.id else self._generate_id() 

230 self._storage[id] = record.copy(deep=True) 

231 return id 

232 

233 def read(self, id: str) -> Record | None: 

234 """Read a record from memory.""" 

235 with self._lock: 

236 record = self._storage.get(id) 

237 return record.copy(deep=True) if record else None 

238 

239 def update(self, id: str, record: Record) -> bool: 

240 """Update a record in memory.""" 

241 with self._lock: 

242 if id in self._storage: 

243 self._storage[id] = record.copy(deep=True) 

244 return True 

245 return False 

246 

247 def delete(self, id: str) -> bool: 

248 """Delete a record from memory.""" 

249 with self._lock: 

250 if id in self._storage: 

251 del self._storage[id] 

252 return True 

253 return False 

254 

255 def exists(self, id: str) -> bool: 

256 """Check if a record exists in memory.""" 

257 with self._lock: 

258 return id in self._storage 

259 

260 def upsert(self, id: str, record: Record) -> str: 

261 """Update or insert a record with the specified ID.""" 

262 with self._lock: 

263 self._storage[id] = record.copy(deep=True) 

264 return id 

265 

266 def search(self, query: Query | ComplexQuery) -> list[Record]: 

267 """Search for records matching the query.""" 

268 # Handle ComplexQuery using base class implementation 

269 if isinstance(query, ComplexQuery): 

270 return self._search_with_complex_query(query) 

271 

272 with self._lock: 

273 results = [] 

274 

275 for id, record in self._storage.items(): 

276 # Apply filters 

277 matches = True 

278 for filter in query.filters: 

279 field_value = record.get_value(filter.field) 

280 if not filter.matches(field_value): 

281 matches = False 

282 break 

283 

284 if matches: 

285 results.append((id, record)) 

286 

287 # Apply sorting 

288 if query.sort_specs: 

289 for sort_spec in reversed(query.sort_specs): 

290 reverse = sort_spec.order.value == "desc" 

291 results.sort(key=lambda x: x[1].get_value(sort_spec.field, ""), reverse=reverse) 

292 

293 # Extract records 

294 records = [record for _, record in results] 

295 

296 # Apply offset and limit 

297 if query.offset_value: 

298 records = records[query.offset_value :] 

299 if query.limit_value: 

300 records = records[: query.limit_value] 

301 

302 # Apply field projection 

303 if query.fields: 

304 projected_records = [] 

305 for record in records: 

306 projected_records.append(record.project(query.fields)) 

307 records = projected_records 

308 

309 # Return deep copies 

310 return [record.copy(deep=True) for record in records] 

311 

312 def _count_all(self) -> int: 

313 """Count all records in memory.""" 

314 with self._lock: 

315 return len(self._storage) 

316 

317 def clear(self) -> int: 

318 """Clear all records from memory.""" 

319 with self._lock: 

320 count = len(self._storage) 

321 self._storage.clear() 

322 return count 

323 

324 def create_batch(self, records: list[Record]) -> list[str]: 

325 """Create multiple records efficiently.""" 

326 with self._lock: 

327 ids = [] 

328 for record in records: 

329 # Use record's ID if it has one, otherwise generate a new one 

330 id = record.id if record.id else self._generate_id() 

331 self._storage[id] = record.copy(deep=True) 

332 ids.append(id) 

333 return ids 

334 

335 def read_batch(self, ids: list[str]) -> list[Record | None]: 

336 """Read multiple records efficiently.""" 

337 with self._lock: 

338 results = [] 

339 for id in ids: 

340 record = self._storage.get(id) 

341 results.append(record.copy(deep=True) if record else None) 

342 return results 

343 

344 def delete_batch(self, ids: list[str]) -> list[bool]: 

345 """Delete multiple records efficiently.""" 

346 with self._lock: 

347 results = [] 

348 for id in ids: 

349 if id in self._storage: 

350 del self._storage[id] 

351 results.append(True) 

352 else: 

353 results.append(False) 

354 return results 

355 

356 def stream_read( 

357 self, 

358 query: Optional[Query] = None, 

359 config: Optional[StreamConfig] = None 

360 ) -> Iterator[Record]: 

361 """Stream records from memory.""" 

362 config = config or StreamConfig() 

363 

364 # Get all matching records 

365 if query: 

366 records = self.search(query) 

367 else: 

368 with self._lock: 

369 records = list(self._storage.values()) 

370 

371 # Yield records in batches 

372 for i in range(0, len(records), config.batch_size): 

373 batch = records[i:i + config.batch_size] 

374 for record in batch: 

375 yield record.copy(deep=True) 

376 

377 def stream_write( 

378 self, 

379 records: Iterator[Record], 

380 config: Optional[StreamConfig] = None 

381 ) -> StreamResult: 

382 """Stream records into memory.""" 

383 # Use the default implementation from mixin 

384 return self._default_stream_write(records, config)