Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/agent_endpoints/databricks_oauth.py: 29%

108 statements  

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

1""" 

2OAuth M2M (client_credentials) support for A2A agents that target Databricks 

3App endpoints. 

4 

5Databricks Apps reject static bearer tokens; they require a short-lived OAuth 

6access token minted from the workspace OIDC token endpoint. When an agent is 

7registered with a ``databricks_oauth`` block in its ``litellm_params``, LiteLLM 

8fetches that token via the client_credentials grant, caches it until shortly 

9before expiry, and attaches it as the outbound ``Authorization`` header on every 

10call the proxy makes to the agent. 

11 

12Config example:: 

13 

14 agents: 

15 - agent_name: my-databricks-app 

16 agent_card_params: 

17 url: https://my-app-1234.aws.databricksapps.com 

18 litellm_params: 

19 databricks_oauth: 

20 client_id: os.environ/DATABRICKS_CLIENT_ID 

21 client_secret: os.environ/DATABRICKS_CLIENT_SECRET 

22 workspace_url: https://dbc-abc123.cloud.databricks.com 

23""" 

24 

25import asyncio 

26import base64 

27import hashlib 

28from collections.abc import Mapping 

29from dataclasses import dataclass 

30from typing import Final 

31 

32import httpx 

33 

34from litellm._logging import verbose_logger 

35from litellm.caching.in_memory_cache import InMemoryCache 

36from litellm.llms.custom_httpx.http_handler import get_async_httpx_client 

37from litellm.secret_managers.main import get_secret_str 

38from litellm.types.llms.custom_http import httpxSpecialProvider 

39 

40DATABRICKS_OAUTH_PARAM: Final = "databricks_oauth" 

41 

42_DEFAULT_SCOPE: Final = "all-apis" 

43_TOKEN_EXPIRY_BUFFER_SECONDS: Final = 60 

44_DEFAULT_TTL_SECONDS: Final = 3600 

45 

46 

47def _resolve_secret(value: object) -> str | None: 

48 """Resolve a config value, expanding ``os.environ/`` references.""" 

49 if not isinstance(value, str): 

50 return None 

51 if value.startswith("os.environ/"): 

52 return get_secret_str(value) 

53 return value 

54 

55 

56def _token_url_from_workspace(workspace_url: str) -> str: 

57 """Build the workspace OIDC token endpoint from a workspace URL.""" 

58 base = workspace_url.strip().rstrip("/") 

59 base = base.removesuffix("/serving-endpoints") 

60 return f"{base}/oidc/v1/token" 

61 

62 

63@dataclass(frozen=True) 

64class DatabricksAppOAuthConfig: 

65 client_id: str 

66 client_secret: str 

67 token_url: str 

68 scope: str 

69 

70 @property 

71 def cache_key(self) -> str: 

72 # Include a digest of the secret so a rotated client_secret yields a new 

73 # key and forces a fresh token instead of serving the stale one. 

74 secret_digest: Final = hashlib.sha256(self.client_secret.encode()).hexdigest()[:16] 

75 return f"{self.token_url}|{self.client_id}|{self.scope}|{secret_digest}" 

76 

77 

78def parse_databricks_oauth_config( 

79 litellm_params: Mapping[str, object] | None, 

80) -> DatabricksAppOAuthConfig | None: 

81 """Build a Databricks App OAuth config from an agent's ``litellm_params``. 

82 

83 Returns ``None`` when the agent has no ``databricks_oauth`` block. Raises 

84 ``ValueError`` when the block is present but incomplete, so misconfiguration 

85 surfaces loudly instead of silently sending an unauthenticated request. 

86 """ 

87 if not litellm_params: 

88 return None 

89 

90 raw: Final = litellm_params.get(DATABRICKS_OAUTH_PARAM) 

91 if raw is None: 

92 return None 

93 if not isinstance(raw, dict): 

94 raise ValueError(f"'{DATABRICKS_OAUTH_PARAM}' must be a mapping of OAuth settings, got {type(raw).__name__}") 

95 

96 client_id: Final = _resolve_secret(raw.get("client_id")) 

97 client_secret: Final = _resolve_secret(raw.get("client_secret")) 

98 workspace_url: Final = _resolve_secret(raw.get("workspace_url")) 

99 

100 missing: Final = [ 

101 name 

102 for name, value in ( 

103 ("client_id", client_id), 

104 ("client_secret", client_secret), 

105 ("workspace_url", workspace_url), 

106 ) 

107 if not value 

108 ] 

109 if missing: 

110 raise ValueError(f"Databricks App OAuth config is missing required field(s): {', '.join(missing)}") 

111 

112 scope: Final = _resolve_secret(raw.get("scope")) or _DEFAULT_SCOPE 

113 

114 return DatabricksAppOAuthConfig( 

115 client_id=client_id, 

116 client_secret=client_secret, 

117 token_url=_token_url_from_workspace(workspace_url), 

118 scope=scope, 

119 ) 

120 

121 

122class DatabricksAppOAuthTokenCache(InMemoryCache): 

123 """In-memory cache for Databricks App OAuth client_credentials tokens. 

124 

125 Keyed by token endpoint + client_id + scope so distinct agents and service 

126 principals never share a token. A per-key ``asyncio.Lock`` collapses 

127 concurrent fetches into a single token request. 

128 """ 

129 

130 def __init__(self) -> None: 

131 super().__init__(default_ttl=_DEFAULT_TTL_SECONDS) 

132 self._locks: dict[str, asyncio.Lock] = {} 

133 

134 def _get_lock(self, cache_key: str) -> asyncio.Lock: 

135 return self._locks.setdefault(cache_key, asyncio.Lock()) 

136 

137 def _remove_key(self, key: str) -> None: 

138 # Drop the per-key lock alongside the cached token so ``_locks`` stays 

139 # bounded by the live key set rather than growing for every key ever seen. 

140 super()._remove_key(key) 

141 self._locks.pop(key, None) 

142 

143 def flush_cache(self) -> None: 

144 super().flush_cache() 

145 self._locks.clear() 

146 

147 async def async_get_token(self, config: DatabricksAppOAuthConfig) -> str: 

148 cache_key: Final = config.cache_key 

149 

150 cached = self.get_cache(cache_key) 

151 if cached is not None: 

152 return cached 

153 

154 async with self._get_lock(cache_key): 

155 cached = self.get_cache(cache_key) 

156 if cached is not None: 

157 return cached 

158 

159 token, ttl = await self._fetch_token(config) 

160 # ttl == 0 means the token's own lifetime is shorter than the 

161 # refresh buffer; skip caching so we never hand out a stale token, 

162 # and drop the lock we just created since no cached entry will ever 

163 # trigger _remove_key to clean it up. 

164 if ttl > 0: 

165 self.set_cache(cache_key, token, ttl=ttl) 

166 else: 

167 self._locks.pop(cache_key, None) 

168 return token 

169 

170 async def _fetch_token(self, config: DatabricksAppOAuthConfig) -> tuple[str, int]: 

171 client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.A2A) 

172 

173 verbose_logger.debug("Fetching Databricks App OAuth token from %s", config.token_url) 

174 

175 basic_auth: Final = base64.b64encode(f"{config.client_id}:{config.client_secret}".encode()).decode() 

176 try: 

177 response: Final = await client.post( 

178 config.token_url, 

179 data={ 

180 "grant_type": "client_credentials", 

181 "scope": config.scope, 

182 }, 

183 headers={ 

184 "Authorization": f"Basic {basic_auth}", 

185 "Content-Type": "application/x-www-form-urlencoded", 

186 }, 

187 ) 

188 except httpx.HTTPStatusError as exc: 

189 raise ValueError( 

190 f"Databricks App OAuth token request failed with status {exc.response.status_code}" 

191 ) from exc 

192 except httpx.HTTPError as exc: 

193 raise ValueError(f"Databricks App OAuth token request failed: {exc}") from exc 

194 

195 body: Final[object] = response.json() 

196 if not isinstance(body, dict): 

197 raise ValueError( 

198 f"Databricks App OAuth token response returned non-object JSON (got {type(body).__name__})" 

199 ) 

200 

201 access_token: Final = body.get("access_token") 

202 if not access_token: 

203 raise ValueError("Databricks App OAuth token response missing 'access_token'") 

204 

205 raw_expires_in: Final = body.get("expires_in") 

206 try: 

207 expires_in = int(raw_expires_in) if raw_expires_in is not None else _DEFAULT_TTL_SECONDS 

208 except (TypeError, ValueError): 

209 expires_in = _DEFAULT_TTL_SECONDS 

210 

211 ttl: Final = max(expires_in - _TOKEN_EXPIRY_BUFFER_SECONDS, 0) 

212 return access_token, ttl 

213 

214 

215databricks_app_oauth_token_cache: Final = DatabricksAppOAuthTokenCache() 

216 

217 

218async def resolve_databricks_app_auth_header( 

219 litellm_params: Mapping[str, object] | None, 

220) -> dict[str, str] | None: 

221 """Return ``{"Authorization": "Bearer <token>"}`` for a Databricks App agent. 

222 

223 Returns ``None`` when the agent is not configured for Databricks App OAuth. 

224 """ 

225 config: Final = parse_databricks_oauth_config(litellm_params) 

226 if config is None: 

227 return None 

228 

229 token: Final = await databricks_app_oauth_token_cache.async_get_token(config) 

230 return {"Authorization": f"Bearer {token}"}