feat(evaluation): track job calls and agent token usage in benchmarks (#406)
Some checks are pending
Pre-commit / run (ubuntu-latest) (push) Waiting to run
Tests ReMe / Unit Tests - py3.11 (push) Waiting to run
Tests ReMe / Unit Tests - py3.12 (push) Waiting to run
Tests ReMe / Unit Tests - py3.13 (push) Waiting to run
Windows Smoke / CLI smoke - py3.11 (push) Waiting to run

* 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 85bf32064d.

* Reapply "fix: exclude stream replies from token accounting"

This reverts commit 6722c24dc5.

* 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 <jiangniurou.xyf@dail-algo011164204033.ET135>
Co-authored-by: jinli.yl <jinli.yl@alibaba-inc.com>
This commit is contained in:
xyf2020 2026-08-04 11:42:18 +08:00 • committed by GitHub
parent 3d487d8d45
commit 6b035c6553
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
22 changed files with 1282 additions and 55 deletions

View file

@ -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

View file

@ -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

View file

@ -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")

View file

@ -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."""

View file

@ -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

View file

@ -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":

View file

@ -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():

View file

@ -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:

View file

@ -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}")

View file

@ -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:

View file

@ -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",
]

View file

@ -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),
)

View file

@ -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"

View file

@ -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",
]

View file

@ -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)

View file

@ -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)

View file

@ -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

View file

@ -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")

View file

@ -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,
},
}

View file

@ -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 ------------------------------------

View file

@ -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

View file

@ -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