mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-08 03:10:24 +00:00
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
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 commit85bf32064d. * Reapply "fix: exclude stream replies from token accounting" This reverts commit6722c24dc5. * 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:
parent
3d487d8d45
commit
6b035c6553
22 changed files with 1282 additions and 55 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
51
reme/schema/token_usage.py
Normal file
51
reme/schema/token_usage.py
Normal 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),
|
||||
)
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
191
reme/utils/evaluation_interface.py
Normal file
191
reme/utils/evaluation_interface.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
142
tests/unit/test_evaluation_interface.py
Normal file
142
tests/unit/test_evaluation_interface.py
Normal 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,
|
||||
},
|
||||
}
|
||||
|
|
@ -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 ------------------------------------
|
||||
|
||||
|
||||
|
|
|
|||
227
tests/unit/test_token_usage.py
Normal file
227
tests/unit/test_token_usage.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue