mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
Merge 1fa770fc5f into e6c4580a31
This commit is contained in:
commit
4b328c2da2
2 changed files with 155 additions and 5 deletions
|
|
@ -1,6 +1,6 @@
|
|||
import json
|
||||
import time
|
||||
from collections.abc import AsyncIterator, Iterator
|
||||
from collections.abc import AsyncIterator, Iterator, Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
|
||||
from httpx._models import Headers, Response
|
||||
|
|
@ -188,6 +188,10 @@ class OllamaConfig(BaseConfig):
|
|||
elif value["type"] == "json_schema":
|
||||
optional_params["format"] = value["json_schema"]["schema"]
|
||||
|
||||
if "functions_unsupported_model" in optional_params and optional_params.get("stream") is True:
|
||||
# functions_unsupported_model means tools are emulated via prompt injection here.
|
||||
optional_params["fake_stream"] = True
|
||||
|
||||
return optional_params
|
||||
|
||||
def _supports_function_calling(self, ollama_model_info: dict) -> bool:
|
||||
|
|
@ -246,7 +250,7 @@ class OllamaConfig(BaseConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: str,
|
||||
encoding: str | None,
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
@ -323,9 +327,11 @@ class OllamaConfig(BaseConfig):
|
|||
model_response.created = int(time.time())
|
||||
model_response.model = "ollama/" + model
|
||||
_prompt: Final = request_data.get("prompt", "")
|
||||
prompt_tokens: Final = response_json.get(
|
||||
"prompt_eval_count",
|
||||
len(encoding.encode(_prompt, disallowed_special=())),
|
||||
_prompt_eval_count: Final = response_json.get("prompt_eval_count")
|
||||
prompt_tokens: Final = (
|
||||
_prompt_eval_count
|
||||
if _prompt_eval_count is not None
|
||||
else (len(encoding.encode(_prompt, disallowed_special=())) if encoding is not None else 0)
|
||||
)
|
||||
completion_tokens: Final = response_json.get(
|
||||
"eval_count", len(response_json.get("message", dict()).get("content", ""))
|
||||
|
|
@ -393,6 +399,26 @@ class OllamaConfig(BaseConfig):
|
|||
|
||||
return data
|
||||
|
||||
def sign_request(
|
||||
self,
|
||||
headers: dict[
|
||||
str, str
|
||||
], # mutable-ok: returned unchanged, must stay assignable to BaseConfig's dict return type
|
||||
optional_params: Mapping[str, object],
|
||||
request_data: Mapping[str, object],
|
||||
api_base: str,
|
||||
api_key: str | None = None,
|
||||
model: str | None = None,
|
||||
stream: bool | None = None,
|
||||
fake_stream: bool | None = None,
|
||||
) -> tuple[dict[str, str], bytes | None]: # mutable-ok: return type must match BaseConfig's dict-shaped contract
|
||||
if fake_stream is True:
|
||||
# /api/generate defaults to streaming when "stream" is absent from the body, but
|
||||
# the shared fake-stream handling drops the key instead of setting it False. Force
|
||||
# it explicitly so the actual request to Ollama is a single non-streaming call.
|
||||
return headers, json.dumps({**request_data, "stream": False}).encode()
|
||||
return headers, None
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
|
|
|
|||
|
|
@ -2,14 +2,18 @@ import json
|
|||
from litellm._uuid import uuid
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.llms.ollama.completion.transformation import (
|
||||
OllamaConfig,
|
||||
OllamaTextCompletionResponseIterator,
|
||||
)
|
||||
from litellm.types.utils import Message, ModelResponse, ModelResponseStream
|
||||
from litellm.utils import get_optional_params
|
||||
|
||||
|
||||
class TestOllamaConfig:
|
||||
|
|
@ -502,3 +506,123 @@ class TestOllamaTextCompletionResponseIterator:
|
|||
assert result["usage"]["prompt_tokens"] == 10
|
||||
assert result["usage"]["completion_tokens"] == 5
|
||||
assert result["usage"]["total_tokens"] == 15
|
||||
|
||||
|
||||
class TestOllamaFakeStreamActivation:
|
||||
def _tools(self):
|
||||
return [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_current_weather",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
def test_tools_and_stream_activate_fake_stream(self):
|
||||
optional_params = get_optional_params(
|
||||
model="llama2",
|
||||
custom_llm_provider="ollama",
|
||||
tools=self._tools(),
|
||||
stream=True,
|
||||
drop_params=True,
|
||||
)
|
||||
|
||||
assert optional_params.get("fake_stream") is True
|
||||
assert optional_params.get("format") == "json"
|
||||
assert "functions_unsupported_model" in optional_params
|
||||
|
||||
def test_tools_without_stream_does_not_activate_fake_stream(self):
|
||||
optional_params = get_optional_params(
|
||||
model="llama2",
|
||||
custom_llm_provider="ollama",
|
||||
tools=self._tools(),
|
||||
stream=False,
|
||||
drop_params=True,
|
||||
)
|
||||
|
||||
assert "fake_stream" not in optional_params
|
||||
|
||||
def test_stream_without_tools_does_not_activate_fake_stream(self):
|
||||
optional_params = get_optional_params(
|
||||
model="llama2",
|
||||
custom_llm_provider="ollama",
|
||||
stream=True,
|
||||
)
|
||||
|
||||
assert "fake_stream" not in optional_params
|
||||
|
||||
|
||||
class TestOllamaFakeStreamToolCalls:
|
||||
def test_tools_stream_true_reconstructs_tool_calls_via_fake_stream(self):
|
||||
"""Test that tools + stream=True routes through fake_stream and yields reconstructed tool_calls."""
|
||||
tool_call_json = {
|
||||
"name": "get_current_weather",
|
||||
"arguments": {"location": "San Francisco"},
|
||||
}
|
||||
mock_ollama_response = {
|
||||
"model": "llama2",
|
||||
"response": json.dumps(tool_call_json),
|
||||
"done": True,
|
||||
"done_reason": "stop",
|
||||
"prompt_eval_count": 42,
|
||||
"eval_count": 16,
|
||||
}
|
||||
|
||||
mock_client = MagicMock(spec=HTTPHandler)
|
||||
mock_client.post.return_value = httpx.Response(
|
||||
status_code=200,
|
||||
content=json.dumps(mock_ollama_response).encode(),
|
||||
request=httpx.Request("POST", "http://127.0.0.1:11434/api/generate"),
|
||||
)
|
||||
|
||||
response = litellm.completion(
|
||||
model="ollama/llama2",
|
||||
api_base="http://127.0.0.1:11434",
|
||||
messages=[{"role": "user", "content": "What is the weather in San Francisco?"}],
|
||||
tools=[
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_current_weather",
|
||||
"description": "Get current weather.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"location": {"type": "string"}},
|
||||
"required": ["location"],
|
||||
},
|
||||
},
|
||||
}
|
||||
],
|
||||
stream=True,
|
||||
drop_params=True,
|
||||
client=mock_client,
|
||||
)
|
||||
|
||||
reassembled_content = ""
|
||||
tool_calls_seen = []
|
||||
finish_reasons = []
|
||||
for chunk in response:
|
||||
delta = chunk.choices[0].delta
|
||||
if delta.content:
|
||||
reassembled_content += delta.content
|
||||
if getattr(delta, "tool_calls", None):
|
||||
tool_calls_seen.extend(delta.tool_calls)
|
||||
if chunk.choices[0].finish_reason:
|
||||
finish_reasons.append(chunk.choices[0].finish_reason)
|
||||
|
||||
assert mock_client.post.call_count == 1
|
||||
request_body = json.loads(mock_client.post.call_args.kwargs["data"])
|
||||
|
||||
assert request_body.get("stream") is False
|
||||
assert "tools" not in request_body
|
||||
|
||||
assert tool_calls_seen, "expected delta.tool_calls to be populated"
|
||||
assert tool_calls_seen[0]["function"]["name"] == "get_current_weather"
|
||||
assert json.loads(tool_calls_seen[0]["function"]["arguments"]) == {
|
||||
"location": "San Francisco"
|
||||
}
|
||||
|
||||
assert finish_reasons == ["tool_calls"]
|
||||
assert json.dumps(tool_call_json) not in reassembled_content
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue