diff --git a/tests/integration/observability/test_grayswan_wire.py b/tests/integration/observability/test_grayswan_wire.py index 589c8d9d340..b14e4a42079 100644 --- a/tests/integration/observability/test_grayswan_wire.py +++ b/tests/integration/observability/test_grayswan_wire.py @@ -106,7 +106,7 @@ def _grayswan_config( return path -def _vendor(violation: float = 0.0): +def _vendor(violation: float = 0.0) -> Callable[[Request], Reply]: def respond(request: Request) -> Reply: assert request.method == "POST" assert request.target == "/cygnal/monitor", request.target @@ -125,7 +125,7 @@ def _serving_model_probe(respond: Callable[[Request], Reply]) -> Callable[[Reque return wrapped -def _chat_provider(message: dict[str, JsonValue]): +def _chat_provider(message: dict[str, JsonValue]) -> Callable[[Request], Reply]: def respond(request: Request) -> Reply: assert request.target == "/chat/completions", request.target return Reply( @@ -577,7 +577,7 @@ def test_post_call_multi_choice_texts_and_tool_calls_stay_split(gateway: Gateway ], body -def _chat_stream_provider(chunks: int): +def _chat_stream_provider(chunks: int) -> Callable[[Request], Reply]: def respond(request: Request) -> Reply: assert request.target == "/chat/completions", request.target frames: Final = tuple( diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py index 83b0be53fee..954c57b5cb6 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py @@ -248,7 +248,7 @@ async def test_run_guardrail_posts_payload(monkeypatch, grayswan_guardrail: Gray def fake_process( response_json: dict, - data: dict | None = None, + data: dict[str, object] | None = None, hook_type: GuardrailEventHooks | None = None, ) -> None: captured["response"] = response_json @@ -598,11 +598,13 @@ def test_ensure_litellm_metadata_noop_when_already_present() -> None: class _CapturingClient: - def __init__(self, payload: dict[str, float] | None = None): + def __init__(self, payload: dict[str, float] | None = None) -> None: self.payload = payload or {"violation": 0.0} self.calls: tuple[Mapping[str, object], ...] = () - async def post(self, *, url: str, headers: dict, json: dict, timeout: float): + async def post( + self, *, url: str, headers: Mapping[str, str], json: Mapping[str, object], timeout: float + ) -> _DummyResponse: self.calls = ( *self.calls, MappingProxyType({"url": url, "headers": headers, "json": json, "timeout": timeout}), @@ -611,7 +613,7 @@ class _CapturingClient: class _LoggingObj: - def __init__(self, call_type): + def __init__(self, call_type: str | None) -> None: self.call_type = call_type