| 12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091929394959697 |
- """Unified normalization, exact dedupe, priority, and budget service."""
- from __future__ import annotations
- import json
- import unicodedata
- from copy import deepcopy
- from typing import Any
- from query_planning.domain import (
- GenerationPlanDraft,
- GenerationRequest,
- GenerationResult,
- GenerationStats,
- QueryCandidate,
- )
- from query_planning.errors import NoQueryCandidatesError
- from query_planning.generators import QueryGeneratorRegistry, default_generator_registry
- def normalize_query_text(value: str) -> tuple[str, str]:
- display_text = " ".join(unicodedata.normalize("NFKC", value or "").strip().split())
- return display_text, display_text.casefold()
- def _unique_dicts(values: list[dict[str, Any]]) -> list[dict[str, Any]]:
- out: list[dict[str, Any]] = []
- seen: set[str] = set()
- for value in values:
- key = json.dumps(value, ensure_ascii=False, sort_keys=True, default=str)
- if key in seen:
- continue
- seen.add(key)
- out.append(deepcopy(value))
- return out
- class UnifiedQueryGenerationService:
- def __init__(self, registry: QueryGeneratorRegistry | None = None) -> None:
- self.registry = registry or default_generator_registry()
- def generate(self, request: GenerationRequest) -> GenerationResult:
- output = self.registry.get(request.generator_kind).generate(request)
- generated_count = len(output.candidates)
- normalized: list[QueryCandidate] = []
- for index, raw in enumerate(output.candidates):
- text, key = normalize_query_text(raw.query_text)
- if not text:
- continue
- candidate = deepcopy(raw)
- candidate.query_text = text
- candidate.dedupe_key = key
- candidate.original_position = raw.original_position if raw.original_position >= 0 else index
- candidate.source_refs = _unique_dicts(candidate.source_refs)
- normalized.append(candidate)
- by_key: dict[str, QueryCandidate] = {}
- for candidate in normalized:
- existing = by_key.get(candidate.dedupe_key)
- if existing is None:
- by_key[candidate.dedupe_key] = candidate
- continue
- existing.priority = max(existing.priority, candidate.priority)
- existing.source_refs = _unique_dicts(existing.source_refs + candidate.source_refs)
- existing.knowledge_need_keys = list(
- dict.fromkeys(existing.knowledge_need_keys + candidate.knowledge_need_keys)
- )
- unique = list(by_key.values())
- for candidate in unique:
- candidate.metadata = deepcopy(candidate.metadata)
- candidate.metadata["origins"] = deepcopy(candidate.source_refs)
- candidate.metadata["priority"] = candidate.priority
- unique.sort(key=lambda candidate: (-candidate.priority, candidate.original_position))
- selected = unique if request.max_queries is None else unique[: request.max_queries]
- if not selected:
- raise NoQueryCandidatesError("query generation produced no selectable query")
- stats = GenerationStats(
- generated_count=generated_count,
- normalized_count=len(normalized),
- unique_count=len(unique),
- selected_count=len(selected),
- dropped_count=generated_count - len(selected),
- )
- plan = GenerationPlanDraft(
- generator_kind=request.generator_kind,
- input_snapshot=deepcopy(request.payload.get("input_snapshot", request.payload)),
- generator_config=deepcopy(output.generator_config),
- metadata=deepcopy(request.metadata),
- )
- return GenerationResult(
- request=request,
- plan=plan,
- knowledge_needs=output.knowledge_needs,
- selected_queries=tuple(selected),
- stats=stats,
- )
|