Compare commits
1
Commits
master
...
f8c29e09bb
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f8c29e09bb |
@@ -28,6 +28,18 @@ CAPTCHA_SECRET_KEY=
|
|||||||
# Self-serve sign-up (default false — set true to allow /user/create/)
|
# Self-serve sign-up (default false — set true to allow /user/create/)
|
||||||
ENABLE_ACCOUNT_REGISTRATION=false
|
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 / finance (optional local — required for checkout + webhooks)
|
||||||
STRIPE_SECRET_KEY=
|
STRIPE_SECRET_KEY=
|
||||||
STRIPE_PUBLISHABLE_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.
|
# public registration; set true in chat_backend_prod.env / chat_backend_beta.env.
|
||||||
ENABLE_ACCOUNT_REGISTRATION=false
|
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 / finance
|
||||||
STRIPE_SECRET_KEY=replace-with-stripe-secret-key
|
STRIPE_SECRET_KEY=replace-with-stripe-secret-key
|
||||||
STRIPE_PUBLISHABLE_KEY=replace-with-stripe-publishable-key
|
STRIPE_PUBLISHABLE_KEY=replace-with-stripe-publishable-key
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ from .models import (
|
|||||||
PromptMetric,
|
PromptMetric,
|
||||||
DocumentWorkspace,
|
DocumentWorkspace,
|
||||||
Document,
|
Document,
|
||||||
|
OAuthIdentity,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Register your models here.
|
# Register your models here.
|
||||||
@@ -145,3 +146,22 @@ admin.site.register(Feedback, FeedbackAdmin)
|
|||||||
|
|
||||||
admin.site.register(DocumentWorkspace, DocumentWorkspaceAdmin)
|
admin.site.register(DocumentWorkspace, DocumentWorkspaceAdmin)
|
||||||
admin.site.register(Document, DocumentAdmin)
|
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')],
|
||||||
|
},
|
||||||
|
),
|
||||||
|
]
|
||||||
@@ -77,6 +77,47 @@ class CustomUser(AbstractUser):
|
|||||||
return f"https://chat.aimloperations.com/set_password?slug={self.slug}"
|
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 = (
|
FEEDBACK_CHOICE = (
|
||||||
("SUBMITTED", "Submitted"),
|
("SUBMITTED", "Submitted"),
|
||||||
("RESOLVED", "Resolved"),
|
("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,
|
ConversationDetailView,
|
||||||
CompanyUsersView,
|
CompanyUsersView,
|
||||||
SetUserPassword,
|
SetUserPassword,
|
||||||
ResetUserPassword,
|
|
||||||
ConversationPreferences,
|
ConversationPreferences,
|
||||||
UserPromptAnalytics,
|
UserPromptAnalytics,
|
||||||
UserConversationAnalytics,
|
UserConversationAnalytics,
|
||||||
@@ -26,6 +25,7 @@ from .views import (
|
|||||||
DocumentUploadView,
|
DocumentUploadView,
|
||||||
DocumentDetailView,
|
DocumentDetailView,
|
||||||
)
|
)
|
||||||
|
from .views_oauth import OAuthCallbackView, OAuthStartView
|
||||||
from rest_framework.routers import DefaultRouter
|
from rest_framework.routers import DefaultRouter
|
||||||
|
|
||||||
|
|
||||||
@@ -34,6 +34,16 @@ urlpatterns = [
|
|||||||
path("token/refresh/", jwt_views.TokenRefreshView.as_view(), name="token_refresh"),
|
path("token/refresh/", jwt_views.TokenRefreshView.as_view(), name="token_refresh"),
|
||||||
path("user/create/", CustomUserCreate.as_view(), name="create_user"),
|
path("user/create/", CustomUserCreate.as_view(), name="create_user"),
|
||||||
path("public/settings/", PublicSettingsView.as_view(), name="public_settings"),
|
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/invite/", CustomUserInvite.as_view(), name="invite_user"),
|
||||||
path("user/reset_password/", reset_password, name="reset_password"),
|
path("user/reset_password/", reset_password, name="reset_password"),
|
||||||
path(
|
path(
|
||||||
|
|||||||
@@ -100,9 +100,12 @@ class PublicSettingsView(APIView):
|
|||||||
authentication_classes = ()
|
authentication_classes = ()
|
||||||
|
|
||||||
def get(self, request):
|
def get(self, request):
|
||||||
|
from .views_oauth import oauth_public_flags
|
||||||
|
|
||||||
return Response(
|
return Response(
|
||||||
{
|
{
|
||||||
"enable_account_registration": settings.ENABLE_ACCOUNT_REGISTRATION,
|
"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()
|
||||||
@@ -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.
|
# control-node secret (chat_backend_<env>.env) when ready for public sign-up.
|
||||||
ENABLE_ACCOUNT_REGISTRATION = env_bool("ENABLE_ACCOUNT_REGISTRATION", False)
|
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)
|
# Finance / Stripe (subscription billing)
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
Reference in New Issue
Block a user