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).
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}"
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user