From f10a1c9935f133c43705a9cab8932eeae3f54bf5 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=85=B1=E7=89=9B=E8=82=89?= Date: Wed, 29 Jul 2026 22:45:53 +0800 Subject: [PATCH] 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 --- reme/components/job/background_job.py | 1 + reme/components/job/base_job.py | 18 +++ reme/components/job/cron_job.py | 1 + reme/components/job/stream_job.py | 1 + reme/steps/benchmark/base/agentic_answer.py | 4 +- reme/utils/__init__.py | 7 +- reme/utils/counter.py | 77 ++++++++++++- tests/unit/test_job.py | 68 +++++++++++ tests/unit/test_utils.py | 119 ++++++++++++++++++++ 9 files changed, 286 insertions(+), 10 deletions(-) diff --git a/reme/components/job/background_job.py b/reme/components/job/background_job.py index fc62662a..2f124bdc 100644 --- a/reme/components/job/background_job.py +++ b/reme/components/job/background_job.py @@ -123,6 +123,7 @@ class BackgroundJob(BaseJob): async def __call__(self, **kwargs) -> Response: """Default body: run steps in order; errors propagate to supervisor.""" + self._record_call(BackgroundJob) merged = {**self.kwargs, **kwargs} context = RuntimeContext(stop_event=self._stop_event, **merged) for step in self._build_steps(): diff --git a/reme/components/job/base_job.py b/reme/components/job/base_job.py index 834e0125..2615cd7c 100644 --- a/reme/components/job/base_job.py +++ b/reme/components/job/base_job.py @@ -7,6 +7,7 @@ from ..component_registry import R from ..runtime_context import RuntimeContext from ...enumeration import ComponentEnum from ...schema import ComponentConfig, Response +from ...utils import global_counter_inc if TYPE_CHECKING: from ...steps import BaseStep @@ -57,8 +58,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: diff --git a/reme/components/job/cron_job.py b/reme/components/job/cron_job.py index 595653e3..49a117d6 100644 --- a/reme/components/job/cron_job.py +++ b/reme/components/job/cron_job.py @@ -47,6 +47,7 @@ class CronJob(BackgroundJob): if self._stop_event.is_set(): break try: + self._record_call(CronJob) await self._execute_steps() except Exception as exc: self.logger.exception(f"Cron job '{self.name}' failed: {exc}") diff --git a/reme/components/job/stream_job.py b/reme/components/job/stream_job.py index 8eb93f92..6de561a0 100644 --- a/reme/components/job/stream_job.py +++ b/reme/components/job/stream_job.py @@ -12,6 +12,7 @@ class StreamJob(BaseJob): async def __call__(self, **kwargs) -> None: """Run steps; emit failures as ERROR chunks, then a terminal DONE marker.""" + self._record_call(StreamJob) merged = {**self.kwargs, **kwargs} context = RuntimeContext(**merged) try: diff --git a/reme/steps/benchmark/base/agentic_answer.py b/reme/steps/benchmark/base/agentic_answer.py index 1423b5e7..cf69bbc5 100644 --- a/reme/steps/benchmark/base/agentic_answer.py +++ b/reme/steps/benchmark/base/agentic_answer.py @@ -5,7 +5,7 @@ import os from ...base_step import BaseStep from ...index._dedup import _ToolContextDedupMixin from ....enumeration import ChunkEnum -from ....utils.counter import global_counter_next +from ....utils.counter import global_counter_inc class BaseAgenticAnswerStep(BaseStep): @@ -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" diff --git a/reme/utils/__init__.py b/reme/utils/__init__.py index 69a16ed6..6fbbe767 100644 --- a/reme/utils/__init__.py +++ b/reme/utils/__init__.py @@ -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", ] diff --git a/reme/utils/counter.py b/reme/utils/counter.py index 204d9cf6..cd4c53d0 100644 --- a/reme/utils/counter.py +++ b/reme/utils/counter.py @@ -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) diff --git a/tests/unit/test_job.py b/tests/unit/test_job.py index 06b13998..3531fcd8 100644 --- a/tests/unit/test_job.py +++ b/tests/unit/test_job.py @@ -18,6 +18,7 @@ from reme.components.job.cron_job import CronJob from reme.components.job.stream_job import StreamJob from reme.components.job import cron_job as cron_job_module from reme.schema import ComponentConfig +from reme.utils import global_counter_get # -- helpers ------------------------------------------------------------------ @@ -132,6 +133,73 @@ def test_stream_job_merges_config_kwargs_into_context(): asyncio.run(run()) +# -- Job call counters ------------------------------------------------------- + + +def test_base_job_records_calls_by_name(): + async def run(): + app_context = SimpleNamespace(metadata={}) + job = BaseJob(name="search", app_context=app_context) + + await job() + await job() + + assert global_counter_get(app_context.metadata, ["__job_counter", "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 ------------------------------------ diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index d8bdcefb..34eff188 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -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