From 78312c192d008b21d58c7cb035f07cb231e75585 Mon Sep 17 00:00:00 2001 From: wjq990112 <45777252+wjq990112@users.noreply.github.com> Date: Sun, 19 Jul 2026 05:09:41 +0000 Subject: [PATCH] fix(moa): preserve custom provider context metadata Preserve compatible custom provider metadata through MoA aggregator context resolution and cover the resolver and compressor paths. --- agent/model_metadata.py | 14 ++- tests/agent/test_model_metadata.py | 173 ++++++++++++++++++++++++++--- 2 files changed, 167 insertions(+), 20 deletions(-) diff --git a/agent/model_metadata.py b/agent/model_metadata.py index a358a6217c54..288083628e02 100644 --- a/agent/model_metadata.py +++ b/agent/model_metadata.py @@ -2201,11 +2201,18 @@ def get_model_context_length( # acting context, so they're ignored here. if (provider or "").strip().lower() == "moa": try: - from hermes_cli.config import load_config + from hermes_cli.config import ( + get_compatible_custom_providers, + load_config, + ) from hermes_cli.moa_config import resolve_moa_preset from hermes_cli.runtime_provider import resolve_runtime_provider - preset = resolve_moa_preset(load_config().get("moa") or {}, model) + config = load_config() + effective_custom_providers = custom_providers + if effective_custom_providers is None: + effective_custom_providers = get_compatible_custom_providers(config) + preset = resolve_moa_preset(config.get("moa") or {}, model) agg = preset.get("aggregator") or {} agg_provider = str(agg.get("provider") or "").strip() agg_model = str(agg.get("model") or "").strip() @@ -2215,7 +2222,8 @@ def get_model_context_length( agg_model, base_url=rt.get("base_url", "") or "", api_key=rt.get("api_key", "") or "", - provider=agg_provider, + provider=rt.get("provider") or agg_provider, + custom_providers=effective_custom_providers, ) except Exception: logger.debug("MoA aggregator context-length resolution failed", exc_info=True) diff --git a/tests/agent/test_model_metadata.py b/tests/agent/test_model_metadata.py index eebef76ec1dd..7a29178a40a0 100644 --- a/tests/agent/test_model_metadata.py +++ b/tests/agent/test_model_metadata.py @@ -1800,27 +1800,31 @@ class TestGrok43StaleCacheGuard: class TestMoAContextLength: """MoA virtual provider resolves context from the aggregator slot, not 256K default.""" - def _write_moa_config(self, home, aggregator): + def _write_moa_config( + self, home, aggregator, custom_providers=None, providers=None + ): import os os.makedirs(home, exist_ok=True) - with open(os.path.join(home, "config.yaml"), "w") as f: - yaml.safe_dump( - { - "moa": { - "default_preset": "p", - "presets": { - "p": { - "enabled": True, - "reference_models": [ - {"provider": "openrouter", "model": "openai/gpt-5.5"} - ], - "aggregator": aggregator, - } - }, + payload = { + "moa": { + "default_preset": "p", + "presets": { + "p": { + "enabled": True, + "reference_models": [ + {"provider": "openrouter", "model": "openai/gpt-5.5"} + ], + "aggregator": aggregator, } }, - f, - ) + } + } + if custom_providers is not None: + payload["custom_providers"] = custom_providers + if providers is not None: + payload["providers"] = providers + with open(os.path.join(home, "config.yaml"), "w") as f: + yaml.safe_dump(payload, f) def test_moa_resolves_from_aggregator(self, tmp_path, monkeypatch): home = str(tmp_path / ".hermes") @@ -1843,3 +1847,138 @@ class TestMoAContextLength: "p", base_url="http://127.0.0.1/v1", provider="moa", config_context_length=500_000 ) assert ctx == 500_000 + + def test_moa_resolves_custom_provider_per_model_context(self, tmp_path, monkeypatch): + home = str(tmp_path / ".hermes") + monkeypatch.setenv("HERMES_HOME", home) + self._write_moa_config( + home, + {"provider": "custom:example", "model": "example-model"}, + custom_providers=[ + { + "name": "example", + "base_url": "http://127.0.0.1:1/v1", + "model": "example-model", + "models": { + "example-model": {"context_length": 777_777}, + }, + } + ], + ) + + ctx = get_model_context_length( + "p", base_url="http://127.0.0.1/v1", provider="moa" + ) + + assert ctx == 777_777 + + def test_moa_resolves_canonical_provider_per_model_context( + self, tmp_path, monkeypatch + ): + home = str(tmp_path / ".hermes") + monkeypatch.setenv("HERMES_HOME", home) + self._write_moa_config( + home, + {"provider": "custom:example", "model": "example-model"}, + providers={ + "example": { + "api": "http://127.0.0.1:1/v1", + "default_model": "example-model", + "models": { + "example-model": {"context_length": 888_888}, + }, + } + }, + ) + + with patch( + "agent.model_metadata._resolve_endpoint_context_length", + return_value=None, + ) as endpoint_probe: + ctx = get_model_context_length( + "p", base_url="http://127.0.0.1/v1", provider="moa" + ) + + assert ctx == 888_888 + endpoint_probe.assert_not_called() + + def test_moa_custom_context_configures_compressor_threshold( + self, tmp_path, monkeypatch + ): + from agent.context_compressor import ContextCompressor + + configured_context = 600_000 + home = str(tmp_path / ".hermes") + monkeypatch.setenv("HERMES_HOME", home) + self._write_moa_config( + home, + {"provider": "custom:example", "model": "example-model"}, + providers={ + "example": { + "api": "http://127.0.0.1:1/v1", + "default_model": "example-model", + "models": { + "example-model": { + "context_length": configured_context, + }, + }, + } + }, + ) + + with patch( + "agent.model_metadata._resolve_endpoint_context_length", + return_value=None, + ) as endpoint_probe: + compressor = ContextCompressor( + model="p", + base_url="http://127.0.0.1/v1", + provider="moa", + threshold_percent=0.50, + quiet_mode=True, + ) + + assert compressor.context_length == configured_context + assert compressor.threshold_tokens == configured_context // 2 + endpoint_probe.assert_not_called() + + def test_moa_preserves_caller_supplied_custom_provider_context( + self, tmp_path, monkeypatch + ): + home = str(tmp_path / ".hermes") + monkeypatch.setenv("HERMES_HOME", home) + self._write_moa_config( + home, + {"provider": "custom:example", "model": "example-model"}, + providers={ + "example": { + "api": "http://127.0.0.1:1/v1", + "default_model": "example-model", + "models": {"example-model": {}}, + } + }, + ) + supplied = [ + { + "name": "example", + "base_url": "http://127.0.0.1:1/v1", + "model": "example-model", + "models": { + "example-model": {"context_length": 999_999}, + }, + } + ] + + with patch( + "agent.model_metadata._resolve_endpoint_context_length", + return_value=None, + ) as endpoint_probe: + ctx = get_model_context_length( + "p", + base_url="http://127.0.0.1/v1", + provider="moa", + custom_providers=supplied, + ) + + assert ctx == 999_999 + endpoint_probe.assert_not_called()