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
« 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."""
3import asyncio
4import json
5import logging
6import time
7import uuid
8from typing import Any, AsyncIterator, Optional
10from dataknobs_config import ConfigurableBase
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
26logger = logging.getLogger(__name__)
28# Global pool manager for Elasticsearch clients
29_client_manager = ConnectionPoolManager()
32class AsyncElasticsearchDatabase(AsyncDatabase, ConfigurableBase):
33 """Native async Elasticsearch database backend with connection pooling."""
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
45 @classmethod
46 def from_config(cls, config: dict) -> "AsyncElasticsearchDatabase":
47 """Create from config dictionary."""
48 return cls(config)
50 async def connect(self) -> None:
51 """Connect to the Elasticsearch database."""
52 if self._connected:
53 return
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 )
63 # Ensure index exists
64 await self._ensure_index()
65 self._connected = True
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
75 def _initialize(self) -> None:
76 """Initialize is handled in connect."""
77 pass
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.")
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 }
96 await self._client.indices.create(
97 index=self.index_name,
98 mappings=mappings
99 )
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.")
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 }
113 for field_name, field_obj in record.fields.items():
114 doc["data"][field_name] = field_obj.value
116 return doc
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", {})
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"]
129 return Record(data=data, metadata=metadata)
131 async def create(self, record: Record) -> str:
132 """Create a new record."""
133 self._check_connection()
134 doc = self._record_to_doc(record)
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
145 response = await self._client.index(**kwargs)
147 return response["_id"]
149 async def create_batch(self, records: list[Record]) -> list[str]:
150 """Create multiple records in batch."""
151 self._check_connection()
153 ids = []
154 operations = []
156 for record in records:
157 doc = self._record_to_doc(record)
158 operations.append({"index": {"_index": self.index_name}})
159 operations.append(doc)
161 if operations:
162 response = await self._client.bulk(
163 operations=operations,
164 refresh=self.refresh
165 )
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"])
172 return ids
174 async def read(self, id: str) -> Record | None:
175 """Read a record by ID."""
176 self._check_connection()
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
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)
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
203 async def delete(self, id: str) -> bool:
204 """Delete a record by ID."""
205 self._check_connection()
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
217 async def exists(self, id: str) -> bool:
218 """Check if a record exists."""
219 self._check_connection()
221 return await self._client.exists(
222 index=self.index_name,
223 id=id
224 )
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)
231 await self._client.index(
232 index=self.index_name,
233 id=id,
234 document=doc,
235 refresh=self.refresh
236 )
238 return id
240 async def search(self, query: Query) -> list[Record]:
241 """Search for records matching the query."""
242 self._check_connection()
244 # Build Elasticsearch query
245 es_query = {"bool": {"must": []}}
247 for filter in query.filters:
248 field_path = f"data.{filter.field}"
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 })
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": {}}
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}})
310 # Build request body
311 body = {"query": es_query}
312 if sort:
313 body["sort"] = sort
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
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 )
328 # Convert to records
329 records = []
330 for hit in response["hits"]["hits"]:
331 record = self._doc_to_record(hit)
333 # Apply field projection if specified
334 if query.fields:
335 record = record.project(query.fields)
337 records.append(record)
339 return records
341 async def _count_all(self) -> int:
342 """Count all records in the database."""
343 self._check_connection()
345 response = await self._client.count(index=self.index_name)
346 return response["count"]
348 async def clear(self) -> int:
349 """Clear all records from the database."""
350 self._check_connection()
352 # Get count before deletion
353 count = await self._count_all()
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 )
362 return response.get("deleted", count)
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()
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}})
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 )
390 scroll_id = response["_scroll_id"]
391 hits = response["hits"]["hits"]
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
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)
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
423 batch = []
424 async for record in records:
425 batch.append(record)
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]
433 continue_processing = await async_process_batch_with_fallback(
434 batch,
435 batch_func,
436 self.create,
437 result,
438 config
439 )
441 if not continue_processing:
442 quitting = True
443 break
445 batch = []
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]
453 await async_process_batch_with_fallback(
454 batch,
455 batch_func,
456 self.create,
457 result,
458 config
459 )
461 result.duration = time.time() - start_time
462 return result
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
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)
476 # Execute bulk
477 await self._client.bulk(
478 operations=operations,
479 refresh=self.refresh
480 )