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

1"""S3 backend implementation with proper connection management.""" 

2 

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 

12 

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 

19 

20logger = logging.getLogger(__name__) 

21 

22 

23class SyncS3Database(SyncDatabase, ConfigurableBase): 

24 """S3-based database backend with proper connection management. 

25  

26 Stores records as JSON objects in S3 with metadata as tags. 

27 """ 

28 

29 def __init__(self, config: Optional[Dict[str, Any]] = None): 

30 """Initialize S3 database configuration. 

31  

32 Args: 

33 config: Configuration dictionary 

34 """ 

35 super().__init__(config) 

36 

37 # Connection state 

38 self.s3_client = None 

39 self._connected = False 

40 

41 # Cache for performance 

42 self._index_cache = {} 

43 self._cache_dirty = True 

44 

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

49 

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) 

58 

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

63 

64 @classmethod 

65 def from_config(cls, config: dict) -> "SyncS3Database": 

66 """Create instance from configuration dictionary.""" 

67 return cls(config) 

68 

69 def connect(self) -> None: 

70 """Connect to S3 service.""" 

71 if self._connected: 

72 return # Already connected 

73 

74 import boto3 

75 from botocore.config import Config as BotoConfig 

76 from botocore.exceptions import ClientError 

77 

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 ) 

84 

85 client_kwargs = { 

86 "config": boto_config, 

87 "use_ssl": not bool(self.endpoint_url) # Disable SSL for local testing 

88 } 

89 

90 if self.endpoint_url: 

91 client_kwargs["endpoint_url"] = self.endpoint_url 

92 

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 

96 

97 if self.aws_session_token: 

98 client_kwargs["aws_session_token"] = self.aws_session_token 

99 

100 # Create S3 client 

101 self.s3_client = boto3.client("s3", **client_kwargs) 

102 self.ClientError = ClientError 

103 

104 # Verify bucket exists or create it 

105 self._ensure_bucket_exists() 

106 

107 self._connected = True 

108 logger.info(f"Connected to S3 with bucket={self.bucket}, prefix={self.prefix}") 

109 

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}") 

117 

118 def _initialize(self) -> None: 

119 """Initialize method - connection setup moved to connect().""" 

120 pass 

121 

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.") 

126 

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 

146 

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" 

150 

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 

156 

157 return { 

158 "data": data, 

159 "metadata": record.metadata or {} 

160 } 

161 

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) 

167 

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

169 """Create a new record in S3.""" 

170 self._check_connection() 

171 

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) 

175 

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() 

182 

183 # Convert record to JSON 

184 obj_data = self._record_to_s3_object(record) 

185 body = json.dumps(obj_data) 

186 

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 ) 

194 

195 # Invalidate cache 

196 self._cache_dirty = True 

197 

198 logger.debug(f"Created record {record_id} at {key}") 

199 return record_id 

200 

201 def read(self, id: str) -> Optional[Record]: 

202 """Read a record from S3.""" 

203 self._check_connection() 

204 

205 key = self._get_object_key(id) 

206 

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 

216 

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

218 """Update an existing record in S3.""" 

219 self._check_connection() 

220 

221 key = self._get_object_key(id) 

222 

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 

232 

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() 

238 

239 # Update the object 

240 obj_data = self._record_to_s3_object(record) 

241 body = json.dumps(obj_data) 

242 

243 self.s3_client.put_object( 

244 Bucket=self.bucket, 

245 Key=key, 

246 Body=body, 

247 ContentType='application/json' 

248 ) 

249 

250 # Invalidate cache 

251 self._cache_dirty = True 

252 

253 logger.debug(f"Updated record {id} at {key}") 

254 return True 

255 

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

257 """Delete a record from S3.""" 

258 self._check_connection() 

259 

260 key = self._get_object_key(id) 

261 

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 

269 

270 # Delete the object 

271 self.s3_client.delete_object(Bucket=self.bucket, Key=key) 

272 

273 # Invalidate cache 

274 self._cache_dirty = True 

275 

276 logger.debug(f"Deleted record {id} at {key}") 

277 return True 

278 

279 def list_all(self) -> List[str]: 

280 """List all record IDs in the database. 

281  

282 Returns: 

283 List of all record IDs 

284 """ 

285 self._check_connection() 

286 record_ids = [] 

287 

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 ) 

294 

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) 

303 

304 logger.debug(f"Listed {len(record_ids)} records from S3") 

305 return record_ids 

306 

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

308 """Check if a record exists in S3.""" 

309 self._check_connection() 

310 

311 key = self._get_object_key(id) 

312 

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 

320 

321 def search(self, query: Query) -> List[Record]: 

322 """Search for records matching the query. 

323  

324 Note: S3 doesn't support complex queries, so we need to list and filter. 

325 """ 

326 self._check_connection() 

327 

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) 

332 

333 for page in pages: 

334 if 'Contents' not in page: 

335 continue 

336 

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) 

344 

345 for future in as_completed(futures): 

346 record = future.result() 

347 if record: 

348 records.append(record) 

349 

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 ) 

358 

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] 

364 

365 # Apply field projection 

366 if query.fields: 

367 records = [r.project(query.fields) for r in records] 

368 

369 return records 

370 

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) 

378 

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 

384 

385 return record 

386 except Exception as e: 

387 logger.warning(f"Error fetching {key}: {e}") 

388 return None 

389 

390 def _count_all(self) -> int: 

391 """Count all records in S3.""" 

392 self._check_connection() 

393 

394 count = 0 

395 paginator = self.s3_client.get_paginator('list_objects_v2') 

396 pages = paginator.paginate(Bucket=self.bucket, Prefix=self.prefix) 

397 

398 for page in pages: 

399 if 'Contents' in page: 

400 count += sum(1 for obj in page['Contents'] if obj['Key'].endswith('.json')) 

401 

402 return count 

403 

404 def clear(self) -> int: 

405 """Clear all records from S3.""" 

406 self._check_connection() 

407 

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) 

412 

413 for page in pages: 

414 if 'Contents' not in page: 

415 continue 

416 

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) 

425 

426 # Clear cache 

427 self._index_cache = {} 

428 self._cache_dirty = True 

429 

430 logger.info(f"Cleared {count} records from S3") 

431 return count 

432 

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() 

441 

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) 

445 

446 batch = [] 

447 for page in pages: 

448 if 'Contents' not in page: 

449 continue 

450 

451 for obj in page['Contents']: 

452 if not obj['Key'].endswith('.json'): 

453 continue 

454 

455 record = self._fetch_and_filter(obj['Key'], query or Query()) 

456 if record: 

457 batch.append(record) 

458 

459 if len(batch) >= config.batch_size: 

460 for r in batch: 

461 yield r 

462 batch = [] 

463 

464 # Yield remaining records 

465 for r in batch: 

466 yield r 

467 

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 

479 

480 batch = [] 

481 for record in records: 

482 batch.append(record) 

483 

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 ) 

494 

495 if not continue_processing: 

496 quitting = True 

497 break 

498 

499 batch = [] 

500 

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 ) 

510 

511 result.duration = time.time() - start_time 

512 return result 

513 

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) 

522 

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 

526 

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() 

535 

536 key = self._get_object_key(record_id) 

537 obj_data = self._record_to_s3_object(record) 

538 body = json.dumps(obj_data) 

539 

540 self.s3_client.put_object( 

541 Bucket=self.bucket, 

542 Key=key, 

543 Body=body, 

544 ContentType='application/json' 

545 ) 

546 

547 

548# Import the native async implementation  

549from .s3_async import AsyncS3Database