395 lines
14 KiB
Python
395 lines
14 KiB
Python
"""Tests for ``EvoScientist.middleware.configurable_model``.
|
|
|
|
Verifies that the middleware reads ``model`` / ``model_provider`` from
|
|
the active ``RunnableConfig.configurable`` (via ``langgraph.config.get_config``)
|
|
and overrides ``request.model`` accordingly, without breaking the no-override
|
|
pass-through path or the per-instance cache.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from contextlib import contextmanager
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
from EvoScientist.middleware.configurable_model import (
|
|
ConfigurableModelMiddleware,
|
|
_read_model_override,
|
|
)
|
|
from tests.conftest import run_async as _run
|
|
|
|
|
|
@contextmanager
|
|
def _patched_config(configurable: dict | object | None):
|
|
"""Patch ``langgraph.config.get_config`` to return a controlled value.
|
|
|
|
Pass ``None`` to simulate "outside a runnable context" (raises
|
|
``RuntimeError`` like the real ``get_config()`` does).
|
|
Pass a dict for ``configurable`` to expose just that key.
|
|
Pass any other object to simulate a malformed config.
|
|
"""
|
|
import langgraph.config as _lg_cfg
|
|
|
|
if configurable is None:
|
|
# Simulate get_config raising outside a runnable context.
|
|
with patch.object(
|
|
_lg_cfg,
|
|
"get_config",
|
|
side_effect=RuntimeError("Called get_config outside of a runnable context"),
|
|
):
|
|
yield
|
|
elif isinstance(configurable, dict):
|
|
with patch.object(
|
|
_lg_cfg,
|
|
"get_config",
|
|
return_value={"configurable": configurable},
|
|
):
|
|
yield
|
|
else:
|
|
with patch.object(
|
|
_lg_cfg,
|
|
"get_config",
|
|
return_value=configurable,
|
|
):
|
|
yield
|
|
|
|
|
|
def _make_request():
|
|
"""Build a minimal ``ModelRequest`` stub.
|
|
|
|
``request.override(model=...)`` returns a new request whose ``model``
|
|
field reflects the override.
|
|
"""
|
|
req = MagicMock()
|
|
|
|
def _override(**kwargs):
|
|
new = MagicMock()
|
|
new.model = kwargs.get("model", req.model)
|
|
return new
|
|
|
|
req.override = MagicMock(side_effect=_override)
|
|
return req
|
|
|
|
|
|
# =============================================================================
|
|
# 1. _read_model_override — input parsing
|
|
# =============================================================================
|
|
|
|
|
|
class TestReadModelOverride:
|
|
"""Verify the helper that pulls (model, provider) from active config."""
|
|
|
|
def test_returns_override_when_both_present(self):
|
|
with _patched_config({"model": "gpt-5", "model_provider": "openai"}):
|
|
assert _read_model_override() == ("gpt-5", "openai")
|
|
|
|
def test_provider_optional(self):
|
|
with _patched_config({"model": "claude-haiku-4-5"}):
|
|
assert _read_model_override() == ("claude-haiku-4-5", None)
|
|
|
|
def test_no_configurable_key(self):
|
|
with _patched_config({}):
|
|
assert _read_model_override() == (None, None)
|
|
|
|
def test_outside_runnable_context(self):
|
|
"""``get_config`` raises outside a runnable — middleware must no-op."""
|
|
with _patched_config(None):
|
|
assert _read_model_override() == (None, None)
|
|
|
|
def test_empty_string_treated_as_absent(self):
|
|
with _patched_config({"model": "", "model_provider": ""}):
|
|
assert _read_model_override() == (None, None)
|
|
|
|
def test_non_string_ignored(self):
|
|
with _patched_config({"model": 42, "model_provider": object()}):
|
|
assert _read_model_override() == (None, None)
|
|
|
|
def test_non_dict_configurable_safe(self):
|
|
with _patched_config({"configurable": "not-a-dict"}):
|
|
# Inner ``configurable`` is the wrong type — patched_config
|
|
# wraps it again so we end up with {"configurable": {"configurable": "..."}}
|
|
# which has no model/model_provider keys → no override.
|
|
assert _read_model_override() == (None, None)
|
|
|
|
def test_non_dict_config_safe(self):
|
|
with _patched_config("garbage"):
|
|
assert _read_model_override() == (None, None)
|
|
|
|
|
|
# =============================================================================
|
|
# 2. ConfigurableModelMiddleware — pass-through behavior
|
|
# =============================================================================
|
|
|
|
|
|
class TestPassThrough:
|
|
"""When no override is present, the middleware must not touch the request."""
|
|
|
|
def test_sync_no_override_passes_request_unchanged(self):
|
|
mw = ConfigurableModelMiddleware()
|
|
req = _make_request()
|
|
sentinel = object()
|
|
handler = MagicMock(return_value=sentinel)
|
|
with _patched_config({}):
|
|
result = mw.wrap_model_call(req, handler)
|
|
assert result is sentinel
|
|
handler.assert_called_once_with(req)
|
|
req.override.assert_not_called()
|
|
|
|
def test_async_no_override_passes_request_unchanged(self):
|
|
mw = ConfigurableModelMiddleware()
|
|
req = _make_request()
|
|
|
|
async def handler(r):
|
|
assert r is req
|
|
return "ok"
|
|
|
|
with _patched_config({}):
|
|
result = _run(mw.awrap_model_call(req, handler))
|
|
assert result == "ok"
|
|
req.override.assert_not_called()
|
|
|
|
def test_outside_runnable_context_passes_through(self):
|
|
mw = ConfigurableModelMiddleware()
|
|
req = _make_request()
|
|
handler = MagicMock(return_value="ok")
|
|
with _patched_config(None):
|
|
mw.wrap_model_call(req, handler)
|
|
handler.assert_called_once_with(req)
|
|
|
|
|
|
# =============================================================================
|
|
# 3. ConfigurableModelMiddleware — override behavior
|
|
# =============================================================================
|
|
|
|
|
|
class TestModelOverride:
|
|
"""When override present, middleware resolves model and overrides request."""
|
|
|
|
def test_sync_override_calls_get_chat_model_and_overrides(self):
|
|
mw = ConfigurableModelMiddleware()
|
|
req = _make_request()
|
|
new_model = MagicMock(name="resolved_chat_model")
|
|
handler = MagicMock(return_value="response")
|
|
|
|
with (
|
|
_patched_config({"model": "gpt-5", "model_provider": "openai"}),
|
|
patch(
|
|
"EvoScientist.llm.get_chat_model", return_value=new_model
|
|
) as mock_get,
|
|
):
|
|
mw.wrap_model_call(req, handler)
|
|
|
|
mock_get.assert_called_once_with(model="gpt-5", provider="openai")
|
|
req.override.assert_called_once_with(model=new_model)
|
|
# Handler must receive the OVERRIDDEN request, not the original.
|
|
called_with = handler.call_args[0][0]
|
|
assert called_with is not req
|
|
assert called_with.model is new_model
|
|
|
|
def test_async_override_path_parity(self):
|
|
mw = ConfigurableModelMiddleware()
|
|
req = _make_request()
|
|
new_model = MagicMock()
|
|
|
|
async def handler(r):
|
|
assert r.model is new_model
|
|
return "ok"
|
|
|
|
with (
|
|
_patched_config(
|
|
{"model": "claude-opus-4-8", "model_provider": "anthropic"}
|
|
),
|
|
patch(
|
|
"EvoScientist.llm.get_chat_model", return_value=new_model
|
|
) as mock_get,
|
|
):
|
|
result = _run(mw.awrap_model_call(req, handler))
|
|
|
|
assert result == "ok"
|
|
mock_get.assert_called_once_with(model="claude-opus-4-8", provider="anthropic")
|
|
|
|
def test_provider_omitted_passed_as_none(self):
|
|
mw = ConfigurableModelMiddleware()
|
|
req = _make_request()
|
|
new_model = MagicMock()
|
|
handler = MagicMock(return_value="ok")
|
|
|
|
with (
|
|
_patched_config({"model": "gpt-5"}),
|
|
patch(
|
|
"EvoScientist.llm.get_chat_model", return_value=new_model
|
|
) as mock_get,
|
|
):
|
|
mw.wrap_model_call(req, handler)
|
|
|
|
mock_get.assert_called_once_with(model="gpt-5", provider=None)
|
|
|
|
|
|
# =============================================================================
|
|
# 4. ConfigurableModelMiddleware — caching
|
|
# =============================================================================
|
|
|
|
|
|
class TestCache:
|
|
"""Two consecutive calls with same (model, provider) should hit cache."""
|
|
|
|
def test_cache_hit_avoids_second_get_chat_model(self):
|
|
mw = ConfigurableModelMiddleware()
|
|
req1 = _make_request()
|
|
req2 = _make_request()
|
|
new_model = MagicMock()
|
|
handler = MagicMock(return_value="ok")
|
|
|
|
with (
|
|
_patched_config({"model": "gpt-5", "model_provider": "openai"}),
|
|
patch(
|
|
"EvoScientist.llm.get_chat_model", return_value=new_model
|
|
) as mock_get,
|
|
):
|
|
mw.wrap_model_call(req1, handler)
|
|
mw.wrap_model_call(req2, handler)
|
|
|
|
# First call resolves via factory, second hits the cache.
|
|
assert mock_get.call_count == 1
|
|
|
|
def test_cache_miss_on_different_provider(self):
|
|
mw = ConfigurableModelMiddleware()
|
|
req = _make_request()
|
|
handler = MagicMock(return_value="ok")
|
|
|
|
with patch(
|
|
"EvoScientist.llm.get_chat_model", side_effect=[MagicMock(), MagicMock()]
|
|
) as mock_get:
|
|
with _patched_config(
|
|
{"model": "claude-sonnet-4-6", "model_provider": "anthropic"}
|
|
):
|
|
mw.wrap_model_call(req, handler)
|
|
with _patched_config(
|
|
{"model": "claude-sonnet-4-6", "model_provider": "custom-anthropic"}
|
|
):
|
|
mw.wrap_model_call(req, handler)
|
|
|
|
assert mock_get.call_count == 2
|
|
|
|
def test_independent_instances_have_independent_caches(self):
|
|
mw_a = ConfigurableModelMiddleware()
|
|
mw_b = ConfigurableModelMiddleware()
|
|
req = _make_request()
|
|
handler = MagicMock(return_value="ok")
|
|
|
|
with (
|
|
_patched_config({"model": "gpt-5", "model_provider": "openai"}),
|
|
patch(
|
|
"EvoScientist.llm.get_chat_model",
|
|
side_effect=[MagicMock(), MagicMock()],
|
|
) as mock_get,
|
|
):
|
|
mw_a.wrap_model_call(req, handler)
|
|
mw_b.wrap_model_call(req, handler)
|
|
|
|
# Different instances must each resolve once.
|
|
assert mock_get.call_count == 2
|
|
|
|
|
|
# =============================================================================
|
|
# 5. Resilience — get_chat_model raising
|
|
# =============================================================================
|
|
|
|
|
|
class TestResolveFailure:
|
|
"""If get_chat_model raises, middleware must fall back to original model."""
|
|
|
|
def test_sync_falls_back_when_resolve_raises(self):
|
|
mw = ConfigurableModelMiddleware()
|
|
req = _make_request()
|
|
handler = MagicMock(return_value="ok")
|
|
|
|
with (
|
|
_patched_config({"model": "doesnotexist", "model_provider": "openai"}),
|
|
patch(
|
|
"EvoScientist.llm.get_chat_model",
|
|
side_effect=ValueError("unknown model"),
|
|
),
|
|
):
|
|
mw.wrap_model_call(req, handler)
|
|
|
|
# Override never happened — handler called with original request.
|
|
handler.assert_called_once_with(req)
|
|
req.override.assert_not_called()
|
|
|
|
def test_async_falls_back_when_resolve_raises(self):
|
|
mw = ConfigurableModelMiddleware()
|
|
req = _make_request()
|
|
|
|
called = []
|
|
|
|
async def handler(r):
|
|
called.append(r)
|
|
return "ok"
|
|
|
|
with (
|
|
_patched_config({"model": "doesnotexist", "model_provider": "openai"}),
|
|
patch(
|
|
"EvoScientist.llm.get_chat_model",
|
|
side_effect=ValueError("unknown model"),
|
|
),
|
|
):
|
|
result = _run(mw.awrap_model_call(req, handler))
|
|
|
|
assert result == "ok"
|
|
assert called == [req]
|
|
|
|
|
|
# =============================================================================
|
|
# 6. Integration — real langgraph contextvar (no get_config mock)
|
|
# =============================================================================
|
|
|
|
|
|
class TestRunnableContextVarIntegration:
|
|
"""Set the actual ``var_child_runnable_config`` contextvar that LangGraph
|
|
populates per node, then verify the middleware reads through to it.
|
|
|
|
This catches breakage of the ``langgraph.config.get_config()`` contract
|
|
that pure-mock tests would miss (e.g. if get_config is moved to a
|
|
different module, or the contextvar mechanism changes).
|
|
"""
|
|
|
|
def test_real_contextvar_drives_override(self):
|
|
"""Without mocking get_config, set the contextvar and verify override."""
|
|
from langchain_core.runnables.config import var_child_runnable_config
|
|
|
|
mw = ConfigurableModelMiddleware()
|
|
req = _make_request()
|
|
new_model = MagicMock(name="resolved")
|
|
handler = MagicMock(return_value="ok")
|
|
|
|
token = var_child_runnable_config.set(
|
|
{"configurable": {"model": "gpt-5.5", "model_provider": "openai"}}
|
|
)
|
|
try:
|
|
with patch(
|
|
"EvoScientist.llm.get_chat_model", return_value=new_model
|
|
) as mock_get:
|
|
mw.wrap_model_call(req, handler)
|
|
finally:
|
|
var_child_runnable_config.reset(token)
|
|
|
|
mock_get.assert_called_once_with(model="gpt-5.5", provider="openai")
|
|
req.override.assert_called_once_with(model=new_model)
|
|
|
|
def test_real_contextvar_unset_passes_through(self):
|
|
"""When no contextvar is set, get_config() raises → no override."""
|
|
from langchain_core.runnables.config import var_child_runnable_config
|
|
|
|
# Defensive: ensure no leftover contextvar from another test.
|
|
token = var_child_runnable_config.set(None)
|
|
try:
|
|
mw = ConfigurableModelMiddleware()
|
|
req = _make_request()
|
|
handler = MagicMock(return_value="ok")
|
|
mw.wrap_model_call(req, handler)
|
|
finally:
|
|
var_child_runnable_config.reset(token)
|
|
|
|
handler.assert_called_once_with(req)
|
|
req.override.assert_not_called()
|