diff --git a/skillopt/model/__init__.py b/skillopt/model/__init__.py index d332997..5a2d1f7 100644 --- a/skillopt/model/__init__.py +++ b/skillopt/model/__init__.py @@ -215,7 +215,7 @@ def chat_optimizer_messages( timeout=timeout, ) if get_optimizer_backend() == "minimax_chat": - return _minimax.chat_target_messages( + return _minimax.chat_optimizer_messages( messages=messages, max_completion_tokens=max_completion_tokens, retries=retries, @@ -540,3 +540,4 @@ def set_optimizer_deployment(deployment: str) -> None: _openai.set_optimizer_deployment(deployment) _claude.set_optimizer_deployment(deployment) _qwen.set_optimizer_deployment(deployment) + _minimax.set_optimizer_deployment(deployment) diff --git a/skillopt/model/minimax_backend.py b/skillopt/model/minimax_backend.py index 7d9a42c..bb888a1 100644 --- a/skillopt/model/minimax_backend.py +++ b/skillopt/model/minimax_backend.py @@ -36,6 +36,9 @@ TARGET_DEPLOYMENT = os.environ.get( "TARGET_DEPLOYMENT", default_model_for_backend("minimax_chat"), ) +# Optimizer role can point at a different MiniMax model than the target; falls +# back to TARGET_DEPLOYMENT when unset so single-model setups are unaffected. +OPTIMIZER_DEPLOYMENT = os.environ.get("OPTIMIZER_DEPLOYMENT", "") _config_lock = threading.Lock() tracker = TokenTracker() @@ -222,7 +225,7 @@ def chat_target( stage: str = "target", reasoning_effort: str | None = None, timeout: float | None = None, -) -> tuple[str, dict[int]]: +) -> tuple[str, dict[str, int]]: del reasoning_effort messages = [{"role": "system", "content": system}, {"role": "user", "content": user}] return _chat_messages_impl( @@ -242,7 +245,7 @@ def chat_optimizer( stage: str = "optimizer", reasoning_effort: str | None = None, timeout: float | None = None, -) -> tuple[str, dict[int]]: +) -> tuple[str, dict[str, int]]: """Optimizer chat call. Backend stores the trained skill; uses the same MiniMax-proxied OpenAI-compat endpoint as `chat_target`. Added in the parallel-training fix; previously missing in skillopt 0.2.0's @@ -256,6 +259,7 @@ def chat_optimizer( max_completion_tokens, retries, stage, + deployment=OPTIMIZER_DEPLOYMENT or None, timeout=timeout, ) @@ -285,6 +289,35 @@ def chat_target_messages( ) +def chat_optimizer_messages( + messages: list[dict[str, Any]], + max_completion_tokens: int = 16384, + retries: int = 5, + stage: str = "optimizer", + reasoning_effort: str | None = None, + *, + tools: list[dict[str, Any]] | None = None, + tool_choice: str | dict[str, Any] | None = None, + return_message: bool = False, + timeout: float | None = None, +) -> tuple[Any, dict[str, int]]: + """Optimizer-role message API. Same endpoint as ``chat_target_messages`` but + honours ``OPTIMIZER_DEPLOYMENT`` so optimizer and target can use distinct + MiniMax models; falls back to the target deployment when unset.""" + del reasoning_effort + return _chat_messages_impl( + messages, + max_completion_tokens, + retries, + stage, + tools=tools, + tool_choice=tool_choice, + return_message=return_message, + deployment=OPTIMIZER_DEPLOYMENT or None, + timeout=timeout, + ) + + def get_token_summary() -> dict[str, dict[str, int]]: return tracker.summary() @@ -300,4 +333,10 @@ def set_reasoning_effort(effort: str | None) -> None: def set_target_deployment(deployment: str) -> None: global TARGET_DEPLOYMENT TARGET_DEPLOYMENT = deployment or default_model_for_backend("minimax_chat") - os.environ["TARGET_DEPLOYMENT"] = TARGET_DEPLOYMENT \ No newline at end of file + os.environ["TARGET_DEPLOYMENT"] = TARGET_DEPLOYMENT + + +def set_optimizer_deployment(deployment: str) -> None: + global OPTIMIZER_DEPLOYMENT + OPTIMIZER_DEPLOYMENT = deployment or default_model_for_backend("minimax_chat") + os.environ["OPTIMIZER_DEPLOYMENT"] = OPTIMIZER_DEPLOYMENT \ No newline at end of file diff --git a/tests/test_minimax_backend.py b/tests/test_minimax_backend.py new file mode 100644 index 0000000..4ed86bb --- /dev/null +++ b/tests/test_minimax_backend.py @@ -0,0 +1,62 @@ +"""Tests for the MiniMax backend, focusing on optimizer/target deployment +routing (regression for the #116 follow-up: optimizer calls must honour +OPTIMIZER_DEPLOYMENT, not silently reuse TARGET_DEPLOYMENT).""" +from __future__ import annotations + +import unittest +from unittest import mock + +from skillopt.model import minimax_backend as mm + + +def _fake_response(_payload, _timeout): + return { + "choices": [{"message": {"content": "ok"}}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + } + + +class TestMiniMaxDeploymentRouting(unittest.TestCase): + def setUp(self): + self._saved = (mm.TARGET_DEPLOYMENT, mm.OPTIMIZER_DEPLOYMENT) + mm.set_target_deployment("target-model") + mm.set_optimizer_deployment("optimizer-model") + + def tearDown(self): + mm.TARGET_DEPLOYMENT, mm.OPTIMIZER_DEPLOYMENT = self._saved + + def _captured_model(self, fn, *args, **kwargs): + seen = {} + + def spy(payload, timeout): + seen["model"] = payload["model"] + return _fake_response(payload, timeout) + + with mock.patch.object(mm, "_post_chat_completion", side_effect=spy): + fn(*args, **kwargs) + return seen["model"] + + def test_target_text_uses_target_deployment(self): + self.assertEqual(self._captured_model(mm.chat_target, "sys", "usr"), "target-model") + + def test_optimizer_text_uses_optimizer_deployment(self): + # The core bug: before the fix this sent "target-model". + self.assertEqual(self._captured_model(mm.chat_optimizer, "sys", "usr"), "optimizer-model") + + def test_optimizer_messages_uses_optimizer_deployment(self): + msgs = [{"role": "user", "content": "hi"}] + self.assertEqual(self._captured_model(mm.chat_optimizer_messages, msgs), "optimizer-model") + + def test_target_messages_uses_target_deployment(self): + msgs = [{"role": "user", "content": "hi"}] + self.assertEqual(self._captured_model(mm.chat_target_messages, msgs), "target-model") + + def test_optimizer_falls_back_to_target_when_unset(self): + # Empty optimizer deployment -> setter fills default; explicitly clear it + # to confirm the `or None` fallback path uses TARGET_DEPLOYMENT. + mm.OPTIMIZER_DEPLOYMENT = "" + self.assertEqual(self._captured_model(mm.chat_optimizer, "sys", "usr"), "target-model") + + +if __name__ == "__main__": + unittest.main()