query_audit.py 5.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149
  1. from __future__ import annotations
  2. import hashlib
  3. import json
  4. from collections import Counter
  5. from dataclasses import dataclass
  6. from typing import Any
  7. from uuid import uuid4
  8. from sqlalchemy import text
  9. from app.core.knowledge.access import KnowledgeAccessContext
  10. from app.core.knowledge.retrieval.contracts import KnowledgeEvidence
  11. @dataclass(frozen=True)
  12. class KnowledgeQueryAuditRecord:
  13. uid: str
  14. query_hash: str
  15. user_id: str | None
  16. roles: tuple[str, ...]
  17. business_domain_uids: tuple[str, ...]
  18. mode: str
  19. retriever_counts: dict[str, int]
  20. cited_points: tuple[dict[str, Any], ...]
  21. degraded_components: tuple[str, ...]
  22. correlation_id: str
  23. latency_ms: int
  24. def normalize_audited_query(query: str) -> str:
  25. return " ".join(str(query).split())
  26. def build_query_audit(
  27. *,
  28. query: str,
  29. context: KnowledgeAccessContext,
  30. mode: str,
  31. evidence: tuple[KnowledgeEvidence, ...] | list[KnowledgeEvidence],
  32. cited_points: tuple[tuple[str, int], ...] = (),
  33. degraded_components: tuple[str, ...] = (),
  34. latency_ms: int,
  35. ) -> KnowledgeQueryAuditRecord:
  36. normalized = normalize_audited_query(query)
  37. retriever_counts: Counter[str] = Counter()
  38. for item in evidence:
  39. retriever_counts.update(set(item.retriever.split("+")))
  40. minimized_points = tuple(
  41. {
  42. "point_key": point_key,
  43. "point_revision": int(point_revision),
  44. }
  45. for point_key, point_revision in dict.fromkeys(cited_points)
  46. )
  47. return KnowledgeQueryAuditRecord(
  48. uid=str(uuid4()),
  49. query_hash=hashlib.sha256(normalized.encode("utf-8")).hexdigest(),
  50. user_id=context.subject_id or None,
  51. roles=tuple(sorted(context.roles)),
  52. business_domain_uids=tuple(sorted(context.business_domain_uids)),
  53. mode=str(mode),
  54. retriever_counts=dict(sorted(retriever_counts.items())),
  55. cited_points=minimized_points,
  56. degraded_components=tuple(sorted(set(degraded_components))),
  57. correlation_id=context.correlation_id,
  58. latency_ms=max(0, int(latency_ms)),
  59. )
  60. class SqlKnowledgeQueryAuditRepository:
  61. def __init__(self, session):
  62. self._session = session
  63. def record(self, record: KnowledgeQueryAuditRecord) -> None:
  64. try:
  65. self._session.execute(
  66. text(
  67. """
  68. INSERT INTO public.knowledge_query_audits (
  69. id, query_hash, user_id, roles,
  70. business_domain_uids, mode, retriever_counts,
  71. cited_points, degraded_components, correlation_id,
  72. latency_ms
  73. ) VALUES (
  74. CAST(:id AS uuid), :query_hash,
  75. CAST(:user_id AS uuid), CAST(:roles AS jsonb),
  76. CAST(:business_domain_uids AS jsonb), :mode,
  77. CAST(:retriever_counts AS jsonb),
  78. CAST(:cited_points AS jsonb),
  79. CAST(:degraded_components AS jsonb),
  80. CAST(:correlation_id AS uuid), :latency_ms
  81. )
  82. """
  83. ),
  84. {
  85. "id": record.uid,
  86. "query_hash": record.query_hash,
  87. "user_id": record.user_id,
  88. "roles": json.dumps(record.roles),
  89. "business_domain_uids": json.dumps(
  90. record.business_domain_uids
  91. ),
  92. "mode": record.mode,
  93. "retriever_counts": json.dumps(record.retriever_counts),
  94. "cited_points": json.dumps(record.cited_points),
  95. "degraded_components": json.dumps(
  96. record.degraded_components
  97. ),
  98. "correlation_id": record.correlation_id,
  99. "latency_ms": record.latency_ms,
  100. },
  101. )
  102. self._session.commit()
  103. except Exception:
  104. self._session.rollback()
  105. raise
  106. def list(self, *, limit: int = 100) -> tuple[dict[str, Any], ...]:
  107. rows = (
  108. self._session.execute(
  109. text(
  110. """
  111. SELECT id::text, query_hash, user_id::text,
  112. roles, business_domain_uids, mode,
  113. retriever_counts, cited_points,
  114. degraded_components, correlation_id::text,
  115. latency_ms, created_at
  116. FROM public.knowledge_query_audits
  117. ORDER BY created_at DESC, id DESC
  118. LIMIT :limit
  119. """
  120. ),
  121. {"limit": max(1, min(int(limit), 200))},
  122. )
  123. .mappings()
  124. .all()
  125. )
  126. return tuple(
  127. {
  128. **dict(row),
  129. "created_at": (
  130. row["created_at"].isoformat()
  131. if row["created_at"] is not None
  132. else None
  133. ),
  134. }
  135. for row in rows
  136. )