refactor(guardrails): hoist function-local imports to module top

Moves the imports this PR added inside functions to the top of their
modules: UserAPIKeyAuth and BaseTranslation in realtime_streaming,
BaseTranslation, StreamingScanKey and LiteLLMProxyRequestSetup in
unified_guardrail (so its annotations no longer need quotes), and the
test-only imports in the unified guardrail, guardrail endpoint, pass-through
and realtime tests. None of them creates a cycle, and import litellm still
does not load the proxy request layer or fastapi

One import stays inside its function: base_translation's
LiteLLMProxyRequestSetup. import litellm loads base_translation before
litellm.Router exists, and litellm_pre_call_utils imports Router, so hoisting
it makes import litellm fail with "cannot import name 'Router' from
'litellm'". It carries a one-line comment saying so
This commit is contained in:
Caduri Katzav 2026-09-29 23:59:33 +03:00
parent 1dbd73d512
commit 88e0c43e7c
7 changed files with 27 additions and 65 deletions

View file

@ -12,7 +12,9 @@ import litellm
from litellm._logging import redact_internal_details_from_client_message, verbose_logger
from litellm.constants import REALTIME_SESSION_FAILURE_LOGGED_KEY, REALTIME_SESSION_SUCCESS_LOGGED_KEY
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig, RealtimeBackend
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.llms.openai import (
OpenAIRealtimeEvents,
OpenAIRealtimeOutputItemDone,
@ -29,7 +31,6 @@ if TYPE_CHECKING:
from websockets.asyncio.client import ClientConnection
from websockets.exceptions import ConnectionClosed
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
CLIENT_CONNECTION_CLASS = ClientConnection
@ -124,9 +125,7 @@ DefaultLoggedRealTimeEventTypes: Final = [
]
def _as_user_api_key_auth(user_api_key_dict: object) -> "UserAPIKeyAuth | None":
from litellm.proxy._types import UserAPIKeyAuth
def _as_user_api_key_auth(user_api_key_dict: object) -> UserAPIKeyAuth | None:
return user_api_key_dict if isinstance(user_api_key_dict, UserAPIKeyAuth) else None
@ -838,7 +837,6 @@ class RealTimeStreaming:
typed user messages and tool outputs use ``pre_call``.
"""
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.types.guardrails import GuardrailEventHooks
if event_hooks is None:

View file

@ -86,6 +86,7 @@ class BaseTranslation(ABC):
"""The authenticated key's identity as prefixed metadata, an allowlist safe to hand to guardrail vendors."""
if user_api_key_dict is None:
return {}
# Lazy: `import litellm` loads this module before litellm.Router exists, and litellm_pre_call_utils imports it
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
return {

View file

@ -20,7 +20,9 @@ from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.api_route_to_call_types import get_call_types_for_route
from litellm.llms import get_guardrail_translation_mapping, load_guardrail_translation_mappings
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation, StreamingScanKey
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import (
CallTypes,
@ -34,10 +36,6 @@ if TYPE_CHECKING:
# Imported lazily at runtime (inside the streaming hook) to avoid a
# module-level cyclic import with litellm.integrations.custom_guardrail.
from litellm.integrations.custom_guardrail import ModifyResponseException
from litellm.llms.base_llm.guardrail_translation.base_translation import (
BaseTranslation,
StreamingScanKey,
)
# Call types that stream JSON-RPC events (A2A); guardrail HTTPException is emitted as in-stream error
A2A_CALL_TYPES: Final = (CallTypes.asend_message, CallTypes.send_message)
@ -56,7 +54,7 @@ class _EndpointTranslation(Protocol):
def process_output_streaming_response(self) -> "Callable[..., Awaitable[object]]": ...
@property
def get_streaming_scan_key(self) -> "Callable[[Sequence[object]], StreamingScanKey | None]": ...
def get_streaming_scan_key(self) -> Callable[[Sequence[object]], StreamingScanKey | None]: ...
@property
def build_block_sse_chunks(self) -> "Callable[..., Sequence[bytes] | None]": ...
@ -71,7 +69,7 @@ def _as_endpoint_translation(translation: _EndpointTranslation) -> _EndpointTran
def resolve_endpoint_translation(
user_api_key_dict: UserAPIKeyAuth, first_response_item: object | None
) -> "tuple[str, BaseTranslation] | None":
) -> tuple[str, BaseTranslation] | None:
"""
Resolve the endpoint guardrail translation for a streamed response: the
request route wins, falling back to inferring the call type from the first
@ -108,7 +106,7 @@ def _held_choices(held_chars_per_choice: Mapping[int, int]) -> frozenset[int]:
return frozenset(idx for idx, held in held_chars_per_choice.items() if held > 0)
def _is_redundant_scan(scan_key: "StreamingScanKey | None", last_scan_key: "StreamingScanKey | None") -> bool:
def _is_redundant_scan(scan_key: StreamingScanKey | None, last_scan_key: StreamingScanKey | None) -> bool:
if scan_key is None:
return False
return scan_key == last_scan_key or scan_key.has_nothing_to_scan
@ -160,11 +158,6 @@ _PROXY_ENRICHED_IDENTITY_FIELDS: Final = frozenset({"user_api_key_auth_metadata"
def _ensure_litellm_metadata(data: dict, user_api_key_dict: UserAPIKeyAuth) -> None:
"""Overwrite the identity fields of data['litellm_metadata'] from the authenticated key, in place."""
from litellm.llms.base_llm.guardrail_translation.base_translation import (
BaseTranslation,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
existing: Final = data.get("litellm_metadata")
if isinstance(existing, dict):
identity: Final = LiteLLMProxyRequestSetup.get_authenticated_identity_metadata(user_api_key_dict)
@ -414,7 +407,7 @@ class UnifiedLLMGuardrails(CustomLogger):
@staticmethod
def _resolve_transform_call_type(
user_api_key_dict: UserAPIKeyAuth,
mappings: Mapping[CallTypes, type["BaseTranslation"]],
mappings: Mapping[CallTypes, type[BaseTranslation]],
) -> str | None:
"""Resolve the call type for the incremental_diff path, or None if the
route is unresolvable / unsupported.
@ -677,7 +670,7 @@ class UnifiedLLMGuardrails(CustomLogger):
call_type: str,
sampling_rate: int,
end_of_stream_only: bool,
mappings: Mapping[CallTypes, type["BaseTranslation"]],
mappings: Mapping[CallTypes, type[BaseTranslation]],
) -> AsyncGenerator[object, None]:
"""Emit guardrail text transformations as new deltas on the stream.

View file

@ -1,11 +1,16 @@
"""Tests for unified guardrail."""
import copy
import json
import logging
from collections.abc import Callable
from types import SimpleNamespace
from typing import TYPE_CHECKING, Final, Literal
from unittest.mock import MagicMock
import httpx
import pytest
from fastapi import Request
import litellm
from litellm.caching import DualCache
@ -23,6 +28,7 @@ from litellm.llms.base_llm.guardrail_translation.utils import (
openai_messages_without_tool,
)
from litellm.llms.base_llm.ocr.transformation import OCRPage, OCRResponse
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.mistral.ocr.guardrail_translation.handler import OCRHandler
from litellm.llms.openai.chat.guardrail_translation.handler import (
OpenAIChatCompletionsHandler,
@ -34,12 +40,15 @@ from litellm.proxy._experimental.mcp_server.guardrail_translation.handler import
MCPGuardrailTranslationHandler,
)
from litellm.proxy._types import LiteLLMRoutes, UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail import (
unified_guardrail as unified_module,
)
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_litellm_data_to_request
from litellm.proxy.utils import ProxyLogging
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.utils import CallTypes, Delta, GenericGuardrailAPIInputs, ModelResponseStream, StreamingChoices
@ -2429,13 +2438,6 @@ class TestGuardrailsSeeAuthenticatedIdentity:
@staticmethod
def _generic_guardrail(vendor_payloads: list[dict[str, object]]) -> CustomGuardrail:
import json
import httpx
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
def vendor(request: httpx.Request) -> httpx.Response:
vendor_payloads.append(json.loads(request.content))
return httpx.Response(200, json={"action": "NONE"})
@ -2472,8 +2474,6 @@ class TestGuardrailsSeeAuthenticatedIdentity:
@pytest.mark.asyncio
async def test_mcp_tool_call_reaches_vendor_with_key_alias(self) -> None:
from litellm.proxy.utils import ProxyLogging
vendor_payloads: list[dict[str, object]] = []
key = UserAPIKeyAuth(api_key="sk-real-caller-key", key_alias="prod-app", team_id="team-prod")
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
@ -2530,14 +2530,6 @@ class TestGuardrailsSeeAuthenticatedIdentity:
) -> None:
"""After the chat-path metadata build, the guardrail hook leaves the proxy's bucket as it was (team metadata
in user_api_key_auth_metadata included) and the vendor gets the logged key, never a raw CLI session token."""
import copy
import json
from unittest.mock import MagicMock
from fastapi import Request
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_litellm_data_to_request
_patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings())
request = MagicMock(spec=Request)
request.url = MagicMock()
@ -2584,8 +2576,6 @@ class TestGuardrailsSeeAuthenticatedIdentity:
@pytest.mark.asyncio
@pytest.mark.parametrize("bucket", ["metadata", "litellm_metadata", None])
async def test_pass_through_cli_session_key_sends_stable_hash(self, monkeypatch, bucket: str | None) -> None:
import json
_patch_translation_mappings(monkeypatch, discover_guardrail_translation_mappings())
vendor_payloads: list[dict[str, object]] = []
key = _cli_session_key("/anthropic/v1/messages")

View file

@ -4,11 +4,13 @@ from datetime import datetime
from typing import Dict, List, Optional
from unittest.mock import AsyncMock
import httpx
import pytest
from fastapi import HTTPException
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_endpoints import (
CreateGuardrailRequest,
@ -33,6 +35,7 @@ from litellm.proxy.guardrails.guardrail_endpoints import (
from litellm.proxy.guardrails.guardrail_endpoints import (
test_custom_code_guardrail as run_custom_code_test_endpoint,
)
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
MOCK_ADMIN_USER = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
from litellm.proxy.guardrails.guardrail_registry import (
@ -1661,11 +1664,6 @@ async def test_apply_guardrail_forwards_real_request_headers_not_caller_supplied
async def test_apply_guardrail_generic_guardrail_api_sends_authenticated_identity_to_vendor(mocker):
"""End to end through a real GenericGuardrailAPI: the vendor payload names the authenticated key even when
the body forges user_api_key_alias and user_api_key_token, which the generic guardrail maps onto the hash."""
import httpx
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
vendor_payloads = []
def vendor(request: httpx.Request) -> httpx.Response:
@ -1715,11 +1713,6 @@ async def test_apply_guardrail_request_route_comes_from_the_key(mocker):
@pytest.mark.asyncio
async def test_apply_guardrail_cli_session_key_sends_stable_hash_to_vendor(mocker):
"""A CLI session key's raw per-login token must never reach the vendor; it gets the stable logged key."""
import httpx
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
raw_session_token = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa"
vendor_payloads = []

View file

@ -26,6 +26,7 @@ from litellm._logging import verbose_proxy_logger
from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS,
@ -48,6 +49,7 @@ from litellm.proxy.pass_through_endpoints.success_handler import (
)
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.types import utils as types_utils
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY,
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
@ -7631,9 +7633,6 @@ async def test_pass_through_request_strips_caller_identity_before_guardrail_hook
control fields and inbound headers reached pre_call_hook guardrails as the caller's identity. The upstream body
is unchanged because these keys never reach it.
"""
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
upstream_bodies = []
def transport_handler(upstream_request: httpx.Request) -> httpx.Response:

View file

@ -17,6 +17,8 @@ from litellm.litellm_core_utils.realtime_streaming import (
client_sent_openai_beta_realtime_header,
)
from litellm.llms.xai.realtime.transformation import XAIRealtimeNormalizer
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.grayswan.grayswan import GraySwanGuardrail
from litellm.types.guardrails import GuardrailEventHooks
@ -3555,11 +3557,6 @@ async def test_provider_bytes_are_sent_raw_after_pacing():
@pytest.mark.asyncio
async def test_realtime_transcript_guardrail_receives_authenticated_identity(monkeypatch: pytest.MonkeyPatch):
"""Transcript guardrails get the session key's identity in litellm_metadata, as the chat path provides it."""
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
received_request_data = []
class IdentityRecordingGuardrail(CustomGuardrail):
@ -3594,11 +3591,6 @@ async def test_realtime_transcript_guardrail_receives_authenticated_identity(mon
@pytest.mark.asyncio
async def test_realtime_grayswan_payload_carries_only_identity(monkeypatch: pytest.MonkeyPatch):
"""Gray Swan forwards litellm_metadata verbatim, so realtime must hand it identity and no key secrets."""
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.grayswan.grayswan import GraySwanGuardrail
from litellm.types.guardrails import GuardrailEventHooks
vendor_payloads: list[dict[str, object]] = []
class RecordingGraySwan(GraySwanGuardrail):
@ -3652,10 +3644,6 @@ async def test_realtime_guardrail_gets_no_identity_from_non_auth_sdk_value(
monkeypatch: pytest.MonkeyPatch, sdk_value: object
):
"""Only a proxy-authenticated UserAPIKeyAuth yields identity; an SDK-supplied value never raises or fakes one."""
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
received_request_data: list[dict[str, object]] = []
class IdentityRecordingGuardrail(CustomGuardrail):