From be6ff471ad45416a03090d3e8274d3cc342fa8b3 Mon Sep 17 00:00:00 2001 From: Ryan Westfall Date: Mon, 27 Jul 2026 06:54:56 -0500 Subject: [PATCH 1/2] 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. --- .env.example | 12 + .env.prod.example | 11 + llm_be/chat_backend/admin.py | 20 + .../migrations/0024_oauthidentity.py | 37 ++ llm_be/chat_backend/models.py | 41 ++ llm_be/chat_backend/oauth.py | 383 ++++++++++++++++++ llm_be/chat_backend/tests/test_oauth.py | 241 +++++++++++ llm_be/chat_backend/urls.py | 12 +- llm_be/chat_backend/views.py | 3 + llm_be/chat_backend/views_oauth.py | 155 +++++++ llm_be/llm_be/settings.py | 13 + 11 files changed, 927 insertions(+), 1 deletion(-) create mode 100644 llm_be/chat_backend/migrations/0024_oauthidentity.py create mode 100644 llm_be/chat_backend/oauth.py create mode 100644 llm_be/chat_backend/tests/test_oauth.py create mode 100644 llm_be/chat_backend/views_oauth.py diff --git a/.env.example b/.env.example index 99d587a..0202b28 100644 --- a/.env.example +++ b/.env.example @@ -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= diff --git a/.env.prod.example b/.env.prod.example index de9d58d..63cbdde 100644 --- a/.env.prod.example +++ b/.env.prod.example @@ -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 diff --git a/llm_be/chat_backend/admin.py b/llm_be/chat_backend/admin.py index 285539a..9760a5f 100644 --- a/llm_be/chat_backend/admin.py +++ b/llm_be/chat_backend/admin.py @@ -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) diff --git a/llm_be/chat_backend/migrations/0024_oauthidentity.py b/llm_be/chat_backend/migrations/0024_oauthidentity.py new file mode 100644 index 0000000..539e82c --- /dev/null +++ b/llm_be/chat_backend/migrations/0024_oauthidentity.py @@ -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')], + }, + ), + ] diff --git a/llm_be/chat_backend/models.py b/llm_be/chat_backend/models.py index ecb934b..60b4e5a 100644 --- a/llm_be/chat_backend/models.py +++ b/llm_be/chat_backend/models.py @@ -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"), diff --git a/llm_be/chat_backend/oauth.py b/llm_be/chat_backend/oauth.py new file mode 100644 index 0000000..9339275 --- /dev/null +++ b/llm_be/chat_backend/oauth.py @@ -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 diff --git a/llm_be/chat_backend/tests/test_oauth.py b/llm_be/chat_backend/tests/test_oauth.py new file mode 100644 index 0000000..662d13c --- /dev/null +++ b/llm_be/chat_backend/tests/test_oauth.py @@ -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() + ) diff --git a/llm_be/chat_backend/urls.py b/llm_be/chat_backend/urls.py index 41418f3..7672d0c 100644 --- a/llm_be/chat_backend/urls.py +++ b/llm_be/chat_backend/urls.py @@ -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//start/", + OAuthStartView.as_view(), + name="oauth_start", + ), + path( + "auth/oauth//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" diff --git a/llm_be/chat_backend/views.py b/llm_be/chat_backend/views.py index 8e82656..d41299a 100644 --- a/llm_be/chat_backend/views.py +++ b/llm_be/chat_backend/views.py @@ -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(), } ) diff --git a/llm_be/chat_backend/views_oauth.py b/llm_be/chat_backend/views_oauth.py new file mode 100644 index 0000000..781d13b --- /dev/null +++ b/llm_be/chat_backend/views_oauth.py @@ -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() diff --git a/llm_be/llm_be/settings.py b/llm_be/llm_be/settings.py index 0538678..79237a7 100644 --- a/llm_be/llm_be/settings.py +++ b/llm_be/llm_be/settings.py @@ -309,6 +309,19 @@ ALLOW_INTERNET_ACCESS = env_bool("ALLOW_INTERNET_ACCESS", True) # control-node secret (chat_backend_.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) # --------------------------------------------------------------------------- -- 2.54.0 From bbe805557367e86a518cad295632b9c1a9aff0e3 Mon Sep 17 00:00:00 2001 From: Ryan Westfall Date: Mon, 27 Jul 2026 07:12:45 -0500 Subject: [PATCH 2/2] Fix migration conflict with password-reset leaves Renumber OAuthIdentity to 0026 after 0025_outbound_email (master already shipped 0024_user_auth_event), and restore missing ResetUserPassword import. --- .../migrations/0024_oauthidentity.py | 37 --------- .../migrations/0026_oauthidentity.py | 75 +++++++++++++++++++ llm_be/chat_backend/urls.py | 1 + 3 files changed, 76 insertions(+), 37 deletions(-) delete mode 100644 llm_be/chat_backend/migrations/0024_oauthidentity.py create mode 100644 llm_be/chat_backend/migrations/0026_oauthidentity.py diff --git a/llm_be/chat_backend/migrations/0024_oauthidentity.py b/llm_be/chat_backend/migrations/0024_oauthidentity.py deleted file mode 100644 index 539e82c..0000000 --- a/llm_be/chat_backend/migrations/0024_oauthidentity.py +++ /dev/null @@ -1,37 +0,0 @@ -# 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')], - }, - ), - ] diff --git a/llm_be/chat_backend/migrations/0026_oauthidentity.py b/llm_be/chat_backend/migrations/0026_oauthidentity.py new file mode 100644 index 0000000..4321e11 --- /dev/null +++ b/llm_be/chat_backend/migrations/0026_oauthidentity.py @@ -0,0 +1,75 @@ +# 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", "0025_outbound_email"), + ] + + 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" + ), + ], + }, + ), + ] diff --git a/llm_be/chat_backend/urls.py b/llm_be/chat_backend/urls.py index 7672d0c..1ad339e 100644 --- a/llm_be/chat_backend/urls.py +++ b/llm_be/chat_backend/urls.py @@ -15,6 +15,7 @@ from .views import ( ConversationDetailView, CompanyUsersView, SetUserPassword, + ResetUserPassword, ConversationPreferences, UserPromptAnalytics, UserConversationAnalytics, -- 2.54.0