diff --git a/rag_eval/reporting/writers.py b/rag_eval/reporting/writers.py index 8793984..6f4fd89 100644 --- a/rag_eval/reporting/writers.py +++ b/rag_eval/reporting/writers.py @@ -45,6 +45,7 @@ def write_run_artifacts(result: EvaluationResult) -> None: "dataset": result.scenario.dataset.path.as_posix(), "valid_samples": len(result.valid_samples), "invalid_samples": len(result.invalid_samples), + "token_usage": result.token_usage, } artifact_paths.metadata_json.write_text( json.dumps(metadata, ensure_ascii=False, indent=2), diff --git a/rag_eval/shared/models.py b/rag_eval/shared/models.py index 677c2c8..25e259f 100644 --- a/rag_eval/shared/models.py +++ b/rag_eval/shared/models.py @@ -153,6 +153,9 @@ class EvaluationResult: valid_samples: list[NormalizedSample] invalid_samples: list[InvalidSample] score_rows: list[dict[str, Any]] + # Token usage grouped by model name: {model: {input_tokens, output_tokens, calls}}. + # Populated by callers via rag_eval.metrics.token_tracker.track_token_usage(). + token_usage: dict[str, dict[str, int]] = field(default_factory=dict) @dataclass(slots=True) diff --git a/tests/test_token_usage_persistence.py b/tests/test_token_usage_persistence.py new file mode 100644 index 0000000..bc04f0a --- /dev/null +++ b/tests/test_token_usage_persistence.py @@ -0,0 +1,75 @@ +"""Tests that EvaluationResult.token_usage is persisted into metadata.json.""" +from __future__ import annotations + +import json +from pathlib import Path + +from rag_eval.reporting.writers import write_run_artifacts +from rag_eval.shared.models import DatasetConfig, EvaluationResult, RuntimeConfig, Scenario + + +def _scenario(tmp_path: Path) -> Scenario: + return Scenario( + scenario_name="token-persist-test", + mode="offline", + dataset=DatasetConfig(path=tmp_path / "dataset.csv"), + judge_model="gpt-5", + embedding_model="embedding-model", + metrics=["faithfulness"], + output_dir=tmp_path / "outputs", + runtime=RuntimeConfig(batch_size=1), + ) + + +def test_evaluation_result_defaults_token_usage_to_empty_dict(tmp_path: Path) -> None: + result = EvaluationResult( + scenario=_scenario(tmp_path), + run_id="run-1", + started_at="t0", + finished_at="t1", + valid_samples=[], + invalid_samples=[], + score_rows=[], + ) + assert result.token_usage == {} + + +def test_write_run_artifacts_persists_token_usage(tmp_path: Path) -> None: + scenario = _scenario(tmp_path) + result = EvaluationResult( + scenario=scenario, + run_id="run-2", + started_at="t0", + finished_at="t1", + valid_samples=[], + invalid_samples=[], + score_rows=[], + token_usage={"gpt-5": {"input_tokens": 100, "output_tokens": 40, "calls": 2}}, + ) + + write_run_artifacts(result) + + metadata_path = scenario.output_dir / "run-2" / "metadata.json" + metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + assert metadata["token_usage"] == { + "gpt-5": {"input_tokens": 100, "output_tokens": 40, "calls": 2} + } + + +def test_write_run_artifacts_writes_empty_token_usage_when_unset(tmp_path: Path) -> None: + scenario = _scenario(tmp_path) + result = EvaluationResult( + scenario=scenario, + run_id="run-3", + started_at="t0", + finished_at="t1", + valid_samples=[], + invalid_samples=[], + score_rows=[], + ) + + write_run_artifacts(result) + + metadata_path = scenario.output_dir / "run-3" / "metadata.json" + metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + assert metadata["token_usage"] == {}