fix nvidia tool message schema
This commit is contained in:
@@ -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",),
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user