fix nvidia tool message schema

This commit is contained in:
pmos69
2026-05-23 01:38:15 +01:00
committed by Riccardo Roveri
parent 00e5a361b6
commit fc05cfb58e
3 changed files with 112 additions and 1 deletions
+26 -1
View File
@@ -1,9 +1,34 @@
"""NVIDIA NIM provider profile."""
import copy
from typing import Any
from providers import register_provider
from providers.base import ProviderProfile
nvidia = ProviderProfile(
class NvidiaProviderProfile(ProviderProfile):
"""NVIDIA NIM accepts a stricter ToolMessage schema than most OpenAI-compatible APIs."""
def prepare_messages(self, messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
needs_sanitize = any(
isinstance(msg, dict)
and msg.get("role") == "tool"
and ("name" in msg or "tool_name" in msg)
for msg in messages
)
if not needs_sanitize:
return messages
sanitized = copy.deepcopy(messages)
for msg in sanitized:
if isinstance(msg, dict) and msg.get("role") == "tool":
msg.pop("name", None)
msg.pop("tool_name", None)
return sanitized
nvidia = NvidiaProviderProfile(
name="nvidia",
aliases=("nvidia-nim",),
env_vars=("NVIDIA_API_KEY",),
+45
View File
@@ -39,6 +39,51 @@ class TestNvidiaProfileWiring:
assert kwargs["model"] == "nvidia/test-model"
def test_nvidia_tool_messages_drop_name_fields(self, transport):
profile = get_provider_profile("nvidia")
msgs = [
{"role": "user", "content": "run a command"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {"name": "terminal", "arguments": "{}"},
}
],
},
{
"role": "tool",
"name": "terminal",
"tool_name": "terminal",
"tool_call_id": "call_1",
"content": "ok",
},
]
kwargs = transport.build_kwargs(
model="mistralai/mistral-large-3-675b-instruct-2512",
messages=msgs,
tools=None,
provider_profile=profile,
max_tokens=None,
max_tokens_param_fn=lambda x: {"max_tokens": x} if x else {},
timeout=300,
reasoning_config=None,
request_overrides=None,
session_id="test",
ollama_num_ctx=None,
)
assert kwargs["messages"][2] == {
"role": "tool",
"tool_call_id": "call_1",
"content": "ok",
}
assert msgs[2]["name"] == "terminal"
assert msgs[2]["tool_name"] == "terminal"
class TestDeepSeekProfileWiring:
def test_deepseek_no_forced_max_tokens(self, transport):
+41
View File
@@ -25,6 +25,47 @@ class TestNvidiaProfile:
assert "nvidia.com" in p.base_url
def test_prepare_messages_strips_tool_result_names(self):
p = get_provider_profile("nvidia")
msgs = [
{"role": "user", "content": "run a command"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {"name": "terminal", "arguments": "{}"},
}
],
},
{
"role": "tool",
"name": "terminal",
"tool_name": "terminal",
"tool_call_id": "call_1",
"content": "ok",
},
]
result = p.prepare_messages(msgs)
assert "name" not in result[2]
assert "tool_name" not in result[2]
assert result[2] == {
"role": "tool",
"tool_call_id": "call_1",
"content": "ok",
}
assert msgs[2]["name"] == "terminal"
assert msgs[2]["tool_name"] == "terminal"
def test_prepare_messages_passthrough_without_tool_result_names(self):
p = get_provider_profile("nvidia")
msgs = [{"role": "tool", "tool_call_id": "call_1", "content": "ok"}]
assert p.prepare_messages(msgs) is msgs
class TestKimiProfile:
def test_temperature_omit(self):