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:
mrinal-berri 2026-10-09 18:51:36 -07:00 • committed by GitHub
parent fb547a0956
commit fa7ad80a62
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 1032 additions and 4 deletions

View file

@ -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 [

View file

@ -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)

View 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

View file

@ -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",
}