job counter

This commit is contained in:
酱牛肉 2026-07-30 12:33:40 +08:00
parent 4e06d59aca
commit e14ff16cac
4 changed files with 90 additions and 17 deletions

View file

@ -252,17 +252,16 @@ async def answer_question_agentic(app, question: str) -> tuple[str, dict]:
Returns (answer, metadata)
"""
from reme.utils.evaluation_interface import check_job_count
from reme.utils.evaluation_interface import track_job_counts
search_calls_before = check_job_count("search", app.context)
query_resp = await app.run_job(
"agentic_answer",
query=question,
)
with track_job_counts(["search"], app.context) as counts:
query_resp = await app.run_job(
"agentic_answer",
query=question,
)
answer = (query_resp.answer or "").strip()
search_calls = check_job_count("search", app.context) - search_calls_before
return answer, {"mode": "agentic", "search_calls": search_calls}
return answer, {"mode": "agentic", "search_calls": counts["search"]}
# ---------------------------------------------------------------------------

View file

@ -259,7 +259,7 @@ async def evaluate_item(item: dict, eval_config: dict, item_index: int, eval_onl
"""
from reme import Application
from reme.config import resolve_app_config
from reme.utils.evaluation_interface import check_job_count
from reme.utils.evaluation_interface import track_job_counts
reme_cfg = eval_config["reme"]
dream_trigger_hour = reme_cfg.get("dream_trigger_hour", 23)
@ -433,13 +433,13 @@ async def evaluate_item(item: dict, eval_config: dict, item_index: int, eval_onl
f"[Item {item_index}] Asking (agentic): {question[:80]}... query_time={query_time}",
)
search_calls_before = check_job_count("search", app.context)
query_resp = await app.run_job(
"agentic_answer",
query=question,
query_time=query_time,
)
agentic_search_calls = check_job_count("search", app.context) - search_calls_before
with track_job_counts(["search"], app.context) as counts:
query_resp = await app.run_job(
"agentic_answer",
query=question,
query_time=query_time,
)
agentic_search_calls = counts["search"]
agentic_response = (query_resp.answer or "").strip()
if not agentic_response:
agentic_response = "(no answer generated)"

View file

@ -28,3 +28,40 @@ def check_job_count(job_name: str, app_context: "ApplicationContext") -> int:
raise TypeError(f"Job '{job_name}' does not inherit from a supported job implementation")
# pylint: disable-next=protected-access
return global_counter_get(app_context.metadata, job._counter_key(entry_class))
class JobCountTracker:
"""Measure registered job calls made while this context is active."""
def __init__(self, job_names: list[str], app_context: "ApplicationContext") -> None:
self.job_names = list(dict.fromkeys(job_names))
self.app_context = app_context
self._start_counts: dict[str, int] = {}
self.counts: dict[str, int] = {}
def __enter__(self) -> dict[str, int]:
self._start_counts = {name: check_job_count(name, self.app_context) for name in self.job_names}
return self.counts
def __exit__(self, exc_type, exc_value, traceback) -> bool:
self.counts.update(
{
name: check_job_count(name, self.app_context) - start_count
for name, start_count in self._start_counts.items()
},
)
return False
def track_job_counts(job_names: list[str], app_context: "ApplicationContext") -> JobCountTracker:
"""Return a context manager that reports call deltas for ``job_names``.
Example:
.. code-block:: python
with track_job_counts(["search"], app.context) as counts:
await app.run_job("agentic_answer", query="...")
assert counts == {"search": 2}
"""
return JobCountTracker(job_names, app_context)

View file

@ -6,7 +6,7 @@ from types import SimpleNamespace
import pytest
from reme.components.job import BaseJob, StreamJob
from reme.utils.evaluation_interface import check_job_count
from reme.utils.evaluation_interface import check_job_count, track_job_counts
def test_check_job_count_reads_registered_base_job_count():
@ -49,3 +49,40 @@ def test_check_job_count_rejects_unknown_job_name():
with pytest.raises(KeyError, match="Job 'missing' not found"):
check_job_count("missing", app_context)
def test_track_job_counts_returns_calls_made_inside_context():
"""The context manager reports only the calls made in its body."""
async def run():
app_context = SimpleNamespace(metadata={}, jobs={})
search = BaseJob(name="search", app_context=app_context)
app_context.jobs[search.name] = search
await search()
with track_job_counts(["search"], app_context) as counts:
await search()
await search()
assert counts == {"search": 2}
asyncio.run(run())
def test_track_job_counts_updates_results_when_body_raises():
"""Calls made before an exception are still included in the delta."""
async def run():
app_context = SimpleNamespace(metadata={}, jobs={})
search = BaseJob(name="search", app_context=app_context)
app_context.jobs[search.name] = search
counts = {}
with pytest.raises(RuntimeError, match="boom"):
with track_job_counts(["search"], app_context) as counts:
await search()
raise RuntimeError("boom")
assert counts == {"search": 1}
asyncio.run(run())