mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
refactor(types): replace Any with proven types in 7 files (#43844)
* refactor(types): replace Any with proven types in 9 files Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): pin xecguard-adjacent guardrail retry, mcp mixed tools, jwt routing and complexity router paths Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(types): keep only live-provable Any removals Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(types): type cache_hit as bool | None on the sync success path Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
7be2983f11
commit
fe9b6fd603
9 changed files with 181 additions and 8 deletions
|
|
@ -112,6 +112,7 @@ from litellm.litellm_core_utils.served_output_texts import (
|
|||
SERVED_OUTPUT_TEXTS_KEY,
|
||||
overlay_served_output_texts,
|
||||
)
|
||||
from litellm.litellm_core_utils.thread_pool_executor import executor
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.llms.base_llm.search.transformation import SearchResponse
|
||||
from litellm.responses.utils import ResponseAPILoggingUtils
|
||||
|
|
@ -175,7 +176,7 @@ from litellm.types.utils import (
|
|||
Usage,
|
||||
)
|
||||
from litellm.types.videos.main import VideoObject
|
||||
from litellm.utils import _get_base_model_from_metadata, executor, print_verbose
|
||||
from litellm.utils import _get_base_model_from_metadata, print_verbose
|
||||
|
||||
from ..integrations.argilla import ArgillaLogger
|
||||
from ..integrations.arize.arize_phoenix import ArizePhoenixLogger
|
||||
|
|
@ -3970,7 +3971,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
result: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
cache_hit: object | None = None,
|
||||
cache_hit: bool | None = None,
|
||||
) -> None:
|
||||
"""
|
||||
Handles calling success callbacks for Async calls.
|
||||
|
|
|
|||
|
|
@ -1104,7 +1104,7 @@ class CustomStreamWrapper:
|
|||
self,
|
||||
chunk: Any,
|
||||
model_response: ModelResponseStream,
|
||||
completion_obj: dict[str, Any],
|
||||
completion_obj: dict[str, object],
|
||||
) -> _ProviderChunkResult:
|
||||
response_obj: dict[str, Any] = {}
|
||||
if (
|
||||
|
|
|
|||
|
|
@ -1679,7 +1679,7 @@ def _create_elicitation_callback():
|
|||
|
||||
|
||||
def _record_mcp_guardrail_evaluations(
|
||||
synthetic_llm_data: dict[str, Any], # mutable-ok: `_sync_guardrail_info_to_logging_obj` takes a concrete dict
|
||||
synthetic_llm_data: dict[str, object], # mutable-ok: `_sync_guardrail_info_to_logging_obj` takes a concrete dict
|
||||
litellm_logging_obj: "LiteLLMLoggingObj | None",
|
||||
) -> None:
|
||||
"""Bridge guardrail decision records off an MCP synthetic request onto the request's logger.
|
||||
|
|
|
|||
|
|
@ -3641,7 +3641,7 @@ async def _get_source_cache_base_spend(
|
|||
) -> float:
|
||||
source_cache_keys: Final = [source_cache_key] if isinstance(source_cache_key, str) else source_cache_key
|
||||
for cache_key in source_cache_keys:
|
||||
source = await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
source: object = await user_api_key_cache.async_get_cache(key=cache_key)
|
||||
if source is None:
|
||||
continue
|
||||
if isinstance(source, dict):
|
||||
|
|
|
|||
|
|
@ -159,7 +159,7 @@ class LiteLLM_Proxy_MCP_Handler:
|
|||
def _parse_mcp_tools(tools: Iterable[Mapping[str, object]] | None) -> SplitTools:
|
||||
items: Final = tuple(tools or ())
|
||||
gateway_tools: Final[list[ToolParam]] = [tool for tool in items if _names_gateway_explicitly(tool)]
|
||||
other_tools: Final[list[Any]] = [tool for tool in items if not _names_gateway_explicitly(tool)]
|
||||
other_tools: Final[list[ToolParam]] = [tool for tool in items if not _names_gateway_explicitly(tool)]
|
||||
return gateway_tools, other_tools
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -2441,7 +2441,7 @@ class ComplexityRouter(CustomLogger):
|
|||
self,
|
||||
prompt: str,
|
||||
system_prompt: str | None = None,
|
||||
request_kwargs: dict[str, Any] | None = None,
|
||||
request_kwargs: Mapping[str, object] | None = None,
|
||||
messages: Sequence[Mapping[str, object]] | None = None,
|
||||
) -> tuple[ComplexityTier | str, float | None]:
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -124,7 +124,7 @@ class LoggingSurface(Protocol):
|
|||
) -> object: ...
|
||||
|
||||
def handle_sync_success_callbacks_for_async_calls(
|
||||
self, result: object, start_time: datetime.datetime, end_time: datetime.datetime, cache_hit: object = None
|
||||
self, result: object, start_time: datetime.datetime, end_time: datetime.datetime, cache_hit: bool | None = None
|
||||
) -> None: ...
|
||||
|
||||
def failure_handler(
|
||||
|
|
|
|||
79
tests/integration/mcp/test_responses_mcp_mixed_tools.py
Normal file
79
tests/integration/mcp/test_responses_mcp_mixed_tools.py
Normal file
|
|
@ -0,0 +1,79 @@
|
|||
import json
|
||||
import uuid
|
||||
from typing import Final
|
||||
|
||||
from integration._support.client import Gateway
|
||||
from integration._support.mcp import mcp_peer, register_mcp, tool_calls
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
|
||||
_FUNCTION_TOOL: Final = {
|
||||
"type": "function",
|
||||
"name": "lookup_weather",
|
||||
"description": "Look up the forecast for a city",
|
||||
"parameters": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]},
|
||||
}
|
||||
|
||||
|
||||
def test_responses_with_gateway_mcp_and_caller_function_tool_hands_both_to_model_and_returns_the_function_call(
|
||||
gateway: Gateway,
|
||||
) -> None:
|
||||
alias: Final = "mix" + uuid.uuid4().hex[:8]
|
||||
upstream_tools: list[tuple[str, ...]] = []
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.target.endswith("/responses"), request.target
|
||||
body: Final = json.loads(request.body)
|
||||
upstream_tools.append(tuple(str(tool.get("name")) for tool in body.get("tools", ())))
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "resp_mixed",
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "completed",
|
||||
"model": "gpt-4o-mini",
|
||||
"output": [
|
||||
{
|
||||
"type": "function_call",
|
||||
"id": "fc_weather",
|
||||
"call_id": "call_weather",
|
||||
"name": "lookup_weather",
|
||||
"arguments": json.dumps({"city": "Paris"}),
|
||||
"status": "completed",
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
with mcp_peer() as peer, wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
server_id: Final = register_mcp(scenario, peer, alias)
|
||||
model: Final = scenario.model(model="openai/responses/gpt-4o-mini", api_base=wire.url + "/v1")
|
||||
key: Final = scenario.key(object_permission={"mcp_servers": [server_id]})
|
||||
peer.drain()
|
||||
response: Final = gateway.client.post(
|
||||
"/v1/responses",
|
||||
headers={"Authorization": f"Bearer {key}"},
|
||||
json={
|
||||
"model": model,
|
||||
"input": "what is the weather in Paris",
|
||||
"tools": [
|
||||
{
|
||||
"type": "mcp",
|
||||
"server_url": "litellm_proxy",
|
||||
"server_label": "litellm",
|
||||
"require_approval": "never",
|
||||
},
|
||||
_FUNCTION_TOOL,
|
||||
],
|
||||
},
|
||||
timeout=90,
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert upstream_tools, "model was never called"
|
||||
assert all("lookup_weather" in names and f"{alias}-add" in names for names in upstream_tools), upstream_tools
|
||||
calls: Final = [item for item in response.json()["output"] if item.get("type") == "function_call"]
|
||||
assert [call["name"] for call in calls] == ["lookup_weather"], response.text
|
||||
assert json.loads(calls[0]["arguments"]) == {"city": "Paris"}
|
||||
assert tool_calls(peer.drain()) == (), "a caller-owned function call must never reach the MCP peer"
|
||||
|
|
@ -0,0 +1,93 @@
|
|||
import json
|
||||
import uuid
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import yaml
|
||||
from integration._support.client import Gateway
|
||||
from integration._support.process import owned_proxy
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
|
||||
|
||||
def _completion(content: str) -> Reply:
|
||||
return Reply(
|
||||
body=json.dumps(
|
||||
{
|
||||
"id": "chatcmpl-" + uuid.uuid4().hex[:8],
|
||||
"object": "chat.completion",
|
||||
"created": 1700000000,
|
||||
"model": "gpt-4o-mini",
|
||||
"choices": [
|
||||
{"index": 0, "message": {"role": "assistant", "content": content}, "finish_reason": "stop"}
|
||||
],
|
||||
"usage": {"prompt_tokens": 9, "completion_tokens": 3, "total_tokens": 12},
|
||||
}
|
||||
).encode()
|
||||
)
|
||||
|
||||
|
||||
def _tier_model(answer: str) -> Callable[[Request], Reply]:
|
||||
def respond(request: Request) -> Reply:
|
||||
if request.method == "GET":
|
||||
return Reply(body=json.dumps({"object": "list", "data": []}).encode())
|
||||
assert request.target == "/chat/completions", request.target
|
||||
return _completion(answer)
|
||||
|
||||
return respond
|
||||
|
||||
|
||||
def test_llm_classifier_verdict_routes_the_request_to_the_classified_tier_model(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
prompt: Final = f"hi there {uuid.uuid4().hex}"
|
||||
|
||||
def classifier(request: Request) -> Reply:
|
||||
if request.method == "GET":
|
||||
return Reply(body=json.dumps({"object": "list", "data": []}).encode())
|
||||
assert request.target == "/chat/completions", request.target
|
||||
body: Final = json.loads(request.body)
|
||||
assert "response_format" in body, body
|
||||
assert [message["role"] for message in body["messages"]] == ["system", "user"], body
|
||||
assert prompt in json.dumps(body["messages"][1]), body
|
||||
return _completion(json.dumps({"tier": "COMPLEX"}))
|
||||
|
||||
with (
|
||||
wire_server(classifier) as judge,
|
||||
wire_server(_tier_model("simple answer")) as simple,
|
||||
wire_server(_tier_model("complex answer")) as complex_tier,
|
||||
):
|
||||
router: Final = "router-" + uuid.uuid4().hex[:8]
|
||||
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
|
||||
config["model_list"] = [
|
||||
{
|
||||
"model_name": name,
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "api_base": url, "api_key": "synthetic"},
|
||||
}
|
||||
for name, url in (("judge", judge.url), ("simple", simple.url), ("complex", complex_tier.url))
|
||||
] + [
|
||||
{
|
||||
"model_name": router,
|
||||
"litellm_params": {
|
||||
"model": "auto_router/complexity_router",
|
||||
"complexity_router_config": {
|
||||
"classifier_type": "llm",
|
||||
"classifier_llm_config": {"model": "judge", "timeout_ms": 20000},
|
||||
"tiers": {"SIMPLE": "simple", "MEDIUM": "simple", "COMPLEX": "complex", "REASONING": "complex"},
|
||||
},
|
||||
},
|
||||
}
|
||||
]
|
||||
path: Final = tmp_path / "complexity_router.yaml"
|
||||
path.write_text(yaml.safe_dump(config))
|
||||
with owned_proxy(gateway, tmp_path, {}, config=path) as candidate:
|
||||
response: Final = candidate.request(
|
||||
"POST", "/v1/chat/completions", {"model": router, "messages": [{"role": "user", "content": prompt}]}
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.json()["choices"][0]["message"]["content"] == "complex answer", response.text
|
||||
assert len([call for call in judge.drain() if call.method == "POST"]) == 1
|
||||
assert [call for call in simple.drain() if call.method == "POST"] == []
|
||||
forwarded: Final = [call for call in complex_tier.drain() if call.method == "POST"]
|
||||
assert len(forwarded) == 1
|
||||
assert prompt in forwarded[0].body.decode()
|
||||
Loading…
Add table
Reference in a new issue