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
« prev ^ index » next coverage.py v7.10.3, created at 2025-08-17 19:59 -0500
1"""In-memory database backend implementation."""
3import asyncio
4import threading
5import time
6import uuid
7from collections import OrderedDict
8from typing import Any, AsyncIterator, Iterator, Optional
10from dataknobs_config import ConfigurableBase
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
19class AsyncMemoryDatabase(AsyncDatabase, AsyncStreamingMixin, ConfigurableBase):
20 """Async in-memory database implementation."""
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()
27 @classmethod
28 def from_config(cls, config: dict) -> "AsyncMemoryDatabase":
29 """Create from config dictionary."""
30 return cls(config)
32 async def connect(self) -> None:
33 """Connect to the database (no-op for memory backend)."""
34 pass
36 def _generate_id(self) -> str:
37 """Generate a unique ID for a record."""
38 return str(uuid.uuid4())
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
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
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
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
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
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
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)
87 async with self._lock:
88 results = []
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
99 if matches:
100 results.append((id, record))
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)
108 # Extract records
109 records = [record for _, record in results]
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]
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
124 # Return deep copies
125 return [record.copy(deep=True) for record in records]
127 async def _count_all(self) -> int:
128 """Count all records in memory."""
129 async with self._lock:
130 return len(self._storage)
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
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
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
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
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()
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())
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)
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)
204class SyncMemoryDatabase(SyncDatabase, StreamingMixin, ConfigurableBase):
205 """Synchronous in-memory database implementation."""
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()
212 @classmethod
213 def from_config(cls, config: dict) -> "SyncMemoryDatabase":
214 """Create from config dictionary."""
215 return cls(config)
217 def connect(self) -> None:
218 """Connect to the database (no-op for memory backend)."""
219 pass
221 def _generate_id(self) -> str:
222 """Generate a unique ID for a record."""
223 return str(uuid.uuid4())
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
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
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
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
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
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
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)
272 with self._lock:
273 results = []
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
284 if matches:
285 results.append((id, record))
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)
293 # Extract records
294 records = [record for _, record in results]
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]
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
309 # Return deep copies
310 return [record.copy(deep=True) for record in records]
312 def _count_all(self) -> int:
313 """Count all records in memory."""
314 with self._lock:
315 return len(self._storage)
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
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
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
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
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()
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())
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)
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)