from decimal import Decimal from unittest.mock import patch from rest_framework.test import APIClient from django.test import TestCase from apps.accounts.models import AdminAuditLog, Invitation, Team, TeamMember, User class AdminInvitationApiTests(TestCase): """Phase 2:平台超管邀请码发放管理(发开团队码 / 列跨团队码 / 撤销)+ 权限 + 审计。""" def setUp(self): self.admin = User.objects.create_user(username="padmin", password="x", is_platform_admin=True) self.normal = User.objects.create_user(username="normal", password="x") self.team = Team.objects.create(name="Normal Team", owner=self.normal) TeamMember.objects.create(team=self.team, user=self.normal, role=TeamMember.Role.OWNER) self.admin_client = APIClient() self.admin_client.force_authenticate(self.admin) self.normal_client = APIClient() self.normal_client.force_authenticate(self.normal) def test_list_requires_platform_admin(self): self.assertEqual(self.normal_client.get("/api/admin/invitations/").status_code, 403) self.assertEqual(self.admin_client.get("/api/admin/invitations/").status_code, 200) def test_create_create_team_code_and_audit(self): r = self.admin_client.post("/api/admin/invitations/", {"email": "x@y.z"}, format="json") self.assertEqual(r.status_code, 201) self.assertEqual(r.data["kind"], Invitation.Kind.CREATE_TEAM) self.assertTrue(r.data["code"].startswith("TEAM-")) self.assertIsNone(r.data["team"]) self.assertEqual(r.data["register_url"], f"/register?invite={r.data['code']}") self.assertTrue(AdminAuditLog.objects.filter(action="invite.issue_create_team").exists()) def test_create_requires_platform_admin(self): self.assertEqual(self.normal_client.post("/api/admin/invitations/", {}, format="json").status_code, 403) def test_list_is_cross_team(self): self.admin_client.post("/api/admin/invitations/", {}, format="json") Invitation.objects.create(kind=Invitation.Kind.JOIN_TEAM, team=self.team, role=TeamMember.Role.MEMBER) r = self.admin_client.get("/api/admin/invitations/") self.assertEqual(r.status_code, 200) self.assertGreaterEqual(r.data["count"], 2) kinds = {i["kind"] for i in r.data["results"]} self.assertIn(Invitation.Kind.CREATE_TEAM, kinds) self.assertIn(Invitation.Kind.JOIN_TEAM, kinds) def test_filter_by_kind(self): self.admin_client.post("/api/admin/invitations/", {}, format="json") Invitation.objects.create(kind=Invitation.Kind.JOIN_TEAM, team=self.team, role=TeamMember.Role.MEMBER) r = self.admin_client.get("/api/admin/invitations/?kind=create_team") self.assertTrue(all(i["kind"] == Invitation.Kind.CREATE_TEAM for i in r.data["results"])) def test_revoke_and_audit(self): created = self.admin_client.post("/api/admin/invitations/", {}, format="json") inv_id = created.data["id"] r = self.admin_client.post(f"/api/admin/invitations/{inv_id}/revoke/") self.assertEqual(r.status_code, 200) self.assertEqual(r.data["status"], Invitation.Status.REVOKED) self.assertTrue(AdminAuditLog.objects.filter(action="invite.revoke").exists()) def test_revoke_requires_platform_admin(self): created = self.admin_client.post("/api/admin/invitations/", {}, format="json") inv_id = created.data["id"] self.assertEqual(self.normal_client.post(f"/api/admin/invitations/{inv_id}/revoke/").status_code, 403) def test_issued_create_team_code_registers_new_team(self): created = self.admin_client.post("/api/admin/invitations/", {}, format="json") code = created.data["code"] reg = APIClient().post( "/api/auth/register/", {"username": "new-via-admin-code", "password": "strong-pass-1", "team_name": "Via Admin", "invite_code": code}, format="json", ) self.assertEqual(reg.status_code, 201) self.assertTrue(Team.objects.filter(name="Via Admin").exists()) # 码标 used self.assertEqual(Invitation.objects.get(code=code).status, Invitation.Status.USED) def test_revoked_code_cannot_register(self): created = self.admin_client.post("/api/admin/invitations/", {}, format="json") code = created.data["code"] self.admin_client.post(f"/api/admin/invitations/{created.data['id']}/revoke/") reg = APIClient().post( "/api/auth/register/", {"username": "blocked-user", "password": "strong-pass-1", "team_name": "Blocked", "invite_code": code}, format="json", ) self.assertEqual(reg.status_code, 400) class AdminTeamUserApiTests(TestCase): """Phase 3:平台团队 + 用户管理(列/详情/启停/强制改密)+ 权限 + 团队停用拦登录。""" def setUp(self): from apps.billing.models import CreditAccount self.admin = User.objects.create_user(username="padmin3", password="x", is_platform_admin=True) self.owner = User.objects.create_user(username="t-owner", password="ownerpass1") self.member = User.objects.create_user(username="t-member", password="memberpass1") self.team = Team.objects.create(name="Acme", owner=self.owner) TeamMember.objects.create(team=self.team, user=self.owner, role=TeamMember.Role.OWNER) TeamMember.objects.create(team=self.team, user=self.member, role=TeamMember.Role.MEMBER) CreditAccount.objects.create(team=self.team, balance="123.4500") self.ac = APIClient() self.ac.force_authenticate(self.admin) self.normal = APIClient() self.normal.force_authenticate(self.owner) def _login(self, username, password): return APIClient().post("/api/auth/login/", {"username": username, "password": password}, format="json") def test_teams_list_permission_and_fields(self): self.assertEqual(self.normal.get("/api/admin/teams/").status_code, 403) r = self.ac.get("/api/admin/teams/") self.assertEqual(r.status_code, 200) row = next(t for t in r.data["results"] if t["name"] == "Acme") self.assertEqual(row["member_count"], 2) self.assertEqual(row["owner_username"], "t-owner") self.assertEqual(row["balance"], "123.4500") def test_team_detail_has_members(self): r = self.ac.get(f"/api/admin/teams/{self.team.id}/") self.assertEqual(r.status_code, 200) self.assertEqual(len(r.data["members"]), 2) def test_team_toggle_blocks_then_restores_member_login(self): self.assertEqual(self._login("t-member", "memberpass1").status_code, 200) tr = self.ac.post(f"/api/admin/teams/{self.team.id}/toggle/") self.assertEqual(tr.status_code, 200) self.assertEqual(tr.data["status"], Team.Status.DISABLED) self.assertEqual(self._login("t-member", "memberpass1").status_code, 400) # 停用后被拒 self.ac.post(f"/api/admin/teams/{self.team.id}/toggle/") # 恢复 self.assertEqual(self._login("t-member", "memberpass1").status_code, 200) self.assertTrue(AdminAuditLog.objects.filter(action="team.toggle_status").exists()) def test_team_toggle_requires_admin(self): self.assertEqual(self.normal.post(f"/api/admin/teams/{self.team.id}/toggle/").status_code, 403) def test_users_list_with_teams(self): self.assertEqual(self.normal.get("/api/admin/users/").status_code, 403) r = self.ac.get("/api/admin/users/") self.assertEqual(r.status_code, 200) member_row = next(u for u in r.data["results"] if u["username"] == "t-member") self.assertEqual(member_row["teams"][0]["team_name"], "Acme") def test_user_toggle_blocks_login(self): r = self.ac.post(f"/api/admin/users/{self.member.id}/toggle/") self.assertEqual(r.status_code, 200) self.assertEqual(r.data["status"], User.Status.DISABLED) self.assertEqual(self._login("t-member", "memberpass1").status_code, 400) self.assertTrue(AdminAuditLog.objects.filter(action="user.toggle_status").exists()) def test_cannot_disable_platform_admin(self): self.assertEqual(self.ac.post(f"/api/admin/users/{self.admin.id}/toggle/").status_code, 400) def test_reset_password(self): r = self.ac.post(f"/api/admin/users/{self.member.id}/reset-password/", {"password": "newpass789"}, format="json") self.assertEqual(r.status_code, 204) self.assertEqual(self._login("t-member", "newpass789").status_code, 200) self.assertEqual(self._login("t-member", "memberpass1").status_code, 400) # 旧密码失效 self.assertTrue(AdminAuditLog.objects.filter(action="user.reset_password").exists()) def test_reset_password_too_short_rejected(self): self.assertEqual( self.ac.post(f"/api/admin/users/{self.member.id}/reset-password/", {"password": "short"}, format="json").status_code, 400, ) def test_user_actions_require_admin(self): self.assertEqual(self.normal.post(f"/api/admin/users/{self.member.id}/toggle/").status_code, 403) self.assertEqual( self.normal.post(f"/api/admin/users/{self.member.id}/reset-password/", {"password": "newpass789"}, format="json").status_code, 403, ) class AdminQualityWordTests(TestCase): """Phase 4:质量词 CRUD + 权限 + 生成侧读配置(无配置回落写死值,零回归)。""" def setUp(self): self.admin = User.objects.create_user(username="padmin4", password="x", is_platform_admin=True) self.normal = User.objects.create_user(username="normal4", password="x") self.ac = APIClient() self.ac.force_authenticate(self.admin) self.nc = APIClient() self.nc.force_authenticate(self.normal) def test_crud_and_permission_and_audit(self): self.assertEqual(self.nc.get("/api/admin/quality-words/").status_code, 403) self.assertEqual(self.nc.post("/api/admin/quality-words/", {"stage": "person", "text": "x"}, format="json").status_code, 403) r = self.ac.post("/api/admin/quality-words/", {"stage": "person", "text": "赛博朋克风"}, format="json") self.assertEqual(r.status_code, 201) wid = r.data["id"] self.assertTrue(AdminAuditLog.objects.filter(action="quality_word.create").exists()) lst = self.ac.get("/api/admin/quality-words/") self.assertEqual(lst.status_code, 200) self.assertTrue(any(w["text"] == "赛博朋克风" for w in lst.data)) up = self.ac.patch(f"/api/admin/quality-words/{wid}/", {"text": "国潮风"}, format="json") self.assertEqual(up.status_code, 200) self.assertEqual(up.data["text"], "国潮风") self.assertEqual(self.ac.delete(f"/api/admin/quality-words/{wid}/").status_code, 204) def test_invalid_stage_rejected(self): r = self.ac.post("/api/admin/quality-words/", {"stage": "bogus", "text": "x"}, format="json") self.assertEqual(r.status_code, 400) def test_person_prompt_reads_template_with_fallback(self): """人物立绘提示词改由「提示词模板」(person_portrait)驱动,取代旧质量词: 种子模板 = 原默认正文(描述填进 {描述});改模板 → 生成立即用新正文。""" from apps.ai.models import PromptTemplate from apps.ai.services import build_person_frontal_prompt # 种子模板 → 含默认尾串,{描述} 被填入 out = build_person_frontal_prompt("小姐姐") self.assertIn("纯色背景", out) self.assertIn("小姐姐", out) # 后台改模板 → 生成用新正文,{描述} 仍替换 PromptTemplate.objects.filter(key="person_portrait").update(template="赛博朋克风{描述}霓虹光") out2 = build_person_frontal_prompt("小姐姐") self.assertIn("赛博朋克风", out2) self.assertIn("霓虹光", out2) self.assertIn("小姐姐", out2) self.assertNotIn("纯色背景", out2) def test_disabled_word_not_used(self): from apps.ai.models import QualityWord from apps.ai.services import quality_words QualityWord.objects.create(stage="video", text="启用词", enabled=True) QualityWord.objects.create(stage="video", text="停用词", enabled=False) words = quality_words("video") self.assertIn("启用词", words) self.assertNotIn("停用词", words) def test_prompt_templates_list_update_permission(self): """提示词模板:列表(6 条,带 label/placeholders)+ 改正文比例 + 权限 + 审计 + 生成侧立即生效。""" from apps.ai.services import build_product_triview_prompt_refs self.assertEqual(self.nc.get("/api/admin/prompt-templates/").status_code, 403) lst = self.ac.get("/api/admin/prompt-templates/") self.assertEqual(lst.status_code, 200) self.assertEqual(len(lst.data), 6) # 视频线 6 条 row = next(r for r in lst.data if r["key"] == "product_triview") self.assertEqual(row["label"], "商品三视图") self.assertIn("商品", row["placeholders"]) up = self.ac.patch( f"/api/admin/prompt-templates/{row['id']}/", {"template": "参考图1是「{商品}」是商品。生成三视图。", "ratio": "portrait"}, format="json", ) self.assertEqual(up.status_code, 200) self.assertTrue(AdminAuditLog.objects.filter(action="prompt_template.update").exists()) class P: title = "杂鱼煲" out = build_product_triview_prompt_refs(P(), "") self.assertEqual(out, "参考图1是「杂鱼煲」是商品。生成三视图。") # 生成侧立即用新模板,且 {商品} 已替换 class AdminAssetReviewTests(TestCase): """Phase 5:跨团队人像审核队列(列/筛/批量送审/轮询)+ 权限。火山调用全 mock。""" def setUp(self): from apps.assets.models import Asset self.admin = User.objects.create_user(username="padmin5", password="x", is_platform_admin=True) self.normal = User.objects.create_user(username="normal5", password="x") self.teamA = Team.objects.create(name="TeamA", owner=self.normal) TeamMember.objects.create(team=self.teamA, user=self.normal, role=TeamMember.Role.OWNER) self.user2 = User.objects.create_user(username="u2-5", password="x") self.teamB = Team.objects.create(name="TeamB", owner=self.user2) self.a_proc = Asset.objects.create(team=self.teamA, name="proc", asset_type="image", category=Asset.Category.PERSON, review_status="processing") self.a_active = Asset.objects.create(team=self.teamA, name="active", asset_type="image", category=Asset.Category.PERSON, review_status="active") self.a_failed = Asset.objects.create(team=self.teamB, name="failed", asset_type="image", category=Asset.Category.PERSON, review_status="failed") self.a_none = Asset.objects.create(team=self.teamA, name="none", asset_type="image", category=Asset.Category.PERSON, review_status="") # scene 属于 REVIEW_CATEGORIES(场景图可能真人出镜,44e9023 扩审核范围),该进队列 Asset.objects.create(team=self.teamA, name="scene", asset_type="image", category=Asset.Category.SCENE, review_status="") # 图片趴上传素材不在 REVIEW_CATEGORIES,不该进队列 Asset.objects.create(team=self.teamA, name="upload", asset_type="image", category=Asset.Category.UPLOAD, review_status="") self.ac = APIClient() self.ac.force_authenticate(self.admin) self.nc = APIClient() self.nc.force_authenticate(self.normal) def test_list_permission_and_review_scope_cross_team(self): self.assertEqual(self.nc.get("/api/admin/asset-reviews/").status_code, 403) r = self.ac.get("/api/admin/asset-reviews/") self.assertEqual(r.status_code, 200) self.assertEqual(r.data["count"], 5) # 4 person + 1 scene(跨 2 团队),范围外 upload 排除 self.assertEqual({row["team_name"] for row in r.data["results"]}, {"TeamA", "TeamB"}) self.assertNotIn("upload", {row["name"] for row in r.data["results"]}) def test_filter_by_status(self): self.assertEqual(self.ac.get("/api/admin/asset-reviews/?review_status=processing").data["count"], 1) self.assertEqual(self.ac.get("/api/admin/asset-reviews/?review_status=active").data["count"], 1) self.assertEqual(self.ac.get("/api/admin/asset-reviews/?review_status=failed").data["count"], 1) self.assertEqual(self.ac.get("/api/admin/asset-reviews/?review_status=none").data["count"], 2) # person a_none + scene def test_submit_calls_review_and_audit(self): with patch("apps.adminpanel.views.submit_asset_for_review") as mock_submit: r = self.ac.post( "/api/admin/asset-reviews/submit/", {"asset_ids": [str(self.a_none.id), str(self.a_failed.id)]}, format="json", ) self.assertEqual(r.status_code, 200) self.assertEqual(r.data["submitted"], 2) self.assertEqual(mock_submit.call_count, 2) self.assertTrue(AdminAuditLog.objects.filter(action="asset_review.submit").exists()) def test_submit_requires_admin(self): with patch("apps.adminpanel.views.submit_asset_for_review") as mock_submit: r = self.nc.post("/api/admin/asset-reviews/submit/", {"asset_ids": [str(self.a_none.id)]}, format="json") self.assertEqual(r.status_code, 403) mock_submit.assert_not_called() def test_poll_only_processing(self): with patch("apps.adminpanel.views.poll_asset_review", return_value="active") as mock_poll: r = self.ac.post("/api/admin/asset-reviews/poll/", {}, format="json") self.assertEqual(r.status_code, 200) self.assertEqual(r.data["polled"], 1) # 仅 a_proc 处于 processing mock_poll.assert_called_once() class AdminTaskMonitorTests(TestCase): """Phase 6:AI 任务监控(列/筛/详情/成本异常)+ 失败重投(celery 全 mock)+ 权限。""" def setUp(self): from apps.ai.models import AITask, ModelConfig, ModelProvider self.AITask = AITask self.admin = User.objects.create_user(username="padmin6", password="x", is_platform_admin=True) self.normal = User.objects.create_user(username="normal6", password="x") self.team = Team.objects.create(name="T6", owner=self.normal) TeamMember.objects.create(team=self.team, user=self.normal, role=TeamMember.Role.OWNER) prov = ModelProvider.objects.create(name="prov6", display_name="P6") self.mc = ModelConfig.objects.create(provider=prov, name="m6", display_name="M6", capability=ModelConfig.Capability.IMAGE) def mk(tt, st, est, act, key): return AITask.objects.create( team=self.team, model_config=self.mc, task_type=tt, status=st, estimated_cost=est, actual_cost=act, idempotency_key=key, ) self.t_ok = mk(AITask.Type.PRODUCT_IMAGE, AITask.Status.SUCCEEDED, "1.0", "1.0", "k-ok") self.t_failed = mk(AITask.Type.PRODUCT_IMAGE, AITask.Status.FAILED, "1.0", "1.0", "k-failed") self.t_anom = mk(AITask.Type.PERSON_IMAGE, AITask.Status.SUCCEEDED, "1.0", "5.0", "k-anom") self.t_script = mk(AITask.Type.SCRIPT_GENERATION, AITask.Status.FAILED, "1.0", "1.0", "k-script") self.ac = APIClient() self.ac.force_authenticate(self.admin) self.nc = APIClient() self.nc.force_authenticate(self.normal) def test_list_permission_and_count(self): self.assertEqual(self.nc.get("/api/admin/tasks/").status_code, 403) r = self.ac.get("/api/admin/tasks/") self.assertEqual(r.status_code, 200) self.assertGreaterEqual(r.data["count"], 4) def test_list_defers_large_payload_columns_but_detail_keeps_them(self): from django.db import connection from django.test.utils import CaptureQueriesContext self.t_ok.request_payload = {"prompt": "x" * 100_000} self.t_ok.response_payload = {"image": "y" * 100_000} self.t_ok.error_message = "z" * 100_000 self.t_ok.save(update_fields=["request_payload", "response_payload", "error_message", "updated_at"]) with CaptureQueriesContext(connection) as captured: response = self.ac.get("/api/admin/tasks/?page_size=10") self.assertEqual(response.status_code, 200) task_selects = [ item["sql"] for item in captured.captured_queries if "FROM \"ai_aitask\"" in item["sql"] and "COUNT(" not in item["sql"] ] self.assertTrue(task_selects) for sql in task_selects: self.assertNotIn("request_payload", sql) self.assertNotIn("response_payload", sql) self.assertNotIn("error_message", sql) detail = self.ac.get(f"/api/admin/tasks/{self.t_ok.id}/") self.assertEqual(detail.status_code, 200) self.assertEqual(detail.data["request_payload"], self.t_ok.request_payload) self.assertEqual(detail.data["response_payload"], self.t_ok.response_payload) self.assertEqual(detail.data["error_message"], self.t_ok.error_message) def test_filter_status_type_anomaly(self): self.assertTrue(all(t["status"] == "failed" for t in self.ac.get("/api/admin/tasks/?status=failed").data["results"])) self.assertTrue(all(t["task_type"] == "person_image" for t in self.ac.get("/api/admin/tasks/?task_type=person_image").data["results"])) ids = {t["id"] for t in self.ac.get("/api/admin/tasks/?anomaly=1").data["results"]} self.assertIn(str(self.t_anom.id), ids) self.assertNotIn(str(self.t_ok.id), ids) def test_cost_anomaly_flag(self): self.assertTrue(self.ac.get(f"/api/admin/tasks/{self.t_anom.id}/").data["cost_anomaly"]) self.assertFalse(self.ac.get(f"/api/admin/tasks/{self.t_ok.id}/").data["cost_anomaly"]) def test_detail_has_payloads(self): from django.utils import timezone from apps.ai.models import AIModelAttempt AIModelAttempt.objects.create( task=self.t_failed, sequence=1, provider=self.mc.provider, model_config=self.mc, provider_name=self.mc.provider.name, provider_display_name=self.mc.provider.display_name, model_name=self.mc.name, model_display_name=self.mc.display_name, public_model_name="AirShelf Image", capability="image", operation="image_generate", status=AIModelAttempt.Status.FAILED, started_at=timezone.now(), error_type="provider_unavailable", ) d = self.ac.get(f"/api/admin/tasks/{self.t_failed.id}/") self.assertEqual(d.status_code, 200) self.assertIn("request_payload", d.data) self.assertIn("response_payload", d.data) self.assertEqual(len(d.data["attempts"]), 1) self.assertEqual(d.data["attempts"][0]["model_name"], self.mc.name) row = next(item for item in self.ac.get("/api/admin/tasks/").data["results"] if item["id"] == str(self.t_failed.id)) self.assertNotIn("attempts", row) def test_retry_dispatch_and_audit(self): with patch("apps.ai.tasks.generate_standalone_image_task.delay") as mock_delay: r = self.ac.post(f"/api/admin/tasks/{self.t_failed.id}/retry/") self.assertEqual(r.status_code, 200) mock_delay.assert_called_once() self.assertTrue(AdminAuditLog.objects.filter(action="task.retry").exists()) def test_retry_non_failed_rejected(self): self.assertEqual(self.ac.post(f"/api/admin/tasks/{self.t_ok.id}/retry/").status_code, 400) def test_retry_unsupported_type_rejected(self): self.assertEqual(self.ac.post(f"/api/admin/tasks/{self.t_script.id}/retry/").status_code, 400) def test_retry_requires_admin(self): self.assertEqual(self.nc.post(f"/api/admin/tasks/{self.t_failed.id}/retry/").status_code, 403) # ── 手动回收僵尸任务(reap):标失败 + 退预留;10 分钟窗口内/非 RESERVED/非超管全拒 ── def _mk_stuck_reserved(self, key: str, minutes_ago: int = 30): """卡死任务工厂:RESERVED + 冻结 20 积分 + 回拨 updated_at(auto_now 只能 queryset.update 绕)。""" from datetime import timedelta from django.utils import timezone from apps.billing.models import CreditAccount from apps.billing.services.ledger import reserve_credit CreditAccount.objects.get_or_create(team=self.team, defaults={"balance": Decimal("1000")}) task = self.AITask.objects.create( team=self.team, model_config=self.mc, task_type=self.AITask.Type.PRODUCT_IMAGE, status=self.AITask.Status.RESERVED, estimated_cost="20", idempotency_key=key, ) reserve_credit(team=self.team, user=self.normal, task=task, amount=Decimal("20")) self.AITask.objects.filter(id=task.id).update(updated_at=timezone.now() - timedelta(minutes=minutes_ago)) task.refresh_from_db() return task def test_reap_stuck_reserved_refunds(self): from apps.billing.models import CreditAccount, CreditReservation task = self._mk_stuck_reserved("k-reap-stuck") acct = CreditAccount.objects.get(team=self.team) self.assertEqual(acct.reserved_balance, Decimal("20")) # 列表侧 reapable 标记(前端据此显示按钮),与端点闸同口径 row = next(t for t in self.ac.get("/api/admin/tasks/?status=reserved").data["results"] if t["id"] == str(task.id)) self.assertTrue(row["reapable"]) r = self.ac.post(f"/api/admin/tasks/{task.id}/reap/") self.assertEqual(r.status_code, 200) self.assertEqual(r.data["status"], "failed") self.assertFalse(r.data["reapable"]) task.refresh_from_db() self.assertEqual(task.status, self.AITask.Status.FAILED) acct.refresh_from_db() self.assertEqual(acct.reserved_balance, Decimal("0")) self.assertEqual(task.credit_reservation.status, CreditReservation.Status.RELEASED) self.assertTrue(AdminAuditLog.objects.filter(action="task.reap").exists()) # 幂等收口:已终态再回收 → 400,不会双退 self.assertEqual(self.ac.post(f"/api/admin/tasks/{task.id}/reap/").status_code, 400) def test_reap_fresh_reserved_rejected(self): task = self._mk_stuck_reserved("k-reap-fresh", minutes_ago=0) r = self.ac.post(f"/api/admin/tasks/{task.id}/reap/") self.assertEqual(r.status_code, 400) task.refresh_from_db() self.assertEqual(task.status, self.AITask.Status.RESERVED) # 新鲜任务不显示回收按钮 row = next(t for t in self.ac.get("/api/admin/tasks/?status=reserved").data["results"] if t["id"] == str(task.id)) self.assertFalse(row["reapable"]) def test_reap_non_reserved_rejected(self): self.assertEqual(self.ac.post(f"/api/admin/tasks/{self.t_ok.id}/reap/").status_code, 400) def test_reap_requires_admin(self): task = self._mk_stuck_reserved("k-reap-perm") self.assertEqual(self.nc.post(f"/api/admin/tasks/{task.id}/reap/").status_code, 403) def _mk_inflight_video(self, *, key: str, provider_task_id: str = "ark-poll-1"): return self.AITask.objects.create( team=self.team, model_config=self.mc, task_type=self.AITask.Type.FREE_VIDEO, status=self.AITask.Status.POLLING, estimated_cost="1.0", actual_cost="0", idempotency_key=key, provider_task_id=provider_task_id, ) def test_poll_inflight_free_video(self): task = self._mk_inflight_video(key="k-poll-fv") def fake_finalize(*, task): task.status = self.AITask.Status.SUCCEEDED task.save(update_fields=["status", "updated_at"]) return task with patch("apps.ai.free_video.finalize_free_video", side_effect=fake_finalize) as mock_fin: r = self.ac.post("/api/admin/tasks/poll/", {}, format="json") self.assertEqual(r.status_code, 200) self.assertEqual(r.data["polled"], 1) self.assertEqual(r.data["statuses"][str(task.id)], "succeeded") mock_fin.assert_called_once() task.refresh_from_db() self.assertEqual(task.status, self.AITask.Status.SUCCEEDED) def test_poll_skips_terminal_and_missing_provider_id(self): self.AITask.objects.create( team=self.team, model_config=self.mc, task_type=self.AITask.Type.FREE_VIDEO, status=self.AITask.Status.POLLING, estimated_cost="1.0", actual_cost="0", idempotency_key="k-poll-nopid", ) with patch("apps.ai.free_video.finalize_free_video") as mock_fin: r = self.ac.post("/api/admin/tasks/poll/", {}, format="json") self.assertEqual(r.status_code, 200) self.assertEqual(r.data["polled"], 0) mock_fin.assert_not_called() def test_poll_task_ids_filter(self): keep = self._mk_inflight_video(key="k-poll-keep", provider_task_id="ark-keep") skip = self._mk_inflight_video(key="k-poll-skip", provider_task_id="ark-skip") with patch("apps.ai.free_video.finalize_free_video", side_effect=lambda *, task: task) as mock_fin: r = self.ac.post("/api/admin/tasks/poll/", {"task_ids": [str(keep.id)]}, format="json") self.assertEqual(r.status_code, 200) self.assertEqual(r.data["polled"], 1) self.assertIn(str(keep.id), r.data["statuses"]) self.assertNotIn(str(skip.id), r.data["statuses"]) mock_fin.assert_called_once() def test_poll_requires_admin(self): self.assertEqual(self.nc.post("/api/admin/tasks/poll/", {}, format="json").status_code, 403) class AdminBillingTests(TestCase): """Phase 7:计费审计(流水浏览/手动调额)+ 4 层额度策略(CRUD + 拦截生效)+ 权限。""" def setUp(self): from apps.billing.models import CreditAccount, CreditLedger self.admin = User.objects.create_user(username="padmin7", password="x", is_platform_admin=True) self.normal = User.objects.create_user(username="normal7", password="x") self.team = Team.objects.create(name="BillTeam", owner=self.normal) TeamMember.objects.create(team=self.team, user=self.normal, role=TeamMember.Role.OWNER) self.account = CreditAccount.objects.create(team=self.team, balance="100.0000") CreditLedger.objects.create(team=self.team, ledger_type=CreditLedger.Type.RECHARGE, amount="100", balance_after="100", reason="init") CreditLedger.objects.create(team=self.team, ledger_type=CreditLedger.Type.CHARGE, amount="10", balance_after="90", reason="gen") self.ac = APIClient() self.ac.force_authenticate(self.admin) self.nc = APIClient() self.nc.force_authenticate(self.normal) def test_ledger_browse_permission_and_filter(self): self.assertEqual(self.nc.get("/api/admin/ledgers/").status_code, 403) r = self.ac.get("/api/admin/ledgers/") self.assertEqual(r.status_code, 200) self.assertGreaterEqual(r.data["count"], 2) charge = self.ac.get("/api/admin/ledgers/?ledger_type=charge") self.assertTrue(all(item["ledger_type"] == "charge" for item in charge.data["results"])) def test_adjust_positive_updates_balance_and_audit(self): r = self.ac.post("/api/admin/ledgers/adjust/", {"team": str(self.team.id), "amount": "50", "reason": "补偿"}, format="json") self.assertEqual(r.status_code, 201) self.assertEqual(r.data["ledger_type"], "adjustment") self.account.refresh_from_db() self.assertEqual(self.account.balance, Decimal("150.0000")) self.assertTrue(AdminAuditLog.objects.filter(action="credit.adjust").exists()) def test_adjust_negative(self): r = self.ac.post("/api/admin/ledgers/adjust/", {"team": str(self.team.id), "amount": "-30", "reason": "扣"}, format="json") self.assertEqual(r.status_code, 201) self.account.refresh_from_db() self.assertEqual(self.account.balance, Decimal("70.0000")) def test_adjust_beyond_balance_rejected(self): r = self.ac.post("/api/admin/ledgers/adjust/", {"team": str(self.team.id), "amount": "-9999", "reason": "x"}, format="json") self.assertEqual(r.status_code, 400) def test_adjust_requires_admin(self): self.assertEqual( self.nc.post("/api/admin/ledgers/adjust/", {"team": str(self.team.id), "amount": "1"}, format="json").status_code, 403, ) def test_quota_policy_crud(self): self.assertEqual(self.nc.get("/api/admin/quota-policies/").status_code, 403) c = self.ac.post("/api/admin/quota-policies/", {"team": str(self.team.id), "per_task_limit": "5", "is_active": True}, format="json") self.assertEqual(c.status_code, 201) pid = c.data["id"] self.assertEqual(self.ac.get(f"/api/admin/quota-policies/?team={self.team.id}").data["count"], 1) up = self.ac.patch(f"/api/admin/quota-policies/{pid}/", {"per_task_limit": "9"}, format="json") self.assertEqual(up.data["per_task_limit"], "9.0000") self.assertEqual(self.ac.delete(f"/api/admin/quota-policies/{pid}/").status_code, 204) def test_quota_enforcement_per_task_and_guard(self): from apps.billing.models import QuotaPolicy from apps.billing.services.ledger import _enforce_quota_policy # 无策略 → 不拦截(零回归) _enforce_quota_policy(team=self.team, project=None, amount=Decimal("999")) policy = QuotaPolicy.objects.create(team=self.team, per_task_limit=Decimal("1"), is_active=True) with self.assertRaises(ValueError): _enforce_quota_policy(team=self.team, project=None, amount=Decimal("5")) _enforce_quota_policy(team=self.team, project=None, amount=Decimal("1")) # 限内放行 # 停用策略 → 不拦截 policy.is_active = False policy.save(update_fields=["is_active"]) _enforce_quota_policy(team=self.team, project=None, amount=Decimal("999")) def test_quota_enforcement_monthly(self): from apps.billing.models import QuotaPolicy from apps.billing.services.ledger import _enforce_quota_policy QuotaPolicy.objects.create(team=self.team, monthly_limit=Decimal("10"), is_active=True) # setUp 本月已有 CHARGE 10 → 再 +1 超 10 → 拦 with self.assertRaises(ValueError): _enforce_quota_policy(team=self.team, project=None, amount=Decimal("1")) class AdminModelProviderTests(TestCase): """Phase 8:模型供应商 / 模型 CRUD + 启停 + 定价 + 设默认(改变 get_default_model)+ 权限。""" def setUp(self): from apps.ai.models import ModelConfig, ModelProvider self.admin = User.objects.create_user(username="padmin8", password="x", is_platform_admin=True) self.normal = User.objects.create_user(username="normal8", password="x") self.prov = ModelProvider.objects.create(name="prov8", display_name="P8", api_key="secret-key") self.m1 = ModelConfig.objects.create(provider=self.prov, name="m8a", display_name="M8A", capability=ModelConfig.Capability.IMAGE) self.m2 = ModelConfig.objects.create(provider=self.prov, name="m8b", display_name="M8B", capability=ModelConfig.Capability.IMAGE) self.ac = APIClient() self.ac.force_authenticate(self.admin) self.nc = APIClient() self.nc.force_authenticate(self.normal) def test_providers_list_permission_and_api_key_hidden(self): self.assertEqual(self.nc.get("/api/admin/providers/").status_code, 403) r = self.ac.get("/api/admin/providers/") self.assertEqual(r.status_code, 200) row = next(p for p in r.data if p["name"] == "prov8") self.assertNotIn("api_key", row) # 密钥不回传 self.assertTrue(row["has_api_key"]) self.assertEqual(row["model_count"], 2) def test_provider_create_key_writeonly_persisted(self): from apps.ai.models import ModelProvider r = self.ac.post("/api/admin/providers/", {"name": "newp8", "display_name": "New", "api_key": "k123"}, format="json") self.assertEqual(r.status_code, 201) self.assertTrue(r.data["has_api_key"]) self.assertNotIn("api_key", r.data) self.assertEqual(ModelProvider.objects.get(name="newp8").api_key, "k123") def test_provider_toggle_status_and_delete(self): r = self.ac.patch(f"/api/admin/providers/{self.prov.id}/", {"status": "disabled"}, format="json") self.assertEqual(r.status_code, 200) self.assertEqual(r.data["status"], "disabled") def test_provider_error_rules_are_validated_on_admin_write(self): valid_metadata = { "error_rules": [ { "provider_code": "NewProviderCode", "contains_all": ["reference token", "expired"], "public_code": "asset_unavailable", "enabled": True, } ] } accepted = self.ac.patch( f"/api/admin/providers/{self.prov.id}/", {"metadata": valid_metadata}, format="json", ) self.assertEqual(accepted.status_code, 200) self.assertEqual(accepted.data["metadata"], valid_metadata) rejected = self.ac.patch( f"/api/admin/providers/{self.prov.id}/", { "metadata": { "error_rules": [ { "contains_all": [], "public_code": "arbitrary_new_code", } ] } }, format="json", ) self.assertEqual(rejected.status_code, 400) self.prov.refresh_from_db() self.assertEqual(self.prov.metadata, valid_metadata) def test_models_list_filter(self): r = self.ac.get(f"/api/admin/models/?provider={self.prov.id}") self.assertEqual(r.status_code, 200) self.assertEqual(len(r.data), 2) cap = self.ac.get("/api/admin/models/?capability=image") self.assertTrue(all(m["capability"] == "image" for m in cap.data)) def test_model_create_update_delete(self): c = self.ac.post( "/api/admin/models/", {"provider": str(self.prov.id), "name": "m8c", "display_name": "C", "capability": "text"}, format="json", ) self.assertEqual(c.status_code, 201) mid = c.data["id"] up = self.ac.patch(f"/api/admin/models/{mid}/", {"unit_price": "2.5", "status": "disabled"}, format="json") self.assertEqual(up.data["unit_price"], "2.5000") self.assertEqual(up.data["status"], "disabled") self.assertEqual(self.ac.delete(f"/api/admin/models/{mid}/").status_code, 204) def test_set_default_changes_get_default_model(self): from apps.ai.services import get_default_model r = self.ac.post(f"/api/admin/models/{self.m2.id}/set-default/") self.assertEqual(r.status_code, 200) self.assertTrue(r.data["is_default"]) self.assertEqual(get_default_model("image").id, self.m2.id) # 改设 m1 → m2 清默认 self.ac.post(f"/api/admin/models/{self.m1.id}/set-default/") self.assertEqual(get_default_model("image").id, self.m1.id) self.m2.refresh_from_db() self.assertFalse(self.m2.is_default) self.assertTrue(AdminAuditLog.objects.filter(action="model.set_default").exists()) def test_write_requires_admin(self): self.assertEqual(self.nc.post("/api/admin/providers/", {"name": "x", "display_name": "x"}, format="json").status_code, 403) self.assertEqual(self.nc.post(f"/api/admin/models/{self.m1.id}/set-default/").status_code, 403) class AdminGovernanceTests(TestCase): """Phase 9:收尾治理(全局项目监控 + 数据完整性体检)+ 权限。""" def setUp(self): from apps.products.models import Product from apps.projects.models import Project self.admin = User.objects.create_user(username="padmin9", password="x", is_platform_admin=True) self.normal = User.objects.create_user(username="normal9", password="x") self.team = Team.objects.create(name="GovTeam", owner=self.normal) TeamMember.objects.create(team=self.team, user=self.normal, role=TeamMember.Role.OWNER) self.product = Product.objects.create(team=self.team, title="无图商品", created_by=self.normal) self.proj_ok = Project.objects.create(team=self.team, name="P-ok", product=self.product, status=Project.Status.SCRIPTING, created_by=self.normal) self.proj_failed = Project.objects.create(team=self.team, name="P-failed", product=self.product, status=Project.Status.FAILED, created_by=self.normal) self.ac = APIClient() self.ac.force_authenticate(self.admin) self.nc = APIClient() self.nc.force_authenticate(self.normal) def test_projects_list_permission_and_filter(self): self.assertEqual(self.nc.get("/api/admin/projects/").status_code, 403) r = self.ac.get("/api/admin/projects/") self.assertEqual(r.status_code, 200) self.assertGreaterEqual(r.data["count"], 2) failed = self.ac.get("/api/admin/projects/?status=failed") self.assertTrue(all(p["status"] == "failed" for p in failed.data["results"])) # 行带团队名 + 商品名 row = next(p for p in r.data["results"] if p["name"] == "P-ok") self.assertEqual(row["team_name"], "GovTeam") self.assertEqual(row["product_title"], "无图商品") def test_integrity_report(self): self.assertEqual(self.nc.get("/api/admin/integrity/").status_code, 403) r = self.ac.get("/api/admin/integrity/") self.assertEqual(r.status_code, 200) checks = {c["key"]: c["count"] for c in r.data["checks"]} self.assertIn("products_without_image", checks) self.assertGreaterEqual(checks["products_without_image"], 1) # 我们建了一个无图商品 self.assertGreaterEqual(checks["failed_projects"], 1) self.assertGreaterEqual(checks["teams_without_account"], 1) # GovTeam 无 credit_account self.assertIn("totals", r.data) self.assertGreaterEqual(r.data["totals"]["projects"], 2)