pipeline.py 2.3 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374
  1. from __future__ import annotations
  2. from collections.abc import Sequence
  3. from typing import Protocol
  4. from app.core.knowledge.access import KnowledgeAccessContext
  5. from app.core.knowledge.retrieval.contracts import KnowledgeEvidence, SearchResult
  6. from app.core.knowledge.retrieval.fusion import reciprocal_rank_fusion
  7. from app.core.knowledge.retrieval.router import route_query
  8. class Retriever(Protocol):
  9. def retrieve(
  10. self,
  11. query: str,
  12. context: KnowledgeAccessContext,
  13. limit: int,
  14. ) -> Sequence[KnowledgeEvidence]: ...
  15. class KnowledgeRetrievalPipeline:
  16. def __init__(
  17. self,
  18. *,
  19. lexical: Retriever,
  20. vector: Retriever,
  21. device: Retriever | None = None,
  22. graph: Retriever | None = None,
  23. lightrag: Retriever | None = None,
  24. ):
  25. self._retrievers = {
  26. "lexical": lexical,
  27. "vector": vector,
  28. "device": device,
  29. "graph": graph,
  30. "lightrag": lightrag,
  31. }
  32. def search(
  33. self,
  34. query: str,
  35. *,
  36. context: KnowledgeAccessContext,
  37. mode: str = "auto",
  38. limit: int = 20,
  39. ) -> SearchResult:
  40. resolved_mode = route_query(query) if mode == "auto" else mode
  41. names = ["lexical", "vector", "device"]
  42. if resolved_mode == "relationship" and self._retrievers["graph"] is not None:
  43. names.append("graph")
  44. if resolved_mode == "global" and self._retrievers["lightrag"] is not None:
  45. names.append("lightrag")
  46. ranked: dict[str, Sequence[KnowledgeEvidence]] = {}
  47. degraded: list[str] = []
  48. for name in names:
  49. retriever = self._retrievers[name]
  50. if retriever is None:
  51. continue
  52. try:
  53. ranked[name] = retriever.retrieve(query, context, limit)
  54. except Exception:
  55. degraded.append(name)
  56. fused = reciprocal_rank_fusion(ranked, limit=limit)
  57. authorized = tuple(
  58. evidence
  59. for evidence in fused
  60. if context.permits_domain(evidence.business_domain_uid)
  61. and evidence.freshness_status != "stale"
  62. )
  63. return SearchResult(
  64. evidence=authorized,
  65. mode=resolved_mode,
  66. degraded_components=tuple(degraded),
  67. )