diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index e303a18962f..ba9b04ab222 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -32,6 +32,7 @@ from litellm.types.llms.openai import ( ChatCompletionFileObject, ChatCompletionFunctionMessage, ChatCompletionImageObject, + ChatCompletionSystemMessage, ChatCompletionTextObject, ChatCompletionToolCallFunctionChunk, ChatCompletionToolMessage, @@ -5130,6 +5131,26 @@ def _bedrock_tools_pt(tools: list, model: str | None = None) -> list[BedrockTool return tool_block_list +def _append_function_prompt(message: ChatCompletionSystemMessage, text: str) -> ChatCompletionSystemMessage: + content: Final = message["content"] + if isinstance(content, str): + return {**message, "content": content + text} + return {**message, "content": [*content, ChatCompletionTextObject(type="text", text=text)]} + + +def function_call_prompt(messages: Sequence[AllMessageValues], function_descriptions: str) -> list[AllMessageValues]: + function_prompt: Final = ( + 'Produce JSON OUTPUT ONLY! Adhere to this format {"name": "function_name", "arguments":{"argument_name": ' + '"argument_value"}} The following functions are available to you:' + function_descriptions + ) + if not any(message["role"] == "system" for message in messages): + return [*messages, ChatCompletionSystemMessage(role="system", content=function_prompt)] + return [ + _append_function_prompt(message, f" {function_prompt}") if message["role"] == "system" else message + for message in messages + ] + + def response_schema_prompt(model: str, response_schema: dict) -> str: """ Decides if a user-defined custom prompt or default needs to be used diff --git a/litellm/llms/ollama/chat/transformation.py b/litellm/llms/ollama/chat/transformation.py index bd624bfbfa0..db97c139e02 100644 --- a/litellm/llms/ollama/chat/transformation.py +++ b/litellm/llms/ollama/chat/transformation.py @@ -24,7 +24,6 @@ from litellm.types.llms.ollama import ( from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionAssistantToolCall, - ChatCompletionToolParam, ChatCompletionUsageBlock, ) from litellm.types.utils import ModelResponse, ModelResponseStream @@ -184,10 +183,8 @@ class OllamaChatConfig(BaseConfig): if param == "tools": optional_params["tools"] = value - if param == "functions": - optional_params["tools"] = tuple( - ChatCompletionToolParam(type="function", function=function) for function in value - ) + if param == "functions" and value: + optional_params["tools"] = [{"type": "function", "function": function} for function in value] non_default_params.pop("tool_choice", None) # causes ollama requests to hang non_default_params.pop("functions", None) # causes ollama requests to hang return optional_params diff --git a/litellm/llms/ollama/common_utils.py b/litellm/llms/ollama/common_utils.py index ed4bab22a84..0111340ca8e 100644 --- a/litellm/llms/ollama/common_utils.py +++ b/litellm/llms/ollama/common_utils.py @@ -11,6 +11,18 @@ class OllamaError(BaseLLMException): super().__init__(status_code=status_code, message=message, headers=headers) +def resolve_ollama_tool_calling_provider( + custom_llm_provider: str, has_tools: bool, add_function_to_prompt: bool +) -> str: + """ + /api/generate has no native tool calling, so ollama/ tool requests go through the ollama_chat + adapter unless add_function_to_prompt opts back into the legacy JSON prompt emulation + """ + if custom_llm_provider == "ollama" and has_tools and not add_function_to_prompt: + return "ollama_chat" + return custom_llm_provider + + def _convert_image(image): """ Convert image to base64 encoded image if not already in base64 format diff --git a/litellm/llms/ollama/completion/transformation.py b/litellm/llms/ollama/completion/transformation.py index a1340ba1952..7a748fb260f 100644 --- a/litellm/llms/ollama/completion/transformation.py +++ b/litellm/llms/ollama/completion/transformation.py @@ -14,6 +14,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( from litellm.litellm_core_utils.prompt_templates.factory import ( convert_to_ollama_image, custom_prompt, + function_call_prompt, ollama_pt, ) from litellm.litellm_core_utils.prompt_templates.image_handling import ( @@ -159,6 +160,9 @@ class OllamaConfig(BaseConfig): "response_format", "max_completion_tokens", "reasoning_effort", + "tools", + "tool_choice", + "functions", ] def map_openai_params( @@ -193,6 +197,9 @@ class OllamaConfig(BaseConfig): optional_params["format"] = "json" elif value["type"] == "json_schema": optional_params["format"] = value["json_schema"]["schema"] + elif param in ("tools", "functions") and value: + optional_params["format"] = "json" + optional_params["prompted_functions"] = "".join(f"\n{function}\n" for function in value) return optional_params @@ -377,6 +384,12 @@ class OllamaConfig(BaseConfig): headers: dict, ) -> dict: custom_prompt_dict: Final = litellm_params.get("custom_prompt_dict") or litellm.custom_prompt_dict + prompted_functions: Final = optional_params.pop("prompted_functions", None) + prompt_messages: Final = ( + function_call_prompt(messages=messages, function_descriptions=prompted_functions) + if isinstance(prompted_functions, str) + else messages + ) text_completion_request: Final = litellm_params.get("text_completion") if model in custom_prompt_dict: @@ -386,12 +399,12 @@ class OllamaConfig(BaseConfig): role_dict=model_prompt_details["roles"], initial_prompt_value=model_prompt_details["initial_prompt_value"], final_prompt_value=model_prompt_details["final_prompt_value"], - messages=messages, + messages=prompt_messages, ) elif text_completion_request: # handle `/completions` requests - ollama_prompt = get_str_from_messages(messages=messages) + ollama_prompt = get_str_from_messages(messages=prompt_messages) else: # handle `/chat/completions` requests - modified_prompt: Final = ollama_pt(model=model, messages=messages) + modified_prompt: Final = ollama_pt(model=model, messages=prompt_messages) if isinstance(modified_prompt, dict): ollama_prompt, images = ( modified_prompt["prompt"], diff --git a/litellm/main.py b/litellm/main.py index 8984f6abfec..6db13ddcdb7 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -220,6 +220,7 @@ from .llms.nvidia_riva.audio_transcription.transformation import ( NvidiaRivaAudioTranscriptionConfig, ) from .llms.oci.chat.transformation import OCIChatConfig +from .llms.ollama.common_utils import resolve_ollama_tool_calling_provider from .llms.ollama.completion import handler as ollama from .llms.oobabooga.chat import oobabooga from .llms.openai.completion.handler import OpenAITextCompletion @@ -5307,11 +5308,11 @@ def completion( GenericLiteLLMParams(**_supplemental_provider_params) if _supplemental_provider_params else None ), ) - if custom_llm_provider == "ollama" and (tools or functions): - custom_llm_provider = "ollama_chat" # rebind-ok: /api/generate has no native tool calling - elif custom_llm_provider == "ollama": - tools = None # rebind-ok: empty tools must not change plain completion behavior - functions = None # rebind-ok: empty functions must not change plain completion behavior + custom_llm_provider = resolve_ollama_tool_calling_provider( # rebind-ok: ollama tools use the chat adapter + custom_llm_provider, + has_tools=True if tools or functions else False, + add_function_to_prompt=litellm.add_function_to_prompt, + ) ## RESPONSES API BRIDGE LOGIC ## - check early and normalize model name responses_api_model_info, model = responses_api_bridge_check( diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 66d10fd1407..d6bac6b45f6 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -20,6 +20,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( _convert_to_bedrock_tool_call_result, anthropic_messages_pt, convert_to_gemini_tool_call_result, + function_call_prompt, make_valid_bedrock_tool_name, ollama_pt, sanitize_messages_for_tool_calling, @@ -3721,3 +3722,37 @@ def test_convert_to_anthropic_tool_invoke_keeps_paired_server_tool_use(): }, server_result, ] + + +FUNCTION_PROMPT_DESCRIPTIONS: Final = "\n{'name': 'graph_stats'}\n" + + +@pytest.mark.parametrize( + ("messages", "expected_system_contents"), + [ + ([{"role": "user", "content": "hi"}], None), + ([{"role": "system", "content": "Be brief."}, {"role": "user", "content": "hi"}], "Be brief. "), + ( + [{"role": "system", "content": [{"type": "text", "text": "Be brief."}]}, {"role": "user", "content": "hi"}], + [{"type": "text", "text": "Be brief."}], + ), + ], +) +def test_function_call_prompt_returns_new_messages(messages, expected_system_contents): + original: Final = json.loads(json.dumps(messages)) + + result: Final = function_call_prompt(messages=messages, function_descriptions=FUNCTION_PROMPT_DESCRIPTIONS) + + assert messages == original + system_messages: Final = [m for m in result if m["role"] == "system"] + assert len(system_messages) == 1 + content: Final = system_messages[0]["content"] + prompt_text: Final = content if isinstance(content, str) else content[-1]["text"] + assert "Produce JSON OUTPUT ONLY" in prompt_text + assert "graph_stats" in prompt_text + if expected_system_contents is None: + assert result[:-1] == original + elif isinstance(expected_system_contents, str): + assert content.startswith(expected_system_contents) + else: + assert content[:-1] == expected_system_contents diff --git a/tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py b/tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py index 9942f44dcca..509b1e80e72 100644 --- a/tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py +++ b/tests/test_litellm/llms/ollama/test_ollama_completion_transformation.py @@ -727,3 +727,37 @@ async def test_ollama_async_native_tools(legacy_functions: bool) -> None: client=handler, ) assert response.choices[0].message.content == "Hello" + + +def test_ollama_add_function_to_prompt_keeps_legacy_json_emulation(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "add_function_to_prompt", True) + requests = [] + + def handle(request: httpx.Request) -> httpx.Response: + requests.append((request.url.path, json.loads(request.content))) + return httpx.Response( + 200, json={"response": '{"name": "graph_stats", "arguments": {}}', "done": True, "prompt_eval_count": 1} + ) + + messages: Final = [ + {"role": "system", "content": "You are a graph assistant."}, + {"role": "user", "content": "How many nodes does the graph have?"}, + ] + + response: Final = litellm.completion( + model="ollama/qwen3.8:27b", + messages=messages, + tools=GRAPH_STATS_TOOLS, + tool_choice="auto", + api_base="http://ollama.example:11434", + client=HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(handle))), + ) + + assert [path for path, _ in requests] == ["/api/generate"] + body: Final = requests[0][1] + assert body["format"] == "json" + assert "Produce JSON OUTPUT ONLY" in body["prompt"] + assert "graph_stats" in body["prompt"] + assert "prompted_functions" not in body["options"] + assert response.choices[0].message.tool_calls[0].function.name == "graph_stats" + assert response.choices[0].finish_reason == "tool_calls"