diff --git a/core/backend/apps/accounts/migrations/0003_invitation.py b/core/backend/apps/accounts/migrations/0003_invitation.py new file mode 100644 index 0000000..935f877 --- /dev/null +++ b/core/backend/apps/accounts/migrations/0003_invitation.py @@ -0,0 +1,38 @@ +# Generated by Django 5.1.15 on 2026-06-19 04:29 + +import apps.accounts.models +import django.db.models.deletion +import uuid +from django.conf import settings +from django.db import migrations, models + + +class Migration(migrations.Migration): + + dependencies = [ + ('accounts', '0002_loginsession_userpreference'), + ] + + operations = [ + migrations.CreateModel( + name='Invitation', + fields=[ + ('id', models.UUIDField(default=uuid.uuid4, editable=False, primary_key=True, serialize=False)), + ('created_at', models.DateTimeField(auto_now_add=True)), + ('updated_at', models.DateTimeField(auto_now=True)), + ('code', models.CharField(default=apps.accounts.models._gen_invite_code, max_length=32, unique=True)), + ('role', models.CharField(choices=[('owner', 'Owner'), ('admin', 'Admin'), ('member', 'Member'), ('viewer', 'Viewer')], default='member', max_length=24)), + ('email', models.EmailField(blank=True, max_length=254)), + ('monthly_credit_limit', models.DecimalField(decimal_places=2, default=0, max_digits=12)), + ('expires_at', models.DateTimeField(default=apps.accounts.models._default_invite_expiry)), + ('used_at', models.DateTimeField(blank=True, null=True)), + ('status', models.CharField(choices=[('pending', 'Pending'), ('used', 'Used'), ('revoked', 'Revoked'), ('expired', 'Expired')], default='pending', max_length=16)), + ('created_by', models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='created_invitations', to=settings.AUTH_USER_MODEL)), + ('team', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='invitations', to='accounts.team')), + ('used_by', models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='used_invitations', to=settings.AUTH_USER_MODEL)), + ], + options={ + 'indexes': [models.Index(fields=['team', 'status'], name='accounts_in_team_id_b48ad1_idx')], + }, + ), + ] diff --git a/core/backend/apps/accounts/models.py b/core/backend/apps/accounts/models.py index 3c8fc35..388a32a 100644 --- a/core/backend/apps/accounts/models.py +++ b/core/backend/apps/accounts/models.py @@ -95,3 +95,53 @@ class TeamMember(TimeStampedModel): def __str__(self) -> str: return f"{self.team} / {self.user} / {self.role}" + + +def _gen_invite_code() -> str: + import secrets + import string + + alphabet = "".join(c for c in (string.ascii_uppercase + string.digits) if c not in "O0I1") + block = lambda: "".join(secrets.choice(alphabet) for _ in range(4)) # noqa: E731 + return f"TEAM-{block()}-{block()}" + + +def _default_invite_expiry(): + from datetime import timedelta + + from django.utils import timezone + + return timezone.now() + timedelta(days=7) + + +class Invitation(TimeStampedModel): + """团队邀请码:超管/团管生成,新用户凭码注册即加入该团队(幂等加成员、不重复发额度)。""" + + class Status(models.TextChoices): + PENDING = "pending", "Pending" + USED = "used", "Used" + REVOKED = "revoked", "Revoked" + EXPIRED = "expired", "Expired" + + team = models.ForeignKey(Team, on_delete=models.CASCADE, related_name="invitations") + code = models.CharField(max_length=32, unique=True, default=_gen_invite_code) + role = models.CharField(max_length=24, choices=TeamMember.Role.choices, default=TeamMember.Role.MEMBER) + email = models.EmailField(blank=True) + monthly_credit_limit = models.DecimalField(max_digits=12, decimal_places=2, default=0) + created_by = models.ForeignKey(User, on_delete=models.SET_NULL, null=True, blank=True, related_name="created_invitations") + expires_at = models.DateTimeField(default=_default_invite_expiry) + used_by = models.ForeignKey(User, on_delete=models.SET_NULL, null=True, blank=True, related_name="used_invitations") + used_at = models.DateTimeField(null=True, blank=True) + status = models.CharField(max_length=16, choices=Status.choices, default=Status.PENDING) + + class Meta: + indexes = [models.Index(fields=["team", "status"])] + + @property + def is_valid(self) -> bool: + from django.utils import timezone + + return self.status == self.Status.PENDING and self.expires_at > timezone.now() + + def __str__(self) -> str: + return f"{self.team} / {self.code} / {self.status}" diff --git a/core/backend/apps/accounts/serializers.py b/core/backend/apps/accounts/serializers.py index 4a23ac4..b391084 100644 --- a/core/backend/apps/accounts/serializers.py +++ b/core/backend/apps/accounts/serializers.py @@ -2,7 +2,7 @@ from rest_framework import serializers from apps.billing.models import CreditAccount, CreditLedger -from .models import LoginSession, Team, TeamMember, User, UserPreference +from .models import Invitation, LoginSession, Team, TeamMember, User, UserPreference class UserSerializer(serializers.ModelSerializer): @@ -59,45 +59,109 @@ class RegisterSerializer(serializers.Serializer): password = serializers.CharField(min_length=8, write_only=True) email = serializers.EmailField(required=False, allow_blank=True) team_name = serializers.CharField(max_length=128, required=False, allow_blank=True) + # 邀请码(可选):有效码 → 加入已有团队;空 → 开新团队当超管。认证定调:不用邮箱注册,username 为主标识。 + invite_code = serializers.CharField(max_length=32, required=False, allow_blank=True) def validate_username(self, value): if User.objects.filter(username=value).exists(): raise serializers.ValidationError("username already exists") return value + def validate_invite_code(self, value): + code = (value or "").strip().upper() + if not code: + return "" + from django.utils import timezone + + from .models import Invitation + + inv = Invitation.objects.filter(code=code).first() + if inv is None: + raise serializers.ValidationError("邀请码不存在") + if inv.status != Invitation.Status.PENDING: + raise serializers.ValidationError("邀请码已被使用或已撤销") + if inv.expires_at <= timezone.now(): + raise serializers.ValidationError("邀请码已过期") + return code + def create(self, validated_data): from decimal import Decimal from django.conf import settings from django.db import transaction + from django.utils import timezone + from .models import Invitation + + invite_code = (validated_data.get("invite_code") or "").strip().upper() with transaction.atomic(): user = User.objects.create_user( username=validated_data["username"], password=validated_data["password"], email=validated_data.get("email", ""), ) - team = Team.objects.create( - name=validated_data.get("team_name") or f"{user.username}'s Team", - owner=user, - ) - TeamMember.objects.create(team=team, user=user, role=TeamMember.Role.OWNER) - trial = Decimal(str(settings.DEFAULT_TRIAL_CREDITS)) - CreditAccount.objects.create(team=team, balance=trial) - # 开户赠送额度必须进流水,否则这笔钱在账单里查无凭证、无法对账(余额从第一条起就可追溯) - if trial > 0: - CreditLedger.objects.create( + if invite_code: + # 凭码加入已有团队:幂等加成员(避开 unique_together 撞 IntegrityError)、标邀请 used、 + # 整段跳过 CreditAccount + ledger(团队级 OneToOne 已存在,二次建必报错、额度也不重发)。 + inv = Invitation.objects.select_for_update().filter(code=invite_code).first() + if inv is None or inv.status != Invitation.Status.PENDING or inv.expires_at <= timezone.now(): + raise serializers.ValidationError({"invite_code": "邀请码已失效"}) + team = inv.team + role = inv.role + if role == TeamMember.Role.OWNER: + role = TeamMember.Role.ADMIN # 不允许 OWNER(Team.owner 唯一) + TeamMember.objects.get_or_create( team=team, user=user, - ledger_type=CreditLedger.Type.RECHARGE, - amount=trial, - balance_after=trial, - reason="新用户试用额度赠送", - metadata={"kind": "trial_grant"}, + defaults={ + "role": role, + "status": TeamMember.Status.ACTIVE, + "monthly_credit_limit": inv.monthly_credit_limit, + }, ) + inv.status = Invitation.Status.USED + inv.used_by = user + inv.used_at = timezone.now() + inv.save(update_fields=["status", "used_by", "used_at", "updated_at"]) + else: + # 无码:开新团队当 owner + 发试用额度(原逻辑) + team = Team.objects.create( + name=validated_data.get("team_name") or f"{user.username}'s Team", + owner=user, + ) + TeamMember.objects.create(team=team, user=user, role=TeamMember.Role.OWNER) + trial = Decimal(str(settings.DEFAULT_TRIAL_CREDITS)) + CreditAccount.objects.create(team=team, balance=trial) + # 开户赠送额度必须进流水,否则这笔钱在账单里查无凭证、无法对账(余额从第一条起就可追溯) + if trial > 0: + CreditLedger.objects.create( + team=team, + user=user, + ledger_type=CreditLedger.Type.RECHARGE, + amount=trial, + balance_after=trial, + reason="新用户试用额度赠送", + metadata={"kind": "trial_grant"}, + ) return {"user": user, "team": team} +class InvitationSerializer(serializers.ModelSerializer): + used_by_username = serializers.CharField(source="used_by.username", read_only=True, default=None) + register_url = serializers.SerializerMethodField() + + class Meta: + model = Invitation + fields = [ + "id", "code", "role", "email", "monthly_credit_limit", "status", + "expires_at", "used_by", "used_by_username", "used_at", "created_at", "register_url", + ] + read_only_fields = fields + + def get_register_url(self, obj): + return f"/register?invite={obj.code}" + + class LoginSerializer(serializers.Serializer): username = serializers.CharField() password = serializers.CharField(write_only=True) diff --git a/core/backend/apps/accounts/tests.py b/core/backend/apps/accounts/tests.py index f0a4a7e..e84f12b 100644 --- a/core/backend/apps/accounts/tests.py +++ b/core/backend/apps/accounts/tests.py @@ -47,3 +47,72 @@ class AuthApiTests(TestCase): self.assertEqual(genesis.first().amount, trial) self.assertEqual(genesis.first().balance_after, trial) # 流水终点 == 账户余额,可对账 + +class InvitationFlowTests(TestCase): + def _register(self, client, username, **extra): + return client.post( + "/api/auth/register/", + {"username": username, "password": "strong-password", **extra}, + format="json", + ) + + def test_invite_join_same_team_no_duplicate_credits(self): + from apps.accounts.models import Invitation + + owner_client = APIClient() + r = self._register(owner_client, "inv-owner", team_name="Owner Team") + self.assertEqual(r.status_code, 201) + team = Team.objects.get(name="Owner Team") + accounts_before = CreditAccount.objects.filter(team=team).count() + ledgers_before = CreditLedger.objects.filter(team=team).count() + + # 超管生成邀请码 + owner_client.credentials(HTTP_AUTHORIZATION=f"Token {r.data['token']}") + gen = owner_client.post( + "/api/auth/team/invitations/", + {"role": "member", "monthly_credit_limit": "500"}, + format="json", + ) + self.assertEqual(gen.status_code, 201) + code = gen.data["code"] + self.assertTrue(code.startswith("TEAM-")) + + # 新用户凭码注册 + r2 = self._register(APIClient(), "inv-joiner", invite_code=code) + self.assertEqual(r2.status_code, 201) + joiner = User.objects.get(username="inv-joiner") + + # 进的是同一团队、角色对、active、额度对 + member = TeamMember.objects.get(team=team, user=joiner) + self.assertEqual(member.role, TeamMember.Role.MEMBER) + self.assertEqual(member.status, TeamMember.Status.ACTIVE) + self.assertEqual(member.monthly_credit_limit, Decimal("500")) + # 没开新团队 + self.assertFalse(Team.objects.filter(owner=joiner).exists()) + # 关键:不重复发额度 —— 团队 CreditAccount/ledger 数量不变 + self.assertEqual(CreditAccount.objects.filter(team=team).count(), accounts_before) + self.assertEqual(CreditLedger.objects.filter(team=team).count(), ledgers_before) + # 邀请标 used + inv = Invitation.objects.get(code=code) + self.assertEqual(inv.status, Invitation.Status.USED) + self.assertEqual(inv.used_by, joiner) + + def test_register_with_unknown_invite_rejected(self): + r = self._register(APIClient(), "bad-inv", invite_code="TEAM-ABCD-EFGH") + self.assertEqual(r.status_code, 400) + self.assertFalse(User.objects.filter(username="bad-inv").exists()) + + def test_invite_generate_requires_manage_permission(self): + owner_client = APIClient() + ro = self._register(owner_client, "perm-owner", team_name="Perm Team") + owner_client.credentials(HTTP_AUTHORIZATION=f"Token {ro.data['token']}") + gen = owner_client.post("/api/auth/team/invitations/", {"role": "member"}, format="json") + code = gen.data["code"] + + joiner_client = APIClient() + rj = self._register(joiner_client, "perm-member", invite_code=code) + # 普通成员尝试生成邀请码 → 403 + joiner_client.credentials(HTTP_AUTHORIZATION=f"Token {rj.data['token']}") + gen2 = joiner_client.post("/api/auth/team/invitations/", {"role": "member"}, format="json") + self.assertEqual(gen2.status_code, 403) + diff --git a/core/backend/apps/accounts/urls.py b/core/backend/apps/accounts/urls.py index 4b2fef0..a9acead 100644 --- a/core/backend/apps/accounts/urls.py +++ b/core/backend/apps/accounts/urls.py @@ -8,8 +8,10 @@ from .views import ( me, preferences, register, + revoke_invitation, revoke_login_session, revoke_other_sessions, + team_invitations, team_member_detail, team_member_password, team_members, @@ -31,4 +33,6 @@ urlpatterns = [ path("team/members/", team_members, name="team-members"), path("team/members//", team_member_detail, name="team-member-detail"), path("team/members//password/", team_member_password, name="team-member-password"), + path("team/invitations/", team_invitations, name="team-invitations"), + path("team/invitations//revoke/", revoke_invitation, name="team-invitation-revoke"), ] diff --git a/core/backend/apps/accounts/views.py b/core/backend/apps/accounts/views.py index a2edd29..512c49a 100644 --- a/core/backend/apps/accounts/views.py +++ b/core/backend/apps/accounts/views.py @@ -12,8 +12,9 @@ from rest_framework.response import Response from apps.common.api import get_current_team -from .models import LoginSession, TeamMember, User, UserPreference +from .models import Invitation, LoginSession, TeamMember, User, UserPreference from .serializers import ( + InvitationSerializer, LoginSerializer, LoginSessionSerializer, RegisterSerializer, @@ -238,6 +239,45 @@ def team_members(request): return Response(TeamMemberSerializer(member).data, status=status.HTTP_201_CREATED) +@api_view(["GET", "POST"]) +@permission_classes([IsAuthenticated]) +def team_invitations(request): + """团队邀请码:GET 列本团队邀请;POST 生成新码。权限=超管/团管(复用 can_manage_team)。""" + team = get_current_team(request.user) + if not can_manage_team(request.user, team): + return Response({"detail": "permission denied"}, status=status.HTTP_403_FORBIDDEN) + if request.method == "GET": + invites = team.invitations.select_related("used_by").order_by("-created_at") + return Response(InvitationSerializer(invites, many=True).data) + role = normalize_member_role(request.data.get("role")) + if role == TeamMember.Role.OWNER: + role = TeamMember.Role.ADMIN # 邀请不允许 OWNER(Team.owner 唯一) + invite = Invitation.objects.create( + team=team, + role=role, + email=str(request.data.get("email") or "").strip(), + monthly_credit_limit=request.data.get("monthly_credit_limit") or request.data.get("monthly") or 0, + created_by=request.user, + ) + return Response(InvitationSerializer(invite).data, status=status.HTTP_201_CREATED) + + +@api_view(["POST"]) +@permission_classes([IsAuthenticated]) +def revoke_invitation(request, invite_id): + """撤销一个待用邀请码。权限=超管/团管。""" + team = get_current_team(request.user) + if not can_manage_team(request.user, team): + return Response({"detail": "permission denied"}, status=status.HTTP_403_FORBIDDEN) + invite = Invitation.objects.filter(team=team, id=invite_id).first() + if invite is None: + return Response({"detail": "not found"}, status=status.HTTP_404_NOT_FOUND) + if invite.status == Invitation.Status.PENDING: + invite.status = Invitation.Status.REVOKED + invite.save(update_fields=["status", "updated_at"]) + return Response(InvitationSerializer(invite).data) + + @api_view(["PATCH", "DELETE"]) @permission_classes([IsAuthenticated]) def team_member_detail(request, member_id):