Coverage for src/dataknobs_data/backends/s3_async.py: 14%
301 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"""Native async S3 backend implementation with aioboto3 and connection pooling."""
3import asyncio
4import json
5import logging
6import time
7import uuid
8from datetime import datetime
9from typing import Any, AsyncIterator, Optional
11from dataknobs_config import ConfigurableBase
13from ..database import AsyncDatabase
14from ..exceptions import DatabaseError
15from ..pooling import ConnectionPoolManager
16from ..pooling.s3 import (
17 S3PoolConfig,
18 create_aioboto3_session,
19 validate_s3_session
20)
21from ..query import Operator, Query, SortOrder
22from ..records import Record
23from ..streaming import StreamConfig, StreamResult
24from ..streaming import async_process_batch_with_fallback
26logger = logging.getLogger(__name__)
28# Global pool manager for S3 sessions
29_session_manager = ConnectionPoolManager()
32class AsyncS3Database(AsyncDatabase, ConfigurableBase):
33 """Native async S3 database backend with aioboto3 and session pooling."""
35 def __init__(self, config: dict[str, Any] | None = None):
36 """Initialize async S3 database."""
37 super().__init__(config)
39 if not config or "bucket" not in config:
40 raise ValueError("S3 backend requires 'bucket' in configuration")
42 self._pool_config = S3PoolConfig.from_dict(config)
43 self._session = None
44 self._connected = False
46 @classmethod
47 def from_config(cls, config: dict) -> "AsyncS3Database":
48 """Create from config dictionary."""
49 return cls(config)
51 async def connect(self) -> None:
52 """Connect to S3 service."""
53 if self._connected:
54 return
56 # Get or create session for current event loop
57 self._session = await _session_manager.get_pool(
58 self._pool_config,
59 create_aioboto3_session,
60 lambda session: validate_s3_session(session, self._pool_config)
61 )
63 self._connected = True
65 async def close(self) -> None:
66 """Close the S3 connection."""
67 if self._connected:
68 self._session = None
69 self._connected = False
71 def _initialize(self) -> None:
72 """Initialize is handled in connect."""
73 pass
75 def _check_connection(self) -> None:
76 """Check if database is connected."""
77 if not self._connected or not self._session:
78 raise RuntimeError("Database not connected. Call connect() first.")
80 def _get_key(self, id: str) -> str:
81 """Get the S3 key for a given record ID."""
82 if self._pool_config.prefix:
83 return f"{self._pool_config.prefix}/{id}.json"
84 return f"{id}.json"
86 def _record_to_s3_object(self, record: Record) -> dict[str, Any]:
87 """Convert a Record to an S3 object."""
88 data = {}
89 for field_name, field_obj in record.fields.items():
90 data[field_name] = field_obj.value
92 return {
93 "data": data,
94 "metadata": record.metadata or {},
95 "created_at": datetime.utcnow().isoformat(),
96 "updated_at": datetime.utcnow().isoformat()
97 }
99 def _s3_object_to_record(self, obj: dict[str, Any]) -> Record:
100 """Convert an S3 object to a Record."""
101 data = obj.get("data", {})
102 metadata = obj.get("metadata", {})
104 # Add timestamps to metadata
105 if "created_at" in obj:
106 metadata["created_at"] = obj["created_at"]
107 if "updated_at" in obj:
108 metadata["updated_at"] = obj["updated_at"]
110 return Record(data=data, metadata=metadata)
112 async def create(self, record: Record) -> str:
113 """Create a new record in S3."""
114 self._check_connection()
116 # Use record's ID if it has one, otherwise generate a new one
117 id = record.id if record.id else str(uuid.uuid4())
118 key = self._get_key(id)
119 obj = self._record_to_s3_object(record)
121 # Add ID to metadata
122 obj["metadata"]["id"] = id
124 async with self._session.client("s3", endpoint_url=self._pool_config.endpoint_url) as s3:
125 await s3.put_object(
126 Bucket=self._pool_config.bucket,
127 Key=key,
128 Body=json.dumps(obj),
129 ContentType="application/json"
130 )
132 return id
134 async def read(self, id: str) -> Record | None:
135 """Read a record from S3."""
136 self._check_connection()
138 key = self._get_key(id)
140 try:
141 async with self._session.client("s3", endpoint_url=self._pool_config.endpoint_url) as s3:
142 response = await s3.get_object(
143 Bucket=self._pool_config.bucket,
144 Key=key
145 )
147 # Read the object body
148 body = await response['Body'].read()
149 obj = json.loads(body)
151 record = self._s3_object_to_record(obj)
152 # Ensure ID is in metadata
153 record.metadata["id"] = id
155 return record
156 except Exception:
157 return None
159 async def update(self, id: str, record: Record) -> bool:
160 """Update an existing record in S3."""
161 self._check_connection()
163 # Check if record exists
164 if not await self.exists(id):
165 return False
167 key = self._get_key(id)
168 obj = self._record_to_s3_object(record)
170 # Preserve ID in metadata
171 obj["metadata"]["id"] = id
173 async with self._session.client("s3", endpoint_url=self._pool_config.endpoint_url) as s3:
174 await s3.put_object(
175 Bucket=self._pool_config.bucket,
176 Key=key,
177 Body=json.dumps(obj),
178 ContentType="application/json"
179 )
181 return True
183 async def delete(self, id: str) -> bool:
184 """Delete a record from S3."""
185 self._check_connection()
187 key = self._get_key(id)
189 try:
190 async with self._session.client("s3", endpoint_url=self._pool_config.endpoint_url) as s3:
191 await s3.delete_object(
192 Bucket=self._pool_config.bucket,
193 Key=key
194 )
195 return True
196 except Exception:
197 return False
199 async def exists(self, id: str) -> bool:
200 """Check if a record exists in S3."""
201 self._check_connection()
203 key = self._get_key(id)
205 try:
206 async with self._session.client("s3", endpoint_url=self._pool_config.endpoint_url) as s3:
207 await s3.head_object(
208 Bucket=self._pool_config.bucket,
209 Key=key
210 )
211 return True
212 except Exception:
213 return False
215 async def upsert(self, id: str, record: Record) -> str:
216 """Update or insert a record with a specific ID."""
217 self._check_connection()
219 key = self._get_key(id)
220 obj = self._record_to_s3_object(record)
222 # Add ID to metadata
223 obj["metadata"]["id"] = id
225 async with self._session.client("s3", endpoint_url=self._pool_config.endpoint_url) as s3:
226 await s3.put_object(
227 Bucket=self._pool_config.bucket,
228 Key=key,
229 Body=json.dumps(obj),
230 ContentType="application/json"
231 )
233 return id
235 async def search(self, query: Query) -> list[Record]:
236 """Search for records matching the query."""
237 self._check_connection()
239 # S3 doesn't support complex queries, so we need to list and filter
240 records = []
242 async with self._session.client("s3", endpoint_url=self._pool_config.endpoint_url) as s3:
243 # List all objects
244 paginator = s3.get_paginator('list_objects_v2')
246 params = {
247 'Bucket': self._pool_config.bucket,
248 }
249 if self._pool_config.prefix:
250 params['Prefix'] = self._pool_config.prefix
252 async for page in paginator.paginate(**params):
253 if 'Contents' not in page:
254 continue
256 # Process each object
257 for obj_summary in page['Contents']:
258 key = obj_summary['Key']
260 # Skip non-JSON files
261 if not key.endswith('.json'):
262 continue
264 # Get the object
265 response = await s3.get_object(
266 Bucket=self._pool_config.bucket,
267 Key=key
268 )
270 body = await response['Body'].read()
271 obj = json.loads(body)
272 record = self._s3_object_to_record(obj)
274 # Extract ID from key
275 id = key.replace(self._pool_config.prefix + '/', '').replace('.json', '')
276 record.metadata["id"] = id
278 # Apply filters
279 if self._matches_filters(record, query.filters):
280 records.append(record)
282 # Apply sorting
283 if query.sort_specs:
284 for sort_spec in reversed(query.sort_specs):
285 reverse = sort_spec.order == SortOrder.DESC
286 records.sort(
287 key=lambda r: r.get_field(sort_spec.field).value if r.get_field(sort_spec.field) else None,
288 reverse=reverse
289 )
291 # Apply offset and limit
292 if query.offset_value:
293 records = records[query.offset_value:]
294 if query.limit_value:
295 records = records[:query.limit_value]
297 # Apply field projection
298 if query.fields:
299 records = [r.project(query.fields) for r in records]
301 return records
303 def _matches_filters(self, record: Record, filters: list) -> bool:
304 """Check if a record matches all filters."""
305 for filter in filters:
306 field = record.get_field(filter.field)
307 if not field:
308 return False
310 value = field.value
312 if filter.operator == Operator.EQ:
313 if value != filter.value:
314 return False
315 elif filter.operator == Operator.NEQ:
316 if value == filter.value:
317 return False
318 elif filter.operator == Operator.GT:
319 if value <= filter.value:
320 return False
321 elif filter.operator == Operator.LT:
322 if value >= filter.value:
323 return False
324 elif filter.operator == Operator.GTE:
325 if value < filter.value:
326 return False
327 elif filter.operator == Operator.LTE:
328 if value > filter.value:
329 return False
330 elif filter.operator == Operator.LIKE:
331 if not str(filter.value) in str(value):
332 return False
333 elif filter.operator == Operator.IN:
334 if value not in filter.value:
335 return False
336 elif filter.operator == Operator.NOT_IN:
337 if value in filter.value:
338 return False
340 return True
342 async def _count_all(self) -> int:
343 """Count all records in the database."""
344 self._check_connection()
346 count = 0
347 async with self._session.client("s3", endpoint_url=self._pool_config.endpoint_url) as s3:
348 paginator = s3.get_paginator('list_objects_v2')
350 params = {
351 'Bucket': self._pool_config.bucket,
352 }
353 if self._pool_config.prefix:
354 params['Prefix'] = self._pool_config.prefix
356 async for page in paginator.paginate(**params):
357 if 'Contents' in page:
358 for obj in page['Contents']:
359 if obj['Key'].endswith('.json'):
360 count += 1
362 return count
364 async def clear(self) -> int:
365 """Clear all records from the database."""
366 self._check_connection()
368 count = 0
369 async with self._session.client("s3", endpoint_url=self._pool_config.endpoint_url) as s3:
370 # List and delete all objects
371 paginator = s3.get_paginator('list_objects_v2')
373 params = {
374 'Bucket': self._pool_config.bucket,
375 }
376 if self._pool_config.prefix:
377 params['Prefix'] = self._pool_config.prefix
379 async for page in paginator.paginate(**params):
380 if 'Contents' not in page:
381 continue
383 # Build delete request
384 objects_to_delete = []
385 for obj in page['Contents']:
386 if obj['Key'].endswith('.json'):
387 objects_to_delete.append({'Key': obj['Key']})
388 count += 1
390 # Delete in batch
391 if objects_to_delete:
392 await s3.delete_objects(
393 Bucket=self._pool_config.bucket,
394 Delete={'Objects': objects_to_delete}
395 )
397 return count
399 async def stream_read(
400 self,
401 query: Optional[Query] = None,
402 config: Optional[StreamConfig] = None
403 ) -> AsyncIterator[Record]:
404 """Stream records from S3."""
405 self._check_connection()
406 config = config or StreamConfig()
408 async with self._session.client("s3", endpoint_url=self._pool_config.endpoint_url) as s3:
409 paginator = s3.get_paginator('list_objects_v2')
411 params = {
412 'Bucket': self._pool_config.bucket,
413 'MaxKeys': config.batch_size
414 }
415 if self._pool_config.prefix:
416 params['Prefix'] = self._pool_config.prefix
418 async for page in paginator.paginate(**params):
419 if 'Contents' not in page:
420 continue
422 for obj_summary in page['Contents']:
423 key = obj_summary['Key']
425 if not key.endswith('.json'):
426 continue
428 # Get the object
429 response = await s3.get_object(
430 Bucket=self._pool_config.bucket,
431 Key=key
432 )
434 body = await response['Body'].read()
435 obj = json.loads(body)
436 record = self._s3_object_to_record(obj)
438 # Extract ID from key
439 id = key.replace(self._pool_config.prefix + '/', '').replace('.json', '')
440 record.metadata["id"] = id
442 # Apply filters if query provided
443 if query and query.filters:
444 if not self._matches_filters(record, query.filters):
445 continue
447 # Apply field projection
448 if query and query.fields:
449 record = record.project(query.fields)
451 yield record
453 async def stream_write(
454 self,
455 records: AsyncIterator[Record],
456 config: Optional[StreamConfig] = None
457 ) -> StreamResult:
458 """Stream records into S3."""
459 self._check_connection()
460 config = config or StreamConfig()
461 result = StreamResult()
462 start_time = time.time()
463 quitting = False
465 batch = []
466 async for record in records:
467 batch.append(record)
469 if len(batch) >= config.batch_size:
470 # Write batch with graceful fallback
471 async def batch_func(b):
472 await self._write_batch(b)
473 return [r.id for r in b]
475 continue_processing = await async_process_batch_with_fallback(
476 batch,
477 batch_func,
478 self.create,
479 result,
480 config
481 )
483 if not continue_processing:
484 quitting = True
485 break
487 batch = []
489 # Write remaining batch
490 if batch and not quitting:
491 async def batch_func(b):
492 await self._write_batch(b)
493 return [r.id for r in b]
495 await async_process_batch_with_fallback(
496 batch,
497 batch_func,
498 self.create,
499 result,
500 config
501 )
503 result.duration = time.time() - start_time
504 return result
506 async def _write_batch(self, records: list[Record]) -> None:
507 """Write a batch of records to S3."""
508 if not records:
509 return
511 async with self._session.client("s3", endpoint_url=self._pool_config.endpoint_url) as s3:
512 # Write each record (S3 doesn't have native batch write)
513 # We could potentially use multipart upload for very large batches
514 tasks = []
515 for record in records:
516 id = str(uuid.uuid4())
517 key = self._get_key(id)
518 obj = self._record_to_s3_object(record)
519 obj["metadata"]["id"] = id
521 task = s3.put_object(
522 Bucket=self._pool_config.bucket,
523 Key=key,
524 Body=json.dumps(obj),
525 ContentType="application/json"
526 )
527 tasks.append(task)
529 # Execute all uploads concurrently
530 await asyncio.gather(*tasks)
532 async def list_all(self) -> list[str]:
533 """List all record IDs in the database."""
534 self._check_connection()
536 ids = []
537 async with self._session.client("s3", endpoint_url=self._pool_config.endpoint_url) as s3:
538 paginator = s3.get_paginator('list_objects_v2')
540 params = {
541 'Bucket': self._pool_config.bucket,
542 }
543 if self._pool_config.prefix:
544 params['Prefix'] = self._pool_config.prefix
546 async for page in paginator.paginate(**params):
547 if 'Contents' not in page:
548 continue
550 for obj in page['Contents']:
551 key = obj['Key']
552 if key.endswith('.json'):
553 # Extract ID from key
554 id = key.replace(self._pool_config.prefix + '/', '').replace('.json', '')
555 ids.append(id)
557 return ids