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:
mateo-berri 2026-08-29 22:08:54 -07:00
parent 70e2f4e68f
commit b0ce17c755
8 changed files with 168 additions and 104 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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