diff --git a/benchmark/beam/run.py b/benchmark/beam/run.py index 89f2e53f..4b292cbf 100644 --- a/benchmark/beam/run.py +++ b/benchmark/beam/run.py @@ -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"]} # --------------------------------------------------------------------------- diff --git a/benchmark/longmemeval/run.py b/benchmark/longmemeval/run.py index 3870a3d6..75676f55 100644 --- a/benchmark/longmemeval/run.py +++ b/benchmark/longmemeval/run.py @@ -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)" diff --git a/reme/utils/evaluation_interface.py b/reme/utils/evaluation_interface.py index 0b46db4c..e681ed73 100644 --- a/reme/utils/evaluation_interface.py +++ b/reme/utils/evaluation_interface.py @@ -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) diff --git a/tests/unit/test_evaluation_interface.py b/tests/unit/test_evaluation_interface.py index b95b3daa..76a4ae67 100644 --- a/tests/unit/test_evaluation_interface.py +++ b/tests/unit/test_evaluation_interface.py @@ -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())