Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/agent_endpoints/kill_switch.py: 49%
108 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 12:01 +0000
1from base64 import b64encode
2from collections.abc import AsyncIterator, Awaitable, Callable, Mapping
3from dataclasses import dataclass
4from datetime import datetime, timezone
5from types import MappingProxyType
6from typing import Final, Protocol, TypeAlias
8import httpx
9from typing_extensions import assert_never
11from litellm._logging import verbose_proxy_logger
12from litellm._uuid import uuid
13from litellm.constants import (
14 AGENT_KILL_SWITCH_RESPONSE_BODY_MAX_CHARS,
15 AGENT_KILL_SWITCH_TIMEOUT_SECONDS,
16 REDACTED_BY_LITELM_STRING,
17)
18from litellm.llms.custom_httpx.http_handler import (
19 get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # its params arg is a bare dict in http_handler
20)
21from litellm.proxy._types import LiteLLM_AuditLogs, LitellmTableNames, UserAPIKeyAuth
22from litellm.proxy.management_helpers.audit_logs import create_audit_log_for_update, get_audit_log_changed_by
23from litellm.types.agents import (
24 AgentKillSwitchApiKeyAuth,
25 AgentKillSwitchAuth,
26 AgentKillSwitchBasicAuth,
27 AgentKillSwitchBearerAuth,
28 AgentKillSwitchConfig,
29 AgentKillSwitchResult,
30)
31from litellm.types.llms.custom_http import httpxSpecialProvider
34def _with_auth(config: AgentKillSwitchConfig, auth: AgentKillSwitchAuth) -> AgentKillSwitchConfig:
35 return AgentKillSwitchConfig(
36 url=config.url,
37 method=config.method,
38 headers=config.headers,
39 query_params=config.query_params,
40 body=config.body,
41 auth=auth,
42 )
45def redact_kill_switch(config: AgentKillSwitchConfig | None) -> AgentKillSwitchConfig | None:
46 if config is None or config.auth is None: 46 ↛ 48line 46 didn't jump to line 48 because the condition on line 46 was always true
47 return config
48 return _with_auth(config, _redact_auth(config.auth))
51def _redact_auth(auth: AgentKillSwitchAuth) -> AgentKillSwitchAuth:
52 match auth:
53 case AgentKillSwitchBearerAuth():
54 return AgentKillSwitchBearerAuth(type="bearer", token=REDACTED_BY_LITELM_STRING)
55 case AgentKillSwitchApiKeyAuth():
56 return AgentKillSwitchApiKeyAuth(
57 type="api_key", header_name=auth.header_name, api_key=REDACTED_BY_LITELM_STRING
58 )
59 case AgentKillSwitchBasicAuth():
60 return AgentKillSwitchBasicAuth(type="basic", username=auth.username, password=REDACTED_BY_LITELM_STRING)
61 case _:
62 assert_never(auth)
65def restore_kill_switch(
66 incoming: AgentKillSwitchConfig | None,
67 existing: AgentKillSwitchConfig | None,
68) -> AgentKillSwitchConfig | None:
69 """Put the stored secret back behind an auth field echoed as the redaction
70 marker; a marker with no stored secret of the same auth type becomes ""."""
71 if incoming is None or incoming.auth is None:
72 return incoming
73 existing_auth: Final = existing.auth if existing is not None else None
74 return _with_auth(incoming, _restore_auth(incoming.auth, existing_auth))
77def _restore_secret(incoming_value: str, existing_value: str | None) -> str:
78 if incoming_value != REDACTED_BY_LITELM_STRING: 78 ↛ 80line 78 didn't jump to line 80 because the condition on line 78 was always true
79 return incoming_value
80 return existing_value if existing_value is not None else ""
83def _restore_auth(incoming: AgentKillSwitchAuth, existing: AgentKillSwitchAuth | None) -> AgentKillSwitchAuth:
84 match incoming:
85 case AgentKillSwitchBearerAuth():
86 stored_token: Final = existing.token if isinstance(existing, AgentKillSwitchBearerAuth) else None
87 return AgentKillSwitchBearerAuth(type="bearer", token=_restore_secret(incoming.token, stored_token))
88 case AgentKillSwitchApiKeyAuth(): 88 ↛ 89line 88 didn't jump to line 89 because the pattern on line 88 never matched
89 stored_key: Final = existing.api_key if isinstance(existing, AgentKillSwitchApiKeyAuth) else None
90 return AgentKillSwitchApiKeyAuth(
91 type="api_key",
92 header_name=incoming.header_name,
93 api_key=_restore_secret(incoming.api_key, stored_key),
94 )
95 case AgentKillSwitchBasicAuth(): 95 ↛ 102line 95 didn't jump to line 102 because the pattern on line 95 always matched
96 stored_password: Final = existing.password if isinstance(existing, AgentKillSwitchBasicAuth) else None
97 return AgentKillSwitchBasicAuth(
98 type="basic",
99 username=incoming.username,
100 password=_restore_secret(incoming.password, stored_password),
101 )
102 case _:
103 assert_never(incoming)
106@dataclass(frozen=True, slots=True)
107class KillSwitchRequest:
108 method: str
109 url: str
110 headers: Mapping[str, str]
111 json_body: Mapping[str, object] | None
114def _auth_headers(auth: AgentKillSwitchAuth | None) -> Mapping[str, str]:
115 match auth:
116 case None:
117 return MappingProxyType({})
118 case AgentKillSwitchBearerAuth():
119 return MappingProxyType({"Authorization": f"Bearer {auth.token}"})
120 case AgentKillSwitchApiKeyAuth():
121 return MappingProxyType({auth.header_name: auth.api_key})
122 case AgentKillSwitchBasicAuth():
123 credentials: Final = b64encode(f"{auth.username}:{auth.password}".encode()).decode()
124 return MappingProxyType({"Authorization": f"Basic {credentials}"})
125 case _:
126 assert_never(auth)
129def build_kill_switch_request(config: AgentKillSwitchConfig) -> KillSwitchRequest:
130 url: Final = httpx.URL(config.url).copy_merge_params(config.query_params)
131 return KillSwitchRequest(
132 method=config.method,
133 url=str(url),
134 headers=MappingProxyType({**config.headers, **_auth_headers(config.auth)}),
135 json_body=config.body,
136 )
139class KillSwitchHttpClient(Protocol):
140 def build_request( 140 ↛ exitline 140 didn't return from function 'build_request' because
141 self,
142 method: str,
143 url: str,
144 *,
145 headers: Mapping[str, str],
146 json: Mapping[str, object] | None,
147 timeout: float,
148 ) -> httpx.Request: ...
150 async def send(self, request: httpx.Request, *, stream: bool, follow_redirects: bool) -> httpx.Response: ... 150 ↛ exitline 150 didn't return from function 'send' because
153def default_kill_switch_http_client() -> KillSwitchHttpClient:
154 return get_async_httpx_client(llm_provider=httpxSpecialProvider.AgentKillSwitch).client
157KillSwitchAuditLogWriter: TypeAlias = Callable[[LiteLLM_AuditLogs], Awaitable[None]] # mutable-ok: Callable params
160def default_kill_switch_audit_log_writer() -> KillSwitchAuditLogWriter:
161 return create_audit_log_for_update
164def build_kill_switch_audit_log(
165 *,
166 result: AgentKillSwitchResult,
167 user_api_key_dict: UserAPIKeyAuth,
168 litellm_proxy_admin_name: str | None,
169) -> LiteLLM_AuditLogs:
170 return LiteLLM_AuditLogs(
171 id=str(uuid.uuid4()),
172 updated_at=datetime.now(timezone.utc),
173 changed_by=get_audit_log_changed_by(
174 litellm_changed_by=None,
175 user_api_key_dict=user_api_key_dict,
176 litellm_proxy_admin_name=litellm_proxy_admin_name,
177 ),
178 changed_by_api_key=user_api_key_dict.api_key,
179 table_name=LitellmTableNames.AGENT_TABLE_NAME,
180 object_id=result.agent_id,
181 action="kill_switch_fired",
182 updated_values=result.model_dump_json(exclude_none=True),
183 )
186async def fire_kill_switch(
187 *,
188 agent_id: str,
189 config: AgentKillSwitchConfig,
190 http_client: KillSwitchHttpClient,
191 timeout: float = AGENT_KILL_SWITCH_TIMEOUT_SECONDS,
192) -> AgentKillSwitchResult:
193 request: Final = build_kill_switch_request(config)
194 reported_url: Final = str(httpx.URL(request.url).copy_with(query=None))
195 verbose_proxy_logger.info("Firing kill switch for agent %s: %s %s", agent_id, request.method, reported_url)
196 try:
197 response: Final = await http_client.send(
198 http_client.build_request(
199 request.method,
200 request.url,
201 headers=request.headers,
202 json=request.json_body,
203 timeout=timeout,
204 ),
205 stream=True,
206 follow_redirects=False,
207 )
208 body: Final = await _read_text_prefix(response, AGENT_KILL_SWITCH_RESPONSE_BODY_MAX_CHARS)
209 except httpx.HTTPError as exc:
210 verbose_proxy_logger.warning("Kill switch for agent %s failed: %s", agent_id, type(exc).__name__)
211 return AgentKillSwitchResult(
212 agent_id=agent_id,
213 url=reported_url,
214 method=config.method,
215 error=type(exc).__name__,
216 )
217 return AgentKillSwitchResult(
218 agent_id=agent_id,
219 url=reported_url,
220 method=config.method,
221 status_code=response.status_code,
222 response_body=body,
223 )
226async def _read_text_prefix(response: httpx.Response, max_chars: int) -> str:
227 try:
228 return await _take_text(response.aiter_text(), max_chars)
229 finally:
230 await response.aclose()
233async def _take_text(chunks: AsyncIterator[str], max_chars: int) -> str:
234 taken = "" # rebind-ok: running prefix of a stream that is abandoned once the cap is hit
235 async for chunk in chunks:
236 taken += chunk # rebind-ok: see above
237 if len(taken) >= max_chars:
238 break
239 return taken[:max_chars]