mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
1dbd73d512
commit
88e0c43e7c
7 changed files with 27 additions and 65 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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 = []
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue