from datetime import timedelta from django.core.cache import cache from django.db.models import Count, Q, Sum from django.utils import timezone from rest_framework import status from rest_framework.decorators import action from rest_framework.pagination import PageNumberPagination from rest_framework.response import Response from rest_framework.viewsets import ModelViewSet class NotificationPagination(PageNumberPagination): # 收件箱滚动加载:每批 10 条,前端可用 ?page_size= 覆盖(上限 100) page_size = 10 page_size_query_param = "page_size" max_page_size = 100 from apps.assets.models import Asset from apps.billing.models import CreditAccount, CreditLedger from apps.common.api import TeamScopedViewSetMixin from apps.projects.models import Project from .models import Notification from .serializers import NotificationSerializer def project_stage_label(project): return { "script": "Stage 1 · 脚本", "base_assets": "Stage 2 · 基础资产", "storyboard": "Stage 3 · 故事板", "video": "Stage 4 · 视频", "export": "Stage 5 · 导出", }.get(project.current_stage, "Stage 1 · 脚本") # AI 任务类型 → 计费消息里的中文阶段标签(脚本/图片/视频/配音都会扣费) TASK_TYPE_LABEL = { "script_generation": "脚本生成", "script_optimization": "脚本优化", "product_image": "商品图生成", "person_image": "模特图生成", "scene_image": "场景图生成", "storyboard": "故事板生成", "video_segment": "视频生成", "voiceover": "配音生成", "export": "成片导出", } def project_priority(project): if project.status == Project.Status.COMPLETED: return Notification.Priority.OK if project.status == Project.Status.FAILED: return Notification.Priority.ERR return Notification.Priority.INFO def _step_time(dt): """处理记录单步时间:今天只显示 HH:MM,跨天加 MM-DD(对齐设计稿的 mono 时间列)。""" if dt is None: return "" local = timezone.localtime(dt) now = timezone.localtime() return local.strftime("%H:%M") if local.date() == now.date() else local.strftime("%m-%d %H:%M") def ensure_team_notifications(team, user): def create_once(dedupe_key, **payload): obj, created = Notification.objects.get_or_create( team=team, recipient=user, dedupe_key=dedupe_key, defaults=payload, ) if created: return # 已存在的旧通知:补齐早期没存的字段(历史消息也能显示),不重建 updates = [] # 1) 处理记录 timeline timeline = (payload.get("metadata") or {}).get("timeline") if timeline and not (obj.metadata or {}).get("timeline"): obj.metadata = {**(obj.metadata or {}), "timeline": timeline} updates.append("metadata") # 2) 费用:早期建的是占位「-」,现在能从计费账本算出真实金额了就回填 new_cost = payload.get("cost_label") if new_cost and new_cost != "-" and obj.cost_label in ("", "-"): obj.cost_label = new_cost updates.append("cost_label") if updates: updates.append("updated_at") obj.save(update_fields=updates) create_once( "system:welcome", notification_type=Notification.Type.SYSTEM, priority=Notification.Priority.INFO, title="团队已接入 AirShelf", brief="真实消息中心已启用,状态会写入 Django 数据库。", body="消息已从演示数据切换为团队级通知表。已读、未读、归档等操作都会持久化保存。", source="Airshelf 系统", stage="系统公告", owner_label="系统", cost_label="-", related_url="settings.html#sec-notify", metadata={"timeline": [[_step_time(timezone.now()), "团队接入 AirShelf,真实消息中心已启用"]]}, ) recent_projects = list( Project.objects.filter(team=team).select_related("product", "created_by").order_by("-updated_at")[:5] ) # 5 个项目的累计扣费一次聚合(原先每个项目各跑一次 aggregate = N+1) spend_by_project = dict( CreditLedger.objects.filter(project__in=recent_projects, ledger_type=CreditLedger.Type.CHARGE) .values("project") .annotate(total=Sum("amount")) .values_list("project", "total") ) for project in recent_projects: product_title = project.product.title if project.product_id else "未绑定商品" # 项目累计花费 = 该项目所有 AI 扣费流水之和(脚本/图片/视频/导出),无扣费则显示「-」 project_spend = spend_by_project.get(project.id) project_cost = f"¥{project_spend:.2f}" if project_spend else "-" create_once( f"project:{project.id}:status:{project.status}:{project.current_stage}", notification_type=Notification.Type.TASK, priority=project_priority(project), title=f"项目「{project.name}」状态更新", brief=f"{product_title} · {project_stage_label(project)} · {project.get_status_display()}", body=f"项目「{project.name}」当前处于 {project_stage_label(project)}。这条消息来自 Django 项目表,刷新后状态会保持一致。", source="视频项目", project=project, stage=project_stage_label(project), owner_label=project.created_by.username if project.created_by_id else "成员", cost_label=project_cost, related_url=f"pipeline.html?project_id={project.id}", metadata={ "status": project.status, "current_stage": project.current_stage, "timeline": [ [_step_time(project.created_at), f"项目「{project.name}」创建"], [_step_time(project.updated_at), f"推进至 {project_stage_label(project)}"], [_step_time(project.updated_at), f"当前状态 · {project.get_status_display()}"], ], }, ) for asset in Asset.objects.filter(team=team).select_related("created_by", "origin_task").order_by("-updated_at")[:3]: # AI 生成的资产带 origin_task,可取其实际扣费;手动上传的资产无扣费 → 「-」 origin = asset.origin_task asset_cost = f"¥{origin.actual_cost:.2f}" if origin and origin.actual_cost else "-" create_once( f"asset:{asset.id}:created", notification_type=Notification.Type.TASK, priority=Notification.Priority.OK, title=f"资产「{asset.name}」已加入资产库", brief=f"{asset.get_category_display()} · {asset.get_asset_type_display()}", body="资产记录来自真实资产表。后续上传、AI 生成、导出成片都可以在这里形成团队通知。", source="资产库", stage="资产入库", owner_label=asset.created_by.username if asset.created_by_id else "成员", cost_label=asset_cost, related_url="library.html", metadata={ "asset_id": str(asset.id), "category": asset.category, "asset_type": asset.asset_type, "timeline": [ [_step_time(asset.created_at), f"{asset.get_category_display()} · {asset.get_asset_type_display()} 生成"], [_step_time(asset.updated_at), "通过基础审核,写入资产库"], ], }, ) # 计费消息:把真实扣费流水(图片/视频/脚本/配音生成)写成「计费」类通知,带真实金额。 # 原先 ensure_team_notifications 完全没读账本,所以计费消息的费用都是空的。 # 计费通知覆盖「消息保留期(90 天)」内的全部扣费流水,而不是固定最近 10 条 —— # 否则计费 tab 计数会被卡在 10,与实际扣费条数严重不符。上限 500 仅作暴走兜底。 charge_window_start = timezone.now() - timedelta(days=90) recent_charges = ( CreditLedger.objects.filter( team=team, ledger_type=CreditLedger.Type.CHARGE, created_at__gte=charge_window_start, ) .select_related("task", "project", "user") .order_by("-created_at")[:500] ) for ledger in recent_charges: task = ledger.task type_label = TASK_TYPE_LABEL.get(task.task_type, "AI 生成") if task else "AI 生成" project_name = ledger.project.name if ledger.project_id else "" amount_text = f"¥{ledger.amount:.2f}" scope = f" · {project_name}" if project_name else "" create_once( f"billing:charge:{ledger.id}", notification_type=Notification.Type.BILLING, priority=Notification.Priority.INFO, title=f"{type_label}已扣费 {amount_text}", brief=f"{ledger.reason or 'AI 任务扣费'}{scope} · 扣费后余额 ¥{ledger.balance_after:.2f}", body=( f"本次{type_label}消耗 {amount_text},来自真实计费账本(CreditLedger)。" f"扣费后团队余额为 ¥{ledger.balance_after:.2f}。明细可在消费页查看完整账本。" ), source="计费中心", project=ledger.project, stage=type_label, owner_label=ledger.user.username if ledger.user_id else "系统", cost_label=amount_text, related_url="account.html", metadata={ "ledger_id": str(ledger.id), "amount": str(ledger.amount), "timeline": [ [_step_time(ledger.created_at), f"{type_label}提交并预留额度"], [_step_time(ledger.created_at), f"实际扣费 {amount_text}"], [_step_time(ledger.created_at), f"扣费后余额 ¥{ledger.balance_after:.2f}"], ], }, ) account, _ = CreditAccount.objects.get_or_create(team=team) if account.balance <= 100: create_once( f"billing:low-balance:{account.id}", notification_type=Notification.Type.BILLING, priority=Notification.Priority.WARN, title="团队余额低于预警线", brief=f"当前余额 ¥{account.balance:.2f},建议及时充值。", body="余额低于 100 元时系统会生成预警通知。充值或调低成员额度后可在消费页查看最新账本。", source="计费中心", stage="余额监控", owner_label="系统", cost_label=f"¥{account.balance:.2f}", related_url="account.html", metadata={ "timeline": [ [_step_time(timezone.now()), f"余额降至 ¥{account.balance:.2f},低于预警线 ¥100"], [_step_time(timezone.now()), "已发送站内预警通知"], ], }, ) class NotificationViewSet(TeamScopedViewSetMixin, ModelViewSet): serializer_class = NotificationSerializer queryset = Notification.objects.select_related("team", "recipient", "project").all() pagination_class = NotificationPagination search_fields = ["title", "brief", "body", "source", "stage"] ordering_fields = ["created_at", "updated_at"] ordering = ["-created_at"] # 团队 + 收件人 + 未归档:分类计数的基准集(不含 tab/未读/搜索过滤) def _recipient_scope(self): queryset = super().get_queryset().filter(archived_at__isnull=True) user = self.request.user return queryset.filter(Q(recipient=user) | Q(recipient__isnull=True)) def get_queryset(self): queryset = self._recipient_scope() notification_type = self.request.query_params.get("type") if notification_type and notification_type not in {"all", "unread"}: queryset = queryset.filter(notification_type=notification_type) if self.request.query_params.get("unread") in {"1", "true", "yes"}: queryset = queryset.filter(is_read=False) return queryset def _refresh_notifications(self, request): """生成/补齐通知:搬到 Celery 后台执行,不再阻塞请求。 - 首次(团队还没有任何通知)同步跑一次,避免初次打开面板是空的; - 之后每团队最多每 60s 触发一次后台刷新,其余请求只读现有数据 → 接口秒回。 """ team = self.get_team() if not Notification.objects.filter(team=team).exists(): ensure_team_notifications(team, request.user) return if cache.add(f"ops:notif-refresh:{team.id}", 1, timeout=60): from .tasks import ensure_team_notifications_task user_id = str(request.user.id) if getattr(request.user, "id", None) else None ensure_team_notifications_task.delay(str(team.id), user_id) def list(self, request, *args, **kwargs): self._refresh_notifications(request) response = super().list(request, *args, **kwargs) data = response.data if isinstance(data, dict): # 分类 chip 计数取绝对总数(忽略当前 tab/搜索),与设计稿一致 base = self._recipient_scope() # 原先 6 次独立 count(每次 list 请求都跑)合成一次条件聚合,接口少 5 个 DB 往返 counts = base.aggregate( all=Count("id"), unread=Count("id", filter=Q(is_read=False)), task=Count("id", filter=Q(notification_type="task")), team=Count("id", filter=Q(notification_type="team")), billing=Count("id", filter=Q(notification_type="billing")), system=Count("id", filter=Q(notification_type="system")), ) unread_count = counts["unread"] data["unread_count"] = unread_count data["type_counts"] = { "all": counts["all"], "unread": counts["unread"], "task": counts["task"], "team": counts["team"], "billing": counts["billing"], "system": counts["system"], } return response def perform_create(self, serializer): serializer.save(team=self.get_team(), recipient=self.request.user) @action(detail=False, methods=["post"], url_path="mark-all-read") def mark_all_read(self, request): now = timezone.now() count = self.get_queryset().filter(is_read=False).update(is_read=True, read_at=now, updated_at=now) return Response({"updated": count, "unread_count": self.get_queryset().filter(is_read=False).count()}) @action(detail=False, methods=["post"], url_path="mark-all-unread") def mark_all_unread(self, request): now = timezone.now() count = self.get_queryset().filter(is_read=True).update(is_read=False, read_at=None, updated_at=now) return Response({"updated": count, "unread_count": self.get_queryset().filter(is_read=False).count()}) @action(detail=True, methods=["post"], url_path="mark-read") def mark_read(self, request, pk=None): notification = self.get_object() notification.mark_read() return Response(self.get_serializer(notification).data) @action(detail=True, methods=["post"], url_path="mark-unread") def mark_unread(self, request, pk=None): notification = self.get_object() notification.mark_unread() return Response(self.get_serializer(notification).data) @action(detail=True, methods=["post"], url_path="archive") def archive(self, request, pk=None): notification = self.get_object() notification.archive() return Response(status=status.HTTP_204_NO_CONTENT)