mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 12:11:45 +00:00
Store execution traces in a new request_traces table
One row per execution holding its span tree as JSONB, linked to Logs rows by request, message, activity and workflow-run ids. message_id cascades so deleting or truncating a conversation drops its traces; a daily beat task enforces TRACES_RETENTION_DAYS.
This commit is contained in:
1 parent
f6269c483f
commit
f3f817b1e8
7 files changed
+686
-10
No files matched your search
@@ -0,0 +1,84 @@
|
||||
"""0037 request_traces — one execution trace per request for the Logs UI.
|
||||
|
||||
Each row is the full span tree of one execution (a chat turn, a tool-approval
|
||||
continuation, a scheduled or webhook run, a search, a graph extraction):
|
||||
agent runs, LLM calls, tool calls, retrieval and embeddings with their
|
||||
timings, stored as a ``spans`` JSONB array. The Logs UI loads a trace whole,
|
||||
so one row per trace keeps the write to a single INSERT and lets retention
|
||||
and conversation deletion remove a trace in one step.
|
||||
|
||||
``message_id`` cascades: deleting a conversation, or truncating it when a
|
||||
turn is superseded, removes that turn's traces with it. The link ids
|
||||
(``request_id``, ``activity_id``, ``workflow_run_id``) carry partial indexes
|
||||
because only the Logs rows that have them look traces up by them.
|
||||
|
||||
``status`` is the only CHECK; span kinds and sources are validated in code so
|
||||
new ones need no migration.
|
||||
|
||||
Idempotent both ways.
|
||||
|
||||
Revision ID: 0037_request_traces
|
||||
Revises: 0036_device_audit_created_idx
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "0037_request_traces"
|
||||
down_revision: Union[str, None] = "0036_device_audit_created_idx"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS request_traces (
|
||||
id UUID PRIMARY KEY,
|
||||
request_id TEXT,
|
||||
message_id UUID REFERENCES conversation_messages(id) ON DELETE CASCADE,
|
||||
conversation_id UUID,
|
||||
activity_id TEXT,
|
||||
workflow_run_id UUID,
|
||||
user_id TEXT,
|
||||
agent_id UUID,
|
||||
source TEXT NOT NULL,
|
||||
name TEXT,
|
||||
status TEXT NOT NULL
|
||||
CONSTRAINT request_traces_status_chk
|
||||
CHECK (status IN ('ok', 'error', 'paused', 'cancelled')),
|
||||
started_at TIMESTAMPTZ NOT NULL,
|
||||
duration_ms INTEGER NOT NULL DEFAULT 0,
|
||||
span_count INTEGER NOT NULL DEFAULT 0,
|
||||
dropped_spans INTEGER NOT NULL DEFAULT 0,
|
||||
summary JSONB NOT NULL DEFAULT '{}'::jsonb,
|
||||
spans JSONB NOT NULL DEFAULT '[]'::jsonb,
|
||||
otel_trace_id TEXT,
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
||||
);
|
||||
"""
|
||||
)
|
||||
op.execute(
|
||||
"""
|
||||
CREATE INDEX IF NOT EXISTS request_traces_request_idx
|
||||
ON request_traces (request_id) WHERE request_id IS NOT NULL;
|
||||
CREATE INDEX IF NOT EXISTS request_traces_message_idx
|
||||
ON request_traces (message_id) WHERE message_id IS NOT NULL;
|
||||
CREATE INDEX IF NOT EXISTS request_traces_activity_idx
|
||||
ON request_traces (activity_id) WHERE activity_id IS NOT NULL;
|
||||
CREATE INDEX IF NOT EXISTS request_traces_workflow_run_idx
|
||||
ON request_traces (workflow_run_id) WHERE workflow_run_id IS NOT NULL;
|
||||
CREATE INDEX IF NOT EXISTS request_traces_user_started_idx
|
||||
ON request_traces (user_id, started_at DESC);
|
||||
CREATE INDEX IF NOT EXISTS request_traces_agent_started_idx
|
||||
ON request_traces (agent_id, started_at DESC) WHERE agent_id IS NOT NULL;
|
||||
CREATE INDEX IF NOT EXISTS request_traces_created_idx
|
||||
ON request_traces (created_at);
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.execute("DROP TABLE IF EXISTS request_traces;")
|
||||
@@ -608,6 +608,11 @@ def setup_periodic_tasks(sender, **kwargs):
|
||||
cleanup_guardrail_events.s(),
|
||||
name="cleanup-guardrail-events",
|
||||
)
|
||||
sender.add_periodic_task(
|
||||
timedelta(hours=24),
|
||||
cleanup_traces.s(),
|
||||
name="cleanup-traces",
|
||||
)
|
||||
sender.add_periodic_task(
|
||||
timedelta(hours=24),
|
||||
cleanup_orphan_memories.s(),
|
||||
@@ -827,6 +832,30 @@ def cleanup_guardrail_events(self):
|
||||
return {"deleted": deleted, "ttl_days": ttl_days}
|
||||
|
||||
|
||||
@celery.task(bind=True, acks_late=False)
|
||||
def cleanup_traces(self):
|
||||
"""Delete ``request_traces`` rows older than ``TRACES_RETENTION_DAYS``.
|
||||
|
||||
Every chat turn, scheduled run and search writes a trace, and each one
|
||||
carries content previews, so the table is bounded by a retention window
|
||||
like the other per-request journals.
|
||||
"""
|
||||
from docsgpt.core.settings import settings
|
||||
if not settings.POSTGRES_URI:
|
||||
return {"deleted": 0, "skipped": "POSTGRES_URI not set"}
|
||||
|
||||
from docsgpt.storage.db.engine import get_engine
|
||||
from docsgpt.storage.db.repositories.request_traces import (
|
||||
RequestTracesRepository,
|
||||
)
|
||||
|
||||
ttl_days = settings.TRACES_RETENTION_DAYS
|
||||
engine = get_engine()
|
||||
with engine.begin() as conn:
|
||||
deleted = RequestTracesRepository(conn).purge_older_than(ttl_days)
|
||||
return {"deleted": deleted, "ttl_days": ttl_days}
|
||||
|
||||
|
||||
@celery.task(bind=True, acks_late=False)
|
||||
def cleanup_orphan_memories(self):
|
||||
"""Sweep orphan memories left by the 0009 FK-to-trigger orphan window.
|
||||
|
||||
@@ -825,6 +825,42 @@ Index(
|
||||
Index("ix_guardrail_events_message", guardrail_events_table.c.message_id)
|
||||
Index("ix_guardrail_events_created", guardrail_events_table.c.created_at)
|
||||
|
||||
# One execution trace per request (chat turn, continuation, scheduled or
|
||||
# webhook run, search, graph extraction): the span tree as a JSONB array,
|
||||
# rendered as a waterfall in the Logs UI. ``message_id`` cascades so deleting
|
||||
# or truncating a conversation removes its traces. Migration 0037.
|
||||
request_traces_table = Table(
|
||||
"request_traces",
|
||||
metadata,
|
||||
Column("id", UUID(as_uuid=True), primary_key=True),
|
||||
Column("request_id", Text),
|
||||
Column(
|
||||
"message_id",
|
||||
UUID(as_uuid=True),
|
||||
ForeignKey("conversation_messages.id", ondelete="CASCADE"),
|
||||
),
|
||||
Column("conversation_id", UUID(as_uuid=True)),
|
||||
Column("activity_id", Text),
|
||||
Column("workflow_run_id", UUID(as_uuid=True)),
|
||||
Column("user_id", Text),
|
||||
Column("agent_id", UUID(as_uuid=True)),
|
||||
Column("source", Text, nullable=False),
|
||||
Column("name", Text),
|
||||
# ok | error | paused | cancelled
|
||||
Column("status", Text, nullable=False),
|
||||
Column("started_at", DateTime(timezone=True), nullable=False),
|
||||
Column("duration_ms", Integer, nullable=False, server_default="0"),
|
||||
Column("span_count", Integer, nullable=False, server_default="0"),
|
||||
Column("dropped_spans", Integer, nullable=False, server_default="0"),
|
||||
Column("summary", JSONB, nullable=False, server_default="{}"),
|
||||
Column("spans", JSONB, nullable=False, server_default="[]"),
|
||||
Column("otel_trace_id", Text),
|
||||
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
)
|
||||
|
||||
Index("request_traces_user_started_idx", request_traces_table.c.user_id, request_traces_table.c.started_at)
|
||||
Index("request_traces_created_idx", request_traces_table.c.created_at)
|
||||
|
||||
tool_call_attempts_table = Table(
|
||||
"tool_call_attempts",
|
||||
metadata,
|
||||
|
||||
@@ -0,0 +1,219 @@
|
||||
"""Repository for ``request_traces``: one stored execution trace per request."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
import json
|
||||
import logging
|
||||
from typing import Any, Dict, Iterable, List, Optional
|
||||
|
||||
from sqlalchemy import Connection, text
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
from docsgpt.storage.db.base_repository import looks_like_uuid, row_to_dict
|
||||
from docsgpt.storage.db.serialization import PGNativeJSONEncoder
|
||||
from docsgpt.utils import strip_null_bytes
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
#: Columns a trace can be looked up by, mapped to their SQL type.
|
||||
REF_FIELDS: Dict[str, str] = {
|
||||
"id": "uuid",
|
||||
"request_id": "text",
|
||||
"message_id": "uuid",
|
||||
"activity_id": "text",
|
||||
"workflow_run_id": "uuid",
|
||||
}
|
||||
|
||||
_SUMMARY_COLUMNS = (
|
||||
"id, request_id, message_id, conversation_id, activity_id, workflow_run_id, "
|
||||
"user_id, agent_id, source, name, status, started_at, duration_ms, "
|
||||
"span_count, dropped_spans, summary, otel_trace_id"
|
||||
)
|
||||
|
||||
|
||||
def _dump_jsonb(value: Any) -> str:
|
||||
return json.dumps(strip_null_bytes(value), cls=PGNativeJSONEncoder)
|
||||
|
||||
|
||||
def _uuid_or_none(value: Optional[str]) -> Optional[str]:
|
||||
return str(value) if value and looks_like_uuid(str(value)) else None
|
||||
|
||||
|
||||
def _started_at(record: Dict[str, Any]) -> datetime.datetime:
|
||||
started_ns = record.get("started_at_ns")
|
||||
if started_ns:
|
||||
return datetime.datetime.fromtimestamp(started_ns / 1e9, tz=datetime.timezone.utc)
|
||||
return datetime.datetime.now(datetime.timezone.utc)
|
||||
|
||||
|
||||
def _scope_clause(user_id: Optional[str], agent_id: Optional[str]) -> tuple[str, Dict[str, Any]]:
|
||||
"""Owner scoping: an owned agent's traces, else the caller's own traces."""
|
||||
if agent_id:
|
||||
return "agent_id = CAST(:scope_agent AS uuid)", {"scope_agent": str(agent_id)}
|
||||
return "user_id = :scope_user", {"scope_user": user_id}
|
||||
|
||||
|
||||
class RequestTracesRepository:
|
||||
def __init__(self, conn: Connection) -> None:
|
||||
self._conn = conn
|
||||
|
||||
def insert(self, record: Dict[str, Any]) -> bool:
|
||||
"""Insert one finished trace (the dict from ``Trace.to_record``).
|
||||
|
||||
A foreign-key failure means the message was deleted before the trace
|
||||
flushed (e.g. the turn was superseded); the trace describes a turn the
|
||||
user discarded, so it is dropped rather than raised. The insert runs
|
||||
in a savepoint so that drop leaves the caller's transaction usable.
|
||||
|
||||
Args:
|
||||
record: The trace record.
|
||||
|
||||
Returns:
|
||||
True when the row was written.
|
||||
"""
|
||||
params = {
|
||||
"id": record["id"],
|
||||
"request_id": record.get("request_id"),
|
||||
"message_id": _uuid_or_none(record.get("message_id")),
|
||||
"conversation_id": _uuid_or_none(record.get("conversation_id")),
|
||||
"activity_id": record.get("activity_id"),
|
||||
"workflow_run_id": _uuid_or_none(record.get("workflow_run_id")),
|
||||
"user_id": record.get("user_id"),
|
||||
"agent_id": _uuid_or_none(record.get("agent_id")),
|
||||
"source": record.get("source") or "unknown",
|
||||
"name": strip_null_bytes(record.get("name")),
|
||||
"status": record.get("status") or "ok",
|
||||
"started_at": _started_at(record),
|
||||
"duration_ms": int(record.get("duration_ms") or 0),
|
||||
"span_count": int(record.get("span_count") or 0),
|
||||
"dropped_spans": int(record.get("dropped_spans") or 0),
|
||||
"summary": _dump_jsonb(record.get("summary") or {}),
|
||||
"spans": _dump_jsonb(record.get("spans") or []),
|
||||
"otel_trace_id": record.get("otel_trace_id"),
|
||||
}
|
||||
statement = text(
|
||||
"""
|
||||
INSERT INTO request_traces
|
||||
(id, request_id, message_id, conversation_id, activity_id,
|
||||
workflow_run_id, user_id, agent_id, source, name, status,
|
||||
started_at, duration_ms, span_count, dropped_spans, summary,
|
||||
spans, otel_trace_id)
|
||||
VALUES
|
||||
(CAST(:id AS uuid), :request_id, CAST(:message_id AS uuid),
|
||||
CAST(:conversation_id AS uuid), :activity_id,
|
||||
CAST(:workflow_run_id AS uuid), :user_id, CAST(:agent_id AS uuid),
|
||||
:source, :name, :status, :started_at, :duration_ms, :span_count,
|
||||
:dropped_spans, CAST(:summary AS jsonb), CAST(:spans AS jsonb),
|
||||
:otel_trace_id)
|
||||
ON CONFLICT (id) DO NOTHING
|
||||
"""
|
||||
)
|
||||
try:
|
||||
with self._conn.begin_nested():
|
||||
result = self._conn.execute(statement, params)
|
||||
except IntegrityError as exc:
|
||||
sqlstate = getattr(getattr(exc, "orig", None), "sqlstate", None)
|
||||
if sqlstate != "23503":
|
||||
raise
|
||||
logger.info("Dropped trace %s: its message no longer exists", record["id"])
|
||||
return False
|
||||
return (result.rowcount or 0) > 0
|
||||
|
||||
def list_by_ref(
|
||||
self,
|
||||
field: str,
|
||||
value: str,
|
||||
*,
|
||||
user_id: Optional[str],
|
||||
agent_id: Optional[str] = None,
|
||||
limit: int = 20,
|
||||
) -> List[dict]:
|
||||
"""Full traces matching ``field = value``, oldest first, owner-scoped.
|
||||
|
||||
Args:
|
||||
field: One of :data:`REF_FIELDS`.
|
||||
value: The id to match.
|
||||
user_id: The caller; used when ``agent_id`` is not given.
|
||||
agent_id: An agent the caller owns; scopes to that agent's traces.
|
||||
limit: Maximum traces returned.
|
||||
|
||||
Returns:
|
||||
Trace dicts including ``spans``; empty for an unknown field or id.
|
||||
"""
|
||||
sql_type = REF_FIELDS.get(field)
|
||||
if sql_type is None or not value:
|
||||
return []
|
||||
if sql_type == "uuid" and not looks_like_uuid(str(value)):
|
||||
return []
|
||||
if not agent_id and not user_id:
|
||||
return []
|
||||
scope, scope_params = _scope_clause(user_id, agent_id)
|
||||
result = self._conn.execute(
|
||||
text(
|
||||
f"""
|
||||
SELECT {_SUMMARY_COLUMNS}, spans FROM request_traces
|
||||
WHERE {field} = CAST(:value AS {sql_type}) AND {scope}
|
||||
ORDER BY started_at
|
||||
LIMIT :limit
|
||||
"""
|
||||
),
|
||||
{"value": str(value), "limit": max(1, min(int(limit), 100)), **scope_params},
|
||||
)
|
||||
return [row_to_dict(row) for row in result.fetchall()]
|
||||
|
||||
def summaries_for_refs(
|
||||
self,
|
||||
refs: Dict[str, Iterable[str]],
|
||||
*,
|
||||
user_id: Optional[str],
|
||||
agent_id: Optional[str] = None,
|
||||
) -> Dict[str, Dict[str, List[dict]]]:
|
||||
"""Trace summaries (no spans) for a page of Logs rows.
|
||||
|
||||
Args:
|
||||
refs: ``{field: [ids...]}`` for fields in :data:`REF_FIELDS`.
|
||||
user_id: The caller; used when ``agent_id`` is not given.
|
||||
agent_id: An agent the caller owns.
|
||||
|
||||
Returns:
|
||||
``{field: {id: [summary, ...]}}`` with summaries oldest first.
|
||||
"""
|
||||
out: Dict[str, Dict[str, List[dict]]] = {}
|
||||
if not agent_id and not user_id:
|
||||
return out
|
||||
scope, scope_params = _scope_clause(user_id, agent_id)
|
||||
for field, values in refs.items():
|
||||
sql_type = REF_FIELDS.get(field)
|
||||
if sql_type is None:
|
||||
continue
|
||||
ids = sorted({str(v) for v in values if v})
|
||||
if sql_type == "uuid":
|
||||
ids = [v for v in ids if looks_like_uuid(v)]
|
||||
if not ids:
|
||||
continue
|
||||
result = self._conn.execute(
|
||||
text(
|
||||
f"""
|
||||
SELECT {_SUMMARY_COLUMNS} FROM request_traces
|
||||
WHERE {field} = ANY(CAST(:ids AS {sql_type}[])) AND {scope}
|
||||
ORDER BY started_at
|
||||
"""
|
||||
),
|
||||
{"ids": ids, **scope_params},
|
||||
)
|
||||
for row in result.fetchall():
|
||||
data = row_to_dict(row)
|
||||
out.setdefault(field, {}).setdefault(str(data[field]), []).append(data)
|
||||
return out
|
||||
|
||||
def purge_older_than(self, days: int) -> int:
|
||||
"""Delete traces older than ``days`` (the retention window)."""
|
||||
result = self._conn.execute(
|
||||
text(
|
||||
"DELETE FROM request_traces "
|
||||
"WHERE created_at < NOW() - CAST(:days || ' days' AS interval)"
|
||||
),
|
||||
{"days": str(max(1, days))},
|
||||
)
|
||||
return result.rowcount or 0
|
||||
@@ -299,7 +299,7 @@ class TestSetupPeriodicTasks:
|
||||
|
||||
setup_periodic_tasks(sender)
|
||||
|
||||
assert sender.add_periodic_task.call_count == 14
|
||||
assert sender.add_periodic_task.call_count == 15
|
||||
|
||||
calls = sender.add_periodic_task.call_args_list
|
||||
|
||||
@@ -326,20 +326,23 @@ class TestSetupPeriodicTasks:
|
||||
# guardrail_events retention sweep (24h)
|
||||
assert calls[8][0][0] == timedelta(hours=24)
|
||||
assert calls[8][1].get("name") == "cleanup-guardrail-events"
|
||||
# orphan memories sweep (24h)
|
||||
# request_traces retention sweep (24h)
|
||||
assert calls[9][0][0] == timedelta(hours=24)
|
||||
assert calls[9][1].get("name") == "cleanup-orphan-memories"
|
||||
assert calls[9][1].get("name") == "cleanup-traces"
|
||||
# orphan memories sweep (24h)
|
||||
assert calls[10][0][0] == timedelta(hours=24)
|
||||
assert calls[10][1].get("name") == "cleanup-orphan-memories"
|
||||
# scheduler dispatcher
|
||||
assert calls[10][1].get("name") == "dispatch-scheduled-runs"
|
||||
assert calls[11][1].get("name") == "dispatch-scheduled-runs"
|
||||
# schedule runs cleanup (24h)
|
||||
assert calls[11][0][0] == timedelta(hours=24)
|
||||
assert calls[11][1].get("name") == "cleanup-schedule-runs"
|
||||
assert calls[12][0][0] == timedelta(hours=24)
|
||||
assert calls[12][1].get("name") == "cleanup-schedule-runs"
|
||||
# sandbox session reaper (60s)
|
||||
assert calls[12][0][0] == timedelta(seconds=60)
|
||||
assert calls[12][1].get("name") == "reap-sandbox-sessions"
|
||||
assert calls[13][0][0] == timedelta(seconds=60)
|
||||
assert calls[13][1].get("name") == "reap-sandbox-sessions"
|
||||
# stale workflow-run reaper (5m)
|
||||
assert calls[13][0][0] == timedelta(seconds=300)
|
||||
assert calls[13][1].get("name") == "reap-stale-workflow-runs"
|
||||
assert calls[14][0][0] == timedelta(seconds=300)
|
||||
assert calls[14][1].get("name") == "reap-stale-workflow-runs"
|
||||
|
||||
|
||||
class TestMcpOauthTask:
|
||||
@@ -679,6 +682,67 @@ class TestCleanupMessageEventsTask:
|
||||
assert [r["sequence_no"] for r in rows] == [1]
|
||||
|
||||
|
||||
class TestCleanupTracesTask:
|
||||
"""Retention janitor for ``request_traces``."""
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_skips_when_postgres_uri_missing(self, monkeypatch):
|
||||
from docsgpt.api.user.tasks import cleanup_traces
|
||||
from docsgpt.core.settings import settings
|
||||
|
||||
monkeypatch.setattr(settings, "POSTGRES_URI", None, raising=False)
|
||||
|
||||
assert cleanup_traces.run() == {"deleted": 0, "skipped": "POSTGRES_URI not set"}
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_deletes_traces_past_retention_window(self, pg_conn, monkeypatch):
|
||||
import time
|
||||
import uuid
|
||||
|
||||
from sqlalchemy import text as _text
|
||||
|
||||
from docsgpt.api.user.tasks import cleanup_traces
|
||||
from docsgpt.core.settings import settings
|
||||
from docsgpt.storage.db.repositories.request_traces import (
|
||||
RequestTracesRepository,
|
||||
)
|
||||
|
||||
repo = RequestTracesRepository(pg_conn)
|
||||
for request_id in ("stale", "fresh"):
|
||||
repo.insert(
|
||||
{
|
||||
"id": str(uuid.uuid4()),
|
||||
"request_id": request_id,
|
||||
"user_id": "u1",
|
||||
"source": "stream",
|
||||
"status": "ok",
|
||||
"started_at_ns": time.time_ns(),
|
||||
}
|
||||
)
|
||||
pg_conn.execute(
|
||||
_text(
|
||||
"UPDATE request_traces SET created_at = now() - interval '45 days' "
|
||||
"WHERE request_id = 'stale'"
|
||||
)
|
||||
)
|
||||
monkeypatch.setattr(settings, "POSTGRES_URI", "postgresql://stub", raising=False)
|
||||
monkeypatch.setattr(settings, "TRACES_RETENTION_DAYS", 30)
|
||||
|
||||
@contextmanager
|
||||
def _fake_begin():
|
||||
yield pg_conn
|
||||
|
||||
fake_engine = MagicMock()
|
||||
fake_engine.begin = _fake_begin
|
||||
|
||||
with patch("docsgpt.storage.db.engine.get_engine", return_value=fake_engine):
|
||||
result = cleanup_traces.run()
|
||||
|
||||
assert result == {"deleted": 1, "ttl_days": 30}
|
||||
remaining = pg_conn.execute(_text("SELECT request_id FROM request_traces")).scalars().all()
|
||||
assert remaining == ["fresh"]
|
||||
|
||||
|
||||
class TestCleanupOrphanMemoriesTask:
|
||||
"""Sweeps orphan memories from the FK-to-trigger orphan window."""
|
||||
|
||||
|
||||
@@ -0,0 +1,179 @@
|
||||
"""Tests for RequestTracesRepository against a real Postgres instance."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
import uuid
|
||||
|
||||
from sqlalchemy import text
|
||||
|
||||
from docsgpt.storage.db.repositories.conversations import ConversationsRepository
|
||||
from docsgpt.storage.db.repositories.request_traces import RequestTracesRepository
|
||||
|
||||
|
||||
def _record(**overrides):
|
||||
base = {
|
||||
"id": str(uuid.uuid4()),
|
||||
"request_id": "req-1",
|
||||
"message_id": None,
|
||||
"conversation_id": None,
|
||||
"activity_id": None,
|
||||
"workflow_run_id": None,
|
||||
"user_id": "u1",
|
||||
"agent_id": None,
|
||||
"source": "stream",
|
||||
"name": "stream",
|
||||
"status": "ok",
|
||||
"started_at_ns": time.time_ns(),
|
||||
"duration_ms": 1234,
|
||||
"span_count": 1,
|
||||
"dropped_spans": 0,
|
||||
"summary": {"llm_calls": 1},
|
||||
"spans": [
|
||||
{
|
||||
"id": "s1",
|
||||
"parent_id": None,
|
||||
"kind": "llm",
|
||||
"name": "chat m",
|
||||
"status": "ok",
|
||||
"offset_ms": 1.0,
|
||||
"duration_ms": 10.0,
|
||||
"attributes": {"gen_ai.request.model": "m"},
|
||||
}
|
||||
],
|
||||
"otel_trace_id": None,
|
||||
}
|
||||
base.update(overrides)
|
||||
return base
|
||||
|
||||
|
||||
def _message(pg_conn, user_id="u1"):
|
||||
convs = ConversationsRepository(pg_conn)
|
||||
conv = convs.create(user_id, "t")
|
||||
msg = convs.reserve_message(
|
||||
str(conv["id"]), prompt="q", placeholder_response="..."
|
||||
)
|
||||
return str(conv["id"]), str(msg["id"])
|
||||
|
||||
|
||||
class TestInsert:
|
||||
def test_round_trip(self, pg_conn):
|
||||
repo = RequestTracesRepository(pg_conn)
|
||||
record = _record()
|
||||
assert repo.insert(record) is True
|
||||
rows = repo.list_by_ref("request_id", "req-1", user_id="u1")
|
||||
assert len(rows) == 1
|
||||
row = rows[0]
|
||||
assert row["id"] == record["id"]
|
||||
assert row["duration_ms"] == 1234
|
||||
assert row["summary"] == {"llm_calls": 1}
|
||||
assert row["spans"][0]["attributes"]["gen_ai.request.model"] == "m"
|
||||
|
||||
def test_strips_nul_bytes(self, pg_conn):
|
||||
repo = RequestTracesRepository(pg_conn)
|
||||
record = _record(
|
||||
spans=[{"id": "s", "name": "a\x00b", "preview": {"result": "x\x00y"}}]
|
||||
)
|
||||
assert repo.insert(record)
|
||||
row = repo.list_by_ref("id", record["id"], user_id="u1")[0]
|
||||
assert row["spans"][0]["name"] == "ab"
|
||||
assert row["spans"][0]["preview"]["result"] == "xy"
|
||||
|
||||
def test_non_uuid_link_ids_become_null(self, pg_conn):
|
||||
repo = RequestTracesRepository(pg_conn)
|
||||
record = _record(agent_id="not-a-uuid", conversation_id="nope")
|
||||
assert repo.insert(record)
|
||||
row = repo.list_by_ref("id", record["id"], user_id="u1")[0]
|
||||
assert row["agent_id"] is None
|
||||
assert row["conversation_id"] is None
|
||||
|
||||
def test_missing_message_drops_row_and_keeps_transaction_usable(self, pg_conn):
|
||||
repo = RequestTracesRepository(pg_conn)
|
||||
assert repo.insert(_record(message_id=str(uuid.uuid4()))) is False
|
||||
# The savepoint rolled back; the outer transaction still works.
|
||||
assert repo.insert(_record(request_id="req-after")) is True
|
||||
|
||||
def test_duplicate_id_is_ignored(self, pg_conn):
|
||||
repo = RequestTracesRepository(pg_conn)
|
||||
record = _record()
|
||||
assert repo.insert(record) is True
|
||||
assert repo.insert(record) is False
|
||||
|
||||
|
||||
class TestScoping:
|
||||
def test_other_users_cannot_read(self, pg_conn):
|
||||
repo = RequestTracesRepository(pg_conn)
|
||||
repo.insert(_record())
|
||||
assert repo.list_by_ref("request_id", "req-1", user_id="u2") == []
|
||||
|
||||
def test_agent_scope_returns_agent_traces_from_any_user(self, pg_conn):
|
||||
repo = RequestTracesRepository(pg_conn)
|
||||
agent_id = str(uuid.uuid4())
|
||||
repo.insert(_record(user_id="someone-else", agent_id=agent_id))
|
||||
rows = repo.list_by_ref("request_id", "req-1", user_id="owner", agent_id=agent_id)
|
||||
assert len(rows) == 1
|
||||
|
||||
def test_unknown_field_and_bad_uuid_return_empty(self, pg_conn):
|
||||
repo = RequestTracesRepository(pg_conn)
|
||||
repo.insert(_record())
|
||||
assert repo.list_by_ref("user_id", "u1", user_id="u1") == []
|
||||
assert repo.list_by_ref("message_id", "not-uuid", user_id="u1") == []
|
||||
assert repo.list_by_ref("request_id", "req-1", user_id=None) == []
|
||||
|
||||
|
||||
class TestMessageLink:
|
||||
def test_rounds_for_a_message_are_ordered(self, pg_conn):
|
||||
conv_id, msg_id = _message(pg_conn)
|
||||
repo = RequestTracesRepository(pg_conn)
|
||||
first = _record(message_id=msg_id, status="paused", started_at_ns=time.time_ns())
|
||||
second = _record(message_id=msg_id, started_at_ns=time.time_ns() + 1_000_000)
|
||||
repo.insert(second)
|
||||
repo.insert(first)
|
||||
rows = repo.list_by_ref("message_id", msg_id, user_id="u1")
|
||||
assert [r["id"] for r in rows] == [first["id"], second["id"]]
|
||||
|
||||
def test_deleting_the_conversation_cascades(self, pg_conn):
|
||||
conv_id, msg_id = _message(pg_conn)
|
||||
repo = RequestTracesRepository(pg_conn)
|
||||
repo.insert(_record(message_id=msg_id, conversation_id=conv_id))
|
||||
ConversationsRepository(pg_conn).delete(conv_id, "u1")
|
||||
count = pg_conn.execute(text("SELECT count(*) FROM request_traces")).scalar()
|
||||
assert count == 0
|
||||
|
||||
|
||||
class TestSummaries:
|
||||
def test_summaries_grouped_by_field_and_id(self, pg_conn):
|
||||
repo = RequestTracesRepository(pg_conn)
|
||||
a = _record(request_id="r-a")
|
||||
b = _record(request_id="r-b", activity_id="act-1")
|
||||
repo.insert(a)
|
||||
repo.insert(b)
|
||||
out = repo.summaries_for_refs(
|
||||
{"request_id": ["r-a", "r-b", "r-missing"], "activity_id": ["act-1"]},
|
||||
user_id="u1",
|
||||
)
|
||||
assert set(out["request_id"]) == {"r-a", "r-b"}
|
||||
assert out["activity_id"]["act-1"][0]["id"] == b["id"]
|
||||
assert "spans" not in out["request_id"]["r-a"][0]
|
||||
|
||||
def test_summaries_are_scoped(self, pg_conn):
|
||||
repo = RequestTracesRepository(pg_conn)
|
||||
repo.insert(_record(request_id="r-a"))
|
||||
assert repo.summaries_for_refs({"request_id": ["r-a"]}, user_id="u2") == {}
|
||||
|
||||
|
||||
class TestPurge:
|
||||
def test_purge_older_than(self, pg_conn):
|
||||
repo = RequestTracesRepository(pg_conn)
|
||||
old = _record(request_id="old")
|
||||
repo.insert(old)
|
||||
repo.insert(_record(request_id="new"))
|
||||
pg_conn.execute(
|
||||
text(
|
||||
"UPDATE request_traces SET created_at = now() - interval '40 days' "
|
||||
"WHERE request_id = 'old'"
|
||||
)
|
||||
)
|
||||
assert repo.purge_older_than(30) == 1
|
||||
assert repo.list_by_ref("request_id", "old", user_id="u1") == []
|
||||
assert len(repo.list_by_ref("request_id", "new", user_id="u1")) == 1
|
||||
@@ -0,0 +1,65 @@
|
||||
"""Migration round-trip test for 0037_request_traces."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import text
|
||||
|
||||
|
||||
pytestmark = pytest.mark.integration
|
||||
|
||||
|
||||
def _alembic_ini() -> Path:
|
||||
return Path(__file__).resolve().parents[3] / "docsgpt" / "alembic.ini"
|
||||
|
||||
|
||||
def _run_alembic(url: str, *args: str) -> None:
|
||||
subprocess.check_call(
|
||||
[sys.executable, "-m", "alembic", "-c", str(_alembic_ini()), *args],
|
||||
timeout=60,
|
||||
env={**os.environ, "POSTGRES_URI": url},
|
||||
)
|
||||
|
||||
|
||||
def _alembic_version(conn) -> str:
|
||||
return conn.execute(text("SELECT version_num FROM alembic_version")).scalar()
|
||||
|
||||
|
||||
def _table_exists(conn, table: str) -> bool:
|
||||
return conn.execute(text("SELECT to_regclass(:t)"), {"t": f"public.{table}"}).scalar() is not None
|
||||
|
||||
|
||||
_0037 = "0037_request_traces"
|
||||
_0036 = "0036_device_audit_created_idx"
|
||||
|
||||
|
||||
class TestMigration0037RoundTrip:
|
||||
def test_head_has_request_traces(self, pg_engine):
|
||||
with pg_engine.connect() as conn:
|
||||
assert _alembic_version(conn) >= _0037
|
||||
assert _table_exists(conn, "request_traces")
|
||||
|
||||
def test_downgrade_drops_then_upgrade_restores(self, pg_engine):
|
||||
url = pg_engine.url.render_as_string(hide_password=False)
|
||||
_run_alembic(url, "downgrade", _0036)
|
||||
with pg_engine.connect() as conn:
|
||||
assert _alembic_version(conn) == _0036
|
||||
assert not _table_exists(conn, "request_traces")
|
||||
_run_alembic(url, "upgrade", "head")
|
||||
with pg_engine.connect() as conn:
|
||||
assert _table_exists(conn, "request_traces")
|
||||
|
||||
def test_status_check_rejects_unknown_status(self, pg_engine):
|
||||
with pytest.raises(Exception):
|
||||
with pg_engine.begin() as conn:
|
||||
conn.execute(
|
||||
text(
|
||||
"INSERT INTO request_traces (id, source, status, started_at) "
|
||||
"VALUES (gen_random_uuid(), 'stream', 'weird', now())"
|
||||
)
|
||||
)
|
||||
Reference in new issue
Block a user