Add Google/Microsoft SSO OAuth for register and sign-in (#24)
Introduce OAuthIdentity storage, start/callback endpoints, JWT handoff to the SPA, and mocked IdP tests so Drive OAuth (#11) can reuse the same model.
This commit is contained in:
@@ -28,6 +28,18 @@ CAPTCHA_SECRET_KEY=
|
||||
# Self-serve sign-up (default false — set true to allow /user/create/)
|
||||
ENABLE_ACCOUNT_REGISTRATION=false
|
||||
|
||||
# OAuth SSO — Google / Microsoft (#24). Leave blank to hide SSO buttons.
|
||||
# Redirect URIs (register in each IdP console):
|
||||
# {OAUTH_CALLBACK_BASE_URL}/api/auth/oauth/google/callback/
|
||||
# {OAUTH_CALLBACK_BASE_URL}/api/auth/oauth/microsoft/callback/
|
||||
GOOGLE_OAUTH_CLIENT_ID=
|
||||
GOOGLE_OAUTH_CLIENT_SECRET=
|
||||
MICROSOFT_OAUTH_CLIENT_ID=
|
||||
MICROSOFT_OAUTH_CLIENT_SECRET=
|
||||
MICROSOFT_OAUTH_TENANT=common
|
||||
# Optional; defaults to request host. Example local: http://127.0.0.1:8001
|
||||
OAUTH_CALLBACK_BASE_URL=http://127.0.0.1:8001
|
||||
|
||||
# Stripe / finance (optional local — required for checkout + webhooks)
|
||||
STRIPE_SECRET_KEY=
|
||||
STRIPE_PUBLISHABLE_KEY=
|
||||
|
||||
@@ -49,6 +49,17 @@ CAPTCHA_SECRET_KEY=replace-with-captcha-secret
|
||||
# public registration; set true in chat_backend_prod.env / chat_backend_beta.env.
|
||||
ENABLE_ACCOUNT_REGISTRATION=false
|
||||
|
||||
# OAuth SSO — Google / Microsoft (#24). Never commit real secrets.
|
||||
# Register redirect URIs:
|
||||
# https://chatbackend.aimloperations.com/api/auth/oauth/google/callback/
|
||||
# https://chatbackend.aimloperations.com/api/auth/oauth/microsoft/callback/
|
||||
GOOGLE_OAUTH_CLIENT_ID=
|
||||
GOOGLE_OAUTH_CLIENT_SECRET=
|
||||
MICROSOFT_OAUTH_CLIENT_ID=
|
||||
MICROSOFT_OAUTH_CLIENT_SECRET=
|
||||
MICROSOFT_OAUTH_TENANT=common
|
||||
OAUTH_CALLBACK_BASE_URL=https://chatbackend.aimloperations.com
|
||||
|
||||
# Stripe / finance
|
||||
STRIPE_SECRET_KEY=replace-with-stripe-secret-key
|
||||
STRIPE_PUBLISHABLE_KEY=replace-with-stripe-publishable-key
|
||||
|
||||
@@ -13,6 +13,7 @@ from .models import (
|
||||
Document,
|
||||
UserAuthEvent,
|
||||
OutboundEmail,
|
||||
OAuthIdentity,
|
||||
)
|
||||
|
||||
|
||||
@@ -232,3 +233,22 @@ admin.site.register(Feedback, FeedbackAdmin)
|
||||
|
||||
admin.site.register(DocumentWorkspace, DocumentWorkspaceAdmin)
|
||||
admin.site.register(Document, DocumentAdmin)
|
||||
|
||||
|
||||
class OAuthIdentityAdmin(admin.ModelAdmin):
|
||||
model = OAuthIdentity
|
||||
list_display = (
|
||||
"provider",
|
||||
"email",
|
||||
"subject",
|
||||
"user",
|
||||
"token_expires_at",
|
||||
"created",
|
||||
"last_modified",
|
||||
)
|
||||
list_filter = ("provider",)
|
||||
search_fields = ("email", "subject", "user__email")
|
||||
readonly_fields = ("created", "last_modified", "raw_profile")
|
||||
|
||||
|
||||
admin.site.register(OAuthIdentity, OAuthIdentityAdmin)
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
# Generated by Django 6.0 on 2026-07-27 11:51
|
||||
|
||||
import django.db.models.deletion
|
||||
import django.utils.timezone
|
||||
from django.conf import settings
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('chat_backend', '0023_promptmetric_tokens_in_promptmetric_tokens_out'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='OAuthIdentity',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('created', models.DateTimeField(default=django.utils.timezone.now)),
|
||||
('last_modified', models.DateTimeField(default=django.utils.timezone.now)),
|
||||
('provider', models.CharField(choices=[('google', 'Google'), ('microsoft', 'Microsoft')], max_length=32)),
|
||||
('subject', models.CharField(help_text='OIDC subject (sub) from the identity provider', max_length=255)),
|
||||
('email', models.EmailField(blank=True, default='', max_length=254)),
|
||||
('access_token', models.TextField(blank=True, default='')),
|
||||
('refresh_token', models.TextField(blank=True, default='')),
|
||||
('token_expires_at', models.DateTimeField(blank=True, null=True)),
|
||||
('scopes', models.TextField(blank=True, default='')),
|
||||
('raw_profile', models.JSONField(blank=True, default=dict)),
|
||||
('user', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='oauth_identities', to=settings.AUTH_USER_MODEL)),
|
||||
],
|
||||
options={
|
||||
'verbose_name_plural': 'OAuth identities',
|
||||
'constraints': [models.UniqueConstraint(fields=('provider', 'subject'), name='uniq_oauth_provider_subject'), models.UniqueConstraint(fields=('provider', 'user'), name='uniq_oauth_provider_user')],
|
||||
},
|
||||
),
|
||||
]
|
||||
@@ -162,6 +162,47 @@ class OutboundEmail(models.Model):
|
||||
return f"{self.subject} → {self.to_email} ({self.status})"
|
||||
|
||||
|
||||
class OAuthIdentity(TimeInfoBase):
|
||||
"""Linked IdP identity + tokens (SSO now; Drive OAuth reuse later — #11)."""
|
||||
|
||||
class Provider(models.TextChoices):
|
||||
GOOGLE = "google", "Google"
|
||||
MICROSOFT = "microsoft", "Microsoft"
|
||||
|
||||
user = models.ForeignKey(
|
||||
CustomUser,
|
||||
on_delete=models.CASCADE,
|
||||
related_name="oauth_identities",
|
||||
)
|
||||
provider = models.CharField(max_length=32, choices=Provider.choices)
|
||||
subject = models.CharField(
|
||||
max_length=255,
|
||||
help_text="OIDC subject (sub) from the identity provider",
|
||||
)
|
||||
email = models.EmailField(blank=True, default="")
|
||||
access_token = models.TextField(blank=True, default="")
|
||||
refresh_token = models.TextField(blank=True, default="")
|
||||
token_expires_at = models.DateTimeField(null=True, blank=True)
|
||||
scopes = models.TextField(blank=True, default="")
|
||||
raw_profile = models.JSONField(default=dict, blank=True)
|
||||
|
||||
class Meta:
|
||||
constraints = [
|
||||
models.UniqueConstraint(
|
||||
fields=["provider", "subject"],
|
||||
name="uniq_oauth_provider_subject",
|
||||
),
|
||||
models.UniqueConstraint(
|
||||
fields=["provider", "user"],
|
||||
name="uniq_oauth_provider_user",
|
||||
),
|
||||
]
|
||||
verbose_name_plural = "OAuth identities"
|
||||
|
||||
def __str__(self):
|
||||
return f"{self.provider}:{self.subject} → {self.user_id}"
|
||||
|
||||
|
||||
FEEDBACK_CHOICE = (
|
||||
("SUBMITTED", "Submitted"),
|
||||
("RESOLVED", "Resolved"),
|
||||
|
||||
@@ -0,0 +1,383 @@
|
||||
"""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, 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"
|
||||
|
||||
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"
|
||||
)
|
||||
|
||||
|
||||
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) -> str:
|
||||
return signing.dumps(
|
||||
{"provider": provider, "intent": intent},
|
||||
salt=STATE_SALT,
|
||||
)
|
||||
|
||||
|
||||
def load_oauth_state(state: str) -> dict[str, str]:
|
||||
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 {"login", "signup"}:
|
||||
raise OAuthError("invalid_state", "Invalid OAuth intent.")
|
||||
return {"provider": provider, "intent": intent}
|
||||
|
||||
|
||||
def _microsoft_tenant() -> str:
|
||||
return settings.MICROSOFT_OAUTH_TENANT or "common"
|
||||
|
||||
|
||||
def build_authorization_url(*, provider: str, redirect_uri: str, state: str) -> str:
|
||||
if not provider_configured(provider):
|
||||
raise OAuthError("provider_not_configured", f"{provider} OAuth is not configured.")
|
||||
|
||||
if provider == OAuthIdentity.Provider.GOOGLE:
|
||||
params = {
|
||||
"client_id": settings.GOOGLE_OAUTH_CLIENT_ID,
|
||||
"redirect_uri": redirect_uri,
|
||||
"response_type": "code",
|
||||
"scope": "openid email profile",
|
||||
"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:
|
||||
params = {
|
||||
"client_id": settings.MICROSOFT_OAUTH_CLIENT_ID,
|
||||
"redirect_uri": redirect_uri,
|
||||
"response_type": "code",
|
||||
"response_mode": "query",
|
||||
"scope": "openid email profile offline_access",
|
||||
"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 _create_sso_user(profile: ProviderProfile) -> CustomUser:
|
||||
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()
|
||||
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
|
||||
@@ -0,0 +1,241 @@
|
||||
"""Tests for Google / Microsoft OAuth SSO (#24)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import patch
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
from django.test import override_settings
|
||||
from django.urls import reverse
|
||||
from rest_framework import status
|
||||
from rest_framework.test import APITestCase
|
||||
from rest_framework_simplejwt.tokens import AccessToken
|
||||
|
||||
from chat_backend.models import CustomUser, OAuthIdentity
|
||||
from chat_backend.oauth import ProviderProfile, dump_oauth_state
|
||||
from chat_backend.tests.factories import make_user
|
||||
|
||||
OAUTH_SETTINGS = {
|
||||
"GOOGLE_OAUTH_CLIENT_ID": "google-client-id",
|
||||
"GOOGLE_OAUTH_CLIENT_SECRET": "google-client-secret",
|
||||
"MICROSOFT_OAUTH_CLIENT_ID": "ms-client-id",
|
||||
"MICROSOFT_OAUTH_CLIENT_SECRET": "ms-client-secret",
|
||||
"MICROSOFT_OAUTH_TENANT": "common",
|
||||
"FRONTEND_BASE_URL": "http://frontend.test",
|
||||
"OAUTH_CALLBACK_BASE_URL": "http://backend.test",
|
||||
"ENABLE_ACCOUNT_REGISTRATION": True,
|
||||
}
|
||||
|
||||
|
||||
def _google_profile(**overrides) -> ProviderProfile:
|
||||
data = dict(
|
||||
provider=OAuthIdentity.Provider.GOOGLE,
|
||||
subject="google-sub-1",
|
||||
email="sso.user@example.com",
|
||||
email_verified=True,
|
||||
first_name="Sso",
|
||||
last_name="User",
|
||||
access_token="access-token",
|
||||
refresh_token="refresh-token",
|
||||
expires_in=3600,
|
||||
scopes="openid email profile",
|
||||
raw={"userinfo": {"sub": "google-sub-1"}},
|
||||
)
|
||||
data.update(overrides)
|
||||
return ProviderProfile(**data)
|
||||
|
||||
|
||||
@override_settings(**OAUTH_SETTINGS)
|
||||
class PublicSettingsOAuthTestCase(APITestCase):
|
||||
def test_exposes_configured_oauth_providers(self):
|
||||
response = self.client.get(reverse("public_settings"))
|
||||
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||
self.assertTrue(response.data["oauth"]["google"])
|
||||
self.assertTrue(response.data["oauth"]["microsoft"])
|
||||
|
||||
@override_settings(GOOGLE_OAUTH_CLIENT_ID="", GOOGLE_OAUTH_CLIENT_SECRET="")
|
||||
def test_hides_unconfigured_provider(self):
|
||||
response = self.client.get(reverse("public_settings"))
|
||||
self.assertFalse(response.data["oauth"]["google"])
|
||||
self.assertTrue(response.data["oauth"]["microsoft"])
|
||||
|
||||
|
||||
@override_settings(**OAUTH_SETTINGS)
|
||||
class OAuthStartTestCase(APITestCase):
|
||||
def test_start_redirects_to_google(self):
|
||||
response = self.client.get(
|
||||
reverse("oauth_start", kwargs={"provider": "google"}),
|
||||
{"intent": "login"},
|
||||
)
|
||||
self.assertEqual(response.status_code, status.HTTP_302_FOUND)
|
||||
location = response["Location"]
|
||||
self.assertIn("accounts.google.com", location)
|
||||
params = parse_qs(urlparse(location).query)
|
||||
self.assertEqual(params["client_id"], ["google-client-id"])
|
||||
self.assertIn("state", params)
|
||||
self.assertTrue(
|
||||
params["redirect_uri"][0].endswith("/api/auth/oauth/google/callback/")
|
||||
)
|
||||
|
||||
def test_start_unknown_provider_404(self):
|
||||
response = self.client.get(
|
||||
reverse("oauth_start", kwargs={"provider": "apple"})
|
||||
)
|
||||
self.assertEqual(response.status_code, status.HTTP_404_NOT_FOUND)
|
||||
|
||||
@override_settings(ENABLE_ACCOUNT_REGISTRATION=False)
|
||||
def test_signup_intent_blocked_when_registration_disabled(self):
|
||||
response = self.client.get(
|
||||
reverse("oauth_start", kwargs={"provider": "google"}),
|
||||
{"intent": "signup"},
|
||||
)
|
||||
self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN)
|
||||
|
||||
|
||||
@override_settings(**OAUTH_SETTINGS)
|
||||
class OAuthCallbackTestCase(APITestCase):
|
||||
def _callback(self, provider="google", intent="login", code="auth-code"):
|
||||
state = dump_oauth_state(provider=provider, intent=intent)
|
||||
return self.client.get(
|
||||
reverse("oauth_callback", kwargs={"provider": provider}),
|
||||
{"code": code, "state": state},
|
||||
)
|
||||
|
||||
def _assert_jwt_redirect(self, response, *, created: bool):
|
||||
self.assertEqual(response.status_code, status.HTTP_302_FOUND)
|
||||
location = response["Location"]
|
||||
self.assertTrue(location.startswith("http://frontend.test/auth/callback/?"))
|
||||
params = parse_qs(urlparse(location).query)
|
||||
self.assertIn("access", params)
|
||||
self.assertIn("refresh", params)
|
||||
self.assertEqual(params["created"], ["1" if created else "0"])
|
||||
access = AccessToken(params["access"][0])
|
||||
return access, params
|
||||
|
||||
@patch("chat_backend.views_oauth.exchange_code_for_profile")
|
||||
def test_callback_creates_user_and_returns_jwt(self, mock_exchange):
|
||||
mock_exchange.return_value = _google_profile()
|
||||
|
||||
response = self._callback(intent="signup")
|
||||
access, params = self._assert_jwt_redirect(response, created=True)
|
||||
self.assertEqual(params["needs_checkout"], ["1"])
|
||||
|
||||
user = CustomUser.objects.get(email="sso.user@example.com")
|
||||
self.assertEqual(int(access["user_id"]), user.id)
|
||||
self.assertFalse(user.has_usable_password())
|
||||
self.assertTrue(user.is_company_manager)
|
||||
self.assertEqual(user.company.name, "sso.user@example.com's workspace")
|
||||
|
||||
identity = OAuthIdentity.objects.get(provider="google", subject="google-sub-1")
|
||||
self.assertEqual(identity.user_id, user.id)
|
||||
self.assertEqual(identity.refresh_token, "refresh-token")
|
||||
mock_exchange.assert_called_once()
|
||||
|
||||
@patch("chat_backend.views_oauth.exchange_code_for_profile")
|
||||
def test_callback_logs_in_existing_identity(self, mock_exchange):
|
||||
user = make_user(email="sso.user@example.com", password="pass12345")
|
||||
OAuthIdentity.objects.create(
|
||||
user=user,
|
||||
provider=OAuthIdentity.Provider.GOOGLE,
|
||||
subject="google-sub-1",
|
||||
email=user.email,
|
||||
)
|
||||
mock_exchange.return_value = _google_profile()
|
||||
|
||||
response = self._callback(intent="login")
|
||||
access, params = self._assert_jwt_redirect(response, created=False)
|
||||
self.assertEqual(params["needs_checkout"], ["0"])
|
||||
self.assertEqual(int(access["user_id"]), user.id)
|
||||
self.assertEqual(CustomUser.objects.filter(email="sso.user@example.com").count(), 1)
|
||||
|
||||
@patch("chat_backend.views_oauth.exchange_code_for_profile")
|
||||
def test_callback_links_existing_email_account(self, mock_exchange):
|
||||
user = make_user(email="sso.user@example.com", password="pass12345")
|
||||
mock_exchange.return_value = _google_profile()
|
||||
|
||||
response = self._callback(intent="login")
|
||||
access, _params = self._assert_jwt_redirect(response, created=False)
|
||||
self.assertEqual(int(access["user_id"]), user.id)
|
||||
|
||||
identity = OAuthIdentity.objects.get(provider="google", subject="google-sub-1")
|
||||
self.assertEqual(identity.user_id, user.id)
|
||||
self.assertEqual(CustomUser.objects.count(), 1)
|
||||
|
||||
@patch("chat_backend.views_oauth.exchange_code_for_profile")
|
||||
def test_callback_rejects_unverified_email(self, mock_exchange):
|
||||
mock_exchange.return_value = _google_profile(email_verified=False)
|
||||
|
||||
response = self._callback(intent="signup")
|
||||
self.assertEqual(response.status_code, status.HTTP_302_FOUND)
|
||||
params = parse_qs(urlparse(response["Location"]).query)
|
||||
self.assertEqual(params["error"], ["email_unverified"])
|
||||
self.assertEqual(CustomUser.objects.count(), 0)
|
||||
self.assertEqual(OAuthIdentity.objects.count(), 0)
|
||||
|
||||
@patch("chat_backend.views_oauth.exchange_code_for_profile")
|
||||
def test_callback_rejects_missing_email(self, mock_exchange):
|
||||
mock_exchange.return_value = _google_profile(email="", email_verified=True)
|
||||
|
||||
response = self._callback(intent="signup")
|
||||
params = parse_qs(urlparse(response["Location"]).query)
|
||||
self.assertEqual(params["error"], ["email_missing"])
|
||||
|
||||
@patch("chat_backend.views_oauth.exchange_code_for_profile")
|
||||
@override_settings(ENABLE_ACCOUNT_REGISTRATION=False)
|
||||
def test_callback_login_without_account_when_registration_off(self, mock_exchange):
|
||||
mock_exchange.return_value = _google_profile()
|
||||
|
||||
response = self._callback(intent="login")
|
||||
params = parse_qs(urlparse(response["Location"]).query)
|
||||
self.assertEqual(params["error"], ["account_not_found"])
|
||||
self.assertEqual(CustomUser.objects.count(), 0)
|
||||
|
||||
@patch("chat_backend.views_oauth.exchange_code_for_profile")
|
||||
def test_link_conflict_when_provider_already_linked_to_other_subject(
|
||||
self, mock_exchange
|
||||
):
|
||||
user = make_user(email="sso.user@example.com", password="pass12345")
|
||||
OAuthIdentity.objects.create(
|
||||
user=user,
|
||||
provider=OAuthIdentity.Provider.GOOGLE,
|
||||
subject="other-google-sub",
|
||||
email=user.email,
|
||||
)
|
||||
mock_exchange.return_value = _google_profile(subject="google-sub-1")
|
||||
|
||||
response = self._callback(intent="login")
|
||||
params = parse_qs(urlparse(response["Location"]).query)
|
||||
self.assertEqual(params["error"], ["link_conflict"])
|
||||
|
||||
def test_provider_access_denied(self):
|
||||
response = self.client.get(
|
||||
reverse("oauth_callback", kwargs={"provider": "google"}),
|
||||
{"error": "access_denied", "error_description": "User cancelled"},
|
||||
)
|
||||
params = parse_qs(urlparse(response["Location"]).query)
|
||||
self.assertEqual(params["error"], ["access_denied"])
|
||||
|
||||
@patch("chat_backend.views_oauth.exchange_code_for_profile")
|
||||
def test_microsoft_callback_creates_user(self, mock_exchange):
|
||||
mock_exchange.return_value = ProviderProfile(
|
||||
provider=OAuthIdentity.Provider.MICROSOFT,
|
||||
subject="ms-oid-1",
|
||||
email="ms.user@example.com",
|
||||
email_verified=True,
|
||||
first_name="Ms",
|
||||
last_name="User",
|
||||
access_token="ms-access",
|
||||
refresh_token="ms-refresh",
|
||||
expires_in=3600,
|
||||
scopes="openid email profile offline_access",
|
||||
raw={},
|
||||
)
|
||||
response = self._callback(provider="microsoft", intent="signup")
|
||||
access, _params = self._assert_jwt_redirect(response, created=True)
|
||||
user = CustomUser.objects.get(email="ms.user@example.com")
|
||||
self.assertEqual(int(access["user_id"]), user.id)
|
||||
self.assertTrue(
|
||||
OAuthIdentity.objects.filter(
|
||||
provider="microsoft", subject="ms-oid-1", user=user
|
||||
).exists()
|
||||
)
|
||||
@@ -15,7 +15,6 @@ from .views import (
|
||||
ConversationDetailView,
|
||||
CompanyUsersView,
|
||||
SetUserPassword,
|
||||
ResetUserPassword,
|
||||
ConversationPreferences,
|
||||
UserPromptAnalytics,
|
||||
UserConversationAnalytics,
|
||||
@@ -25,6 +24,7 @@ from .views import (
|
||||
DocumentUploadView,
|
||||
DocumentDetailView,
|
||||
)
|
||||
from .views_oauth import OAuthCallbackView, OAuthStartView
|
||||
from rest_framework.routers import DefaultRouter
|
||||
|
||||
|
||||
@@ -33,6 +33,16 @@ urlpatterns = [
|
||||
path("token/refresh/", jwt_views.TokenRefreshView.as_view(), name="token_refresh"),
|
||||
path("user/create/", CustomUserCreate.as_view(), name="create_user"),
|
||||
path("public/settings/", PublicSettingsView.as_view(), name="public_settings"),
|
||||
path(
|
||||
"auth/oauth/<str:provider>/start/",
|
||||
OAuthStartView.as_view(),
|
||||
name="oauth_start",
|
||||
),
|
||||
path(
|
||||
"auth/oauth/<str:provider>/callback/",
|
||||
OAuthCallbackView.as_view(),
|
||||
name="oauth_callback",
|
||||
),
|
||||
path("user/invite/", CustomUserInvite.as_view(), name="invite_user"),
|
||||
path(
|
||||
"user/reset_password/", ResetUserPassword.as_view(), name="reset_password"
|
||||
|
||||
@@ -107,9 +107,12 @@ class PublicSettingsView(APIView):
|
||||
authentication_classes = ()
|
||||
|
||||
def get(self, request):
|
||||
from .views_oauth import oauth_public_flags
|
||||
|
||||
return Response(
|
||||
{
|
||||
"enable_account_registration": settings.ENABLE_ACCOUNT_REGISTRATION,
|
||||
"oauth": oauth_public_flags(),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@@ -0,0 +1,155 @@
|
||||
"""OAuth SSO start + callback views (#24)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from urllib.parse import urlencode
|
||||
|
||||
from django.conf import settings
|
||||
from django.http import HttpResponseRedirect
|
||||
from django.urls import reverse
|
||||
from rest_framework import permissions, status
|
||||
from rest_framework.response import Response
|
||||
from rest_framework.views import APIView
|
||||
from rest_framework_simplejwt.tokens import RefreshToken
|
||||
|
||||
from .models import OAuthIdentity
|
||||
from .oauth import (
|
||||
OAuthError,
|
||||
build_authorization_url,
|
||||
configured_providers,
|
||||
dump_oauth_state,
|
||||
exchange_code_for_profile,
|
||||
load_oauth_state,
|
||||
provider_configured,
|
||||
resolve_user_from_profile,
|
||||
upsert_identity,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _callback_redirect_uri(request, provider: str) -> str:
|
||||
"""Absolute backend callback URL registered with the IdP."""
|
||||
path = reverse("oauth_callback", kwargs={"provider": provider})
|
||||
base = (settings.OAUTH_CALLBACK_BASE_URL or "").rstrip("/")
|
||||
if base:
|
||||
return f"{base}{path}"
|
||||
return request.build_absolute_uri(path)
|
||||
|
||||
|
||||
def _frontend_callback_url(**params: str) -> str:
|
||||
base = settings.FRONTEND_BASE_URL.rstrip("/")
|
||||
query = urlencode({k: v for k, v in params.items() if v is not None and v != ""})
|
||||
return f"{base}/auth/callback/?{query}"
|
||||
|
||||
|
||||
def _redirect_error(code: str, message: str = "") -> HttpResponseRedirect:
|
||||
return HttpResponseRedirect(
|
||||
_frontend_callback_url(error=code, error_description=message or code)
|
||||
)
|
||||
|
||||
|
||||
class OAuthStartView(APIView):
|
||||
"""Redirect the browser to Google / Microsoft authorize URL."""
|
||||
|
||||
permission_classes = (permissions.AllowAny,)
|
||||
authentication_classes = ()
|
||||
http_method_names = ["get"]
|
||||
|
||||
def get(self, request, provider: str):
|
||||
provider = (provider or "").lower()
|
||||
if provider not in OAuthIdentity.Provider.values:
|
||||
return Response(
|
||||
{"detail": "Unsupported OAuth provider."},
|
||||
status=status.HTTP_404_NOT_FOUND,
|
||||
)
|
||||
if not provider_configured(provider):
|
||||
return Response(
|
||||
{"detail": f"{provider} OAuth is not configured."},
|
||||
status=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
)
|
||||
|
||||
intent = (request.query_params.get("intent") or "login").lower()
|
||||
if intent not in {"login", "signup"}:
|
||||
return Response(
|
||||
{"detail": "intent must be 'login' or 'signup'."},
|
||||
status=status.HTTP_400_BAD_REQUEST,
|
||||
)
|
||||
if intent == "signup" and not settings.ENABLE_ACCOUNT_REGISTRATION:
|
||||
return Response(
|
||||
{"detail": "Account registration is disabled."},
|
||||
status=status.HTTP_403_FORBIDDEN,
|
||||
)
|
||||
|
||||
state = dump_oauth_state(provider=provider, intent=intent)
|
||||
redirect_uri = _callback_redirect_uri(request, provider)
|
||||
try:
|
||||
auth_url = build_authorization_url(
|
||||
provider=provider, redirect_uri=redirect_uri, state=state
|
||||
)
|
||||
except OAuthError as exc:
|
||||
return Response({"detail": exc.message}, status=status.HTTP_400_BAD_REQUEST)
|
||||
|
||||
return HttpResponseRedirect(auth_url)
|
||||
|
||||
|
||||
class OAuthCallbackView(APIView):
|
||||
"""IdP redirect target — exchange code, create/link user, send JWTs to FE."""
|
||||
|
||||
permission_classes = (permissions.AllowAny,)
|
||||
authentication_classes = ()
|
||||
http_method_names = ["get"]
|
||||
|
||||
def get(self, request, provider: str):
|
||||
provider = (provider or "").lower()
|
||||
if provider not in OAuthIdentity.Provider.values:
|
||||
return _redirect_error("invalid_provider", "Unsupported OAuth provider.")
|
||||
|
||||
error = request.query_params.get("error")
|
||||
if error:
|
||||
description = request.query_params.get("error_description") or error
|
||||
code = "access_denied" if error == "access_denied" else "provider_error"
|
||||
return _redirect_error(code, description)
|
||||
|
||||
code = request.query_params.get("code")
|
||||
state = request.query_params.get("state")
|
||||
if not code or not state:
|
||||
return _redirect_error("missing_code", "Missing OAuth code or state.")
|
||||
|
||||
try:
|
||||
state_data = load_oauth_state(state)
|
||||
if state_data["provider"] != provider:
|
||||
raise OAuthError("invalid_state", "OAuth provider mismatch.")
|
||||
|
||||
redirect_uri = _callback_redirect_uri(request, provider)
|
||||
profile = exchange_code_for_profile(
|
||||
provider=provider, code=code, redirect_uri=redirect_uri
|
||||
)
|
||||
user, created = resolve_user_from_profile(
|
||||
profile=profile, intent=state_data["intent"]
|
||||
)
|
||||
# Refresh stored tokens on every successful login.
|
||||
upsert_identity(user, profile)
|
||||
except OAuthError as exc:
|
||||
logger.info("OAuth callback failed (%s): %s", exc.code, exc.message)
|
||||
return _redirect_error(exc.code, exc.message)
|
||||
except Exception:
|
||||
logger.exception("Unexpected OAuth callback failure")
|
||||
return _redirect_error("server_error", "Unexpected OAuth error.")
|
||||
|
||||
refresh = RefreshToken.for_user(user)
|
||||
needs_checkout = "1" if created else "0"
|
||||
return HttpResponseRedirect(
|
||||
_frontend_callback_url(
|
||||
access=str(refresh.access_token),
|
||||
refresh=str(refresh),
|
||||
created="1" if created else "0",
|
||||
needs_checkout=needs_checkout,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def oauth_public_flags() -> dict:
|
||||
"""Feature flags for /public/settings/."""
|
||||
return configured_providers()
|
||||
@@ -309,6 +309,19 @@ ALLOW_INTERNET_ACCESS = env_bool("ALLOW_INTERNET_ACCESS", True)
|
||||
# control-node secret (chat_backend_<env>.env) when ready for public sign-up.
|
||||
ENABLE_ACCOUNT_REGISTRATION = env_bool("ENABLE_ACCOUNT_REGISTRATION", False)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# OAuth SSO (Google / Microsoft) — #24
|
||||
# ---------------------------------------------------------------------------
|
||||
GOOGLE_OAUTH_CLIENT_ID = env("GOOGLE_OAUTH_CLIENT_ID", "") or ""
|
||||
GOOGLE_OAUTH_CLIENT_SECRET = env("GOOGLE_OAUTH_CLIENT_SECRET", "") or ""
|
||||
MICROSOFT_OAUTH_CLIENT_ID = env("MICROSOFT_OAUTH_CLIENT_ID", "") or ""
|
||||
MICROSOFT_OAUTH_CLIENT_SECRET = env("MICROSOFT_OAUTH_CLIENT_SECRET", "") or ""
|
||||
# Azure AD tenant: "common" (personal + work), "organizations", or a tenant ID.
|
||||
MICROSOFT_OAUTH_TENANT = env("MICROSOFT_OAUTH_TENANT", "common") or "common"
|
||||
# Public backend origin for IdP redirect URIs (e.g. https://chatbackend.aimloperations.com).
|
||||
# When empty, callback URLs are built from the incoming request.
|
||||
OAUTH_CALLBACK_BASE_URL = (env("OAUTH_CALLBACK_BASE_URL", "") or "").rstrip("/")
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Finance / Stripe (subscription billing)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user