From b79b01e38afbe740530efb294c98e116d2bf9f6c Mon Sep 17 00:00:00 2001 From: Tin Chi Lo Date: Tue, 21 Jul 2026 13:59:23 -0700 Subject: [PATCH] fix(proxy): translate custom tool grammar formats and tool_choice across API surfaces Cursor's ApplyPatch is a grammar-constrained custom tool; the Responses surface carries the grammar flat while chat completions wraps the same fields in a grammar object, so the nested envelope from the previous commit still 400d at OpenAI (tools[N].custom.format.grammar). Adds a shared flat to nested format helper pair in prompt_templates/common_utils used by the cursor messages arm and the chat-to-responses bridge, nests flat Responses-style tool_choice objects on the cursor arm, flattens chat custom tool_choice on the chat-to-responses bridge, and maps custom tool_choice to function tool_choice on the responses-to-chat bridge to match that bridge's custom-to-function tool downgrade --- .../transformation.py | 29 +++--- .../prompt_templates/common_utils.py | 25 +++++ .../proxy/response_api_endpoints/endpoints.py | 28 +++++- .../transformation.py | 6 ++ ...responses_transformation_transformation.py | 48 ++++++++++ ...ore_utils_prompt_templates_common_utils.py | 45 +++++++++ .../response_api_endpoints/test_endpoints.py | 96 ++++++++++++++++++- .../test_litellm_completion_responses.py | 21 ++++ 8 files changed, 282 insertions(+), 16 deletions(-) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 75cee42dd55..e1842a62d56 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -155,17 +155,20 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): pass def _normalize_tool_choice_for_responses_api(self, tool_choice: Any) -> Any: - """Chat tool_choice uses function.name; Responses API expects top-level name.""" - if not isinstance(tool_choice, dict) or tool_choice.get("type") != "function": + """Chat tool_choice nests the name under function/custom; Responses API expects top-level name.""" + if not isinstance(tool_choice, dict): + return tool_choice + choice_type = tool_choice.get("type") + if choice_type not in ("function", "custom"): return tool_choice if isinstance(tool_choice.get("name"), str) and tool_choice.get("name"): - # Return only Responses shape so stray chat ``function`` key is not sent upstream. - return {"type": "function", "name": tool_choice["name"]} - fn = tool_choice.get("function") - if isinstance(fn, dict): - fn_name = fn.get("name") - if isinstance(fn_name, str) and fn_name: - return {"type": "function", "name": fn_name} + # Return only Responses shape so stray chat ``function``/``custom`` keys are not sent upstream. + return {"type": choice_type, "name": tool_choice["name"]} + nested = tool_choice.get(choice_type) + if isinstance(nested, dict): + nested_name = nested.get("name") + if isinstance(nested_name, str) and nested_name: + return {"type": choice_type, "name": nested_name} return tool_choice def _handle_raw_dict_response_item(self, item: Dict[str, Any], index: int) -> Tuple[Optional[Any], int]: @@ -896,6 +899,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): ) ) elif tool.get("type") == "custom" and isinstance(tool.get("custom"), dict): + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + convert_custom_tool_format_to_responses_shape, + ) + custom_payload = tool["custom"] flat_custom: CustomToolParam = { "type": "custom", @@ -903,8 +910,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): } if custom_payload.get("description") is not None: flat_custom["description"] = custom_payload["description"] - if custom_payload.get("format") is not None: - flat_custom["format"] = custom_payload["format"] + if isinstance(custom_payload.get("format"), dict): + flat_custom["format"] = convert_custom_tool_format_to_responses_shape(custom_payload["format"]) responses_tools.append(flat_custom) else: responses_tools.append(tool) # type: ignore diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index c43089950ee..3a7a710c6a9 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -1252,6 +1252,31 @@ def is_function_call(optional_params: dict) -> bool: return False +def convert_custom_tool_format_to_chat_shape(format_obj: dict) -> dict: + """ + Responses API grammar formats are flat ({"type": "grammar", "definition", "syntax"}); + Chat Completions wraps the same fields in a "grammar" object. Text formats are + identical on both surfaces and pass through, as does anything unrecognized. + """ + if format_obj.get("type") == "grammar" and "grammar" not in format_obj: + return { + "type": "grammar", + "grammar": {k: format_obj[k] for k in ("definition", "syntax") if k in format_obj}, + } + return format_obj + + +def convert_custom_tool_format_to_responses_shape(format_obj: dict) -> dict: + """ + Inverse of convert_custom_tool_format_to_chat_shape: unwrap the Chat Completions + "grammar" object into the flat Responses API grammar shape. + """ + grammar = format_obj.get("grammar") + if format_obj.get("type") == "grammar" and isinstance(grammar, dict): + return {"type": "grammar", **{k: grammar[k] for k in ("definition", "syntax") if k in grammar}} + return format_obj + + def get_file_ids_from_messages(messages: List[AllMessageValues]) -> List[str]: """ Gets file ids from messages diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index f9b2cc79f73..7b64bcda7ce 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -29,10 +29,17 @@ _FLAT_FUNCTION_TOOL_KEYS = ("name", "description", "parameters", "strict") def _nest_flat_chat_tool(tool: object) -> object: + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + convert_custom_tool_format_to_chat_shape, + ) + if not isinstance(tool, dict) or "name" not in tool: return tool if tool.get("type") == "custom" and "custom" not in tool: - return {"type": "custom", "custom": {k: tool[k] for k in _FLAT_CUSTOM_TOOL_KEYS if k in tool}} + payload = {k: tool[k] for k in _FLAT_CUSTOM_TOOL_KEYS if k in tool} + if isinstance(payload.get("format"), dict): + payload = {**payload, "format": convert_custom_tool_format_to_chat_shape(payload["format"])} + return {"type": "custom", "custom": payload} if tool.get("type") == "function" and "function" not in tool: return {"type": "function", "function": {k: tool[k] for k in _FLAT_FUNCTION_TOOL_KEYS if k in tool}} return tool @@ -42,6 +49,16 @@ def _nest_flat_chat_tools(tools: list) -> list: return [_nest_flat_chat_tool(tool) for tool in tools] +def _nest_flat_chat_tool_choice(tool_choice: object) -> object: + if not isinstance(tool_choice, dict) or "name" not in tool_choice: + return tool_choice + if tool_choice.get("type") == "custom" and "custom" not in tool_choice: + return {"type": "custom", "custom": {"name": tool_choice["name"]}} + if tool_choice.get("type") == "function" and "function" not in tool_choice: + return {"type": "function", "function": {"name": tool_choice["name"]}} + return tool_choice + + @router.post( "/v1/responses", dependencies=[Depends(user_api_key_auth)], @@ -391,10 +408,17 @@ async def cursor_chat_completions( # Genuine chat completions body (Cursor sends these for models whose BYOK it # already fixed); delegate so behavior matches /chat/completions exactly tools = data.get("tools") + tool_choice = data.get("tool_choice") + normalized: dict = {} if isinstance(tools, list): nested_tools = _nest_flat_chat_tools(tools) if nested_tools != tools: - _safe_set_request_parsed_body(request=request, parsed_body={**data, "tools": nested_tools}) + normalized["tools"] = nested_tools + nested_tool_choice = _nest_flat_chat_tool_choice(tool_choice) + if nested_tool_choice != tool_choice: + normalized["tool_choice"] = nested_tool_choice + if normalized: + _safe_set_request_parsed_body(request=request, parsed_body={**data, **normalized}) return await chat_completion( request=request, fastapi_response=fastapi_response, diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 6b1ca3564e3..176274d236f 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -162,6 +162,12 @@ class LiteLLMCompletionResponsesConfig: if function_name: return {"type": "function", "function": {"name": function_name}} return "required" + elif tool_choice_type == "custom": + custom = tool_choice.get("custom") + custom_name = tool_choice.get("name") or (custom.get("name") if isinstance(custom, dict) else None) + if custom_name: + return {"type": "function", "function": {"name": custom_name}} + return "required" # Return as-is for unknown formats return tool_choice diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index 767d1631649..64112ed43e8 100644 --- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -2475,6 +2475,15 @@ def test_map_optional_params_tool_choice_chat_nested_to_responses_api(): {"type": "function", "name": "foo"}, ), ({"type": "required"}, {"type": "required"}), + ( + {"type": "custom", "custom": {"name": "ApplyPatch"}}, + {"type": "custom", "name": "ApplyPatch"}, + ), + ( + {"type": "custom", "name": "ApplyPatch"}, + {"type": "custom", "name": "ApplyPatch"}, + ), + ({"type": "custom"}, {"type": "custom"}), ], ) def test_normalize_tool_choice_for_responses_api(tool_choice, expected): @@ -3249,3 +3258,42 @@ def test_convert_tools_to_responses_format_flattens_custom_tool_without_optional handler = LiteLLMResponsesTransformationHandler() converted = handler._convert_tools_to_responses_format([{"type": "custom", "custom": {"name": "Minimal"}}]) assert converted[0] == {"type": "custom", "name": "Minimal"} + + +def test_convert_tools_to_responses_format_unwraps_nested_grammar_format(): + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + LiteLLMResponsesTransformationHandler, + ) + + handler = LiteLLMResponsesTransformationHandler() + converted = handler._convert_tools_to_responses_format( + [ + { + "type": "custom", + "custom": { + "name": "ApplyPatch", + "format": { + "type": "grammar", + "grammar": {"definition": "start: patch", "syntax": "lark"}, + }, + }, + } + ] + ) + assert converted[0] == { + "type": "custom", + "name": "ApplyPatch", + "format": {"type": "grammar", "definition": "start: patch", "syntax": "lark"}, + } + + +def test_convert_tools_to_responses_format_text_format_passes_through(): + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + LiteLLMResponsesTransformationHandler, + ) + + handler = LiteLLMResponsesTransformationHandler() + converted = handler._convert_tools_to_responses_format( + [{"type": "custom", "custom": {"name": "A", "format": {"type": "text"}}}] + ) + assert converted[0] == {"type": "custom", "name": "A", "format": {"type": "text"}} diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py index 1b1db634ed2..3728cc80323 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_common_utils.py @@ -721,3 +721,48 @@ class TestUnpackLegacyDefs: out = unpack_legacy_defs(schema) assert "components" not in out assert out["properties"]["r0"]["properties"]["p0"] == {"type": "string"} + + +class TestCustomToolFormatShapeConversion: + def test_flat_grammar_to_chat_shape(self): + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + convert_custom_tool_format_to_chat_shape, + ) + + assert convert_custom_tool_format_to_chat_shape( + {"type": "grammar", "definition": "start: patch", "syntax": "lark"} + ) == {"type": "grammar", "grammar": {"definition": "start: patch", "syntax": "lark"}} + + def test_nested_grammar_to_responses_shape(self): + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + convert_custom_tool_format_to_responses_shape, + ) + + assert convert_custom_tool_format_to_responses_shape( + {"type": "grammar", "grammar": {"definition": "start: patch", "syntax": "regex"}} + ) == {"type": "grammar", "definition": "start: patch", "syntax": "regex"} + + def test_both_directions_are_idempotent_and_pass_text_through(self): + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + convert_custom_tool_format_to_chat_shape, + convert_custom_tool_format_to_responses_shape, + ) + + flat = {"type": "grammar", "definition": "d", "syntax": "lark"} + nested = {"type": "grammar", "grammar": {"definition": "d", "syntax": "lark"}} + text = {"type": "text"} + assert convert_custom_tool_format_to_chat_shape(nested) == nested + assert convert_custom_tool_format_to_responses_shape(flat) == flat + assert convert_custom_tool_format_to_chat_shape(text) == text + assert convert_custom_tool_format_to_responses_shape(text) == text + assert convert_custom_tool_format_to_chat_shape(convert_custom_tool_format_to_responses_shape(nested)) == nested + + def test_unrecognized_formats_pass_through(self): + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + convert_custom_tool_format_to_chat_shape, + convert_custom_tool_format_to_responses_shape, + ) + + for weird in ({}, {"type": "grammar"}, {"type": "future_format", "x": 1}): + assert convert_custom_tool_format_to_chat_shape(dict(weird)) in (weird, {"type": "grammar", "grammar": {}}) + assert convert_custom_tool_format_to_responses_shape(dict(weird)) == weird diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py index 4ef338b35e4..86fafa40811 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py @@ -981,9 +981,18 @@ class TestCursorMessagesArmToolNormalization: "type": "function", "function": {"name": "read_file", "parameters": {"type": "object"}}, }, - {"type": "custom", "name": "ApplyPatch", "description": "V4A patch"}, + { + "type": "custom", + "name": "ApplyPatch", + "description": "V4A patch", + "format": { + "type": "grammar", + "definition": "start: patch", + "syntax": "lark", + }, + }, ], - "tool_choice": "required", + "tool_choice": {"type": "custom", "name": "ApplyPatch"}, }, headers={"Authorization": "Bearer sk-1234"}, ) @@ -993,8 +1002,19 @@ class TestCursorMessagesArmToolNormalization: assert response.status_code == 200 assert seen["body"]["tools"] == [ {"type": "function", "function": {"name": "read_file", "parameters": {"type": "object"}}}, - {"type": "custom", "custom": {"name": "ApplyPatch", "description": "V4A patch"}}, + { + "type": "custom", + "custom": { + "name": "ApplyPatch", + "description": "V4A patch", + "format": { + "type": "grammar", + "grammar": {"definition": "start: patch", "syntax": "lark"}, + }, + }, + }, ] + assert seen["body"]["tool_choice"] == {"type": "custom", "custom": {"name": "ApplyPatch"}} assert seen["body"]["messages"] == [{"role": "user", "content": "use ApplyPatch"}] @pytest.mark.asyncio @@ -1030,3 +1050,73 @@ class TestCursorMessagesArmToolNormalization: assert response.status_code == 200 assert seen["body"]["tools"] == body["tools"] assert seen["body"]["messages"] == body["messages"] + + +class TestNestFlatChatToolGrammarFormat: + def test_flat_grammar_format_is_wrapped_for_chat(self): + from litellm.proxy.response_api_endpoints.endpoints import _nest_flat_chat_tools + + result = _nest_flat_chat_tools( + [ + { + "type": "custom", + "name": "ApplyPatch", + "description": "V4A patch", + "format": {"type": "grammar", "definition": "start: patch", "syntax": "lark"}, + } + ] + ) + assert result == [ + { + "type": "custom", + "custom": { + "name": "ApplyPatch", + "description": "V4A patch", + "format": { + "type": "grammar", + "grammar": {"definition": "start: patch", "syntax": "lark"}, + }, + }, + } + ] + + def test_flat_text_format_is_copied_unchanged(self): + from litellm.proxy.response_api_endpoints.endpoints import _nest_flat_chat_tools + + result = _nest_flat_chat_tools( + [{"type": "custom", "name": "A", "format": {"type": "text"}}] + ) + assert result == [{"type": "custom", "custom": {"name": "A", "format": {"type": "text"}}}] + + +class TestNestFlatChatToolChoice: + def test_flat_custom_tool_choice_is_nested(self): + from litellm.proxy.response_api_endpoints.endpoints import _nest_flat_chat_tool_choice + + assert _nest_flat_chat_tool_choice({"type": "custom", "name": "ApplyPatch"}) == { + "type": "custom", + "custom": {"name": "ApplyPatch"}, + } + + def test_flat_function_tool_choice_is_nested(self): + from litellm.proxy.response_api_endpoints.endpoints import _nest_flat_chat_tool_choice + + assert _nest_flat_chat_tool_choice({"type": "function", "name": "f"}) == { + "type": "function", + "function": {"name": "f"}, + } + + def test_non_flat_tool_choice_values_pass_through(self): + from litellm.proxy.response_api_endpoints.endpoints import _nest_flat_chat_tool_choice + + for unchanged in ( + "auto", + "required", + None, + {"type": "custom", "custom": {"name": "x"}}, + {"type": "function", "function": {"name": "f"}}, + {"type": "auto"}, + {"name": "typeless"}, + 42, + ): + assert _nest_flat_chat_tool_choice(unchanged) == unchanged diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py index d8e3f495ced..3f3f51f0d3f 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -959,6 +959,27 @@ class TestToolChoiceTransformation: ) assert result == {"type": "function", "function": {"name": "get_weather"}} + def test_transform_tool_choice_custom_follows_function_downgrade(self): + """ + This bridge downgrades custom tools to function tools + (convert_custom_tool_to_function_tool), so a custom tool_choice must become a + function tool_choice naming the same tool or it references a tool type absent + from the converted request. + """ + flat = LiteLLMCompletionResponsesConfig._transform_tool_choice( + {"type": "custom", "name": "ApplyPatch"} + ) + assert flat == {"type": "function", "function": {"name": "ApplyPatch"}} + + nested = LiteLLMCompletionResponsesConfig._transform_tool_choice( + {"type": "custom", "custom": {"name": "ApplyPatch"}} + ) + assert nested == {"type": "function", "function": {"name": "ApplyPatch"}} + + def test_transform_tool_choice_custom_without_name_falls_back_to_required(self): + result = LiteLLMCompletionResponsesConfig._transform_tool_choice({"type": "custom"}) + assert result == "required" + def test_transform_tool_choice_function_without_name_falls_back_to_required(self): """A function-type dict with no name still falls back to required""" result = LiteLLMCompletionResponsesConfig._transform_tool_choice(