service.py 3.8 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091929394959697
  1. """Unified normalization, exact dedupe, priority, and budget service."""
  2. from __future__ import annotations
  3. import json
  4. import unicodedata
  5. from copy import deepcopy
  6. from typing import Any
  7. from query_planning.domain import (
  8. GenerationPlanDraft,
  9. GenerationRequest,
  10. GenerationResult,
  11. GenerationStats,
  12. QueryCandidate,
  13. )
  14. from query_planning.errors import NoQueryCandidatesError
  15. from query_planning.generators import QueryGeneratorRegistry, default_generator_registry
  16. def normalize_query_text(value: str) -> tuple[str, str]:
  17. display_text = " ".join(unicodedata.normalize("NFKC", value or "").strip().split())
  18. return display_text, display_text.casefold()
  19. def _unique_dicts(values: list[dict[str, Any]]) -> list[dict[str, Any]]:
  20. out: list[dict[str, Any]] = []
  21. seen: set[str] = set()
  22. for value in values:
  23. key = json.dumps(value, ensure_ascii=False, sort_keys=True, default=str)
  24. if key in seen:
  25. continue
  26. seen.add(key)
  27. out.append(deepcopy(value))
  28. return out
  29. class UnifiedQueryGenerationService:
  30. def __init__(self, registry: QueryGeneratorRegistry | None = None) -> None:
  31. self.registry = registry or default_generator_registry()
  32. def generate(self, request: GenerationRequest) -> GenerationResult:
  33. output = self.registry.get(request.generator_kind).generate(request)
  34. generated_count = len(output.candidates)
  35. normalized: list[QueryCandidate] = []
  36. for index, raw in enumerate(output.candidates):
  37. text, key = normalize_query_text(raw.query_text)
  38. if not text:
  39. continue
  40. candidate = deepcopy(raw)
  41. candidate.query_text = text
  42. candidate.dedupe_key = key
  43. candidate.original_position = raw.original_position if raw.original_position >= 0 else index
  44. candidate.source_refs = _unique_dicts(candidate.source_refs)
  45. normalized.append(candidate)
  46. by_key: dict[str, QueryCandidate] = {}
  47. for candidate in normalized:
  48. existing = by_key.get(candidate.dedupe_key)
  49. if existing is None:
  50. by_key[candidate.dedupe_key] = candidate
  51. continue
  52. existing.priority = max(existing.priority, candidate.priority)
  53. existing.source_refs = _unique_dicts(existing.source_refs + candidate.source_refs)
  54. existing.knowledge_need_keys = list(
  55. dict.fromkeys(existing.knowledge_need_keys + candidate.knowledge_need_keys)
  56. )
  57. unique = list(by_key.values())
  58. for candidate in unique:
  59. candidate.metadata = deepcopy(candidate.metadata)
  60. candidate.metadata["origins"] = deepcopy(candidate.source_refs)
  61. candidate.metadata["priority"] = candidate.priority
  62. unique.sort(key=lambda candidate: (-candidate.priority, candidate.original_position))
  63. selected = unique if request.max_queries is None else unique[: request.max_queries]
  64. if not selected:
  65. raise NoQueryCandidatesError("query generation produced no selectable query")
  66. stats = GenerationStats(
  67. generated_count=generated_count,
  68. normalized_count=len(normalized),
  69. unique_count=len(unique),
  70. selected_count=len(selected),
  71. dropped_count=generated_count - len(selected),
  72. )
  73. plan = GenerationPlanDraft(
  74. generator_kind=request.generator_kind,
  75. input_snapshot=deepcopy(request.payload.get("input_snapshot", request.payload)),
  76. generator_config=deepcopy(output.generator_config),
  77. metadata=deepcopy(request.metadata),
  78. )
  79. return GenerationResult(
  80. request=request,
  81. plan=plan,
  82. knowledge_needs=output.knowledge_needs,
  83. selected_queries=tuple(selected),
  84. stats=stats,
  85. )