diff --git a/.gitignore b/.gitignore index bfdb408..5726b3d 100644 --- a/.gitignore +++ b/.gitignore @@ -68,7 +68,10 @@ dmypy.json .ruff_cache/ # Logs -logs/ +logs/* +!logs/traces/ +logs/traces/* +!logs/traces/viewer.html *.log # Database diff --git a/logs/traces/viewer.html b/logs/traces/viewer.html new file mode 100644 index 0000000..88c46bf --- /dev/null +++ b/logs/traces/viewer.html @@ -0,0 +1,1213 @@ + + + + + + Tatlock Trace Viewer + + + +
+

Tatlock Trace Viewer

+
+
+ + +
+
+ +
+ +
+
+
+

Recent Traces

+ +
+
+
+
+ + +
+
+ + +
+
+
+ +
+
+
+
+
+ Loading traces... +
+
+
+ +
+
+

Timeline

+
+
Router
+
Steward
+
Tatlock
+
Expert
+
Tool
+
+
+
+
+
📊
+

Select a trace from the list
or drop a JSON file here

+
+
+
+ +
+
+

Details

+
+
+
+

Select a span to view details

+
+
+
+
+ + + + + + diff --git a/src/core/tracing.py b/src/core/tracing.py new file mode 100644 index 0000000..ce56396 --- /dev/null +++ b/src/core/tracing.py @@ -0,0 +1,434 @@ +""" +Lightweight request tracing for local development. + +Captures the full request flow through Tatlock's multi-agent architecture +as structured JSON traces for debugging and optimization. + +Enable via DEBUG=true environment variable. + +Traces are written to logs/traces/{trace_id}.json +View with logs/traces/viewer.html +""" +from contextlib import asynccontextmanager +from contextvars import ContextVar +from dataclasses import dataclass, field +from datetime import datetime, timezone +from enum import Enum +from pathlib import Path +from typing import Any + +import json +import secrets + +from src.core.logging_config import get_logger + +logger = get_logger(__name__) + + +class SpanType(str, Enum): + """Types of traced operations.""" + ROUTER = "router" + STEWARD = "steward" + TATLOCK = "tatlock" + EXPERT = "expert" + TOOL = "tool" + + +class SpanStatus(str, Enum): + """Span completion status.""" + OK = "ok" + ERROR = "error" + + +@dataclass +class Span: + """A single traced operation.""" + span_id: str + name: str + type: SpanType + start_time: datetime + parent_id: str | None = None + end_time: datetime | None = None + status: SpanStatus = SpanStatus.OK + metadata: dict[str, Any] = field(default_factory=dict) + details: dict[str, Any] = field(default_factory=dict) + children: list[str] = field(default_factory=list) + error: str | None = None + + @property + def duration_ms(self) -> float | None: + """Calculate duration in milliseconds.""" + if self.end_time and self.start_time: + return (self.end_time - self.start_time).total_seconds() * 1000 + return None + + def to_dict(self) -> dict[str, Any]: + """Convert span to dictionary for JSON serialization.""" + result = { + "span_id": self.span_id, + "parent_id": self.parent_id, + "name": self.name, + "type": self.type.value, + "start_time": self.start_time.isoformat(), + "end_time": self.end_time.isoformat() if self.end_time else None, + "duration_ms": round(self.duration_ms, 2) if self.duration_ms else None, + "status": self.status.value, + "metadata": self.metadata if self.metadata else None, + } + # Only include non-empty optional fields + if self.details: + result["details"] = self.details + if self.children: + result["children"] = self.children + if self.error: + result["error"] = self.error + return {k: v for k, v in result.items() if v is not None} + + +@dataclass +class Trace: + """Complete trace of a request.""" + trace_id: str + conversation_id: str | None + user: str + timestamp: datetime + request: dict[str, Any] + spans: list[Span] = field(default_factory=list) + response: dict[str, Any] | None = None + status: str = "in_progress" + + @property + def total_duration_ms(self) -> float | None: + """Calculate total trace duration from span timings.""" + if not self.spans: + return None + start = min(s.start_time for s in self.spans) + ends = [s.end_time for s in self.spans if s.end_time] + if not ends: + return None + end = max(ends) + return (end - start).total_seconds() * 1000 + + def to_dict(self) -> dict[str, Any]: + """Convert trace to dictionary for JSON serialization.""" + return { + "trace_id": self.trace_id, + "conversation_id": self.conversation_id, + "user": self.user, + "timestamp": self.timestamp.isoformat(), + "total_duration_ms": round(self.total_duration_ms, 2) if self.total_duration_ms else None, + "status": self.status, + "request": self.request, + "response": self.response, + "spans": [s.to_dict() for s in self.spans], + } + + +# ContextVar for async-safe trace propagation +_current_trace: ContextVar[Trace | None] = ContextVar("current_trace", default=None) +_current_span: ContextVar[Span | None] = ContextVar("current_span", default=None) + + +def tracing_enabled() -> bool: + """Check if tracing is enabled (requires DEBUG=true).""" + from src.core.config import config + return config.DEBUG + + +def _generate_id(prefix: str = "") -> str: + """Generate unique ID with optional prefix.""" + return f"{prefix}{secrets.token_hex(8)}" + + +def start_trace( + conversation_id: str | None, + user: str, + request: dict[str, Any], +) -> Trace | None: + """ + Start a new trace for a request. + + Args: + conversation_id: Conversation identifier + user: User identifier + request: Request data (should include preview and full) + + Returns: + Trace object if tracing enabled, None otherwise + """ + if not tracing_enabled(): + return None + + trace = Trace( + trace_id=_generate_id("trace_"), + conversation_id=conversation_id, + user=user, + timestamp=datetime.now(timezone.utc), + request=request, + ) + _current_trace.set(trace) + + logger.debug("trace_started", trace_id=trace.trace_id, user=user) + return trace + + +def get_current_trace() -> Trace | None: + """Get the current trace from context.""" + return _current_trace.get() + + +def get_current_span() -> Span | None: + """Get the current span from context.""" + return _current_span.get() + + +def start_span( + name: str, + span_type: SpanType, + metadata: dict[str, Any] | None = None, + details: dict[str, Any] | None = None, +) -> Span | None: + """ + Start a new span within the current trace. + + Args: + name: Span name (e.g., "steward_analysis") + span_type: Type of operation + metadata: Quick-access metadata (shown in timeline) + details: Expandable details (prompts, full responses) + + Returns: + Span object if tracing enabled, None otherwise + """ + trace = get_current_trace() + if not trace: + return None + + parent = get_current_span() + span = Span( + span_id=_generate_id("span_"), + name=name, + type=span_type, + start_time=datetime.now(timezone.utc), + parent_id=parent.span_id if parent else None, + metadata=metadata or {}, + details=details or {}, + ) + + # Add to parent's children list + if parent: + parent.children.append(span.span_id) + + trace.spans.append(span) + _current_span.set(span) + + logger.debug( + "span_started", + span_id=span.span_id, + name=name, + type=span_type.value, + parent_id=span.parent_id, + ) + return span + + +def end_span( + span: Span | None = None, + status: SpanStatus = SpanStatus.OK, + metadata_update: dict[str, Any] | None = None, + details_update: dict[str, Any] | None = None, + error: str | None = None, +) -> None: + """ + End a span and restore parent as current. + + Args: + span: Span to end (defaults to current span) + status: Completion status + metadata_update: Additional metadata to merge + details_update: Additional details to merge + error: Error message if failed + """ + if span is None: + span = get_current_span() + if not span: + return + + span.end_time = datetime.now(timezone.utc) + span.status = status + if error: + span.error = error + span.status = SpanStatus.ERROR + if metadata_update: + span.metadata.update(metadata_update) + if details_update: + span.details.update(details_update) + + # Restore parent span as current + trace = get_current_trace() + if trace and span.parent_id: + parent = next((s for s in trace.spans if s.span_id == span.parent_id), None) + _current_span.set(parent) + else: + _current_span.set(None) + + logger.debug( + "span_ended", + span_id=span.span_id, + duration_ms=span.duration_ms, + status=status.value, + ) + + +def end_trace( + response: dict[str, Any] | None = None, + status: str = "completed", +) -> str | None: + """ + End the current trace and write to file. + + Args: + response: Response data to include + status: Final trace status ("completed" or "error") + + Returns: + Path to trace file if written, None otherwise + """ + trace = get_current_trace() + if not trace: + return None + + trace.response = response + trace.status = status + + # Write trace to file + trace_path = _write_trace(trace) + + # Clear context + _current_trace.set(None) + _current_span.set(None) + + logger.info( + "trace_completed", + trace_id=trace.trace_id, + total_duration_ms=round(trace.total_duration_ms, 2) if trace.total_duration_ms else None, + span_count=len(trace.spans), + path=str(trace_path) if trace_path else None, + ) + + return str(trace_path) if trace_path else None + + +def _write_trace(trace: Trace) -> Path | None: + """Write trace to JSON file.""" + try: + # Ensure traces directory exists + traces_dir = Path("logs/traces") + traces_dir.mkdir(parents=True, exist_ok=True) + + # Write trace file + trace_path = traces_dir / f"{trace.trace_id}.json" + with open(trace_path, "w") as f: + json.dump(trace.to_dict(), f, indent=2, default=str) + + return trace_path + + except Exception as e: + logger.error("trace_write_failed", error=str(e), trace_id=trace.trace_id) + return None + + +@asynccontextmanager +async def trace_span( + name: str, + span_type: SpanType, + metadata: dict[str, Any] | None = None, + details: dict[str, Any] | None = None, +): + """ + Async context manager for tracing a span. + + Automatically handles start/end timing and error capture. + + Usage: + async with trace_span("steward_analysis", SpanType.STEWARD) as span: + result = await analyze_request(...) + if span: + span.metadata["result_count"] = len(result) + + Args: + name: Span name + span_type: Type of operation + metadata: Initial metadata + details: Initial details (expandable in viewer) + + Yields: + Span object or None if tracing disabled + """ + span = start_span(name, span_type, metadata, details) + try: + yield span + except Exception as e: + end_span(span, SpanStatus.ERROR, error=str(e)) + raise + else: + end_span(span, SpanStatus.OK) + + +def add_tool_spans_from_messages(messages: list[Any], parent_span: Span | None = None) -> None: + """ + Extract tool calls from PydanticAI result messages and add as child spans. + + Call this after an agent.run() to capture tool-level timing retroactively. + Note: Since we don't have actual timing, we estimate based on sequence. + + Args: + messages: List from result.new_messages() + parent_span: Parent span to attach tool spans to + """ + trace = get_current_trace() + if not trace or not parent_span: + return + + # Import PydanticAI message types + try: + from pydantic_ai.messages import ModelRequest, ModelResponse, ToolCallPart, ToolReturnPart + except ImportError: + return + + # Track tool calls and their returns + tool_calls: dict[str, dict[str, Any]] = {} + + for msg in messages: + if isinstance(msg, ModelResponse): + for part in msg.parts: + if isinstance(part, ToolCallPart): + tool_calls[part.tool_call_id] = { + "name": part.tool_name, + "args": part.args if hasattr(part, 'args') else {}, + } + elif isinstance(msg, ModelRequest): + for part in msg.parts: + if isinstance(part, ToolReturnPart): + if part.tool_call_id in tool_calls: + tool_info = tool_calls[part.tool_call_id] + # Create a span for this tool call + span = Span( + span_id=_generate_id("span_"), + name=tool_info["name"], + type=SpanType.TOOL, + start_time=parent_span.start_time, # Approximate + end_time=parent_span.end_time or datetime.now(timezone.utc), + parent_id=parent_span.span_id, + status=SpanStatus.OK, + metadata={ + "tool_name": tool_info["name"], + "args_preview": str(tool_info.get("args", {}))[:100], + }, + details={ + "args": tool_info.get("args", {}), + "result": part.content[:2000] if isinstance(part.content, str) else str(part.content)[:2000], + }, + ) + parent_span.children.append(span.span_id) + trace.spans.append(span) diff --git a/src/core/tracing_router.py b/src/core/tracing_router.py new file mode 100644 index 0000000..9d3a00f --- /dev/null +++ b/src/core/tracing_router.py @@ -0,0 +1,153 @@ +""" +Trace viewer router. + +Serves the trace viewer UI and trace files when tracing is enabled. +Only available when DEBUG=true. +""" +from pathlib import Path + +from fastapi import APIRouter, HTTPException +from fastapi.responses import HTMLResponse, JSONResponse + +from src.core.config import config +from src.core.logging_config import get_logger + +logger = get_logger(__name__) + +router = APIRouter(prefix="/traces", tags=["traces"]) + +TRACES_DIR = Path("logs/traces") +VIEWER_PATH = TRACES_DIR / "viewer.html" + + +def tracing_enabled() -> bool: + """Check if tracing is enabled.""" + return config.DEBUG + + +@router.get("", response_class=HTMLResponse) +async def get_trace_viewer(): + """ + Serve the trace viewer UI. + + Returns the standalone HTML viewer for browsing traces. + """ + if not tracing_enabled(): + raise HTTPException(status_code=404, detail="Tracing not enabled") + + if not VIEWER_PATH.exists(): + raise HTTPException(status_code=404, detail="Viewer not found") + + return HTMLResponse(content=VIEWER_PATH.read_text()) + + +@router.get("/list") +async def list_traces( + limit: int = 50, + since_minutes: int | None = None, + status: str | None = None, + search: str | None = None, +): + """ + List available trace files. + + Returns most recent traces first, with basic metadata. + + Args: + limit: Maximum number of traces to return (default 50) + since_minutes: Only return traces from the last N minutes + status: Filter by status (completed, error, streaming) + search: Search in request preview text + """ + if not tracing_enabled(): + raise HTTPException(status_code=404, detail="Tracing not enabled") + + if not TRACES_DIR.exists(): + return {"traces": [], "total": 0} + + import json + from datetime import datetime, timezone, timedelta + + # Calculate cutoff time if filtering by time + cutoff_time = None + if since_minutes: + cutoff_time = datetime.now(timezone.utc) - timedelta(minutes=since_minutes) + + # Get all trace files, sorted by modification time (newest first) + trace_files = sorted( + TRACES_DIR.glob("trace_*.json"), + key=lambda p: p.stat().st_mtime, + reverse=True, + ) + + traces = [] + for path in trace_files: + if len(traces) >= limit: + break + + try: + with open(path) as f: + data = json.load(f) + + # Parse timestamp for filtering + trace_timestamp = data.get("timestamp") + if cutoff_time and trace_timestamp: + try: + ts = datetime.fromisoformat(trace_timestamp.replace('Z', '+00:00')) + if ts < cutoff_time: + continue + except (ValueError, TypeError): + pass + + # Filter by status + trace_status = data.get("status", "") + if status and trace_status != status: + continue + + # Filter by search text + request_preview = data.get("request", {}).get("input_preview", "") + if search and search.lower() not in request_preview.lower(): + continue + + traces.append({ + "trace_id": data.get("trace_id"), + "timestamp": trace_timestamp, + "user": data.get("user"), + "status": trace_status, + "total_duration_ms": data.get("total_duration_ms"), + "span_count": len(data.get("spans", [])), + "request_preview": request_preview[:100], + }) + except Exception as e: + logger.warning("trace_list_parse_error", path=str(path), error=str(e)) + + return {"traces": traces, "total": len(traces)} + + +@router.get("/{trace_id}") +async def get_trace(trace_id: str): + """ + Get a specific trace by ID. + + Returns the full trace JSON. + """ + if not tracing_enabled(): + raise HTTPException(status_code=404, detail="Tracing not enabled") + + # Sanitize trace_id to prevent path traversal + if not trace_id.startswith("trace_") or "/" in trace_id or "\\" in trace_id: + raise HTTPException(status_code=400, detail="Invalid trace ID") + + trace_path = TRACES_DIR / f"{trace_id}.json" + + if not trace_path.exists(): + raise HTTPException(status_code=404, detail="Trace not found") + + try: + import json + with open(trace_path) as f: + data = json.load(f) + return JSONResponse(content=data) + except Exception as e: + logger.error("trace_read_error", trace_id=trace_id, error=str(e)) + raise HTTPException(status_code=500, detail="Failed to read trace")