Add Google/Microsoft SSO OAuth for register and sign-in (#24)
CI / test (pull_request) Successful in 10s
Unit Tests / test (pull_request) Successful in 10s

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:
2026-07-27 06:54:56 -05:00
parent 30ce3d048d
commit f8c29e09bb
11 changed files with 927 additions and 1 deletions
+20
View File
@@ -11,6 +11,7 @@ from .models import (
PromptMetric,
DocumentWorkspace,
Document,
OAuthIdentity,
)
# Register your models here.
@@ -145,3 +146,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')],
},
),
]
+41
View File
@@ -77,6 +77,47 @@ class CustomUser(AbstractUser):
return f"https://chat.aimloperations.com/set_password?slug={self.slug}"
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"),
+383
View File
@@ -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
+241
View File
@@ -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()
)
+11 -1
View File
@@ -15,7 +15,6 @@ from .views import (
ConversationDetailView,
CompanyUsersView,
SetUserPassword,
ResetUserPassword,
ConversationPreferences,
UserPromptAnalytics,
UserConversationAnalytics,
@@ -26,6 +25,7 @@ from .views import (
DocumentUploadView,
DocumentDetailView,
)
from .views_oauth import OAuthCallbackView, OAuthStartView
from rest_framework.routers import DefaultRouter
@@ -34,6 +34,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/", reset_password, name="reset_password"),
path(
+3
View File
@@ -100,9 +100,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(),
}
)
+155
View File
@@ -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()
+13
View File
@@ -301,6 +301,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)
# ---------------------------------------------------------------------------