diff --git a/abevalflow/llm_client.py b/abevalflow/llm_client.py index c6339dc..f01d7a1 100644 --- a/abevalflow/llm_client.py +++ b/abevalflow/llm_client.py @@ -13,11 +13,24 @@ import logging import os +from dataclasses import dataclass from typing import TYPE_CHECKING if TYPE_CHECKING: from openai import OpenAI + +@dataclass +class LLMResult: + """Chat completion result with token usage.""" + + content: str + prompt_tokens: int + completion_tokens: int + total_tokens: int + model: str + + logger = logging.getLogger(__name__) DEFAULT_BASE_URL = "http://litellm.ab-eval-flow.svc:4000/v1" @@ -45,15 +58,15 @@ def get_model() -> str: return os.environ.get("LLM_MODEL", DEFAULT_MODEL) -def chat_completion( +def chat_completion_with_usage( messages: list[dict[str, str]], *, model: str | None = None, temperature: float = 0.3, max_tokens: int = 4096, **kwargs, -) -> str: - """Send a chat completion request and return the assistant message content. +) -> LLMResult: + """Send a chat completion request and return content with token usage. Raises on API errors so callers can handle retries at a higher level. """ @@ -75,6 +88,38 @@ def chat_completion( **kwargs, ) - content = response.choices[0].message.content - logger.info("chat_completion ← %d chars", len(content) if content else 0) - return content or "" + content = response.choices[0].message.content or "" + usage = response.usage + prompt_tokens = usage.prompt_tokens if usage else 0 + completion_tokens = usage.completion_tokens if usage else 0 + + logger.info( + "chat_completion ← %d chars, tokens: %d prompt + %d completion", + len(content), + prompt_tokens, + completion_tokens, + ) + + return LLMResult( + content=content, + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=prompt_tokens + completion_tokens, + model=resolved_model, + ) + + +def chat_completion( + messages: list[dict[str, str]], + *, + model: str | None = None, + temperature: float = 0.3, + max_tokens: int = 4096, + **kwargs, +) -> str: + """Send a chat completion request and return the assistant message content. + + For token usage, use ``chat_completion_with_usage()`` instead. + """ + result = chat_completion_with_usage(messages, model=model, temperature=temperature, max_tokens=max_tokens, **kwargs) + return result.content diff --git a/abevalflow/observability/__init__.py b/abevalflow/observability/__init__.py new file mode 100644 index 0000000..fc63c44 --- /dev/null +++ b/abevalflow/observability/__init__.py @@ -0,0 +1,9 @@ +"""Observability layer for ABEvalFlow pipeline metrics and tracing.""" + +from abevalflow.observability.context import MetricsContext, TimingRecord, TokenUsage + +__all__ = [ + "MetricsContext", + "TimingRecord", + "TokenUsage", +] diff --git a/abevalflow/observability/context.py b/abevalflow/observability/context.py new file mode 100644 index 0000000..0061686 --- /dev/null +++ b/abevalflow/observability/context.py @@ -0,0 +1,136 @@ +"""Metrics context for accumulating token usage and timing across a pipeline run. + +Append-only during execution, then serialized and persisted at the end. +Per-gate token buckets keyed by phase/gate name for future parallelization. +""" + +from __future__ import annotations + +import logging +import time +from datetime import UTC, datetime +from pathlib import Path + +from pydantic import BaseModel, Field + +logger = logging.getLogger(__name__) + +CHECKPOINT_FILENAME = "_metrics_checkpoint.json" + + +class TokenUsage(BaseModel): + """Token counts for a single phase or gate.""" + + prompt_tokens: int = 0 + completion_tokens: int = 0 + total_tokens: int = 0 + model_name: str | None = None + call_count: int = 0 + + def accumulate(self, prompt: int, completion: int, model: str | None = None) -> None: + self.prompt_tokens += prompt + self.completion_tokens += completion + self.total_tokens += prompt + completion + self.call_count += 1 + if model: + self.model_name = model + + +class TimingRecord(BaseModel): + """Duration record for a pipeline phase.""" + + name: str + start_time: float = Field(description="Unix timestamp") + end_time: float | None = None + duration_ms: int | None = None + + def stop(self) -> None: + self.end_time = time.time() + self.duration_ms = int((self.end_time - self.start_time) * 1000) + + +class MetricsContext(BaseModel): + """Accumulates token usage and timing across a pipeline run. + + Token usage is bucketed by phase/gate name to support concurrent execution. + """ + + run_id: str = "" + submission_name: str = "" + model_name: str | None = None + start_time: datetime = Field(default_factory=lambda: datetime.now(UTC)) + timings: dict[str, TimingRecord] = Field(default_factory=dict) + token_usage: dict[str, TokenUsage] = Field(default_factory=dict) + + def record_tokens( + self, + phase_name: str, + prompt_tokens: int, + completion_tokens: int, + model: str | None = None, + ) -> None: + if phase_name not in self.token_usage: + self.token_usage[phase_name] = TokenUsage() + self.token_usage[phase_name].accumulate(prompt_tokens, completion_tokens, model) + if model and not self.model_name: + self.model_name = model + + def start_timing(self, phase_name: str) -> None: + self.timings[phase_name] = TimingRecord(name=phase_name, start_time=time.time()) + + def stop_timing(self, phase_name: str) -> None: + if phase_name in self.timings: + self.timings[phase_name].stop() + + @property + def total_prompt_tokens(self) -> int: + return sum(u.prompt_tokens for u in self.token_usage.values()) + + @property + def total_completion_tokens(self) -> int: + return sum(u.completion_tokens for u in self.token_usage.values()) + + @property + def total_tokens(self) -> int: + return sum(u.total_tokens for u in self.token_usage.values()) + + @property + def llm_calls_count(self) -> int: + return sum(u.call_count for u in self.token_usage.values()) + + def timing_ms(self, phase_name: str) -> int | None: + rec = self.timings.get(phase_name) + return rec.duration_ms if rec else None + + def checkpoint(self, workspace_path: Path) -> None: + path = workspace_path / CHECKPOINT_FILENAME + path.write_text(self.model_dump_json(indent=2)) + logger.info("Metrics checkpoint written to %s", path) + + @classmethod + def load_checkpoint(cls, workspace_path: Path) -> MetricsContext | None: + path = workspace_path / CHECKPOINT_FILENAME + if not path.exists(): + return None + try: + return cls.model_validate_json(path.read_bytes()) + except Exception: + logger.warning("Failed to load metrics checkpoint from %s", path, exc_info=True) + return None + + def to_observability_dict(self) -> dict: + """Convert to kwargs for ObservabilityMetricsRow.""" + return { + "submission_name": self.submission_name, + "model_name": self.model_name, + "pipeline_duration_ms": self.timing_ms("pipeline"), + "prepare_duration_ms": self.timing_ms("prepare"), + "test_duration_ms": self.timing_ms("test"), + "evaluate_duration_ms": self.timing_ms("evaluate"), + "analyze_duration_ms": self.timing_ms("analyze"), + "store_duration_ms": self.timing_ms("store"), + "total_prompt_tokens": self.total_prompt_tokens or None, + "total_completion_tokens": self.total_completion_tokens or None, + "total_tokens": self.total_tokens or None, + "llm_calls_count": self.llm_calls_count or None, + } diff --git a/abevalflow/observability/decorators.py b/abevalflow/observability/decorators.py new file mode 100644 index 0000000..286422e --- /dev/null +++ b/abevalflow/observability/decorators.py @@ -0,0 +1,31 @@ +"""Observability decorators for timing and tracing.""" + +from __future__ import annotations + +import functools +import logging +import time +from collections.abc import Callable +from typing import Any + +logger = logging.getLogger(__name__) + + +def timed_gate(func: Callable[..., Any]) -> Callable[..., Any]: + """Decorator that logs gate execution time in milliseconds. + + Attaches ``_duration_ms`` to the return value if it has that attribute, + otherwise logs only. + """ + + @functools.wraps(func) + def wrapper(*args: Any, **kwargs: Any) -> Any: + start = time.time() + result = func(*args, **kwargs) + duration_ms = int((time.time() - start) * 1000) + logger.info("Gate %s executed in %dms", func.__qualname__, duration_ms) + if hasattr(result, "_duration_ms"): + result._duration_ms = duration_ms + return result + + return wrapper diff --git a/pipeline/tasks/phases/test.yaml b/pipeline/tasks/phases/test.yaml index 38e65ce..ae06fb6 100644 --- a/pipeline/tasks/phases/test.yaml +++ b/pipeline/tasks/phases/test.yaml @@ -460,17 +460,17 @@ spec: echo "Quality review: passed=$PASSED recommendation=$REC" fi - # Step 6: Finalize results + # Step 6: Finalize results and write metrics checkpoint - name: finalize image: registry.access.redhat.com/ubi9/python-311:9.6 script: | #!/usr/bin/env bash set -euo pipefail echo "=== TEST PHASE: Finalize ===" - + SECURITY_PASSED=$(cat "$(results.security-passed.path)") QUALITY_PASSED=$(cat "$(results.quality-passed.path)") - + # Overall pass if both passed if [ "$SECURITY_PASSED" = "true" ] && [ "$QUALITY_PASSED" = "true" ]; then echo -n "true" > "$(results.tests-passed.path)" @@ -479,3 +479,19 @@ spec: echo -n "false" > "$(results.tests-passed.path)" echo "Tests FAILED (security=$SECURITY_PASSED, quality=$QUALITY_PASSED)" fi + + # Write metrics checkpoint with token data from quality review + PIPELINE_DIR="$(workspaces.source.path)/_pipeline" + REPORT_DIR="$(workspaces.source.path)/reports/$(params.submission-name)" + REVIEW_FILE="$(workspaces.source.path)/_ai_review.json" + mkdir -p "$REPORT_DIR" + + pip install --quiet --no-cache-dir pydantic 2>&1 | tail -1 + export PYTHONPATH="$PIPELINE_DIR" + + python3 "$PIPELINE_DIR/scripts/write_metrics_checkpoint.py" \ + --run-id "$(params.pipeline-run-name)" \ + --submission-name "$(params.submission-name)" \ + --report-dir "$REPORT_DIR" \ + --review-file "$REVIEW_FILE" \ + 2>&1 || echo "Warning: metrics checkpoint failed (non-blocking)" diff --git a/scripts/store_results.py b/scripts/store_results.py index 2430242..0ce5656 100644 --- a/scripts/store_results.py +++ b/scripts/store_results.py @@ -32,12 +32,14 @@ GateResultRow, MCPCheckerRun, MCPCheckerTask, + ObservabilityMetricsRow, ScorecardRow, SecurityScan, Trial, ) from abevalflow.db.observer import discover_observers, notify_observers from abevalflow.mcpchecker_report import MCPCheckerResult +from abevalflow.observability.context import MetricsContext from abevalflow.report import AnalysisResult from abevalflow.scorecard import Scorecard @@ -347,6 +349,8 @@ def store( ).scalar_one_or_none() if existing is not None: + # Note: early return skips scorecard/metrics persistence for + # pre-existing runs. Backfill via scripts/backfill_scorecards.py. logger.warning( "Run %s already exists (id=%s) — skipping", effective_run_id, @@ -392,6 +396,20 @@ def store( except Exception: logger.warning("Failed to persist scorecard — continuing without", exc_info=True) + metrics_ctx = MetricsContext.load_checkpoint(report_dir) + if metrics_ctx and metrics_ctx.total_tokens > 0: + try: + with session.begin_nested(): + session.add( + ObservabilityMetricsRow( + pipeline_run_id=effective_run_id, + **metrics_ctx.to_observability_dict(), + ) + ) + logger.info("Observability metrics queued: tokens=%s", metrics_ctx.total_tokens) + except Exception: + logger.warning("Failed to persist observability metrics — continuing without", exc_info=True) + try: session.commit() except IntegrityError as e: diff --git a/scripts/test_quality_review.py b/scripts/test_quality_review.py index 9c8b423..b7607e2 100644 --- a/scripts/test_quality_review.py +++ b/scripts/test_quality_review.py @@ -390,7 +390,7 @@ def review_submission(submission_dir: Path) -> dict: ) engine = None - response_text = llm_client.chat_completion( + llm_result = llm_client.chat_completion_with_usage( messages=[ {"role": "system", "content": system_prompt}, {"role": "user", "content": user_prompt}, @@ -398,6 +398,7 @@ def review_submission(submission_dir: Path) -> dict: temperature=0.2, max_tokens=4096, ) + response_text = llm_result.content try: assessment = json.loads(response_text) @@ -410,7 +411,15 @@ def review_submission(submission_dir: Path) -> dict: else: raise ValueError("LLM review response is not valid JSON") - return _normalize_assessment(assessment, engine=engine) + assessment = _normalize_assessment(assessment, engine=engine) + assessment["token_usage"] = { + "prompt_tokens": llm_result.prompt_tokens, + "completion_tokens": llm_result.completion_tokens, + "total_tokens": llm_result.total_tokens, + "model": llm_result.model, + } + + return assessment def main(argv: list[str] | None = None) -> int: diff --git a/scripts/write_metrics_checkpoint.py b/scripts/write_metrics_checkpoint.py new file mode 100644 index 0000000..6b4f951 --- /dev/null +++ b/scripts/write_metrics_checkpoint.py @@ -0,0 +1,70 @@ +"""Write a metrics checkpoint from quality review token data. + +Reads _ai_review.json for token usage and writes a MetricsContext +checkpoint to the report directory. + +Usage:: + + python scripts/write_metrics_checkpoint.py \\ + --run-id \\ + --submission-name \\ + --report-dir /workspace/source/reports/ \\ + --review-file /workspace/source/_ai_review.json +""" + +from __future__ import annotations + +import argparse +import json +import logging +from pathlib import Path + +from abevalflow.observability.context import MetricsContext + +logger = logging.getLogger(__name__) + + +def main() -> None: + logging.basicConfig(level=logging.INFO, format="%(message)s") + + parser = argparse.ArgumentParser(description="Write metrics checkpoint") + parser.add_argument("--run-id", required=True) + parser.add_argument("--submission-name", required=True) + parser.add_argument("--report-dir", type=Path, required=True) + parser.add_argument("--review-file", type=Path, default=None) + args = parser.parse_args() + + ctx = MetricsContext( + run_id=args.run_id, + submission_name=args.submission_name, + ) + + if args.review_file and args.review_file.exists(): + try: + review = json.loads(args.review_file.read_text()) + usage = review.get("token_usage", {}) + if usage: + ctx.record_tokens( + "quality_review", + usage.get("prompt_tokens", 0), + usage.get("completion_tokens", 0), + usage.get("model"), + ) + logger.info("Recorded quality review tokens: %s", usage) + else: + logger.info("No token_usage in review file") + except Exception as e: + logger.warning("Could not read review tokens: %s", e) + else: + logger.info("No review file found, skipping token capture") + + if ctx.total_tokens > 0: + args.report_dir.mkdir(parents=True, exist_ok=True) + ctx.checkpoint(args.report_dir) + logger.info("Metrics checkpoint written to %s", args.report_dir) + else: + logger.info("No token data collected, skipping checkpoint") + + +if __name__ == "__main__": + main() diff --git a/tests/test_llm_client.py b/tests/test_llm_client.py index d37a011..882ee25 100644 --- a/tests/test_llm_client.py +++ b/tests/test_llm_client.py @@ -7,6 +7,7 @@ import pytest from abevalflow import llm_client +from abevalflow.llm_client import LLMResult class TestResolveConfig: @@ -104,3 +105,67 @@ def test_passes_temperature_and_max_tokens(self, mock_get_client: MagicMock) -> call_kwargs = mock_client.chat.completions.create.call_args assert call_kwargs.kwargs["temperature"] == 0.7 assert call_kwargs.kwargs["max_tokens"] == 2048 + + +class TestChatCompletionWithUsage: + @patch("abevalflow.llm_client.get_client") + def test_returns_llm_result(self, mock_get_client: MagicMock) -> None: + mock_client = MagicMock() + mock_response = MagicMock() + mock_response.choices = [MagicMock()] + mock_response.choices[0].message.content = "Hello!" + mock_response.usage.prompt_tokens = 100 + mock_response.usage.completion_tokens = 50 + mock_client.chat.completions.create.return_value = mock_response + mock_get_client.return_value = mock_client + + result = llm_client.chat_completion_with_usage( + [{"role": "user", "content": "hi"}], + model="test-model", + ) + + assert isinstance(result, LLMResult) + assert result.content == "Hello!" + assert result.prompt_tokens == 100 + assert result.completion_tokens == 50 + assert result.total_tokens == 150 + assert result.model == "test-model" + + @patch("abevalflow.llm_client.get_client") + def test_handles_no_usage(self, mock_get_client: MagicMock) -> None: + mock_client = MagicMock() + mock_response = MagicMock() + mock_response.choices = [MagicMock()] + mock_response.choices[0].message.content = "Hello!" + mock_response.usage = None + mock_client.chat.completions.create.return_value = mock_response + mock_get_client.return_value = mock_client + + result = llm_client.chat_completion_with_usage( + [{"role": "user", "content": "hi"}], + model="test-model", + ) + + assert result.content == "Hello!" + assert result.prompt_tokens == 0 + assert result.completion_tokens == 0 + assert result.total_tokens == 0 + + @patch("abevalflow.llm_client.get_client") + def test_backward_compatible_chat_completion(self, mock_get_client: MagicMock) -> None: + mock_client = MagicMock() + mock_response = MagicMock() + mock_response.choices = [MagicMock()] + mock_response.choices[0].message.content = "Hello!" + mock_response.usage.prompt_tokens = 100 + mock_response.usage.completion_tokens = 50 + mock_client.chat.completions.create.return_value = mock_response + mock_get_client.return_value = mock_client + + result = llm_client.chat_completion( + [{"role": "user", "content": "hi"}], + model="test-model", + ) + + assert isinstance(result, str) + assert result == "Hello!" diff --git a/tests/test_metrics_context.py b/tests/test_metrics_context.py new file mode 100644 index 0000000..990ab8a --- /dev/null +++ b/tests/test_metrics_context.py @@ -0,0 +1,142 @@ +"""Tests for abevalflow.observability.context — MetricsContext, TokenUsage, TimingRecord.""" + +from __future__ import annotations + +import json +from pathlib import Path + +from abevalflow.observability.context import CHECKPOINT_FILENAME, MetricsContext, TimingRecord, TokenUsage + + +class TestTokenUsage: + def test_defaults(self) -> None: + t = TokenUsage() + assert t.prompt_tokens == 0 + assert t.completion_tokens == 0 + assert t.total_tokens == 0 + assert t.call_count == 0 + + def test_accumulate(self) -> None: + t = TokenUsage() + t.accumulate(100, 50, "claude-sonnet") + assert t.prompt_tokens == 100 + assert t.completion_tokens == 50 + assert t.total_tokens == 150 + assert t.call_count == 1 + assert t.model_name == "claude-sonnet" + + def test_accumulate_multiple(self) -> None: + t = TokenUsage() + t.accumulate(100, 50, "claude-sonnet") + t.accumulate(200, 80, "claude-sonnet") + assert t.prompt_tokens == 300 + assert t.completion_tokens == 130 + assert t.total_tokens == 430 + assert t.call_count == 2 + + +class TestTimingRecord: + def test_stop_computes_duration(self) -> None: + rec = TimingRecord(name="test", start_time=1000.0) + rec.end_time = 1002.5 + rec.duration_ms = int((rec.end_time - rec.start_time) * 1000) + assert rec.duration_ms == 2500 + + def test_stop_method(self) -> None: + import time + + rec = TimingRecord(name="test", start_time=time.time()) + rec.stop() + assert rec.end_time is not None + assert rec.duration_ms is not None + assert rec.duration_ms >= 0 + + +class TestMetricsContext: + def test_record_tokens(self) -> None: + ctx = MetricsContext(run_id="run-1", submission_name="test-skill") + ctx.record_tokens("quality_review", 500, 200, "claude-sonnet") + assert ctx.total_prompt_tokens == 500 + assert ctx.total_completion_tokens == 200 + assert ctx.total_tokens == 700 + assert ctx.llm_calls_count == 1 + assert ctx.model_name == "claude-sonnet" + + def test_record_tokens_multiple_phases(self) -> None: + ctx = MetricsContext(run_id="run-1", submission_name="test-skill") + ctx.record_tokens("quality_review", 500, 200, "claude-sonnet") + ctx.record_tokens("security_scan", 300, 100, "claude-sonnet") + assert ctx.total_prompt_tokens == 800 + assert ctx.total_completion_tokens == 300 + assert ctx.total_tokens == 1100 + assert ctx.llm_calls_count == 2 + + def test_record_tokens_same_phase(self) -> None: + ctx = MetricsContext(run_id="run-1", submission_name="test-skill") + ctx.record_tokens("quality_review", 500, 200) + ctx.record_tokens("quality_review", 300, 150) + assert ctx.token_usage["quality_review"].prompt_tokens == 800 + assert ctx.token_usage["quality_review"].call_count == 2 + assert ctx.total_tokens == 1150 + + def test_timing(self) -> None: + ctx = MetricsContext() + ctx.timings["test"] = TimingRecord(name="test", start_time=1000.0, end_time=1005.0, duration_ms=5000) + assert ctx.timing_ms("test") == 5000 + assert ctx.timing_ms("missing") is None + + def test_checkpoint_round_trip(self, tmp_path: Path) -> None: + ctx = MetricsContext(run_id="run-1", submission_name="test-skill", model_name="claude-sonnet") + ctx.record_tokens("quality_review", 500, 200, "claude-sonnet") + ctx.timings["test"] = TimingRecord(name="test", start_time=1000.0, end_time=1005.0, duration_ms=5000) + ctx.checkpoint(tmp_path) + + loaded = MetricsContext.load_checkpoint(tmp_path) + assert loaded is not None + assert loaded.run_id == "run-1" + assert loaded.submission_name == "test-skill" + assert loaded.total_tokens == 700 + assert loaded.timing_ms("test") == 5000 + + def test_checkpoint_file_written(self, tmp_path: Path) -> None: + ctx = MetricsContext(run_id="run-1") + ctx.checkpoint(tmp_path) + path = tmp_path / CHECKPOINT_FILENAME + assert path.exists() + data = json.loads(path.read_text()) + assert data["run_id"] == "run-1" + + def test_load_checkpoint_missing(self, tmp_path: Path) -> None: + result = MetricsContext.load_checkpoint(tmp_path) + assert result is None + + def test_load_checkpoint_invalid(self, tmp_path: Path) -> None: + path = tmp_path / CHECKPOINT_FILENAME + path.write_text("not valid json{{{") + result = MetricsContext.load_checkpoint(tmp_path) + assert result is None + + def test_to_observability_dict(self) -> None: + ctx = MetricsContext(run_id="run-1", submission_name="test-skill", model_name="claude-sonnet") + ctx.record_tokens("quality_review", 500, 200, "claude-sonnet") + ctx.timings["pipeline"] = TimingRecord(name="pipeline", start_time=0, end_time=10, duration_ms=10000) + ctx.timings["test"] = TimingRecord(name="test", start_time=0, end_time=5, duration_ms=5000) + + d = ctx.to_observability_dict() + assert d["submission_name"] == "test-skill" + assert d["model_name"] == "claude-sonnet" + assert d["total_prompt_tokens"] == 500 + assert d["total_completion_tokens"] == 200 + assert d["total_tokens"] == 700 + assert d["pipeline_duration_ms"] == 10000 + assert d["test_duration_ms"] == 5000 + assert d["evaluate_duration_ms"] is None + assert d["llm_calls_count"] == 1 + + def test_empty_context_observability_dict(self) -> None: + ctx = MetricsContext() + d = ctx.to_observability_dict() + assert d["total_prompt_tokens"] is None + assert d["total_completion_tokens"] is None + assert d["total_tokens"] is None + assert d["llm_calls_count"] is None diff --git a/tests/test_quality_review_aeh.py b/tests/test_quality_review_aeh.py index 93f7d94..7bf5684 100644 --- a/tests/test_quality_review_aeh.py +++ b/tests/test_quality_review_aeh.py @@ -7,6 +7,7 @@ import yaml +from abevalflow.llm_client import LLMResult from abevalflow.schemas import SubmissionMetadata from scripts.test_quality_review import ( _advisory_aeh_missing_files, @@ -94,8 +95,14 @@ def test_aeh_review_uses_aeh_prompt_and_stays_passed(self, tmp_path: Path): } payload = __import__("json").dumps(fake) with patch( - "scripts.test_quality_review.llm_client.chat_completion", - return_value=payload, + "scripts.test_quality_review.llm_client.chat_completion_with_usage", + return_value=LLMResult( + content=payload, + prompt_tokens=1, + completion_tokens=1, + total_tokens=2, + model="test", + ), ): assessment = review_submission(sub) assert assessment["engine"] == "aeh" diff --git a/tests/test_store_results_observability.py b/tests/test_store_results_observability.py index ec9af03..adda186 100644 --- a/tests/test_store_results_observability.py +++ b/tests/test_store_results_observability.py @@ -12,6 +12,7 @@ Base, EvaluationRun, GateResultRow, + ObservabilityMetricsRow, ScorecardRow, ) from abevalflow.gates.base import GateMode, GateResult, GateType @@ -260,3 +261,43 @@ def test_scorecard_unique_constraint(self, db_url: str, session_factory) -> None ) with pytest.raises(IntegrityError): session.commit() + + def test_store_with_metrics_checkpoint(self, tmp_path: Path, db_url: str, session_factory) -> None: + from abevalflow.observability.context import MetricsContext, TimingRecord + + result = _sample_result() + report_dir = _write_report(tmp_path, result) + scorecard = _sample_scorecard() + _write_scorecard(report_dir, scorecard) + + ctx = MetricsContext( + run_id="tekton-run-metrics", + submission_name="my-submission", + model_name="claude-sonnet", + ) + ctx.record_tokens("quality_review", 500, 200, "claude-sonnet") + ctx.timings["pipeline"] = TimingRecord(name="pipeline", start_time=0, end_time=10, duration_ms=10000) + ctx.checkpoint(report_dir) + + ok = store(report_dir, db_url, "tekton-run-metrics") + assert ok is True + + with session_factory() as session: + metrics = session.execute(select(ObservabilityMetricsRow)).scalar_one() + assert metrics.submission_name == "my-submission" + assert metrics.model_name == "claude-sonnet" + assert metrics.total_prompt_tokens == 500 + assert metrics.total_completion_tokens == 200 + assert metrics.total_tokens == 700 + assert metrics.pipeline_duration_ms == 10000 + assert metrics.llm_calls_count == 1 + + def test_store_without_metrics_checkpoint(self, tmp_path: Path, db_url: str, session_factory) -> None: + result = _sample_result() + report_dir = _write_report(tmp_path, result) + + ok = store(report_dir, db_url, "tekton-run-no-metrics") + assert ok is True + + with session_factory() as session: + assert session.execute(select(ObservabilityMetricsRow)).scalar_one_or_none() is None