"""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 TopicTableQueryGenerator(_PayloadQueryGenerator): """Consumes candidates already derived from a topic snapshot by the LLM planner.""" kind = GeneratorKind.TOPIC_TABLE 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(), TopicTableQueryGenerator()) )