mirror of
https://github.com/affaan-m/ECC.git
synced 2026-09-20 16:47:59 +02:00
fix(llm): forward max_tokens to Ollama num_predict
This commit is contained in:
@@ -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"})
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user