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:
kshitijk4poor
2026-08-01 14:52:14 +05:30
committed by kshitij
parent a1ff62a139
commit 4c2d0c7fd8
3 changed files with 115 additions and 118 deletions

View File

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

View File

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

View File

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