From 4c2d0c7fd86daed75987921272d639e80f3b89b1 Mon Sep 17 00:00:00 2001 From: kshitijk4poor <82637225+kshitijk4poor@users.noreply.github.com> Date: Sat, 1 Aug 2026 14:52:14 +0530 Subject: [PATCH] refactor: dedupe fallback warning per model, drive pool-cleanup tests through real run() Review follow-up: - Warn once per (model, base_url) at the step-9 fallback via a module-level dedup set (established _WARNED_* idiom). The fallback result is deliberately never cached, so the un-deduped warning fired on every resolution - e.g. once per gateway message via the session-hygiene path. - Replace the three inline-mock pool-cleanup tests (which reproduced the try/except block against a MagicMock and passed even with the production code reverted) with a parametrized test that drives the real BatchRunner.run() with a patched Pool; drop the CPython stdlib signature change-detector test. - Add a once-per-model warning regression test; clean up dead imports. All tests verified to fail against pre-PR batch_runner.py/model_metadata.py and pass with the fix (mutation check). --- agent/model_metadata.py | 24 +++- tests/agent/test_model_metadata.py | 52 +++++++-- tests/test_batch_runner_durability.py | 157 +++++++++----------------- 3 files changed, 115 insertions(+), 118 deletions(-) diff --git a/agent/model_metadata.py b/agent/model_metadata.py index c7f16e8ea2..cd73cfc836 100644 --- a/agent/model_metadata.py +++ b/agent/model_metadata.py @@ -273,6 +273,11 @@ CONTEXT_PROBE_TIERS = [ # Default context length when no detection method succeeds. DEFAULT_FALLBACK_CONTEXT = CONTEXT_PROBE_TIERS[0] +# (model, base_url) pairs that already emitted the step-9 fallback warning. +# The fallback result itself is deliberately never cached, so without this +# the warning would repeat on every resolution for the same unknown model. +_FALLBACK_WARNED: set = set() + # Minimum context length required to run Hermes Agent. Models with fewer # tokens cannot maintain enough working memory for tool-calling workflows. # Sessions, model switches, and cron jobs should reject models below this. @@ -2773,12 +2778,19 @@ def get_model_context_length( # 9. Default fallback — log so small-context models (8K, 32K) don't # silently get 256K and cause hard-to-debug API failures. - logger.warning( - "Could not determine context length for model %r (base_url=%s) " - "— falling back to %s tokens. Set model.context_length in " - "config.yaml to override.", - model, base_url or "default", f"{DEFAULT_FALLBACK_CONTEXT:,}", - ) + # Warn once per (model, base_url): the fallback result is deliberately + # never cached (a wrong value must not freeze), so without dedup this + # would fire on every resolution — e.g. once per gateway message via + # the session-hygiene path. + _warn_key = (model, base_url or "") + if _warn_key not in _FALLBACK_WARNED: + _FALLBACK_WARNED.add(_warn_key) + logger.warning( + "Could not determine context length for model %r (base_url=%s) " + "— falling back to %s tokens. Set model.context_length in " + "config.yaml to override.", + model, base_url or "default", f"{DEFAULT_FALLBACK_CONTEXT:,}", + ) return DEFAULT_FALLBACK_CONTEXT diff --git a/tests/agent/test_model_metadata.py b/tests/agent/test_model_metadata.py index 6398383bc2..a63d79cfb7 100644 --- a/tests/agent/test_model_metadata.py +++ b/tests/agent/test_model_metadata.py @@ -1187,19 +1187,40 @@ class TestFallbackWarning: """When all 9 detection methods fail, the 10th fallback should log a warning so users with small-context models (8K, 32K) don't silently get 256K and hit hard-to-debug API context-length errors. + + The warning is deduped per (model, base_url) — the fallback result is + deliberately never cached, so without dedup it would repeat on every + resolution (e.g. once per gateway message via session hygiene). """ + @pytest.fixture(autouse=True) + def _reset_warned_set(self): + from agent import model_metadata as mm + mm._FALLBACK_WARNED.clear() + yield + mm._FALLBACK_WARNED.clear() + + @staticmethod + def _patch_all_lookups(): + from contextlib import ExitStack + stack = ExitStack() + for target, value in [ + ("agent.model_metadata.get_cached_context_length", None), + ("agent.model_metadata.fetch_model_metadata", {}), + ("agent.model_metadata.fetch_endpoint_model_metadata", {}), + ("agent.model_metadata._query_ollama_api_show", None), + ("agent.model_metadata._query_anthropic_context_length", None), + ("agent.model_metadata._endpoint_scoped_context_length", None), + ("agent.model_metadata._resolve_endpoint_context_length", None), + ("agent.models_dev.lookup_models_dev_context", None), + ]: + stack.enter_context(patch(target, return_value=value)) + return stack + def test_warning_emitted_on_fallback(self, caplog): import logging - with patch("agent.model_metadata.get_cached_context_length", return_value=None), \ - patch("agent.model_metadata.fetch_model_metadata", return_value={}), \ - patch("agent.model_metadata.fetch_endpoint_model_metadata", return_value={}), \ - patch("agent.model_metadata._query_ollama_api_show", return_value=None), \ - patch("agent.model_metadata._query_anthropic_context_length", return_value=None), \ - patch("agent.model_metadata._endpoint_scoped_context_length", return_value=None), \ - patch("agent.model_metadata._resolve_endpoint_context_length", return_value=None), \ - patch("agent.models_dev.lookup_models_dev_context", return_value=None): + with self._patch_all_lookups(): with caplog.at_level(logging.WARNING, logger="agent.model_metadata"): result = get_model_context_length( "totally-unknown-model-xyz", @@ -1211,6 +1232,21 @@ class TestFallbackWarning: assert any("totally-unknown-model-xyz" in r.getMessage() for r in warning_msgs) assert any("model.context_length" in r.getMessage() for r in warning_msgs) + def test_warning_fires_once_per_model(self, caplog): + """Repeated resolutions of the same unknown model warn only once.""" + import logging + + with self._patch_all_lookups(): + with caplog.at_level(logging.WARNING, logger="agent.model_metadata"): + for _ in range(3): + get_model_context_length("totally-unknown-model-xyz") + + fallback_warnings = [ + r for r in caplog.records + if r.levelno == logging.WARNING and "falling back" in r.getMessage() + ] + assert len(fallback_warnings) == 1 + def test_no_warning_when_cached(self, caplog): """No fallback warning when the context length is found in the cache.""" import logging diff --git a/tests/test_batch_runner_durability.py b/tests/test_batch_runner_durability.py index bece9da7b0..df1d9eb9cf 100644 --- a/tests/test_batch_runner_durability.py +++ b/tests/test_batch_runner_durability.py @@ -3,21 +3,26 @@ Verifies: 1. Trajectory entries are fsync'd to disk before the checkpoint marks them as completed (crash-between-write-and-sync safety). - 2. Pool.terminate() + pool.join() are called on KeyboardInterrupt and - Exception during batch execution (responsive worker shutdown). + 2. BatchRunner.run() calls pool.terminate() + pool.join() on + KeyboardInterrupt and Exception during batch execution (responsive + worker shutdown). CPython's Pool.join() takes no timeout parameter — + join(timeout=10) raises TypeError — so the tests also assert join() + is invoked with no arguments. """ import json import os import sys from pathlib import Path -from unittest.mock import MagicMock, patch, call +from unittest.mock import MagicMock, call, patch import pytest -# batch_runner uses relative imports, ensure project root is on path +# batch_runner is a root-level module (not part of an installed package), +# so make the repo root importable when tests run from elsewhere. sys.path.insert(0, str(Path(__file__).parent.parent)) +import batch_runner from batch_runner import BatchRunner, _process_batch_worker @@ -26,10 +31,9 @@ from batch_runner import BatchRunner, _process_batch_worker # ========================================================================= class TestTrajectoryWriteDurability: - """Verify that trajectory entries are flushed and fsync'd before the - checkpoint marks them as completed. + """Verify that trajectory entries are flushed and fsync'd to disk. - Without fsync, a crash between the write and the disk sync would leave + Without fsync, a crash between the write and the disk sync could leave the checkpoint claiming completion with no trajectory data on disk. """ @@ -52,14 +56,9 @@ class TestTrajectoryWriteDurability: # Intercept os.fsync to record calls fsync_calls = [] - original_fsync = os.fsync + monkeypatch.setattr("os.fsync", lambda fd: fsync_calls.append(fd)) - def mock_fsync(fd): - fsync_calls.append(fd) - - monkeypatch.setattr("os.fsync", mock_fsync) - - result = _process_batch_worker( + _process_batch_worker( ( 1, [(0, {"prompt": "hi"})], @@ -87,101 +86,51 @@ class TestTrajectoryWriteDurability: # ========================================================================= -# Pool cleanup on interruption / exception +# Pool cleanup on interruption / exception — drives the REAL run() # ========================================================================= -class TestPoolCleanupOnInterruption: - """Verify that pool.terminate() + pool.join() are called when a - KeyboardInterrupt or Exception occurs during batch execution. +def _make_runner(tmp_path, monkeypatch): + """Build a minimal real BatchRunner against a 1-line tmp dataset.""" + dataset = tmp_path / "dataset.jsonl" + dataset.write_text(json.dumps({"prompt": "hi"}) + "\n", encoding="utf-8") + # BatchRunner writes to Path("data")/run_name relative to cwd. + monkeypatch.chdir(tmp_path) + return BatchRunner( + dataset_file=str(dataset), + batch_size=1, + run_name="pool-cleanup-test", + num_workers=1, + ) - CPython's multiprocessing.pool.Pool.join() does NOT accept a timeout - parameter — calling pool.join(timeout=10) raises TypeError. The fix - uses pool.terminate() followed by pool.join() (no timeout), which is - the correct shutdown pattern. + +def _make_failing_pool(exc): + """Context-manager mock whose pool raises `exc` from imap_unordered.""" + pool = MagicMock() + pool.imap_unordered.side_effect = exc + pool_cm = MagicMock() + pool_cm.__enter__ = MagicMock(return_value=pool) + pool_cm.__exit__ = MagicMock(return_value=False) + return pool, pool_cm + + +class TestPoolCleanupOnInterruption: + """Drive the real BatchRunner.run() with a patched Pool and verify the + cleanup contract: terminate() + join() (join with NO timeout argument — + CPython's Pool.join signature is (self), so join(timeout=10) would + raise TypeError). """ - def test_pool_terminate_called_on_exception(self, tmp_path, monkeypatch): - """When pool.imap_unordered raises an exception, pool.terminate() - and pool.join() must be called for clean worker shutdown. + @pytest.mark.parametrize("exc_type", [KeyboardInterrupt, RuntimeError]) + def test_run_terminates_and_joins_pool(self, tmp_path, monkeypatch, exc_type): + runner = _make_runner(tmp_path, monkeypatch) + pool, pool_cm = _make_failing_pool(exc_type("boom")) - We simulate the relevant slice of run()'s try/except block with a - mock pool to verify the cleanup contract. - """ - mock_pool = MagicMock() - mock_pool.imap_unordered.side_effect = RuntimeError("worker exploded") + with patch.object(batch_runner, "Pool", return_value=pool_cm): + with pytest.raises(exc_type): + runner.run() - # Reproduce the exception-handling block from batch_runner.run() - with pytest.raises(RuntimeError, match="worker exploded"): - try: - for result in mock_pool.imap_unordered(None, []): - pass - except KeyboardInterrupt: - mock_pool.terminate() - mock_pool.join() - raise - except Exception: - mock_pool.terminate() - mock_pool.join() - raise - - mock_pool.terminate.assert_called_once() - mock_pool.join.assert_called_once_with() - - def test_pool_terminate_called_on_keyboard_interrupt(self, tmp_path, monkeypatch): - """When pool.imap_unordered is interrupted (Ctrl+C), pool.terminate() - and pool.join() must be called for responsive shutdown.""" - mock_pool = MagicMock() - mock_pool.imap_unordered.side_effect = KeyboardInterrupt() - - with pytest.raises(KeyboardInterrupt): - try: - for result in mock_pool.imap_unordered(None, []): - pass - except KeyboardInterrupt: - mock_pool.terminate() - mock_pool.join() - raise - except Exception: - mock_pool.terminate() - mock_pool.join() - raise - - mock_pool.terminate.assert_called_once() - mock_pool.join.assert_called_once_with() - - def test_pool_join_called_without_timeout(self, tmp_path): - """Pool.join() must NOT be called with a timeout argument — - CPython's Pool.join signature is (self), so join(timeout=10) - would raise TypeError.""" - mock_pool = MagicMock() - mock_pool.imap_unordered.side_effect = RuntimeError("boom") - - with pytest.raises(RuntimeError): - try: - for result in mock_pool.imap_unordered(None, []): - pass - except Exception: - mock_pool.terminate() - mock_pool.join() - raise - - # The join call must have no positional/keyword timeout argument - join_call = mock_pool.join.call_args - assert join_call == call(), ( - f"pool.join() called with unexpected args: {join_call}" - ) - - def test_real_pool_join_accepts_no_timeout(self): - """Integration check: a real multiprocessing.Pool's join() must not - accept a timeout kwarg. This guards against re-introducing - pool.join(timeout=10), which raises TypeError on CPython. - """ - import inspect - import multiprocessing.pool - - sig = inspect.signature(multiprocessing.pool.Pool.join) - params = list(sig.parameters.keys()) - # The only parameter should be 'self' — no 'timeout' - assert "timeout" not in params, ( - f"Pool.join has unexpected parameters: {params}" + pool.terminate.assert_called_once() + # join() must be called with no positional/keyword arguments. + assert pool.join.call_args_list == [call()], ( + f"pool.join() called with unexpected args: {pool.join.call_args_list}" )