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

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 

7 

8import httpx 

9from typing_extensions import assert_never 

10 

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 

32 

33 

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 ) 

43 

44 

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

49 

50 

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) 

63 

64 

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

75 

76 

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

81 

82 

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) 

104 

105 

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 

112 

113 

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) 

127 

128 

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 ) 

137 

138 

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: ... 

149 

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

151 

152 

153def default_kill_switch_http_client() -> KillSwitchHttpClient: 

154 return get_async_httpx_client(llm_provider=httpxSpecialProvider.AgentKillSwitch).client 

155 

156 

157KillSwitchAuditLogWriter: TypeAlias = Callable[[LiteLLM_AuditLogs], Awaitable[None]] # mutable-ok: Callable params 

158 

159 

160def default_kill_switch_audit_log_writer() -> KillSwitchAuditLogWriter: 

161 return create_audit_log_for_update 

162 

163 

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 ) 

184 

185 

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 ) 

224 

225 

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

231 

232 

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]