fix(llm): forward max_tokens to Ollama num_predict

This commit is contained in:
Aniruddha Adak
2026-09-17 02:33:20 +05:30
parent 8321021c54
commit 6502cf24bf
2 changed files with 60 additions and 1 deletions
+6 -1
View File
@@ -70,8 +70,13 @@ class OllamaProvider(LLMProvider):
"messages": [msg.to_dict() for msg in input.messages],
"stream": False,
}
options: dict[str, Any] = {}
if input.temperature != 1.0:
payload["options"] = {"temperature": input.temperature}
options["temperature"] = input.temperature
if input.max_tokens is not None:
options["num_predict"] = input.max_tokens
if options:
payload["options"] = options
data = json.dumps(payload).encode("utf-8")
req = urllib.request.Request(url, data=data, headers={"Content-Type": "application/json"})
+54
View File
@@ -1,3 +1,6 @@
import json
import urllib.request
from io import BytesIO
from types import SimpleNamespace
import pytest
@@ -5,6 +8,7 @@ import pytest
from llm.core.types import LLMInput, Message, Role, ToolDefinition
from llm.providers.claude import ClaudeProvider
from llm.providers.constants import EMPTY_FILTERED_RESPONSE_ERROR
from llm.providers.ollama import OllamaProvider
from llm.providers.openai import OpenAIProvider
@@ -114,6 +118,56 @@ def test_openai_provider_allows_missing_usage():
assert output.usage is None
@pytest.mark.parametrize(
("max_tokens", "temperature", "expected_options"),
[
(128, 1.0, {"num_predict": 128}),
(128, 0.2, {"temperature": 0.2, "num_predict": 128}),
(128, 0.0, {"temperature": 0.0, "num_predict": 128}),
(0, 1.0, {"num_predict": 0}),
(None, 1.0, {}),
(None, 0.2, {"temperature": 0.2}),
],
)
def test_ollama_provider_serializes_generation_options(
monkeypatch, max_tokens, temperature, expected_options
):
requests = []
def fake_urlopen(request, timeout):
requests.append((request, timeout))
return BytesIO(b'{"message": {"content": "ok"}, "done_reason": "stop"}')
monkeypatch.setattr(urllib.request, "urlopen", fake_urlopen)
provider = OllamaProvider(base_url="http://localhost:11434", default_model="llama3.2")
output = provider.generate(
LLMInput(
messages=[Message(role=Role.USER, content="hi")],
max_tokens=max_tokens,
temperature=temperature,
)
)
assert len(requests) == 1
request, timeout = requests[0]
expected_payload = {
"model": "llama3.2",
"messages": [{"role": "user", "content": "hi"}],
"stream": False,
}
if expected_options:
expected_payload["options"] = expected_options
assert json.loads(request.data) == expected_payload
assert request.full_url == "http://localhost:11434/api/chat"
assert request.get_method() == "POST"
assert request.get_header("Content-type") == "application/json"
assert timeout == 60
assert output.content == "ok"
assert output.model == "llama3.2"
assert output.stop_reason == "stop"
def test_claude_provider_serializes_tools_for_messages_api():
provider = ClaudeProvider(api_key="test")
client = _AnthropicClient()