Skip to content

Commit 2d8f3d4

Browse files
committed
refactor(oauth-config): use shared authorization attempts
1 parent b140fcf commit 2d8f3d4

4 files changed

Lines changed: 345 additions & 46 deletions

File tree

backend/onyx/auth/oauth_token_manager.py

Lines changed: 24 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -276,8 +276,14 @@ def is_token_expired(cls, token_data: dict[str, Any]) -> bool:
276276
# Add 60 second buffer to avoid race conditions
277277
return int(time.time()) + 60 >= expires_at
278278

279-
def exchange_code_for_token(self, code: str, redirect_uri: str) -> dict[str, Any]:
280-
"""Exchange authorization code for access token"""
279+
def exchange_code_for_token(
280+
self,
281+
code: str,
282+
redirect_uri: str,
283+
*,
284+
code_verifier: str | None = None,
285+
) -> dict[str, Any]:
286+
"""Exchange an authorization code, including a PKCE verifier when supplied."""
281287
if (
282288
self.oauth_config.client_id is None
283289
or self.oauth_config.client_secret is None
@@ -287,22 +293,33 @@ def exchange_code_for_token(self, code: str, redirect_uri: str) -> dict[str, Any
287293
)
288294

289295
return exchange_oauth_code_for_token(
290-
self._flow_params(self.oauth_config), code, redirect_uri
296+
self.flow_params(self.oauth_config),
297+
code,
298+
redirect_uri,
299+
code_verifier=code_verifier,
291300
)
292301

293302
@staticmethod
294303
def build_authorization_url(
295-
oauth_config: OAuthConfig, redirect_uri: str, state: str
304+
oauth_config: OAuthConfig,
305+
redirect_uri: str,
306+
state: str,
307+
*,
308+
code_challenge: str | None = None,
296309
) -> str:
297-
"""Build OAuth authorization URL"""
310+
"""Build an authorization URL, including a PKCE challenge when supplied."""
298311
if oauth_config.client_id is None:
299312
raise ValueError("OAuth client_id is required to build authorization URL")
300313
return build_oauth_authorization_url(
301-
OAuthTokenManager._flow_params(oauth_config), redirect_uri, state
314+
OAuthTokenManager.flow_params(oauth_config),
315+
redirect_uri,
316+
state,
317+
code_challenge=code_challenge,
302318
)
303319

304320
@staticmethod
305-
def _flow_params(oauth_config: OAuthConfig) -> OAuthFlowParams:
321+
def flow_params(oauth_config: OAuthConfig) -> OAuthFlowParams:
322+
"""Return the protocol inputs represented by an OAuthConfig."""
306323
if oauth_config.client_id is None:
307324
raise ValueError("OAuth client_id is required")
308325
client_secret = (

backend/onyx/server/features/oauth_config/api.py

Lines changed: 100 additions & 38 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,18 @@
11
"""API endpoints for OAuth configuration management."""
22

3+
import secrets
4+
35
from fastapi import APIRouter, Depends, HTTPException
6+
from pydantic import BaseModel, ConfigDict
47
from sqlalchemy.orm import Session
58

6-
from onyx.auth.oauth_token_manager import OAuthTokenManager
9+
from onyx.auth.oauth_token_manager import (
10+
OAuthTokenManager,
11+
conflicting_authorization_params,
12+
)
713
from onyx.auth.permissions import has_global_permission, require_permission
14+
from onyx.auth.pkce import generate_pkce_pair
15+
from onyx.cache.factory import get_cache_backend
816
from onyx.configs.app_configs import WEB_DOMAIN
917
from onyx.db.engine.sql_engine import get_session
1018
from onyx.db.enums import Permission
@@ -21,9 +29,15 @@
2129
)
2230
from onyx.error_handling.error_codes import OnyxErrorCode
2331
from onyx.error_handling.exceptions import OnyxError
24-
from onyx.federated_connectors.oauth_utils import (
25-
generate_oauth_state,
26-
verify_oauth_state,
32+
from onyx.oauth.authorization_attempt import (
33+
AuthorizationAttemptStore,
34+
canonical_json_fingerprint,
35+
generate_authorization_state,
36+
)
37+
from onyx.oauth.models import (
38+
OAuthConfigurationFingerprint,
39+
PKCECodeVerifier,
40+
SafeOAuthReturnPath,
2741
)
2842
from onyx.server.features.oauth_config.models import (
2943
OAuthCallbackResponse,
@@ -40,6 +54,47 @@
4054
admin_router = APIRouter(prefix="/admin/oauth-config")
4155
router = APIRouter(prefix="/oauth-config")
4256

57+
_OAUTH_CALLBACK_PATH = "/oauth-config/callback"
58+
59+
60+
class _OAuthConfigAttemptPayload(BaseModel):
61+
model_config = ConfigDict(extra="forbid", frozen=True)
62+
63+
oauth_config_id: int
64+
return_path: SafeOAuthReturnPath
65+
configuration_fingerprint: OAuthConfigurationFingerprint
66+
code_verifier: PKCECodeVerifier
67+
68+
69+
_AUTHORIZATION_ATTEMPTS = AuthorizationAttemptStore(
70+
cache_backend_provider=lambda: get_cache_backend(),
71+
namespace="oauth-config",
72+
payload_type=_OAuthConfigAttemptPayload,
73+
)
74+
75+
76+
def _oauth_callback_url() -> str:
77+
return f"{WEB_DOMAIN}{_OAUTH_CALLBACK_PATH}"
78+
79+
80+
def _oauth_config_fingerprint(oauth_config: OAuthConfig) -> str:
81+
return canonical_json_fingerprint(
82+
{
83+
"redirect_uri": _oauth_callback_url(),
84+
"flow": OAuthTokenManager.flow_params(oauth_config).model_dump(mode="json"),
85+
},
86+
)
87+
88+
89+
def _validate_additional_authorization_params(oauth_config: OAuthConfig) -> None:
90+
reserved = conflicting_authorization_params(oauth_config.additional_params)
91+
if reserved:
92+
raise OnyxError(
93+
OnyxErrorCode.INVALID_INPUT,
94+
"OAuth additional parameters cannot override: "
95+
f"{', '.join(sorted(reserved))}",
96+
)
97+
4398

4499
def _oauth_config_to_snapshot(
45100
oauth_config: OAuthConfig, db_session: Session
@@ -215,18 +270,25 @@ def initiate_oauth_flow(
215270
detail=f"OAuth config with id {request.oauth_config_id} not found",
216271
)
217272

218-
# Generate state parameter and store in Redis
219-
state = generate_oauth_state(
220-
federated_connector_id=request.oauth_config_id,
221-
user_id=str(user.id),
222-
redirect_uri=request.return_path,
223-
additional_data={"oauth_config_id": request.oauth_config_id},
224-
)
273+
_validate_additional_authorization_params(oauth_config)
274+
code_verifier, code_challenge = generate_pkce_pair()
275+
state = generate_authorization_state()
225276

226-
# Build authorization URL
227-
redirect_uri = f"{WEB_DOMAIN}/oauth-config/callback"
228277
authorization_url = OAuthTokenManager.build_authorization_url(
229-
oauth_config, redirect_uri, state
278+
oauth_config,
279+
_oauth_callback_url(),
280+
state,
281+
code_challenge=code_challenge,
282+
)
283+
_AUTHORIZATION_ATTEMPTS.store(
284+
owner_id=str(user.id),
285+
state=state,
286+
payload=_OAuthConfigAttemptPayload(
287+
oauth_config_id=oauth_config.id,
288+
return_path=request.return_path,
289+
configuration_fingerprint=_oauth_config_fingerprint(oauth_config),
290+
code_verifier=code_verifier,
291+
),
230292
)
231293

232294
return OAuthInitiateResponse(authorization_url=authorization_url, state=state)
@@ -245,39 +307,39 @@ def handle_oauth_callback(
245307
Exchanges the authorization code for an access token and stores it.
246308
Accepts code and state as query parameters (standard OAuth flow).
247309
"""
248-
try:
249-
# Verify state and retrieve session data
250-
session = verify_oauth_state(state)
251-
252-
# Verify the user_id matches
253-
if str(user.id) != session.user_id:
254-
raise HTTPException(
255-
status_code=403, detail="User mismatch in OAuth callback"
256-
)
257-
258-
# Extract oauth_config_id from session (stored during initiate)
259-
oauth_config_id = session.federated_connector_id
260-
261-
# Get OAuth config
262-
oauth_config = get_oauth_config(oauth_config_id, db_session)
263-
if not oauth_config:
264-
raise HTTPException(
265-
status_code=404,
266-
detail=f"OAuth config with id {oauth_config_id} not found",
267-
)
310+
attempt = _AUTHORIZATION_ATTEMPTS.consume(owner_id=str(user.id), state=state)
311+
payload = attempt.payload
312+
313+
oauth_config = get_oauth_config(payload.oauth_config_id, db_session)
314+
if not oauth_config:
315+
raise OnyxError(
316+
OnyxErrorCode.NOT_FOUND,
317+
f"OAuth config with id {payload.oauth_config_id} not found",
318+
)
319+
if not secrets.compare_digest(
320+
payload.configuration_fingerprint,
321+
_oauth_config_fingerprint(oauth_config),
322+
):
323+
raise OnyxError(
324+
OnyxErrorCode.INVALID_INPUT,
325+
"OAuth configuration changed while authorization was pending.",
326+
)
268327

328+
try:
269329
# Exchange code for token
270-
redirect_uri = f"{WEB_DOMAIN}/oauth-config/callback"
271330
token_manager = OAuthTokenManager(oauth_config, user.id, db_session)
272-
token_data = token_manager.exchange_code_for_token(code, redirect_uri)
331+
token_data = token_manager.exchange_code_for_token(
332+
code,
333+
_oauth_callback_url(),
334+
code_verifier=payload.code_verifier,
335+
)
273336

274337
# Store token
275338
upsert_user_oauth_token(oauth_config.id, user.id, token_data, db_session)
276339

277340
# Return success with redirect
278-
return_path = session.redirect_uri or "/chat"
279341
return OAuthCallbackResponse(
280-
redirect_url=return_path,
342+
redirect_url=payload.return_path,
281343
)
282344

283345
except ValueError as e:

backend/onyx/server/features/oauth_config/models.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,8 @@
33

44
from pydantic import BaseModel
55

6+
from onyx.oauth.models import SafeOAuthReturnPath
7+
68

79
class OAuthConfigCreate(BaseModel):
810
name: str
@@ -40,7 +42,7 @@ class OAuthConfigSnapshot(BaseModel):
4042

4143
class OAuthInitiateRequest(BaseModel):
4244
oauth_config_id: int
43-
return_path: str = "/chat" # Where to redirect after OAuth flow
45+
return_path: SafeOAuthReturnPath = "/chat"
4446

4547

4648
class OAuthInitiateResponse(BaseModel):

0 commit comments

Comments
 (0)