Files
cmhub/apps/licensing/tests.py
T

628 lines
24 KiB
Python

from datetime import timedelta
from decimal import Decimal
from concurrent.futures import ThreadPoolExecutor
from django.db import close_old_connections
from django.test import TestCase, TransactionTestCase, override_settings
from django.urls import reverse
from django.utils import timezone
from rest_framework.test import APIClient
from apps.licensing.models import (
ClientDevice,
DeviceBindingAudit,
DeviceCredential,
DeviceSession,
LegacyMigrationGrant,
LicenseEvent,
LicenseSeat,
MigrationRequest,
SoftwareEntitlement,
SoftwarePlan,
)
from apps.licensing.services import (
LicensingError,
assign_license_seat,
confirm_migration_request,
create_legacy_migration_grant,
create_migration_request,
grant_software_entitlement,
record_device_heartbeat,
release_license_seat,
renew_software_entitlement,
revoke_software_entitlement,
register_device,
)
from apps.users.models import ApiKey, User, UserWallet
class DeviceRegistrationApiTests(TestCase):
def setUp(self):
self.user = User.objects.create_user(
username="device-user",
email="device@example.com",
password="test-password",
)
UserWallet.objects.create(user=self.user, points_balance=20)
self.api_key, self.raw_api_key = ApiKey.create_for_user(self.user, name="desktop")
self.client = APIClient()
self.register_url = reverse("api-client-device-register")
self.heartbeat_url = reverse("api-client-device-heartbeat")
self.payload = {
"product_code": ClientDevice.ProductCode.CMSHOPEE,
"device_id": "v1:installed-device-abcdef",
"device_id_version": "v1",
"installation_public_key": "test-installation-public-key",
"platform": ClientDevice.Platform.WINDOWS,
"client_version": "0.1.0",
}
def register_device(self):
self.client.credentials(HTTP_AUTHORIZATION=f"Bearer {self.raw_api_key}")
return self.client.post(self.register_url, self.payload, format="json")
def test_register_requires_api_key(self):
response = self.client.post(self.register_url, self.payload, format="json")
self.assertEqual(response.status_code, 401)
self.assertEqual(response.data["error"]["code"], "unauthorized")
self.assertFalse(ClientDevice.objects.exists())
def test_register_creates_hashed_device_and_one_time_session_token(self):
response = self.register_device()
self.assertEqual(response.status_code, 201)
self.assertEqual(response.data["device"]["product_code"], "cmshopee")
self.assertIn("device_session_token", response.data)
self.assertNotIn("installation_public_key", response.data)
self.assertNotIn("device_id", response.data)
device = ClientDevice.objects.get(user=self.user)
session = DeviceSession.objects.get(device=device)
raw_token = response.data["device_session_token"]
self.assertNotEqual(device.device_fingerprint, self.payload["device_id"])
self.assertNotEqual(device.public_key_fingerprint, self.payload["installation_public_key"])
self.assertNotEqual(session.token_hash, raw_token)
self.assertTrue(session.matches_token(raw_token))
self.assertEqual(
DeviceBindingAudit.objects.filter(
device=device,
action=DeviceBindingAudit.Action.REGISTERED,
).count(),
1,
)
def test_repeat_registration_reuses_device_and_rotates_session(self):
first = self.register_device()
second = self.register_device()
self.assertEqual(first.status_code, 201)
self.assertEqual(second.status_code, 200)
self.assertEqual(ClientDevice.objects.filter(user=self.user).count(), 1)
device = ClientDevice.objects.get(user=self.user)
self.assertEqual(
DeviceSession.objects.filter(device=device, revoked_at__isnull=True).count(),
1,
)
self.assertEqual(DeviceSession.objects.filter(device=device).count(), 2)
self.assertNotEqual(
first.data["device_session_token"],
second.data["device_session_token"],
)
@override_settings(DEVICE_SESSION_TTL_SECONDS=172800)
def test_heartbeat_uses_device_session_and_activity_is_throttled(self):
registration = self.register_device()
token = registration.data["device_session_token"]
self.client.credentials(HTTP_X_DEVICE_SESSION=token)
response = self.client.post(self.heartbeat_url, {}, format="json")
self.assertEqual(response.status_code, 200)
self.assertFalse(response.data["activity_updated"])
device = ClientDevice.objects.get(user=self.user)
session = DeviceSession.objects.get(device=device, revoked_at__isnull=True)
updated = record_device_heartbeat(
session,
now=device.last_seen_at + timedelta(days=1, seconds=1),
)
self.assertTrue(updated)
device.refresh_from_db()
self.assertEqual(
DeviceBindingAudit.objects.filter(
device=device,
action=DeviceBindingAudit.Action.HEARTBEAT,
).count(),
1,
)
def test_invalid_or_revoked_device_session_is_rejected(self):
self.client.credentials(HTTP_X_DEVICE_SESSION="dvs_cmhub_invalid")
invalid = self.client.post(self.heartbeat_url, {}, format="json")
self.assertEqual(invalid.status_code, 401)
self.assertEqual(invalid.data["error"]["code"], "device_session_invalid")
registration = self.register_device()
device = ClientDevice.objects.get(user=self.user)
device.status = ClientDevice.Status.REVOKED
device.save(update_fields=("status",))
self.client.credentials(HTTP_X_DEVICE_SESSION=registration.data["device_session_token"])
revoked = self.client.post(self.heartbeat_url, {}, format="json")
self.assertEqual(revoked.status_code, 403)
self.assertEqual(revoked.data["error"]["code"], "device_revoked")
def test_existing_balance_api_remains_compatible_without_device_header(self):
self.client.credentials(HTTP_AUTHORIZATION=f"Bearer {self.raw_api_key}")
response = self.client.get(reverse("api-balance"))
self.assertEqual(response.status_code, 200)
self.assertEqual(response.data["points_balance"], 20)
@override_settings(DEVICE_ACTIVITY_UPDATE_SECONDS=60)
def test_heartbeat_updates_when_device_is_stale(self):
registration = self.register_device()
device = ClientDevice.objects.get(user=self.user)
stale_time = timezone.now() - timedelta(minutes=2)
ClientDevice.objects.filter(pk=device.pk).update(last_seen_at=stale_time)
self.client.credentials(HTTP_X_DEVICE_SESSION=registration.data["device_session_token"])
response = self.client.post(self.heartbeat_url, {}, format="json")
self.assertEqual(response.status_code, 200)
self.assertTrue(response.data["activity_updated"])
device.refresh_from_db()
self.assertGreater(device.last_seen_at, stale_time)
@override_settings(DEVICE_ACTIVITY_UPDATE_SECONDS=60)
def test_repeat_registration_updates_stale_device_activity(self):
self.register_device()
device = ClientDevice.objects.get(user=self.user)
stale_time = timezone.now() - timedelta(minutes=2)
ClientDevice.objects.filter(pk=device.pk).update(last_seen_at=stale_time)
response = self.register_device()
self.assertEqual(response.status_code, 200)
device.refresh_from_db()
self.assertGreater(device.last_seen_at, stale_time)
class SoftwareEntitlementServiceTests(TestCase):
def setUp(self):
self.user = User.objects.create_user(
username="entitlement-user",
email="entitlement@example.com",
password="test-password",
)
self.operator = User.objects.create_user(
username="entitlement-operator",
email="operator@example.com",
password="test-password",
is_staff=True,
)
self.plan = SoftwarePlan.objects.create(
product_code=ClientDevice.ProductCode.CMSHOPEE,
name="月度套餐",
duration_days=30,
price=Decimal("19.90"),
device_limit=1,
grace_days=3,
)
def create_device(self, suffix):
device_id = f"v1:entitlement-device-{suffix}"
public_key = f"entitlement-public-key-{suffix}"
return ClientDevice.objects.create(
user=self.user,
product_code=ClientDevice.ProductCode.CMSHOPEE,
device_id_version="v1",
device_fingerprint=ClientDevice.fingerprint_device_id("v1", device_id),
public_key_fingerprint=ClientDevice.fingerprint_public_key(public_key),
platform=ClientDevice.Platform.WINDOWS,
client_version="0.1.0",
)
def grant_entitlement(self, **kwargs):
return grant_software_entitlement(
user=self.user,
plan=self.plan,
reason="客服人工授予",
actor=self.operator,
**kwargs,
)
def test_grant_snapshots_plan_creates_fixed_seats_and_audit_event(self):
entitlement = self.grant_entitlement()
self.plan.name = "已调整套餐"
self.plan.price = Decimal("29.90")
self.plan.device_limit = 3
self.plan.save()
entitlement.refresh_from_db()
self.assertEqual(entitlement.plan_name, "月度套餐")
self.assertEqual(entitlement.plan_price, Decimal("19.90"))
self.assertEqual(entitlement.plan_device_limit, 1)
self.assertEqual(LicenseSeat.objects.filter(entitlement=entitlement).count(), 1)
event = LicenseEvent.objects.get(entitlement=entitlement)
self.assertEqual(event.action, LicenseEvent.Action.GRANTED)
self.assertEqual(event.reason, "客服人工授予")
self.assertEqual(event.actor, self.operator)
def test_renew_extends_from_later_of_now_or_existing_expiry(self):
now = timezone.now()
entitlement = self.grant_entitlement(starts_at=now - timedelta(days=10))
previous_expiry = entitlement.expires_at
renewed = renew_software_entitlement(
entitlement=entitlement,
reason="用户续订",
actor=self.operator,
now=now,
)
self.assertEqual(renewed.expires_at, previous_expiry + timedelta(days=30))
self.assertEqual(renewed.grace_expires_at, renewed.expires_at + timedelta(days=3))
self.assertEqual(renewed.status, SoftwareEntitlement.Status.ACTIVE)
self.assertTrue(
LicenseEvent.objects.filter(
entitlement=entitlement,
action=LicenseEvent.Action.RENEWED,
reason="用户续订",
).exists()
)
def test_assign_release_and_revoke_are_audited_and_never_exceed_seat_limit(self):
entitlement = self.grant_entitlement()
first_device = self.create_device("one")
second_device = self.create_device("two")
seat = assign_license_seat(
entitlement=entitlement,
device=first_device,
reason="首次绑定",
actor=self.operator,
)
repeated = assign_license_seat(
entitlement=entitlement,
device=first_device,
reason="重复绑定",
actor=self.operator,
)
self.assertEqual(seat.pk, repeated.pk)
with self.assertRaisesRegex(LicensingError, "设备席位已用完"):
assign_license_seat(
entitlement=entitlement,
device=second_device,
reason="超额绑定",
actor=self.operator,
)
released = release_license_seat(
seat=seat,
reason="客服解绑",
actor=self.operator,
)
assigned_again = assign_license_seat(
entitlement=entitlement,
device=second_device,
reason="重新绑定",
actor=self.operator,
)
self.assertIsNone(released.device_id)
self.assertEqual(assigned_again.device_id, second_device.id)
self.assertEqual(
LicenseEvent.objects.filter(
entitlement=entitlement,
action=LicenseEvent.Action.SEAT_ASSIGNED,
).count(),
2,
)
self.assertTrue(
LicenseEvent.objects.filter(
entitlement=entitlement,
action=LicenseEvent.Action.SEAT_RELEASED,
).exists()
)
revoked = revoke_software_entitlement(
entitlement=entitlement,
reason="退款撤销",
actor=self.operator,
)
self.assertEqual(revoked.status, SoftwareEntitlement.Status.REVOKED)
with self.assertRaisesRegex(LicensingError, "软件权益当前不可用"):
assign_license_seat(
entitlement=entitlement,
device=first_device,
reason="撤销后绑定",
actor=self.operator,
)
def test_manual_operations_require_reason(self):
with self.assertRaisesRegex(LicensingError, "必须填写操作原因"):
grant_software_entitlement(user=self.user, plan=self.plan, reason="")
class SoftwareEntitlementAdminTests(TestCase):
def setUp(self):
self.operator = User.objects.create_user(
username="licensing-admin",
email="licensing-admin@example.com",
password="test-password",
is_staff=True,
)
self.user = User.objects.create_user(
username="licensing-target",
email="licensing-target@example.com",
password="test-password",
)
self.plan = SoftwarePlan.objects.create(
product_code=ClientDevice.ProductCode.CMSHOPEE,
name="后台套餐",
duration_days=30,
price=Decimal("9.90"),
device_limit=2,
)
self.grant_url = reverse("admin:licensing_softwareentitlement_grant")
def test_admin_grant_requires_staff_and_reason_then_writes_event(self):
anonymous = self.client.get(self.grant_url)
self.assertEqual(anonymous.status_code, 302)
self.client.force_login(self.operator)
missing_reason = self.client.post(
self.grant_url,
{"user": self.user.pk, "plan": self.plan.pk, "reason": ""},
)
self.assertEqual(missing_reason.status_code, 200)
self.assertFalse(SoftwareEntitlement.objects.exists())
response = self.client.post(
self.grant_url,
{"user": self.user.pk, "plan": self.plan.pk, "reason": "后台补偿"},
)
self.assertEqual(response.status_code, 302)
entitlement = SoftwareEntitlement.objects.get(user=self.user)
self.assertTrue(
LicenseEvent.objects.filter(
entitlement=entitlement,
action=LicenseEvent.Action.GRANTED,
reason="后台补偿",
actor=self.operator,
).exists()
)
class LicenseSeatConcurrencyTests(TransactionTestCase):
def setUp(self):
self.user = User.objects.create_user(
username="seat-concurrency-user",
email="seat-concurrency@example.com",
password="test-password",
)
self.plan = SoftwarePlan.objects.create(
product_code=ClientDevice.ProductCode.CMSHOPEE,
name="单席位套餐",
duration_days=30,
price=Decimal("9.90"),
device_limit=1,
)
self.entitlement = grant_software_entitlement(
user=self.user,
plan=self.plan,
reason="并发测试授予",
)
self.first_device = self.create_device("first")
self.second_device = self.create_device("second")
def create_device(self, suffix):
return ClientDevice.objects.create(
user=self.user,
product_code=ClientDevice.ProductCode.CMSHOPEE,
device_id_version="v1",
device_fingerprint=ClientDevice.fingerprint_device_id("v1", f"concurrent-{suffix}"),
public_key_fingerprint=ClientDevice.fingerprint_public_key(f"key-{suffix}"),
platform=ClientDevice.Platform.WINDOWS,
client_version="0.1.0",
)
def test_concurrent_assignments_do_not_exceed_fixed_seat_limit(self):
def assign(device_id):
close_old_connections()
try:
entitlement = SoftwareEntitlement.objects.get(pk=self.entitlement.pk)
device = ClientDevice.objects.get(pk=device_id)
seat = assign_license_seat(
entitlement=entitlement,
device=device,
reason="并发绑定",
)
return ("assigned", seat.device_id)
except LicensingError as exc:
return (exc.code, None)
finally:
close_old_connections()
with ThreadPoolExecutor(max_workers=2) as executor:
outcomes = list(
executor.map(assign, (self.first_device.pk, self.second_device.pk))
)
self.assertEqual(sum(outcome[0] == "assigned" for outcome in outcomes), 1)
self.assertEqual(sum(outcome[0] == "seat_limit_reached" for outcome in outcomes), 1)
seat = LicenseSeat.objects.get(entitlement=self.entitlement)
self.assertIn(seat.device_id, {self.first_device.pk, self.second_device.pk})
class LegacyMigrationFlowTests(TestCase):
def setUp(self):
self.user = User.objects.create_user(
username="migration-user",
email="migration@example.com",
password="test-password",
)
self.other_user = User.objects.create_user(
username="migration-other",
email="migration-other@example.com",
password="test-password",
)
self.plan = SoftwarePlan.objects.create(
product_code=ClientDevice.ProductCode.CMSHOPEE,
name="存量迁移套餐",
duration_days=30,
price=Decimal("1.00"),
device_limit=1,
)
self.api_key, self.raw_api_key = ApiKey.create_for_user(self.user, name="legacy")
self.device, self.device_session_token = self.register_device_session()
self.client = APIClient()
def register_device_session(self):
result = register_device(
user=self.user,
api_key=self.api_key,
product_code=ClientDevice.ProductCode.CMSHOPEE,
device_id_version="v1",
device_id="migration-device-id",
public_key="migration-device-public-key",
platform=ClientDevice.Platform.WINDOWS,
client_version="0.2.0",
)
return result.device, result.session_token
def grant_migration(self):
return create_legacy_migration_grant(
user=self.user,
plan=self.plan,
reason="历史付费用户迁移",
eligibility_snapshot={"legacy_customer_id": "legacy-001"},
)
def api_headers(self):
return {
"HTTP_AUTHORIZATION": f"Bearer {self.raw_api_key}",
"HTTP_X_DEVICE_SESSION": self.device_session_token,
}
def test_request_confirm_poll_and_revoke_flow_keeps_credential_hashed(self):
self.grant_migration()
create_response = self.client.post(
reverse("api-client-migration-request-create"),
{},
format="json",
**self.api_headers(),
)
self.assertEqual(create_response.status_code, 201)
raw_credential = create_response.data["device_credential_token"]
self.assertTrue(create_response.data["confirmation_url"].endswith(create_response.data["request_id"]))
migration_request = MigrationRequest.objects.get(
request_id=create_response.data["request_id"]
)
self.assertNotEqual(migration_request.credential_token_hash, raw_credential)
pending = self.client.get(
reverse(
"api-client-migration-request-detail",
args=(migration_request.request_id,),
),
**self.api_headers(),
)
self.assertEqual(pending.status_code, 200)
self.assertEqual(pending.data["status"], MigrationRequest.Status.PENDING)
self.assertNotIn("device_credential_token", pending.data)
self.client.force_login(self.user)
confirm_url = reverse("portal-migration-confirm", args=(migration_request.request_id,))
self.assertEqual(self.client.get(confirm_url).status_code, 200)
self.assertEqual(self.client.post(confirm_url).status_code, 302)
self.assertEqual(self.client.post(confirm_url).status_code, 302)
credential = DeviceCredential.objects.get(migration_request=migration_request)
self.assertNotEqual(credential.token_hash, raw_credential)
self.assertEqual(credential.token_hash, MigrationRequest.hash_credential_token(raw_credential))
self.assertEqual(credential.seat.device_id, self.device.id)
self.assertEqual(DeviceCredential.objects.count(), 1)
self.assertEqual(
LicenseEvent.objects.filter(
action=LicenseEvent.Action.CREDENTIAL_ISSUED,
).count(),
1,
)
confirmed = self.client.get(
reverse(
"api-client-migration-request-detail",
args=(migration_request.request_id,),
),
**self.api_headers(),
)
self.assertEqual(confirmed.status_code, 200)
self.assertEqual(confirmed.data["status"], MigrationRequest.Status.CONFIRMED)
revoke = self.client.post(
reverse("portal-device-credential-revoke", args=(credential.pk,))
)
self.assertEqual(revoke.status_code, 302)
credential.refresh_from_db()
credential.seat.refresh_from_db()
self.assertIsNotNone(credential.revoked_at)
self.assertIsNone(credential.seat.device_id)
self.assertTrue(
LicenseEvent.objects.filter(
action=LicenseEvent.Action.CREDENTIAL_REVOKED,
reason="用户自助解绑设备",
).exists()
)
def test_request_requires_eligible_current_device_and_web_confirmation_same_user(self):
no_grant = self.client.post(
reverse("api-client-migration-request-create"),
{},
format="json",
**self.api_headers(),
)
self.assertEqual(no_grant.status_code, 403)
self.assertEqual(no_grant.data["error"]["code"], "migration_not_eligible")
self.grant_migration()
missing_device_session = self.client.post(
reverse("api-client-migration-request-create"),
{},
format="json",
HTTP_AUTHORIZATION=f"Bearer {self.raw_api_key}",
)
self.assertEqual(missing_device_session.status_code, 401)
self.assertEqual(missing_device_session.data["error"]["code"], "device_session_required")
created = self.client.post(
reverse("api-client-migration-request-create"),
{},
format="json",
**self.api_headers(),
)
migration_request = MigrationRequest.objects.get(request_id=created.data["request_id"])
self.client.force_login(self.other_user)
confirmation = self.client.get(
reverse("portal-migration-confirm", args=(migration_request.request_id,))
)
self.assertEqual(confirmation.status_code, 403)
self.assertFalse(DeviceCredential.objects.exists())
def test_expired_request_and_duplicate_migration_grant_are_rejected(self):
self.grant_migration()
with self.assertRaisesRegex(LicensingError, "已有此产品的迁移资格"):
self.grant_migration()
migration_request, _raw_token = create_migration_request(
user=self.user,
device=self.device,
now=timezone.now() - timedelta(minutes=20),
)
with self.assertRaisesRegex(LicensingError, "迁移请求已过期"):
confirm_migration_request(
request_id=migration_request.request_id,
user=self.user,
now=timezone.now(),
)
self.assertFalse(DeviceCredential.objects.exists())