mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-09 03:20:54 +00:00
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
This commit is contained in:
parent
550317c3bf
commit
f10a1c9935
9 changed files with 286 additions and 10 deletions
|
|
@ -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(BackgroundJob)
|
||||
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,25 @@ 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, entry_class: type["BaseJob"]) -> None:
|
||||
"""Increment this job's application-lifetime call counter.
|
||||
|
||||
The counter path groups calls by the built-in job implementation that
|
||||
accepted the call, then preserves any project-specific subclasses
|
||||
between that implementation and the concrete job class.
|
||||
"""
|
||||
metadata = getattr(self.app_context, "metadata", None)
|
||||
if not isinstance(metadata, dict):
|
||||
return
|
||||
|
||||
mro = type(self).mro()
|
||||
entry_index = mro.index(entry_class)
|
||||
subclass_path = [cls.__name__ for cls in mro[:entry_index]]
|
||||
global_counter_inc(metadata, ["__job_counter", entry_class.__name__, *subclass_path, self.name])
|
||||
|
||||
async def __call__(self, **kwargs) -> Response:
|
||||
"""Run all steps in order, capturing any failure into the response."""
|
||||
self._record_call(BaseJob)
|
||||
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(CronJob)
|
||||
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(StreamJob)
|
||||
merged = {**self.kwargs, **kwargs}
|
||||
context = RuntimeContext(**merged)
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
@ -46,7 +46,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,7 @@ 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_get, global_counter_get_all, global_counter_inc
|
||||
|
||||
__all__ = [
|
||||
"hash_text",
|
||||
|
|
@ -38,5 +38,8 @@ __all__ = [
|
|||
"batch_cosine_similarity",
|
||||
"estimate_token_count",
|
||||
"AsStateHandler",
|
||||
"global_counter_next",
|
||||
"global_counter_add",
|
||||
"global_counter_get",
|
||||
"global_counter_get_all",
|
||||
"global_counter_inc",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""Thread-safe monotonic counter tree utility for shared application state."""
|
||||
|
||||
import copy
|
||||
import threading
|
||||
from typing import Any
|
||||
|
||||
|
|
@ -7,12 +8,13 @@ COUNTER_TREE_KEY = "_counter_tree"
|
|||
COUNTER_LOCK_KEY = "_counter_tree_lock"
|
||||
|
||||
|
||||
def global_counter_next(metadata: dict[str, Any], key: list[str]) -> int:
|
||||
"""Return the next monotonic value for ``key``, starting at 1.
|
||||
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
|
||||
|
|
@ -40,6 +42,69 @@ def global_counter_next(metadata: dict[str, Any], key: list[str]) -> int:
|
|||
tmp = {"value": 0, "children": {}}
|
||||
node["children"][part] = tmp
|
||||
node = tmp
|
||||
res = node["value"] + 1
|
||||
node["value"] = res
|
||||
res = node["value"]
|
||||
node["value"] = res + val
|
||||
return res
|
||||
|
||||
|
||||
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 = metadata.get(COUNTER_LOCK_KEY)
|
||||
if lock is None:
|
||||
lock = threading.Lock()
|
||||
metadata[COUNTER_LOCK_KEY] = lock
|
||||
|
||||
with lock:
|
||||
tree = metadata.get(COUNTER_TREE_KEY)
|
||||
if tree is None:
|
||||
return 0
|
||||
|
||||
node: dict[str, Any] | None = tree
|
||||
for part in key:
|
||||
assert isinstance(part, str)
|
||||
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 = metadata.get(COUNTER_LOCK_KEY)
|
||||
if lock is None:
|
||||
lock = threading.Lock()
|
||||
metadata[COUNTER_LOCK_KEY] = lock
|
||||
|
||||
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)
|
||||
|
|
|
|||
|
|
@ -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", "BaseJob", "search"]) == 2
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
def test_stream_job_records_concrete_subclass_path():
|
||||
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", "StreamJob", "ProjectStreamJob", "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", "BackgroundJob", "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", "CronJob", "nightly"]) == 1
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
# -- BaseJob._start requires app_context ------------------------------------
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -2,11 +2,19 @@
|
|||
|
||||
import asyncio
|
||||
import sys
|
||||
import threading
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from reme.utils import common_utils
|
||||
from reme.utils.counter import (
|
||||
COUNTER_TREE_KEY,
|
||||
global_counter_add,
|
||||
global_counter_get,
|
||||
global_counter_get_all,
|
||||
global_counter_inc,
|
||||
)
|
||||
from reme.utils.similarity_utils import batch_cosine_similarity, cosine_similarity
|
||||
|
||||
|
||||
|
|
@ -65,3 +73,114 @@ 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_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