mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
* Fixed Json.dumps in JSON Schema Validation Error * Added Response Schema to Ollama chat for structured response * Added Test cases * refactor(ollama): remove redundant response_format check The response_format parameter conversion is already handled in utils.py's get_optional_params function, making the duplicate check in ollama_chat.py unnecessary. This change removes the redundant code while maintaining the same functionality.
This commit is contained in:
parent
b7fc72628c
commit
0eb0cf4515
3 changed files with 112 additions and 13 deletions
|
|
@ -17,7 +17,7 @@ def validate_schema(schema: dict, response: str):
|
|||
response_dict = json.loads(response)
|
||||
except json.JSONDecodeError:
|
||||
raise JSONSchemaValidationError(
|
||||
model="", llm_provider="", raw_response=response, schema=response
|
||||
model="", llm_provider="", raw_response=response, schema=json.dumps(schema)
|
||||
)
|
||||
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ from typing import Any, List, Optional, Union
|
|||
import aiohttp
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
import inspect
|
||||
|
||||
import litellm
|
||||
from litellm import verbose_logger
|
||||
|
|
@ -141,6 +142,13 @@ class OllamaChatConfig(OpenAIGPTConfig):
|
|||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
value = non_default_params["response_format"]
|
||||
if inspect.isclass(value) and issubclass(value, BaseModel):
|
||||
non_default_params["response_format"] = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {"schema": value.model_json_schema()}
|
||||
}
|
||||
|
||||
for param, value in non_default_params.items():
|
||||
if param == "max_tokens" or param == "max_completion_tokens":
|
||||
optional_params["num_predict"] = value
|
||||
|
|
@ -156,13 +164,13 @@ class OllamaChatConfig(OpenAIGPTConfig):
|
|||
optional_params["repeat_penalty"] = value
|
||||
if param == "stop":
|
||||
optional_params["stop"] = value
|
||||
if param == "response_format" and value["type"] == "json_object":
|
||||
if param == "response_format" and isinstance(value, dict) and value.get("type") == "json_object":
|
||||
optional_params["format"] = "json"
|
||||
if param == "response_format" and value["type"] == "json_schema":
|
||||
optional_params["format"] = value["json_schema"]["schema"]
|
||||
if param == "response_format" and isinstance(value, dict) and value.get("type") == "json_schema":
|
||||
if value.get("json_schema") and value["json_schema"].get("schema"):
|
||||
optional_params["format"] = value["json_schema"]["schema"]
|
||||
### FUNCTION CALLING LOGIC ###
|
||||
if param == "tools":
|
||||
# ollama actually supports json output
|
||||
## CHECK IF MODEL SUPPORTS TOOL CALLING ##
|
||||
try:
|
||||
model_info = litellm.get_model_info(
|
||||
|
|
@ -185,14 +193,23 @@ class OllamaChatConfig(OpenAIGPTConfig):
|
|||
][0]["function"]["name"]
|
||||
|
||||
if param == "functions":
|
||||
# ollama actually supports json output
|
||||
optional_params["format"] = "json"
|
||||
litellm.add_function_to_prompt = (
|
||||
True # so that main.py adds the function call to the prompt
|
||||
)
|
||||
optional_params["functions_unsupported_model"] = non_default_params.get(
|
||||
"functions"
|
||||
)
|
||||
## CHECK IF MODEL SUPPORTS TOOL CALLING ##
|
||||
try:
|
||||
model_info = litellm.get_model_info(
|
||||
model=model, custom_llm_provider="ollama"
|
||||
)
|
||||
if model_info.get("supports_function_calling") is True:
|
||||
optional_params["tools"] = value
|
||||
else:
|
||||
raise Exception
|
||||
except Exception:
|
||||
optional_params["format"] = "json"
|
||||
litellm.add_function_to_prompt = (
|
||||
True # so that main.py adds the function call to the prompt
|
||||
)
|
||||
optional_params["functions_unsupported_model"] = non_default_params.get(
|
||||
"functions"
|
||||
)
|
||||
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
|
||||
|
|
|
|||
82
tests/litellm/llms/ollama/test_ollama_chat_transformation.py
Normal file
82
tests/litellm/llms/ollama/test_ollama_chat_transformation.py
Normal file
|
|
@ -0,0 +1,82 @@
|
|||
import os
|
||||
import sys
|
||||
import pytest
|
||||
from pydantic import BaseModel
|
||||
import inspect
|
||||
|
||||
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '../../../../..')))
|
||||
|
||||
from litellm.llms.ollama_chat import OllamaChatConfig
|
||||
|
||||
class TestEvent(BaseModel):
|
||||
name: str
|
||||
value: int
|
||||
|
||||
class TestOllamaChatConfigResponseFormat:
|
||||
def test_map_openai_params_with_pydantic_model(self):
|
||||
config = OllamaChatConfig()
|
||||
|
||||
non_default_params = {
|
||||
"response_format": TestEvent
|
||||
}
|
||||
optional_params = {}
|
||||
|
||||
expected_schema_structure = TestEvent.model_json_schema()
|
||||
|
||||
config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model="ollama_chat/test-model",
|
||||
drop_params=False
|
||||
)
|
||||
|
||||
assert "format" in optional_params, "Transformed 'format' key not found in optional_params"
|
||||
|
||||
transformed_format = optional_params["format"]
|
||||
|
||||
assert transformed_format == expected_schema_structure, \
|
||||
f"Transformed schema does not match expected. Got: {transformed_format}, Expected: {expected_schema_structure}"
|
||||
|
||||
def test_map_openai_params_with_dict_json_schema(self):
|
||||
config = OllamaChatConfig()
|
||||
|
||||
direct_schema = TestEvent.model_json_schema()
|
||||
response_format_dict = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {"schema": direct_schema}
|
||||
}
|
||||
|
||||
non_default_params = {
|
||||
"response_format": response_format_dict
|
||||
}
|
||||
optional_params = {}
|
||||
|
||||
config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model="ollama_chat/test-model",
|
||||
drop_params=False
|
||||
)
|
||||
|
||||
assert "format" in optional_params
|
||||
assert optional_params["format"] == direct_schema, \
|
||||
f"Schema from dict did not pass through correctly. Got: {optional_params['format']}, Expected: {direct_schema}"
|
||||
|
||||
def test_map_openai_params_with_json_object(self):
|
||||
config = OllamaChatConfig()
|
||||
|
||||
non_default_params = {
|
||||
"response_format": {"type": "json_object"}
|
||||
}
|
||||
optional_params = {}
|
||||
|
||||
config.map_openai_params(
|
||||
non_default_params=non_default_params,
|
||||
optional_params=optional_params,
|
||||
model="ollama_chat/test-model",
|
||||
drop_params=False
|
||||
)
|
||||
|
||||
assert "format" in optional_params
|
||||
assert optional_params["format"] == "json", \
|
||||
f"Expected 'json' for type 'json_object', got: {optional_params['format']}"
|
||||
Loading…
Add table
Reference in a new issue