mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(databricks): keep #/$defs refs in json_schema response_format (#45659)
* fix(databricks): keep #/$defs refs in json_schema response_format for non-Claude models Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(databricks): type the json_schema response_format override Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(databricks): integration cells for json_schema $defs refs across endpoints Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(databricks): split stacked loop in streaming json_schema cell Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(databricks): mark BaseConfig dict contract on json_schema override Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(databricks): cover json_schema refs on every endpoint, edge refs, sad paths and an outage burst --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
fb547a0956
commit
fa7ad80a62
4 changed files with 1032 additions and 4 deletions
|
|
@ -7,7 +7,7 @@ from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequenc
|
|||
from typing import TYPE_CHECKING, Any, Final, Literal, cast, overload
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
|
||||
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
|
||||
|
|
@ -21,6 +21,9 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
|||
strip_name_from_message,
|
||||
)
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
from litellm.llms.base_llm.base_utils import (
|
||||
type_to_response_format_param, # pyright: ignore[reportUnknownVariableType] # base_utils helper returns an untyped dict
|
||||
)
|
||||
from litellm.types.llms.anthropic import AllAnthropicToolsValues
|
||||
from litellm.types.llms.databricks import (
|
||||
AllDatabricksContentValues,
|
||||
|
|
@ -59,6 +62,10 @@ from ...anthropic.chat.transformation import (
|
|||
from ...openai_like.chat.transformation import OpenAILikeChatConfig
|
||||
from ..common_utils import DatabricksBase, DatabricksException
|
||||
|
||||
_RESPONSE_FORMAT_ADAPTER: ( # mutable-ok: mirrors the dict return contract of get_json_schema_from_pydantic_object
|
||||
Final[TypeAdapter[dict[str, object] | None]]
|
||||
) = TypeAdapter(dict[str, object] | None)
|
||||
|
||||
|
||||
def _is_bare_assistant_message(message_dict: Mapping[str, object]) -> bool:
|
||||
"""Databricks rejects assistant messages with neither content nor tool calls, e.g. a replayed
|
||||
|
|
@ -191,6 +198,14 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig):
|
|||
def get_config(cls, *, model: str | None = None):
|
||||
return super().get_config()
|
||||
|
||||
def get_json_schema_from_pydantic_object(
|
||||
self,
|
||||
response_format: ( # mutable-ok: matches BaseConfig override signature
|
||||
type[BaseModel] | dict[str, object] | None
|
||||
),
|
||||
) -> dict[str, object] | None: # mutable-ok: BaseConfig contract returns a dict
|
||||
return _RESPONSE_FORMAT_ADAPTER.validate_python(type_to_response_format_param(response_format=response_format))
|
||||
|
||||
def get_required_params(self) -> list[ProviderField]:
|
||||
"""For a given provider, return it's required fields with a description"""
|
||||
return [
|
||||
|
|
|
|||
|
|
@ -1,7 +1,9 @@
|
|||
import itertools
|
||||
import json
|
||||
import uuid
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Final, Literal, TypeAlias
|
||||
|
||||
import pytest
|
||||
from integration._support.client import Gateway, eventually
|
||||
|
|
@ -12,6 +14,10 @@ from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter
|
|||
_BACKEND: Final = "databricks-glm-5-2"
|
||||
_API_KEY: Final = "synthetic-databricks-key"
|
||||
_PROMPT: Final = "Summarise the cached briefing in one sentence."
|
||||
_JSON_SCHEMA_BACKEND: Final = "databricks-qwen35-122b-a10b"
|
||||
_CLAUDE_BACKEND: Final = "databricks-claude-haiku-4-5"
|
||||
_JSON_SCHEMA_PROMPT: Final = "Return the requested JSON."
|
||||
_JSON_CONTENT: Final = '{"p":{"name":"Ada"}}'
|
||||
_PROVIDER_USAGE: Final[Mapping[str, JsonValue]] = {
|
||||
"prompt_tokens": 12011,
|
||||
"completion_tokens": 8,
|
||||
|
|
@ -20,6 +26,67 @@ _PROVIDER_USAGE: Final[Mapping[str, JsonValue]] = {
|
|||
"cache_creation_input_tokens": 0,
|
||||
}
|
||||
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
|
||||
_PERSON_SCHEMA: Final[Mapping[str, JsonValue]] = {
|
||||
"$defs": {
|
||||
"Person": {
|
||||
"type": "object",
|
||||
"properties": {"name": {"type": "string"}},
|
||||
"required": ["name"],
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
"properties": {"p": {"$ref": "#/$defs/Person"}},
|
||||
"required": ["p"],
|
||||
}
|
||||
_JSON_SCHEMA_RESPONSE_FORMAT: Final[Mapping[str, JsonValue]] = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {"name": "P", "strict": True, "schema": _PERSON_SCHEMA},
|
||||
}
|
||||
_MESSAGES_PERSON_SCHEMA: Final[Mapping[str, JsonValue]] = {
|
||||
"$defs": {
|
||||
"Person": {
|
||||
"type": "object",
|
||||
"properties": {"name": {"type": "string"}},
|
||||
"required": ["name"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
"properties": {"p": {"$ref": "#/$defs/Person"}},
|
||||
"required": ["p"],
|
||||
"additionalProperties": False,
|
||||
}
|
||||
_MESSAGES_RESPONSE_FORMAT: Final[Mapping[str, JsonValue]] = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": "structured_output",
|
||||
"schema": _MESSAGES_PERSON_SCHEMA,
|
||||
"strict": True,
|
||||
},
|
||||
}
|
||||
_JSON_TOOL_CALL_MESSAGE: Final[dict[str, JsonValue]] = {
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": "json-tool-call",
|
||||
"type": "function",
|
||||
"function": {"name": "json_tool_call", "arguments": _JSON_CONTENT},
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def _chat_completion_response(model: str, message: dict[str, JsonValue], finish_reason: str) -> bytes:
|
||||
response: Final = {
|
||||
"id": "databricks-json-schema-response",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": model,
|
||||
"choices": [{"index": 0, "message": message, "finish_reason": finish_reason}],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
}
|
||||
return json.dumps(response).encode()
|
||||
|
||||
|
||||
class _PromptTokensDetails(BaseModel):
|
||||
|
|
@ -52,12 +119,17 @@ class _Chunk(BaseModel):
|
|||
usage: _Usage | None = None
|
||||
|
||||
|
||||
def _frame(identity: str, choices: list[Mapping[str, object]], usage: Mapping[str, JsonValue] | None = None) -> bytes:
|
||||
def _frame(
|
||||
identity: str,
|
||||
choices: list[Mapping[str, object]],
|
||||
usage: Mapping[str, JsonValue] | None = None,
|
||||
model: str = _BACKEND,
|
||||
) -> bytes:
|
||||
value: Final = {
|
||||
"id": identity,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 1,
|
||||
"model": _BACKEND,
|
||||
"model": model,
|
||||
"choices": choices,
|
||||
**({} if usage is None else {"usage": usage}),
|
||||
}
|
||||
|
|
@ -127,3 +199,672 @@ def test_databricks_stream_final_usage_chunk_reaches_client_and_spend_log(gatewa
|
|||
seconds=70,
|
||||
)
|
||||
assert (rows[0]["prompt_tokens"], rows[0]["completion_tokens"], rows[0]["total_tokens"]) == (12011, 8, 12019)
|
||||
|
||||
|
||||
def test_databricks_json_schema_refs_are_preserved_in_chat_completions(gateway: Gateway) -> None:
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
assert request.target == "/chat/completions"
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
assert body["model"] == _JSON_SCHEMA_BACKEND
|
||||
assert body["messages"] == [{"role": "user", "content": _JSON_SCHEMA_PROMPT}]
|
||||
assert body["stream"] is False
|
||||
assert body["response_format"] == _JSON_SCHEMA_RESPONSE_FORMAT
|
||||
return Reply(
|
||||
body=_chat_completion_response(
|
||||
_JSON_SCHEMA_BACKEND,
|
||||
{"role": "assistant", "content": _JSON_CONTENT},
|
||||
"stop",
|
||||
)
|
||||
)
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"databricks/{_JSON_SCHEMA_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": _JSON_SCHEMA_PROMPT}],
|
||||
"response_format": _JSON_SCHEMA_RESPONSE_FORMAT,
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
body: Final = _JSON_OBJECT.validate_json(response.content)
|
||||
assert body["choices"] == [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": _JSON_CONTENT},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
]
|
||||
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]
|
||||
|
||||
|
||||
def test_databricks_streaming_json_schema_refs_are_preserved(gateway: Gateway) -> None:
|
||||
identity: Final = f"databricks-json-schema-stream-{uuid.uuid4().hex}"
|
||||
frames: Final = (
|
||||
_frame(
|
||||
identity,
|
||||
[{"index": 0, "delta": {"role": "assistant", "content": '{"p":'}, "finish_reason": None}],
|
||||
model=_JSON_SCHEMA_BACKEND,
|
||||
),
|
||||
_frame(
|
||||
identity,
|
||||
[{"index": 0, "delta": {"content": '{"name":"Ada"}}'}, "finish_reason": None}],
|
||||
model=_JSON_SCHEMA_BACKEND,
|
||||
),
|
||||
_frame(identity, [{"index": 0, "delta": {}, "finish_reason": "stop"}], model=_JSON_SCHEMA_BACKEND),
|
||||
b"data: [DONE]\n\n",
|
||||
)
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
assert request.target == "/chat/completions"
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
assert body["model"] == _JSON_SCHEMA_BACKEND
|
||||
assert body["messages"] == [{"role": "user", "content": _JSON_SCHEMA_PROMPT}]
|
||||
assert body["stream"] is True
|
||||
assert body["response_format"] == _JSON_SCHEMA_RESPONSE_FORMAT
|
||||
return Reply(content_type="text/event-stream", chunks=frames)
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"databricks/{_JSON_SCHEMA_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
with gateway.client.stream(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
json={
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": _JSON_SCHEMA_PROMPT}],
|
||||
"response_format": _JSON_SCHEMA_RESPONSE_FORMAT,
|
||||
"stream": True,
|
||||
},
|
||||
headers={"Authorization": f"Bearer {gateway.key}"},
|
||||
) as response:
|
||||
assert response.status_code == 200, response.read()
|
||||
lines: Final = tuple(line for line in response.iter_lines() if line.startswith("data: "))
|
||||
assert lines[-1] == "data: [DONE]", lines
|
||||
chunks: Final = tuple(_Chunk.model_validate_json(line.removeprefix("data: ")) for line in lines[:-1])
|
||||
choices: Final = tuple(itertools.chain.from_iterable(chunk.choices for chunk in chunks))
|
||||
assert "".join(choice.delta.content or "" for choice in choices) == _JSON_CONTENT
|
||||
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]
|
||||
|
||||
|
||||
def test_databricks_claude_json_schema_refs_are_preserved_in_tool_parameters(gateway: Gateway) -> None:
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
assert request.target == "/chat/completions"
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
assert body["model"] == _CLAUDE_BACKEND
|
||||
assert body["messages"] == [{"role": "user", "content": _JSON_SCHEMA_PROMPT}]
|
||||
assert "response_format" not in body
|
||||
assert body["tools"] == [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {"name": "json_tool_call", "parameters": _PERSON_SCHEMA},
|
||||
}
|
||||
]
|
||||
assert body["tool_choice"] == {"type": "function", "function": {"name": "json_tool_call"}}
|
||||
return Reply(body=_chat_completion_response(_CLAUDE_BACKEND, _JSON_TOOL_CALL_MESSAGE, "tool_calls"))
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"databricks/{_CLAUDE_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": _JSON_SCHEMA_PROMPT}],
|
||||
"response_format": _JSON_SCHEMA_RESPONSE_FORMAT,
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
body: Final = _JSON_OBJECT.validate_json(response.content)
|
||||
assert body["choices"] == [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": _JSON_CONTENT},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
]
|
||||
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]
|
||||
|
||||
|
||||
def test_databricks_responses_json_schema_refs_are_preserved_in_chat_bridge(gateway: Gateway) -> None:
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
assert request.target == "/chat/completions"
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
assert body["model"] == _JSON_SCHEMA_BACKEND
|
||||
assert body["messages"] == [{"role": "user", "content": _JSON_SCHEMA_PROMPT}]
|
||||
assert body["response_format"] == {
|
||||
"type": "json_schema",
|
||||
"json_schema": {"name": "P", "schema": _PERSON_SCHEMA, "strict": True},
|
||||
}
|
||||
return Reply(
|
||||
body=_chat_completion_response(
|
||||
_JSON_SCHEMA_BACKEND,
|
||||
{"role": "assistant", "content": _JSON_CONTENT},
|
||||
"stop",
|
||||
)
|
||||
)
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"databricks/{_JSON_SCHEMA_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/responses",
|
||||
{
|
||||
"model": model,
|
||||
"input": _JSON_SCHEMA_PROMPT,
|
||||
"text": {
|
||||
"format": {
|
||||
"type": "json_schema",
|
||||
"name": "P",
|
||||
"strict": True,
|
||||
"schema": _PERSON_SCHEMA,
|
||||
}
|
||||
},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]
|
||||
|
||||
|
||||
def test_databricks_messages_json_schema_refs_are_preserved_for_qwen(gateway: Gateway) -> None:
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
assert request.target == "/chat/completions"
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
assert body["model"] == _JSON_SCHEMA_BACKEND
|
||||
assert body["response_format"] == _MESSAGES_RESPONSE_FORMAT
|
||||
return Reply(
|
||||
body=_chat_completion_response(
|
||||
_JSON_SCHEMA_BACKEND,
|
||||
{"role": "assistant", "content": _JSON_CONTENT},
|
||||
"stop",
|
||||
)
|
||||
)
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"databricks/{_JSON_SCHEMA_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/messages",
|
||||
{
|
||||
"model": model,
|
||||
"max_tokens": 64,
|
||||
"messages": [{"role": "user", "content": _JSON_SCHEMA_PROMPT}],
|
||||
"output_format": {"type": "json_schema", "schema": _PERSON_SCHEMA},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
body: Final = _JSON_OBJECT.validate_json(response.content)
|
||||
assert body["content"] == [{"type": "text", "text": _JSON_CONTENT}]
|
||||
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]
|
||||
|
||||
|
||||
def test_databricks_messages_claude_json_schema_refs_are_preserved_in_tool_parameters(gateway: Gateway) -> None:
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
assert request.target == "/chat/completions"
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
assert body["model"] == _CLAUDE_BACKEND
|
||||
assert "response_format" not in body
|
||||
assert body["tools"] == [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {"name": "json_tool_call", "parameters": _MESSAGES_PERSON_SCHEMA},
|
||||
}
|
||||
]
|
||||
assert body["tool_choice"] == {"type": "function", "function": {"name": "json_tool_call"}}
|
||||
return Reply(body=_chat_completion_response(_CLAUDE_BACKEND, _JSON_TOOL_CALL_MESSAGE, "tool_calls"))
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"databricks/{_CLAUDE_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/messages",
|
||||
{
|
||||
"model": model,
|
||||
"max_tokens": 64,
|
||||
"messages": [{"role": "user", "content": _JSON_SCHEMA_PROMPT}],
|
||||
"output_format": {"type": "json_schema", "schema": _PERSON_SCHEMA},
|
||||
},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
body: Final = _JSON_OBJECT.validate_json(response.content)
|
||||
assert body["content"] == [{"type": "text", "text": _JSON_CONTENT}]
|
||||
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]
|
||||
|
||||
|
||||
_Endpoint: TypeAlias = Literal["chat", "responses", "messages"]
|
||||
_PATHS: Final[Mapping[_Endpoint, str]] = {
|
||||
"chat": "/v1/chat/completions",
|
||||
"responses": "/v1/responses",
|
||||
"messages": "/v1/messages",
|
||||
}
|
||||
_OUTAGE: Final = "upstream-outage"
|
||||
_OUTAGE_MESSAGE: Final = "The service is temporarily unavailable"
|
||||
_EXPECTED_JSON: Final[JsonValue] = _JSON_OBJECT.validate_json(_JSON_CONTENT)
|
||||
_CONTENT_DELTAS: Final[tuple[Mapping[str, JsonValue], ...]] = (
|
||||
{"role": "assistant", "content": '{"p":'},
|
||||
{"content": '{"name":"Ada"}}'},
|
||||
)
|
||||
_TOOL_CALL_DELTAS: Final[tuple[Mapping[str, JsonValue], ...]] = (
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{
|
||||
"index": 0,
|
||||
"id": "json-tool-call",
|
||||
"type": "function",
|
||||
"function": {"name": "json_tool_call", "arguments": ""},
|
||||
}
|
||||
],
|
||||
},
|
||||
{"tool_calls": [{"index": 0, "function": {"arguments": _JSON_CONTENT}}]},
|
||||
)
|
||||
|
||||
|
||||
class _ChatMessage(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
content: str | None = None
|
||||
|
||||
|
||||
class _ChatChoice(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
message: _ChatMessage
|
||||
|
||||
|
||||
class _ChatCompletion(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
choices: tuple[_ChatChoice, ...]
|
||||
|
||||
|
||||
class _OutputPart(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
type: str
|
||||
text: str | None = None
|
||||
|
||||
|
||||
class _OutputItem(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
type: str
|
||||
content: tuple[_OutputPart, ...] = ()
|
||||
name: str | None = None
|
||||
arguments: str | None = None
|
||||
|
||||
|
||||
class _ResponsesBody(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
output: tuple[_OutputItem, ...]
|
||||
|
||||
|
||||
class _ResponsesEvent(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
type: str
|
||||
response: _ResponsesBody | None = None
|
||||
|
||||
|
||||
class _MessagesBlock(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
type: str
|
||||
text: str | None = None
|
||||
|
||||
|
||||
class _MessagesBody(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
content: tuple[_MessagesBlock, ...]
|
||||
|
||||
|
||||
class _MessagesDelta(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
type: str | None = None
|
||||
text: str | None = None
|
||||
partial_json: str | None = None
|
||||
|
||||
|
||||
class _MessagesEvent(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
type: str
|
||||
delta: _MessagesDelta | None = None
|
||||
|
||||
|
||||
class _FunctionDelta(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
arguments: str | None = None
|
||||
|
||||
|
||||
class _ToolCallDelta(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
function: _FunctionDelta | None = None
|
||||
|
||||
|
||||
class _StreamDelta(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
content: str | None = None
|
||||
tool_calls: tuple[_ToolCallDelta, ...] | None = None
|
||||
|
||||
|
||||
class _StreamChoice(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
delta: _StreamDelta
|
||||
|
||||
|
||||
class _StreamChunk(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
choices: tuple[_StreamChoice, ...]
|
||||
|
||||
|
||||
def _json_schema_format(schema: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
|
||||
return {"type": "json_schema", "json_schema": {"name": "P", "strict": True, "schema": dict(schema)}}
|
||||
|
||||
|
||||
def _client_body(
|
||||
endpoint: _Endpoint,
|
||||
model: str,
|
||||
stream: bool,
|
||||
prompt: str = _JSON_SCHEMA_PROMPT,
|
||||
) -> dict[str, JsonValue]:
|
||||
match endpoint:
|
||||
case "chat":
|
||||
return {
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": prompt}],
|
||||
"response_format": _json_schema_format(_PERSON_SCHEMA),
|
||||
"stream": stream,
|
||||
}
|
||||
case "responses":
|
||||
return {
|
||||
"model": model,
|
||||
"input": prompt,
|
||||
"text": {
|
||||
"format": {"type": "json_schema", "name": "P", "strict": True, "schema": dict(_PERSON_SCHEMA)}
|
||||
},
|
||||
"stream": stream,
|
||||
}
|
||||
case "messages":
|
||||
return {
|
||||
"model": model,
|
||||
"max_tokens": 64,
|
||||
"messages": [{"role": "user", "content": prompt}],
|
||||
"output_format": {"type": "json_schema", "schema": dict(_PERSON_SCHEMA)},
|
||||
"stream": stream,
|
||||
}
|
||||
|
||||
|
||||
def _expected_outbound(endpoint: _Endpoint, backend: str) -> dict[str, JsonValue]:
|
||||
schema: Final = dict(_MESSAGES_PERSON_SCHEMA if endpoint == "messages" else _PERSON_SCHEMA)
|
||||
if "claude" in backend:
|
||||
return _JSON_OBJECT.validate_python(
|
||||
{
|
||||
"tools": [{"type": "function", "function": {"name": "json_tool_call", "parameters": schema}}],
|
||||
"tool_choice": {"type": "function", "function": {"name": "json_tool_call"}},
|
||||
}
|
||||
)
|
||||
name: Final = "structured_output" if endpoint == "messages" else "P"
|
||||
return _JSON_OBJECT.validate_python(
|
||||
{"response_format": {"type": "json_schema", "json_schema": {"name": name, "schema": schema, "strict": True}}}
|
||||
)
|
||||
|
||||
|
||||
def _assert_outbound(request: Request, endpoint: _Endpoint, backend: str, stream: bool) -> None:
|
||||
assert (request.method, request.target) == ("POST", "/chat/completions")
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
expected: Final = _expected_outbound(endpoint, backend)
|
||||
assert {key: body.get(key) for key in expected} == expected
|
||||
assert (body["model"], body.get("stream") is True) == (backend, stream)
|
||||
assert "claude" not in backend or "response_format" not in body
|
||||
|
||||
|
||||
def _upstream_reply(backend: str, stream: bool) -> Reply:
|
||||
claude: Final = "claude" in backend
|
||||
finish_reason: Final = "tool_calls" if claude else "stop"
|
||||
if not stream:
|
||||
message: Final[dict[str, JsonValue]] = (
|
||||
_JSON_TOOL_CALL_MESSAGE if claude else {"role": "assistant", "content": _JSON_CONTENT}
|
||||
)
|
||||
return Reply(body=_chat_completion_response(backend, message, finish_reason))
|
||||
identity: Final = f"databricks-json-schema-{uuid.uuid4().hex}"
|
||||
deltas: Final = _TOOL_CALL_DELTAS if claude else _CONTENT_DELTAS
|
||||
frames: Final = (
|
||||
*(_frame(identity, [{"index": 0, "delta": delta, "finish_reason": None}], model=backend) for delta in deltas),
|
||||
_frame(identity, [{"index": 0, "delta": {}, "finish_reason": finish_reason}], model=backend),
|
||||
b"data: [DONE]\n\n",
|
||||
)
|
||||
return Reply(content_type="text/event-stream", chunks=frames)
|
||||
|
||||
|
||||
def _outage_or_reply(request: Request) -> Reply:
|
||||
if _OUTAGE.encode() in request.body:
|
||||
error: Final = {"error_code": "TEMPORARILY_UNAVAILABLE", "message": _OUTAGE_MESSAGE}
|
||||
return Reply(status=503, body=json.dumps(error).encode())
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
return _upstream_reply(str(body["model"]), body.get("stream") is True)
|
||||
|
||||
|
||||
def _item_text(item: _OutputItem) -> str:
|
||||
match item.type:
|
||||
case "message":
|
||||
return "".join(part.text or "" for part in item.content)
|
||||
case "function_call" if item.name == "json_tool_call":
|
||||
return item.arguments or ""
|
||||
case _:
|
||||
return ""
|
||||
|
||||
|
||||
def _responses_text(body: _ResponsesBody) -> str:
|
||||
return "".join(_item_text(item) for item in body.output)
|
||||
|
||||
|
||||
def _chat_delta_text(delta: _StreamDelta) -> str:
|
||||
calls: Final = delta.tool_calls or ()
|
||||
return (delta.content or "") + "".join(call.function.arguments or "" for call in calls if call.function is not None)
|
||||
|
||||
|
||||
def _messages_delta_text(delta: _MessagesDelta) -> str:
|
||||
match delta.type:
|
||||
case "text_delta":
|
||||
return delta.text or ""
|
||||
case "input_json_delta":
|
||||
return delta.partial_json or ""
|
||||
case _:
|
||||
return ""
|
||||
|
||||
|
||||
def _final_text(endpoint: _Endpoint, content: bytes) -> str:
|
||||
match endpoint:
|
||||
case "chat":
|
||||
return _ChatCompletion.model_validate_json(content).choices[0].message.content or ""
|
||||
case "responses":
|
||||
return _responses_text(_ResponsesBody.model_validate_json(content))
|
||||
case "messages":
|
||||
return "".join(block.text or "" for block in _MessagesBody.model_validate_json(content).content)
|
||||
|
||||
|
||||
def _streamed_text(endpoint: _Endpoint, data: tuple[str, ...]) -> str:
|
||||
match endpoint:
|
||||
case "chat":
|
||||
assert data[-1] == "[DONE]", data
|
||||
chunks: Final = tuple(_StreamChunk.model_validate_json(line) for line in data[:-1])
|
||||
choices: Final = tuple(itertools.chain.from_iterable(chunk.choices for chunk in chunks))
|
||||
return "".join(_chat_delta_text(choice.delta) for choice in choices)
|
||||
case "responses":
|
||||
events: Final = tuple(_ResponsesEvent.model_validate_json(line) for line in data if line != "[DONE]")
|
||||
completed: Final = tuple(event.response for event in events if event.type == "response.completed")
|
||||
assert len(completed) == 1 and completed[0] is not None, data
|
||||
return _responses_text(completed[0])
|
||||
case "messages":
|
||||
message_events: Final = tuple(_MessagesEvent.model_validate_json(line) for line in data)
|
||||
assert message_events[-1].type == "message_stop", data
|
||||
deltas: Final = tuple(event.delta for event in message_events if event.delta is not None)
|
||||
return "".join(_messages_delta_text(delta) for delta in deltas)
|
||||
|
||||
|
||||
def _outcome(status: int, text: str) -> tuple[int, JsonValue]:
|
||||
if status == 200:
|
||||
return status, _JSON_OBJECT.validate_json(text)
|
||||
return status, _OUTAGE_MESSAGE if _OUTAGE_MESSAGE in text else text
|
||||
|
||||
|
||||
def _call(gateway: Gateway, endpoint: _Endpoint, body: Mapping[str, JsonValue]) -> tuple[int, str]:
|
||||
if body["stream"] is not True:
|
||||
response: Final = gateway.request("POST", _PATHS[endpoint], body)
|
||||
return response.status_code, _final_text(endpoint, response.content) if response.is_success else response.text
|
||||
with gateway.client.stream(
|
||||
"POST", _PATHS[endpoint], json=dict(body), headers={"Authorization": f"Bearer {gateway.key}"}
|
||||
) as streamed:
|
||||
if not streamed.is_success:
|
||||
return streamed.status_code, streamed.read().decode()
|
||||
data: Final = tuple(line.removeprefix("data: ") for line in streamed.iter_lines() if line.startswith("data: "))
|
||||
return 200, _streamed_text(endpoint, data)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("endpoint", "backend", "stream"),
|
||||
[
|
||||
pytest.param("chat", _CLAUDE_BACKEND, True, id="chat-stream-claude"),
|
||||
pytest.param("responses", _CLAUDE_BACKEND, False, id="responses-claude"),
|
||||
pytest.param("responses", _JSON_SCHEMA_BACKEND, True, id="responses-stream-qwen"),
|
||||
pytest.param("responses", _CLAUDE_BACKEND, True, id="responses-stream-claude"),
|
||||
pytest.param("messages", _JSON_SCHEMA_BACKEND, True, id="messages-stream-qwen"),
|
||||
pytest.param("messages", _CLAUDE_BACKEND, True, id="messages-stream-claude"),
|
||||
],
|
||||
)
|
||||
def test_databricks_json_schema_refs_are_preserved_on_every_endpoint(
|
||||
gateway: Gateway, endpoint: _Endpoint, backend: str, stream: bool
|
||||
) -> None:
|
||||
with wire_server(lambda _request: _upstream_reply(backend, stream)) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"databricks/{backend}", api_base=wire.url, api_key=_API_KEY)
|
||||
status, text = _call(gateway, endpoint, _client_body(endpoint, model, stream))
|
||||
received: Final = wire.drain()
|
||||
assert len(received) == 1, received
|
||||
_assert_outbound(received[0], endpoint, backend, stream)
|
||||
assert _outcome(status, text) == (200, _EXPECTED_JSON), text
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"response_format",
|
||||
[
|
||||
pytest.param(
|
||||
_json_schema_format(
|
||||
{
|
||||
"definitions": {
|
||||
"Person": {"type": "object", "properties": {"name": {"type": "string"}}, "required": ["name"]}
|
||||
},
|
||||
"type": "object",
|
||||
"properties": {"p": {"$ref": "#/definitions/Person"}},
|
||||
"required": ["p"],
|
||||
}
|
||||
),
|
||||
id="definitions-ref",
|
||||
),
|
||||
pytest.param(
|
||||
_json_schema_format(
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"name": {"type": "string"},
|
||||
"children": {"type": "array", "items": {"$ref": "#"}},
|
||||
},
|
||||
"required": ["name", "children"],
|
||||
}
|
||||
),
|
||||
id="recursive-root-ref",
|
||||
),
|
||||
pytest.param(
|
||||
_json_schema_format({"type": "object", "properties": {"$ref": {"type": "string"}}, "required": ["$ref"]}),
|
||||
id="property-named-ref",
|
||||
),
|
||||
pytest.param(
|
||||
_json_schema_format({"type": "object", "properties": {"name": {"type": "string"}}, "required": ["name"]}),
|
||||
id="flat-schema",
|
||||
),
|
||||
pytest.param({"type": "json_object"}, id="json-object"),
|
||||
],
|
||||
)
|
||||
def test_databricks_response_format_reaches_upstream_as_sent(
|
||||
gateway: Gateway, response_format: dict[str, JsonValue]
|
||||
) -> None:
|
||||
with (
|
||||
wire_server(lambda _request: _upstream_reply(_JSON_SCHEMA_BACKEND, False)) as wire,
|
||||
gateway.scenario() as scenario,
|
||||
):
|
||||
model: Final = scenario.model(model=f"databricks/{_JSON_SCHEMA_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": _JSON_SCHEMA_PROMPT}],
|
||||
"response_format": response_format,
|
||||
},
|
||||
)
|
||||
received: Final = wire.drain()
|
||||
assert response.status_code == 200, response.text
|
||||
assert _final_text("chat", response.content) == _JSON_CONTENT
|
||||
assert [(request.method, request.target) for request in received] == [("POST", "/chat/completions")]
|
||||
assert _JSON_OBJECT.validate_json(received[0].body)["response_format"] == response_format
|
||||
|
||||
|
||||
def test_databricks_json_schema_without_schema_returns_the_upstream_error(gateway: Gateway) -> None:
|
||||
response_format: Final[dict[str, JsonValue]] = {"type": "json_schema", "json_schema": {"name": "P", "strict": True}}
|
||||
upstream_error: Final = {"error_code": "INVALID_PARAMETER_VALUE", "message": "json_schema.schema is required"}
|
||||
|
||||
def respond(_request: Request) -> Reply:
|
||||
return Reply(status=400, body=json.dumps(upstream_error).encode())
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"databricks/{_JSON_SCHEMA_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
response: Final = gateway.request(
|
||||
"POST",
|
||||
"/v1/chat/completions",
|
||||
{
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": _JSON_SCHEMA_PROMPT}],
|
||||
"response_format": response_format,
|
||||
},
|
||||
)
|
||||
received: Final = wire.drain()
|
||||
assert (response.status_code, "json_schema.schema is required" in response.text) == (400, True), response.text
|
||||
assert [(request.method, request.target) for request in received] == [("POST", "/chat/completions")]
|
||||
assert _JSON_OBJECT.validate_json(received[0].body)["response_format"] == response_format
|
||||
|
||||
|
||||
def _assert_burst_outbound(request: Request) -> None:
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
endpoint: Final[_Endpoint] = "messages" if body["response_format"] == _MESSAGES_RESPONSE_FORMAT else "chat"
|
||||
_assert_outbound(request, endpoint, _JSON_SCHEMA_BACKEND, body.get("stream") is True)
|
||||
|
||||
|
||||
def _burst_request(index: int) -> tuple[_Endpoint, bool, bool]:
|
||||
endpoints: Final[tuple[_Endpoint, ...]] = ("chat", "responses", "messages")
|
||||
return endpoints[index % 3], index % 2 == 1, index % 5 == 2
|
||||
|
||||
|
||||
def test_databricks_json_schema_burst_through_an_upstream_outage_keeps_every_ref(gateway: Gateway) -> None:
|
||||
plan: Final = tuple(_burst_request(index) for index in range(24))
|
||||
recovery: Final[tuple[_Endpoint, ...]] = ("chat", "responses", "messages")
|
||||
with wire_server(_outage_or_reply) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(model=f"databricks/{_JSON_SCHEMA_BACKEND}", api_base=wire.url, api_key=_API_KEY)
|
||||
endpoints: Final = tuple(endpoint for endpoint, _stream, _down in plan)
|
||||
bodies: Final = tuple(
|
||||
_client_body(endpoint, model, stream, f"{_JSON_SCHEMA_PROMPT} {index} {_OUTAGE if down else ''}")
|
||||
for index, (endpoint, stream, down) in enumerate(plan)
|
||||
)
|
||||
with ThreadPoolExecutor(max_workers=12) as pool:
|
||||
results: Final = tuple(pool.map(_call, itertools.repeat(gateway), endpoints, bodies))
|
||||
recovered: Final = tuple(
|
||||
_call(gateway, endpoint, _client_body(endpoint, model, False)) for endpoint in recovery
|
||||
)
|
||||
received: Final = wire.drain()
|
||||
assert len(received) == len(plan) + len(recovery), received
|
||||
for request in received:
|
||||
_assert_burst_outbound(request)
|
||||
expected: Final = tuple(
|
||||
(503, _OUTAGE_MESSAGE) if down else (200, _EXPECTED_JSON) for _endpoint, _stream, down in plan
|
||||
)
|
||||
assert tuple(_outcome(*result) for result in results) == expected
|
||||
assert tuple(_outcome(*result) for result in recovered) == ((200, _EXPECTED_JSON),) * len(recovery)
|
||||
|
|
|
|||
186
tests/integration/sdk/test_databricks_json_schema_refs_sdk.py
Normal file
186
tests/integration/sdk/test_databricks_json_schema_refs_sdk.py
Normal file
|
|
@ -0,0 +1,186 @@
|
|||
import json
|
||||
import uuid
|
||||
from collections.abc import Iterator, Mapping
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from pydantic import BaseModel, ConfigDict, Field, JsonValue
|
||||
|
||||
import litellm
|
||||
|
||||
_QWEN_BACKEND: Final = "databricks-qwen35-122b-a10b"
|
||||
_CLAUDE_BACKEND: Final = "databricks-claude-haiku-4-5"
|
||||
_API_KEY: Final = "synthetic-databricks-key"
|
||||
_ORDER_JSON: Final = '{"buyer":{"name":"Ada"},"seller":{"name":"Bo"}}'
|
||||
_MESSAGES: Final[tuple[Mapping[str, str], ...]] = ({"role": "user", "content": "Return the order as JSON."},)
|
||||
|
||||
|
||||
class Person(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
name: str
|
||||
|
||||
|
||||
class Order(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
buyer: Person
|
||||
seller: Person
|
||||
|
||||
|
||||
class _StrictObject(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
additional_properties: bool | None = Field(default=None, alias="additionalProperties")
|
||||
|
||||
|
||||
class _OrderSchema(_StrictObject):
|
||||
defs: Mapping[str, _StrictObject] = Field(alias="$defs")
|
||||
|
||||
|
||||
class _JsonSchema(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
name: str
|
||||
strict: bool | None = None
|
||||
json_schema: dict[str, JsonValue] = Field(alias="schema")
|
||||
|
||||
|
||||
class _ResponseFormat(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
type: str
|
||||
json_schema: _JsonSchema
|
||||
|
||||
|
||||
class _Function(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
name: str
|
||||
parameters: dict[str, JsonValue]
|
||||
|
||||
|
||||
class _Tool(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
function: _Function
|
||||
|
||||
|
||||
class _Outbound(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
model: str
|
||||
response_format: _ResponseFormat | None = None
|
||||
tools: tuple[_Tool, ...] = ()
|
||||
|
||||
|
||||
class _CallerMessage(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
content: str
|
||||
|
||||
|
||||
class _CallerChoice(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
message: _CallerMessage
|
||||
|
||||
|
||||
class _CallerCompletion(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="ignore")
|
||||
choices: tuple[_CallerChoice, ...]
|
||||
|
||||
|
||||
def _completion(backend: str) -> Reply:
|
||||
claude: Final = "claude" in backend
|
||||
tool_call: Final[dict[str, JsonValue]] = {
|
||||
"id": "json-tool-call",
|
||||
"type": "function",
|
||||
"function": {"name": "json_tool_call", "arguments": _ORDER_JSON},
|
||||
}
|
||||
message: Final[dict[str, JsonValue]] = (
|
||||
{"role": "assistant", "content": None, "tool_calls": [tool_call]}
|
||||
if claude
|
||||
else {"role": "assistant", "content": _ORDER_JSON}
|
||||
)
|
||||
body: Final[dict[str, JsonValue]] = {
|
||||
"id": f"chatcmpl-{uuid.uuid4().hex}",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": backend,
|
||||
"choices": [{"index": 0, "message": message, "finish_reason": "tool_calls" if claude else "stop"}],
|
||||
"usage": {"prompt_tokens": 9, "completion_tokens": 7, "total_tokens": 16},
|
||||
}
|
||||
return Reply(body=json.dumps(body).encode())
|
||||
|
||||
|
||||
def _refs(node: JsonValue) -> Iterator[str]:
|
||||
if isinstance(node, list):
|
||||
for item in node:
|
||||
yield from _refs(item)
|
||||
return
|
||||
if not isinstance(node, dict):
|
||||
return
|
||||
for key, value in node.items():
|
||||
if key == "$ref" and isinstance(value, str):
|
||||
yield value
|
||||
else:
|
||||
yield from _refs(value)
|
||||
|
||||
|
||||
def _sent_schema(backend: str, received: tuple[Request, ...]) -> dict[str, JsonValue]:
|
||||
assert [(request.method, request.target) for request in received] == [("POST", "/chat/completions")]
|
||||
outbound: Final = _Outbound.model_validate_json(received[0].body)
|
||||
assert outbound.model == backend
|
||||
if "claude" in backend:
|
||||
assert outbound.response_format is None
|
||||
assert [tool.function.name for tool in outbound.tools] == ["json_tool_call"]
|
||||
return outbound.tools[0].function.parameters
|
||||
assert outbound.response_format is not None
|
||||
assert (outbound.response_format.type, outbound.response_format.json_schema.name) == ("json_schema", "Order")
|
||||
assert outbound.response_format.json_schema.strict is True
|
||||
return outbound.response_format.json_schema.json_schema
|
||||
|
||||
|
||||
def _assert_strict_order_schema(schema: dict[str, JsonValue]) -> None:
|
||||
assert set(_refs(schema)) == {"#/$defs/Person"}, schema
|
||||
parsed: Final = _OrderSchema.model_validate(schema)
|
||||
assert (parsed.additional_properties, parsed.defs["Person"].additional_properties) == (False, False), schema
|
||||
|
||||
|
||||
def _caller_order(response: litellm.ModelResponse) -> Order:
|
||||
return Order.model_validate_json(
|
||||
_CallerCompletion.model_validate_json(response.model_dump_json()).choices[0].message.content
|
||||
)
|
||||
|
||||
|
||||
_BACKENDS: Final = (
|
||||
pytest.param(_QWEN_BACKEND, id="qwen"),
|
||||
pytest.param(_CLAUDE_BACKEND, id="claude"),
|
||||
)
|
||||
_EXPECTED_ORDER: Final = Order(buyer=Person(name="Ada"), seller=Person(name="Bo"))
|
||||
|
||||
|
||||
@pytest.mark.parametrize("backend", _BACKENDS)
|
||||
def test_databricks_pydantic_response_format_sends_strict_defs_refs(backend: str) -> None:
|
||||
with wire_server(lambda _request: _completion(backend)) as wire:
|
||||
response: Final = litellm.completion(
|
||||
model=f"databricks/{backend}",
|
||||
api_base=wire.url,
|
||||
api_key=_API_KEY,
|
||||
messages=[dict(message) for message in _MESSAGES],
|
||||
response_format=Order,
|
||||
num_retries=0,
|
||||
)
|
||||
received: Final = wire.drain()
|
||||
_assert_strict_order_schema(_sent_schema(backend, received))
|
||||
assert isinstance(response, litellm.ModelResponse)
|
||||
assert _caller_order(response) == _EXPECTED_ORDER
|
||||
|
||||
|
||||
@pytest.mark.parametrize("backend", _BACKENDS)
|
||||
async def test_databricks_pydantic_response_format_sends_strict_defs_refs_async(backend: str) -> None:
|
||||
with wire_server(lambda _request: _completion(backend)) as wire:
|
||||
response: Final = await litellm.acompletion(
|
||||
model=f"databricks/{backend}",
|
||||
api_base=wire.url,
|
||||
api_key=_API_KEY,
|
||||
messages=[dict(message) for message in _MESSAGES],
|
||||
response_format=Order,
|
||||
num_retries=0,
|
||||
)
|
||||
received: Final = wire.drain()
|
||||
_assert_strict_order_schema(_sent_schema(backend, received))
|
||||
assert isinstance(response, litellm.ModelResponse)
|
||||
assert _caller_order(response) == _EXPECTED_ORDER
|
||||
|
|
@ -1,3 +1,4 @@
|
|||
import copy
|
||||
import json
|
||||
from collections.abc import Iterator
|
||||
from typing import Final
|
||||
|
|
@ -7,6 +8,7 @@ import httpx
|
|||
import pytest
|
||||
import respx
|
||||
from fastapi.testclient import TestClient
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
import litellm
|
||||
from litellm.caching.llm_caching_handler import LLMClientCache
|
||||
|
|
@ -32,6 +34,25 @@ DATABRICKS_API_BASE: Final = "https://my.workspace.cloud.databricks.com/serving-
|
|||
DATABRICKS_API_KEY: Final = "dapimykey"
|
||||
DATABRICKS_CHAT_COMPLETIONS_URL: Final = f"{DATABRICKS_API_BASE}/chat/completions"
|
||||
DATABRICKS_EMBEDDINGS_URL: Final = f"{DATABRICKS_API_BASE}/embeddings"
|
||||
JSON_SCHEMA_RESPONSE_FORMAT: Final = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": "P",
|
||||
"strict": True,
|
||||
"schema": {
|
||||
"$defs": {
|
||||
"Person": {
|
||||
"type": "object",
|
||||
"properties": {"name": {"type": "string"}},
|
||||
"required": ["name"],
|
||||
}
|
||||
},
|
||||
"type": "object",
|
||||
"properties": {"p": {"$ref": "#/$defs/Person"}},
|
||||
"required": ["p"],
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
|
|
@ -2088,3 +2109,68 @@ def mock_chat_streaming_response_chunks() -> List[str]:
|
|||
}
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
(
|
||||
"databricks-meta-llama-3-3-70b-instruct",
|
||||
"databricks-qwen35-122b-a10b",
|
||||
"databricks-gpt-oss-120b",
|
||||
),
|
||||
)
|
||||
def test_databricks_non_claude_json_schema_preserves_refs(model: str) -> None:
|
||||
optional_params: Final = litellm.utils.get_optional_params(
|
||||
model=model,
|
||||
custom_llm_provider="databricks",
|
||||
response_format=copy.deepcopy(JSON_SCHEMA_RESPONSE_FORMAT),
|
||||
)
|
||||
|
||||
assert optional_params["response_format"] == JSON_SCHEMA_RESPONSE_FORMAT
|
||||
|
||||
|
||||
def test_databricks_claude_json_schema_preserves_refs_in_tool_parameters() -> None:
|
||||
optional_params: Final = litellm.utils.get_optional_params(
|
||||
model="databricks-claude-haiku-4-5",
|
||||
custom_llm_provider="databricks",
|
||||
response_format=copy.deepcopy(JSON_SCHEMA_RESPONSE_FORMAT),
|
||||
)
|
||||
|
||||
assert "response_format" not in optional_params
|
||||
assert optional_params["tools"] == [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "json_tool_call",
|
||||
"parameters": JSON_SCHEMA_RESPONSE_FORMAT["json_schema"]["schema"],
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
|
||||
def test_databricks_pydantic_json_schema_preserves_nested_refs() -> None:
|
||||
class Person(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
name: str
|
||||
|
||||
class Order(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
|
||||
person: Person
|
||||
|
||||
optional_params: Final = litellm.utils.get_optional_params(
|
||||
model="databricks-meta-llama-3-3-70b-instruct",
|
||||
custom_llm_provider="databricks",
|
||||
response_format=Order,
|
||||
)
|
||||
schema: Final = optional_params["response_format"]["json_schema"]["schema"]
|
||||
|
||||
assert schema["properties"]["person"] == {"$ref": "#/$defs/Person"}
|
||||
assert schema["$defs"]["Person"] == {
|
||||
"additionalProperties": False,
|
||||
"properties": {"name": {"title": "Name", "type": "string"}},
|
||||
"required": ["name"],
|
||||
"title": "Person",
|
||||
"type": "object",
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue