Coverage for src/qdrant_loader/core/qdrant_manager.py: 69%
238 statements
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-20 10:15 +0000
« prev ^ index » next coverage.py v7.15.0, created at 2026-07-20 10:15 +0000
1import asyncio
2from typing import Any, cast
3from urllib.parse import urlparse
5from qdrant_client import QdrantClient
6from qdrant_client.http import models
7from qdrant_client.http.models import (
8 Distance,
9 VectorParams,
10)
11from qdrant_loader_core.config import (
12 CollectionVectorCapabilities,
13 SparseRuntimeConfig,
14 parse_collection_capabilities,
15)
16from qdrant_loader_core.sparse import get_sparse_encoder
18from ..config import Settings, get_global_config, get_settings
19from ..utils.logging import LoggingConfig
21logger = LoggingConfig.get_logger(__name__)
24class QdrantConnectionError(Exception):
25 """Custom exception for Qdrant connection errors."""
27 def __init__(
28 self, message: str, original_error: str | None = None, url: str | None = None
29 ):
30 self.message = message
31 self.original_error = original_error
32 self.url = url
33 super().__init__(self.message)
36class QdrantManager:
37 def __init__(self, settings: Settings | None = None):
38 """Initialize the qDrant manager.
40 Args:
41 settings: The application settings
42 """
43 self.settings = settings or get_settings()
44 self.client = None
45 self.collection_name = self.settings.qdrant_collection_name
46 self.logger = LoggingConfig.get_logger(__name__)
47 self.batch_size = get_global_config().embedding.batch_size
48 self.sparse_runtime = self._resolve_sparse_runtime_config()
49 self._collection_vector_capabilities: CollectionVectorCapabilities | None = None
50 self._sparse_fallback_warning_emitted = False
51 self.connect()
53 def _is_api_key_present(self) -> bool:
54 """
55 Check if a valid API key is present.
56 Returns True if the API key is a non-empty string that is not 'None' or 'null'.
57 """
58 api_key = self.settings.qdrant_api_key
59 if not api_key: # Catches None, empty string, etc.
60 return False
61 return api_key.lower() not in ["none", "null"]
63 def _resolve_sparse_runtime_config(self) -> SparseRuntimeConfig:
64 try:
65 llm = getattr(get_global_config(), "llm", None) or {}
66 except Exception as e:
67 self.logger.warning(
68 "Failed to read global LLM config for sparse runtime; using defaults",
69 error=str(e),
70 exc_info=True,
71 )
72 llm = {}
73 global_config = {"llm": llm} if isinstance(llm, dict) else {}
74 return SparseRuntimeConfig.from_global_config(global_config)
76 def _get_collection_vector_capabilities(self) -> CollectionVectorCapabilities:
77 if self._collection_vector_capabilities is not None:
78 return self._collection_vector_capabilities
80 client = self._ensure_client_connected()
81 try:
82 info = client.get_collection(collection_name=self.collection_name)
83 except Exception as e:
84 # Don't cache: a transient outage would otherwise pin every
85 # subsequent upsert to dense-only payload shape even after Qdrant
86 # becomes reachable again, mismatching a hybrid collection schema.
87 self.logger.warning(
88 "Failed to inspect Qdrant collection schema; assuming dense-only",
89 collection=self.collection_name,
90 error=str(e),
91 )
92 return CollectionVectorCapabilities()
94 self._collection_vector_capabilities = parse_collection_capabilities(
95 info, self.sparse_runtime
96 )
97 return self._collection_vector_capabilities
99 def _dense_query_using(self) -> str | None:
100 caps = self._get_collection_vector_capabilities()
101 if caps.has_named_dense:
102 return self.sparse_runtime.dense_vector_name
103 return None
105 def _sparse_upsert_enabled(self) -> bool:
106 if not self.sparse_runtime.enabled:
107 return False
109 caps = self._get_collection_vector_capabilities()
110 if caps.has_named_dense and caps.has_sparse:
111 return True
113 if not self._sparse_fallback_warning_emitted:
114 self.logger.warning(
115 "Sparse vectors requested but collection schema does not support them; falling back to dense-only upserts",
116 collection=self.collection_name,
117 dense_vector_name=self.sparse_runtime.dense_vector_name,
118 sparse_vector_name=self.sparse_runtime.sparse_vector_name,
119 )
120 self._sparse_fallback_warning_emitted = True
121 return False
123 def build_point_vector(self, dense_embedding: list[float], text: str) -> object:
124 """Build the point vector payload for upsert.
126 Three shapes are possible depending on the live collection schema:
127 - dense+sparse named dict (hybrid-ready collection),
128 - dense-only named dict (legacy named-vector collection),
129 - raw dense list (legacy unnamed collection).
130 """
131 if self._sparse_upsert_enabled():
132 return self._build_hybrid_payload(dense_embedding, text)
133 return self._build_dense_payload(dense_embedding)
135 def _build_dense_payload(self, dense_embedding: list[float]) -> object:
136 """Return dense-only payload using the named-vector shape if the collection requires it."""
137 if self._dense_query_using() is not None:
138 return {self.sparse_runtime.dense_vector_name: dense_embedding}
139 return dense_embedding
141 def _build_hybrid_payload(self, dense_embedding: list[float], text: str) -> object:
142 """Return dense+sparse payload, with a dense-only fallback on encode failure."""
143 try:
144 sparse = get_sparse_encoder(self.sparse_runtime.model).encode_document(text)
145 except Exception as e:
146 self.logger.warning(
147 "Failed to generate sparse vectors; falling back to dense-only upsert",
148 error=str(e),
149 )
150 return self._build_dense_payload(dense_embedding)
152 if sparse.is_empty():
153 return {self.sparse_runtime.dense_vector_name: dense_embedding}
154 return {
155 self.sparse_runtime.dense_vector_name: dense_embedding,
156 self.sparse_runtime.sparse_vector_name: models.SparseVector(
157 indices=sparse.indices, values=sparse.values
158 ),
159 }
161 def connect(self) -> None:
162 """Establish connection to qDrant server."""
163 try:
164 # Ensure HTTPS is used when API key is present, but only for non-local URLs
165 url = self.settings.qdrant_url
166 api_key = (
167 self.settings.qdrant_api_key if self._is_api_key_present() else None
168 )
170 if api_key:
171 parsed_url = urlparse(url)
172 # Only force HTTPS for non-local URLs
173 if parsed_url.scheme != "https" and not any(
174 host in parsed_url.netloc for host in ["localhost", "127.0.0.1"]
175 ):
176 url = url.replace("http://", "https://", 1)
177 self.logger.warning("Forcing HTTPS connection due to API key usage")
179 try:
180 self.client = QdrantClient(
181 url=url,
182 api_key=api_key,
183 timeout=60, # 60 seconds timeout
184 )
185 self.logger.debug("Successfully connected to qDrant")
186 except Exception as e:
187 raise QdrantConnectionError(
188 "Failed to connect to qDrant: Connection error",
189 original_error=str(e),
190 url=url,
191 ) from e
193 except Exception as e:
194 raise QdrantConnectionError(
195 "Failed to connect to qDrant: Unexpected error",
196 original_error=str(e),
197 url=url,
198 ) from e
200 def _ensure_client_connected(self) -> QdrantClient:
201 """Ensure the client is connected before performing operations."""
202 if self.client is None:
203 raise QdrantConnectionError(
204 "Qdrant client is not connected. Please call connect() first."
205 )
206 return cast(QdrantClient, self.client)
208 async def assert_collection_accessible(self) -> None:
209 """Validate that the configured collection is reachable.
211 Raises when the collection does not exist or Qdrant is unavailable.
212 """
213 client = self._ensure_client_connected()
214 await asyncio.to_thread(
215 client.get_collection, collection_name=self.collection_name
216 )
218 # Indexes that filter-based deletes (delete_points_by_document_id, etc.)
219 # depend on. A silent failure here would surface later as a confusing
220 # Qdrant error mid-delete, so these must raise instead of just warning.
221 _REQUIRED_PAYLOAD_INDEXES = frozenset(
222 {"document_id", "project_id", "source_type", "source"}
223 )
225 def _ensure_payload_indexes(self, client: QdrantClient) -> None:
226 """Ensure all required payload indexes exist on the collection.
228 Safe to call on both new and existing collections — Qdrant ignores
229 duplicate index creation requests.
230 """
231 indexes_to_create = [
232 ("document_id", {"type": "keyword"}),
233 ("project_id", {"type": "keyword"}),
234 ("source_type", {"type": "keyword"}),
235 ("source", {"type": "keyword"}),
236 ("title", {"type": "keyword"}),
237 ("created_at", {"type": "keyword"}),
238 ("updated_at", {"type": "keyword"}),
239 ("is_attachment", {"type": "bool"}),
240 ("parent_document_id", {"type": "keyword"}),
241 ("original_file_type", {"type": "keyword"}),
242 ("is_converted", {"type": "bool"}),
243 ]
245 created_indexes = []
246 failed_indexes = []
248 for field_name, field_schema in indexes_to_create:
249 try:
250 client.create_payload_index(
251 collection_name=self.collection_name,
252 field_name=field_name,
253 field_schema=field_schema, # type: ignore
254 )
255 created_indexes.append(field_name)
256 self.logger.debug(f"Ensured payload index for field: {field_name}")
257 except Exception as e:
258 failed_indexes.append((field_name, str(e)))
259 self.logger.warning(
260 f"Failed to create index for {field_name}", error=str(e)
261 )
263 if failed_indexes:
264 self.logger.warning(
265 "Some indexes failed to create but collection is functional",
266 failed_details=failed_indexes,
267 )
268 required_failures = {
269 name: err
270 for name, err in failed_indexes
271 if name in self._REQUIRED_PAYLOAD_INDEXES
272 }
273 if required_failures:
274 raise RuntimeError(
275 "Failed to create required payload index(es): "
276 f"{required_failures}. These are required for filter-based "
277 "deletes; refusing to continue with a collection in this state."
278 )
280 self.logger.info(
281 "Payload indexes ensured",
282 collection=self.collection_name,
283 created_indexes=created_indexes,
284 failed_indexes=[name for name, _ in failed_indexes] or None,
285 )
287 def create_collection(self) -> None:
288 """Create a new collection if it doesn't exist."""
289 try:
290 client = self._ensure_client_connected()
291 # Check if collection already exists
292 collections = client.get_collections()
293 if any(c.name == self.collection_name for c in collections.collections):
294 self.logger.info(f"Collection {self.collection_name} already exists")
295 self._ensure_payload_indexes(client)
296 return
298 # Get vector size from unified LLM settings first, then legacy embedding.
299 # global_config.llm is a plain dict, so the unified vector size must be
300 # resolved through llm_settings (which parses global.llm.embeddings),
301 # mirroring EmbeddingService.get_embedding_dimension().
302 vector_size: int | None = None
303 try:
304 vs = self.settings.llm_settings.embeddings.vector_size
305 except AttributeError:
306 vs = None
307 if vs is not None:
308 try:
309 vector_size = int(vs)
310 except (TypeError, ValueError) as exc:
311 raise ValueError(
312 "Invalid global.llm.embeddings.vector_size; expected an integer"
313 ) from exc
314 # Qdrant requires a strictly positive dimensionality; 0/negative
315 # values are rejected at collection creation, so fail early here.
316 if vector_size <= 0:
317 raise ValueError(
318 "Invalid global.llm.embeddings.vector_size; "
319 "expected a positive integer"
320 )
322 if vector_size is None:
323 try:
324 legacy_vs = get_global_config().embedding.vector_size
325 except AttributeError:
326 legacy_vs = None
327 if legacy_vs is not None:
328 try:
329 vector_size = int(legacy_vs)
330 except (TypeError, ValueError) as exc:
331 raise ValueError(
332 "Invalid embedding.vector_size; expected an integer"
333 ) from exc
334 if vector_size <= 0:
335 raise ValueError(
336 "Invalid embedding.vector_size; "
337 "expected a positive integer"
338 )
340 if vector_size is None:
341 self.logger.warning(
342 "No vector_size specified in config; falling back to 1024 (deprecated default). Set global.llm.embeddings.vector_size."
343 )
344 vector_size = 1024
346 # sparse.enabled is a strict declaration. If True, the collection
347 # is created with a sparse vector; failures propagate. If False,
348 # dense-only. Operators on Qdrant servers that don't support sparse
349 # vectors must set sparse.enabled=false explicitly.
350 dense_params = VectorParams(size=vector_size, distance=Distance.COSINE)
351 if self.sparse_runtime.enabled:
352 client.create_collection(
353 collection_name=self.collection_name,
354 vectors_config={
355 self.sparse_runtime.dense_vector_name: dense_params
356 },
357 sparse_vectors_config={
358 self.sparse_runtime.sparse_vector_name: models.SparseVectorParams()
359 },
360 )
361 self.logger.info(
362 "Created Qdrant collection with dense+sparse vectors",
363 collection=self.collection_name,
364 dense_vector_name=self.sparse_runtime.dense_vector_name,
365 sparse_vector_name=self.sparse_runtime.sparse_vector_name,
366 sparse_model=self.sparse_runtime.model,
367 )
368 else:
369 client.create_collection(
370 collection_name=self.collection_name,
371 vectors_config=dense_params,
372 )
374 self._collection_vector_capabilities = CollectionVectorCapabilities(
375 has_named_dense=self.sparse_runtime.enabled,
376 has_sparse=self.sparse_runtime.enabled,
377 )
379 self._ensure_payload_indexes(client)
381 self.logger.info(
382 f"Collection {self.collection_name} created with indexes",
383 )
384 except Exception as e:
385 self.logger.error("Failed to create collection", error=str(e))
386 raise
388 async def upsert_points(self, points: list[models.PointStruct]) -> None:
389 """Upsert points into the collection.
391 Args:
392 points: List of points to upsert
393 """
394 self.logger.debug(
395 "Upserting points",
396 extra={"point_count": len(points), "collection": self.collection_name},
397 )
399 try:
400 client = self._ensure_client_connected()
401 await asyncio.to_thread(
402 client.upsert, collection_name=self.collection_name, points=points
403 )
404 self.logger.debug(
405 "Successfully upserted points",
406 extra={"point_count": len(points), "collection": self.collection_name},
407 )
408 except Exception as e:
409 self.logger.error(
410 "Failed to upsert points",
411 extra={
412 "error": str(e),
413 "point_count": len(points),
414 "collection": self.collection_name,
415 },
416 )
417 raise
419 def search(
420 self, query_vector: list[float], limit: int = 5
421 ) -> list[models.ScoredPoint]:
422 """Search for similar vectors in the collection."""
423 try:
424 client = self._ensure_client_connected()
425 query_kwargs: dict[str, Any] = {
426 "collection_name": self.collection_name,
427 "query": query_vector,
428 "limit": limit,
429 }
430 using = self._dense_query_using()
431 if using:
432 query_kwargs["using"] = using
433 # Use query_points API (qdrant-client 1.10+)
434 query_response = client.query_points(**query_kwargs)
435 return query_response.points
436 except Exception as e:
437 logger.error("Failed to search collection", error=str(e))
438 raise
440 def search_with_project_filter(
441 self, query_vector: list[float], project_ids: list[str], limit: int = 5
442 ) -> list[models.ScoredPoint]:
443 """Search for similar vectors in the collection with project filtering.
445 Args:
446 query_vector: Query vector for similarity search
447 project_ids: List of project IDs to filter by
448 limit: Maximum number of results to return
450 Returns:
451 List of scored points matching the query and project filter
452 """
453 try:
454 client = self._ensure_client_connected()
456 # Build project filter
457 project_filter = models.Filter(
458 must=[
459 models.FieldCondition(
460 key="project_id", match=models.MatchAny(any=project_ids)
461 )
462 ]
463 )
465 query_kwargs: dict[str, Any] = {
466 "collection_name": self.collection_name,
467 "query": query_vector,
468 "query_filter": project_filter,
469 "limit": limit,
470 }
471 using = self._dense_query_using()
472 if using:
473 query_kwargs["using"] = using
474 # Use query_points API (qdrant-client 1.10+)
475 query_response = client.query_points(**query_kwargs)
476 return query_response.points
477 except Exception as e:
478 logger.error(
479 "Failed to search collection with project filter",
480 error=str(e),
481 project_ids=project_ids,
482 )
483 raise
485 def get_project_collections(self) -> dict[str, str]:
486 """Get mapping of project IDs to their collection names.
488 Returns:
489 Dictionary mapping project_id to collection_name
490 """
491 try:
492 client = self._ensure_client_connected()
494 # Scroll through all points to get unique project-collection mappings
495 scroll_result = client.scroll(
496 collection_name=self.collection_name,
497 limit=10000, # Large limit to get all unique projects
498 with_payload=True,
499 with_vectors=False,
500 )
502 project_collections = {}
503 for point in scroll_result[0]:
504 if point.payload:
505 project_id = point.payload.get("project_id")
506 collection_name = point.payload.get("collection_name")
507 if project_id and collection_name:
508 project_collections[project_id] = collection_name
510 return project_collections
511 except Exception as e:
512 logger.error("Failed to get project collections", error=str(e))
513 raise
515 def delete_collection(self) -> None:
516 """Delete the collection."""
517 try:
518 client = self._ensure_client_connected()
519 client.delete_collection(collection_name=self.collection_name)
520 logger.debug("Collection deleted", collection=self.collection_name)
521 except Exception as e:
522 logger.error("Failed to delete collection", error=str(e))
523 raise
525 async def delete_points_by_document_id(self, document_ids: list[str]) -> None:
526 """Delete points from the collection by document ID.
528 Args:
529 document_ids: List of document IDs to delete
530 """
531 self.logger.debug(
532 "Deleting points by document ID",
533 extra={
534 "document_count": len(document_ids),
535 "collection": self.collection_name,
536 },
537 )
539 try:
540 client = self._ensure_client_connected()
541 await asyncio.to_thread(
542 client.delete,
543 collection_name=self.collection_name,
544 points_selector=models.Filter(
545 must=[
546 models.FieldCondition(
547 key="document_id", match=models.MatchAny(any=document_ids)
548 )
549 ]
550 ),
551 )
552 self.logger.debug(
553 "Successfully deleted points",
554 extra={
555 "document_count": len(document_ids),
556 "collection": self.collection_name,
557 },
558 )
559 except Exception as e:
560 self.logger.error(
561 "Failed to delete points",
562 extra={
563 "error": str(e),
564 "document_count": len(document_ids),
565 "collection": self.collection_name,
566 },
567 )
568 raise