Files
chat_backend/llm_be/chat_backend/oauth.py
T
westfarn 311e19abe6
CI / test (pull_request) Successful in 10s
Unit Tests / test (pull_request) Successful in 11s
Reject Drive link when user has no company (#55).
Avoid NotNullViolation on DriveConnection.company_id by returning a
no_company OAuth error before token exchange / upsert.
2026-08-02 05:35:28 -05:00

460 lines
16 KiB
Python

"""Google / Microsoft OIDC helpers for SSO (#24)."""
from __future__ import annotations
import logging
from dataclasses import dataclass
from datetime import timedelta
from typing import Any
from urllib.parse import urlencode
import httpx
import jwt
from django.conf import settings
from django.core import signing
from django.utils import timezone
from .models import Company, CustomUser, DriveConnection, OAuthIdentity
logger = logging.getLogger(__name__)
STATE_SALT = "chat_backend.oauth.state"
STATE_MAX_AGE_SECONDS = 600
GOOGLE_AUTH_URL = "https://accounts.google.com/o/oauth2/v2/auth"
GOOGLE_TOKEN_URL = "https://oauth2.googleapis.com/token"
GOOGLE_USERINFO_URL = "https://openidconnect.googleapis.com/v1/userinfo"
GOOGLE_DRIVE_READONLY_SCOPE = "https://www.googleapis.com/auth/drive.readonly"
MICROSOFT_AUTH_URL_TMPL = (
"https://login.microsoftonline.com/{tenant}/oauth2/v2.0/authorize"
)
MICROSOFT_TOKEN_URL_TMPL = (
"https://login.microsoftonline.com/{tenant}/oauth2/v2.0/token"
)
MICROSOFT_DRIVE_PERSONAL_SCOPE = "Files.Read"
MICROSOFT_DRIVE_COMPANY_SCOPE = "Files.Read.All Sites.Read.All"
# OAuth intents (#24 login/signup; #47 Drive linking).
LOGIN_INTENTS = {"login", "signup"}
DRIVE_LINK_INTENTS = {"link_drive", "link_company_drive"}
VALID_INTENTS = LOGIN_INTENTS | DRIVE_LINK_INTENTS
class OAuthError(Exception):
"""User-facing OAuth failure with a stable error code for the FE."""
def __init__(self, code: str, message: str = ""):
self.code = code
self.message = message or code
super().__init__(self.message)
@dataclass(frozen=True)
class ProviderProfile:
provider: str
subject: str
email: str
email_verified: bool
first_name: str
last_name: str
access_token: str
refresh_token: str
expires_in: int | None
scopes: str
raw: dict[str, Any]
def provider_configured(provider: str) -> bool:
if provider == OAuthIdentity.Provider.GOOGLE:
return bool(settings.GOOGLE_OAUTH_CLIENT_ID and settings.GOOGLE_OAUTH_CLIENT_SECRET)
if provider == OAuthIdentity.Provider.MICROSOFT:
return bool(
settings.MICROSOFT_OAUTH_CLIENT_ID and settings.MICROSOFT_OAUTH_CLIENT_SECRET
)
return False
def configured_providers() -> dict[str, bool]:
return {
OAuthIdentity.Provider.GOOGLE: provider_configured(OAuthIdentity.Provider.GOOGLE),
OAuthIdentity.Provider.MICROSOFT: provider_configured(
OAuthIdentity.Provider.MICROSOFT
),
}
def dump_oauth_state(
*, provider: str, intent: str, user_id: int | None = None
) -> str:
payload: dict[str, Any] = {"provider": provider, "intent": intent}
if user_id is not None:
payload["user_id"] = user_id
return signing.dumps(payload, salt=STATE_SALT)
def load_oauth_state(state: str) -> dict[str, Any]:
try:
data = signing.loads(state, salt=STATE_SALT, max_age=STATE_MAX_AGE_SECONDS)
except signing.BadSignature as exc:
raise OAuthError("invalid_state", "OAuth state is invalid or expired.") from exc
provider = data.get("provider")
intent = data.get("intent") or "login"
if provider not in OAuthIdentity.Provider.values:
raise OAuthError("invalid_state", "Unknown OAuth provider in state.")
if intent not in VALID_INTENTS:
raise OAuthError("invalid_state", "Invalid OAuth intent.")
result: dict[str, Any] = {"provider": provider, "intent": intent}
if intent in DRIVE_LINK_INTENTS:
user_id = data.get("user_id")
if not user_id:
raise OAuthError(
"invalid_state", "OAuth state is missing the linking user."
)
result["user_id"] = user_id
return result
def _microsoft_tenant() -> str:
return settings.MICROSOFT_OAUTH_TENANT or "common"
def build_authorization_url(
*, provider: str, redirect_uri: str, state: str, intent: str = "login"
) -> str:
if not provider_configured(provider):
raise OAuthError("provider_not_configured", f"{provider} OAuth is not configured.")
if provider == OAuthIdentity.Provider.GOOGLE:
scope = "openid email profile"
if intent in DRIVE_LINK_INTENTS:
scope = f"{scope} {GOOGLE_DRIVE_READONLY_SCOPE}"
params = {
"client_id": settings.GOOGLE_OAUTH_CLIENT_ID,
"redirect_uri": redirect_uri,
"response_type": "code",
"scope": scope,
"state": state,
"access_type": "offline",
"prompt": "select_account consent",
"include_granted_scopes": "true",
}
return f"{GOOGLE_AUTH_URL}?{urlencode(params)}"
if provider == OAuthIdentity.Provider.MICROSOFT:
scope = "openid email profile offline_access"
if intent == "link_drive":
scope = f"{scope} {MICROSOFT_DRIVE_PERSONAL_SCOPE}"
elif intent == "link_company_drive":
scope = f"{scope} {MICROSOFT_DRIVE_COMPANY_SCOPE}"
params = {
"client_id": settings.MICROSOFT_OAUTH_CLIENT_ID,
"redirect_uri": redirect_uri,
"response_type": "code",
"response_mode": "query",
"scope": scope,
"state": state,
"prompt": "select_account",
}
auth_url = MICROSOFT_AUTH_URL_TMPL.format(tenant=_microsoft_tenant())
return f"{auth_url}?{urlencode(params)}"
raise OAuthError("invalid_provider", f"Unsupported provider: {provider}")
def _decode_id_token_claims(id_token: str | None) -> dict[str, Any]:
if not id_token:
return {}
# Signature verified via TLS token endpoint + client secret; claims are trusted.
return jwt.decode(
id_token,
options={"verify_signature": False, "verify_aud": False},
)
def exchange_code_for_profile(
*, provider: str, code: str, redirect_uri: str
) -> ProviderProfile:
if provider == OAuthIdentity.Provider.GOOGLE:
return _exchange_google(code=code, redirect_uri=redirect_uri)
if provider == OAuthIdentity.Provider.MICROSOFT:
return _exchange_microsoft(code=code, redirect_uri=redirect_uri)
raise OAuthError("invalid_provider", f"Unsupported provider: {provider}")
def _exchange_google(*, code: str, redirect_uri: str) -> ProviderProfile:
with httpx.Client(timeout=20.0) as client:
token_response = client.post(
GOOGLE_TOKEN_URL,
data={
"code": code,
"client_id": settings.GOOGLE_OAUTH_CLIENT_ID,
"client_secret": settings.GOOGLE_OAUTH_CLIENT_SECRET,
"redirect_uri": redirect_uri,
"grant_type": "authorization_code",
},
)
if token_response.status_code >= 400:
logger.warning("Google token exchange failed: %s", token_response.text)
raise OAuthError("token_exchange_failed", "Google token exchange failed.")
token_data = token_response.json()
access_token = token_data.get("access_token") or ""
if not access_token:
raise OAuthError("token_exchange_failed", "Google did not return an access token.")
userinfo_response = client.get(
GOOGLE_USERINFO_URL,
headers={"Authorization": f"Bearer {access_token}"},
)
if userinfo_response.status_code >= 400:
logger.warning("Google userinfo failed: %s", userinfo_response.text)
raise OAuthError("profile_fetch_failed", "Could not load Google profile.")
profile = userinfo_response.json()
claims = _decode_id_token_claims(token_data.get("id_token"))
email = (profile.get("email") or claims.get("email") or "").strip().lower()
email_verified = bool(
profile.get("email_verified", claims.get("email_verified", False))
)
subject = str(profile.get("sub") or claims.get("sub") or "").strip()
if not subject:
raise OAuthError("profile_incomplete", "Google profile missing subject.")
return ProviderProfile(
provider=OAuthIdentity.Provider.GOOGLE,
subject=subject,
email=email,
email_verified=email_verified,
first_name=(profile.get("given_name") or claims.get("given_name") or "").strip(),
last_name=(profile.get("family_name") or claims.get("family_name") or "").strip(),
access_token=access_token,
refresh_token=token_data.get("refresh_token") or "",
expires_in=_as_int(token_data.get("expires_in")),
scopes=token_data.get("scope") or "openid email profile",
raw={"userinfo": profile, "id_token_claims": claims},
)
def _exchange_microsoft(*, code: str, redirect_uri: str) -> ProviderProfile:
token_url = MICROSOFT_TOKEN_URL_TMPL.format(tenant=_microsoft_tenant())
with httpx.Client(timeout=20.0) as client:
token_response = client.post(
token_url,
data={
"code": code,
"client_id": settings.MICROSOFT_OAUTH_CLIENT_ID,
"client_secret": settings.MICROSOFT_OAUTH_CLIENT_SECRET,
"redirect_uri": redirect_uri,
"grant_type": "authorization_code",
"scope": "openid email profile offline_access",
},
)
if token_response.status_code >= 400:
logger.warning("Microsoft token exchange failed: %s", token_response.text)
raise OAuthError("token_exchange_failed", "Microsoft token exchange failed.")
token_data = token_response.json()
claims = _decode_id_token_claims(token_data.get("id_token"))
email = (
claims.get("email")
or claims.get("preferred_username")
or claims.get("upn")
or ""
)
email = str(email).strip().lower()
# Microsoft issues verified tenant emails; treat presence as verified when claim missing.
email_verified = bool(claims.get("email_verified", True if email else False))
subject = str(claims.get("oid") or claims.get("sub") or "").strip()
if not subject:
raise OAuthError("profile_incomplete", "Microsoft profile missing subject.")
name = (claims.get("name") or "").strip()
first_name = (claims.get("given_name") or "").strip()
last_name = (claims.get("family_name") or "").strip()
if not first_name and name:
parts = name.split(" ", 1)
first_name = parts[0]
last_name = parts[1] if len(parts) > 1 else ""
return ProviderProfile(
provider=OAuthIdentity.Provider.MICROSOFT,
subject=subject,
email=email,
email_verified=email_verified,
first_name=first_name,
last_name=last_name,
access_token=token_data.get("access_token") or "",
refresh_token=token_data.get("refresh_token") or "",
expires_in=_as_int(token_data.get("expires_in")),
scopes=token_data.get("scope") or "openid email profile offline_access",
raw={"id_token_claims": claims},
)
def _as_int(value: Any) -> int | None:
try:
return int(value) if value is not None else None
except (TypeError, ValueError):
return None
def _token_expiry(expires_in: int | None):
if not expires_in:
return None
return timezone.now() + timedelta(seconds=expires_in)
def upsert_identity(user: CustomUser, profile: ProviderProfile) -> OAuthIdentity:
identity = OAuthIdentity.objects.filter(
provider=profile.provider, subject=profile.subject
).first()
expires_at = _token_expiry(profile.expires_in)
if identity:
identity.user = user
identity.email = profile.email
identity.access_token = profile.access_token
if profile.refresh_token:
identity.refresh_token = profile.refresh_token
identity.token_expires_at = expires_at
identity.scopes = profile.scopes
identity.raw_profile = profile.raw
identity.save()
return identity
return OAuthIdentity.objects.create(
user=user,
provider=profile.provider,
subject=profile.subject,
email=profile.email,
access_token=profile.access_token,
refresh_token=profile.refresh_token or "",
token_expires_at=expires_at,
scopes=profile.scopes,
raw_profile=profile.raw,
)
def upsert_drive_connection(
*, user: CustomUser, kind: str, profile: ProviderProfile
) -> DriveConnection:
"""Create/refresh a DriveConnection from a link_drive/link_company_drive callback (#47)."""
if not user.company_id:
raise OAuthError(
"no_company",
"A company is required before connecting a Drive account.",
)
expires_at = _token_expiry(profile.expires_in)
lookup_user = user if kind == DriveConnection.Kind.PERSONAL else None
connection = DriveConnection.objects.filter(
company=user.company,
provider=profile.provider,
kind=kind,
user=lookup_user,
).first()
if connection is None:
connection = DriveConnection(
company=user.company,
provider=profile.provider,
kind=kind,
user=lookup_user,
)
connection.access_token = profile.access_token
if profile.refresh_token:
connection.refresh_token = profile.refresh_token
connection.token_expires_at = expires_at
connection.scopes = profile.scopes
connection.external_account_email = profile.email
connection.is_active = True
connection.last_sync_error = ""
connection.save()
return connection
def _create_sso_user(profile: ProviderProfile) -> CustomUser:
from finance.services.plans import try_redeem_backer_email
company = Company.objects.create(
name=f"{profile.email}'s workspace",
state="NA",
zipcode="00000",
address="N/A",
)
user = CustomUser(
username=profile.email,
email=profile.email,
first_name=profile.first_name,
last_name=profile.last_name,
company=company,
is_company_manager=True,
)
user.set_unusable_password()
user.save()
try_redeem_backer_email(user)
return user
def resolve_link_user(user_id: int) -> CustomUser:
"""Load the authenticated user a Drive-link callback should attach to."""
user = CustomUser.objects.filter(pk=user_id, deleted=False).first()
if user is None:
raise OAuthError("user_not_found", "Linking user account was not found.")
return user
def resolve_user_from_profile(
*, profile: ProviderProfile, intent: str
) -> tuple[CustomUser, bool]:
"""
Map IdP profile → CustomUser.
Returns (user, created).
"""
if not profile.email:
raise OAuthError("email_missing", "Email was not provided by the identity provider.")
if not profile.email_verified:
raise OAuthError("email_unverified", "Email from the identity provider is not verified.")
existing_identity = (
OAuthIdentity.objects.select_related("user")
.filter(provider=profile.provider, subject=profile.subject)
.first()
)
if existing_identity:
return existing_identity.user, False
email_user = (
CustomUser.objects.filter(email__iexact=profile.email).first()
or CustomUser.objects.filter(username__iexact=profile.email).first()
)
if email_user:
# Same provider already linked to a different subject → unsafe collision.
other = (
OAuthIdentity.objects.filter(provider=profile.provider, user=email_user)
.exclude(subject=profile.subject)
.first()
)
if other:
raise OAuthError(
"link_conflict",
"This email is already linked to a different identity for this provider.",
)
upsert_identity(email_user, profile)
return email_user, False
# New account path
allow_create = settings.ENABLE_ACCOUNT_REGISTRATION
if intent == "signup" and not allow_create:
raise OAuthError("registration_disabled", "Account registration is disabled.")
if intent == "login" and not allow_create:
raise OAuthError(
"account_not_found",
"No account exists for this email. Contact your administrator.",
)
if not allow_create:
raise OAuthError("registration_disabled", "Account registration is disabled.")
user = _create_sso_user(profile)
upsert_identity(user, profile)
return user, True