feat(token-tracking): accumulate token usage across session_async calls
Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
@@ -0,0 +1,84 @@
|
|||||||
|
"""Tests that session-grouped async scoring accumulates token usage across calls."""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
|
||||||
|
from webapp.models import ScoreRequest
|
||||||
|
from webapp.services.session_score_manager import SessionScoreJobManager
|
||||||
|
|
||||||
|
|
||||||
|
def _wait_for_call_count(mgr: SessionScoreJobManager, session_id: str, expected: int, timeout: float = 2.0):
|
||||||
|
deadline = time.monotonic() + timeout
|
||||||
|
while time.monotonic() < deadline:
|
||||||
|
session = mgr.get_session(session_id)
|
||||||
|
if session is not None and session.call_count >= expected:
|
||||||
|
all_done = all(j.status in ("completed", "failed") for j in session.jobs)
|
||||||
|
if all_done:
|
||||||
|
return session
|
||||||
|
time.sleep(0.02)
|
||||||
|
raise TimeoutError(f"session {session_id} did not reach {expected} completed calls in time")
|
||||||
|
|
||||||
|
|
||||||
|
def test_session_accumulates_token_usage_across_calls(tmp_path, monkeypatch):
|
||||||
|
from rag_eval.metrics.token_tracker import get_current_tracker
|
||||||
|
|
||||||
|
mgr = SessionScoreJobManager(
|
||||||
|
output_dir=tmp_path / "score-session",
|
||||||
|
index_dir=tmp_path / "score-session-jobs",
|
||||||
|
max_workers=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
call_usages = iter([(100, 40), (30, 10)])
|
||||||
|
|
||||||
|
def _fake_score(**kwargs):
|
||||||
|
tracker = get_current_tracker()
|
||||||
|
input_tok, output_tok = next(call_usages)
|
||||||
|
if tracker is not None:
|
||||||
|
tracker.record("gpt-5", input_tok, output_tok)
|
||||||
|
return {m: 0.9 for m in kwargs["metrics"]}
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"webapp.services.inline_scorer.inline_scorer.score", _fake_score
|
||||||
|
)
|
||||||
|
|
||||||
|
request = ScoreRequest(question="q?", answer="a.", metrics=["answer_relevancy"])
|
||||||
|
|
||||||
|
_, run_id = mgr.submit("session-token-test", request)
|
||||||
|
_wait_for_call_count(mgr, "session-token-test", 1)
|
||||||
|
mgr.submit("session-token-test", request)
|
||||||
|
_wait_for_call_count(mgr, "session-token-test", 2)
|
||||||
|
|
||||||
|
run_dir = tmp_path / "score-session" / run_id
|
||||||
|
metadata = json.loads((run_dir / "metadata.json").read_text(encoding="utf-8"))
|
||||||
|
assert metadata["token_usage"] == {
|
||||||
|
"gpt-5": {"input_tokens": 130, "output_tokens": 50, "calls": 2}
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def test_session_first_call_writes_token_usage_from_scratch(tmp_path, monkeypatch):
|
||||||
|
from rag_eval.metrics.token_tracker import get_current_tracker
|
||||||
|
|
||||||
|
mgr = SessionScoreJobManager(
|
||||||
|
output_dir=tmp_path / "score-session",
|
||||||
|
index_dir=tmp_path / "score-session-jobs",
|
||||||
|
max_workers=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _fake_score(**kwargs):
|
||||||
|
tracker = get_current_tracker()
|
||||||
|
if tracker is not None:
|
||||||
|
tracker.record("gpt-5", 50, 20)
|
||||||
|
return {m: 0.9 for m in kwargs["metrics"]}
|
||||||
|
|
||||||
|
monkeypatch.setattr(
|
||||||
|
"webapp.services.inline_scorer.inline_scorer.score", _fake_score
|
||||||
|
)
|
||||||
|
|
||||||
|
request = ScoreRequest(question="q?", answer="a.", metrics=["answer_relevancy"])
|
||||||
|
_, run_id = mgr.submit("session-first-call-test", request)
|
||||||
|
_wait_for_call_count(mgr, "session-first-call-test", 1)
|
||||||
|
|
||||||
|
run_dir = tmp_path / "score-session" / run_id
|
||||||
|
metadata = json.loads((run_dir / "metadata.json").read_text(encoding="utf-8"))
|
||||||
|
assert metadata["token_usage"] == {"gpt-5": {"input_tokens": 50, "output_tokens": 20, "calls": 1}}
|
||||||
@@ -193,6 +193,7 @@ class SessionScoreJobManager:
|
|||||||
|
|
||||||
# Lazy imports — keep web server bootable if ragas is not installed.
|
# Lazy imports — keep web server bootable if ragas is not installed.
|
||||||
from rag_eval.advisor import run_advisor
|
from rag_eval.advisor import run_advisor
|
||||||
|
from rag_eval.metrics.token_tracker import track_token_usage
|
||||||
from rag_eval.metrics.weights import compute_weighted_score
|
from rag_eval.metrics.weights import compute_weighted_score
|
||||||
from rag_eval.reporting.writers import write_run_artifacts
|
from rag_eval.reporting.writers import write_run_artifacts
|
||||||
from rag_eval.settings import EvaluationSettings
|
from rag_eval.settings import EvaluationSettings
|
||||||
@@ -215,6 +216,7 @@ class SessionScoreJobManager:
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
# --- Scoring (can run concurrently for the same session) ----------
|
# --- Scoring (can run concurrently for the same session) ----------
|
||||||
|
with track_token_usage() as usage_tracker:
|
||||||
if effective:
|
if effective:
|
||||||
raw_scores = inline_scorer.score(
|
raw_scores = inline_scorer.score(
|
||||||
question=request.question,
|
question=request.question,
|
||||||
@@ -250,6 +252,14 @@ class SessionScoreJobManager:
|
|||||||
run_dir = self._output_dir / run_id
|
run_dir = self._output_dir / run_id
|
||||||
run_dir.mkdir(parents=True, exist_ok=True)
|
run_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
|
# Merge this call's token usage into the session's running total, so
|
||||||
|
# repeated calls accumulate instead of overwriting (mirrors the
|
||||||
|
# scores.csv append-only accumulation below).
|
||||||
|
existing_metadata = self._read_metadata(run_dir)
|
||||||
|
merged_token_usage = usage_tracker.merge_into(
|
||||||
|
existing_metadata.get("token_usage", {})
|
||||||
|
)
|
||||||
|
|
||||||
# Read all existing rows, then append the new one
|
# Read all existing rows, then append the new one
|
||||||
existing_rows = self._read_score_rows(run_dir)
|
existing_rows = self._read_score_rows(run_dir)
|
||||||
call_number = len(existing_rows) + 1
|
call_number = len(existing_rows) + 1
|
||||||
@@ -312,6 +322,7 @@ class SessionScoreJobManager:
|
|||||||
valid_samples=valid_samples,
|
valid_samples=valid_samples,
|
||||||
invalid_samples=[],
|
invalid_samples=[],
|
||||||
score_rows=all_rows,
|
score_rows=all_rows,
|
||||||
|
token_usage=merged_token_usage,
|
||||||
)
|
)
|
||||||
|
|
||||||
write_run_artifacts(result)
|
write_run_artifacts(result)
|
||||||
@@ -376,6 +387,16 @@ class SessionScoreJobManager:
|
|||||||
except (OSError, ValueError):
|
except (OSError, ValueError):
|
||||||
return []
|
return []
|
||||||
|
|
||||||
|
def _read_metadata(self, run_dir: Path) -> dict[str, Any]:
|
||||||
|
"""Read this session's existing metadata.json, returning {} if absent/unreadable."""
|
||||||
|
metadata_path = run_dir / "metadata.json"
|
||||||
|
if not metadata_path.is_file():
|
||||||
|
return {}
|
||||||
|
try:
|
||||||
|
return json.loads(metadata_path.read_text(encoding="utf-8"))
|
||||||
|
except (OSError, ValueError):
|
||||||
|
return {}
|
||||||
|
|
||||||
def _read_metric_means(self, run_dir: Path) -> dict[str, float | None]:
|
def _read_metric_means(self, run_dir: Path) -> dict[str, float | None]:
|
||||||
"""Compute per-metric means from the session's scores.csv."""
|
"""Compute per-metric means from the session's scores.csv."""
|
||||||
scores_path = run_dir / "scores.csv"
|
scores_path = run_dir / "scores.csv"
|
||||||
|
|||||||
Reference in New Issue
Block a user