feat(moa): per-reference-model max_tokens override

MoA reference_max_tokens is preset-level — one cap for all reference
models. When mixing a verbose model with a terse one, a single cap is
either too tight for the terse model or too loose for the verbose one.

Now each reference slot can optionally carry its own max_tokens:

  reference_models:
    - provider: openrouter
      model: deepseek/deepseek-v4-pro
      max_tokens: ***        # per-slot cap, overrides preset-level
    - provider: openai-codex
      model: gpt-5.5
      # no max_tokens → falls back to preset-level reference_max_tokens

_clean_slot (moa_config.py) preserves an optional max_tokens field on
the slot dict, coerced via _coerce_int_or_none. _run_reference
(moa_loop.py) reads slot-level max_tokens first, falling back to the
preset-level cap passed by the caller. Slots without the field are
unaffected — backward compatible.

Type hints on slot-handling functions updated from dict[str, str] to
dict[str, Any] to reflect the now-heterogeneous slot shape.
This commit is contained in:
Rain 2026-07-07 18:18:23 +02:00 • committed by Teknium
parent ead9d7b256
commit bc7212cf93
5 changed files with 178 additions and 5 deletions

View file

@ -0,0 +1,82 @@
"""Tests for per-slot max_tokens in MoA reference calls.
Verifies that a ``max_tokens`` field on a reference slot dict takes
precedence over the preset-level ``reference_max_tokens``, and that
slot-level max_tokens=None falls back to the preset-level cap.
"""
from __future__ import annotations
from unittest.mock import MagicMock, patch
import pytest
class TestRunReferenceSlotMaxTokens:
"""_run_reference should prefer slot-level max_tokens over preset-level."""
def test_slot_max_tokens_overrides_preset_level(self):
"""When slot has max_tokens, it overrides the preset-level cap."""
from agent.moa_loop import _run_reference
captured_kwargs: dict = {}
def fake_call_llm(**kwargs):
captured_kwargs.update(kwargs)
mock_resp = MagicMock()
mock_resp.choices = [MagicMock(message=MagicMock(content="advice"))]
mock_resp.usage = None
return mock_resp
slot = {"provider": "openrouter", "model": "deepseek/deepseek-v4-pro", "max_tokens": 600}
with patch("agent.moa_loop._slot_runtime", return_value={"provider": "openrouter", "model": "deepseek/deepseek-v4-pro"}), \
patch("agent.moa_loop.call_llm", side_effect=fake_call_llm), \
patch("agent.moa_loop._maybe_apply_moa_cache_control", side_effect=lambda msgs, rt: msgs):
_run_reference(slot, [{"role": "user", "content": "hi"}], max_tokens=2000)
assert captured_kwargs.get("max_tokens") == 600
def test_slot_max_tokens_absent_falls_back_to_preset(self):
"""When slot has no max_tokens, the preset-level cap is used."""
from agent.moa_loop import _run_reference
captured_kwargs: dict = {}
def fake_call_llm(**kwargs):
captured_kwargs.update(kwargs)
mock_resp = MagicMock()
mock_resp.choices = [MagicMock(message=MagicMock(content="advice"))]
mock_resp.usage = None
return mock_resp
slot = {"provider": "openrouter", "model": "deepseek/deepseek-v4-pro"}
with patch("agent.moa_loop._slot_runtime", return_value={"provider": "openrouter", "model": "deepseek/deepseek-v4-pro"}), \
patch("agent.moa_loop.call_llm", side_effect=fake_call_llm), \
patch("agent.moa_loop._maybe_apply_moa_cache_control", side_effect=lambda msgs, rt: msgs):
_run_reference(slot, [{"role": "user", "content": "hi"}], max_tokens=2000)
assert captured_kwargs.get("max_tokens") == 2000
def test_both_none_means_uncapped(self):
"""When neither slot nor preset has max_tokens, it's None (uncapped)."""
from agent.moa_loop import _run_reference
captured_kwargs: dict = {}
def fake_call_llm(**kwargs):
captured_kwargs.update(kwargs)
mock_resp = MagicMock()
mock_resp.choices = [MagicMock(message=MagicMock(content="advice"))]
mock_resp.usage = None
return mock_resp
slot = {"provider": "openrouter", "model": "deepseek/deepseek-v4-pro"}
with patch("agent.moa_loop._slot_runtime", return_value={"provider": "openrouter", "model": "deepseek/deepseek-v4-pro"}), \
patch("agent.moa_loop.call_llm", side_effect=fake_call_llm), \
patch("agent.moa_loop._maybe_apply_moa_cache_control", side_effect=lambda msgs, rt: msgs):
_run_reference(slot, [{"role": "user", "content": "hi"}], max_tokens=None)
assert captured_kwargs.get("max_tokens") is None

View file

@ -461,3 +461,79 @@ def test_validate_moa_payload_rejects_non_dict():
assert validate_moa_payload(None)
assert validate_moa_payload([1, 2])
assert validate_moa_payload({"presets": {"p": "not-a-dict"}})
# ── Per-slot max_tokens ────────────────────────────────────────────────────
def test_slot_max_tokens_preserved():
"""A max_tokens field on a reference slot survives normalization."""
cfg = normalize_moa_config(
{
"presets": {
"p": {
"reference_models": [
{"provider": "openrouter", "model": "deepseek/deepseek-v4-pro", "max_tokens": 600},
{"provider": "openai-codex", "model": "gpt-5.5"},
],
"aggregator": {"provider": "openrouter", "model": "anthropic/claude-opus-4.8"},
}
}
}
)
refs = cfg["presets"]["p"]["reference_models"]
assert refs[0] == {"provider": "openrouter", "model": "deepseek/deepseek-v4-pro", "max_tokens": 600}
assert refs[1] == {"provider": "openai-codex", "model": "gpt-5.5"}
def test_slot_max_tokens_coerced_from_string():
"""Hand-edited YAML string '600' coerces to int on a slot."""
cfg = normalize_moa_config(
{
"presets": {
"p": {
"reference_models": [
{"provider": "openrouter", "model": "deepseek/deepseek-v4-pro", "max_tokens": "600"},
],
}
}
}
)
refs = cfg["presets"]["p"]["reference_models"]
assert refs[0]["max_tokens"] == 600
def test_slot_max_tokens_invalid_dropped():
"""Non-positive / non-numeric slot max_tokens is dropped (slot kept)."""
for bad in (0, -5, "abc", "", None):
cfg = normalize_moa_config(
{
"presets": {
"p": {
"reference_models": [
{"provider": "openrouter", "model": "deepseek/deepseek-v4-pro", "max_tokens": bad},
],
}
}
}
)
ref = cfg["presets"]["p"]["reference_models"][0]
assert "max_tokens" not in ref, bad
assert ref["provider"] == "openrouter"
def test_slot_max_tokens_absent_by_default():
"""Slots without max_tokens don't get the field — backward compat."""
cfg = normalize_moa_config(
{
"presets": {
"p": {
"reference_models": [
{"provider": "openrouter", "model": "deepseek/deepseek-v4-pro"},
],
}
}
}
)
ref = cfg["presets"]["p"]["reference_models"][0]
assert "max_tokens" not in ref