mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(utils): OpenAI tool names — ContextVar mapping, provider-gated sanitize, Greptile
- Per-request ContextVar (no global InMemoryCache) for sanitized↔original - Sanitize only for OpenAI-style providers (whitelist in openai_tool_name_mapping) - begin_openai_tool_name_mapping_scope in completion; restore only hits current map - Tests: skip when disabled, sanitize when enabled, context reset, convert restore Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
d99d7dd9ec
commit
4676d54171
6 changed files with 149 additions and 69 deletions
|
|
@ -72,7 +72,7 @@ def _normalize_images_for_message(
|
|||
|
||||
|
||||
def _map_tool_call_dict_openai_names_to_user(tc: dict) -> dict:
|
||||
"""Restore client tool names when LiteLLM sanitized outbound tools for OpenAI."""
|
||||
"""Restore client tool names when this completion rewrote them for OpenAI."""
|
||||
from litellm.litellm_core_utils.openai_tool_name_mapping import (
|
||||
restore_openai_tool_name_for_user,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,36 +1,76 @@
|
|||
"""
|
||||
Mapping from OpenAI-safe tool names back to client-supplied names.
|
||||
Per-request mapping from OpenAI-safe tool names back to client-supplied names.
|
||||
|
||||
OpenAI requires tools[].function.name to match ^[a-zA-Z0-9_-]+$. LiteLLM sanitizes
|
||||
outbound requests and stores sanitized -> original in an in-memory cache (same pattern
|
||||
as litellm.bedrock_tool_name_mappings / make_valid_bedrock_tool_name).
|
||||
OpenAI's Chat Completions API requires tools[].function.name to match
|
||||
``^[a-zA-Z0-9_-]+$``. For providers that enforce this, we rewrite outbound tool
|
||||
names and keep sanitized -> original in a ContextVar dict for this completion only
|
||||
(never a process-wide cache).
|
||||
|
||||
Restore in ``convert_dict_to_response`` only affects names present in the current
|
||||
request's mapping, so other providers and concurrent requests are unaffected.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
import contextvars
|
||||
from typing import Dict, Optional
|
||||
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
_CTX: contextvars.ContextVar[Optional[Dict[str, str]]] = contextvars.ContextVar(
|
||||
"litellm_openai_tool_name_mapping", default=None
|
||||
)
|
||||
|
||||
# Mirrors bedrock_tool_name_mappings in llms/bedrock/chat/invoke_handler.py
|
||||
openai_tool_name_mappings: InMemoryCache = InMemoryCache(
|
||||
max_size_in_memory=50, default_ttl=600
|
||||
# Providers where the upstream Chat Completions API applies OpenAI tool-name validation.
|
||||
_OPENAI_TOOL_NAME_VALIDATION_PROVIDERS = frozenset(
|
||||
{
|
||||
"openai",
|
||||
"azure",
|
||||
"azure_ai",
|
||||
"custom_openai",
|
||||
"text-completion-openai",
|
||||
"groq",
|
||||
"deepinfra",
|
||||
"together_ai",
|
||||
"fireworks_ai",
|
||||
"nvidia_nim",
|
||||
"github_copilot",
|
||||
"perplexity",
|
||||
"xai",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def get_openai_tool_name(response_tool_name: str) -> str:
|
||||
"""
|
||||
If LiteLLM sanitized the outbound tool name, map the API response name back to the original.
|
||||
def should_sanitize_openai_tool_names(litellm_provider: str) -> bool:
|
||||
return litellm_provider in _OPENAI_TOOL_NAME_VALIDATION_PROVIDERS
|
||||
|
||||
Same idea as get_bedrock_tool_name for Bedrock toolSpec names.
|
||||
"""
|
||||
if response_tool_name in openai_tool_name_mappings.cache_dict:
|
||||
response_tool_name = openai_tool_name_mappings.cache_dict[response_tool_name]
|
||||
|
||||
def begin_openai_tool_name_mapping_scope() -> None:
|
||||
"""Reset mapping for this completion (call once at the start of completion())."""
|
||||
_CTX.set({})
|
||||
|
||||
|
||||
def _store(sanitized: str, original: str) -> None:
|
||||
if sanitized == original:
|
||||
return
|
||||
m = _CTX.get()
|
||||
if m is None:
|
||||
m = {}
|
||||
_CTX.set(m)
|
||||
m[sanitized] = original
|
||||
|
||||
|
||||
def get_openai_tool_name(response_tool_name: str) -> str:
|
||||
m = _CTX.get()
|
||||
if m is not None and response_tool_name in m:
|
||||
return m[response_tool_name]
|
||||
return response_tool_name
|
||||
|
||||
|
||||
def restore_openai_tool_name_for_user(sanitized_name: Optional[str]) -> Optional[str]:
|
||||
"""Nullable wrapper used when normalizing tool_calls from responses."""
|
||||
if sanitized_name is None:
|
||||
return None
|
||||
return get_openai_tool_name(sanitized_name)
|
||||
|
||||
|
||||
def store_openai_tool_name_mapping(sanitized: str, original: str) -> None:
|
||||
"""Record a rewrite when ``validate_and_fix_openai_tools`` sanitizes a name."""
|
||||
_store(sanitized, original)
|
||||
|
|
|
|||
|
|
@ -2,6 +2,9 @@ import base64
|
|||
import time
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union, cast
|
||||
|
||||
from litellm.litellm_core_utils.openai_tool_name_mapping import (
|
||||
restore_openai_tool_name_for_user,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
ChatCompletionAssistantContentValue,
|
||||
ChatCompletionAudioDelta,
|
||||
|
|
|
|||
|
|
@ -160,6 +160,10 @@ from litellm.utils import (
|
|||
|
||||
from ._logging import verbose_logger
|
||||
from .caching.caching import disable_cache, enable_cache, update_cache
|
||||
from .litellm_core_utils.openai_tool_name_mapping import (
|
||||
begin_openai_tool_name_mapping_scope,
|
||||
should_sanitize_openai_tool_names,
|
||||
)
|
||||
from .litellm_core_utils.core_helpers import safe_deep_copy
|
||||
from .litellm_core_utils.fallback_utils import (
|
||||
async_completion_with_fallbacks,
|
||||
|
|
@ -1160,7 +1164,24 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
raise ValueError("model param not passed in.")
|
||||
# validate messages
|
||||
messages = validate_and_fix_openai_messages(messages=messages)
|
||||
tools = validate_and_fix_openai_tools(tools=tools)
|
||||
begin_openai_tool_name_mapping_scope()
|
||||
try:
|
||||
_, _litellm_provider_for_tools, _, _ = get_llm_provider(
|
||||
model=model,
|
||||
custom_llm_provider=kwargs.get("custom_llm_provider"),
|
||||
api_base=kwargs.get("api_base") or base_url,
|
||||
api_key=kwargs.get("api_key") or api_key,
|
||||
litellm_params=kwargs.get("litellm_params"),
|
||||
)
|
||||
_sanitize_openai_fn_tool_names = should_sanitize_openai_tool_names(
|
||||
_litellm_provider_for_tools
|
||||
)
|
||||
except Exception:
|
||||
_sanitize_openai_fn_tool_names = False
|
||||
tools = validate_and_fix_openai_tools(
|
||||
tools=tools,
|
||||
sanitize_openai_function_tool_names=_sanitize_openai_fn_tool_names,
|
||||
)
|
||||
# validate tool_choice
|
||||
tool_choice = validate_chat_completion_tool_choice(tool_choice=tool_choice)
|
||||
# validate optional params
|
||||
|
|
|
|||
|
|
@ -7880,7 +7880,6 @@ _OPENAI_FUNCTION_TOOL_NAME_MAX_LENGTH = 64
|
|||
def _sanitize_openai_function_tool_name(name: str, index: int) -> str:
|
||||
"""
|
||||
Normalize function.name to match OpenAI's pattern ^[a-zA-Z0-9_-]+$ and length cap.
|
||||
See: OpenAI Chat Completions tools[].function.name validation.
|
||||
"""
|
||||
if name is None or (isinstance(name, str) and not str(name).strip()):
|
||||
return f"litellm_unnamed_tool_{index}"
|
||||
|
|
@ -7891,7 +7890,6 @@ def _sanitize_openai_function_tool_name(name: str, index: int) -> str:
|
|||
|
||||
|
||||
def _make_unique_openai_tool_name(base: str, used_names: set[str]) -> str:
|
||||
"""Disambiguate sanitized names when multiple tools map to the same string."""
|
||||
candidate = base
|
||||
n = 0
|
||||
while candidate in used_names:
|
||||
|
|
@ -7910,7 +7908,7 @@ def _maybe_fix_openai_function_tool_name(
|
|||
tool_dict: dict, index: int, used_names: set[str]
|
||||
) -> None:
|
||||
from litellm.litellm_core_utils.openai_tool_name_mapping import (
|
||||
openai_tool_name_mappings,
|
||||
store_openai_tool_name_mapping,
|
||||
)
|
||||
|
||||
fn = tool_dict.get("function")
|
||||
|
|
@ -7927,19 +7925,32 @@ def _maybe_fix_openai_function_tool_name(
|
|||
unique = _make_unique_openai_tool_name(base, used_names)
|
||||
fn["name"] = unique
|
||||
if unique != raw_original:
|
||||
openai_tool_name_mappings.set_cache(key=unique, value=raw_original)
|
||||
store_openai_tool_name_mapping(unique, raw_original)
|
||||
|
||||
|
||||
def validate_and_fix_openai_tools(tools: Optional[List]) -> Optional[List[dict]]:
|
||||
def validate_and_fix_openai_tools(
|
||||
tools: Optional[List],
|
||||
*,
|
||||
sanitize_openai_function_tool_names: bool = False,
|
||||
) -> Optional[List[dict]]:
|
||||
"""
|
||||
Ensure tools is List[dict] and not List[BaseModel].
|
||||
|
||||
Sanitizes OpenAI function tool names to match ^[a-zA-Z0-9_-]+$ (max 64 chars),
|
||||
disambiguates collisions, and does not mutate caller-provided dicts.
|
||||
When ``sanitize_openai_function_tool_names`` is True (OpenAI-compatible
|
||||
providers only), rewrites function names to ``^[a-zA-Z0-9_-]+$`` (max 64 chars).
|
||||
"""
|
||||
if tools is None:
|
||||
return tools
|
||||
new_tools: List[dict] = []
|
||||
if not sanitize_openai_function_tool_names:
|
||||
new_tools: List[dict] = []
|
||||
for tool in tools:
|
||||
if isinstance(tool, BaseModel):
|
||||
new_tools.append(tool.model_dump())
|
||||
elif isinstance(tool, dict):
|
||||
new_tools.append(tool)
|
||||
return new_tools
|
||||
|
||||
new_tools = []
|
||||
used_names: set[str] = set()
|
||||
for idx, tool in enumerate(tools):
|
||||
if isinstance(tool, BaseModel):
|
||||
|
|
|
|||
|
|
@ -3985,27 +3985,7 @@ class TestValidateAndFixThinkingParam:
|
|||
assert "budget_tokens" not in thinking
|
||||
|
||||
|
||||
def test_validate_and_fix_openai_tools_sanitizes_invalid_names():
|
||||
from litellm.utils import validate_and_fix_openai_tools
|
||||
|
||||
tools_in = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "invalid.name",
|
||||
"description": "x",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
}
|
||||
]
|
||||
original_name = tools_in[0]["function"]["name"]
|
||||
out = validate_and_fix_openai_tools(tools=tools_in)
|
||||
assert out is not None
|
||||
assert out[0]["function"]["name"] == "invalid_name"
|
||||
assert original_name == "invalid.name"
|
||||
|
||||
|
||||
def test_validate_and_fix_openai_tools_dedupes_colliding_names():
|
||||
def test_validate_and_fix_openai_tools_skips_name_rewrite_when_disabled():
|
||||
from litellm.utils import validate_and_fix_openai_tools
|
||||
|
||||
tools_in = [
|
||||
|
|
@ -4015,50 +3995,74 @@ def test_validate_and_fix_openai_tools_dedupes_colliding_names():
|
|||
"name": "a.b",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
},
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "a/b",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
},
|
||||
}
|
||||
]
|
||||
out = validate_and_fix_openai_tools(tools=tools_in)
|
||||
assert out[0]["function"]["name"] == "a_b"
|
||||
assert out[1]["function"]["name"] == "a_b_1"
|
||||
|
||||
|
||||
def test_openai_tool_name_mapping_restore_roundtrip():
|
||||
from litellm.litellm_core_utils.openai_tool_name_mapping import (
|
||||
restore_openai_tool_name_for_user,
|
||||
out = validate_and_fix_openai_tools(
|
||||
tools=tools_in, sanitize_openai_function_tool_names=False
|
||||
)
|
||||
assert out is not None
|
||||
assert out[0]["function"]["name"] == "a.b"
|
||||
|
||||
|
||||
def test_validate_and_fix_openai_tools_sanitizes_when_enabled():
|
||||
from litellm.utils import validate_and_fix_openai_tools
|
||||
|
||||
tools_in = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "invalid.name.at.index.71",
|
||||
"name": "invalid.name",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
}
|
||||
]
|
||||
out = validate_and_fix_openai_tools(tools=tools_in)
|
||||
assert out[0]["function"]["name"] == "invalid_name_at_index_71"
|
||||
original_name = tools_in[0]["function"]["name"]
|
||||
out = validate_and_fix_openai_tools(
|
||||
tools=tools_in, sanitize_openai_function_tool_names=True
|
||||
)
|
||||
assert out is not None
|
||||
assert out[0]["function"]["name"] == "invalid_name"
|
||||
assert original_name == "invalid.name"
|
||||
|
||||
begin_openai_tool_name_mapping_scope,
|
||||
)
|
||||
from litellm.utils import validate_and_fix_openai_tools
|
||||
|
||||
begin_openai_tool_name_mapping_scope()
|
||||
validate_and_fix_openai_tools(
|
||||
tools=[
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "invalid.name.at.index.71",
|
||||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
}
|
||||
],
|
||||
sanitize_openai_function_tool_names=True,
|
||||
)
|
||||
assert (
|
||||
restore_openai_tool_name_for_user("invalid_name_at_index_71")
|
||||
== "invalid.name.at.index.71"
|
||||
)
|
||||
begin_openai_tool_name_mapping_scope()
|
||||
assert (
|
||||
restore_openai_tool_name_for_user("invalid_name_at_index_71")
|
||||
== "invalid_name_at_index_71"
|
||||
)
|
||||
|
||||
|
||||
def test_convert_to_model_response_restores_openai_tool_call_names():
|
||||
def test_convert_to_model_response_restores_openai_tool_call_names_when_mapped():
|
||||
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
|
||||
convert_to_model_response_object,
|
||||
)
|
||||
from litellm.litellm_core_utils.openai_tool_name_mapping import (
|
||||
begin_openai_tool_name_mapping_scope,
|
||||
)
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.utils import validate_and_fix_openai_tools
|
||||
|
||||
begin_openai_tool_name_mapping_scope()
|
||||
validate_and_fix_openai_tools(
|
||||
tools=[
|
||||
{
|
||||
|
|
@ -4068,7 +4072,8 @@ def test_convert_to_model_response_restores_openai_tool_call_names():
|
|||
"parameters": {"type": "object", "properties": {}},
|
||||
},
|
||||
}
|
||||
]
|
||||
],
|
||||
sanitize_openai_function_tool_names=True,
|
||||
)
|
||||
response_object = {
|
||||
"id": "test",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue