Skip to content

Commit d0e5e68

Browse files
authored
fix: handle max_tokens for NVIDIA MiniMax M3 (#9209)
1 parent ab4de04 commit d0e5e68

2 files changed

Lines changed: 86 additions & 7 deletions

File tree

astrbot/core/provider/sources/openai_source.py

Lines changed: 21 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -404,10 +404,26 @@ def _ollama_disable_thinking_enabled(self) -> bool:
404404
return value.strip().lower() in {"1", "true", "yes", "on"}
405405
return bool(value)
406406

407-
def _apply_provider_specific_extra_body_overrides(
408-
self, extra_body: dict[str, Any]
407+
def _apply_provider_specific_request_overrides(
408+
self,
409+
payloads: dict[str, Any],
410+
extra_body: dict[str, Any],
409411
) -> None:
410-
if self.provider_config.get("provider") != "ollama":
412+
provider = self.provider_config.get("provider")
413+
model = str(payloads.get("model", "")).lower()
414+
415+
# NVIDIA's hosted MiniMax M3 endpoint can return empty choices when
416+
# max_tokens is omitted (#9206). Scope the compatibility default to
417+
# that model; other NVIDIA models have different token limits.
418+
if (
419+
provider == "nvidia"
420+
and model == "minimaxai/minimax-m3"
421+
and "max_tokens" not in payloads
422+
and "max_tokens" not in extra_body
423+
):
424+
payloads["max_tokens"] = 8192
425+
426+
if provider != "ollama":
411427
return
412428
if not self._ollama_disable_thinking_enabled():
413429
return
@@ -543,7 +559,7 @@ async def _query(
543559
custom_extra_body = self.provider_config.get("custom_extra_body", {})
544560
if isinstance(custom_extra_body, dict):
545561
extra_body.update(custom_extra_body)
546-
self._apply_provider_specific_extra_body_overrides(extra_body)
562+
self._apply_provider_specific_request_overrides(payloads, extra_body)
547563

548564
model = payloads.get("model", "").lower()
549565

@@ -603,7 +619,7 @@ async def _query_stream(
603619
to_del.append(key)
604620
for key in to_del:
605621
del payloads[key]
606-
self._apply_provider_specific_extra_body_overrides(extra_body)
622+
self._apply_provider_specific_request_overrides(payloads, extra_body)
607623

608624
self._sanitize_assistant_messages(payloads)
609625

tests/test_openai_source.py

Lines changed: 65 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1272,7 +1272,7 @@ async def fake_resolve_media_ref_to_base64_data(
12721272

12731273

12741274
@pytest.mark.asyncio
1275-
async def test_apply_provider_specific_extra_body_overrides_disables_ollama_thinking():
1275+
async def test_apply_provider_specific_request_overrides_disables_ollama_thinking():
12761276
provider = _make_provider(
12771277
{
12781278
"provider": "ollama",
@@ -1287,7 +1287,7 @@ async def test_apply_provider_specific_extra_body_overrides_disables_ollama_thin
12871287
"temperature": 0.2,
12881288
}
12891289

1290-
provider._apply_provider_specific_extra_body_overrides(extra_body)
1290+
provider._apply_provider_specific_request_overrides({}, extra_body)
12911291

12921292
assert extra_body["reasoning_effort"] == "none"
12931293
assert "reasoning" not in extra_body
@@ -1297,6 +1297,69 @@ async def test_apply_provider_specific_extra_body_overrides_disables_ollama_thin
12971297
await provider.terminate()
12981298

12991299

1300+
@pytest.mark.asyncio
1301+
async def test_provider_specific_request_overrides_sets_minimax_m3_max_tokens():
1302+
provider = _make_provider({"provider": "nvidia"})
1303+
try:
1304+
payloads = {"model": "minimaxai/minimax-m3"}
1305+
extra_body = {"temperature": 0.2}
1306+
1307+
provider._apply_provider_specific_request_overrides(payloads, extra_body)
1308+
1309+
assert payloads["max_tokens"] == 8192
1310+
assert extra_body == {"temperature": 0.2}
1311+
finally:
1312+
await provider.terminate()
1313+
1314+
1315+
@pytest.mark.asyncio
1316+
async def test_minimax_m3_max_tokens_preserves_custom_extra_body_value():
1317+
provider = _make_provider({"provider": "nvidia"})
1318+
try:
1319+
payloads = {"model": "minimaxai/minimax-m3"}
1320+
extra_body = {"max_tokens": 4096}
1321+
1322+
provider._apply_provider_specific_request_overrides(payloads, extra_body)
1323+
1324+
assert "max_tokens" not in payloads
1325+
assert extra_body["max_tokens"] == 4096
1326+
finally:
1327+
await provider.terminate()
1328+
1329+
1330+
@pytest.mark.asyncio
1331+
async def test_minimax_m3_max_tokens_preserves_standard_payload_value():
1332+
provider = _make_provider({"provider": "nvidia"})
1333+
try:
1334+
payloads = {
1335+
"model": "minimaxai/minimax-m3",
1336+
"max_tokens": 2048,
1337+
}
1338+
extra_body = {}
1339+
1340+
provider._apply_provider_specific_request_overrides(payloads, extra_body)
1341+
1342+
assert payloads["max_tokens"] == 2048
1343+
assert extra_body == {}
1344+
finally:
1345+
await provider.terminate()
1346+
1347+
1348+
@pytest.mark.asyncio
1349+
async def test_nvidia_request_does_not_set_max_tokens_for_other_models():
1350+
provider = _make_provider({"provider": "nvidia"})
1351+
try:
1352+
payloads = {"model": "nvidia/usdcode"}
1353+
extra_body = {}
1354+
1355+
provider._apply_provider_specific_request_overrides(payloads, extra_body)
1356+
1357+
assert "max_tokens" not in payloads
1358+
assert "max_tokens" not in extra_body
1359+
finally:
1360+
await provider.terminate()
1361+
1362+
13001363
@pytest.mark.asyncio
13011364
async def test_query_injects_reasoning_effort_none_for_ollama(monkeypatch):
13021365
provider = _make_provider(

0 commit comments

Comments
 (0)