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:
酱牛肉 2026-07-29 22:45:53 +08:00
parent 550317c3bf
commit f10a1c9935
9 changed files with 286 additions and 10 deletions

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(BackgroundJob)
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,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:

View file

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

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(StreamJob)
merged = {**self.kwargs, **kwargs}
context = RuntimeContext(**merged)
try:

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

View file

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

View file

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

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

View file

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