mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(gigachat): generic env-credential passthrough fallback plus type hardening
- forward unrouted /gigachat/* requests with env credentials like other passthrough providers (the old fallback returned 400 on any request without a routed model, /gigachat/models included) - fix basedpyright budget breaches across the gigachat provider, common_request_processing, and llm_passthrough_endpoints with real narrowing, no new suppressions - add regression tests for the fallback target, auth header, and model-less endpoints
This commit is contained in:
parent
70e2f4e68f
commit
b0ce17c755
8 changed files with 168 additions and 104 deletions
|
|
@ -14,7 +14,7 @@ from collections.abc import Callable, Mapping, Sequence
|
|||
from datetime import datetime as dt_object
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType, TracebackType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast
|
||||
|
||||
from httpx import Response
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -86,7 +86,6 @@ from litellm.litellm_core_utils.redact_messages import (
|
|||
redact_streaming_responses_for_custom_logger,
|
||||
)
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
|
||||
from litellm.llms.base_llm.search.transformation import SearchResponse
|
||||
from litellm.responses.utils import ResponseAPILoggingUtils
|
||||
from litellm.types.agents import LiteLLMSendMessageResponse
|
||||
|
|
@ -1603,24 +1602,26 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
|
||||
def _response_cost_calculator(
|
||||
self,
|
||||
result: ModelResponse
|
||||
| ModelResponseStream
|
||||
| EmbeddingResponse
|
||||
| ImageResponse
|
||||
| TranscriptionResponse
|
||||
| TextCompletionResponse
|
||||
| HttpxBinaryResponseContent
|
||||
| RerankResponse
|
||||
| Batch
|
||||
| FineTuningJob
|
||||
| ResponsesAPIResponse
|
||||
| ResponseCompletedEvent
|
||||
| OpenAIFileObject
|
||||
| LiteLLMRealtimeStreamLoggingObject
|
||||
| OpenAIModerationResponse
|
||||
| SearchResponse
|
||||
| dict
|
||||
| list,
|
||||
result: Union[
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
EmbeddingResponse,
|
||||
ImageResponse,
|
||||
TranscriptionResponse,
|
||||
TextCompletionResponse,
|
||||
HttpxBinaryResponseContent,
|
||||
RerankResponse,
|
||||
Batch,
|
||||
FineTuningJob,
|
||||
ResponsesAPIResponse,
|
||||
ResponseCompletedEvent,
|
||||
OpenAIFileObject,
|
||||
LiteLLMRealtimeStreamLoggingObject,
|
||||
OpenAIModerationResponse,
|
||||
"SearchResponse",
|
||||
dict,
|
||||
list,
|
||||
],
|
||||
cache_hit: bool | None = None,
|
||||
litellm_model_name: str | None = None,
|
||||
router_model_id: str | None = None,
|
||||
|
|
@ -6262,7 +6263,7 @@ def _get_traceback_str_for_error(error_str: str) -> str:
|
|||
from decimal import Decimal
|
||||
|
||||
# used for unit testing
|
||||
from typing import Any, Optional
|
||||
from typing import Any, Optional, Union
|
||||
|
||||
|
||||
def create_dummy_standard_logging_payload() -> StandardLoggingPayload:
|
||||
|
|
|
|||
|
|
@ -53,20 +53,22 @@ class GigaChatModelResponseIterator:
|
|||
finish_reason: str | None = chunk_finish_reason
|
||||
|
||||
# Handle function_call in stream
|
||||
if chunk_finish_reason == "function_call" and delta.get("function_call"):
|
||||
func_call: Final = delta["function_call"]
|
||||
args_raw: Final = func_call.get("arguments") or {}
|
||||
raw_function_call: Final = delta.get("function_call")
|
||||
if chunk_finish_reason == "function_call" and isinstance(raw_function_call, Mapping) and raw_function_call:
|
||||
func_call: Final[Mapping[str, object]] = raw_function_call
|
||||
args_raw: Final[object] = func_call.get("arguments") or {}
|
||||
args_str: str # rebind-ok: conditionally assigned from dict or str
|
||||
if isinstance(args_raw, dict):
|
||||
args_str = json.dumps(args_raw, ensure_ascii=False) # rebind-ok: build from dict
|
||||
else:
|
||||
args_str = str(args_raw)
|
||||
|
||||
name_raw: Final = func_call.get("name")
|
||||
tool_use = ChatCompletionToolCallChunk(
|
||||
id=f"call_{uuid.uuid4().hex[:24]}",
|
||||
type="function",
|
||||
function=ChatCompletionToolCallFunctionChunk(
|
||||
name=func_call.get("name", ""),
|
||||
name=name_raw if isinstance(name_raw, str) else "",
|
||||
arguments=args_str,
|
||||
),
|
||||
index=0,
|
||||
|
|
|
|||
|
|
@ -167,19 +167,21 @@ class GigaChatConfig(BaseConfig):
|
|||
pass
|
||||
elif param == "tools":
|
||||
# Convert tools to functions format
|
||||
optional_params["functions"] = self._convert_tools_to_functions(value)
|
||||
if isinstance(value, Sequence):
|
||||
optional_params["functions"] = self._convert_tools_to_functions(value)
|
||||
elif param == "tool_choice":
|
||||
# Map OpenAI tool_choice to GigaChat function_call
|
||||
mapped_choice = self._map_tool_choice(value)
|
||||
if mapped_choice is not None:
|
||||
optional_params["function_call"] = mapped_choice
|
||||
if isinstance(value, (str, Mapping)):
|
||||
mapped_choice = self._map_tool_choice(value)
|
||||
if mapped_choice is not None:
|
||||
optional_params["function_call"] = mapped_choice
|
||||
elif param == "functions":
|
||||
optional_params["functions"] = value
|
||||
elif param == "function_call":
|
||||
optional_params["function_call"] = value
|
||||
elif param == "response_format":
|
||||
# Handle structured output via function calling
|
||||
if value.get("type") == "json_schema":
|
||||
if isinstance(value, Mapping) and value.get("type") == "json_schema":
|
||||
json_schema = value.get("json_schema", {})
|
||||
schema_name = json_schema.get("name", "structured_output")
|
||||
schema = json_schema.get("schema", {})
|
||||
|
|
@ -190,9 +192,15 @@ class GigaChatConfig(BaseConfig):
|
|||
"parameters": schema,
|
||||
}
|
||||
|
||||
if "functions" not in optional_params:
|
||||
optional_params["functions"] = [] # mutable-ok: list for httpx
|
||||
optional_params["functions"].append(function_def)
|
||||
existing_functions = optional_params.get("functions")
|
||||
optional_params["functions"] = [
|
||||
*(
|
||||
existing_functions
|
||||
if isinstance(existing_functions, Sequence) and not isinstance(existing_functions, str)
|
||||
else ()
|
||||
),
|
||||
function_def,
|
||||
]
|
||||
optional_params["function_call"] = {"name": schema_name} # mutable-ok: request payload
|
||||
optional_params["_structured_output"] = True
|
||||
|
||||
|
|
@ -246,8 +254,9 @@ class GigaChatConfig(BaseConfig):
|
|||
# OpenAI format: {"type": "function", "function": {"name": "func_name"}}
|
||||
# GigaChat format: {"name": "func_name"}
|
||||
if tool_choice.get("type") == "function":
|
||||
func_name: Final = tool_choice.get("function", {}).get("name")
|
||||
if func_name:
|
||||
function_spec: Final = tool_choice.get("function")
|
||||
func_name: Final = function_spec.get("name") if isinstance(function_spec, Mapping) else None
|
||||
if isinstance(func_name, str) and func_name:
|
||||
return {"name": func_name}
|
||||
|
||||
# Default to None (don't set function_call)
|
||||
|
|
@ -317,7 +326,7 @@ class GigaChatConfig(BaseConfig):
|
|||
giga_messages: Final = self._transform_messages(messages)
|
||||
|
||||
# Build request
|
||||
request_data: Final = {
|
||||
request_data: Final[dict[str, object]] = {
|
||||
"model": model.replace("gigachat/", ""),
|
||||
"messages": giga_messages,
|
||||
}
|
||||
|
|
@ -407,7 +416,7 @@ class GigaChatConfig(BaseConfig):
|
|||
messages: list[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: "tiktoken.Encoding | None",
|
||||
encoding: tiktoken.Encoding | None,
|
||||
api_key: str | None = None,
|
||||
json_mode: bool | None = None,
|
||||
) -> ModelResponse:
|
||||
|
|
|
|||
|
|
@ -92,16 +92,19 @@ class GigaChatPassthroughConfig(BasePassthroughConfig):
|
|||
if provider_chat_config is None:
|
||||
raise ValueError(f"No provider config found for model: {model}")
|
||||
|
||||
raw_messages: Final = request_data.get("messages")
|
||||
litellm_model_response: Final = provider_chat_config.transform_response(
|
||||
model=model,
|
||||
messages=request_data.get("messages", []), # mutable-ok: empty list default for transform_response
|
||||
messages=list(raw_messages)
|
||||
if isinstance(raw_messages, list)
|
||||
else [], # mutable-ok: transform_response wants a list
|
||||
raw_response=httpx_response,
|
||||
model_response=ModelResponse(),
|
||||
logging_obj=logging_obj,
|
||||
optional_params={}, # mutable-ok: empty dict kwarg for transform_response
|
||||
litellm_params={}, # mutable-ok: empty dict kwarg for transform_response
|
||||
api_key="",
|
||||
request_data=request_data,
|
||||
request_data=dict(request_data), # mutable-ok: transform_response wants a dict
|
||||
encoding=encoding,
|
||||
)
|
||||
|
||||
|
|
@ -124,7 +127,7 @@ class GigaChatPassthroughConfig(BasePassthroughConfig):
|
|||
logging_obj=logging_obj,
|
||||
optional_params={}, # mutable-ok: empty dict kwarg for transform_embedding_response
|
||||
api_key="",
|
||||
request_data=request_data,
|
||||
request_data=dict(request_data), # mutable-ok: transform_embedding_response wants a dict
|
||||
litellm_params={}, # mutable-ok: empty dict kwarg for transform_embedding_response
|
||||
)
|
||||
)
|
||||
|
|
@ -172,7 +175,9 @@ class GigaChatPassthroughConfig(BasePassthroughConfig):
|
|||
)
|
||||
translated_chunk = gigachat_iterator.chunk_parser(chunk=message)
|
||||
|
||||
if isinstance(translated_chunk, dict) and generic_chunk_has_all_required_fields(translated_chunk): # pyright: ignore[reportUnnecessaryIsInstance] # runtime guard for patched chunk_parser
|
||||
if isinstance(translated_chunk, dict) and generic_chunk_has_all_required_fields( # pyright: ignore[reportUnnecessaryIsInstance] # runtime guard for patched chunk_parser
|
||||
dict(translated_chunk)
|
||||
):
|
||||
chunk_obj = convert_generic_chunk_to_model_response_stream(
|
||||
translated_chunk # pyright: ignore[reportArgumentType] # validated TypedDict
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2457,12 +2457,15 @@ class ProxyBaseLLMRequestProcessing:
|
|||
logging_obj._on_deferred_stream_complete = _on_deferred_native_stream_complete
|
||||
|
||||
if route_type == "allm_passthrough_route":
|
||||
streaming_headers = custom_headers # rebind-ok: initial assignment before header merge
|
||||
if hasattr(response, "headers"):
|
||||
streaming_headers = ProxyBaseLLMRequestProcessing._merge_passthrough_streaming_headers( # rebind-ok: merge result replaces initial assignment
|
||||
response_headers=getattr(response, "headers", None),
|
||||
upstream_response_headers: Final = getattr(response, "headers", None)
|
||||
streaming_headers: Final = (
|
||||
ProxyBaseLLMRequestProcessing._merge_passthrough_streaming_headers(
|
||||
response_headers=upstream_response_headers,
|
||||
custom_headers=custom_headers,
|
||||
)
|
||||
if upstream_response_headers is not None
|
||||
else custom_headers
|
||||
)
|
||||
|
||||
# Check if response is an async generator
|
||||
if self._is_streaming_response(response):
|
||||
|
|
|
|||
|
|
@ -2824,7 +2824,6 @@ async def gigachat_proxy_route(
|
|||
"""
|
||||
[Docs](https://docs.litellm.ai/docs/pass_through/gigachat)
|
||||
"""
|
||||
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
|
||||
from litellm.proxy.proxy_server import (
|
||||
general_settings,
|
||||
llm_router,
|
||||
|
|
@ -2876,55 +2875,35 @@ async def gigachat_proxy_route(
|
|||
version=version,
|
||||
)
|
||||
|
||||
# Fall back to existing implementation for direct GigaChat models
|
||||
verbose_proxy_logger.debug(
|
||||
"Gigachat passthrough: Using direct Gigachat model '%s' for endpoint '%s'", model, endpoint
|
||||
)
|
||||
|
||||
data: dict[
|
||||
str, Any
|
||||
] = {} # mutable-ok: request body mutated in place by proxy pipeline; pyright: ignore[reportExplicitAny] # Any needed for proxy pipeline flexibility
|
||||
from litellm.llms.gigachat.authenticator import get_access_token
|
||||
from litellm.llms.gigachat.utils import GIGACHAT_BASE_URL
|
||||
|
||||
data["method"] = request.method
|
||||
data["endpoint"] = endpoint
|
||||
data["json"] = request_body
|
||||
data["custom_llm_provider"] = "gigachat"
|
||||
base_target_url: Final = get_secret_str("GIGACHAT_API_BASE") or GIGACHAT_BASE_URL
|
||||
request_path: Final = httpx.URL(endpoint).path
|
||||
encoded_endpoint: Final = request_path if request_path.startswith("/") else f"/{request_path}"
|
||||
|
||||
client: Final = get_async_httpx_client(
|
||||
llm_provider=LlmProviders.GIGACHAT,
|
||||
params={ # mutable-ok: httpx client params
|
||||
"timeout": httpx.Timeout(timeout=600.0, connect=5.0),
|
||||
"ssl_verify": False,
|
||||
},
|
||||
base_url: Final = httpx.URL(base_target_url)
|
||||
updated_url: Final = base_url.copy_with(
|
||||
path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, encoded_endpoint)
|
||||
)
|
||||
data["client"] = client
|
||||
|
||||
base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data)
|
||||
is_streaming_request: Final = await is_streaming_request_fn(request)
|
||||
|
||||
try:
|
||||
return await base_llm_response_processor.base_passthrough_process_llm_request(
|
||||
request=request,
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
llm_router=llm_router,
|
||||
general_settings=general_settings,
|
||||
proxy_config=proxy_config,
|
||||
select_data_generator=select_data_generator,
|
||||
model=model,
|
||||
user_model=user_model,
|
||||
user_temperature=user_temperature,
|
||||
user_request_timeout=user_request_timeout,
|
||||
user_max_tokens=user_max_tokens,
|
||||
user_api_base=user_api_base,
|
||||
version=version,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # Safe catch-all for handle exception
|
||||
raise await base_llm_response_processor._handle_llm_api_exception(
|
||||
e=e,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
endpoint_func: Final = create_pass_through_route(
|
||||
endpoint=endpoint,
|
||||
target=str(updated_url),
|
||||
custom_headers={"Authorization": f"Bearer {get_access_token()}"},
|
||||
is_streaming_request=is_streaming_request,
|
||||
)
|
||||
return await endpoint_func(
|
||||
request,
|
||||
fastapi_response,
|
||||
user_api_key_dict,
|
||||
)
|
||||
|
||||
|
||||
async def handle_gigachat_passthrough_router_model(
|
||||
|
|
|
|||
|
|
@ -92,7 +92,8 @@ class TestGetOpenaiCompatibleProviderInfo:
|
|||
assert api_base == "https://api.example.com"
|
||||
assert api_key == "test-key"
|
||||
|
||||
def test_resolves_api_base_when_none(self):
|
||||
def test_resolves_api_base_when_none(self, monkeypatch):
|
||||
monkeypatch.delenv("GIGACHAT_API_BASE", raising=False)
|
||||
provider, api_base, api_key = self.config._get_openai_compatible_provider_info(
|
||||
api_base=None, api_key="key"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2133,37 +2133,101 @@ class TestGigachatProxyRoute:
|
|||
return_value=False,
|
||||
)
|
||||
@patch( # test-quality-ok: patching litellm internal for unit test isolation
|
||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.base_passthrough_process_llm_request",
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_streaming_request_fn",
|
||||
new_callable=AsyncMock,
|
||||
return_value=False,
|
||||
)
|
||||
async def test_gigachat_proxy_route_fallback_to_http_pass_through(
|
||||
@patch( # test-quality-ok: patching litellm internal for unit test isolation
|
||||
"litellm.llms.gigachat.authenticator.get_access_token",
|
||||
return_value="gigachat-test-token",
|
||||
)
|
||||
async def test_gigachat_proxy_route_fallback_forwards_to_gigachat_api(
|
||||
self,
|
||||
mock_base_passthrough,
|
||||
mock_get_token,
|
||||
mock_is_streaming,
|
||||
mock_is_router,
|
||||
mock_get_body,
|
||||
monkeypatch,
|
||||
):
|
||||
monkeypatch.delenv("GIGACHAT_API_BASE", raising=False)
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_fastapi_response = MagicMock(spec=Response)
|
||||
mock_user_api_key_dict = MagicMock()
|
||||
|
||||
expected_response = Response(
|
||||
content=b'{"response": "success"}',
|
||||
status_code=200,
|
||||
media_type="application/json",
|
||||
)
|
||||
mock_base_passthrough.return_value = expected_response
|
||||
captured_kwargs = {}
|
||||
|
||||
result = await gigachat_proxy_route(
|
||||
endpoint="/chat/completions",
|
||||
request=mock_request,
|
||||
fastapi_response=mock_fastapi_response,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
async def fake_endpoint(request, fastapi_response, user_api_key_dict):
|
||||
return Response(content=b'{"response": "success"}', status_code=200)
|
||||
|
||||
def fake_create_pass_through_route(**kwargs):
|
||||
captured_kwargs.update(kwargs)
|
||||
return fake_endpoint
|
||||
|
||||
with patch( # test-quality-ok: patching litellm internal for unit test isolation
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
|
||||
side_effect=fake_create_pass_through_route,
|
||||
):
|
||||
result = await gigachat_proxy_route(
|
||||
endpoint="/chat/completions",
|
||||
request=mock_request,
|
||||
fastapi_response=mock_fastapi_response,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
assert isinstance(result, Response)
|
||||
assert result.status_code == 200
|
||||
assert result.body == b'{"response": "success"}'
|
||||
mock_base_passthrough.assert_awaited_once()
|
||||
assert captured_kwargs["target"] == "https://gigachat.devices.sberbank.ru/api/v1/chat/completions"
|
||||
assert captured_kwargs["custom_headers"] == {"Authorization": "Bearer gigachat-test-token"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@patch( # test-quality-ok: patching litellm internal for unit test isolation
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
|
||||
return_value={},
|
||||
)
|
||||
@patch( # test-quality-ok: patching litellm internal for unit test isolation
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_streaming_request_fn",
|
||||
new_callable=AsyncMock,
|
||||
return_value=False,
|
||||
)
|
||||
@patch( # test-quality-ok: patching litellm internal for unit test isolation
|
||||
"litellm.llms.gigachat.authenticator.get_access_token",
|
||||
return_value="gigachat-test-token",
|
||||
)
|
||||
async def test_gigachat_proxy_route_models_endpoint_without_model(
|
||||
self,
|
||||
mock_get_token,
|
||||
mock_is_streaming,
|
||||
mock_get_body,
|
||||
monkeypatch,
|
||||
):
|
||||
monkeypatch.delenv("GIGACHAT_API_BASE", raising=False)
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_fastapi_response = MagicMock(spec=Response)
|
||||
mock_user_api_key_dict = MagicMock()
|
||||
|
||||
captured_kwargs = {}
|
||||
|
||||
async def fake_endpoint(request, fastapi_response, user_api_key_dict):
|
||||
return Response(content=b'{"data": []}', status_code=200)
|
||||
|
||||
def fake_create_pass_through_route(**kwargs):
|
||||
captured_kwargs.update(kwargs)
|
||||
return fake_endpoint
|
||||
|
||||
with patch( # test-quality-ok: patching litellm internal for unit test isolation
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route",
|
||||
side_effect=fake_create_pass_through_route,
|
||||
):
|
||||
result = await gigachat_proxy_route(
|
||||
endpoint="models",
|
||||
request=mock_request,
|
||||
fastapi_response=mock_fastapi_response,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
|
||||
assert isinstance(result, Response)
|
||||
assert result.status_code == 200
|
||||
assert captured_kwargs["target"] == "https://gigachat.devices.sberbank.ru/api/v1/models"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_allm_passthrough_streaming_preserves_upstream_headers(self):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue