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
« 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.
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.
12Config example::
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"""
25import asyncio
26import base64
27import hashlib
28from collections.abc import Mapping
29from dataclasses import dataclass
30from typing import Final
32import httpx
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
40DATABRICKS_OAUTH_PARAM: Final = "databricks_oauth"
42_DEFAULT_SCOPE: Final = "all-apis"
43_TOKEN_EXPIRY_BUFFER_SECONDS: Final = 60
44_DEFAULT_TTL_SECONDS: Final = 3600
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
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"
63@dataclass(frozen=True)
64class DatabricksAppOAuthConfig:
65 client_id: str
66 client_secret: str
67 token_url: str
68 scope: str
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}"
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``.
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
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__}")
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"))
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)}")
112 scope: Final = _resolve_secret(raw.get("scope")) or _DEFAULT_SCOPE
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 )
122class DatabricksAppOAuthTokenCache(InMemoryCache):
123 """In-memory cache for Databricks App OAuth client_credentials tokens.
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 """
130 def __init__(self) -> None:
131 super().__init__(default_ttl=_DEFAULT_TTL_SECONDS)
132 self._locks: dict[str, asyncio.Lock] = {}
134 def _get_lock(self, cache_key: str) -> asyncio.Lock:
135 return self._locks.setdefault(cache_key, asyncio.Lock())
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)
143 def flush_cache(self) -> None:
144 super().flush_cache()
145 self._locks.clear()
147 async def async_get_token(self, config: DatabricksAppOAuthConfig) -> str:
148 cache_key: Final = config.cache_key
150 cached = self.get_cache(cache_key)
151 if cached is not None:
152 return cached
154 async with self._get_lock(cache_key):
155 cached = self.get_cache(cache_key)
156 if cached is not None:
157 return cached
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
170 async def _fetch_token(self, config: DatabricksAppOAuthConfig) -> tuple[str, int]:
171 client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.A2A)
173 verbose_logger.debug("Fetching Databricks App OAuth token from %s", config.token_url)
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
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 )
201 access_token: Final = body.get("access_token")
202 if not access_token:
203 raise ValueError("Databricks App OAuth token response missing 'access_token'")
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
211 ttl: Final = max(expires_in - _TOKEN_EXPIRY_BUFFER_SECONDS, 0)
212 return access_token, ttl
215databricks_app_oauth_token_cache: Final = DatabricksAppOAuthTokenCache()
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.
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
229 token: Final = await databricks_app_oauth_token_cache.async_get_token(config)
230 return {"Authorization": f"Bearer {token}"}