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

1"""Native async S3 backend implementation with aioboto3 and connection pooling.""" 

2 

3import asyncio 

4import json 

5import logging 

6import time 

7import uuid 

8from datetime import datetime 

9from typing import Any, AsyncIterator, Optional 

10 

11from dataknobs_config import ConfigurableBase 

12 

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 

25 

26logger = logging.getLogger(__name__) 

27 

28# Global pool manager for S3 sessions 

29_session_manager = ConnectionPoolManager() 

30 

31 

32class AsyncS3Database(AsyncDatabase, ConfigurableBase): 

33 """Native async S3 database backend with aioboto3 and session pooling.""" 

34 

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

36 """Initialize async S3 database.""" 

37 super().__init__(config) 

38 

39 if not config or "bucket" not in config: 

40 raise ValueError("S3 backend requires 'bucket' in configuration") 

41 

42 self._pool_config = S3PoolConfig.from_dict(config) 

43 self._session = None 

44 self._connected = False 

45 

46 @classmethod 

47 def from_config(cls, config: dict) -> "AsyncS3Database": 

48 """Create from config dictionary.""" 

49 return cls(config) 

50 

51 async def connect(self) -> None: 

52 """Connect to S3 service.""" 

53 if self._connected: 

54 return 

55 

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 ) 

62 

63 self._connected = True 

64 

65 async def close(self) -> None: 

66 """Close the S3 connection.""" 

67 if self._connected: 

68 self._session = None 

69 self._connected = False 

70 

71 def _initialize(self) -> None: 

72 """Initialize is handled in connect.""" 

73 pass 

74 

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

79 

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" 

85 

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 

91 

92 return { 

93 "data": data, 

94 "metadata": record.metadata or {}, 

95 "created_at": datetime.utcnow().isoformat(), 

96 "updated_at": datetime.utcnow().isoformat() 

97 } 

98 

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

103 

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

109 

110 return Record(data=data, metadata=metadata) 

111 

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

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

114 self._check_connection() 

115 

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) 

120 

121 # Add ID to metadata 

122 obj["metadata"]["id"] = id 

123 

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 ) 

131 

132 return id 

133 

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

135 """Read a record from S3.""" 

136 self._check_connection() 

137 

138 key = self._get_key(id) 

139 

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 ) 

146 

147 # Read the object body 

148 body = await response['Body'].read() 

149 obj = json.loads(body) 

150 

151 record = self._s3_object_to_record(obj) 

152 # Ensure ID is in metadata 

153 record.metadata["id"] = id 

154 

155 return record 

156 except Exception: 

157 return None 

158 

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

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

161 self._check_connection() 

162 

163 # Check if record exists 

164 if not await self.exists(id): 

165 return False 

166 

167 key = self._get_key(id) 

168 obj = self._record_to_s3_object(record) 

169 

170 # Preserve ID in metadata 

171 obj["metadata"]["id"] = id 

172 

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 ) 

180 

181 return True 

182 

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

184 """Delete a record from S3.""" 

185 self._check_connection() 

186 

187 key = self._get_key(id) 

188 

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 

198 

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

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

201 self._check_connection() 

202 

203 key = self._get_key(id) 

204 

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 

214 

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

216 """Update or insert a record with a specific ID.""" 

217 self._check_connection() 

218 

219 key = self._get_key(id) 

220 obj = self._record_to_s3_object(record) 

221 

222 # Add ID to metadata 

223 obj["metadata"]["id"] = id 

224 

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 ) 

232 

233 return id 

234 

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

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

237 self._check_connection() 

238 

239 # S3 doesn't support complex queries, so we need to list and filter 

240 records = [] 

241 

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

245 

246 params = { 

247 'Bucket': self._pool_config.bucket, 

248 } 

249 if self._pool_config.prefix: 

250 params['Prefix'] = self._pool_config.prefix 

251 

252 async for page in paginator.paginate(**params): 

253 if 'Contents' not in page: 

254 continue 

255 

256 # Process each object 

257 for obj_summary in page['Contents']: 

258 key = obj_summary['Key'] 

259 

260 # Skip non-JSON files 

261 if not key.endswith('.json'): 

262 continue 

263 

264 # Get the object 

265 response = await s3.get_object( 

266 Bucket=self._pool_config.bucket, 

267 Key=key 

268 ) 

269 

270 body = await response['Body'].read() 

271 obj = json.loads(body) 

272 record = self._s3_object_to_record(obj) 

273 

274 # Extract ID from key 

275 id = key.replace(self._pool_config.prefix + '/', '').replace('.json', '') 

276 record.metadata["id"] = id 

277 

278 # Apply filters 

279 if self._matches_filters(record, query.filters): 

280 records.append(record) 

281 

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 ) 

290 

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] 

296 

297 # Apply field projection 

298 if query.fields: 

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

300 

301 return records 

302 

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 

309 

310 value = field.value 

311 

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 

339 

340 return True 

341 

342 async def _count_all(self) -> int: 

343 """Count all records in the database.""" 

344 self._check_connection() 

345 

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

349 

350 params = { 

351 'Bucket': self._pool_config.bucket, 

352 } 

353 if self._pool_config.prefix: 

354 params['Prefix'] = self._pool_config.prefix 

355 

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 

361 

362 return count 

363 

364 async def clear(self) -> int: 

365 """Clear all records from the database.""" 

366 self._check_connection() 

367 

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

372 

373 params = { 

374 'Bucket': self._pool_config.bucket, 

375 } 

376 if self._pool_config.prefix: 

377 params['Prefix'] = self._pool_config.prefix 

378 

379 async for page in paginator.paginate(**params): 

380 if 'Contents' not in page: 

381 continue 

382 

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 

389 

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 ) 

396 

397 return count 

398 

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

407 

408 async with self._session.client("s3", endpoint_url=self._pool_config.endpoint_url) as s3: 

409 paginator = s3.get_paginator('list_objects_v2') 

410 

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 

417 

418 async for page in paginator.paginate(**params): 

419 if 'Contents' not in page: 

420 continue 

421 

422 for obj_summary in page['Contents']: 

423 key = obj_summary['Key'] 

424 

425 if not key.endswith('.json'): 

426 continue 

427 

428 # Get the object 

429 response = await s3.get_object( 

430 Bucket=self._pool_config.bucket, 

431 Key=key 

432 ) 

433 

434 body = await response['Body'].read() 

435 obj = json.loads(body) 

436 record = self._s3_object_to_record(obj) 

437 

438 # Extract ID from key 

439 id = key.replace(self._pool_config.prefix + '/', '').replace('.json', '') 

440 record.metadata["id"] = id 

441 

442 # Apply filters if query provided 

443 if query and query.filters: 

444 if not self._matches_filters(record, query.filters): 

445 continue 

446 

447 # Apply field projection 

448 if query and query.fields: 

449 record = record.project(query.fields) 

450 

451 yield record 

452 

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 

464 

465 batch = [] 

466 async for record in records: 

467 batch.append(record) 

468 

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] 

474 

475 continue_processing = await async_process_batch_with_fallback( 

476 batch, 

477 batch_func, 

478 self.create, 

479 result, 

480 config 

481 ) 

482 

483 if not continue_processing: 

484 quitting = True 

485 break 

486 

487 batch = [] 

488 

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] 

494 

495 await async_process_batch_with_fallback( 

496 batch, 

497 batch_func, 

498 self.create, 

499 result, 

500 config 

501 ) 

502 

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

504 return result 

505 

506 async def _write_batch(self, records: list[Record]) -> None: 

507 """Write a batch of records to S3.""" 

508 if not records: 

509 return 

510 

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 

520 

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) 

528 

529 # Execute all uploads concurrently 

530 await asyncio.gather(*tasks) 

531 

532 async def list_all(self) -> list[str]: 

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

534 self._check_connection() 

535 

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

539 

540 params = { 

541 'Bucket': self._pool_config.bucket, 

542 } 

543 if self._pool_config.prefix: 

544 params['Prefix'] = self._pool_config.prefix 

545 

546 async for page in paginator.paginate(**params): 

547 if 'Contents' not in page: 

548 continue 

549 

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) 

556 

557 return ids