mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(ollama): satisfy lint budgets and simplify streaming tool call normalization
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
fa7b3f74d9
commit
f32dd6a6ea
2 changed files with 45 additions and 36 deletions
|
|
@ -341,8 +341,8 @@ class OllamaChatConfig(BaseConfig):
|
|||
if "error" in response_json:
|
||||
raise OllamaError(
|
||||
message=str(response_json["error"]),
|
||||
status_code=raw_response.status_code if raw_response.status_code >= 400 else 400,
|
||||
headers=dict(raw_response.headers),
|
||||
status_code=max(raw_response.status_code, 400),
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
|
||||
## RESPONSE OBJECT
|
||||
|
|
@ -387,7 +387,11 @@ class OllamaChatConfig(BaseConfig):
|
|||
model_response.choices[0].message = message
|
||||
model_response.choices[0].finish_reason = "tool_calls"
|
||||
else:
|
||||
_message: Final = litellm.Message(**(response_json_message or {}))
|
||||
_message: Final = (
|
||||
litellm.Message(**response_json_message)
|
||||
if response_json_message is not None
|
||||
else litellm.Message(content=None)
|
||||
)
|
||||
model_response.choices[0].message = _message
|
||||
# Set finish_reason to "tool_calls" when tool_calls are present
|
||||
# Fixes: https://github.com/BerriAI/litellm/issues/18922
|
||||
|
|
@ -396,9 +400,8 @@ class OllamaChatConfig(BaseConfig):
|
|||
model_response.created = int(time.time())
|
||||
model_response.model = "ollama_chat/" + model
|
||||
prompt_tokens = response_json.get("prompt_eval_count", litellm.token_counter(messages=messages))
|
||||
completion_tokens: Final = response_json.get("eval_count") or litellm.token_counter(
|
||||
text=(response_json_message or {}).get("content") or ""
|
||||
)
|
||||
_message_content: Final = response_json_message.get("content") if response_json_message is not None else None
|
||||
completion_tokens: Final = response_json.get("eval_count") or litellm.token_counter(text=_message_content or "")
|
||||
setattr(
|
||||
model_response,
|
||||
"usage",
|
||||
|
|
@ -431,6 +434,21 @@ class OllamaChatCompletionResponseIterator(BaseModelResponseIterator):
|
|||
finished_reasoning_content: bool = False
|
||||
stream_tool_call_count: int = 0
|
||||
|
||||
def _assign_tool_call_index_and_id(
|
||||
self,
|
||||
tool_call: dict, # mutable-ok: normalizes the provider chunk's tool call dict in place
|
||||
) -> None:
|
||||
function: Final = tool_call.get("function")
|
||||
if function is None:
|
||||
return
|
||||
function_index: Final = function.pop("index", None)
|
||||
if tool_call.get("index") is None:
|
||||
tool_call["index"] = function_index if function_index is not None else self.stream_tool_call_count
|
||||
self.stream_tool_call_count = max(self.stream_tool_call_count + 1, tool_call["index"] + 1)
|
||||
function_args: Final = function.get("arguments")
|
||||
if function_args is not None and len(function_args) > 0 and self._is_function_call_complete(function_args):
|
||||
tool_call["id"] = str(uuid.uuid4())
|
||||
|
||||
def _is_function_call_complete(self, function_args: str | dict) -> bool:
|
||||
if isinstance(function_args, dict):
|
||||
return True
|
||||
|
|
@ -476,23 +494,14 @@ class OllamaChatCompletionResponseIterator(BaseModelResponseIterator):
|
|||
raise OllamaError(
|
||||
message=str(chunk["error"]),
|
||||
status_code=400,
|
||||
headers={"Content-Type": "application/json"},
|
||||
headers=Headers(),
|
||||
)
|
||||
|
||||
# process tool calls - if complete function arg - add id to tool call
|
||||
tool_calls: Final = chunk["message"].get("tool_calls")
|
||||
if tool_calls is not None:
|
||||
for tool_call in tool_calls:
|
||||
function = tool_call.get("function") or {}
|
||||
function_index = function.pop("index", None)
|
||||
if tool_call.get("index") is None:
|
||||
tool_call["index"] = function_index if function_index is not None else self.stream_tool_call_count
|
||||
self.stream_tool_call_count = max(self.stream_tool_call_count + 1, tool_call["index"] + 1)
|
||||
function_args = function.get("arguments")
|
||||
if function_args is not None and len(function_args) > 0:
|
||||
is_function_call_complete = self._is_function_call_complete(function_args)
|
||||
if is_function_call_complete:
|
||||
tool_call["id"] = str(uuid.uuid4())
|
||||
self._assign_tool_call_index_and_id(tool_call)
|
||||
|
||||
# PROCESS REASONING CONTENT
|
||||
reasoning_content: str | None = None
|
||||
|
|
|
|||
|
|
@ -20,7 +20,9 @@ from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
|||
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUsageBlock
|
||||
from litellm.types.utils import (
|
||||
ChatCompletionDeltaToolCall,
|
||||
Delta,
|
||||
Function,
|
||||
GenericStreamingChunk,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
|
|
@ -456,25 +458,23 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
|
|||
def _handle_string_chunk(self, str_line: str) -> GenericStreamingChunk | ModelResponseStream:
|
||||
return self.chunk_parser(json.loads(str_line))
|
||||
|
||||
def _parse_buffered_function_call(self) -> list[dict] | None:
|
||||
def _parse_buffered_function_call(self) -> ChatCompletionDeltaToolCall | None:
|
||||
if self.buffered_json_content is None:
|
||||
return None
|
||||
try:
|
||||
parsed = json.loads(self.buffered_json_content)
|
||||
parsed: Final = json.loads(self.buffered_json_content)
|
||||
except json.JSONDecodeError:
|
||||
return None
|
||||
if isinstance(parsed, dict) and "name" in parsed and "arguments" in parsed:
|
||||
return [
|
||||
{
|
||||
"id": f"call_{uuid.uuid4()}",
|
||||
"index": 0,
|
||||
"function": {
|
||||
"name": parsed["name"],
|
||||
"arguments": json.dumps(parsed["arguments"]),
|
||||
},
|
||||
"type": "function",
|
||||
}
|
||||
]
|
||||
return ChatCompletionDeltaToolCall(
|
||||
id=f"call_{uuid.uuid4()}",
|
||||
index=0,
|
||||
type="function",
|
||||
function=Function(
|
||||
name=parsed["name"],
|
||||
arguments=json.dumps(parsed["arguments"]),
|
||||
),
|
||||
)
|
||||
return None
|
||||
|
||||
def chunk_parser(self, chunk: dict) -> GenericStreamingChunk | ModelResponseStream:
|
||||
|
|
@ -499,13 +499,13 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
|
|||
completion_tokens=eval_count,
|
||||
total_tokens=prompt_eval_count + eval_count,
|
||||
)
|
||||
tool_calls: Final = self._parse_buffered_function_call()
|
||||
if tool_calls is not None:
|
||||
buffered_tool_call: Final = self._parse_buffered_function_call()
|
||||
if buffered_tool_call is not None:
|
||||
return ModelResponseStream(
|
||||
choices=[
|
||||
choices=[ # mutable-ok: ModelResponseStream only accepts a list of choices
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(content=None, tool_calls=tool_calls),
|
||||
delta=Delta(content=None, tool_calls=(buffered_tool_call,)),
|
||||
finish_reason="tool_calls",
|
||||
)
|
||||
],
|
||||
|
|
@ -513,7 +513,7 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
|
|||
)
|
||||
if self.buffered_json_content is not None:
|
||||
return ModelResponseStream(
|
||||
choices=[
|
||||
choices=[ # mutable-ok: ModelResponseStream only accepts a list of choices
|
||||
StreamingChoices(
|
||||
index=0,
|
||||
delta=Delta(content=self.buffered_json_content),
|
||||
|
|
@ -537,7 +537,7 @@ class OllamaTextCompletionResponseIterator(BaseModelResponseIterator):
|
|||
):
|
||||
self.buffered_json_content = (self.buffered_json_content or "") + text
|
||||
return ModelResponseStream(
|
||||
choices=[StreamingChoices(index=0, delta=Delta())],
|
||||
choices=[StreamingChoices(index=0, delta=Delta())], # mutable-ok: requires a list of choices
|
||||
usage=None,
|
||||
)
|
||||
reasoning_content: str | None = None
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue