From fa7ad80a622afb8fe64dce9365abd4358ae834b1 Mon Sep 17 00:00:00 2001 From: mrinal-berri Date: Fri, 9 Oct 2026 18:51:36 -0700 Subject: [PATCH] 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> --- .../llms/databricks/chat/transformation.py | 17 +- .../providers/test_databricks_chat_wire.py | 747 +++++++++++++++++- .../test_databricks_json_schema_refs_sdk.py | 186 +++++ .../test_databricks_chat_transformation.py | 86 ++ 4 files changed, 1032 insertions(+), 4 deletions(-) create mode 100644 tests/integration/sdk/test_databricks_json_schema_refs_sdk.py diff --git a/litellm/llms/databricks/chat/transformation.py b/litellm/llms/databricks/chat/transformation.py index 89b614d6258..50d3b82a989 100644 --- a/litellm/llms/databricks/chat/transformation.py +++ b/litellm/llms/databricks/chat/transformation.py @@ -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 [ diff --git a/tests/integration/providers/test_databricks_chat_wire.py b/tests/integration/providers/test_databricks_chat_wire.py index 614382a77f0..f6192d32c31 100644 --- a/tests/integration/providers/test_databricks_chat_wire.py +++ b/tests/integration/providers/test_databricks_chat_wire.py @@ -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) diff --git a/tests/integration/sdk/test_databricks_json_schema_refs_sdk.py b/tests/integration/sdk/test_databricks_json_schema_refs_sdk.py new file mode 100644 index 00000000000..009679fa4b5 --- /dev/null +++ b/tests/integration/sdk/test_databricks_json_schema_refs_sdk.py @@ -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 diff --git a/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py b/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py index c8e92145ab3..a366019d310 100644 --- a/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py +++ b/tests/unit/llms/databricks/chat/test_databricks_chat_transformation.py @@ -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", + }