| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112 |
- """Generator strategies and registry."""
- from __future__ import annotations
- from collections.abc import Iterable
- from typing import Any
- from query_planning.domain import (
- GenerationRequest,
- GeneratorKind,
- GeneratorOutput,
- KnowledgeNeedDraft,
- QueryCandidate,
- )
- from query_planning.errors import UnsupportedGeneratorError
- from query_planning.ports import QueryGenerator
- def _candidate(value: Any, position: int, kind: GeneratorKind) -> QueryCandidate:
- if isinstance(value, QueryCandidate):
- if value.original_position == 0 and position:
- value.original_position = position
- return value
- if isinstance(value, str):
- return QueryCandidate(
- query_text=value,
- original_position=position,
- source_refs=[{"generator_kind": kind.value}],
- )
- if not isinstance(value, dict):
- return QueryCandidate(query_text="", original_position=position)
- source_refs = value.get("source_refs") or []
- if not source_refs:
- source_refs = [{"generator_kind": kind.value}]
- filter_reason = (
- value.get("filter_reason")
- if "filter_reason" in value
- else value.get("reason")
- )
- return QueryCandidate(
- query_text=str(value.get("query_text") or value.get("query") or ""),
- axes=dict(value.get("axes") or value.get("parts") or {}),
- priority=int(value.get("priority") or 0),
- original_position=int(value.get("original_position", position)),
- source_refs=[dict(ref) for ref in source_refs if isinstance(ref, dict)],
- knowledge_need_keys=[str(key) for key in value.get("knowledge_need_keys") or []],
- metadata=dict(value.get("metadata") or {}),
- filter_reason=filter_reason,
- )
- def _need(value: Any) -> KnowledgeNeedDraft | None:
- if isinstance(value, KnowledgeNeedDraft):
- return value
- if not isinstance(value, dict):
- return None
- need_key = str(value.get("need_key") or "").strip()
- if not need_key:
- return None
- return KnowledgeNeedDraft(
- need_key=need_key,
- decision_context=str(value.get("decision_context") or ""),
- unknown_information=str(value.get("unknown_information") or ""),
- source_ref=dict(value.get("source_ref") or {}),
- priority=int(value.get("priority") or 0),
- metadata=dict(value.get("metadata") or {}),
- )
- class _PayloadQueryGenerator:
- kind: GeneratorKind
- def generate(self, request: GenerationRequest) -> GeneratorOutput:
- raw_candidates: Iterable[Any] = request.payload.get("candidates") or ()
- candidates = tuple(_candidate(value, index, self.kind) for index, value in enumerate(raw_candidates))
- needs = tuple(
- need
- for raw in request.payload.get("knowledge_needs") or ()
- if (need := _need(raw)) is not None
- )
- return GeneratorOutput(
- candidates=candidates,
- knowledge_needs=needs,
- generator_config=dict(request.payload.get("generator_config") or {}),
- )
- class CartesianQueryGenerator(_PayloadQueryGenerator):
- kind = GeneratorKind.CARTESIAN
- class ManualQueryGenerator(_PayloadQueryGenerator):
- kind = GeneratorKind.MANUAL
- class QueryGeneratorRegistry:
- def __init__(self, generators: Iterable[QueryGenerator] = ()) -> None:
- self._generators: dict[GeneratorKind, QueryGenerator] = {}
- for generator in generators:
- self.register(generator)
- def register(self, generator: QueryGenerator) -> None:
- self._generators[generator.kind] = generator
- def get(self, kind: GeneratorKind) -> QueryGenerator:
- try:
- return self._generators[kind]
- except KeyError as exc:
- raise UnsupportedGeneratorError(f"query generator is not registered: {kind.value}") from exc
- def default_generator_registry() -> QueryGeneratorRegistry:
- return QueryGeneratorRegistry((CartesianQueryGenerator(), ManualQueryGenerator()))
|