From 6b035c6553784a47514388ae6e549c893bf5cd5e Mon Sep 17 00:00:00 2001 From: xyf2020 <75460675+xyf2020@users.noreply.github.com> Date: Tue, 4 Aug 2026 11:42:18 +0800 Subject: [PATCH] feat(evaluation): track job calls and agent token usage in benchmarks (#406) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(counter): extend counter tree utils and record job call statistics - replace global_counter_next with fetch-and-add style global_counter_add/inc, plus read-only global_counter_get and global_counter_get_all - record per-job call counts in app_context.metadata via BaseJob._record_call, covering background/cron/stream jobs - update agentic_answer step and utils exports; add unit tests for job counting and counter utils * feat(evaluation): add check_job_count interface and report search calls in benchmarks - Extract _counter_key from BaseJob._record_call for reusable counter lookup - Add reme.utils.evaluation_interface.check_job_count read-only helper - Track and report average search calls per query in beam and longmemeval benchmarks * job counter * token消耗量统计 * benchmark输出完整token消耗统计 * benchmark统计输出改用标准差 - beam/longmemeval 的工具调用与 token 统计由方差改为标准差输出 - 修复 lint: 局部变量遮蔽 importlib.metadata、补充测试 docstring - black 格式化 * fix(evaluation): preserve complete token usage metrics * fix: exclude stream replies from token accounting * Revert "fix: exclude stream replies from token accounting" This reverts commit 85bf32064d1205b11a9f993a9eaca9513f6071df. * Reapply "fix: exclude stream replies from token accounting" This reverts commit 6722c24dc577b116238b757618ffbd59d6c3cb40. * support agent scope 2.0.5 * feat: support injection_config to disable runtime state injection in benchmarks - Add InjectionConfig passthrough in AsAgentWrapper.reply() - Disable inject_runtime_state in BaseAgenticAnswerStep to avoid wall-clock time conflicting with benchmark query_time anchors - Disable inject_runtime_state in beam/lme llm_judge calls * feat: agentscope dual-version compat & benchmark improvements - Add version_tuple utility for semantic version comparison - AsAgentWrapper: version-aware InjectionConfig, max_iters doubling, and token usage collection (reply vs reply_stream) for AS>=2.0.5/<2.0.5 - Default inject_runtime_state=False in wrapper to avoid benchmark time-anchor conflicts; remove per-callsite injection_config overrides - longmemeval run.py: support question_ids filter in dataset config - Fix unused import in test_evaluation_interface; format fixes * chore: remove temporary flip-test benchmark config * revert: pin agentscope to 2.0.4.post1 and drop dual-version compat * fix(evaluation): clarify usage semantics and atomic counters --------- Co-authored-by: sa-buc Co-authored-by: jinli.yl --- benchmark/beam/run.py | 57 ++++- benchmark/longmemeval/run.py | 57 ++++- .../agent_wrapper/as_agent_wrapper.py | 38 ++- .../agent_wrapper/base_agent_wrapper.py | 21 +- .../agent_wrapper/cc_agent_wrapper.py | 30 ++- .../agent_wrapper/codex_agent_wrapper.py | 21 +- reme/components/job/background_job.py | 1 + reme/components/job/base_job.py | 8 + reme/components/job/cron_job.py | 1 + reme/components/job/stream_job.py | 1 + reme/schema/__init__.py | 2 + reme/schema/token_usage.py | 51 ++++ reme/steps/benchmark/base/agentic_answer.py | 6 +- reme/utils/__init__.py | 14 +- reme/utils/counter.py | 137 +++++++++-- reme/utils/evaluation_interface.py | 191 +++++++++++++++ tests/unit/test_cc_agent_wrapper.py | 6 +- tests/unit/test_codex_agent_wrapper.py | 64 +++++ tests/unit/test_evaluation_interface.py | 142 +++++++++++ tests/unit/test_job.py | 68 ++++++ tests/unit/test_token_usage.py | 227 ++++++++++++++++++ tests/unit/test_utils.py | 194 +++++++++++++++ 22 files changed, 1282 insertions(+), 55 deletions(-) create mode 100644 reme/schema/token_usage.py create mode 100644 reme/utils/evaluation_interface.py create mode 100644 tests/unit/test_evaluation_interface.py create mode 100644 tests/unit/test_token_usage.py diff --git a/benchmark/beam/run.py b/benchmark/beam/run.py index 45c74f10..97c620e5 100644 --- a/benchmark/beam/run.py +++ b/benchmark/beam/run.py @@ -252,13 +252,26 @@ async def answer_question_agentic(app, question: str) -> tuple[str, dict]: Returns (answer, metadata) """ - query_resp = await app.run_job( - "agentic_answer", - query=question, - ) + from reme.utils.evaluation_interface import track_agent_token_usage, track_job_counts + + with ( + track_job_counts(["search"], app.context) as tool_counts, + track_agent_token_usage( + ["bench"], + app.context, + ) as token_usages, + ): + query_resp = await app.run_job( + "agentic_answer", + query=question, + ) answer = (query_resp.answer or "").strip() - return answer, {"mode": "agentic"} + return answer, { + "mode": "agentic", + "tool_counts": tool_counts, + "token_usage": token_usages["bench"], + } # --------------------------------------------------------------------------- @@ -448,6 +461,10 @@ async def evaluate_case(eval_config: dict, case_id: str, eval_only: bool = False if not agentic_answer: agentic_answer = "(no answer generated)" logger.info(f"[Case {case_id}] Agentic answer: {agentic_answer[:200]}...") + logger.info( + f"[Case {case_id}] Agentic tool calls: {agentic_meta.get('tool_counts', {})}", + ) + logger.info(f"[Case {case_id}] Bench token usage: {agentic_meta.get('token_usage', {})}") # Judge agentic answer logger.info(f"[Case {case_id}] Judging agentic ({q_type})...") @@ -678,6 +695,8 @@ def main( # pylint: disable=too-many-statements type_binary_scores: dict[str, list[float]] = {} all_scores: list[float] = [] all_binary_scores: list[float] = [] + all_tool_call_totals: list[int] = [] + all_token_usages: list[dict[str, int | None]] = [] for case_result in results: if "error" in case_result: @@ -700,6 +719,9 @@ def main( # pylint: disable=too-many-statements type_binary_scores[qtype].append(binary_score) all_scores.append(score) all_binary_scores.append(binary_score) + metadata = q.get("agentic_metadata", {}) + all_tool_call_totals.append(sum(metadata.get("tool_counts", {}).values())) + all_token_usages.append(metadata.get("token_usage", {})) print("\n ── AGENTIC ──") if all_scores: @@ -713,6 +735,16 @@ def main( # pylint: disable=too-many-statements binary_overall = sum(all_binary_scores) / len(all_binary_scores) if all_binary_scores else 0 print(f" {'-'*38}") print(f" {'OVERALL':<40s}: {overall:.3f} binary={binary_overall:.3f} ({len(all_scores)} Qs)") + tool_call_mean, tool_call_std = _mean_and_std(all_tool_call_totals) + print(f" Tool calls/query: mean={tool_call_mean:.2f} std={tool_call_std:.2f}") + print(" Bench reported tokens/query:") + for metric in _TOKEN_USAGE_METRICS: + values = [usage[metric] for usage in all_token_usages if usage.get(metric) is not None] + if values: + mean, std = _mean_and_std(values) + print(f" {metric}: mean={mean:.2f} std={std:.2f}") + else: + print(f" {metric}: unavailable") else: print(" (no results)") @@ -752,6 +784,21 @@ def main( # pylint: disable=too-many-statements print("=" * 70 + "\n") +_TOKEN_USAGE_METRICS = ( + "input_tokens", + "output_tokens", + "total_tokens", +) + + +def _mean_and_std(values: list[int]) -> tuple[float, float]: + """Return population mean and standard deviation for one per-question metric.""" + if not values: + return 0.0, 0.0 + mean = sum(values) / len(values) + return mean, (sum((value - mean) ** 2 for value in values) / len(values)) ** 0.5 + + if __name__ == "__main__": import argparse diff --git a/benchmark/longmemeval/run.py b/benchmark/longmemeval/run.py index f974b3a5..baf386ae 100644 --- a/benchmark/longmemeval/run.py +++ b/benchmark/longmemeval/run.py @@ -259,6 +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 track_agent_token_usage, track_job_counts reme_cfg = eval_config["reme"] dream_trigger_hour = reme_cfg.get("dream_trigger_hour", 23) @@ -432,16 +433,23 @@ 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}", ) - query_resp = await app.run_job( - "agentic_answer", - query=question, - query_time=query_time, - ) + with ( + track_job_counts(["search"], app.context) as tool_counts, + track_agent_token_usage( + ["bench"], + app.context, + ) as token_usages, + ): + query_resp = await app.run_job("agentic_answer", query=question, query_time=query_time) + agentic_tool_counts = tool_counts + agentic_token_usage = token_usages["bench"] agentic_response = (query_resp.answer or "").strip() if not agentic_response: agentic_response = "(no answer generated)" logger.info(f"[Item {item_index}] Agentic response: {agentic_response[:200]}...") + logger.info(f"[Item {item_index}] Agentic tool calls: {agentic_tool_counts}") + logger.info(f"[Item {item_index}] Bench token usage: {agentic_token_usage}") # ── Phase 5: Judge agentic response (via answer_judge_step) ────────── logger.info(f"[Item {item_index}] Judging agentic (binary, type={item['question_type']})...") @@ -464,6 +472,8 @@ async def evaluate_item(item: dict, eval_config: dict, item_index: int, eval_onl "ground_truth": item["answer"], "agentic_response": agentic_response, "agentic_judgment": agentic_judgment, + "agentic_tool_counts": agentic_tool_counts, + "agentic_token_usage": agentic_token_usage, "sessions_ingested": len(sorted_sessions), "dreams_triggered": len(dream_dates_triggered), } @@ -561,6 +571,16 @@ def main( f"Filtered by question_types={question_types}: {before_filter} -> {len(items_with_idx)} items", ) + # Filter by question_id if specified + question_ids = dataset_cfg.get("question_ids") or [] + if question_ids: + qid_set = set(question_ids) + before_filter = len(items_with_idx) + items_with_idx = [(idx, item) for idx, item in items_with_idx if item.get("question_id") in qid_set] + logger.info( + f"Filtered by question_ids ({len(qid_set)} ids): {before_filter} -> {len(items_with_idx)} items", + ) + logger.info( "Evaluating %d item(s) starting from index %d%s", len(items_with_idx), @@ -711,6 +731,18 @@ def _print_summary(results: list[dict], start_time: float) -> None: # Agentic stats print("\n ── Agentic (ReAct) ──") print(f" Overall accuracy: {agentic_correct}/{total} ({100*agentic_correct/total:.1f}%)") + tool_call_totals = [sum(r.get("agentic_tool_counts", {}).values()) for r in results] + tool_call_mean, tool_call_std = _mean_and_std(tool_call_totals) + print(f" Tool calls/query: mean={tool_call_mean:.2f} std={tool_call_std:.2f}") + token_usages = [r.get("agentic_token_usage", {}) for r in results] + print(" Bench reported tokens/query:") + for metric in _TOKEN_USAGE_METRICS: + values = [usage[metric] for usage in token_usages if usage.get(metric) is not None] + if values: + mean, std = _mean_and_std(values) + print(f" {metric}: mean={mean:.2f} std={std:.2f}") + else: + print(f" {metric}: unavailable") print(" Per-type accuracy:") for qtype, stats in sorted(agentic_type_stats.items()): acc = 100 * stats["correct"] / stats["total"] if stats["total"] else 0 @@ -724,6 +756,21 @@ def _print_summary(results: list[dict], start_time: float) -> None: print("=" * 60 + "\n") +_TOKEN_USAGE_METRICS = ( + "input_tokens", + "output_tokens", + "total_tokens", +) + + +def _mean_and_std(values: list[int]) -> tuple[float, float]: + """Return population mean and standard deviation for one per-query metric.""" + if not values: + return 0.0, 0.0 + mean = sum(values) / len(values) + return mean, (sum((value - mean) ** 2 for value in values) / len(values)) ** 0.5 + + if __name__ == "__main__": import argparse diff --git a/reme/components/agent_wrapper/as_agent_wrapper.py b/reme/components/agent_wrapper/as_agent_wrapper.py index f1f8e9ec..3b539bf5 100644 --- a/reme/components/agent_wrapper/as_agent_wrapper.py +++ b/reme/components/agent_wrapper/as_agent_wrapper.py @@ -59,7 +59,7 @@ from .base_agent_wrapper import BaseAgentWrapper from ..as_llm import BaseAsLLM from ..component_registry import R from ...enumeration import ChunkEnum -from ...schema import StreamChunk +from ...schema import StreamChunk, TokenUsage from ...utils import AsStateHandler if TYPE_CHECKING: @@ -149,6 +149,11 @@ class AsAgentWrapper(BaseAgentWrapper): SDK_PACKAGE = "agentscope" + @staticmethod + def _agentscope_usage(usage: Any) -> TokenUsage: + """Normalize AgentScope's portable input/output usage.""" + return TokenUsage.from_provider(usage) + def __init__(self, as_llm: str = "default", session_retention_days: int = 10, **kwargs): super().__init__(**kwargs) self.as_llm = self.bind(as_llm, BaseAsLLM, optional=False) @@ -344,9 +349,9 @@ class AsAgentWrapper(BaseAgentWrapper): agent, inputs = await self._build_agent(inputs, **kwargs) await agent.observe(inputs) - await agent.reply() + last_msg = await agent.reply() + usage = self._agentscope_usage(last_msg.usage) if last_msg.usage is not None else None await self._dump_state(agent.state) - last_msg = agent.state.context[-1] result = { "session_id": agent.state.session_id, @@ -354,6 +359,9 @@ class AsAgentWrapper(BaseAgentWrapper): "result": last_msg.get_text_content(), } + if usage is None: + self.logger.error("AgentScope did not return token usage; token accounting is unavailable for this reply.") + output_schema: dict | None = kwargs.get("output_schema") if output_schema is not None: assert self.as_llm is not None, "AsAgentWrapper requires a bound as_llm component with a valid model." @@ -365,6 +373,17 @@ class AsAgentWrapper(BaseAgentWrapper): tool_choice=ToolChoice(mode="auto"), ) result["structured_output"] = res.content + if res.usage is None: + usage = None + self.logger.error( + "AgentScope did not return structured-output token usage; token accounting is unavailable.", + ) + elif usage is not None: + usage = TokenUsage.combine([usage, self._agentscope_usage(res.usage)]) + + result["usage"] = usage.model_dump() if usage is not None else None + if usage is not None: + self._record_token_usage(usage) return result @@ -453,13 +472,16 @@ class AsAgentWrapper(BaseAgentWrapper): if isinstance(event, ModelCallStartEvent): return cls._chunk(ChunkEnum.USAGE, chunk="", metadata={"model_name": getattr(event, "model_name", None)}) if isinstance(event, ModelCallEndEvent): - usage = {"input_tokens": event.input_tokens, "output_tokens": event.output_tokens} + usage = TokenUsage(input_tokens=event.input_tokens, output_tokens=event.output_tokens) return cls._chunk( ChunkEnum.USAGE, - chunk=json.dumps(usage), - input_tokens=event.input_tokens, - output_tokens=event.output_tokens, - metadata={"model_name": getattr(event, "model_name", None)}, + chunk=json.dumps(usage.model_dump()), + input_tokens=usage.input_tokens, + output_tokens=usage.output_tokens, + metadata={ + "model_name": getattr(event, "model_name", None), + "usage": usage.model_dump(), + }, ) if isinstance(event, ExceedMaxItersEvent): return cls._chunk(ChunkEnum.ERROR, chunk="Exceeded max iterations") diff --git a/reme/components/agent_wrapper/base_agent_wrapper.py b/reme/components/agent_wrapper/base_agent_wrapper.py index 81960268..69f1ace6 100644 --- a/reme/components/agent_wrapper/base_agent_wrapper.py +++ b/reme/components/agent_wrapper/base_agent_wrapper.py @@ -11,7 +11,8 @@ from pydantic import BaseModel from ..base_component import BaseComponent from ..outbound_proxy import BaseOutboundProxy from ...enumeration import ChunkEnum, ComponentEnum -from ...schema import StreamChunk +from ...schema import StreamChunk, TokenUsage +from ...utils import global_counter_add_many if TYPE_CHECKING: from ..job.base_job import BaseJob @@ -22,6 +23,7 @@ class BaseAgentWrapper(BaseComponent): component_type = ComponentEnum.AGENT_WRAPPER SDK_PACKAGE: ClassVar[str | None] = None + TOKEN_COUNTER_PREFIX: ClassVar[str] = "__token_counter" def __init__( self, @@ -227,6 +229,23 @@ class BaseAgentWrapper(BaseComponent): """Create a StreamChunk with a short backend-friendly call site.""" return StreamChunk(chunk_type=chunk_type, **kwargs) + def _record_token_usage(self, usage: TokenUsage) -> None: + """Add one completed invocation to the application token tree.""" + if self.app_context is None: + return + counters = getattr(self.app_context, "metadata", None) + if not isinstance(counters, dict): + return + prefix = (self.TOKEN_COUNTER_PREFIX, self.name) + global_counter_add_many( + counters, + { + (*prefix, "input_tokens"): usage.input_tokens, + (*prefix, "output_tokens"): usage.output_tokens, + (*prefix, "total_tokens"): usage.total_tokens, + }, + ) + @abstractmethod async def reply(self, inputs: Any, **kwargs) -> dict: """Send inputs to the agent and return a dict with session_id and last_message.""" diff --git a/reme/components/agent_wrapper/cc_agent_wrapper.py b/reme/components/agent_wrapper/cc_agent_wrapper.py index 064b832c..068175df 100644 --- a/reme/components/agent_wrapper/cc_agent_wrapper.py +++ b/reme/components/agent_wrapper/cc_agent_wrapper.py @@ -11,7 +11,7 @@ from typing import Any, TYPE_CHECKING from .base_agent_wrapper import BaseAgentWrapper from ..component_registry import R from ...enumeration import ChunkEnum -from ...schema import StreamChunk +from ...schema import StreamChunk, TokenUsage if TYPE_CHECKING: from claude_agent_sdk import AssistantMessage, ResultMessage, UserMessage @@ -36,6 +36,11 @@ class CcAgentWrapper(BaseAgentWrapper): DEFAULT_DISALLOWED_TOOLS = ["WebSearch"] MCP_SERVER_NAME = "mcp_server" + @staticmethod + def _claude_usage(usage: dict[str, Any] | None) -> TokenUsage: + """Normalize Claude CLI's input/output usage.""" + return TokenUsage.from_provider(usage or {}) + @property def session_path(self) -> Path: """Directory used for persisted Claude Code sessions.""" @@ -248,12 +253,17 @@ class CcAgentWrapper(BaseAgentWrapper): if event_type == "message_delta": delta = raw.get("delta", {}) usage = raw.get("usage", {}) + normalized = cls._claude_usage(usage) return cls._chunk( ChunkEnum.USAGE, session_id=session_id, - chunk=json.dumps(usage), - output_tokens=usage.get("output_tokens"), - metadata={"stop_reason": delta.get("stop_reason")}, + chunk=json.dumps(normalized.model_dump()), + input_tokens=normalized.input_tokens, + output_tokens=normalized.output_tokens, + metadata={ + "stop_reason": delta.get("stop_reason"), + "usage": normalized.model_dump(), + }, ) if event_type == "message_stop": @@ -389,14 +399,16 @@ class CcAgentWrapper(BaseAgentWrapper): """Convert the SDK terminal result into usage and error chunks.""" session_id = msg.session_id or "" usage = msg.usage or {} + normalized = cls._claude_usage(usage) chunks = [ cls._chunk( ChunkEnum.USAGE, session_id=session_id, - chunk=json.dumps(usage), - input_tokens=usage.get("input_tokens"), - output_tokens=usage.get("output_tokens"), + chunk=json.dumps(normalized.model_dump()), + input_tokens=normalized.input_tokens, + output_tokens=normalized.output_tokens, metadata={ + "usage": normalized.model_dump(), "duration_ms": msg.duration_ms, "duration_api_ms": msg.duration_api_ms, "stop_reason": msg.stop_reason, @@ -443,11 +455,15 @@ class CcAgentWrapper(BaseAgentWrapper): if last_msg is None: raise ValueError("No message received from Claude Code.") + usage = self._claude_usage(last_msg.usage) if last_msg.usage is not None else None result = { "session_id": last_msg.session_id or "", "last_message": asdict(last_msg), "result": last_msg.result, + "usage": usage.model_dump() if usage is not None else None, } + if usage is not None: + self._record_token_usage(usage) if kwargs.get("output_schema") is not None: result["structured_output"] = last_msg.structured_output return result diff --git a/reme/components/agent_wrapper/codex_agent_wrapper.py b/reme/components/agent_wrapper/codex_agent_wrapper.py index bd97203b..882fbfbf 100644 --- a/reme/components/agent_wrapper/codex_agent_wrapper.py +++ b/reme/components/agent_wrapper/codex_agent_wrapper.py @@ -18,7 +18,7 @@ from typing import Any, TYPE_CHECKING from .base_agent_wrapper import BaseAgentWrapper from ..component_registry import R from ...enumeration import ChunkEnum -from ...schema import StreamChunk +from ...schema import StreamChunk, TokenUsage if TYPE_CHECKING: from openai_codex import AsyncCodex, AsyncThread, CodexConfig, RunInput @@ -79,6 +79,11 @@ class CodexAgentWrapper(BaseAgentWrapper): }, ) + @staticmethod + def _codex_usage(usage: Any) -> TokenUsage: + """Normalize Codex's full-turn input/output usage snapshot.""" + return TokenUsage.from_provider(usage) + # pylint: disable=too-many-arguments def __init__( self, @@ -421,7 +426,12 @@ class CodexAgentWrapper(BaseAgentWrapper): "last_message": final_response, "result": final_response, "turn": self._serialize(result), + "usage": None, } + if (raw_usage := getattr(result, "usage", None)) is not None: + usage = self._codex_usage(raw_usage.last) + response["usage"] = usage.model_dump() + self._record_token_usage(usage) if kwargs.get("output_schema") is not None: try: response["structured_output"] = json.loads(final_response) @@ -508,13 +518,14 @@ class CodexAgentWrapper(BaseAgentWrapper): ] if method == "thread/tokenUsage/updated": usage = payload.token_usage.last - data = cls._serialize(usage) + normalized = cls._codex_usage(usage) return [ make_chunk( ChunkEnum.USAGE, - chunk=data, - input_tokens=usage.input_tokens, - output_tokens=usage.output_tokens, + chunk=normalized.model_dump(), + input_tokens=normalized.input_tokens, + output_tokens=normalized.output_tokens, + metadata={"usage": normalized.model_dump()}, ), ] if method == "error": diff --git a/reme/components/job/background_job.py b/reme/components/job/background_job.py index fc62662a..662840e9 100644 --- a/reme/components/job/background_job.py +++ b/reme/components/job/background_job.py @@ -123,6 +123,7 @@ class BackgroundJob(BaseJob): async def __call__(self, **kwargs) -> Response: """Default body: run steps in order; errors propagate to supervisor.""" + self._record_call() merged = {**self.kwargs, **kwargs} context = RuntimeContext(stop_event=self._stop_event, **merged) for step in self._build_steps(): diff --git a/reme/components/job/base_job.py b/reme/components/job/base_job.py index 834e0125..5133111d 100644 --- a/reme/components/job/base_job.py +++ b/reme/components/job/base_job.py @@ -7,6 +7,7 @@ from ..component_registry import R from ..runtime_context import RuntimeContext from ...enumeration import ComponentEnum from ...schema import ComponentConfig, Response +from ...utils import global_counter_inc if TYPE_CHECKING: from ...steps import BaseStep @@ -57,8 +58,15 @@ class BaseJob(BaseComponent): # dict(params) copies kwargs so steps cannot mutate the shared spec. return [step_cls(**dict(params)) for step_cls, params in self.step_specs] + def _record_call(self) -> None: + """Increment this job's application-lifetime call counter.""" + metadata = getattr(self.app_context, "metadata", None) + if isinstance(metadata, dict): + global_counter_inc(metadata, ["__job_counter", self.name]) + async def __call__(self, **kwargs) -> Response: """Run all steps in order, capturing any failure into the response.""" + self._record_call() merged = {**self.kwargs, **kwargs} context = RuntimeContext(**merged) try: diff --git a/reme/components/job/cron_job.py b/reme/components/job/cron_job.py index 595653e3..a72a7878 100644 --- a/reme/components/job/cron_job.py +++ b/reme/components/job/cron_job.py @@ -47,6 +47,7 @@ class CronJob(BackgroundJob): if self._stop_event.is_set(): break try: + self._record_call() await self._execute_steps() except Exception as exc: self.logger.exception(f"Cron job '{self.name}' failed: {exc}") diff --git a/reme/components/job/stream_job.py b/reme/components/job/stream_job.py index 8eb93f92..c48bda75 100644 --- a/reme/components/job/stream_job.py +++ b/reme/components/job/stream_job.py @@ -12,6 +12,7 @@ class StreamJob(BaseJob): async def __call__(self, **kwargs) -> None: """Run steps; emit failures as ERROR chunks, then a terminal DONE marker.""" + self._record_call() merged = {**self.kwargs, **kwargs} context = RuntimeContext(**merged) try: diff --git a/reme/schema/__init__.py b/reme/schema/__init__.py index 2498137c..fc919f66 100644 --- a/reme/schema/__init__.py +++ b/reme/schema/__init__.py @@ -40,6 +40,7 @@ from .file_node import FileNode from .request import Request from .response import Response from .stream_chunk import StreamChunk +from .token_usage import TokenUsage __all__ = [ "ApplicationConfig", @@ -83,5 +84,6 @@ __all__ = [ "Response", "SelectedPaper", "StreamChunk", + "TokenUsage", "TopicSelectionOutput", ] diff --git a/reme/schema/token_usage.py b/reme/schema/token_usage.py new file mode 100644 index 00000000..df178152 --- /dev/null +++ b/reme/schema/token_usage.py @@ -0,0 +1,51 @@ +"""Backend-neutral token accounting contracts.""" + +from typing import Any + +from pydantic import BaseModel, Field, model_validator + + +class TokenUsage(BaseModel): + """Portable token usage reported for one completed agent invocation. + + Only the provider's top-level input and output counters are retained. + Provider-specific cache and reasoning breakdowns are intentionally + excluded, so values may not share identical billing semantics across + providers. ``total_tokens`` is always their derived sum. + """ + + input_tokens: int = Field(default=0, ge=0) + output_tokens: int = Field(default=0, ge=0) + total_tokens: int = Field(default=0, ge=0) + + @model_validator(mode="after") + def _set_total(self) -> "TokenUsage": + self.total_tokens = self.input_tokens + self.output_tokens + return self + + @classmethod + def from_provider( + cls, + usage: Any, + ) -> "TokenUsage": + """Keep only a provider's portable top-level input/output counters.""" + + def get(*names: str) -> int | None: + for name in names: + value = usage.get(name) if isinstance(usage, dict) else getattr(usage, name, None) + if value is not None: + return int(value) + return None + + return cls( + input_tokens=get("input_tokens", "prompt_tokens") or 0, + output_tokens=get("output_tokens", "completion_tokens") or 0, + ) + + @classmethod + def combine(cls, usages: list["TokenUsage"]) -> "TokenUsage": + """Combine completed model calls into one full-invocation usage.""" + return cls( + input_tokens=sum(item.input_tokens for item in usages), + output_tokens=sum(item.output_tokens for item in usages), + ) diff --git a/reme/steps/benchmark/base/agentic_answer.py b/reme/steps/benchmark/base/agentic_answer.py index 1423b5e7..27b9f26c 100644 --- a/reme/steps/benchmark/base/agentic_answer.py +++ b/reme/steps/benchmark/base/agentic_answer.py @@ -5,7 +5,7 @@ import os from ...base_step import BaseStep from ...index._dedup import _ToolContextDedupMixin from ....enumeration import ChunkEnum -from ....utils.counter import global_counter_next +from ....utils.counter import global_counter_inc class BaseAgenticAnswerStep(BaseStep): @@ -25,6 +25,8 @@ class BaseAgenticAnswerStep(BaseStep): The agent's final answer text. """ + # Reasoning-round budget. The AgentScope wrapper converts it to the + # backend's iteration-counting semantics; other backends ignore it. MAX_ITERATION = 10 TOOL_CONTEXT_PREFIX: str = "content_agentic_answer" @@ -46,7 +48,7 @@ class BaseAgenticAnswerStep(BaseStep): if self.app_context is not None: tool_context_id = ( f"{self.TOOL_CONTEXT_PREFIX}_{os.getpid()}_" - f"{global_counter_next(self.app_context.metadata, [self.TOOL_CONTEXT_PREFIX])}" + f"{global_counter_inc(self.app_context.metadata, [self.TOOL_CONTEXT_PREFIX])}" ) else: tool_context_id = f"{self.TOOL_CONTEXT_PREFIX}_{os.getpid()}_local" diff --git a/reme/utils/__init__.py b/reme/utils/__init__.py index 69a16ed6..a41e66e9 100644 --- a/reme/utils/__init__.py +++ b/reme/utils/__init__.py @@ -15,7 +15,13 @@ from .service_utils import find_reme, locate_reme, precheck_start, cli_find_reme from .similarity_utils import cosine_similarity, batch_cosine_similarity from .token_utils import estimate_token_count from .agent_state_io import AsStateHandler -from .counter import global_counter_next +from .counter import ( + global_counter_add, + global_counter_add_many, + global_counter_get, + global_counter_get_all, + global_counter_inc, +) __all__ = [ "hash_text", @@ -38,5 +44,9 @@ __all__ = [ "batch_cosine_similarity", "estimate_token_count", "AsStateHandler", - "global_counter_next", + "global_counter_add", + "global_counter_add_many", + "global_counter_get", + "global_counter_get_all", + "global_counter_inc", ] diff --git a/reme/utils/counter.py b/reme/utils/counter.py index 204d9cf6..a2f578af 100644 --- a/reme/utils/counter.py +++ b/reme/utils/counter.py @@ -1,18 +1,80 @@ """Thread-safe monotonic counter tree utility for shared application state.""" +import copy import threading +from collections.abc import Mapping from typing import Any COUNTER_TREE_KEY = "_counter_tree" COUNTER_LOCK_KEY = "_counter_tree_lock" +_COUNTER_INIT_LOCK = threading.Lock() -def global_counter_next(metadata: dict[str, Any], key: list[str]) -> int: - """Return the next monotonic value for ``key``, starting at 1. +def _get_counter_lock(metadata: dict[str, Any]) -> Any: + """Return the metadata-scoped lock, creating it once when needed.""" + lock = metadata.get(COUNTER_LOCK_KEY) + if lock is not None: + return lock + + # Two threads may reach the first counter operation concurrently. Guard + # initialization so they cannot install and then use different locks. + with _COUNTER_INIT_LOCK: + lock = metadata.get(COUNTER_LOCK_KEY) + if lock is None: + lock = threading.Lock() + metadata[COUNTER_LOCK_KEY] = lock + return lock + + +def global_counter_add_many( + metadata: dict[str, Any], + updates: Mapping[tuple[str, ...], int], +) -> dict[tuple[str, ...], int]: + """Atomically fetch-and-add multiple counter paths. + + All paths are validated before the counter tree is mutated. The returned + mapping contains each path's value immediately before its increment. + """ + normalized = dict(updates) + for path, value in normalized.items(): + if not isinstance(path, tuple) or not all(isinstance(part, str) for part in path): + raise TypeError("counter paths must be tuples of strings") + if not isinstance(value, int): + raise TypeError("counter increments must be integers") + if not normalized: + return {} + + lock = _get_counter_lock(metadata) + with lock: + tree = metadata.get(COUNTER_TREE_KEY) + if tree is None: + tree = {"value": 0, "children": {}} + metadata[COUNTER_TREE_KEY] = tree + + nodes: dict[tuple[str, ...], dict[str, Any]] = {} + for path in normalized: + node = tree + for part in path: + child = node["children"].get(part) + if child is None: + child = {"value": 0, "children": {}} + node["children"][part] = child + node = child + nodes[path] = node + + previous = {path: node["value"] for path, node in nodes.items()} + for path, value in normalized.items(): + nodes[path]["value"] += value + return previous + + +def global_counter_add(metadata: dict[str, Any], key: list[str], val: int) -> int: + """Fetch-and-add: return the old value for ``key``, then add ``val`` to it. Walks the counter tree stored in ``metadata`` along ``key``, creating - missing nodes on the way, then increments and returns the target node's - counter. An empty ``key`` increments the root node, which serves as a + missing nodes on the way, then returns the target node's current counter + value and adds ``val`` to it. Counters start at 0, so the first call + returns 0. An empty ``key`` targets the root node, which serves as a process-wide thread-safe global counter. The counter tree (``{"value": 0, "children": {}}``) and its @@ -21,25 +83,62 @@ def global_counter_next(metadata: dict[str, Any], key: list[str]) -> int: If they are missing they are created lazily so the function is safe to call with a plain ``dict``. """ - lock = metadata.get(COUNTER_LOCK_KEY) - if lock is None: - lock = threading.Lock() - metadata[COUNTER_LOCK_KEY] = lock + path = tuple(key) + return global_counter_add_many(metadata, {path: val})[path] + + +def global_counter_inc(metadata: dict[str, Any], key: list[str]) -> int: + """Fetch-and-increment: return the old value for ``key``, then add 1. + + Counters start at 0, so the first call returns 0. See + :func:`global_counter_add` for details on the counter tree layout. + """ + return global_counter_add(metadata, key, 1) + + +def global_counter_get(metadata: dict[str, Any], key: list[str]) -> int: + """Return the current value for ``key`` without modifying the tree. + + Unlike :func:`global_counter_add`, missing nodes are never created; a + path that does not exist yet is reported as 0, matching the value the + node would hold right before its first increment. + """ + lock = _get_counter_lock(metadata) with lock: tree = metadata.get(COUNTER_TREE_KEY) if tree is None: - tree = {"value": 0, "children": {}} - metadata[COUNTER_TREE_KEY] = tree + return 0 - node: dict[str, Any] = tree + node: dict[str, Any] | None = tree for part in key: assert isinstance(part, str) - tmp = node["children"].get(part, None) - if tmp is None: - tmp = {"value": 0, "children": {}} - node["children"][part] = tmp - node = tmp - res = node["value"] + 1 - node["value"] = res - return res + node = node["children"].get(part) + if node is None: + return 0 + return node["value"] + + +def global_counter_get_all(metadata: dict[str, Any], key: list[str]) -> dict[str, Any] | None: + """Return a deep copy of the subtree rooted at ``key``, or ``None``. + + Walks the counter tree along ``key`` without creating missing nodes and + returns a deep copy of the node found there (``{"value": ..., "children": + ...}``), so callers can inspect it without racing concurrent updates. + Returns ``None`` when the tree or any part of ``key`` does not exist. + An empty ``key`` returns a copy of the whole tree. + """ + lock = _get_counter_lock(metadata) + + with lock: + tree = metadata.get(COUNTER_TREE_KEY) + if tree is None: + return None + + node: dict[str, Any] | None = tree + for part in key: + assert isinstance(part, str) + node = node["children"].get(part) + if node is None: + return None + return copy.deepcopy(node) diff --git a/reme/utils/evaluation_interface.py b/reme/utils/evaluation_interface.py new file mode 100644 index 00000000..1f6819d4 --- /dev/null +++ b/reme/utils/evaluation_interface.py @@ -0,0 +1,191 @@ +"""Read-only evaluation helpers for application job execution statistics. + +These helpers take before/after snapshots of application-lifetime counters. +They are intentionally not thread-safe request attribution: overlapping calls +in the same Application contribute to each other's deltas. They are intended +for the benchmark utilities, where each tracked evaluation runs without other +work sharing its Application instance. +""" + +from typing import TYPE_CHECKING + +from .counter import global_counter_get, global_counter_get_all + +if TYPE_CHECKING: + from ..components.application_context import ApplicationContext + + +_TOKEN_METRICS = ( + "input_tokens", + "output_tokens", + "total_tokens", +) + + +def check_job_count(job_name: str, app_context: "ApplicationContext") -> int: + """Return the application-lifetime execution count for a registered job. + + ``app_context`` scopes the lookup because ReMe does not maintain a global + current Application instance. Unknown job names use the same ``KeyError`` + contract as :meth:`Application.run_job`. + """ + if job_name not in app_context.jobs: + raise KeyError(f"Job '{job_name}' not found") + return global_counter_get(app_context.metadata, ["__job_counter", job_name]) + + +class JobCountTracker: + """Measure registered job calls made while this context is active. + + Not thread-safe for per-request attribution; see the module docstring. + """ + + 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) + + +def check_agent_token_count( + agent_name: str, + app_context: "ApplicationContext", + metric: str = "total_tokens", +) -> int: + """Return one application-lifetime token metric for an agent wrapper. + + The agent name is the configured ``agent_wrapper`` component name (for + example ``"bench"``), and ``metric`` is one leaf in ReMe's token counter + tree, such as ``input_tokens`` or ``total_tokens``. + """ + return global_counter_get(app_context.metadata, ["__token_counter", agent_name, metric]) + + +def check_agent_token_usage(agent_name: str, app_context: "ApplicationContext") -> dict[str, int | None]: + """Return all token metrics currently accumulated for one agent wrapper. + + A metric remains ``None`` until the backend has reported it. This keeps + unavailable usage distinct from zero. + """ + tree = global_counter_get_all(app_context.metadata, ["__token_counter", agent_name]) + children = tree.get("children", {}) if tree is not None else {} + usage: dict[str, int | None] = {} + for metric in _TOKEN_METRICS: + node = children.get(metric) + usage[metric] = node["value"] if node is not None else None + return usage + + +class AgentTokenCountTracker: + """Measure one token metric for agent wrappers during a context block. + + Not thread-safe for per-request attribution; see the module docstring. + """ + + def __init__( + self, + agent_names: list[str], + app_context: "ApplicationContext", + metric: str = "total_tokens", + ) -> None: + self.agent_names = list(dict.fromkeys(agent_names)) + self.app_context = app_context + self.metric = metric + self._start_counts: dict[str, int] = {} + self.counts: dict[str, int] = {} + + def __enter__(self) -> dict[str, int]: + self._start_counts = { + name: check_agent_token_count(name, self.app_context, self.metric) for name in self.agent_names + } + return self.counts + + def __exit__(self, exc_type, exc_value, traceback) -> bool: + self.counts.update( + { + name: check_agent_token_count(name, self.app_context, self.metric) - start_count + for name, start_count in self._start_counts.items() + }, + ) + return False + + +def track_agent_token_counts( + agent_names: list[str], + app_context: "ApplicationContext", + metric: str = "total_tokens", +) -> AgentTokenCountTracker: + """Return a context manager that reports agent token deltas. + + Example: + + .. code-block:: python + + with track_agent_token_counts(["bench"], app.context) as counts: + await app.run_job("agentic_answer", query="...") + assert counts["bench"] > 0 + """ + return AgentTokenCountTracker(agent_names, app_context, metric) + + +class AgentTokenUsageTracker: + """Measure all token metrics for agent wrappers during a context block. + + Not thread-safe for per-request attribution; see the module docstring. + """ + + def __init__(self, agent_names: list[str], app_context: "ApplicationContext") -> None: + self.agent_names = list(dict.fromkeys(agent_names)) + self.app_context = app_context + self._start_usage: dict[str, dict[str, int | None]] = {} + self.usages: dict[str, dict[str, int | None]] = {} + + def __enter__(self) -> dict[str, dict[str, int | None]]: + self._start_usage = {name: check_agent_token_usage(name, self.app_context) for name in self.agent_names} + return self.usages + + def __exit__(self, exc_type, exc_value, traceback) -> bool: + for name in self.agent_names: + end_usage = check_agent_token_usage(name, self.app_context) + delta: dict[str, int | None] = {} + for metric in _TOKEN_METRICS: + current = end_usage[metric] + start = self._start_usage[name][metric] + delta[metric] = None if current is None else current - (start or 0) + self.usages[name] = delta + return False + + +def track_agent_token_usage( + agent_names: list[str], + app_context: "ApplicationContext", +) -> AgentTokenUsageTracker: + """Return a context manager that reports full per-agent token usage deltas.""" + return AgentTokenUsageTracker(agent_names, app_context) diff --git a/tests/unit/test_cc_agent_wrapper.py b/tests/unit/test_cc_agent_wrapper.py index cf9ac0ee..c123faa7 100644 --- a/tests/unit/test_cc_agent_wrapper.py +++ b/tests/unit/test_cc_agent_wrapper.py @@ -473,13 +473,17 @@ async def test_reply_stream_emits_one_reply_end_for_normal_sdk_lifecycle(tmp_pat usage={"input_tokens": 1, "output_tokens": 2}, ) + wrapper = _wrapper(tmp_path) + recorded_usages = [] monkeypatch.setattr("claude_agent_sdk.query", query) - chunks = [chunk async for chunk in _wrapper(tmp_path).reply_stream("hello")] + monkeypatch.setattr(wrapper, "_record_token_usage", recorded_usages.append) + chunks = [chunk async for chunk in wrapper.reply_stream("hello")] assert sum(chunk.chunk_type == ChunkEnum.REPLY_END for chunk in chunks) == 1 delta_usage = next(chunk for chunk in chunks if chunk.metadata.get("stop_reason") == "end_turn") assert delta_usage.chunk_type == ChunkEnum.USAGE assert delta_usage.output_tokens == 2 + assert not recorded_usages @pytest.mark.asyncio diff --git a/tests/unit/test_codex_agent_wrapper.py b/tests/unit/test_codex_agent_wrapper.py index b32a3a13..139887c0 100644 --- a/tests/unit/test_codex_agent_wrapper.py +++ b/tests/unit/test_codex_agent_wrapper.py @@ -862,6 +862,70 @@ async def test_reply_stream_interrupts_turn_when_consumer_closes_early(tmp_path, assert stream_closed +@pytest.mark.asyncio +async def test_reply_stream_does_not_record_token_usage(tmp_path, monkeypatch): + wrapper, _job = _wrapper(tmp_path, auth_mode="oauth") + recorded_usages = [] + usage = TokenUsageBreakdown( + cachedInputTokens=0, + inputTokens=3, + outputTokens=5, + reasoningOutputTokens=0, + totalTokens=8, + ) + + class FakeTurn: + id = "turn-1" + + async def stream(self): + yield SimpleNamespace( + method="thread/tokenUsage/updated", + payload=SimpleNamespace(token_usage=SimpleNamespace(last=usage)), + ) + yield SimpleNamespace( + method="turn/completed", + payload=SimpleNamespace( + turn=SimpleNamespace( + id=self.id, + status=SimpleNamespace(value="completed"), + duration_ms=1, + error=None, + ), + ), + ) + + async def interrupt(self): + raise AssertionError("Completed turns must not be interrupted") + + class FakeThread: + id = "thread-1" + + async def turn(self, _inputs, **_kwargs): + return FakeTurn() + + class FakeCodex: + def __init__(self, _config): + pass + + async def account(self): + return SimpleNamespace(account=SimpleNamespace()) + + async def close(self): + return None + + async def thread_start(self, **_kwargs): + return FakeThread() + + monkeypatch.setattr("reme.components.agent_wrapper.codex_agent_wrapper.AsyncCodex", FakeCodex) + monkeypatch.setattr(wrapper, "_record_token_usage", recorded_usages.append) + + chunks = [chunk async for chunk in wrapper.reply_stream("answer")] + await wrapper.close() + + assert any(chunk.chunk_type == ChunkEnum.USAGE for chunk in chunks) + assert not recorded_usages + + @pytest.mark.asyncio async def test_close_waits_for_active_turn(tmp_path, monkeypatch): wrapper, _job = _wrapper(tmp_path, auth_mode="oauth") diff --git a/tests/unit/test_evaluation_interface.py b/tests/unit/test_evaluation_interface.py new file mode 100644 index 00000000..fd58b2a9 --- /dev/null +++ b/tests/unit/test_evaluation_interface.py @@ -0,0 +1,142 @@ +"""Tests for read-only job execution count evaluation helpers.""" + +import asyncio +from types import SimpleNamespace + +import pytest + +from reme.components.job import BaseJob, StreamJob +from reme.utils import global_counter_add +from reme.utils.evaluation_interface import ( + check_agent_token_count, + check_job_count, + track_agent_token_usage, + track_agent_token_counts, + track_job_counts, +) + + +def test_check_job_count_reads_registered_base_job_count(): + """check_job_count returns the number of completed BaseJob invocations.""" + + async def run(): + app_context = SimpleNamespace(metadata={}, jobs={}) + job = BaseJob(name="search", app_context=app_context) + app_context.jobs[job.name] = job + + await job() + await job() + + assert check_job_count("search", app_context) == 2 + + asyncio.run(run()) + + +def test_check_job_count_reads_custom_job_by_name(): + """check_job_count reads subclassed job counters without depending on inheritance.""" + + async def run(): + class ProjectStreamJob(StreamJob): + """Project-specific StreamJob subclass used to exercise MRO lookup.""" + + app_context = SimpleNamespace(metadata={}, jobs={}) + job = ProjectStreamJob(name="chat", app_context=app_context) + app_context.jobs[job.name] = job + + await job(stream_queue=asyncio.Queue()) + + assert check_job_count("chat", app_context) == 1 + + asyncio.run(run()) + + +def test_check_job_count_rejects_unknown_job_name(): + """Unknown job names raise KeyError, matching Application.run_job.""" + app_context = SimpleNamespace(metadata={}, jobs={}) + + 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()) + + +def test_track_agent_token_counts_returns_delta_for_one_agent(): + """Token tracking mirrors job-count tracking over the token counter tree.""" + app_context = SimpleNamespace(metadata={}, jobs={}) + global_counter_add(app_context.metadata, ["__token_counter", "bench", "total_tokens"], 10) + + with track_agent_token_counts(["bench"], app_context) as counts: + global_counter_add(app_context.metadata, ["__token_counter", "bench", "total_tokens"], 25) + + assert counts == {"bench": 25} + assert check_agent_token_count("bench", app_context) == 35 + + +def test_track_agent_token_usage_reports_only_supported_metrics(): + """Detailed usage tracking reports the shared input/output contract.""" + app_context = SimpleNamespace(metadata={}, jobs={}) + for metric, value in (("input_tokens", 10), ("output_tokens", 5), ("total_tokens", 15)): + global_counter_add(app_context.metadata, ["__token_counter", "bench", metric], value) + + with track_agent_token_usage(["bench"], app_context) as usages: + for metric, value in (("input_tokens", 20), ("output_tokens", 7), ("total_tokens", 27)): + global_counter_add(app_context.metadata, ["__token_counter", "bench", metric], value) + + assert usages == { + "bench": { + "input_tokens": 20, + "output_tokens": 7, + "total_tokens": 27, + }, + } + + +def test_track_agent_token_usage_keeps_unavailable_usage_as_none(): + """A backend that reports no usage remains unavailable to benchmarks.""" + app_context = SimpleNamespace(metadata={}, jobs={}) + + with track_agent_token_usage(["bench"], app_context) as usages: + pass + + assert usages == { + "bench": { + "input_tokens": None, + "output_tokens": None, + "total_tokens": None, + }, + } diff --git a/tests/unit/test_job.py b/tests/unit/test_job.py index 06b13998..b1749975 100644 --- a/tests/unit/test_job.py +++ b/tests/unit/test_job.py @@ -18,6 +18,7 @@ from reme.components.job.cron_job import CronJob from reme.components.job.stream_job import StreamJob from reme.components.job import cron_job as cron_job_module from reme.schema import ComponentConfig +from reme.utils import global_counter_get # -- helpers ------------------------------------------------------------------ @@ -132,6 +133,73 @@ def test_stream_job_merges_config_kwargs_into_context(): asyncio.run(run()) +# -- Job call counters ------------------------------------------------------- + + +def test_base_job_records_calls_by_name(): + async def run(): + app_context = SimpleNamespace(metadata={}) + job = BaseJob(name="search", app_context=app_context) + + await job() + await job() + + assert global_counter_get(app_context.metadata, ["__job_counter", "search"]) == 2 + + asyncio.run(run()) + + +def test_stream_job_subclass_records_calls_by_job_name(): + async def run(): + class ProjectStreamJob(StreamJob): + pass + + app_context = SimpleNamespace(metadata={}) + job = ProjectStreamJob(name="chat", app_context=app_context) + + await job(stream_queue=asyncio.Queue()) + + assert global_counter_get(app_context.metadata, ["__job_counter", "chat"]) == 1 + + asyncio.run(run()) + + +def test_background_job_records_calls_by_name(): + async def run(): + app_context = SimpleNamespace(metadata={}) + job = BackgroundJob(name="watch", app_context=app_context) + + await job() + + assert global_counter_get(app_context.metadata, ["__job_counter", "watch"]) == 1 + + asyncio.run(run()) + + +def test_cron_job_records_each_triggered_execution(): + async def run(): + app_context = SimpleNamespace(metadata={}) + job = CronJob(name="nightly", cron="* * * * *", app_context=app_context) + job._stop_event = asyncio.Event() + waits = 0 + + async def wait_once(_delay): + nonlocal waits + waits += 1 + if waits > 1: + job._stop_event.set() + + job._wait_or_stop = wait_once + job._next_fire_delay = lambda: 0.0 + job._build_steps = lambda: [] + + await job() + + assert global_counter_get(app_context.metadata, ["__job_counter", "nightly"]) == 1 + + asyncio.run(run()) + + # -- BaseJob._start requires app_context ------------------------------------ diff --git a/tests/unit/test_token_usage.py b/tests/unit/test_token_usage.py new file mode 100644 index 00000000..cd505815 --- /dev/null +++ b/tests/unit/test_token_usage.py @@ -0,0 +1,227 @@ +"""Tests for unified agent token accounting.""" + +from types import SimpleNamespace + +from agentscope.model._model_usage import ChatUsage +import pytest + +from reme.components.agent_wrapper import AsAgentWrapper, BaseAgentWrapper +from reme.components.application_context import ApplicationContext +from reme.schema import TokenUsage +from reme.utils import global_counter_get_all + + +class _UsageWrapper(BaseAgentWrapper): + async def reply(self, inputs, **kwargs): + raise NotImplementedError + + +def test_provider_usage_keeps_only_input_and_output_tokens(): + """Provider-specific cache and reasoning details are not persisted.""" + usage = TokenUsage.from_provider( + { + "input_tokens": 10, + "output_tokens": 4, + "cache_read_input_tokens": 20, + "cache_creation_input_tokens": 30, + "reasoning_output_tokens": 2, + }, + ) + + assert usage.model_dump() == { + "input_tokens": 10, + "output_tokens": 4, + "total_tokens": 14, + } + + +def test_provider_usage_uses_reported_input_without_cache_adjustment(): + """Reported input is kept unchanged across backend-specific usage shapes.""" + usage = TokenUsage.from_provider( + { + "input_tokens": 60, + "output_tokens": 4, + "cached_input_tokens": 20, + "reasoning_output_tokens": 2, + }, + ) + + assert usage.input_tokens == 60 + assert usage.total_tokens == 64 + + +def test_provider_usage_accepts_prompt_and_completion_aliases(): + """Portable OpenAI-style aliases normalize to the common counters.""" + usage = TokenUsage.from_provider({"prompt_tokens": 3, "completion_tokens": 2}) + + assert usage.model_dump() == { + "input_tokens": 3, + "output_tokens": 2, + "total_tokens": 5, + } + + +def test_total_tokens_is_always_derived_from_input_and_output(): + """Caller-supplied totals cannot violate the portable usage invariant.""" + usage = TokenUsage(input_tokens=3, output_tokens=2, total_tokens=999) + + assert usage.total_tokens == 5 + + +def test_agentscope_usage_keeps_portable_input_and_output(tmp_path): + """AgentScope usage has the same input/output-only contract.""" + context = ApplicationContext(workspace_dir=str(tmp_path)) + wrapper = AsAgentWrapper(name="research", as_llm="", app_context=context) + usage = ChatUsage( + input_tokens=10, + output_tokens=4, + time=0.0, + cache_input_tokens=20, + cache_creation_input_tokens=30, + ) + + assert wrapper._agentscope_usage(usage).model_dump() == { # pylint: disable=protected-access + "input_tokens": 10, + "output_tokens": 4, + "total_tokens": 14, + } + + +def test_combined_usage_sums_all_model_calls(): + """One wrapper invocation records aggregate input/output usage.""" + usage = TokenUsage.combine( + [ + TokenUsage(input_tokens=10, output_tokens=4), + TokenUsage(input_tokens=5, output_tokens=2), + ], + ) + + assert usage.model_dump() == { + "input_tokens": 15, + "output_tokens": 6, + "total_tokens": 21, + } + + +def test_token_counter_is_a_per_agent_metric_tree(tmp_path): + """Recorded usage accumulates input, output, and total tokens per agent.""" + context = ApplicationContext(workspace_dir=str(tmp_path)) + wrapper = _UsageWrapper(name="research", app_context=context) + wrapper._record_token_usage( # pylint: disable=protected-access + TokenUsage(input_tokens=10, output_tokens=4), + ) + wrapper._record_token_usage(TokenUsage(input_tokens=3, output_tokens=2)) # pylint: disable=protected-access + + assert global_counter_get_all(context.metadata, ["__token_counter", "research"]) == { + "value": 0, + "children": { + "input_tokens": {"value": 13, "children": {}}, + "output_tokens": {"value": 6, "children": {}}, + "total_tokens": {"value": 19, "children": {}}, + }, + } + + +@pytest.mark.asyncio +async def test_agentscope_reply_records_final_message_usage(tmp_path, monkeypatch): + """AgentScope 2.0.4.post1 reports aggregate usage on the final message.""" + context = ApplicationContext(workspace_dir=str(tmp_path)) + wrapper = AsAgentWrapper(name="research", as_llm="", app_context=context) + message = SimpleNamespace( + # Match the AgentScope 2.0.4.post1 accumulation test: 10/20 + 5/8. + usage=SimpleNamespace(input_tokens=15, output_tokens=28), + model_dump=lambda: {"text": "answer"}, + get_text_content=lambda: "answer", + ) + agent = SimpleNamespace( + state=SimpleNamespace(session_id="session-1", context=[message]), + observe=lambda _inputs: _async_none(), + reply=lambda: _async_value(message), + reply_stream=_unexpected_reply_stream, + ) + + async def build_agent(inputs, **_kwargs): + return agent, inputs + + monkeypatch.setattr(wrapper, "_build_agent", build_agent) + monkeypatch.setattr(wrapper, "_dump_state", _async_none) + + result = await wrapper.reply("hello") + + assert result["usage"] == {"input_tokens": 15, "output_tokens": 28, "total_tokens": 43} + assert ( + global_counter_get_all(context.metadata, ["__token_counter", "research"])["children"]["total_tokens"]["value"] + == 43 + ) + + +@pytest.mark.asyncio +async def test_agentscope_reply_without_usage_leaves_token_counters_unset(tmp_path, monkeypatch): + """Replies without any usage information remain visibly unavailable.""" + context = ApplicationContext(workspace_dir=str(tmp_path)) + wrapper = AsAgentWrapper(name="research", as_llm="", app_context=context) + message = SimpleNamespace( + usage=None, + model_dump=lambda: {"text": "answer"}, + get_text_content=lambda: "answer", + ) + agent = SimpleNamespace( + state=SimpleNamespace(session_id="session-1", context=[message]), + observe=lambda _inputs: _async_none(), + reply=lambda: _async_value(message), + reply_stream=_unexpected_reply_stream, + ) + + async def build_agent(inputs, **_kwargs): + return agent, inputs + + monkeypatch.setattr(wrapper, "_build_agent", build_agent) + monkeypatch.setattr(wrapper, "_dump_state", _async_none) + + result = await wrapper.reply("hello") + + assert result["usage"] is None + assert "__token_counter" not in context.metadata + + +async def _async_none(*_args, **_kwargs): + return None + + +async def _async_value(value): + return value + + +def _unexpected_reply_stream(*_args, **_kwargs): + raise AssertionError("Non-streaming replies must use Agent.reply()") + + +@pytest.mark.asyncio +async def test_agentscope_stream_reply_does_not_record_token_usage(tmp_path, monkeypatch): + """Only non-streaming AgentScope replies contribute to token accounting.""" + context = ApplicationContext(workspace_dir=str(tmp_path)) + wrapper = AsAgentWrapper(name="research", as_llm="", app_context=context) + + class FakeAgent: + """Minimal AgentScope stream double.""" + + state = type("State", (), {"session_id": "session-1"})() + + async def reply_stream(self, inputs): + """Yield no events for the supplied input.""" + if inputs is None: + yield None + + async def build_agent(inputs, **_kwargs): + """Build the minimal stream double.""" + return FakeAgent(), inputs + + async def dump_state(_state): + """Avoid durable state writes in this accounting test.""" + return None + + monkeypatch.setattr(wrapper, "_build_agent", build_agent) + monkeypatch.setattr(wrapper, "_dump_state", dump_state) + + assert [chunk async for chunk in wrapper.reply_stream("hello")] == [] + assert "__token_counter" not in context.metadata diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index d8bdcefb..6b21eef1 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -2,11 +2,21 @@ import asyncio import sys +import threading import numpy as np import pytest from reme.utils import common_utils +from reme.utils.counter import ( + COUNTER_LOCK_KEY, + COUNTER_TREE_KEY, + global_counter_add, + global_counter_add_many, + global_counter_get, + global_counter_get_all, + global_counter_inc, +) from reme.utils.similarity_utils import batch_cosine_similarity, cosine_similarity @@ -65,3 +75,187 @@ def test_mock_reme_server_uses_reme_entrypoint(monkeypatch): asyncio.run(run()) assert captured["cmd"][:4] == [sys.executable, "-m", "reme.reme", "start"] + + +def test_inc_returns_old_value_starting_at_zero(): + """First call returns 0, then values increase by 1 per call.""" + metadata: dict = {} + + assert global_counter_inc(metadata, ["a"]) == 0 + assert global_counter_inc(metadata, ["a"]) == 1 + assert global_counter_inc(metadata, ["a"]) == 2 + + +def test_add_returns_old_value_and_adds_val(): + """``add`` is fetch-and-add: old value out, ``val`` added in.""" + metadata: dict = {} + + assert global_counter_add(metadata, ["a"], 10) == 0 + assert global_counter_add(metadata, ["a"], 5) == 10 + assert global_counter_inc(metadata, ["a"]) == 15 + assert global_counter_get(metadata, ["a"]) == 16 + + +def test_add_many_returns_old_values_and_updates_all_paths(): + """``add_many`` updates sibling metrics under one counter-tree lock.""" + metadata: dict = {} + global_counter_add(metadata, ["usage", "input"], 2) + + previous = global_counter_add_many( + metadata, + { + ("usage", "input"): 3, + ("usage", "output"): 4, + ("usage", "total"): 7, + }, + ) + + assert previous == { + ("usage", "input"): 2, + ("usage", "output"): 0, + ("usage", "total"): 0, + } + assert global_counter_get(metadata, ["usage", "input"]) == 5 + assert global_counter_get(metadata, ["usage", "output"]) == 4 + assert global_counter_get(metadata, ["usage", "total"]) == 7 + + +def test_add_many_holds_the_counter_lock_once_for_the_batch(): + """A multi-metric update is one critical section, not several writes.""" + + class CountingLock: + """Lock test double that counts entered critical sections.""" + + def __init__(self): + self.lock = threading.Lock() + self.entries = 0 + + def __enter__(self): + self.lock.acquire() + self.entries += 1 + return self + + def __exit__(self, *_args): + self.lock.release() + + lock = CountingLock() + metadata = {COUNTER_LOCK_KEY: lock} + + global_counter_add_many( + metadata, + { + ("usage", "input"): 3, + ("usage", "output"): 4, + ("usage", "total"): 7, + }, + ) + + assert lock.entries == 1 + + +def test_add_many_validates_every_update_before_mutating(): + """One invalid update leaves every valid counter unchanged.""" + metadata: dict = {} + + with pytest.raises(TypeError, match="increments must be integers"): + global_counter_add_many( + metadata, + { + ("usage", "input"): 3, + ("usage", "output"): "invalid", + }, + ) + + assert COUNTER_TREE_KEY not in metadata + + +def test_counters_are_isolated_by_key_path(): + """Sibling and nested keys, plus the root, hold independent counters.""" + metadata: dict = {} + + assert global_counter_inc(metadata, ["a"]) == 0 + assert global_counter_inc(metadata, ["b"]) == 0 + assert global_counter_inc(metadata, ["a", "child"]) == 0 + assert global_counter_inc(metadata, []) == 0 + + assert global_counter_get(metadata, ["a"]) == 1 + assert global_counter_get(metadata, ["b"]) == 1 + assert global_counter_get(metadata, ["a", "child"]) == 1 + assert global_counter_get(metadata, []) == 1 + + +def test_get_does_not_create_missing_nodes(): + """``get`` reports 0 for missing paths and leaves the tree untouched.""" + metadata: dict = {} + + assert global_counter_get(metadata, ["missing"]) == 0 + assert COUNTER_TREE_KEY not in metadata + + global_counter_inc(metadata, ["a"]) + assert global_counter_get(metadata, ["a", "missing"]) == 0 + assert "missing" not in metadata[COUNTER_TREE_KEY]["children"]["a"]["children"] + + +def test_get_all_returns_none_for_missing_key(): + """``get_all`` returns None when the tree or the path does not exist.""" + metadata: dict = {} + + assert global_counter_get_all(metadata, []) is None + assert global_counter_get_all(metadata, ["missing"]) is None + + global_counter_inc(metadata, ["a"]) + assert global_counter_get_all(metadata, ["missing"]) is None + assert global_counter_get_all(metadata, ["a", "missing"]) is None + + +def test_get_all_returns_subtree_and_whole_tree(): + """``get_all`` returns the node at ``key``; an empty key returns the root.""" + metadata: dict = {} + global_counter_add(metadata, ["a"], 2) + global_counter_add(metadata, ["a", "child"], 3) + + subtree = global_counter_get_all(metadata, ["a"]) + assert subtree == {"value": 2, "children": {"child": {"value": 3, "children": {}}}} + + root = global_counter_get_all(metadata, []) + assert root["value"] == 0 + assert root["children"]["a"] == subtree + + +def test_get_all_returns_deep_copy(): + """Mutating the returned subtree must not affect the live counter tree.""" + metadata: dict = {} + global_counter_add(metadata, ["a", "child"], 3) + + subtree = global_counter_get_all(metadata, ["a"]) + subtree["value"] = 999 + subtree["children"]["child"]["value"] = 999 + subtree["children"]["extra"] = {"value": 1, "children": {}} + + assert global_counter_get(metadata, ["a"]) == 0 + assert global_counter_get(metadata, ["a", "child"]) == 3 + assert global_counter_get_all(metadata, ["a", "extra"]) is None + + +def test_concurrent_inc_yields_unique_values(): + """Parallel increments on one key never return duplicate values.""" + metadata: dict = {} + results: list[int] = [] + results_lock = threading.Lock() + calls_per_thread = 200 + + def worker(): + for _ in range(calls_per_thread): + value = global_counter_inc(metadata, ["shared"]) + with results_lock: + results.append(value) + + threads = [threading.Thread(target=worker) for _ in range(8)] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + + total = len(threads) * calls_per_thread + assert sorted(results) == list(range(total)) + assert global_counter_get(metadata, ["shared"]) == total