diff --git a/.env.example b/.env.example index 3dc6010..99d587a 100644 --- a/.env.example +++ b/.env.example @@ -25,6 +25,9 @@ EMAIL_USE_TLS=true # Captcha (optional local) CAPTCHA_SECRET_KEY= +# Self-serve sign-up (default false — set true to allow /user/create/) +ENABLE_ACCOUNT_REGISTRATION=false + # 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 36f7440..de9d58d 100644 --- a/.env.prod.example +++ b/.env.prod.example @@ -45,6 +45,10 @@ EMAIL_USE_TLS=true # Captcha CAPTCHA_SECRET_KEY=replace-with-captcha-secret +# Self-serve account registration (sign-up). Keep false until ready to open +# public registration; set true in chat_backend_prod.env / chat_backend_beta.env. +ENABLE_ACCOUNT_REGISTRATION=false + # Stripe / finance STRIPE_SECRET_KEY=replace-with-stripe-secret-key STRIPE_PUBLISHABLE_KEY=replace-with-stripe-publishable-key diff --git a/README.md b/README.md index d6a9b10..6747fff 100644 --- a/README.md +++ b/README.md @@ -90,6 +90,7 @@ with `COMPOSE_DATABASE_URL` if needed. | `OLLAMA_MODEL` / `OLLAMA_EMBED_MODEL` | from `DEBUG` | optional | Override model names | | `EMAIL_HOST_*` | empty | yes (prod/beta) | SMTP2GO | | `CAPTCHA_SECRET_KEY` | empty | recommended | | +| `ENABLE_ACCOUNT_REGISTRATION` | `false` | optional | Self-serve sign-up; keep false until ready | | `STRIPE_SECRET_KEY` / `STRIPE_PUBLISHABLE_KEY` / `STRIPE_WEBHOOK_SECRET` | empty | yes for billing | Stripe API + webhook | | `STRIPE_PRICE_ID` | empty | optional | Pre-created Price; else `$10/mo` from settings | | `FRONTEND_BASE_URL` | `http://localhost:3000` | set in prod | Checkout success/cancel base | diff --git a/llm_be/chat_backend/serializers.py b/llm_be/chat_backend/serializers.py index 882b190..9a6947a 100644 --- a/llm_be/chat_backend/serializers.py +++ b/llm_be/chat_backend/serializers.py @@ -54,13 +54,55 @@ class CustomUserSerializer(serializers.ModelSerializer): fields = "__all__" extra_kwargs = {"password": {"write_only": True}} - # def create(self, validated_data): - # password = validated_data.pop('password',None) - # instance = self.Meta.model(**validated_data) - # if password is not None: - # instance.set_password(password) - # instance.save() - # return instance + +class SelfServeRegistrationSerializer(serializers.Serializer): + """Minimal payload for public self-serve sign-up (gated by settings).""" + + email = serializers.EmailField(required=True) + password = serializers.CharField(min_length=8, write_only=True) + first_name = serializers.CharField( + required=False, allow_blank=True, max_length=150, default="" + ) + last_name = serializers.CharField( + required=False, allow_blank=True, max_length=150, default="" + ) + company_name = serializers.CharField( + required=False, allow_blank=True, max_length=256, default="" + ) + + def validate_email(self, value: str) -> str: + email = value.strip().lower() + if CustomUser.objects.filter(email__iexact=email).exists(): + raise serializers.ValidationError("A user with this email already exists.") + if CustomUser.objects.filter(username__iexact=email).exists(): + raise serializers.ValidationError("A user with this email already exists.") + return email + + def create(self, validated_data): + email = validated_data["email"] + password = validated_data["password"] + first_name = (validated_data.get("first_name") or "").strip() + last_name = (validated_data.get("last_name") or "").strip() + company_name = (validated_data.get("company_name") or "").strip() + if not company_name: + company_name = f"{email}'s workspace" + + company = Company.objects.create( + name=company_name, + state="NA", + zipcode="00000", + address="N/A", + ) + user = CustomUser.objects.create_user( + username=email, + email=email, + password=password, + first_name=first_name, + last_name=last_name, + company=company, + is_company_manager=True, + ) + return user class ConversationSerializer(serializers.ModelSerializer): diff --git a/llm_be/chat_backend/tests/test_views_users.py b/llm_be/chat_backend/tests/test_views_users.py index acdfa6a..5d1108e 100644 --- a/llm_be/chat_backend/tests/test_views_users.py +++ b/llm_be/chat_backend/tests/test_views_users.py @@ -95,14 +95,71 @@ class TokenTestCase(APITestCase): self.assertEqual(self.client.get(url).status_code, status.HTTP_200_OK) +class PublicSettingsTestCase(APITestCase): + def test_returns_registration_flag(self): + with self.settings(ENABLE_ACCOUNT_REGISTRATION=False): + response = self.client.get(reverse("public_settings")) + + self.assertEqual(response.status_code, status.HTTP_200_OK) + self.assertEqual(response.data["enable_account_registration"], False) + + class UserCreateTestCase(APITestCase): + def test_registration_disabled_by_default(self): + response = self.client.post( + reverse("create_user"), + {"email": "new@example.com", "password": "securepass1"}, + format="json", + ) + + self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN) + self.assertEqual(CustomUser.objects.count(), 0) + def test_invalid_payload_returns_errors(self): - response = self.client.post(reverse("create_user"), {}, format="json") + with self.settings(ENABLE_ACCOUNT_REGISTRATION=True): + response = self.client.post(reverse("create_user"), {}, format="json") self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) self.assertIn("email", response.data) self.assertEqual(CustomUser.objects.count(), 0) + def test_creates_user_company_and_returns_tokens(self): + with self.settings(ENABLE_ACCOUNT_REGISTRATION=True): + response = self.client.post( + reverse("create_user"), + { + "email": "New.User@Example.com", + "password": "securepass1", + "first_name": "New", + "last_name": "User", + "company_name": "New Co", + }, + format="json", + ) + + self.assertEqual(response.status_code, status.HTTP_201_CREATED) + self.assertEqual(response.data["email"], "new.user@example.com") + self.assertIn("access", response.data) + self.assertIn("refresh", response.data) + + user = CustomUser.objects.get(email="new.user@example.com") + self.assertTrue(user.check_password("securepass1")) + self.assertTrue(user.is_company_manager) + self.assertEqual(user.company.name, "New Co") + self.assertEqual(user.username, "new.user@example.com") + + def test_duplicate_email_rejected(self): + make_user(email="taken@example.com") + with self.settings(ENABLE_ACCOUNT_REGISTRATION=True): + response = self.client.post( + reverse("create_user"), + {"email": "taken@example.com", "password": "securepass1"}, + format="json", + ) + + self.assertEqual(response.status_code, status.HTTP_400_BAD_REQUEST) + self.assertIn("email", response.data) + class SetPasswordTestCase(APITestCase): def setUp(self): diff --git a/llm_be/chat_backend/urls.py b/llm_be/chat_backend/urls.py index 9f32270..510b75a 100644 --- a/llm_be/chat_backend/urls.py +++ b/llm_be/chat_backend/urls.py @@ -7,6 +7,7 @@ from .views import ( CustomUserInvite, LogoutAndBlacklistRefreshTokenForUserView, CustomUserGet, + PublicSettingsView, is_authenticated, AnnouncmentView, FeedbackView, @@ -32,6 +33,7 @@ urlpatterns = [ path("token/obtain/", CustomObtainTokenView.as_view(), name="token_create"), 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("user/invite/", CustomUserInvite.as_view(), name="invite_user"), path("user/reset_password/", reset_password, name="reset_password"), path( diff --git a/llm_be/chat_backend/views.py b/llm_be/chat_backend/views.py index 02874c3..d589bee 100644 --- a/llm_be/chat_backend/views.py +++ b/llm_be/chat_backend/views.py @@ -6,6 +6,7 @@ from rest_framework import permissions, status from .serializers import ( MyTokenObtainPairSerializer, CustomUserSerializer, + SelfServeRegistrationSerializer, BasicUserSerializer, AnnouncmentSerializer, CompanySerializer, @@ -92,18 +93,49 @@ class CustomObtainTokenView(TokenObtainPairView): serializer_class = MyTokenObtainPairSerializer +class PublicSettingsView(APIView): + """Non-secret feature flags / public config for the SPA.""" + + permission_classes = (permissions.AllowAny,) + authentication_classes = () + + def get(self, request): + return Response( + { + "enable_account_registration": settings.ENABLE_ACCOUNT_REGISTRATION, + } + ) + + class CustomUserCreate(APIView): + """Self-serve registration. Gated by ENABLE_ACCOUNT_REGISTRATION (default off).""" + permission_classes = (permissions.AllowAny,) authentication_classes = () def post(self, request, format="json"): - serializer = CustomUserSerializer(data=request.data) - if serializer.is_valid(): - user = serializer.save() - if user: - json = serializer.data - return Response(json, status=status.HTTP_201_CREATED) - return Response(serializer.errors, status=status.HTTP_400_BAD_REQUEST) + if not settings.ENABLE_ACCOUNT_REGISTRATION: + return Response( + {"detail": "Account registration is disabled."}, + status=status.HTTP_403_FORBIDDEN, + ) + + serializer = SelfServeRegistrationSerializer(data=request.data) + if not serializer.is_valid(): + return Response(serializer.errors, status=status.HTTP_400_BAD_REQUEST) + + user = serializer.save() + refresh = RefreshToken.for_user(user) + return Response( + { + "email": user.email, + "first_name": user.first_name, + "last_name": user.last_name, + "access": str(refresh.access_token), + "refresh": str(refresh), + }, + status=status.HTTP_201_CREATED, + ) def send_invite_email(slug, email_to_invite): diff --git a/llm_be/llm_be/settings.py b/llm_be/llm_be/settings.py index 12164bd..7bda6ac 100644 --- a/llm_be/llm_be/settings.py +++ b/llm_be/llm_be/settings.py @@ -297,6 +297,10 @@ os.makedirs(directory_path, exist_ok=True) ALLOW_IMAGE_GENERATION = env_bool("ALLOW_IMAGE_GENERATION", False) ALLOW_INTERNET_ACCESS = env_bool("ALLOW_INTERNET_ACCESS", True) +# Self-serve account registration (sign-up page). Default off — enable via +# control-node secret (chat_backend_.env) when ready for public sign-up. +ENABLE_ACCOUNT_REGISTRATION = env_bool("ENABLE_ACCOUNT_REGISTRATION", False) + # --------------------------------------------------------------------------- # Finance / Stripe (subscription billing) # ---------------------------------------------------------------------------