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

1import asyncio 

2from typing import Any, cast 

3from urllib.parse import urlparse 

4 

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 

17 

18from ..config import Settings, get_global_config, get_settings 

19from ..utils.logging import LoggingConfig 

20 

21logger = LoggingConfig.get_logger(__name__) 

22 

23 

24class QdrantConnectionError(Exception): 

25 """Custom exception for Qdrant connection errors.""" 

26 

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) 

34 

35 

36class QdrantManager: 

37 def __init__(self, settings: Settings | None = None): 

38 """Initialize the qDrant manager. 

39 

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

52 

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

62 

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) 

75 

76 def _get_collection_vector_capabilities(self) -> CollectionVectorCapabilities: 

77 if self._collection_vector_capabilities is not None: 

78 return self._collection_vector_capabilities 

79 

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

93 

94 self._collection_vector_capabilities = parse_collection_capabilities( 

95 info, self.sparse_runtime 

96 ) 

97 return self._collection_vector_capabilities 

98 

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 

104 

105 def _sparse_upsert_enabled(self) -> bool: 

106 if not self.sparse_runtime.enabled: 

107 return False 

108 

109 caps = self._get_collection_vector_capabilities() 

110 if caps.has_named_dense and caps.has_sparse: 

111 return True 

112 

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 

122 

123 def build_point_vector(self, dense_embedding: list[float], text: str) -> object: 

124 """Build the point vector payload for upsert. 

125 

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) 

134 

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 

140 

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) 

151 

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 } 

160 

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 ) 

169 

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

178 

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 

192 

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 

199 

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) 

207 

208 async def assert_collection_accessible(self) -> None: 

209 """Validate that the configured collection is reachable. 

210 

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 ) 

217 

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 ) 

224 

225 def _ensure_payload_indexes(self, client: QdrantClient) -> None: 

226 """Ensure all required payload indexes exist on the collection. 

227 

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 ] 

244 

245 created_indexes = [] 

246 failed_indexes = [] 

247 

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 ) 

262 

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 ) 

279 

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 ) 

286 

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 

297 

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 ) 

321 

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 ) 

339 

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 

345 

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 ) 

373 

374 self._collection_vector_capabilities = CollectionVectorCapabilities( 

375 has_named_dense=self.sparse_runtime.enabled, 

376 has_sparse=self.sparse_runtime.enabled, 

377 ) 

378 

379 self._ensure_payload_indexes(client) 

380 

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 

387 

388 async def upsert_points(self, points: list[models.PointStruct]) -> None: 

389 """Upsert points into the collection. 

390 

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 ) 

398 

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 

418 

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 

439 

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. 

444 

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 

449 

450 Returns: 

451 List of scored points matching the query and project filter 

452 """ 

453 try: 

454 client = self._ensure_client_connected() 

455 

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 ) 

464 

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 

484 

485 def get_project_collections(self) -> dict[str, str]: 

486 """Get mapping of project IDs to their collection names. 

487 

488 Returns: 

489 Dictionary mapping project_id to collection_name 

490 """ 

491 try: 

492 client = self._ensure_client_connected() 

493 

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 ) 

501 

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 

509 

510 return project_collections 

511 except Exception as e: 

512 logger.error("Failed to get project collections", error=str(e)) 

513 raise 

514 

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 

524 

525 async def delete_points_by_document_id(self, document_ids: list[str]) -> None: 

526 """Delete points from the collection by document ID. 

527 

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 ) 

538 

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