Coverage for src/dataknobs_data/backends/elasticsearch_async.py: 17%

237 statements  

« prev     ^ index     » next       coverage.py v7.10.3, created at 2025-08-17 19:59 -0500

1"""Native async Elasticsearch backend implementation with connection pooling.""" 

2 

3import asyncio 

4import json 

5import logging 

6import time 

7import uuid 

8from typing import Any, AsyncIterator, Optional 

9 

10from dataknobs_config import ConfigurableBase 

11 

12from ..database import AsyncDatabase 

13from ..exceptions import DatabaseError 

14from ..pooling import ConnectionPoolManager 

15from ..pooling.elasticsearch import ( 

16 ElasticsearchPoolConfig, 

17 create_async_elasticsearch_client, 

18 validate_elasticsearch_client, 

19 close_elasticsearch_client 

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 Elasticsearch clients 

29_client_manager = ConnectionPoolManager() 

30 

31 

32class AsyncElasticsearchDatabase(AsyncDatabase, ConfigurableBase): 

33 """Native async Elasticsearch database backend with connection pooling.""" 

34 

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

36 """Initialize async Elasticsearch database.""" 

37 super().__init__(config) 

38 config = config or {} 

39 self._pool_config = ElasticsearchPoolConfig.from_dict(config) 

40 self.index_name = self._pool_config.index 

41 self.refresh = config.get("refresh", True) 

42 self._client = None 

43 self._connected = False 

44 

45 @classmethod 

46 def from_config(cls, config: dict) -> "AsyncElasticsearchDatabase": 

47 """Create from config dictionary.""" 

48 return cls(config) 

49 

50 async def connect(self) -> None: 

51 """Connect to the Elasticsearch database.""" 

52 if self._connected: 

53 return 

54 

55 # Get or create client for current event loop 

56 self._client = await _client_manager.get_pool( 

57 self._pool_config, 

58 create_async_elasticsearch_client, 

59 validate_elasticsearch_client, 

60 close_elasticsearch_client 

61 ) 

62 

63 # Ensure index exists 

64 await self._ensure_index() 

65 self._connected = True 

66 

67 async def close(self) -> None: 

68 """Close the database connection.""" 

69 if self._connected: 

70 # Note: The client is managed by the pool manager, so we don't close it here 

71 # Just mark as disconnected 

72 self._client = None 

73 self._connected = False 

74 

75 def _initialize(self) -> None: 

76 """Initialize is handled in connect.""" 

77 pass 

78 

79 async def _ensure_index(self) -> None: 

80 """Ensure the index exists with proper mappings.""" 

81 if not self._client: 

82 raise RuntimeError("Database not connected. Call connect() first.") 

83 

84 # Check if index exists 

85 if not await self._client.indices.exists(index=self.index_name): 

86 # Create index with mappings 

87 mappings = { 

88 "properties": { 

89 "data": {"type": "object", "enabled": True}, 

90 "metadata": {"type": "object", "enabled": True}, 

91 "created_at": {"type": "date"}, 

92 "updated_at": {"type": "date"} 

93 } 

94 } 

95 

96 await self._client.indices.create( 

97 index=self.index_name, 

98 mappings=mappings 

99 ) 

100 

101 def _check_connection(self) -> None: 

102 """Check if database is connected.""" 

103 if not self._connected or not self._client: 

104 raise RuntimeError("Database not connected. Call connect() first.") 

105 

106 def _record_to_doc(self, record: Record) -> dict[str, Any]: 

107 """Convert a Record to an Elasticsearch document.""" 

108 doc = { 

109 "data": {}, 

110 "metadata": record.metadata or {} 

111 } 

112 

113 for field_name, field_obj in record.fields.items(): 

114 doc["data"][field_name] = field_obj.value 

115 

116 return doc 

117 

118 def _doc_to_record(self, doc: dict[str, Any]) -> Record: 

119 """Convert an Elasticsearch document to a Record.""" 

120 data = doc.get("_source", {}).get("data", {}) 

121 metadata = doc.get("_source", {}).get("metadata", {}) 

122 

123 # Add document ID to metadata 

124 if "_id" in doc: 

125 metadata["id"] = doc["_id"] 

126 if "_score" in doc: 

127 metadata["_score"] = doc["_score"] 

128 

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

130 

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

132 """Create a new record.""" 

133 self._check_connection() 

134 doc = self._record_to_doc(record) 

135 

136 # Create document with explicit ID if record has one 

137 kwargs = { 

138 "index": self.index_name, 

139 "document": doc, 

140 "refresh": self.refresh 

141 } 

142 if record.id: 

143 kwargs["id"] = record.id 

144 

145 response = await self._client.index(**kwargs) 

146 

147 return response["_id"] 

148 

149 async def create_batch(self, records: list[Record]) -> list[str]: 

150 """Create multiple records in batch.""" 

151 self._check_connection() 

152 

153 ids = [] 

154 operations = [] 

155 

156 for record in records: 

157 doc = self._record_to_doc(record) 

158 operations.append({"index": {"_index": self.index_name}}) 

159 operations.append(doc) 

160 

161 if operations: 

162 response = await self._client.bulk( 

163 operations=operations, 

164 refresh=self.refresh 

165 ) 

166 

167 # Extract IDs from response 

168 for item in response.get("items", []): 

169 if "index" in item and "_id" in item["index"]: 

170 ids.append(item["index"]["_id"]) 

171 

172 return ids 

173 

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

175 """Read a record by ID.""" 

176 self._check_connection() 

177 

178 try: 

179 response = await self._client.get( 

180 index=self.index_name, 

181 id=id 

182 ) 

183 return self._doc_to_record(response) 

184 except Exception: 

185 return None 

186 

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

188 """Update an existing record.""" 

189 self._check_connection() 

190 doc = self._record_to_doc(record) 

191 

192 try: 

193 await self._client.update( 

194 index=self.index_name, 

195 id=id, 

196 doc=doc, 

197 refresh=self.refresh 

198 ) 

199 return True 

200 except Exception: 

201 return False 

202 

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

204 """Delete a record by ID.""" 

205 self._check_connection() 

206 

207 try: 

208 await self._client.delete( 

209 index=self.index_name, 

210 id=id, 

211 refresh=self.refresh 

212 ) 

213 return True 

214 except Exception: 

215 return False 

216 

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

218 """Check if a record exists.""" 

219 self._check_connection() 

220 

221 return await self._client.exists( 

222 index=self.index_name, 

223 id=id 

224 ) 

225 

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

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

228 self._check_connection() 

229 doc = self._record_to_doc(record) 

230 

231 await self._client.index( 

232 index=self.index_name, 

233 id=id, 

234 document=doc, 

235 refresh=self.refresh 

236 ) 

237 

238 return id 

239 

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

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

242 self._check_connection() 

243 

244 # Build Elasticsearch query 

245 es_query = {"bool": {"must": []}} 

246 

247 for filter in query.filters: 

248 field_path = f"data.{filter.field}" 

249 

250 if filter.operator == Operator.EQ: 

251 # For string values, use keyword field for exact matching 

252 if isinstance(filter.value, str): 

253 field_path = f"{field_path}.keyword" 

254 es_query["bool"]["must"].append({"term": {field_path: filter.value}}) 

255 elif filter.operator == Operator.NEQ: 

256 es_query["bool"]["must_not"] = es_query["bool"].get("must_not", []) 

257 es_query["bool"]["must_not"].append({"term": {field_path: filter.value}}) 

258 elif filter.operator == Operator.GT: 

259 es_query["bool"]["must"].append({"range": {field_path: {"gt": filter.value}}}) 

260 elif filter.operator == Operator.LT: 

261 es_query["bool"]["must"].append({"range": {field_path: {"lt": filter.value}}}) 

262 elif filter.operator == Operator.GTE: 

263 es_query["bool"]["must"].append({"range": {field_path: {"gte": filter.value}}}) 

264 elif filter.operator == Operator.LTE: 

265 es_query["bool"]["must"].append({"range": {field_path: {"lte": filter.value}}}) 

266 elif filter.operator == Operator.LIKE: 

267 es_query["bool"]["must"].append({"wildcard": {field_path: f"*{filter.value}*"}}) 

268 elif filter.operator == Operator.IN: 

269 es_query["bool"]["must"].append({"terms": {field_path: filter.value}}) 

270 elif filter.operator == Operator.NOT_IN: 

271 es_query["bool"]["must_not"] = es_query["bool"].get("must_not", []) 

272 es_query["bool"]["must_not"].append({"terms": {field_path: filter.value}}) 

273 elif filter.operator == Operator.BETWEEN: 

274 # Use Elasticsearch's native range query for efficient BETWEEN 

275 if isinstance(filter.value, (list, tuple)) and len(filter.value) == 2: 

276 lower, upper = filter.value 

277 es_query["bool"]["must"].append({ 

278 "range": { 

279 field_path: { 

280 "gte": lower, 

281 "lte": upper 

282 } 

283 } 

284 }) 

285 elif filter.operator == Operator.NOT_BETWEEN: 

286 # NOT BETWEEN using must_not with range 

287 if isinstance(filter.value, (list, tuple)) and len(filter.value) == 2: 

288 lower, upper = filter.value 

289 es_query["bool"]["must_not"] = es_query["bool"].get("must_not", []) 

290 es_query["bool"]["must_not"].append({ 

291 "range": { 

292 field_path: { 

293 "gte": lower, 

294 "lte": upper 

295 } 

296 } 

297 }) 

298 

299 # If no filters, use match_all 

300 if not es_query["bool"]["must"] and "must_not" not in es_query["bool"]: 

301 es_query = {"match_all": {}} 

302 

303 # Build sort 

304 sort = [] 

305 if query.sort_specs: 

306 for sort_spec in query.sort_specs: 

307 direction = "desc" if sort_spec.order == SortOrder.DESC else "asc" 

308 sort.append({f"data.{sort_spec.field}": {"order": direction}}) 

309 

310 # Build request body 

311 body = {"query": es_query} 

312 if sort: 

313 body["sort"] = sort 

314 

315 # Add size and from for pagination 

316 size = query.limit_value if query.limit_value else 10000 

317 from_param = query.offset_value if query.offset_value else 0 

318 

319 # Execute search 

320 response = await self._client.search( 

321 index=self.index_name, 

322 query=es_query, 

323 sort=sort if sort else None, 

324 size=size, 

325 from_=from_param 

326 ) 

327 

328 # Convert to records 

329 records = [] 

330 for hit in response["hits"]["hits"]: 

331 record = self._doc_to_record(hit) 

332 

333 # Apply field projection if specified 

334 if query.fields: 

335 record = record.project(query.fields) 

336 

337 records.append(record) 

338 

339 return records 

340 

341 async def _count_all(self) -> int: 

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

343 self._check_connection() 

344 

345 response = await self._client.count(index=self.index_name) 

346 return response["count"] 

347 

348 async def clear(self) -> int: 

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

350 self._check_connection() 

351 

352 # Get count before deletion 

353 count = await self._count_all() 

354 

355 # Delete by query - delete all documents 

356 response = await self._client.delete_by_query( 

357 index=self.index_name, 

358 query={"match_all": {}}, 

359 refresh=self.refresh 

360 ) 

361 

362 return response.get("deleted", count) 

363 

364 async def stream_read( 

365 self, 

366 query: Optional[Query] = None, 

367 config: Optional[StreamConfig] = None 

368 ) -> AsyncIterator[Record]: 

369 """Stream records from Elasticsearch using scroll API.""" 

370 self._check_connection() 

371 config = config or StreamConfig() 

372 

373 # Build query 

374 es_query = {"match_all": {}} 

375 if query and query.filters: 

376 es_query = {"bool": {"must": []}} 

377 for filter in query.filters: 

378 field_path = f"data.{filter.field}" 

379 if filter.operator == Operator.EQ: 

380 es_query["bool"]["must"].append({"term": {field_path: filter.value}}) 

381 

382 # Initial search with scroll 

383 response = await self._client.search( 

384 index=self.index_name, 

385 query=es_query, 

386 scroll="2m", 

387 size=config.batch_size 

388 ) 

389 

390 scroll_id = response["_scroll_id"] 

391 hits = response["hits"]["hits"] 

392 

393 try: 

394 while hits: 

395 for hit in hits: 

396 record = self._doc_to_record(hit) 

397 if query and query.fields: 

398 record = record.project(query.fields) 

399 yield record 

400 

401 # Get next batch 

402 response = await self._client.scroll( 

403 scroll_id=scroll_id, 

404 scroll="2m" 

405 ) 

406 hits = response["hits"]["hits"] 

407 finally: 

408 # Clear scroll 

409 await self._client.clear_scroll(scroll_id=scroll_id) 

410 

411 async def stream_write( 

412 self, 

413 records: AsyncIterator[Record], 

414 config: Optional[StreamConfig] = None 

415 ) -> StreamResult: 

416 """Stream records into Elasticsearch using bulk API.""" 

417 self._check_connection() 

418 config = config or StreamConfig() 

419 result = StreamResult() 

420 start_time = time.time() 

421 quitting = False 

422 

423 batch = [] 

424 async for record in records: 

425 batch.append(record) 

426 

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

428 # Write batch with graceful fallback 

429 async def batch_func(b): 

430 await self._write_batch(b) 

431 return [r.id for r in b] 

432 

433 continue_processing = await async_process_batch_with_fallback( 

434 batch, 

435 batch_func, 

436 self.create, 

437 result, 

438 config 

439 ) 

440 

441 if not continue_processing: 

442 quitting = True 

443 break 

444 

445 batch = [] 

446 

447 # Write remaining batch 

448 if batch and not quitting: 

449 async def batch_func(b): 

450 await self._write_batch(b) 

451 return [r.id for r in b] 

452 

453 await async_process_batch_with_fallback( 

454 batch, 

455 batch_func, 

456 self.create, 

457 result, 

458 config 

459 ) 

460 

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

462 return result 

463 

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

465 """Write a batch of records using bulk API.""" 

466 if not records: 

467 return 

468 

469 # Build bulk operations 

470 operations = [] 

471 for record in records: 

472 doc = self._record_to_doc(record) 

473 operations.append({"index": {"_index": self.index_name}}) 

474 operations.append(doc) 

475 

476 # Execute bulk 

477 await self._client.bulk( 

478 operations=operations, 

479 refresh=self.refresh 

480 )