tracing.py 6.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190
  1. """Best-effort tracing helpers for the formal creation pipeline."""
  2. from __future__ import annotations
  3. import hashlib
  4. import logging
  5. import os
  6. import time
  7. from dataclasses import dataclass, field
  8. from pathlib import Path
  9. from typing import Any, Callable, Protocol
  10. from uuid import UUID
  11. from core.config import CreationDbConfig, env_value, load_env_file
  12. from core.text_limits import clip_text
  13. logger = logging.getLogger(__name__)
  14. @dataclass(frozen=True)
  15. class TraceConfig:
  16. enabled: bool = True
  17. capture_request: bool = True
  18. capture_response: bool = True
  19. max_payload_chars: int = 50_000
  20. @classmethod
  21. def from_env(cls, env_file: str | Path = ".env") -> "TraceConfig":
  22. file_env = load_env_file(env_file)
  23. return cls(
  24. enabled=_env_bool(env_value("CK_TRACE_ENABLED", file_env, "true")),
  25. capture_request=_env_bool(env_value("CK_TRACE_CAPTURE_REQUEST", file_env, "true")),
  26. capture_response=_env_bool(env_value("CK_TRACE_CAPTURE_RESPONSE", file_env, "true")),
  27. max_payload_chars=_env_int(env_value("CK_TRACE_MAX_PAYLOAD_CHARS", file_env, "50000"), 50_000),
  28. )
  29. @dataclass(frozen=True)
  30. class TraceContext:
  31. pipeline_run_id: UUID | None = None
  32. pipeline_job_id: UUID | None = None
  33. acquisition_run_id: UUID | None = None
  34. acquisition_job_id: UUID | None = None
  35. query_id: UUID | None = None
  36. item_id: UUID | None = None
  37. decode_job_id: UUID | None = None
  38. decode_result_id: UUID | None = None
  39. payload_draft_id: UUID | None = None
  40. ingest_record_id: UUID | None = None
  41. platform: str | None = None
  42. stage: str | None = None
  43. substage: str | None = None
  44. metadata: dict[str, Any] = field(default_factory=dict)
  45. def child(self, **updates: Any) -> "TraceContext":
  46. values = {
  47. "pipeline_run_id": self.pipeline_run_id,
  48. "pipeline_job_id": self.pipeline_job_id,
  49. "acquisition_run_id": self.acquisition_run_id,
  50. "acquisition_job_id": self.acquisition_job_id,
  51. "query_id": self.query_id,
  52. "item_id": self.item_id,
  53. "decode_job_id": self.decode_job_id,
  54. "decode_result_id": self.decode_result_id,
  55. "payload_draft_id": self.payload_draft_id,
  56. "ingest_record_id": self.ingest_record_id,
  57. "platform": self.platform,
  58. "stage": self.stage,
  59. "substage": self.substage,
  60. "metadata": dict(self.metadata),
  61. }
  62. metadata = updates.pop("metadata", None)
  63. values.update(updates)
  64. if metadata:
  65. values["metadata"].update(metadata)
  66. return TraceContext(**values)
  67. def as_payload(self) -> dict[str, Any]:
  68. out: dict[str, Any] = {
  69. "pipeline_run_id": self.pipeline_run_id,
  70. "pipeline_job_id": self.pipeline_job_id,
  71. "acquisition_run_id": self.acquisition_run_id,
  72. "acquisition_job_id": self.acquisition_job_id,
  73. "query_id": self.query_id,
  74. "item_id": self.item_id,
  75. "decode_job_id": self.decode_job_id,
  76. "decode_result_id": self.decode_result_id,
  77. "payload_draft_id": self.payload_draft_id,
  78. "ingest_record_id": self.ingest_record_id,
  79. "platform": self.platform,
  80. "stage": self.stage,
  81. "substage": self.substage,
  82. **self.metadata,
  83. }
  84. return {key: str(value) if isinstance(value, UUID) else value for key, value in out.items() if value is not None}
  85. class TraceWriter(Protocol):
  86. def event(self, *, context: TraceContext, stage: str, event_type: str, **kwargs: Any) -> None:
  87. ...
  88. def candidate_hit(self, *, context: TraceContext, **kwargs: Any) -> None:
  89. ...
  90. def llm_call(self, *, context: TraceContext, stage: str, substage: str | None = None, **kwargs: Any) -> None:
  91. ...
  92. class NoopTraceWriter:
  93. def event(self, *, context: TraceContext, stage: str, event_type: str, **kwargs: Any) -> None:
  94. return None
  95. def candidate_hit(self, *, context: TraceContext, **kwargs: Any) -> None:
  96. return None
  97. def llm_call(self, *, context: TraceContext, stage: str, substage: str | None = None, **kwargs: Any) -> None:
  98. return None
  99. def new_trace_writer(db_config: CreationDbConfig, *, env_file: str | Path = ".env") -> TraceWriter:
  100. config = TraceConfig.from_env(env_file)
  101. if not config.enabled:
  102. return NoopTraceWriter()
  103. from pipeline.postgres import PostgresTraceWriter
  104. return PostgresTraceWriter(db_config=db_config, config=config)
  105. def hash_prompt(value: str | None) -> str | None:
  106. if not value:
  107. return None
  108. return hashlib.sha256(value.encode("utf-8")).hexdigest()
  109. def redact_headers(headers: dict[str, Any] | None) -> dict[str, Any]:
  110. if not headers:
  111. return {}
  112. out = dict(headers)
  113. for key in list(out):
  114. if key.lower() in {"authorization", "api-key", "x-api-key"}:
  115. out[key] = "<redacted>"
  116. return out
  117. def sanitize_payload(value: Any, *, max_chars: int) -> Any:
  118. if value is None or isinstance(value, (bool, int, float)):
  119. return value
  120. if isinstance(value, UUID):
  121. return str(value)
  122. if isinstance(value, str):
  123. if value.startswith("data:") and ";base64," in value[:80]:
  124. head, _, data = value.partition(",")
  125. digest = hashlib.sha256(data.encode("utf-8")).hexdigest() if data else ""
  126. return {
  127. "omitted": "base64_data_url",
  128. "media_type": head[:80],
  129. "chars": len(value),
  130. "sha256": digest,
  131. }
  132. return clip_text(value, max_chars)
  133. if isinstance(value, dict):
  134. return {
  135. str(key): sanitize_payload(item, max_chars=max_chars)
  136. for key, item in value.items()
  137. }
  138. if isinstance(value, (list, tuple)):
  139. return [sanitize_payload(item, max_chars=max_chars) for item in value]
  140. if hasattr(value, "model_dump"):
  141. return sanitize_payload(value.model_dump(mode="json"), max_chars=max_chars)
  142. return clip_text(str(value), max_chars)
  143. def timed_ms(start: float) -> int:
  144. return max(0, int((time.perf_counter() - start) * 1000))
  145. def _env_bool(value: str) -> bool:
  146. return value.strip().lower() not in {"0", "false", "no", "off"}
  147. def _env_int(value: str, default: int) -> int:
  148. try:
  149. parsed = int(value)
  150. except (TypeError, ValueError):
  151. logger.warning("invalid trace integer config %r; using %s", value, default)
  152. return default
  153. return parsed if parsed > 0 else default
  154. def current_env_file(default: str = ".env") -> str:
  155. return os.getenv("CK_ENV_FILE", default)