Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/_experimental/mcp_server/faults/classify.py: 19%

49 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1"""The single place that reads upstream OAuth/DCR failure responses. 

2 

3Every accessor here is total: an upstream that lies about its content encoding, sends an undecodable 

4body, or omits the spec fields yields a classified fault, never an exception. Nothing outside this 

5module should touch a failed upstream response's body. 

6""" 

7 

8from __future__ import annotations 

9 

10from typing import Final 

11 

12import httpx 

13 

14from litellm._logging import verbose_logger 

15from litellm.proxy._experimental.mcp_server.faults.types import ( 

16 GATEWAY_CAPABILITY_CODES, 

17 GATEWAY_CREDENTIAL_CODES, 

18 MAX_WIRE_FIELD_CHARS, 

19 CallerRejected, 

20 CredentialSource, 

21 GatewayRejected, 

22 UpstreamOAuthFault, 

23 UpstreamProtocolFault, 

24 UpstreamRegistrationRefused, 

25 UpstreamReportedFault, 

26) 

27 

28 

29def _safe_text(response: httpx.Response) -> str: 

30 try: 

31 return response.text 

32 except Exception: 

33 return "" 

34 

35 

36def _safe_json(response: httpx.Response) -> object: 

37 try: 

38 return response.json() 

39 except Exception: 

40 return None 

41 

42 

43def _bounded_field(value: object) -> str | None: 

44 if not isinstance(value, str) or not value: 

45 return None 

46 return value[:MAX_WIRE_FIELD_CHARS] 

47 

48 

49def _log_out_of_contract(endpoint_kind: str, response: httpx.Response, log_context: str) -> None: 

50 verbose_logger.warning( 

51 "MCP upstream %s endpoint (%s) returned HTTP %s outside the OAuth error contract (first %s chars): %s", 

52 endpoint_kind, 

53 log_context, 

54 response.status_code, 

55 MAX_WIRE_FIELD_CHARS, 

56 _safe_text(response)[:MAX_WIRE_FIELD_CHARS], 

57 ) 

58 

59 

60def _classify_oauth_error_code( 

61 code: str, 

62 description: str | None, 

63 error_uri: str | None, 

64 credential_source: CredentialSource, 

65 log_context: str, 

66) -> UpstreamOAuthFault: 

67 """Blame assignment for a contract-conformant OAuth error code, shared by the token and DCR 

68 classifiers. Codes by which the upstream blames itself keep that blame; ``invalid_target`` is a 

69 gateway configuration gap (the RFC 8707 resource indicator this server sends, or fails to send) 

70 no matter whose credentials were presented; credential-indicting codes follow the credential 

71 source; everything else, including codes we do not recognize, is the caller's to act on. The 

72 upstream's HTTP status is deliberately never consulted: status derives from this classification 

73 at render time, which is what keeps status and code from contradicting each other.""" 

74 if code == "server_error" or code == "temporarily_unavailable": 

75 return UpstreamReportedFault(code=code) 

76 if code in GATEWAY_CAPABILITY_CODES: 

77 verbose_logger.warning( 

78 "MCP server %s: the upstream authorization server rejected the request with " 

79 "invalid_target, meaning it did not accept the RFC 8707 resource indicator for this " 

80 "request. Set upstream_resource on this server to the exact resource identifier the " 

81 "authorization server expects (or to 'auto' to send the server's own canonical url); " 

82 "if it is already set and the authorization server does not support resource " 

83 "indicators, unset it and express the target audience through scopes instead", 

84 log_context, 

85 ) 

86 return GatewayRejected(code=code) 

87 if credential_source == "gateway_stored" and code in GATEWAY_CREDENTIAL_CODES: 

88 verbose_logger.warning( 

89 "MCP server %s: upstream authorization server rejected the gateway's configured client " 

90 "credentials (%s): %s", 

91 log_context, 

92 code, 

93 description or "<no description>", 

94 ) 

95 return GatewayRejected(code=code) 

96 return CallerRejected(code=code, description=description, error_uri=error_uri) 

97 

98 

99def classify_upstream_token_rejection( 

100 response: httpx.Response, 

101 credential_source: CredentialSource, 

102 log_context: str, 

103) -> UpstreamOAuthFault: 

104 """Classify a token-endpoint rejection into exactly one fault: a body with an RFC 6749 §5.2 

105 ``error`` field goes through blame assignment (:func:`_classify_oauth_error_code`); anything 

106 without a usable ``error`` field is an upstream protocol fault.""" 

107 parsed: Final = _safe_json(response) 

108 fields: Final = parsed if isinstance(parsed, dict) else {} 

109 code: Final = _bounded_field(fields.get("error")) 

110 if code is None: 

111 _log_out_of_contract("token", response, log_context) 

112 return UpstreamProtocolFault(note=f"upstream token endpoint returned HTTP {response.status_code}") 

113 return _classify_oauth_error_code( 

114 code, 

115 description=_bounded_field(fields.get("error_description")), 

116 error_uri=_bounded_field(fields.get("error_uri")), 

117 credential_source=credential_source, 

118 log_context=log_context, 

119 ) 

120 

121 

122def classify_upstream_dcr_rejection(response: httpx.Response, log_context: str) -> UpstreamOAuthFault: 

123 """Classify a dynamic-client-registration rejection. RFC 7591 §3.2.2 errors carry 

124 ``error`` / ``error_description`` and go through the same blame assignment as token errors 

125 (registration sends no client credentials, so credential codes stay caller-actionable); anything 

126 without a usable ``error`` field is a registration refusal for 401/403 and a protocol fault otherwise.""" 

127 parsed: Final = _safe_json(response) 

128 fields: Final = parsed if isinstance(parsed, dict) else {} 

129 code: Final = _bounded_field(fields.get("error")) 

130 if code is None: 

131 if response.status_code == 401 or response.status_code == 403: 

132 return UpstreamRegistrationRefused(status_code=response.status_code) 

133 _log_out_of_contract("registration", response, log_context) 

134 return UpstreamProtocolFault(note=f"upstream registration failed with HTTP {response.status_code}") 

135 return _classify_oauth_error_code( 

136 code, 

137 description=_bounded_field(fields.get("error_description")), 

138 error_uri=None, 

139 credential_source="caller_supplied", 

140 log_context=log_context, 

141 )