generators.py 4.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120
  1. """Generator strategies and registry."""
  2. from __future__ import annotations
  3. from collections.abc import Iterable
  4. from typing import Any
  5. from query_planning.domain import (
  6. GenerationRequest,
  7. GeneratorKind,
  8. GeneratorOutput,
  9. KnowledgeNeedDraft,
  10. QueryCandidate,
  11. )
  12. from query_planning.errors import UnsupportedGeneratorError
  13. from query_planning.ports import QueryGenerator
  14. def _candidate(value: Any, position: int, kind: GeneratorKind) -> QueryCandidate:
  15. if isinstance(value, QueryCandidate):
  16. if value.original_position == 0 and position:
  17. value.original_position = position
  18. return value
  19. if isinstance(value, str):
  20. return QueryCandidate(
  21. query_text=value,
  22. original_position=position,
  23. source_refs=[{"generator_kind": kind.value}],
  24. )
  25. if not isinstance(value, dict):
  26. return QueryCandidate(query_text="", original_position=position)
  27. source_refs = value.get("source_refs") or []
  28. if not source_refs:
  29. source_refs = [{"generator_kind": kind.value}]
  30. filter_reason = (
  31. value.get("filter_reason")
  32. if "filter_reason" in value
  33. else value.get("reason")
  34. )
  35. return QueryCandidate(
  36. query_text=str(value.get("query_text") or value.get("query") or ""),
  37. axes=dict(value.get("axes") or value.get("parts") or {}),
  38. priority=int(value.get("priority") or 0),
  39. original_position=int(value.get("original_position", position)),
  40. source_refs=[dict(ref) for ref in source_refs if isinstance(ref, dict)],
  41. knowledge_need_keys=[str(key) for key in value.get("knowledge_need_keys") or []],
  42. metadata=dict(value.get("metadata") or {}),
  43. filter_reason=filter_reason,
  44. )
  45. def _need(value: Any) -> KnowledgeNeedDraft | None:
  46. if isinstance(value, KnowledgeNeedDraft):
  47. return value
  48. if not isinstance(value, dict):
  49. return None
  50. need_key = str(value.get("need_key") or "").strip()
  51. if not need_key:
  52. return None
  53. return KnowledgeNeedDraft(
  54. need_key=need_key,
  55. decision_context=str(value.get("decision_context") or ""),
  56. unknown_information=str(value.get("unknown_information") or ""),
  57. source_ref=dict(value.get("source_ref") or {}),
  58. priority=int(value.get("priority") or 0),
  59. metadata=dict(value.get("metadata") or {}),
  60. )
  61. class _PayloadQueryGenerator:
  62. kind: GeneratorKind
  63. def generate(self, request: GenerationRequest) -> GeneratorOutput:
  64. raw_candidates: Iterable[Any] = request.payload.get("candidates") or ()
  65. candidates = tuple(_candidate(value, index, self.kind) for index, value in enumerate(raw_candidates))
  66. needs = tuple(
  67. need
  68. for raw in request.payload.get("knowledge_needs") or ()
  69. if (need := _need(raw)) is not None
  70. )
  71. return GeneratorOutput(
  72. candidates=candidates,
  73. knowledge_needs=needs,
  74. generator_config=dict(request.payload.get("generator_config") or {}),
  75. )
  76. class CartesianQueryGenerator(_PayloadQueryGenerator):
  77. kind = GeneratorKind.CARTESIAN
  78. class ManualQueryGenerator(_PayloadQueryGenerator):
  79. kind = GeneratorKind.MANUAL
  80. class TopicTableQueryGenerator(_PayloadQueryGenerator):
  81. """Consumes candidates already derived from a topic snapshot by the LLM planner."""
  82. kind = GeneratorKind.TOPIC_TABLE
  83. class QueryGeneratorRegistry:
  84. def __init__(self, generators: Iterable[QueryGenerator] = ()) -> None:
  85. self._generators: dict[GeneratorKind, QueryGenerator] = {}
  86. for generator in generators:
  87. self.register(generator)
  88. def register(self, generator: QueryGenerator) -> None:
  89. self._generators[generator.kind] = generator
  90. def get(self, kind: GeneratorKind) -> QueryGenerator:
  91. try:
  92. return self._generators[kind]
  93. except KeyError as exc:
  94. raise UnsupportedGeneratorError(f"query generator is not registered: {kind.value}") from exc
  95. def default_generator_registry() -> QueryGeneratorRegistry:
  96. return QueryGeneratorRegistry(
  97. (CartesianQueryGenerator(), ManualQueryGenerator(), TopicTableQueryGenerator())
  98. )