test(grayswan): type the test helper parameters

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-30 21:26:30 +00:00
parent b53a8623a0
commit ad27fc4356
2 changed files with 9 additions and 7 deletions

View file

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

View file

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