Coverage for src/dataknobs_data/backends/s3.py: 14%
305 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"""S3 backend implementation with proper connection management."""
3import asyncio
4import json
5import logging
6import time
7from typing import Any, AsyncIterator, Dict, List, Optional, Iterator
8from uuid import uuid4
9import hashlib
10from concurrent.futures import ThreadPoolExecutor, as_completed
11from datetime import datetime
13from dataknobs_config import ConfigurableBase
14from dataknobs_data.records import Record
15from dataknobs_data.query import Query
16from dataknobs_data.database import AsyncDatabase, SyncDatabase
17from dataknobs_data.streaming import StreamConfig, StreamResult
18from dataknobs_data.streaming import process_batch_with_fallback
20logger = logging.getLogger(__name__)
23class SyncS3Database(SyncDatabase, ConfigurableBase):
24 """S3-based database backend with proper connection management.
26 Stores records as JSON objects in S3 with metadata as tags.
27 """
29 def __init__(self, config: Optional[Dict[str, Any]] = None):
30 """Initialize S3 database configuration.
32 Args:
33 config: Configuration dictionary
34 """
35 super().__init__(config)
37 # Connection state
38 self.s3_client = None
39 self._connected = False
41 # Cache for performance
42 self._index_cache = {}
43 self._cache_dirty = True
45 # Store configuration for later connection
46 self.bucket = self.config.get("bucket")
47 if not self.bucket:
48 raise ValueError("S3 bucket name is required in configuration")
50 # Optional configuration with defaults
51 self.prefix = self.config.get("prefix", "records/").rstrip("/") + "/"
52 self.region = self.config.get("region", "us-east-1")
53 self.endpoint_url = self.config.get("endpoint_url")
54 self.max_workers = self.config.get("max_workers", 10)
55 self.multipart_threshold = self.config.get("multipart_threshold", 8 * 1024 * 1024)
56 self.multipart_chunksize = self.config.get("multipart_chunksize", 8 * 1024 * 1024)
57 self.max_retries = self.config.get("max_retries", 3)
59 # AWS credentials (will use environment/IAM role if not provided)
60 self.aws_access_key_id = self.config.get("access_key_id")
61 self.aws_secret_access_key = self.config.get("secret_access_key")
62 self.aws_session_token = self.config.get("session_token")
64 @classmethod
65 def from_config(cls, config: dict) -> "SyncS3Database":
66 """Create instance from configuration dictionary."""
67 return cls(config)
69 def connect(self) -> None:
70 """Connect to S3 service."""
71 if self._connected:
72 return # Already connected
74 import boto3
75 from botocore.config import Config as BotoConfig
76 from botocore.exceptions import ClientError
78 # Configure boto3 client
79 boto_config = BotoConfig(
80 region_name=self.region,
81 max_pool_connections=self.max_workers,
82 retries={'max_attempts': self.max_retries}
83 )
85 client_kwargs = {
86 "config": boto_config,
87 "use_ssl": not bool(self.endpoint_url) # Disable SSL for local testing
88 }
90 if self.endpoint_url:
91 client_kwargs["endpoint_url"] = self.endpoint_url
93 if self.aws_access_key_id and self.aws_secret_access_key:
94 client_kwargs["aws_access_key_id"] = self.aws_access_key_id
95 client_kwargs["aws_secret_access_key"] = self.aws_secret_access_key
97 if self.aws_session_token:
98 client_kwargs["aws_session_token"] = self.aws_session_token
100 # Create S3 client
101 self.s3_client = boto3.client("s3", **client_kwargs)
102 self.ClientError = ClientError
104 # Verify bucket exists or create it
105 self._ensure_bucket_exists()
107 self._connected = True
108 logger.info(f"Connected to S3 with bucket={self.bucket}, prefix={self.prefix}")
110 def close(self) -> None:
111 """Close the S3 connection."""
112 if self.s3_client:
113 # S3 client doesn't need explicit closing, but clear cache
114 self._index_cache = {}
115 self._connected = False
116 logger.info(f"Closed S3 connection to bucket={self.bucket}")
118 def _initialize(self) -> None:
119 """Initialize method - connection setup moved to connect()."""
120 pass
122 def _check_connection(self) -> None:
123 """Check if S3 client is connected."""
124 if not self._connected or not self.s3_client:
125 raise RuntimeError("S3 not connected. Call connect() first.")
127 def _ensure_bucket_exists(self):
128 """Ensure the S3 bucket exists, create if necessary."""
129 try:
130 self.s3_client.head_bucket(Bucket=self.bucket)
131 logger.debug(f"Bucket {self.bucket} exists")
132 except self.ClientError as e:
133 error_code = e.response['Error']['Code']
134 if error_code == '404':
135 # Bucket doesn't exist, create it
136 logger.info(f"Creating bucket {self.bucket}")
137 if self.region == 'us-east-1':
138 self.s3_client.create_bucket(Bucket=self.bucket)
139 else:
140 self.s3_client.create_bucket(
141 Bucket=self.bucket,
142 CreateBucketConfiguration={'LocationConstraint': self.region}
143 )
144 else:
145 raise
147 def _get_object_key(self, record_id: str) -> str:
148 """Generate S3 object key for a record ID."""
149 return f"{self.prefix}{record_id}.json"
151 def _record_to_s3_object(self, record: Record) -> Dict[str, Any]:
152 """Convert a Record to S3 object data."""
153 data = {}
154 for field_name, field_obj in record.fields.items():
155 data[field_name] = field_obj.value
157 return {
158 "data": data,
159 "metadata": record.metadata or {}
160 }
162 def _s3_object_to_record(self, obj_data: Dict[str, Any]) -> Record:
163 """Convert S3 object data to a Record."""
164 data = obj_data.get("data", {})
165 metadata = obj_data.get("metadata", {})
166 return Record(data=data, metadata=metadata)
168 def create(self, record: Record) -> str:
169 """Create a new record in S3."""
170 self._check_connection()
172 # Use record's ID if it has one, otherwise generate a new one
173 record_id = record.id if record.id else str(uuid4())
174 key = self._get_object_key(record_id)
176 # Set metadata
177 record.metadata = record.metadata or {}
178 record.metadata["id"] = record_id
179 now = datetime.utcnow()
180 record.metadata["created_at"] = now.isoformat()
181 record.metadata["updated_at"] = now.isoformat()
183 # Convert record to JSON
184 obj_data = self._record_to_s3_object(record)
185 body = json.dumps(obj_data)
187 # Store in S3
188 self.s3_client.put_object(
189 Bucket=self.bucket,
190 Key=key,
191 Body=body,
192 ContentType='application/json'
193 )
195 # Invalidate cache
196 self._cache_dirty = True
198 logger.debug(f"Created record {record_id} at {key}")
199 return record_id
201 def read(self, id: str) -> Optional[Record]:
202 """Read a record from S3."""
203 self._check_connection()
205 key = self._get_object_key(id)
207 try:
208 response = self.s3_client.get_object(Bucket=self.bucket, Key=key)
209 body = response['Body'].read()
210 obj_data = json.loads(body)
211 return self._s3_object_to_record(obj_data)
212 except self.ClientError as e:
213 if e.response['Error']['Code'] == 'NoSuchKey':
214 return None
215 raise
217 def update(self, id: str, record: Record) -> bool:
218 """Update an existing record in S3."""
219 self._check_connection()
221 key = self._get_object_key(id)
223 # Check if exists and get existing metadata
224 try:
225 response = self.s3_client.get_object(Bucket=self.bucket, Key=key)
226 existing_data = json.loads(response['Body'].read())
227 existing_metadata = existing_data.get("metadata", {})
228 except self.ClientError as e:
229 if e.response['Error']['Code'] == 'NoSuchKey':
230 return False
231 raise
233 # Preserve and update metadata
234 record.metadata = record.metadata or {}
235 record.metadata["id"] = id
236 record.metadata["created_at"] = existing_metadata.get("created_at", datetime.utcnow().isoformat())
237 record.metadata["updated_at"] = datetime.utcnow().isoformat()
239 # Update the object
240 obj_data = self._record_to_s3_object(record)
241 body = json.dumps(obj_data)
243 self.s3_client.put_object(
244 Bucket=self.bucket,
245 Key=key,
246 Body=body,
247 ContentType='application/json'
248 )
250 # Invalidate cache
251 self._cache_dirty = True
253 logger.debug(f"Updated record {id} at {key}")
254 return True
256 def delete(self, id: str) -> bool:
257 """Delete a record from S3."""
258 self._check_connection()
260 key = self._get_object_key(id)
262 # Check if exists
263 try:
264 self.s3_client.head_object(Bucket=self.bucket, Key=key)
265 except self.ClientError as e:
266 if e.response['Error']['Code'] == '404':
267 return False
268 raise
270 # Delete the object
271 self.s3_client.delete_object(Bucket=self.bucket, Key=key)
273 # Invalidate cache
274 self._cache_dirty = True
276 logger.debug(f"Deleted record {id} at {key}")
277 return True
279 def list_all(self) -> List[str]:
280 """List all record IDs in the database.
282 Returns:
283 List of all record IDs
284 """
285 self._check_connection()
286 record_ids = []
288 # Use paginator for large buckets
289 paginator = self.s3_client.get_paginator('list_objects_v2')
290 page_iterator = paginator.paginate(
291 Bucket=self.bucket,
292 Prefix=self.prefix
293 )
295 for page in page_iterator:
296 if 'Contents' in page:
297 for obj in page['Contents']:
298 key = obj['Key']
299 # Extract record ID from key
300 if key.startswith(self.prefix) and key.endswith('.json'):
301 record_id = key[len(self.prefix):-5] # Remove prefix and .json
302 record_ids.append(record_id)
304 logger.debug(f"Listed {len(record_ids)} records from S3")
305 return record_ids
307 def exists(self, id: str) -> bool:
308 """Check if a record exists in S3."""
309 self._check_connection()
311 key = self._get_object_key(id)
313 try:
314 self.s3_client.head_object(Bucket=self.bucket, Key=key)
315 return True
316 except self.ClientError as e:
317 if e.response['Error']['Code'] == '404':
318 return False
319 raise
321 def search(self, query: Query) -> List[Record]:
322 """Search for records matching the query.
324 Note: S3 doesn't support complex queries, so we need to list and filter.
325 """
326 self._check_connection()
328 # List all objects with the prefix
329 records = []
330 paginator = self.s3_client.get_paginator('list_objects_v2')
331 pages = paginator.paginate(Bucket=self.bucket, Prefix=self.prefix)
333 for page in pages:
334 if 'Contents' not in page:
335 continue
337 # Fetch objects in parallel
338 with ThreadPoolExecutor(max_workers=self.max_workers) as executor:
339 futures = []
340 for obj in page['Contents']:
341 if obj['Key'].endswith('.json'):
342 future = executor.submit(self._fetch_and_filter, obj['Key'], query)
343 futures.append(future)
345 for future in as_completed(futures):
346 record = future.result()
347 if record:
348 records.append(record)
350 # Apply sorting if specified
351 if query.sort_specs:
352 for sort_spec in reversed(query.sort_specs):
353 reverse = sort_spec.order.value == "desc"
354 records.sort(
355 key=lambda r: r.get_value(sort_spec.field, ""),
356 reverse=reverse
357 )
359 # Apply offset and limit
360 if query.offset_value:
361 records = records[query.offset_value:]
362 if query.limit_value:
363 records = records[:query.limit_value]
365 # Apply field projection
366 if query.fields:
367 records = [r.project(query.fields) for r in records]
369 return records
371 def _fetch_and_filter(self, key: str, query: Query) -> Optional[Record]:
372 """Fetch an object and apply query filters."""
373 try:
374 response = self.s3_client.get_object(Bucket=self.bucket, Key=key)
375 body = response['Body'].read()
376 obj_data = json.loads(body)
377 record = self._s3_object_to_record(obj_data)
379 # Apply filters
380 for filter in query.filters:
381 field_value = record.get_value(filter.field)
382 if not filter.matches(field_value):
383 return None
385 return record
386 except Exception as e:
387 logger.warning(f"Error fetching {key}: {e}")
388 return None
390 def _count_all(self) -> int:
391 """Count all records in S3."""
392 self._check_connection()
394 count = 0
395 paginator = self.s3_client.get_paginator('list_objects_v2')
396 pages = paginator.paginate(Bucket=self.bucket, Prefix=self.prefix)
398 for page in pages:
399 if 'Contents' in page:
400 count += sum(1 for obj in page['Contents'] if obj['Key'].endswith('.json'))
402 return count
404 def clear(self) -> int:
405 """Clear all records from S3."""
406 self._check_connection()
408 # List and delete all objects
409 count = 0
410 paginator = self.s3_client.get_paginator('list_objects_v2')
411 pages = paginator.paginate(Bucket=self.bucket, Prefix=self.prefix)
413 for page in pages:
414 if 'Contents' not in page:
415 continue
417 # Delete in batches
418 objects = [{'Key': obj['Key']} for obj in page['Contents'] if obj['Key'].endswith('.json')]
419 if objects:
420 self.s3_client.delete_objects(
421 Bucket=self.bucket,
422 Delete={'Objects': objects}
423 )
424 count += len(objects)
426 # Clear cache
427 self._index_cache = {}
428 self._cache_dirty = True
430 logger.info(f"Cleared {count} records from S3")
431 return count
433 def stream_read(
434 self,
435 query: Optional[Query] = None,
436 config: Optional[StreamConfig] = None
437 ) -> Iterator[Record]:
438 """Stream records from S3."""
439 self._check_connection()
440 config = config or StreamConfig()
442 # List objects and stream them
443 paginator = self.s3_client.get_paginator('list_objects_v2')
444 pages = paginator.paginate(Bucket=self.bucket, Prefix=self.prefix)
446 batch = []
447 for page in pages:
448 if 'Contents' not in page:
449 continue
451 for obj in page['Contents']:
452 if not obj['Key'].endswith('.json'):
453 continue
455 record = self._fetch_and_filter(obj['Key'], query or Query())
456 if record:
457 batch.append(record)
459 if len(batch) >= config.batch_size:
460 for r in batch:
461 yield r
462 batch = []
464 # Yield remaining records
465 for r in batch:
466 yield r
468 def stream_write(
469 self,
470 records: Iterator[Record],
471 config: Optional[StreamConfig] = None
472 ) -> StreamResult:
473 """Stream records into S3."""
474 self._check_connection()
475 config = config or StreamConfig()
476 result = StreamResult()
477 start_time = time.time()
478 quitting = False
480 batch = []
481 for record in records:
482 batch.append(record)
484 if len(batch) >= config.batch_size:
485 # Write batch with graceful fallback
486 # Use lambda wrapper for _write_batch
487 continue_processing = process_batch_with_fallback(
488 batch,
489 lambda b: self._write_batch(b) or [r.id for r in b], # _write_batch returns None, we need IDs
490 self.create,
491 result,
492 config
493 )
495 if not continue_processing:
496 quitting = True
497 break
499 batch = []
501 # Write remaining batch
502 if batch and not quitting:
503 process_batch_with_fallback(
504 batch,
505 lambda b: self._write_batch(b) or [r.id for r in b],
506 self.create,
507 result,
508 config
509 )
511 result.duration = time.time() - start_time
512 return result
514 def _write_batch(self, records: List[Record]) -> None:
515 """Write a batch of records to S3."""
516 with ThreadPoolExecutor(max_workers=self.max_workers) as executor:
517 futures = []
518 for record in records:
519 record_id = str(uuid4())
520 future = executor.submit(self._write_single, record_id, record)
521 futures.append(future)
523 # Wait for all writes to complete
524 for future in as_completed(futures):
525 future.result() # This will raise if there was an error
527 def _write_single(self, record_id: str, record: Record) -> None:
528 """Write a single record to S3."""
529 # Set metadata
530 record.metadata = record.metadata or {}
531 record.metadata["id"] = record_id
532 now = datetime.utcnow()
533 record.metadata["created_at"] = now.isoformat()
534 record.metadata["updated_at"] = now.isoformat()
536 key = self._get_object_key(record_id)
537 obj_data = self._record_to_s3_object(record)
538 body = json.dumps(obj_data)
540 self.s3_client.put_object(
541 Bucket=self.bucket,
542 Key=key,
543 Body=body,
544 ContentType='application/json'
545 )
548# Import the native async implementation
549from .s3_async import AsyncS3Database