From 6502cf24bf891dd7d74a99498ceb6be9b42b313a Mon Sep 17 00:00:00 2001 From: Aniruddha Adak Date: Thu, 17 Sep 2026 02:33:20 +0530 Subject: [PATCH] fix(llm): forward max_tokens to Ollama num_predict --- src/llm/providers/ollama.py | 7 ++++- tests/test_provider_tools.py | 54 ++++++++++++++++++++++++++++++++++++ 2 files changed, 60 insertions(+), 1 deletion(-) diff --git a/src/llm/providers/ollama.py b/src/llm/providers/ollama.py index 2f83338d0..257d17900 100644 --- a/src/llm/providers/ollama.py +++ b/src/llm/providers/ollama.py @@ -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"}) diff --git a/tests/test_provider_tools.py b/tests/test_provider_tools.py index 4c9c76f92..ba4c09415 100644 --- a/tests/test_provider_tools.py +++ b/tests/test_provider_tools.py @@ -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()