"""OAuth flow logic — callback, start, disconnect, refresh.""" from __future__ import annotations import base64 import hashlib import json import logging import os import secrets import time from typing import Any, Optional from urllib.parse import urlencode import httpx from backend.ports import BACKEND_DEV_PORT from fastapi import HTTPException, Query from fastapi.responses import HTMLResponse from backend.apps.tools_lib.oauth_providers import _resolve_oauth_provider from backend.apps.tools_lib.routes import _store logger = logging.getLogger(__name__) _pending_oauth: dict[str, str] = {} _pending_pkce: dict[str, str] = {} def _get_store(): return _store async def oauth_callback(code: str = Query(...), state: str = Query("")): tool_id = _pending_oauth.pop(state, None) if not tool_id: tool_id = _pending_oauth.pop(state.split(":")[-1] if ":" in state else state, None) if not tool_id: return HTMLResponse("
{resp.text}", status_code=400)
tokens = resp.json()
access_token = tokens.get("access_token", "")
if provider.token_response_path and not access_token:
obj = tokens
for part in provider.token_response_path.split("."):
obj = obj.get(part, {}) if isinstance(obj, dict) else ""
if isinstance(obj, str) and obj:
access_token = obj
tool.oauth_tokens = {
"access_token": access_token,
"refresh_token": tokens.get("refresh_token", ""),
"token_expiry": time.time() + tokens.get("expires_in", 3600),
}
for response_path, env_var in provider.extra_token_fields.items():
obj_val: Any = tokens
for part in response_path.split("."):
obj_val = obj_val.get(part, "") if isinstance(obj_val, dict) else ""
if obj_val:
tool.oauth_tokens[env_var] = str(obj_val)
tool.auth_status = "connected"
if access_token and provider.userinfo_url:
try:
async with httpx.AsyncClient(timeout=10.0) as info_client:
info_resp = await info_client.get(
provider.userinfo_url,
headers={"Authorization": f"Bearer {access_token}"},
)
if info_resp.status_code == 200:
tool.connected_account_email = info_resp.json().get(provider.userinfo_field)
except Exception as e:
logger.warning(f"Failed to fetch userinfo for {tool.oauth_provider or 'google'}: {e}")
if (tool.oauth_provider or "google") == "notion" and not tool.connected_account_email:
workspace_name = tokens.get("workspace_name")
if workspace_name:
tool.connected_account_email = workspace_name
_get_store().save(tool)
return HTMLResponse("""
You can close this window.
""") async def oauth_disconnect(tool_id: str): tool = _get_store().load(tool_id) access_token = tool.oauth_tokens.get("access_token") if access_token: provider = _resolve_oauth_provider(tool) revoke_url = provider.revoke_url or "https://oauth2.googleapis.com/revoke" try: async with httpx.AsyncClient(timeout=10.0) as client: await client.post( revoke_url, params={"token": access_token}, headers={"Content-Type": "application/x-www-form-urlencoded"}, ) except Exception as e: logger.warning(f"Failed to revoke token for tool {tool.id}: {e}") tool.oauth_tokens = {} tool.auth_status = "configured" tool.connected_account_email = None _get_store().save(tool) return {"ok": True, "tool": tool.model_dump()} async def oauth_start(tool_id: str): tool = _get_store().load(tool_id) provider = _resolve_oauth_provider(tool) client_id = os.environ.get(provider.client_id_env, "") if not client_id: raise HTTPException(status_code=400, detail=f"{provider.client_id_env} not set in backend .env") _port = os.environ.get("OPENSWARM_PORT", str(BACKEND_DEV_PORT)) redirect_uri = f"http://localhost:{_port}/api/tools/oauth/callback" provider_key = tool.oauth_provider or "google" state = f"{provider_key}:{tool_id}" _pending_oauth[state] = tool_id params = { "client_id": client_id, "redirect_uri": redirect_uri, "response_type": "code", "state": state, **provider.extra_auth_params, } if provider.scopes: params["scope"] = " ".join(provider.scopes) if provider.pkce_required: code_verifier = secrets.token_urlsafe(64) code_challenge = base64.urlsafe_b64encode( hashlib.sha256(code_verifier.encode()).digest() ).rstrip(b"=").decode() params["code_challenge"] = code_challenge params["code_challenge_method"] = "S256" _pending_pkce[state] = code_verifier auth_url = f"{provider.auth_url}?{urlencode(params)}" return {"auth_url": auth_url} async def refresh_oauth_token(tool) -> Optional[str]: """Refresh an expired OAuth token. Returns the fresh access_token or None.""" if tool.auth_type != "oauth2": return None refresh_token = tool.oauth_tokens.get("refresh_token") if not refresh_token: return None expiry = tool.oauth_tokens.get("token_expiry", 0) if time.time() < expiry - 60: return tool.oauth_tokens.get("access_token") provider = _resolve_oauth_provider(tool) client_id = os.environ.get(provider.client_id_env, "") client_secret = os.environ.get(provider.client_secret_env, "") if not client_id or not client_secret: return None try: async with httpx.AsyncClient(timeout=15.0) as client: resp = await client.post(provider.token_url, data={ "client_id": client_id, "client_secret": client_secret, "refresh_token": refresh_token, "grant_type": "refresh_token", }) if resp.status_code == 200: data = resp.json() new_token = data["access_token"] tool.oauth_tokens["access_token"] = new_token tool.oauth_tokens["token_expiry"] = time.time() + data.get("expires_in", 3600) if not tool.connected_account_email and provider.userinfo_url: try: async with httpx.AsyncClient(timeout=10.0) as info_client: info_resp = await info_client.get( provider.userinfo_url, headers={"Authorization": f"Bearer {new_token}"}, ) if info_resp.status_code == 200: tool.connected_account_email = info_resp.json().get(provider.userinfo_field) except Exception: pass _get_store().save(tool) return new_token except Exception as e: logger.warning(f"OAuth token refresh failed for tool {tool.id}: {e}") return None refresh_google_token = refresh_oauth_token