Compare commits
10 Commits
yuetsh
...
a1b51ebb9e
| Author | SHA1 | Date | |
|---|---|---|---|
| a1b51ebb9e | |||
| a9a6b87fef | |||
| 2d3588c755 | |||
| a2bfc28ac7 | |||
| 6aac767641 | |||
| 73af9d96b2 | |||
| 8a2fa11afc | |||
| 3f1c7250bd | |||
| bd0a7f30f8 | |||
| 8a043d2ffa |
@@ -1,9 +1,4 @@
|
||||
venv
|
||||
.venv
|
||||
.idea
|
||||
.git
|
||||
.DS_Store
|
||||
__pycache__
|
||||
*.pyc
|
||||
.ruff_cache
|
||||
.pytest_cache
|
||||
|
||||
10
.flake8
Normal file
10
.flake8
Normal file
@@ -0,0 +1,10 @@
|
||||
[flake8]
|
||||
exclude =
|
||||
xss_filter.py,
|
||||
*/migrations/,
|
||||
*settings.py
|
||||
*/apps.py
|
||||
venv/
|
||||
max-line-length = 180
|
||||
inline-quotes = "
|
||||
no-accept-encodings = True
|
||||
12
.github/issue_template.md
vendored
Normal file
12
.github/issue_template.md
vendored
Normal file
@@ -0,0 +1,12 @@
|
||||
在提交issue之前请
|
||||
|
||||
- 认真阅读文档 http://docs.onlinejudge.me/#/
|
||||
- 搜索和查看历史issues
|
||||
- 安全类问题请不要在 GitHub 上公布,请发送邮件到 `admin@qduoj.com`,根据漏洞危害程度发送红包感谢。
|
||||
|
||||
然后提交issue请写清楚下列事项
|
||||
|
||||
- 进行什么操作的时候遇到了什么问题,最好能有复现步骤
|
||||
- 错误提示是什么,如果看不到错误提示,请去data文件夹查看相应log文件。大段的错误提示请包在代码块标记里面。
|
||||
- 你尝试修复问题的操作
|
||||
- 页面问题请写清浏览器版本,尽量有截图
|
||||
32
.github/workflows/deploy.yml
vendored
32
.github/workflows/deploy.yml
vendored
@@ -1,32 +0,0 @@
|
||||
name: Deploy
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- yuetsh
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
jobs:
|
||||
deploy:
|
||||
runs-on: ubuntu-latest
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
include:
|
||||
- name: debian
|
||||
remote_port: 22
|
||||
script: /root/OJDeploy/backend.sh
|
||||
- name: school
|
||||
remote_port: 8822
|
||||
script: /root/OJ/backend.sh
|
||||
steps:
|
||||
- name: Deploy to ${{ matrix.name }}
|
||||
uses: appleboy/ssh-action@v1
|
||||
with:
|
||||
host: ${{ secrets.HOST }}
|
||||
port: ${{ matrix.remote_port }}
|
||||
username: root
|
||||
key: ${{ secrets.KEY }}
|
||||
script: sh ${{ matrix.script }}
|
||||
54
.github/workflows/release.yml
vendored
Normal file
54
.github/workflows/release.yml
vendored
Normal file
@@ -0,0 +1,54 @@
|
||||
name: Release build
|
||||
|
||||
on:
|
||||
push:
|
||||
tags:
|
||||
- v**
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
build:
|
||||
runs-on: ubuntu-latest
|
||||
environment: release
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
steps:
|
||||
- name: Docker metadata
|
||||
id: metadata
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
images: |
|
||||
registry.cn-hongkong.aliyuncs.com/oj-image/backend
|
||||
tags: |
|
||||
type=semver,pattern={{version}}
|
||||
type=semver,pattern={{major}}.{{minor}}
|
||||
type=semver,pattern={{major}},enable=${{ !startsWith(github.ref, 'refs/tags/v0.') }}
|
||||
|
||||
- name: Set up QEMU
|
||||
uses: docker/setup-qemu-action@v3
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v3
|
||||
|
||||
- name: Login to GitHub Container Registry
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
registry: ghcr.io
|
||||
username: ${{ github.actor }}
|
||||
password: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Login to Aliyun Container Registry
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
registry: registry.cn-hongkong.aliyuncs.com
|
||||
username: ${{ secrets.ALIYUN_ACR_USERNAME }}
|
||||
password: ${{ secrets.ALIYUN_ACR_PASSWORD }}
|
||||
|
||||
- name: Build and push
|
||||
uses: docker/build-push-action@v5
|
||||
with:
|
||||
push: true
|
||||
platforms: linux/amd64,linux/arm64
|
||||
tags: ${{ steps.metadata.outputs.tags }}
|
||||
annotations: ${{ steps.metadata.outputs.annotations }}
|
||||
144
CLAUDE.md
144
CLAUDE.md
@@ -1,144 +0,0 @@
|
||||
# CLAUDE.md
|
||||
|
||||
This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository.
|
||||
|
||||
## Project Overview
|
||||
|
||||
**OnlineJudge** is the backend for an Online Judge platform. Built with Django 6 + Django REST Framework, PostgreSQL, Redis, Django Channels (WebSocket), and Dramatiq (async task queue). Python 3.12+, managed with `uv`.
|
||||
|
||||
## Commands
|
||||
|
||||
```bash
|
||||
# Development
|
||||
python dev.py # Start dev server: Django on :8000 + Daphne WebSocket on :8001
|
||||
python manage.py runserver # HTTP only (no WebSocket support)
|
||||
python manage.py migrate # Apply database migrations
|
||||
python manage.py makemigrations # Create new migrations
|
||||
|
||||
# Dependencies
|
||||
uv sync # Install dependencies from uv.lock
|
||||
uv add <package> # Add a dependency
|
||||
|
||||
# Testing
|
||||
python manage.py test # Run all tests
|
||||
python manage.py test account # Run tests for a single app
|
||||
python manage.py test account.tests.TestClassName # Run a single test class
|
||||
|
||||
# Linting
|
||||
ruff check . # Lint (E, F, I rules, 180-char line length)
|
||||
ruff format . # Format (double quotes)
|
||||
|
||||
## Testing Policy
|
||||
|
||||
Do not write tests
|
||||
|
||||
# Initial setup
|
||||
python manage.py inituser --username admin --password <pw> --action create_super_admin
|
||||
python manage.py inituser --username admin --password <pw> --action reset
|
||||
```
|
||||
|
||||
## Architecture
|
||||
|
||||
### App Modules
|
||||
|
||||
Each Django app follows the same structure:
|
||||
```
|
||||
<app>/
|
||||
├── models.py # Django models
|
||||
├── serializers.py # DRF serializers
|
||||
├── views/
|
||||
│ ├── oj.py # User-facing API views
|
||||
│ └── admin.py # Admin API views
|
||||
└── urls/
|
||||
├── oj.py # User-facing URL patterns
|
||||
└── admin.py # Admin URL patterns
|
||||
```
|
||||
|
||||
Apps: `account`, `problem`, `submission`, `contest`, `ai`, `flowchart`, `problemset`, `class_pk`, `announcement`, `tutorial`, `message`, `reaction`, `conf`, `options`, `judge`
|
||||
|
||||
`utils/` is itself a Django app (listed in `INSTALLED_APPS`) — not just a helpers package. It provides `RichTextField` (XSS-sanitized `TextField`), `APIError`, the base `APIView`, caching, WebSocket helpers, and the `inituser` management command. Import shared utilities from `utils.*`.
|
||||
|
||||
### URL Routing
|
||||
|
||||
All routes are registered in `oj/urls.py`:
|
||||
- `api/` — user-facing endpoints
|
||||
- `api/admin/` — admin-only endpoints
|
||||
|
||||
WebSocket routing is in `oj/routing.py`.
|
||||
|
||||
### Settings Structure
|
||||
|
||||
- `oj/settings.py` — base configuration (imports dev or production settings based on `OJ_ENV`)
|
||||
- `oj/dev_settings.py` — development overrides (imported when `OJ_ENV != "production"`)
|
||||
- `oj/production_settings.py` — production overrides
|
||||
|
||||
### Base APIView & View Patterns
|
||||
|
||||
`utils/api/api.py` provides the custom base classes and decorators used by **all** views:
|
||||
|
||||
- **`APIView`** — base class for all views (not DRF's `APIView`). Key methods:
|
||||
- `self.success(data)` — returns `{"error": null, "data": data}`
|
||||
- `self.error(msg)` — returns `{"error": "error", "data": msg}`
|
||||
- `self.paginate_data(request, query_set, serializer)` — offset/limit pagination
|
||||
- `self.invalid_serializer(serializer)` — standard validation error response
|
||||
- **`CSRFExemptAPIView`** — same as `APIView` but CSRF-exempt
|
||||
- **`@validate_serializer(SerializerClass)`** — decorator for view methods that validates `request.data` against a serializer before the method runs. On success, `request.data` is replaced with validated data.
|
||||
|
||||
Typical view method pattern:
|
||||
```python
|
||||
@validate_serializer(CreateProblemSerializer)
|
||||
@super_admin_required
|
||||
def post(self, request):
|
||||
# request.data is already validated
|
||||
return self.success(...)
|
||||
```
|
||||
|
||||
### Authentication & Permissions
|
||||
|
||||
`account/decorators.py` provides decorators used on view methods:
|
||||
- `@login_required` / `@admin_role_required` / `@super_admin_required`
|
||||
- `@problem_permission_required`
|
||||
- `@check_contest_permission(check_type)` — validates contest access, sets `self.contest`
|
||||
- `ensure_created_by(obj, user)` — helper that raises `APIError` if user doesn't own the object
|
||||
|
||||
### Judge System
|
||||
|
||||
- `judge/dispatcher.py` — dispatches submissions to the judge sandbox (JudgeServer)
|
||||
- `judge/tasks.py` — Dramatiq async tasks for judging
|
||||
- `judge/languages.py` — language configurations (compile/run commands, limits)
|
||||
|
||||
Judge status codes are defined in `submission/models.py` (`JudgeStatus` class, codes -2 to 8) and must match the frontend's `utils/constants.ts`.
|
||||
|
||||
### Site Configuration (SysOptions)
|
||||
|
||||
`options/options.py` provides `SysOptions` — a metaclass-based system for site-wide configuration stored in the database with thread-local caching. Access settings like `SysOptions.smtp_config`, `SysOptions.languages`, etc.
|
||||
|
||||
### WebSocket (Channels)
|
||||
|
||||
`submission/consumers.py` — WebSocket consumer for real-time submission status updates. Uses `channels-redis` as the channel layer backend. Push updates via `utils/websocket.py:push_submission_update()`.
|
||||
|
||||
### Caching
|
||||
|
||||
Redis-backed via `django-redis`. Cache keys use MD5 hashing for consistency. See `utils/cache.py`.
|
||||
|
||||
### AI Integration
|
||||
|
||||
`utils/openai.py` — OpenAI client wrapper configured to work with OpenAI-compatible APIs (e.g., DeepSeek). Used by `ai/` app for submission analysis.
|
||||
|
||||
### Data Directory
|
||||
|
||||
Test cases and submission outputs are stored in a separate data directory (configured in settings, not in the repo). The `data/` directory in the repo contains configuration templates and `secret.key`.
|
||||
|
||||
## Key Domain Concepts
|
||||
|
||||
| Concept | Details |
|
||||
|---|---|
|
||||
| Problem types | ACM (binary accept/reject) vs OI (partial scoring) |
|
||||
| Judge statuses | COMPILE_ERROR(-2), WRONG_ANSWER(-1), ACCEPTED(0), CPU_TLE(1), REAL_TLE(2), MLE(3), RE(4), SE(5), PENDING(6), JUDGING(7), PARTIALLY_ACCEPTED(8) |
|
||||
| User roles | Regular / Admin / Super Admin |
|
||||
| Contest types | Public vs Password Protected |
|
||||
| Supported languages | C, C++, Python3, Java, JavaScript, Golang |
|
||||
|
||||
## Related Repository
|
||||
|
||||
The frontend is at `../ojnext` — a Vue 3 + Rsbuild project. See its CLAUDE.md for frontend details.
|
||||
41
Dockerfile
41
Dockerfile
@@ -1,42 +1,29 @@
|
||||
FROM python:3.13-slim
|
||||
FROM python:3.12.2-alpine
|
||||
ARG TARGETARCH
|
||||
ARG TARGETVARIANT
|
||||
|
||||
ENV OJ_ENV=production
|
||||
RUN sed -i 's/dl-cdn.alpinelinux.org/mirrors.ustc.edu.cn/g' /etc/apk/repositories
|
||||
|
||||
ENV OJ_ENV production
|
||||
WORKDIR /app
|
||||
|
||||
COPY ./deploy/requirements.txt /app/deploy/
|
||||
|
||||
RUN --mount=type=cache,target=/var/cache/apt,sharing=locked,id=apt-cache-$TARGETARCH$TARGETVARIANT-final \
|
||||
--mount=type=cache,target=/root/.cache/pip,id=pip-cache-$TARGETARCH$TARGETVARIANT-final \
|
||||
# psycopg2: libpg-dev
|
||||
# pillow: libjpeg-turbo-dev zlib-dev freetype-dev
|
||||
RUN --mount=type=cache,target=/etc/apk/cache,id=apk-cahce-$TARGETARCH$TARGETVARIANT-final \
|
||||
--mount=type=cache,target=/root/.cache/pip,id=pip-cahce-$TARGETARCH$TARGETVARIANT-final \
|
||||
<<EOS
|
||||
set -ex
|
||||
pip config set global.index-url https://pypi.tuna.tsinghua.edu.cn/simple
|
||||
if [ -f /etc/apt/sources.list.d/debian.sources ]; then
|
||||
sed -i 's|deb.debian.org|mirrors.tuna.tsinghua.edu.cn|g; s|security.debian.org|mirrors.tuna.tsinghua.edu.cn|g' /etc/apt/sources.list.d/debian.sources
|
||||
fi
|
||||
if [ -f /etc/apt/sources.list ]; then
|
||||
sed -i 's|deb.debian.org|mirrors.tuna.tsinghua.edu.cn|g; s|security.debian.org|mirrors.tuna.tsinghua.edu.cn|g' /etc/apt/sources.list
|
||||
fi
|
||||
apt-get update
|
||||
# libpq / libjpeg 都在 wheel 里自带(psycopg_binary.libs、pillow.libs),不装系统版。
|
||||
# zlib 由 base 镜像的 CPython 带着,Pillow 用的是它。
|
||||
apt-get install -y --no-install-recommends \
|
||||
ca-certificates \
|
||||
clang-format \
|
||||
passwd \
|
||||
supervisor
|
||||
pip config set global.index-url https://mirrors.ustc.edu.cn/pypi/web/simple
|
||||
apk add gcc libc-dev python3-dev libpq libpq-dev libjpeg-turbo libjpeg-turbo-dev zlib zlib-dev freetype freetype-dev supervisor openssl nginx curl unzip
|
||||
pip install -r /app/deploy/requirements.txt
|
||||
rm -rf /var/lib/apt/lists/*
|
||||
apk del gcc libc-dev python3-dev libpq-dev libjpeg-turbo-dev zlib-dev freetype-dev
|
||||
EOS
|
||||
|
||||
# Caddy 官方镜像里是静态链接的 Go 二进制,直接拷进 slim 就能跑,不需要额外依赖。
|
||||
COPY --from=caddy:2-alpine /usr/bin/caddy /usr/bin/caddy
|
||||
|
||||
COPY --chmod=755 ./ /app/
|
||||
COPY ./ /app/
|
||||
RUN mkdir -p /app/dist/
|
||||
RUN chmod -R u=rwX,go=rX ./ && chmod +x ./deploy/entrypoint.sh
|
||||
|
||||
HEALTHCHECK --interval=30s --timeout=3s --start-period=5s --retries=3 \
|
||||
CMD python3 /app/deploy/health_check.py
|
||||
HEALTHCHECK --interval=5s CMD [ "/usr/local/bin/python3", "/app/deploy/health_check.py" ]
|
||||
EXPOSE 8000
|
||||
ENTRYPOINT [ "/app/deploy/entrypoint.sh" ]
|
||||
|
||||
@@ -1,85 +1,60 @@
|
||||
import functools
|
||||
import hashlib
|
||||
import inspect
|
||||
import time
|
||||
|
||||
from contest.models import Contest, ContestStatus, ContestType
|
||||
from problem.models import Problem
|
||||
from utils.api import APIError, JSONResponse
|
||||
from contest.models import Contest, ContestType, ContestStatus, ContestRuleType
|
||||
from utils.api import JSONResponse, APIError
|
||||
from utils.constants import CONTEST_PASSWORD_SESSION_KEY
|
||||
|
||||
from .models import ProblemPermission
|
||||
|
||||
|
||||
class BasePermissionDecorator(object):
|
||||
def __init__(self, func):
|
||||
self.func = func
|
||||
functools.update_wrapper(self, func)
|
||||
|
||||
def __get__(self, obj, obj_type):
|
||||
if inspect.iscoroutinefunction(self.func):
|
||||
return functools.partial(self._async_call, obj)
|
||||
return functools.partial(self.__call__, obj)
|
||||
|
||||
def error(self, data, err="permission-denied"):
|
||||
return JSONResponse.response({"error": err, "data": data})
|
||||
|
||||
def _permission_error(self, request):
|
||||
if not request.user.is_authenticated:
|
||||
return self.error("请先登录", err="login-required")
|
||||
return self.error("权限不足", err="permission-denied")
|
||||
def error(self, data):
|
||||
return JSONResponse.response({"error": "permission-denied", "data": data})
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
request = args[1]
|
||||
self.request = args[1]
|
||||
|
||||
if self.check_permission(request):
|
||||
if request.user.is_disabled:
|
||||
return self.error("账号已禁用")
|
||||
if self.check_permission():
|
||||
if self.request.user.is_disabled:
|
||||
return self.error("Your account is disabled")
|
||||
return self.func(*args, **kwargs)
|
||||
else:
|
||||
return self._permission_error(request)
|
||||
return self.error("Please login first")
|
||||
|
||||
async def _async_call(self, *args, **kwargs):
|
||||
request = args[1]
|
||||
|
||||
if self.check_permission(request):
|
||||
if request.user.is_disabled:
|
||||
return self.error("账号已禁用")
|
||||
return await self.func(*args, **kwargs)
|
||||
return self._permission_error(request)
|
||||
|
||||
def check_permission(self, request):
|
||||
def check_permission(self):
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
class login_required(BasePermissionDecorator):
|
||||
def check_permission(self, request):
|
||||
return request.user.is_authenticated
|
||||
def check_permission(self):
|
||||
return self.request.user.is_authenticated
|
||||
|
||||
|
||||
class super_admin_required(BasePermissionDecorator):
|
||||
def check_permission(self, request):
|
||||
user = request.user
|
||||
def check_permission(self):
|
||||
user = self.request.user
|
||||
return user.is_authenticated and user.is_super_admin()
|
||||
|
||||
|
||||
class teacher_admin_required(BasePermissionDecorator):
|
||||
def check_permission(self, request):
|
||||
user = request.user
|
||||
return user.is_authenticated and user.is_teacher_or_above()
|
||||
|
||||
|
||||
class admin_role_required(BasePermissionDecorator):
|
||||
def check_permission(self, request):
|
||||
user = request.user
|
||||
def check_permission(self):
|
||||
user = self.request.user
|
||||
return user.is_authenticated and user.is_admin_role()
|
||||
|
||||
|
||||
class problem_permission_required(admin_role_required):
|
||||
def check_permission(self, request):
|
||||
if not super().check_permission(request):
|
||||
def check_permission(self):
|
||||
if not super(problem_permission_required, self).check_permission():
|
||||
return False
|
||||
if request.user.problem_permission == ProblemPermission.NONE:
|
||||
if self.request.user.problem_permission == ProblemPermission.NONE:
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -116,44 +91,48 @@ def check_contest_permission(check_type="details"):
|
||||
若通过验证,在view中可通过self.contest获得该contest
|
||||
"""
|
||||
|
||||
def _get_contest_id(request):
|
||||
return request.data.get("contest_id") or request.GET.get("contest_id")
|
||||
|
||||
def _check_access(self, request, user):
|
||||
if not user.is_authenticated:
|
||||
return self.error("请先登录", err="login-required")
|
||||
|
||||
if user.is_contest_admin(self.contest):
|
||||
return None
|
||||
|
||||
if self.contest.contest_type == ContestType.PASSWORD_PROTECTED_CONTEST:
|
||||
if not check_contest_password(request.session.get(CONTEST_PASSWORD_SESSION_KEY, {}).get(str(self.contest.id)), self.contest.password):
|
||||
return self.error("Wrong password or password expired")
|
||||
|
||||
if self.contest.status == ContestStatus.CONTEST_NOT_START and check_type != "details":
|
||||
return self.error("Contest has not started yet.")
|
||||
|
||||
return None
|
||||
|
||||
def decorator(func):
|
||||
@functools.wraps(func)
|
||||
async def _wrapper(*args, **kwargs):
|
||||
def _check_permission(*args, **kwargs):
|
||||
self = args[0]
|
||||
request = args[1]
|
||||
contest_id = _get_contest_id(request)
|
||||
user = request.user
|
||||
if request.data.get("contest_id"):
|
||||
contest_id = request.data["contest_id"]
|
||||
else:
|
||||
contest_id = request.GET.get("contest_id")
|
||||
if not contest_id:
|
||||
return self.error("Parameter error, contest_id is required")
|
||||
|
||||
try:
|
||||
self.contest = await Contest.objects.select_related("created_by").aget(id=contest_id, visible=True)
|
||||
# use self.contest to avoid query contest again in view.
|
||||
self.contest = Contest.objects.select_related("created_by").get(id=contest_id, visible=True)
|
||||
except Contest.DoesNotExist:
|
||||
return self.error("Contest %s doesn't exist" % contest_id)
|
||||
error = _check_access(self, request, request.user)
|
||||
if error:
|
||||
return error
|
||||
return await func(*args, **kwargs)
|
||||
|
||||
return _wrapper
|
||||
# Anonymous
|
||||
if not user.is_authenticated:
|
||||
return self.error("Please login first.")
|
||||
|
||||
# creator or owner
|
||||
if user.is_contest_admin(self.contest):
|
||||
return func(*args, **kwargs)
|
||||
|
||||
if self.contest.contest_type == ContestType.PASSWORD_PROTECTED_CONTEST:
|
||||
# password error
|
||||
if not check_contest_password(request.session.get(CONTEST_PASSWORD_SESSION_KEY, {}).get(self.contest.id), self.contest.password):
|
||||
return self.error("Wrong password or password expired")
|
||||
|
||||
# regular user get contest problems, ranks etc. before contest started
|
||||
if self.contest.status == ContestStatus.CONTEST_NOT_START and check_type != "details":
|
||||
return self.error("Contest has not started yet.")
|
||||
|
||||
# check does user have permission to get ranks, submissions in OI Contest
|
||||
if self.contest.status == ContestStatus.CONTEST_UNDERWAY and self.contest.rule_type == ContestRuleType.OI:
|
||||
if not self.contest.real_time_rank and (check_type == "ranks" or check_type == "submissions"):
|
||||
return self.error(f"No permission to get {check_type}")
|
||||
|
||||
return func(*args, **kwargs)
|
||||
return _check_permission
|
||||
return decorator
|
||||
|
||||
|
||||
|
||||
@@ -1,60 +0,0 @@
|
||||
from django.core.management.base import BaseCommand
|
||||
|
||||
from account.models import UserProfile
|
||||
from problem.models import Problem
|
||||
from submission.models import JudgeStatus
|
||||
|
||||
ACCEPTED_STATUSES = {JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED}
|
||||
|
||||
|
||||
class Command(BaseCommand):
|
||||
help = "从用户 Profile 中移除已被删除的题目记录,并同步修正 accepted_number"
|
||||
|
||||
def add_arguments(self, parser):
|
||||
parser.add_argument("--dry-run", action="store_true", help="只检查,不写入数据库")
|
||||
|
||||
def handle(self, *args, **options):
|
||||
dry_run = options["dry_run"]
|
||||
|
||||
# 所有现存非比赛题目的 PK 集合
|
||||
existing_ids = set(Problem.objects.filter(contest__isnull=True).values_list("id", flat=True))
|
||||
self.stdout.write(f"现存题库题目数: {len(existing_ids)}")
|
||||
|
||||
profiles = UserProfile.objects.select_related("user").exclude(acm_problems_status={})
|
||||
total = profiles.count()
|
||||
self.stdout.write(f"检查用户数: {total}{'(dry-run 模式)' if dry_run else ''}")
|
||||
|
||||
fixed_count = 0
|
||||
for profile in profiles:
|
||||
problems = profile.acm_problems_status.get("problems", {})
|
||||
if not problems:
|
||||
continue
|
||||
|
||||
stale_keys = [k for k in problems if int(k) not in existing_ids]
|
||||
if not stale_keys:
|
||||
continue
|
||||
|
||||
removed_accepted = sum(1 for k in stale_keys if problems[k].get("status") in ACCEPTED_STATUSES)
|
||||
|
||||
stale_display = [problems[k].get("_id", k) for k in stale_keys]
|
||||
self.stdout.write(
|
||||
f" 用户 {profile.user.username} | 删除 {len(stale_keys)} 题: {', '.join(stale_display)}{f' | 其中已AC {removed_accepted} 题' if removed_accepted else ''}"
|
||||
)
|
||||
|
||||
if dry_run:
|
||||
continue
|
||||
|
||||
for k in stale_keys:
|
||||
del profile.acm_problems_status["problems"][k]
|
||||
|
||||
if removed_accepted:
|
||||
# 防止 accepted_number 变为负数
|
||||
profile.accepted_number = max(0, profile.accepted_number - removed_accepted)
|
||||
|
||||
profile.save(update_fields=["acm_problems_status", "accepted_number"])
|
||||
fixed_count += 1
|
||||
|
||||
if dry_run:
|
||||
self.stdout.write(self.style.WARNING("dry-run 完成,未写入任何数据"))
|
||||
else:
|
||||
self.stdout.write(self.style.SUCCESS(f"完成,共修复 {fixed_count} 个用户 Profile"))
|
||||
@@ -1,10 +1,10 @@
|
||||
from django.conf import settings
|
||||
from django.db import connection
|
||||
from django.utils.deprecation import MiddlewareMixin
|
||||
from django.utils.timezone import now
|
||||
from django.utils.deprecation import MiddlewareMixin
|
||||
|
||||
from account.models import User
|
||||
from utils.api import JSONResponse
|
||||
from account.models import User
|
||||
|
||||
|
||||
class APITokenAuthMiddleware(MiddlewareMixin):
|
||||
@@ -37,10 +37,8 @@ class AdminRoleRequiredMiddleware(MiddlewareMixin):
|
||||
def process_request(self, request):
|
||||
path = request.path_info
|
||||
if path.startswith("/admin/") or path.startswith("/api/admin/"):
|
||||
if not request.user.is_authenticated:
|
||||
return JSONResponse.response({"error": "login-required", "data": "请先登录"})
|
||||
if not request.user.is_admin_role():
|
||||
return JSONResponse.response({"error": "permission-denied", "data": "权限不足"})
|
||||
if not (request.user.is_authenticated and request.user.is_admin_role()):
|
||||
return JSONResponse.response({"error": "login-required", "data": "Please login in first"})
|
||||
|
||||
|
||||
class LogSqlMiddleware(MiddlewareMixin):
|
||||
|
||||
@@ -1,18 +0,0 @@
|
||||
# Generated by Django 5.2.3 on 2025-09-19 06:11
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('account', '0001_initial'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='userprofile',
|
||||
name='class_name',
|
||||
field=models.TextField(null=True),
|
||||
),
|
||||
]
|
||||
@@ -1,22 +0,0 @@
|
||||
# Generated by Django 5.2.3 on 2025-09-19 06:14
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('account', '0002_userprofile_class_name'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RemoveField(
|
||||
model_name='userprofile',
|
||||
name='class_name',
|
||||
),
|
||||
migrations.AddField(
|
||||
model_name='user',
|
||||
name='class_name',
|
||||
field=models.TextField(null=True),
|
||||
),
|
||||
]
|
||||
@@ -1,22 +0,0 @@
|
||||
# Generated by Django 6.0.4 on 2026-05-09 08:18
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
dependencies = [
|
||||
("account", "0003_remove_userprofile_class_name_user_class_name"),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name="user",
|
||||
name="admin_type",
|
||||
field=models.TextField(choices=[("Regular User", "Regular User"), ("Admin", "Admin"), ("Super Admin", "Super Admin")], default="Regular User"),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name="user",
|
||||
name="problem_permission",
|
||||
field=models.TextField(choices=[("None", "None"), ("Own", "Own"), ("All", "All")], default="None"),
|
||||
),
|
||||
]
|
||||
@@ -1,58 +0,0 @@
|
||||
# Generated by Django 6.0.4 on 2026-05-09 11:53
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('account', '0004_alter_user_admin_type_alter_user_problem_permission'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='user',
|
||||
name='is_disabled',
|
||||
field=models.BooleanField(db_default=False, default=False),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='user',
|
||||
name='open_api',
|
||||
field=models.BooleanField(db_default=False, default=False),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='user',
|
||||
name='session_keys',
|
||||
field=models.JSONField(db_default=models.Value([], output_field=models.JSONField()), default=list),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='user',
|
||||
name='two_factor_auth',
|
||||
field=models.BooleanField(db_default=False, default=False),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='userprofile',
|
||||
name='accepted_number',
|
||||
field=models.IntegerField(db_default=0, default=0),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='userprofile',
|
||||
name='acm_problems_status',
|
||||
field=models.JSONField(db_default=models.Value({}, output_field=models.JSONField()), default=dict),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='userprofile',
|
||||
name='oi_problems_status',
|
||||
field=models.JSONField(db_default=models.Value({}, output_field=models.JSONField()), default=dict),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='userprofile',
|
||||
name='submission_number',
|
||||
field=models.IntegerField(db_default=0, default=0),
|
||||
),
|
||||
migrations.AlterField(
|
||||
model_name='userprofile',
|
||||
name='total_score',
|
||||
field=models.BigIntegerField(db_default=0, default=0),
|
||||
),
|
||||
]
|
||||
@@ -1,18 +0,0 @@
|
||||
# Generated by Django 6.0.4 on 2026-06-03 00:08
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('account', '0005_alter_user_is_disabled_alter_user_open_api_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='user',
|
||||
name='admin_type',
|
||||
field=models.TextField(choices=[('Regular User', 'Regular User'), ('Student Admin', 'Student Admin'), ('Teacher Admin', 'Teacher Admin'), ('Super Admin', 'Super Admin')], default='Regular User'),
|
||||
),
|
||||
]
|
||||
@@ -1,27 +0,0 @@
|
||||
# Generated by Django 6.0.4 on 2026-06-03 00:08
|
||||
|
||||
from django.db import migrations
|
||||
|
||||
|
||||
def rename_admin_to_student_admin(apps, schema_editor):
|
||||
User = apps.get_model("account", "User")
|
||||
User.objects.filter(admin_type="Admin").update(admin_type="Student Admin")
|
||||
|
||||
|
||||
def rename_student_admin_to_admin(apps, schema_editor):
|
||||
User = apps.get_model("account", "User")
|
||||
User.objects.filter(admin_type="Student Admin").update(admin_type="Admin")
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
("account", "0006_alter_user_admin_type"),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RunPython(
|
||||
rename_admin_to_student_admin,
|
||||
rename_student_admin_to_admin,
|
||||
),
|
||||
]
|
||||
@@ -1,21 +0,0 @@
|
||||
# Generated by Django 6.0.4 on 2026-07-04 17:32
|
||||
|
||||
from django.db import migrations
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('account', '0007_rename_admin_to_student_admin'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RemoveField(
|
||||
model_name='userprofile',
|
||||
name='oi_problems_status',
|
||||
),
|
||||
migrations.RemoveField(
|
||||
model_name='userprofile',
|
||||
name='total_score',
|
||||
),
|
||||
]
|
||||
@@ -1,21 +0,0 @@
|
||||
# Generated by Django 6.0.4 on 2026-08-06 04:22
|
||||
|
||||
from django.db import migrations
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('account', '0008_remove_userprofile_oi_problems_status_and_more'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RemoveField(
|
||||
model_name='user',
|
||||
name='tfa_token',
|
||||
),
|
||||
migrations.RemoveField(
|
||||
model_name='user',
|
||||
name='two_factor_auth',
|
||||
),
|
||||
]
|
||||
@@ -1,21 +0,0 @@
|
||||
# Generated by Django 6.0.4 on 2026-08-06 05:15
|
||||
|
||||
from django.db import migrations
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('account', '0009_remove_user_tfa_token_remove_user_two_factor_auth'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.RemoveField(
|
||||
model_name='user',
|
||||
name='reset_password_token',
|
||||
),
|
||||
migrations.RemoveField(
|
||||
model_name='user',
|
||||
name='reset_password_token_expire_time',
|
||||
),
|
||||
]
|
||||
@@ -1,21 +1,19 @@
|
||||
from django.conf import settings
|
||||
from django.contrib.auth.models import AbstractBaseUser
|
||||
from django.conf import settings
|
||||
from django.db import models
|
||||
|
||||
from utils.models import JSONField
|
||||
|
||||
|
||||
class AdminType(models.TextChoices):
|
||||
REGULAR_USER = "Regular User", "Regular User"
|
||||
STUDENT_ADMIN = "Student Admin", "Student Admin"
|
||||
TEACHER_ADMIN = "Teacher Admin", "Teacher Admin"
|
||||
SUPER_ADMIN = "Super Admin", "Super Admin"
|
||||
class AdminType(object):
|
||||
REGULAR_USER = "Regular User"
|
||||
ADMIN = "Admin"
|
||||
SUPER_ADMIN = "Super Admin"
|
||||
|
||||
|
||||
class ProblemPermission(models.TextChoices):
|
||||
NONE = "None", "None"
|
||||
OWN = "Own", "Own"
|
||||
ALL = "All", "All"
|
||||
class ProblemPermission(object):
|
||||
NONE = "None"
|
||||
OWN = "Own"
|
||||
ALL = "All"
|
||||
|
||||
|
||||
class UserManager(models.Manager):
|
||||
@@ -24,59 +22,50 @@ class UserManager(models.Manager):
|
||||
def get_by_natural_key(self, username):
|
||||
return self.get(**{f"{self.model.USERNAME_FIELD}__iexact": username})
|
||||
|
||||
async def aget_by_natural_key(self, username):
|
||||
return await self.aget(**{f"{self.model.USERNAME_FIELD}__iexact": username})
|
||||
|
||||
|
||||
class User(AbstractBaseUser):
|
||||
username = models.TextField(unique=True)
|
||||
class_name = models.TextField(null=True)
|
||||
email = models.TextField(null=True)
|
||||
create_time = models.DateTimeField(auto_now_add=True, null=True)
|
||||
# One of UserType
|
||||
admin_type = models.TextField(default=AdminType.REGULAR_USER, choices=AdminType.choices)
|
||||
problem_permission = models.TextField(default=ProblemPermission.NONE, choices=ProblemPermission.choices)
|
||||
admin_type = models.TextField(default=AdminType.REGULAR_USER)
|
||||
problem_permission = models.TextField(default=ProblemPermission.NONE)
|
||||
reset_password_token = models.TextField(null=True)
|
||||
reset_password_token_expire_time = models.DateTimeField(null=True)
|
||||
# SSO auth token
|
||||
auth_token = models.TextField(null=True)
|
||||
session_keys = JSONField(default=list, db_default=models.Value([], output_field=models.JSONField()))
|
||||
two_factor_auth = models.BooleanField(default=False)
|
||||
tfa_token = models.TextField(null=True)
|
||||
session_keys = JSONField(default=list)
|
||||
# open api key
|
||||
open_api = models.BooleanField(default=False, db_default=False)
|
||||
open_api = models.BooleanField(default=False)
|
||||
open_api_appkey = models.TextField(null=True)
|
||||
is_disabled = models.BooleanField(default=False, db_default=False)
|
||||
raw_password = models.CharField(max_length=20, null=True, blank=True, verbose_name="明文密码")
|
||||
is_disabled = models.BooleanField(default=False)
|
||||
raw_password = models.CharField(
|
||||
max_length=20, null=True, blank=True, verbose_name="明文密码"
|
||||
)
|
||||
|
||||
USERNAME_FIELD = "username"
|
||||
REQUIRED_FIELDS = []
|
||||
|
||||
objects = UserManager()
|
||||
|
||||
def is_regular_user(self):
|
||||
return self.admin_type == AdminType.REGULAR_USER
|
||||
|
||||
def is_student_admin(self):
|
||||
return self.admin_type == AdminType.STUDENT_ADMIN
|
||||
|
||||
def is_teacher_admin(self):
|
||||
return self.admin_type == AdminType.TEACHER_ADMIN
|
||||
def is_admin(self):
|
||||
return self.admin_type == AdminType.ADMIN
|
||||
|
||||
def is_super_admin(self):
|
||||
return self.admin_type == AdminType.SUPER_ADMIN
|
||||
|
||||
def is_admin_role(self):
|
||||
return self.admin_type in [
|
||||
AdminType.STUDENT_ADMIN,
|
||||
AdminType.TEACHER_ADMIN,
|
||||
AdminType.SUPER_ADMIN,
|
||||
]
|
||||
|
||||
def is_teacher_or_above(self):
|
||||
return self.admin_type in [AdminType.TEACHER_ADMIN, AdminType.SUPER_ADMIN]
|
||||
return self.admin_type in [AdminType.ADMIN, AdminType.SUPER_ADMIN]
|
||||
|
||||
def can_mgmt_all_problem(self):
|
||||
return self.problem_permission == ProblemPermission.ALL
|
||||
|
||||
def is_contest_admin(self, contest):
|
||||
return self.is_authenticated and (contest.created_by == self or self.admin_type == AdminType.SUPER_ADMIN)
|
||||
return self.is_authenticated and (
|
||||
contest.created_by == self or self.admin_type == AdminType.SUPER_ADMIN
|
||||
)
|
||||
|
||||
def set_password(self, raw_password):
|
||||
super().set_password(raw_password)
|
||||
@@ -103,7 +92,9 @@ class UserProfile(models.Model):
|
||||
# }
|
||||
# }
|
||||
# }
|
||||
acm_problems_status = JSONField(default=dict, db_default=models.Value({}, output_field=models.JSONField()))
|
||||
acm_problems_status = JSONField(default=dict)
|
||||
# like acm_problems_status, merely add "score" field
|
||||
oi_problems_status = JSONField(default=dict)
|
||||
|
||||
real_name = models.TextField(null=True)
|
||||
avatar = models.TextField(default=f"{settings.AVATAR_URI_PREFIX}/default.png")
|
||||
@@ -113,16 +104,25 @@ class UserProfile(models.Model):
|
||||
school = models.TextField(null=True)
|
||||
major = models.TextField(null=True)
|
||||
language = models.TextField(null=True)
|
||||
accepted_number = models.IntegerField(default=0, db_default=0)
|
||||
submission_number = models.IntegerField(default=0, db_default=0)
|
||||
# for ACM
|
||||
accepted_number = models.IntegerField(default=0)
|
||||
# for OI
|
||||
total_score = models.BigIntegerField(default=0)
|
||||
submission_number = models.IntegerField(default=0)
|
||||
|
||||
def add_accepted_problem_number(self):
|
||||
self.accepted_number = models.F("accepted_number") + 1
|
||||
self.save(update_fields=["accepted_number"])
|
||||
self.save()
|
||||
|
||||
def add_submission_number(self):
|
||||
self.submission_number = models.F("submission_number") + 1
|
||||
self.save(update_fields=["submission_number"])
|
||||
self.save()
|
||||
|
||||
# 计算总分时, 应先减掉上次该题所得分数, 然后再加上本次所得分数
|
||||
def add_score(self, this_time_score, last_time_score=None):
|
||||
last_time_score = last_time_score or 0
|
||||
self.total_score = models.F("total_score") - last_time_score + this_time_score
|
||||
self.save()
|
||||
|
||||
class Meta:
|
||||
db_table = "user_profile"
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from django import forms
|
||||
|
||||
from utils.api import UsernameSerializer, serializers
|
||||
from utils.api import serializers, UsernameSerializer
|
||||
|
||||
from .models import AdminType, ProblemPermission, User, UserProfile
|
||||
|
||||
@@ -8,16 +8,45 @@ from .models import AdminType, ProblemPermission, User, UserProfile
|
||||
class UserLoginSerializer(serializers.Serializer):
|
||||
username = serializers.CharField()
|
||||
password = serializers.CharField()
|
||||
tfa_code = serializers.CharField(required=False, allow_blank=True)
|
||||
|
||||
|
||||
class UsernameOrEmailCheckSerializer(serializers.Serializer):
|
||||
username = serializers.CharField(required=False)
|
||||
email = serializers.EmailField(required=False)
|
||||
|
||||
|
||||
class UserRegisterSerializer(serializers.Serializer):
|
||||
username = serializers.CharField(max_length=32)
|
||||
password = serializers.CharField(min_length=6)
|
||||
email = serializers.EmailField(max_length=64)
|
||||
captcha = serializers.CharField()
|
||||
|
||||
|
||||
class UserChangePasswordSerializer(serializers.Serializer):
|
||||
old_password = serializers.CharField()
|
||||
new_password = serializers.CharField(min_length=6)
|
||||
tfa_code = serializers.CharField(required=False, allow_blank=True)
|
||||
|
||||
|
||||
class UserChangeEmailSerializer(serializers.Serializer):
|
||||
password = serializers.CharField()
|
||||
new_email = serializers.EmailField(max_length=64)
|
||||
tfa_code = serializers.CharField(required=False, allow_blank=True)
|
||||
|
||||
|
||||
class GenerateUserSerializer(serializers.Serializer):
|
||||
prefix = serializers.CharField(max_length=16, allow_blank=True)
|
||||
suffix = serializers.CharField(max_length=16, allow_blank=True)
|
||||
number_from = serializers.IntegerField()
|
||||
number_to = serializers.IntegerField()
|
||||
password_length = serializers.IntegerField(max_value=16, default=8)
|
||||
|
||||
|
||||
class ImportUserSerializer(serializers.Serializer):
|
||||
users = serializers.ListField(child=serializers.ListField(child=serializers.CharField(max_length=64)))
|
||||
users = serializers.ListField(
|
||||
child=serializers.ListField(child=serializers.CharField(max_length=64))
|
||||
)
|
||||
|
||||
|
||||
class UserAdminSerializer(serializers.ModelSerializer):
|
||||
@@ -34,10 +63,10 @@ class UserAdminSerializer(serializers.ModelSerializer):
|
||||
"real_name",
|
||||
"create_time",
|
||||
"last_login",
|
||||
"two_factor_auth",
|
||||
"open_api",
|
||||
"is_disabled",
|
||||
"raw_password",
|
||||
"class_name",
|
||||
]
|
||||
|
||||
def get_real_name(self, obj):
|
||||
@@ -61,9 +90,9 @@ class UserSerializer(serializers.ModelSerializer):
|
||||
"problem_permission",
|
||||
"create_time",
|
||||
"last_login",
|
||||
"two_factor_auth",
|
||||
"open_api",
|
||||
"is_disabled",
|
||||
"class_name",
|
||||
]
|
||||
|
||||
|
||||
@@ -87,13 +116,19 @@ class EditUserSerializer(serializers.Serializer):
|
||||
id = serializers.IntegerField()
|
||||
username = serializers.CharField(max_length=32)
|
||||
real_name = serializers.CharField(max_length=32, allow_blank=True, allow_null=True)
|
||||
password = serializers.CharField(min_length=6, allow_blank=True, required=False, default=None)
|
||||
password = serializers.CharField(
|
||||
min_length=6, allow_blank=True, required=False, default=None
|
||||
)
|
||||
email = serializers.EmailField(max_length=64)
|
||||
admin_type = serializers.ChoiceField(choices=AdminType.choices)
|
||||
problem_permission = serializers.ChoiceField(choices=ProblemPermission.choices)
|
||||
admin_type = serializers.ChoiceField(
|
||||
choices=(AdminType.REGULAR_USER, AdminType.ADMIN, AdminType.SUPER_ADMIN)
|
||||
)
|
||||
problem_permission = serializers.ChoiceField(
|
||||
choices=(ProblemPermission.NONE, ProblemPermission.OWN, ProblemPermission.ALL)
|
||||
)
|
||||
open_api = serializers.BooleanField()
|
||||
two_factor_auth = serializers.BooleanField()
|
||||
is_disabled = serializers.BooleanField()
|
||||
class_name = serializers.CharField(required=False, allow_null=True, allow_blank=True)
|
||||
|
||||
|
||||
class EditUserProfileSerializer(serializers.Serializer):
|
||||
@@ -107,10 +142,33 @@ class EditUserProfileSerializer(serializers.Serializer):
|
||||
language = serializers.CharField(max_length=32, allow_blank=True, required=False)
|
||||
|
||||
|
||||
class ApplyResetPasswordSerializer(serializers.Serializer):
|
||||
email = serializers.EmailField()
|
||||
captcha = serializers.CharField()
|
||||
|
||||
|
||||
class ResetPasswordSerializer(serializers.Serializer):
|
||||
token = serializers.CharField()
|
||||
password = serializers.CharField(min_length=6)
|
||||
captcha = serializers.CharField()
|
||||
|
||||
|
||||
class SSOSerializer(serializers.Serializer):
|
||||
token = serializers.CharField()
|
||||
|
||||
|
||||
class TwoFactorAuthCodeSerializer(serializers.Serializer):
|
||||
code = serializers.IntegerField()
|
||||
|
||||
|
||||
class ImageUploadForm(forms.Form):
|
||||
image = forms.FileField()
|
||||
|
||||
|
||||
class FileUploadForm(forms.Form):
|
||||
file = forms.FileField()
|
||||
|
||||
|
||||
class RankInfoSerializer(serializers.ModelSerializer):
|
||||
user = UsernameSerializer()
|
||||
|
||||
|
||||
22
account/tasks.py
Normal file
22
account/tasks.py
Normal file
@@ -0,0 +1,22 @@
|
||||
import logging
|
||||
import dramatiq
|
||||
|
||||
from options.options import SysOptions
|
||||
from utils.shortcuts import send_email, DRAMATIQ_WORKER_ARGS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dramatiq.actor(**DRAMATIQ_WORKER_ARGS(max_retries=3))
|
||||
def send_email_async(from_name, to_email, to_name, subject, content):
|
||||
if not SysOptions.smtp_config:
|
||||
return
|
||||
try:
|
||||
send_email(smtp_config=SysOptions.smtp_config,
|
||||
from_name=from_name,
|
||||
to_email=to_email,
|
||||
to_name=to_name,
|
||||
subject=subject,
|
||||
content=content)
|
||||
except Exception as e:
|
||||
logger.exception(e)
|
||||
31
account/templates/reset_password_email.html
Normal file
31
account/templates/reset_password_email.html
Normal file
@@ -0,0 +1,31 @@
|
||||
<div>
|
||||
<table cellpadding="0" align="center"
|
||||
style="overflow:hidden;background:#fff;margin:0 auto;text-align:left;position:relative;font-size:14px; font-family:'lucida Grande',Verdana;line-height:1.5;box-shadow:0 0 3px #ccc;border:1px solid #ccc;border-radius:5px;border-collapse:collapse;">
|
||||
<tbody>
|
||||
<tr>
|
||||
<th valign="middle"
|
||||
style="height:38px;color:#fff; font-size:14px;line-height:38px; font-weight:bold;text-align:left;padding:10px 24px 6px; border-bottom:1px solid #467ec3;background:#518bcb;border-radius:5px 5px 0 0;">
|
||||
{{ website_name }}</th>
|
||||
</tr>
|
||||
<tr>
|
||||
<td>
|
||||
<div style="padding:20px 35px 40px;">
|
||||
<h2 style="font-weight:bold;margin-bottom:5px;font-size:14px;">Hello, {{ username }}:</h2>
|
||||
<p style="margin-top:20px">
|
||||
Please click <a href="{{ link }}">{{ link }}</a> to reset your password in 20 minutes.
|
||||
</p>
|
||||
<p style="margin-top:20px">
|
||||
To protect your account, please do not use simple passwords.
|
||||
</p>
|
||||
<p style="margin-top:20px">
|
||||
If you still have any questions, please contract system administrator.
|
||||
</p>
|
||||
<p style="margin-left:2em;"></p>
|
||||
<p style="text-indent:0;text-align:right;">{{ website_name }}</p>
|
||||
</div>
|
||||
</td>
|
||||
</tr>
|
||||
|
||||
</tbody>
|
||||
</table>
|
||||
</div>
|
||||
646
account/tests.py
Normal file
646
account/tests.py
Normal file
@@ -0,0 +1,646 @@
|
||||
import time
|
||||
|
||||
from unittest import mock
|
||||
from datetime import timedelta
|
||||
from copy import deepcopy
|
||||
|
||||
from django.contrib import auth
|
||||
from django.utils.timezone import now
|
||||
from otpauth import OtpAuth
|
||||
|
||||
from utils.api.tests import APIClient, APITestCase
|
||||
from utils.shortcuts import rand_str
|
||||
from options.options import SysOptions
|
||||
|
||||
from .models import AdminType, ProblemPermission, User
|
||||
from utils.constants import ContestRuleType
|
||||
|
||||
|
||||
class PermissionDecoratorTest(APITestCase):
|
||||
def setUp(self):
|
||||
self.regular_user = User.objects.create(username="regular_user")
|
||||
self.admin = User.objects.create(username="admin")
|
||||
self.super_admin = User.objects.create(username="super_admin")
|
||||
self.request = mock.MagicMock()
|
||||
self.request.user.is_authenticated = mock.MagicMock()
|
||||
|
||||
def test_login_required(self):
|
||||
self.request.user.is_authenticated.return_value = False
|
||||
|
||||
def test_admin_required(self):
|
||||
pass
|
||||
|
||||
def test_super_admin_required(self):
|
||||
pass
|
||||
|
||||
|
||||
class DuplicateUserCheckAPITest(APITestCase):
|
||||
def setUp(self):
|
||||
user = self.create_user("test", "test123", login=False)
|
||||
user.email = "test@test.com"
|
||||
user.save()
|
||||
self.url = self.reverse("check_username_or_email")
|
||||
|
||||
def test_duplicate_username(self):
|
||||
resp = self.client.post(self.url, data={"username": "test"})
|
||||
data = resp.data["data"]
|
||||
self.assertEqual(data["username"], True)
|
||||
resp = self.client.post(self.url, data={"username": "Test"})
|
||||
self.assertEqual(resp.data["data"]["username"], True)
|
||||
|
||||
def test_ok_username(self):
|
||||
resp = self.client.post(self.url, data={"username": "test1"})
|
||||
data = resp.data["data"]
|
||||
self.assertFalse(data["username"])
|
||||
|
||||
def test_duplicate_email(self):
|
||||
resp = self.client.post(self.url, data={"email": "test@test.com"})
|
||||
self.assertEqual(resp.data["data"]["email"], True)
|
||||
resp = self.client.post(self.url, data={"email": "Test@Test.com"})
|
||||
self.assertTrue(resp.data["data"]["email"])
|
||||
|
||||
def test_ok_email(self):
|
||||
resp = self.client.post(self.url, data={"email": "aa@test.com"})
|
||||
self.assertFalse(resp.data["data"]["email"])
|
||||
|
||||
|
||||
class TFARequiredCheckAPITest(APITestCase):
|
||||
def setUp(self):
|
||||
self.url = self.reverse("tfa_required_check")
|
||||
self.create_user("test", "test123", login=False)
|
||||
|
||||
def test_not_required_tfa(self):
|
||||
resp = self.client.post(self.url, data={"username": "test"})
|
||||
self.assertSuccess(resp)
|
||||
self.assertEqual(resp.data["data"]["result"], False)
|
||||
|
||||
def test_required_tfa(self):
|
||||
user = User.objects.first()
|
||||
user.two_factor_auth = True
|
||||
user.save()
|
||||
resp = self.client.post(self.url, data={"username": "test"})
|
||||
self.assertEqual(resp.data["data"]["result"], True)
|
||||
|
||||
|
||||
class UserLoginAPITest(APITestCase):
|
||||
def setUp(self):
|
||||
self.username = self.password = "test"
|
||||
self.user = self.create_user(username=self.username, password=self.password, login=False)
|
||||
self.login_url = self.reverse("user_login_api")
|
||||
|
||||
def _set_tfa(self):
|
||||
self.user.two_factor_auth = True
|
||||
tfa_token = rand_str(32)
|
||||
self.user.tfa_token = tfa_token
|
||||
self.user.save()
|
||||
return tfa_token
|
||||
|
||||
def test_login_with_correct_info(self):
|
||||
response = self.client.post(self.login_url,
|
||||
data={"username": self.username, "password": self.password})
|
||||
self.assertDictEqual(response.data, {"error": None, "data": "Succeeded"})
|
||||
|
||||
user = auth.get_user(self.client)
|
||||
self.assertTrue(user.is_authenticated)
|
||||
|
||||
def test_login_with_correct_info_upper_username(self):
|
||||
resp = self.client.post(self.login_url, data={"username": self.username.upper(), "password": self.password})
|
||||
self.assertDictEqual(resp.data, {"error": None, "data": "Succeeded"})
|
||||
user = auth.get_user(self.client)
|
||||
self.assertTrue(user.is_authenticated)
|
||||
|
||||
def test_login_with_wrong_info(self):
|
||||
response = self.client.post(self.login_url,
|
||||
data={"username": self.username, "password": "invalid_password"})
|
||||
self.assertDictEqual(response.data, {"error": "error", "data": "Invalid username or password"})
|
||||
|
||||
user = auth.get_user(self.client)
|
||||
self.assertFalse(user.is_authenticated)
|
||||
|
||||
def test_tfa_login(self):
|
||||
token = self._set_tfa()
|
||||
code = OtpAuth(token).totp()
|
||||
if len(str(code)) < 6:
|
||||
code = (6 - len(str(code))) * "0" + str(code)
|
||||
response = self.client.post(self.login_url,
|
||||
data={"username": self.username,
|
||||
"password": self.password,
|
||||
"tfa_code": code})
|
||||
self.assertDictEqual(response.data, {"error": None, "data": "Succeeded"})
|
||||
|
||||
user = auth.get_user(self.client)
|
||||
self.assertTrue(user.is_authenticated)
|
||||
|
||||
def test_tfa_login_wrong_code(self):
|
||||
self._set_tfa()
|
||||
response = self.client.post(self.login_url,
|
||||
data={"username": self.username,
|
||||
"password": self.password,
|
||||
"tfa_code": "qqqqqq"})
|
||||
self.assertDictEqual(response.data, {"error": "error", "data": "Invalid two factor verification code"})
|
||||
|
||||
user = auth.get_user(self.client)
|
||||
self.assertFalse(user.is_authenticated)
|
||||
|
||||
def test_tfa_login_without_code(self):
|
||||
self._set_tfa()
|
||||
response = self.client.post(self.login_url,
|
||||
data={"username": self.username,
|
||||
"password": self.password})
|
||||
self.assertDictEqual(response.data, {"error": "error", "data": "tfa_required"})
|
||||
|
||||
user = auth.get_user(self.client)
|
||||
self.assertFalse(user.is_authenticated)
|
||||
|
||||
def test_user_disabled(self):
|
||||
self.user.is_disabled = True
|
||||
self.user.save()
|
||||
resp = self.client.post(self.login_url, data={"username": self.username,
|
||||
"password": self.password})
|
||||
self.assertDictEqual(resp.data, {"error": "error", "data": "Your account has been disabled"})
|
||||
|
||||
|
||||
class CaptchaTest(APITestCase):
|
||||
def _set_captcha(self, session):
|
||||
captcha = rand_str(4)
|
||||
session["_django_captcha_key"] = captcha
|
||||
session["_django_captcha_expires_time"] = int(time.time()) + 30
|
||||
session.save()
|
||||
return captcha
|
||||
|
||||
|
||||
class UserRegisterAPITest(CaptchaTest):
|
||||
def setUp(self):
|
||||
self.client = APIClient()
|
||||
self.register_url = self.reverse("user_register_api")
|
||||
self.captcha = rand_str(4)
|
||||
|
||||
self.data = {"username": "test_user", "password": "testuserpassword",
|
||||
"real_name": "real_name", "email": "test@qduoj.com",
|
||||
"captcha": self._set_captcha(self.client.session)}
|
||||
|
||||
def test_website_config_limit(self):
|
||||
SysOptions.allow_register = False
|
||||
resp = self.client.post(self.register_url, data=self.data)
|
||||
self.assertDictEqual(resp.data, {"error": "error", "data": "Register function has been disabled by admin"})
|
||||
|
||||
def test_invalid_captcha(self):
|
||||
self.data["captcha"] = "****"
|
||||
response = self.client.post(self.register_url, data=self.data)
|
||||
self.assertDictEqual(response.data, {"error": "error", "data": "Invalid captcha"})
|
||||
|
||||
self.data.pop("captcha")
|
||||
response = self.client.post(self.register_url, data=self.data)
|
||||
self.assertTrue(response.data["error"] is not None)
|
||||
|
||||
def test_register_with_correct_info(self):
|
||||
response = self.client.post(self.register_url, data=self.data)
|
||||
self.assertDictEqual(response.data, {"error": None, "data": "Succeeded"})
|
||||
|
||||
def test_username_already_exists(self):
|
||||
self.test_register_with_correct_info()
|
||||
|
||||
self.data["captcha"] = self._set_captcha(self.client.session)
|
||||
self.data["email"] = "test1@qduoj.com"
|
||||
response = self.client.post(self.register_url, data=self.data)
|
||||
self.assertDictEqual(response.data, {"error": "error", "data": "Username already exists"})
|
||||
|
||||
def test_email_already_exists(self):
|
||||
self.test_register_with_correct_info()
|
||||
|
||||
self.data["captcha"] = self._set_captcha(self.client.session)
|
||||
self.data["username"] = "test_user1"
|
||||
response = self.client.post(self.register_url, data=self.data)
|
||||
self.assertDictEqual(response.data, {"error": "error", "data": "Email already exists"})
|
||||
|
||||
|
||||
class SessionManagementAPITest(APITestCase):
|
||||
def setUp(self):
|
||||
self.create_user("test", "test123")
|
||||
self.url = self.reverse("session_management_api")
|
||||
# launch a request to provide session data
|
||||
login_url = self.reverse("user_login_api")
|
||||
self.client.post(login_url, data={"username": "test", "password": "test123"})
|
||||
|
||||
def test_get_sessions(self):
|
||||
resp = self.client.get(self.url)
|
||||
self.assertSuccess(resp)
|
||||
data = resp.data["data"]
|
||||
self.assertEqual(len(data), 1)
|
||||
|
||||
# def test_delete_session_key(self):
|
||||
# resp = self.client.delete(self.url + "?session_key=" + self.session_key)
|
||||
# self.assertSuccess(resp)
|
||||
|
||||
def test_delete_session_with_invalid_key(self):
|
||||
resp = self.client.delete(self.url + "?session_key=aaaaaaaaaa")
|
||||
self.assertDictEqual(resp.data, {"error": "error", "data": "Invalid session_key"})
|
||||
|
||||
|
||||
class UserProfileAPITest(APITestCase):
|
||||
def setUp(self):
|
||||
self.url = self.reverse("user_profile_api")
|
||||
|
||||
def test_get_profile_without_login(self):
|
||||
resp = self.client.get(self.url)
|
||||
self.assertDictEqual(resp.data, {"error": None, "data": None})
|
||||
|
||||
def test_get_profile(self):
|
||||
self.create_user("test", "test123")
|
||||
resp = self.client.get(self.url)
|
||||
self.assertSuccess(resp)
|
||||
|
||||
def test_update_profile(self):
|
||||
self.create_user("test", "test123")
|
||||
update_data = {"real_name": "zemal", "submission_number": 233, "language": "en-US"}
|
||||
resp = self.client.put(self.url, data=update_data)
|
||||
self.assertSuccess(resp)
|
||||
data = resp.data["data"]
|
||||
self.assertEqual(data["real_name"], "zemal")
|
||||
self.assertEqual(data["submission_number"], 0)
|
||||
self.assertEqual(data["language"], "en-US")
|
||||
|
||||
|
||||
class TwoFactorAuthAPITest(APITestCase):
|
||||
def setUp(self):
|
||||
self.url = self.reverse("two_factor_auth_api")
|
||||
self.create_user("test", "test123")
|
||||
|
||||
def _get_tfa_code(self):
|
||||
user = User.objects.first()
|
||||
code = OtpAuth(user.tfa_token).totp()
|
||||
if len(str(code)) < 6:
|
||||
code = (6 - len(str(code))) * "0" + str(code)
|
||||
return code
|
||||
|
||||
def test_get_image(self):
|
||||
resp = self.client.get(self.url)
|
||||
self.assertSuccess(resp)
|
||||
|
||||
def test_open_tfa_with_invalid_code(self):
|
||||
self.test_get_image()
|
||||
resp = self.client.post(self.url, data={"code": "000000"})
|
||||
self.assertDictEqual(resp.data, {"error": "error", "data": "Invalid code"})
|
||||
|
||||
def test_open_tfa_with_correct_code(self):
|
||||
self.test_get_image()
|
||||
code = self._get_tfa_code()
|
||||
resp = self.client.post(self.url, data={"code": code})
|
||||
self.assertSuccess(resp)
|
||||
user = User.objects.first()
|
||||
self.assertEqual(user.two_factor_auth, True)
|
||||
|
||||
def test_close_tfa_with_invalid_code(self):
|
||||
self.test_open_tfa_with_correct_code()
|
||||
resp = self.client.post(self.url, data={"code": "000000"})
|
||||
self.assertDictEqual(resp.data, {"error": "error", "data": "Invalid code"})
|
||||
|
||||
def test_close_tfa_with_correct_code(self):
|
||||
self.test_open_tfa_with_correct_code()
|
||||
code = self._get_tfa_code()
|
||||
resp = self.client.put(self.url, data={"code": code})
|
||||
self.assertSuccess(resp)
|
||||
user = User.objects.first()
|
||||
self.assertEqual(user.two_factor_auth, False)
|
||||
|
||||
|
||||
@mock.patch("account.views.oj.send_email_async.send")
|
||||
class ApplyResetPasswordAPITest(CaptchaTest):
|
||||
def setUp(self):
|
||||
self.create_user("test", "test123", login=False)
|
||||
user = User.objects.first()
|
||||
user.email = "test@oj.com"
|
||||
user.save()
|
||||
self.url = self.reverse("apply_reset_password_api")
|
||||
self.data = {"email": "test@oj.com", "captcha": self._set_captcha(self.client.session)}
|
||||
|
||||
def _refresh_captcha(self):
|
||||
self.data["captcha"] = self._set_captcha(self.client.session)
|
||||
|
||||
def test_apply_reset_password(self, send_email_send):
|
||||
resp = self.client.post(self.url, data=self.data)
|
||||
self.assertSuccess(resp)
|
||||
send_email_send.assert_called()
|
||||
|
||||
def test_apply_reset_password_twice_in_20_mins(self, send_email_send):
|
||||
self.test_apply_reset_password()
|
||||
send_email_send.reset_mock()
|
||||
self._refresh_captcha()
|
||||
resp = self.client.post(self.url, data=self.data)
|
||||
self.assertDictEqual(resp.data, {"error": "error", "data": "You can only reset password once per 20 minutes"})
|
||||
send_email_send.assert_not_called()
|
||||
|
||||
def test_apply_reset_password_again_after_20_mins(self, send_email_send):
|
||||
self.test_apply_reset_password()
|
||||
user = User.objects.first()
|
||||
user.reset_password_token_expire_time = now() - timedelta(minutes=21)
|
||||
user.save()
|
||||
self._refresh_captcha()
|
||||
self.test_apply_reset_password()
|
||||
|
||||
|
||||
class ResetPasswordAPITest(CaptchaTest):
|
||||
def setUp(self):
|
||||
self.create_user("test", "test123", login=False)
|
||||
self.url = self.reverse("reset_password_api")
|
||||
user = User.objects.first()
|
||||
user.reset_password_token = "online_judge?"
|
||||
user.reset_password_token_expire_time = now() + timedelta(minutes=20)
|
||||
user.save()
|
||||
self.data = {"token": user.reset_password_token,
|
||||
"captcha": self._set_captcha(self.client.session),
|
||||
"password": "test456"}
|
||||
|
||||
def test_reset_password_with_correct_token(self):
|
||||
resp = self.client.post(self.url, data=self.data)
|
||||
self.assertSuccess(resp)
|
||||
self.assertTrue(self.client.login(username="test", password="test456"))
|
||||
|
||||
def test_reset_password_with_invalid_token(self):
|
||||
self.data["token"] = "aaaaaaaaaaa"
|
||||
resp = self.client.post(self.url, data=self.data)
|
||||
self.assertDictEqual(resp.data, {"error": "error", "data": "Token does not exist"})
|
||||
|
||||
def test_reset_password_with_expired_token(self):
|
||||
user = User.objects.first()
|
||||
user.reset_password_token_expire_time = now() - timedelta(seconds=30)
|
||||
user.save()
|
||||
resp = self.client.post(self.url, data=self.data)
|
||||
self.assertDictEqual(resp.data, {"error": "error", "data": "Token has expired"})
|
||||
|
||||
|
||||
class UserChangeEmailAPITest(APITestCase):
|
||||
def setUp(self):
|
||||
self.url = self.reverse("user_change_email_api")
|
||||
self.user = self.create_user("test", "test123")
|
||||
self.new_mail = "test@oj.com"
|
||||
self.data = {"password": "test123", "new_email": self.new_mail}
|
||||
|
||||
def test_change_email_success(self):
|
||||
resp = self.client.post(self.url, data=self.data)
|
||||
self.assertSuccess(resp)
|
||||
|
||||
def test_wrong_password(self):
|
||||
self.data["password"] = "aaaa"
|
||||
resp = self.client.post(self.url, data=self.data)
|
||||
self.assertDictEqual(resp.data, {"error": "error", "data": "Wrong password"})
|
||||
|
||||
def test_duplicate_email(self):
|
||||
u = self.create_user("aa", "bb", login=False)
|
||||
u.email = self.new_mail
|
||||
u.save()
|
||||
resp = self.client.post(self.url, data=self.data)
|
||||
self.assertDictEqual(resp.data, {"error": "error", "data": "The email is owned by other account"})
|
||||
|
||||
|
||||
class UserChangePasswordAPITest(APITestCase):
|
||||
def setUp(self):
|
||||
self.url = self.reverse("user_change_password_api")
|
||||
|
||||
# Create user at first
|
||||
self.username = "test_user"
|
||||
self.old_password = "testuserpassword"
|
||||
self.new_password = "new_password"
|
||||
self.user = self.create_user(username=self.username, password=self.old_password, login=False)
|
||||
|
||||
self.data = {"old_password": self.old_password, "new_password": self.new_password}
|
||||
|
||||
def _get_tfa_code(self):
|
||||
user = User.objects.first()
|
||||
code = OtpAuth(user.tfa_token).totp()
|
||||
if len(str(code)) < 6:
|
||||
code = (6 - len(str(code))) * "0" + str(code)
|
||||
return code
|
||||
|
||||
def test_login_required(self):
|
||||
response = self.client.post(self.url, data=self.data)
|
||||
self.assertEqual(response.data, {"error": "permission-denied", "data": "Please login first"})
|
||||
|
||||
def test_valid_ola_password(self):
|
||||
self.assertTrue(self.client.login(username=self.username, password=self.old_password))
|
||||
response = self.client.post(self.url, data=self.data)
|
||||
self.assertEqual(response.data, {"error": None, "data": "Succeeded"})
|
||||
self.assertTrue(self.client.login(username=self.username, password=self.new_password))
|
||||
|
||||
def test_invalid_old_password(self):
|
||||
self.assertTrue(self.client.login(username=self.username, password=self.old_password))
|
||||
self.data["old_password"] = "invalid"
|
||||
response = self.client.post(self.url, data=self.data)
|
||||
self.assertEqual(response.data, {"error": "error", "data": "Invalid old password"})
|
||||
|
||||
def test_tfa_code_required(self):
|
||||
self.user.two_factor_auth = True
|
||||
self.user.tfa_token = "tfa_token"
|
||||
self.user.save()
|
||||
self.assertTrue(self.client.login(username=self.username, password=self.old_password))
|
||||
self.data["tfa_code"] = rand_str(6)
|
||||
resp = self.client.post(self.url, data=self.data)
|
||||
self.assertEqual(resp.data, {"error": "error", "data": "Invalid two factor verification code"})
|
||||
|
||||
self.data["tfa_code"] = self._get_tfa_code()
|
||||
resp = self.client.post(self.url, data=self.data)
|
||||
self.assertSuccess(resp)
|
||||
|
||||
|
||||
class UserRankAPITest(APITestCase):
|
||||
def setUp(self):
|
||||
self.url = self.reverse("user_rank_api")
|
||||
self.create_user("test1", "test123", login=False)
|
||||
self.create_user("test2", "test123", login=False)
|
||||
test1 = User.objects.get(username="test1")
|
||||
profile1 = test1.userprofile
|
||||
profile1.submission_number = 10
|
||||
profile1.accepted_number = 10
|
||||
profile1.total_score = 240
|
||||
profile1.save()
|
||||
|
||||
test2 = User.objects.get(username="test2")
|
||||
profile2 = test2.userprofile
|
||||
profile2.submission_number = 15
|
||||
profile2.accepted_number = 10
|
||||
profile2.total_score = 700
|
||||
profile2.save()
|
||||
|
||||
def test_get_acm_rank(self):
|
||||
resp = self.client.get(self.url, data={"rule": ContestRuleType.ACM})
|
||||
self.assertSuccess(resp)
|
||||
data = resp.data["data"]["results"]
|
||||
self.assertEqual(data[0]["user"]["username"], "test1")
|
||||
self.assertEqual(data[1]["user"]["username"], "test2")
|
||||
|
||||
def test_get_oi_rank(self):
|
||||
resp = self.client.get(self.url, data={"rule": ContestRuleType.OI})
|
||||
self.assertSuccess(resp)
|
||||
data = resp.data["data"]["results"]
|
||||
self.assertEqual(data[0]["user"]["username"], "test2")
|
||||
self.assertEqual(data[1]["user"]["username"], "test1")
|
||||
|
||||
def test_admin_role_filted(self):
|
||||
self.create_admin("admin", "admin123")
|
||||
admin = User.objects.get(username="admin")
|
||||
profile1 = admin.userprofile
|
||||
profile1.submission_number = 20
|
||||
profile1.accepted_number = 5
|
||||
profile1.total_score = 300
|
||||
profile1.save()
|
||||
resp = self.client.get(self.url, data={"rule": ContestRuleType.ACM})
|
||||
self.assertSuccess(resp)
|
||||
self.assertEqual(len(resp.data["data"]), 2)
|
||||
|
||||
resp = self.client.get(self.url, data={"rule": ContestRuleType.OI})
|
||||
self.assertSuccess(resp)
|
||||
self.assertEqual(len(resp.data["data"]), 2)
|
||||
|
||||
|
||||
class ProfileProblemDisplayIDRefreshAPITest(APITestCase):
|
||||
def setUp(self):
|
||||
pass
|
||||
|
||||
|
||||
class AdminUserTest(APITestCase):
|
||||
def setUp(self):
|
||||
self.user = self.create_super_admin(login=True)
|
||||
self.username = self.password = "test"
|
||||
self.regular_user = self.create_user(username=self.username, password=self.password, login=False)
|
||||
self.url = self.reverse("user_admin_api")
|
||||
self.data = {"id": self.regular_user.id, "username": self.username, "real_name": "test_name",
|
||||
"email": "test@qq.com", "admin_type": AdminType.REGULAR_USER,
|
||||
"problem_permission": ProblemPermission.OWN, "open_api": True,
|
||||
"two_factor_auth": False, "is_disabled": False}
|
||||
|
||||
def test_user_list(self):
|
||||
response = self.client.get(self.url)
|
||||
self.assertSuccess(response)
|
||||
|
||||
def test_edit_user_successfully(self):
|
||||
response = self.client.put(self.url, data=self.data)
|
||||
self.assertSuccess(response)
|
||||
resp_data = response.data["data"]
|
||||
self.assertEqual(resp_data["username"], self.username)
|
||||
self.assertEqual(resp_data["email"], "test@qq.com")
|
||||
self.assertEqual(resp_data["open_api"], True)
|
||||
self.assertEqual(resp_data["two_factor_auth"], False)
|
||||
self.assertEqual(resp_data["is_disabled"], False)
|
||||
self.assertEqual(resp_data["problem_permission"], ProblemPermission.NONE)
|
||||
|
||||
self.assertTrue(self.regular_user.check_password("test"))
|
||||
|
||||
def test_edit_user_password(self):
|
||||
data = self.data
|
||||
new_password = "testpassword"
|
||||
data["password"] = new_password
|
||||
response = self.client.put(self.url, data=data)
|
||||
self.assertSuccess(response)
|
||||
user = User.objects.get(id=self.regular_user.id)
|
||||
self.assertFalse(user.check_password(self.password))
|
||||
self.assertTrue(user.check_password(new_password))
|
||||
|
||||
def test_edit_user_tfa(self):
|
||||
data = self.data
|
||||
self.assertIsNone(self.regular_user.tfa_token)
|
||||
data["two_factor_auth"] = True
|
||||
response = self.client.put(self.url, data=data)
|
||||
self.assertSuccess(response)
|
||||
resp_data = response.data["data"]
|
||||
# if `tfa_token` is None, a new value will be generated
|
||||
self.assertTrue(resp_data["two_factor_auth"])
|
||||
token = User.objects.get(id=self.regular_user.id).tfa_token
|
||||
self.assertIsNotNone(token)
|
||||
|
||||
response = self.client.put(self.url, data=data)
|
||||
self.assertSuccess(response)
|
||||
resp_data = response.data["data"]
|
||||
# if `tfa_token` is not None, the value is not changed
|
||||
self.assertTrue(resp_data["two_factor_auth"])
|
||||
self.assertEqual(User.objects.get(id=self.regular_user.id).tfa_token, token)
|
||||
|
||||
def test_edit_user_openapi(self):
|
||||
data = self.data
|
||||
self.assertIsNone(self.regular_user.open_api_appkey)
|
||||
data["open_api"] = True
|
||||
response = self.client.put(self.url, data=data)
|
||||
self.assertSuccess(response)
|
||||
resp_data = response.data["data"]
|
||||
# if `open_api_appkey` is None, a new value will be generated
|
||||
self.assertTrue(resp_data["open_api"])
|
||||
key = User.objects.get(id=self.regular_user.id).open_api_appkey
|
||||
self.assertIsNotNone(key)
|
||||
|
||||
response = self.client.put(self.url, data=data)
|
||||
self.assertSuccess(response)
|
||||
resp_data = response.data["data"]
|
||||
# if `openapi_app_key` is not None, the value is not changed
|
||||
self.assertTrue(resp_data["open_api"])
|
||||
self.assertEqual(User.objects.get(id=self.regular_user.id).open_api_appkey, key)
|
||||
|
||||
def test_import_users(self):
|
||||
data = {"users": [["user1", "pass1", "eami1@e.com", "user1"],
|
||||
["user2", "pass3", "eamil3@e.com", "user2"]]
|
||||
}
|
||||
resp = self.client.post(self.url, data)
|
||||
self.assertSuccess(resp)
|
||||
# successfully created 2 users
|
||||
self.assertEqual(User.objects.all().count(), 4)
|
||||
|
||||
def test_import_duplicate_user(self):
|
||||
data = {"users": [["user1", "pass1", "eami1@e.com", "user1"],
|
||||
["user1", "pass1", "eami1@e.com", "user1"]]
|
||||
}
|
||||
resp = self.client.post(self.url, data)
|
||||
self.assertFailed(resp, "DETAIL: Key (username)=(user1) already exists.")
|
||||
# no user is created
|
||||
self.assertEqual(User.objects.all().count(), 2)
|
||||
|
||||
def test_delete_users(self):
|
||||
self.test_import_users()
|
||||
user_ids = User.objects.filter(username__in=["user1", "user2"]).values_list("id", flat=True)
|
||||
user_ids = ",".join([str(id) for id in user_ids])
|
||||
resp = self.client.delete(self.url + "?id=" + user_ids)
|
||||
self.assertSuccess(resp)
|
||||
self.assertEqual(User.objects.all().count(), 2)
|
||||
|
||||
|
||||
class GenerateUserAPITest(APITestCase):
|
||||
def setUp(self):
|
||||
self.create_super_admin()
|
||||
self.url = self.reverse("generate_user_api")
|
||||
self.data = {
|
||||
"number_from": 100, "number_to": 105,
|
||||
"prefix": "pre", "suffix": "suf",
|
||||
"default_email": "test@test.com",
|
||||
"password_length": 8
|
||||
}
|
||||
|
||||
def test_error_case(self):
|
||||
data = deepcopy(self.data)
|
||||
data["prefix"] = "t" * 16
|
||||
data["suffix"] = "s" * 14
|
||||
resp = self.client.post(self.url, data=data)
|
||||
self.assertEqual(resp.data["data"], "Username should not more than 32 characters")
|
||||
|
||||
data2 = deepcopy(self.data)
|
||||
data2["number_from"] = 106
|
||||
resp = self.client.post(self.url, data=data2)
|
||||
self.assertEqual(resp.data["data"], "Start number must be lower than end number")
|
||||
|
||||
@mock.patch("account.views.admin.xlsxwriter.Workbook")
|
||||
def test_generate_user_success(self, mock_workbook):
|
||||
resp = self.client.post(self.url, data=self.data)
|
||||
self.assertSuccess(resp)
|
||||
mock_workbook.assert_called()
|
||||
|
||||
|
||||
class OpenAPIAppkeyAPITest(APITestCase):
|
||||
def setUp(self):
|
||||
self.user = self.create_super_admin()
|
||||
self.url = self.reverse("open_api_appkey_api")
|
||||
|
||||
def test_reset_appkey(self):
|
||||
resp = self.client.post(self.url, data={})
|
||||
self.assertFailed(resp)
|
||||
|
||||
self.user.open_api = True
|
||||
self.user.save()
|
||||
resp = self.client.post(self.url, data={})
|
||||
self.assertSuccess(resp)
|
||||
self.assertEqual(resp.data["data"]["appkey"], User.objects.get(username=self.user.username).open_api_appkey)
|
||||
@@ -1,8 +1,8 @@
|
||||
from django.urls import path
|
||||
|
||||
from ..views.admin import ResetUserPasswordAPI, UserAdminAPI
|
||||
from ..views.admin import UserAdminAPI, GenerateUserAPI
|
||||
|
||||
urlpatterns = [
|
||||
path("user", UserAdminAPI.as_view()),
|
||||
path("reset_password", ResetUserPasswordAPI.as_view()),
|
||||
path("generate_user", GenerateUserAPI.as_view()),
|
||||
]
|
||||
|
||||
@@ -1,27 +1,54 @@
|
||||
from django.urls import path
|
||||
|
||||
from ..views.oj import (
|
||||
AvatarUploadAPI,
|
||||
ApplyResetPasswordAPI,
|
||||
ResetPasswordAPI,
|
||||
UserChangePasswordAPI,
|
||||
Metrics,
|
||||
ProfileProblemDisplayIDRefreshAPI,
|
||||
UserActivityRankAPI,
|
||||
UserRegisterAPI,
|
||||
UserChangeEmailAPI,
|
||||
UserLoginAPI,
|
||||
UserLogoutAPI,
|
||||
UserProblemRankAPI,
|
||||
UsernameOrEmailCheck,
|
||||
AvatarUploadAPI,
|
||||
TwoFactorAuthAPI,
|
||||
UserProfileAPI,
|
||||
UserRankAPI,
|
||||
UserRegisterAPI,
|
||||
UserActivityRankAPI,
|
||||
CheckTFARequiredAPI,
|
||||
SessionManagementAPI,
|
||||
ProfileProblemDisplayIDRefreshAPI,
|
||||
OpenAPIAppkeyAPI,
|
||||
SSOAPI,
|
||||
)
|
||||
|
||||
from utils.captcha.views import CaptchaAPIView
|
||||
|
||||
urlpatterns = [
|
||||
path("login", UserLoginAPI.as_view()),
|
||||
path("logout", UserLogoutAPI.as_view()),
|
||||
path("register", UserRegisterAPI.as_view()),
|
||||
path("change_password", UserChangePasswordAPI.as_view()),
|
||||
path("change_email", UserChangeEmailAPI.as_view()),
|
||||
path("apply_reset_password", ApplyResetPasswordAPI.as_view()),
|
||||
path("reset_password", ResetPasswordAPI.as_view()),
|
||||
path("captcha", CaptchaAPIView.as_view()),
|
||||
path("check_username_or_email", UsernameOrEmailCheck.as_view()),
|
||||
path("profile", UserProfileAPI.as_view(), name="user_profile_api"),
|
||||
path("profile/fresh_display_id", ProfileProblemDisplayIDRefreshAPI.as_view()),
|
||||
path("metrics", Metrics.as_view()),
|
||||
path("upload_avatar", AvatarUploadAPI.as_view()),
|
||||
path("tfa_required", CheckTFARequiredAPI.as_view()),
|
||||
path(
|
||||
"two_factor_auth",
|
||||
TwoFactorAuthAPI.as_view(),
|
||||
),
|
||||
path("user_rank", UserRankAPI.as_view()),
|
||||
path("user_activity_rank", UserActivityRankAPI.as_view()),
|
||||
path("user_problem_rank", UserProblemRankAPI.as_view()),
|
||||
path("sessions", SessionManagementAPI.as_view()),
|
||||
path(
|
||||
"open_api_appkey",
|
||||
OpenAPIAppkeyAPI.as_view(),
|
||||
),
|
||||
path("sso", SSOAPI.as_view()),
|
||||
]
|
||||
|
||||
@@ -1,37 +1,24 @@
|
||||
import os
|
||||
import re
|
||||
import xlsxwriter
|
||||
|
||||
from django.db import transaction, IntegrityError
|
||||
from django.db.models import Q
|
||||
from django.http import HttpResponse
|
||||
from django.contrib.auth.hashers import make_password
|
||||
from django.db import IntegrityError, transaction
|
||||
from django.db.models import F, Q
|
||||
from django.utils.crypto import get_random_string
|
||||
|
||||
from submission.models import Submission
|
||||
from utils.api import APIError, APIView, validate_serializer
|
||||
from utils.shortcuts import CLASS_NAME_MAX_DIGITS, CLASS_NAME_MIN_DIGITS, is_valid_class_name, rand_str
|
||||
from utils.api import APIView, validate_serializer
|
||||
from utils.shortcuts import rand_str
|
||||
|
||||
from ..decorators import super_admin_required
|
||||
from ..models import AdminType, ProblemPermission, User, UserProfile
|
||||
from ..serializers import (
|
||||
EditUserSerializer,
|
||||
ImportUserSerializer,
|
||||
UserAdminSerializer,
|
||||
GenerateUserSerializer,
|
||||
)
|
||||
|
||||
|
||||
# ks251XXX 或者 ks2510XX 返回 251 或者 2510。
|
||||
# 不以 ks+数字 开头的(管理员、教师账号)返回 None。
|
||||
# 位数不对就直接报错,不猜——猜错会把 class_name 存歪,
|
||||
# 而剥前缀显示姓名、班级下拉、统计页都依赖它准确。
|
||||
# 这里先用 \d+ 抓全再判位数,不能直接用 CLASS_NAME_RE 匹配:
|
||||
# 那样 ks251001 会"匹配成功"并悄悄取前 4 位,正是要避免的猜测。
|
||||
def get_class_name(username):
|
||||
result = re.match(r"ks(\d+)", username)
|
||||
if not result:
|
||||
return None
|
||||
class_name = result.group(1)
|
||||
if not is_valid_class_name(class_name):
|
||||
raise APIError(f"用户名 {username} 的班级号 {class_name} 是 {len(class_name)} 位,必须是 {CLASS_NAME_MIN_DIGITS}~{CLASS_NAME_MAX_DIGITS} 位数字")
|
||||
return class_name
|
||||
from ..serializers import ImportUserSerializer
|
||||
|
||||
|
||||
class UserAdminAPI(APIView):
|
||||
@@ -53,14 +40,18 @@ class UserAdminAPI(APIView):
|
||||
password=make_password(user_data[1]),
|
||||
email=user_data[2],
|
||||
raw_password=user_data[1],
|
||||
class_name=get_class_name(user_data[0]),
|
||||
)
|
||||
)
|
||||
|
||||
try:
|
||||
with transaction.atomic():
|
||||
ret = User.objects.bulk_create(user_list)
|
||||
UserProfile.objects.bulk_create([UserProfile(user=ret[i], real_name=data[i][3]) for i in range(len(ret))])
|
||||
UserProfile.objects.bulk_create(
|
||||
[
|
||||
UserProfile(user=ret[i], real_name=data[i][3])
|
||||
for i in range(len(ret))
|
||||
]
|
||||
)
|
||||
return self.success()
|
||||
except IntegrityError as e:
|
||||
# Extract detail from exception message
|
||||
@@ -79,22 +70,27 @@ class UserAdminAPI(APIView):
|
||||
user = User.objects.get(id=data["id"])
|
||||
except User.DoesNotExist:
|
||||
return self.error("User does not exist")
|
||||
if User.objects.filter(username=data["username"].lower()).exclude(id=user.id).exists():
|
||||
if (
|
||||
User.objects.filter(username=data["username"].lower())
|
||||
.exclude(id=user.id)
|
||||
.exists()
|
||||
):
|
||||
return self.error("Username already exists")
|
||||
if User.objects.filter(email=data["email"].lower()).exclude(id=user.id).exists():
|
||||
if (
|
||||
User.objects.filter(email=data["email"].lower())
|
||||
.exclude(id=user.id)
|
||||
.exists()
|
||||
):
|
||||
return self.error("Email already exists")
|
||||
|
||||
pre_username = user.username
|
||||
user.username = data["username"].lower()
|
||||
user.class_name = get_class_name(data["username"])
|
||||
user.email = data["email"].lower()
|
||||
user.admin_type = data["admin_type"]
|
||||
user.is_disabled = data["is_disabled"]
|
||||
|
||||
if data["admin_type"] == AdminType.STUDENT_ADMIN:
|
||||
user.problem_permission = data["problem_permission"] or ProblemPermission.OWN
|
||||
elif data["admin_type"] == AdminType.TEACHER_ADMIN:
|
||||
user.problem_permission = data["problem_permission"] or ProblemPermission.OWN
|
||||
if data["admin_type"] == AdminType.ADMIN:
|
||||
user.problem_permission = data["problem_permission"]
|
||||
elif data["admin_type"] == AdminType.SUPER_ADMIN:
|
||||
user.problem_permission = ProblemPermission.ALL
|
||||
else:
|
||||
@@ -111,9 +107,20 @@ class UserAdminAPI(APIView):
|
||||
user.open_api_appkey = None
|
||||
user.open_api = data["open_api"]
|
||||
|
||||
if data["two_factor_auth"]:
|
||||
# Avoid reset user tfa_token after saving changes
|
||||
if not user.two_factor_auth:
|
||||
user.tfa_token = rand_str()
|
||||
else:
|
||||
user.tfa_token = None
|
||||
|
||||
user.two_factor_auth = data["two_factor_auth"]
|
||||
|
||||
user.save()
|
||||
if pre_username != user.username:
|
||||
Submission.objects.filter(username=pre_username).update(username=user.username)
|
||||
Submission.objects.filter(username=pre_username).update(
|
||||
username=user.username
|
||||
)
|
||||
|
||||
UserProfile.objects.filter(user=user).update(real_name=data["real_name"])
|
||||
return self.success(UserAdminSerializer(user).data)
|
||||
@@ -131,25 +138,20 @@ class UserAdminAPI(APIView):
|
||||
return self.error("User does not exist")
|
||||
return self.success(UserAdminSerializer(user).data)
|
||||
|
||||
# 获取排序参数
|
||||
order_by = request.GET.get("order_by", "")
|
||||
user = User.objects.all().order_by("-create_time")
|
||||
|
||||
# 根据排序参数设置排序规则
|
||||
if order_by == "-last_login":
|
||||
# 最近登录,将 None 值放在最后
|
||||
user = User.objects.all().order_by(F("last_login").desc(nulls_last=True))
|
||||
else:
|
||||
# 默认按创建时间倒序
|
||||
user = User.objects.all().order_by("-create_time")
|
||||
is_admin = request.GET.get("admin", "0")
|
||||
|
||||
type = request.GET.get("type", "")
|
||||
|
||||
if type:
|
||||
user = user.filter(admin_type=type)
|
||||
if is_admin == "1":
|
||||
user = user.exclude(admin_type=AdminType.REGULAR_USER)
|
||||
|
||||
keyword = request.GET.get("keyword", None)
|
||||
if keyword:
|
||||
user = user.filter(Q(username__icontains=keyword) | Q(userprofile__real_name__icontains=keyword) | Q(email__icontains=keyword))
|
||||
user = user.filter(
|
||||
Q(username__icontains=keyword)
|
||||
| Q(userprofile__real_name__icontains=keyword)
|
||||
| Q(email__icontains=keyword)
|
||||
)
|
||||
return self.success(self.paginate_data(request, user, UserAdminSerializer))
|
||||
|
||||
@super_admin_required
|
||||
@@ -164,25 +166,76 @@ class UserAdminAPI(APIView):
|
||||
return self.success()
|
||||
|
||||
|
||||
class ResetUserPasswordAPI(APIView):
|
||||
class GenerateUserAPI(APIView):
|
||||
@super_admin_required
|
||||
def get(self, request):
|
||||
"""
|
||||
download users excel
|
||||
"""
|
||||
file_id = request.GET.get("file_id")
|
||||
if not file_id:
|
||||
return self.error("Invalid Parameter, file_id is required")
|
||||
if not re.match(r"^[a-zA-Z0-9]+$", file_id):
|
||||
return self.error("Illegal file_id")
|
||||
file_path = f"/tmp/{file_id}.xlsx"
|
||||
if not os.path.isfile(file_path):
|
||||
return self.error("File does not exist")
|
||||
with open(file_path, "rb") as f:
|
||||
raw_data = f.read()
|
||||
os.remove(file_path)
|
||||
response = HttpResponse(raw_data)
|
||||
response["Content-Disposition"] = "attachment; filename=users.xlsx"
|
||||
response["Content-Type"] = "application/xlsx"
|
||||
return response
|
||||
|
||||
@validate_serializer(GenerateUserSerializer)
|
||||
@super_admin_required
|
||||
def post(self, request):
|
||||
"""
|
||||
重置用户密码为随机6位数字(不包括0)
|
||||
Generate User
|
||||
"""
|
||||
data = request.data
|
||||
user_id = data["id"]
|
||||
number_max_length = max(
|
||||
len(str(data["number_from"])), len(str(data["number_to"]))
|
||||
)
|
||||
if number_max_length + len(data["prefix"]) + len(data["suffix"]) > 32:
|
||||
return self.error("Username should not more than 32 characters")
|
||||
if data["number_from"] > data["number_to"]:
|
||||
return self.error("Start number must be lower than end number")
|
||||
|
||||
file_id = rand_str(8)
|
||||
filename = f"/tmp/{file_id}.xlsx"
|
||||
workbook = xlsxwriter.Workbook(filename)
|
||||
worksheet = workbook.add_worksheet()
|
||||
worksheet.set_column("A:B", 20)
|
||||
worksheet.write("A1", "Username")
|
||||
worksheet.write("B1", "Password")
|
||||
i = 1
|
||||
|
||||
user_list = []
|
||||
for number in range(data["number_from"], data["number_to"] + 1):
|
||||
raw_password = rand_str(data["password_length"])
|
||||
user = User(
|
||||
username=f"{data['prefix']}{number}{data['suffix']}",
|
||||
password=make_password(raw_password),
|
||||
)
|
||||
user.raw_password = raw_password
|
||||
user_list.append(user)
|
||||
|
||||
try:
|
||||
user = User.objects.get(id=user_id)
|
||||
except User.DoesNotExist:
|
||||
return self.error("User does not exist")
|
||||
|
||||
# 生成6位随机数字密码(不包括0)
|
||||
new_password = get_random_string(6, allowed_chars="123456789")
|
||||
|
||||
# 设置新密码
|
||||
user.set_password(new_password)
|
||||
user.save()
|
||||
|
||||
return self.success(new_password)
|
||||
with transaction.atomic():
|
||||
ret = User.objects.bulk_create(user_list)
|
||||
UserProfile.objects.bulk_create(
|
||||
[UserProfile(user=user) for user in ret]
|
||||
)
|
||||
for item in user_list:
|
||||
worksheet.write_string(i, 0, item.username)
|
||||
worksheet.write_string(i, 1, item.raw_password)
|
||||
i += 1
|
||||
workbook.close()
|
||||
return self.success({"file_id": file_id})
|
||||
except IntegrityError as e:
|
||||
# Extract detail from exception message
|
||||
# duplicate key value violates unique constraint "user_username_key"
|
||||
# DETAIL: Key (username)=(root11) already exists.
|
||||
return self.error(str(e).split("\n")[1])
|
||||
|
||||
@@ -1,36 +1,54 @@
|
||||
import asyncio
|
||||
import os
|
||||
from datetime import timedelta
|
||||
from importlib import import_module
|
||||
|
||||
from django.conf import settings
|
||||
from django.contrib import auth
|
||||
from django.template.loader import render_to_string
|
||||
from django.utils.decorators import method_decorator
|
||||
from django.utils.timezone import now
|
||||
from django.views.decorators.csrf import ensure_csrf_cookie, csrf_exempt
|
||||
from django.db.models import Count, Q
|
||||
from django.utils import timezone
|
||||
from django.utils.decorators import method_decorator
|
||||
from django.views.decorators.csrf import ensure_csrf_cookie
|
||||
|
||||
from options.options import SysOptions
|
||||
import qrcode
|
||||
from otpauth import TOTP
|
||||
|
||||
from problem.models import Problem
|
||||
from submission.models import JudgeStatus, Submission
|
||||
from utils.api import AsyncAPIView, validate_serializer
|
||||
from utils.async_helpers import async_cache_get, async_cache_set
|
||||
from utils.constants import CacheKey
|
||||
from utils.shortcuts import datetime2str, rand_str
|
||||
|
||||
from submission.models import Submission, JudgeStatus
|
||||
from utils.constants import ContestRuleType
|
||||
from options.options import SysOptions
|
||||
from utils.api import APIView, validate_serializer, CSRFExemptAPIView
|
||||
from utils.captcha import Captcha
|
||||
from utils.shortcuts import rand_str, img2base64, datetime2str
|
||||
from ..decorators import login_required
|
||||
from ..models import AdminType, User, UserProfile
|
||||
from ..models import User, UserProfile, AdminType
|
||||
from ..serializers import (
|
||||
ApplyResetPasswordSerializer,
|
||||
ResetPasswordSerializer,
|
||||
UserChangePasswordSerializer,
|
||||
UserLoginSerializer,
|
||||
UserRegisterSerializer,
|
||||
UsernameOrEmailCheckSerializer,
|
||||
RankInfoSerializer,
|
||||
UserChangeEmailSerializer,
|
||||
SSOSerializer,
|
||||
)
|
||||
from ..serializers import (
|
||||
TwoFactorAuthCodeSerializer,
|
||||
UserProfileSerializer,
|
||||
EditUserProfileSerializer,
|
||||
ImageUploadForm,
|
||||
RankInfoSerializer,
|
||||
UserLoginSerializer,
|
||||
UserProfileSerializer,
|
||||
UserRegisterSerializer,
|
||||
)
|
||||
from ..tasks import send_email_async
|
||||
|
||||
|
||||
class UserProfileAPI(AsyncAPIView):
|
||||
class UserProfileAPI(APIView):
|
||||
@method_decorator(ensure_csrf_cookie)
|
||||
async def get(self, request, **kwargs):
|
||||
def get(self, request, **kwargs):
|
||||
"""
|
||||
判断是否登录, 若登录返回用户信息
|
||||
"""
|
||||
user = request.user
|
||||
if not user.is_authenticated:
|
||||
return self.success()
|
||||
@@ -38,51 +56,56 @@ class UserProfileAPI(AsyncAPIView):
|
||||
username = request.GET.get("username")
|
||||
try:
|
||||
if username:
|
||||
user = await User.objects.aget(username=username, is_disabled=False)
|
||||
user = User.objects.get(username=username, is_disabled=False)
|
||||
else:
|
||||
user = request.user
|
||||
# api返回的是自己的信息,可以返real_name
|
||||
show_real_name = True
|
||||
except User.DoesNotExist:
|
||||
return self.error("User does not exist")
|
||||
profile = await UserProfile.objects.select_related("user").aget(user=user)
|
||||
return self.success(UserProfileSerializer(profile, show_real_name=show_real_name).data)
|
||||
return self.success(
|
||||
UserProfileSerializer(user.userprofile, show_real_name=show_real_name).data
|
||||
)
|
||||
|
||||
@login_required
|
||||
@validate_serializer(EditUserProfileSerializer)
|
||||
async def put(self, request):
|
||||
@login_required
|
||||
def put(self, request):
|
||||
data = request.data
|
||||
user_profile = await UserProfile.objects.select_related("user").aget(user=request.user)
|
||||
user_profile = request.user.userprofile
|
||||
for k, v in data.items():
|
||||
setattr(user_profile, k, v)
|
||||
await user_profile.asave()
|
||||
return self.success(UserProfileSerializer(user_profile, show_real_name=True).data)
|
||||
|
||||
|
||||
class Metrics(AsyncAPIView):
|
||||
async def get(self, request):
|
||||
userid = request.GET.get("userid")
|
||||
qs = Submission.objects.filter(user_id=userid, contest_id__isnull=True)
|
||||
count, latest, first = await asyncio.gather(
|
||||
qs.acount(),
|
||||
qs.order_by("-create_time").afirst(),
|
||||
qs.order_by("create_time").afirst(),
|
||||
)
|
||||
if count == 0 or not latest or not first:
|
||||
return self.error("暂无提交")
|
||||
user_profile.save()
|
||||
return self.success(
|
||||
{
|
||||
"now": datetime2str(timezone.now()),
|
||||
"latest": datetime2str(latest.create_time),
|
||||
"first": datetime2str(first.create_time),
|
||||
}
|
||||
UserProfileSerializer(user_profile, show_real_name=True).data
|
||||
)
|
||||
|
||||
|
||||
class AvatarUploadAPI(AsyncAPIView):
|
||||
class Metrics(APIView):
|
||||
def get(self, request):
|
||||
userid = request.GET.get("userid")
|
||||
submissions = Submission.objects.filter(user_id=userid, contest_id__isnull=True)
|
||||
if submissions.count() == 0:
|
||||
return self.error("暂无提交")
|
||||
else:
|
||||
latest_submission = submissions.first()
|
||||
last_submission = submissions.last()
|
||||
if last_submission and latest_submission:
|
||||
return self.success(
|
||||
{
|
||||
"now": datetime2str(timezone.now()),
|
||||
"latest": datetime2str(latest_submission.create_time),
|
||||
"first": datetime2str(last_submission.create_time),
|
||||
}
|
||||
)
|
||||
else:
|
||||
return self.error("暂无提交")
|
||||
|
||||
|
||||
class AvatarUploadAPI(APIView):
|
||||
request_parsers = ()
|
||||
|
||||
@login_required
|
||||
async def post(self, request):
|
||||
def post(self, request):
|
||||
form = ImageUploadForm(request.POST, request.FILES)
|
||||
if form.is_valid():
|
||||
avatar = form.cleaned_data["image"]
|
||||
@@ -98,156 +121,412 @@ class AvatarUploadAPI(AsyncAPIView):
|
||||
with open(os.path.join(settings.AVATAR_UPLOAD_DIR, name), "wb") as img:
|
||||
for chunk in avatar:
|
||||
img.write(chunk)
|
||||
user_profile = await UserProfile.objects.aget(user=request.user)
|
||||
user_profile = request.user.userprofile
|
||||
|
||||
user_profile.avatar = f"{settings.AVATAR_URI_PREFIX}/{name}"
|
||||
await user_profile.asave()
|
||||
user_profile.save()
|
||||
return self.success("Succeeded")
|
||||
|
||||
|
||||
class UserLoginAPI(AsyncAPIView):
|
||||
@validate_serializer(UserLoginSerializer)
|
||||
async def post(self, request):
|
||||
class TwoFactorAuthAPI(APIView):
|
||||
@login_required
|
||||
def get(self, request):
|
||||
"""
|
||||
Get QR code
|
||||
"""
|
||||
user = request.user
|
||||
if user.two_factor_auth:
|
||||
return self.error("2FA is already turned on")
|
||||
token = rand_str()
|
||||
user.tfa_token = token
|
||||
user.save()
|
||||
|
||||
label = f"{SysOptions.website_name_shortcut}:{user.username}"
|
||||
image = qrcode.make(
|
||||
TOTP(token).to_uri(
|
||||
"totp", label, SysOptions.website_name.replace(" ", "")
|
||||
)
|
||||
)
|
||||
return self.success(img2base64(image))
|
||||
|
||||
@login_required
|
||||
@validate_serializer(TwoFactorAuthCodeSerializer)
|
||||
def post(self, request):
|
||||
"""
|
||||
Open 2FA
|
||||
"""
|
||||
code = request.data["code"]
|
||||
user = request.user
|
||||
if TOTP(user.tfa_token).verify(code):
|
||||
user.two_factor_auth = True
|
||||
user.save()
|
||||
return self.success("Succeeded")
|
||||
else:
|
||||
return self.error("Invalid code")
|
||||
|
||||
@login_required
|
||||
@validate_serializer(TwoFactorAuthCodeSerializer)
|
||||
def put(self, request):
|
||||
code = request.data["code"]
|
||||
user = request.user
|
||||
if not user.two_factor_auth:
|
||||
return self.error("2FA is already turned off")
|
||||
if TOTP(user.tfa_token).verify(code):
|
||||
user.two_factor_auth = False
|
||||
user.save()
|
||||
return self.success("Succeeded")
|
||||
else:
|
||||
return self.error("Invalid code")
|
||||
|
||||
|
||||
class CheckTFARequiredAPI(APIView):
|
||||
@validate_serializer(UsernameOrEmailCheckSerializer)
|
||||
def post(self, request):
|
||||
"""
|
||||
Check TFA is required
|
||||
"""
|
||||
data = request.data
|
||||
user = await auth.aauthenticate(username=data["username"], password=data["password"])
|
||||
result = False
|
||||
if data.get("username"):
|
||||
try:
|
||||
user = User.objects.get(username=data["username"])
|
||||
result = user.two_factor_auth
|
||||
except User.DoesNotExist:
|
||||
pass
|
||||
return self.success({"result": result})
|
||||
|
||||
|
||||
class UserLoginAPI(APIView):
|
||||
@validate_serializer(UserLoginSerializer)
|
||||
def post(self, request):
|
||||
"""
|
||||
User login api
|
||||
"""
|
||||
data = request.data
|
||||
user = auth.authenticate(username=data["username"], password=data["password"])
|
||||
# None is returned if username or password is wrong
|
||||
if user:
|
||||
if user.is_disabled:
|
||||
return self.error("Your account has been disabled")
|
||||
prev_login = user.last_login
|
||||
await auth.alogin(request, user)
|
||||
request.session["prev_login"] = datetime2str(prev_login) if prev_login else ""
|
||||
return self.success("Succeeded")
|
||||
if not user.two_factor_auth:
|
||||
auth.login(request, user)
|
||||
return self.success("Succeeded")
|
||||
|
||||
# `tfa_code` not in post data
|
||||
if user.two_factor_auth and "tfa_code" not in data:
|
||||
return self.error("tfa_required")
|
||||
|
||||
if TOTP(user.tfa_token).verify(data["tfa_code"]):
|
||||
auth.login(request, user)
|
||||
return self.success("Succeeded")
|
||||
else:
|
||||
return self.error("Invalid two factor verification code")
|
||||
else:
|
||||
return self.error("Invalid username or password")
|
||||
|
||||
|
||||
class UserLogoutAPI(AsyncAPIView):
|
||||
async def get(self, request):
|
||||
await auth.alogout(request)
|
||||
class UserLogoutAPI(APIView):
|
||||
def get(self, request):
|
||||
auth.logout(request)
|
||||
return self.success()
|
||||
|
||||
|
||||
class UserRegisterAPI(AsyncAPIView):
|
||||
class UsernameOrEmailCheck(APIView):
|
||||
@validate_serializer(UsernameOrEmailCheckSerializer)
|
||||
def post(self, request):
|
||||
"""
|
||||
check username or email is duplicate
|
||||
"""
|
||||
data = request.data
|
||||
# True means already exist.
|
||||
result = {"username": False, "email": False}
|
||||
if data.get("username"):
|
||||
result["username"] = User.objects.filter(
|
||||
username=data["username"].lower()
|
||||
).exists()
|
||||
if data.get("email"):
|
||||
result["email"] = User.objects.filter(email=data["email"].lower()).exists()
|
||||
return self.success(result)
|
||||
|
||||
|
||||
class UserRegisterAPI(APIView):
|
||||
@validate_serializer(UserRegisterSerializer)
|
||||
async def post(self, request):
|
||||
if not await SysOptions.aget("allow_register"):
|
||||
def post(self, request):
|
||||
"""
|
||||
User register api
|
||||
"""
|
||||
if not SysOptions.allow_register:
|
||||
return self.error("Register function has been disabled by admin")
|
||||
|
||||
data = request.data
|
||||
data["username"] = data["username"].lower()
|
||||
data["email"] = data["email"].lower()
|
||||
if await User.objects.filter(username=data["username"]).aexists():
|
||||
captcha = Captcha(request)
|
||||
if not captcha.check(data["captcha"]):
|
||||
return self.error("Invalid captcha")
|
||||
if User.objects.filter(username=data["username"]).exists():
|
||||
return self.error("Username already exists")
|
||||
if await User.objects.filter(email=data["email"]).aexists():
|
||||
if User.objects.filter(email=data["email"]).exists():
|
||||
return self.error("Email already exists")
|
||||
user = await User.objects.acreate(username=data["username"], email=data["email"])
|
||||
user = User.objects.create(username=data["username"], email=data["email"])
|
||||
user.set_password(data["password"])
|
||||
await user.asave()
|
||||
await UserProfile.objects.acreate(user=user)
|
||||
user.save()
|
||||
UserProfile.objects.create(user=user)
|
||||
return self.success("Succeeded")
|
||||
|
||||
|
||||
class UserRankAPI(AsyncAPIView):
|
||||
async def get(self, request):
|
||||
class UserChangeEmailAPI(APIView):
|
||||
@validate_serializer(UserChangeEmailSerializer)
|
||||
@login_required
|
||||
def post(self, request):
|
||||
data = request.data
|
||||
user = auth.authenticate(
|
||||
username=request.user.username, password=data["password"]
|
||||
)
|
||||
if user:
|
||||
if user.two_factor_auth:
|
||||
if "tfa_code" not in data:
|
||||
return self.error("tfa_required")
|
||||
if not TOTP(user.tfa_token).verify(data["tfa_code"]):
|
||||
return self.error("Invalid two factor verification code")
|
||||
data["new_email"] = data["new_email"].lower()
|
||||
if User.objects.filter(email=data["new_email"]).exists():
|
||||
return self.error("The email is owned by other account")
|
||||
user.email = data["new_email"]
|
||||
user.save()
|
||||
return self.success("Succeeded")
|
||||
else:
|
||||
return self.error("Wrong password")
|
||||
|
||||
|
||||
class UserChangePasswordAPI(APIView):
|
||||
@validate_serializer(UserChangePasswordSerializer)
|
||||
@login_required
|
||||
def post(self, request):
|
||||
"""
|
||||
User change password api
|
||||
"""
|
||||
data = request.data
|
||||
username = request.user.username
|
||||
user = auth.authenticate(username=username, password=data["old_password"])
|
||||
if user:
|
||||
if user.two_factor_auth:
|
||||
if "tfa_code" not in data:
|
||||
return self.error("tfa_required")
|
||||
if not TOTP(user.tfa_token).verify(data["tfa_code"]):
|
||||
return self.error("Invalid two factor verification code")
|
||||
user.set_password(data["new_password"])
|
||||
user.save()
|
||||
return self.success("Succeeded")
|
||||
else:
|
||||
return self.error("Invalid old password")
|
||||
|
||||
|
||||
class ApplyResetPasswordAPI(APIView):
|
||||
@validate_serializer(ApplyResetPasswordSerializer)
|
||||
def post(self, request):
|
||||
if request.user.is_authenticated:
|
||||
return self.error("You have already logged in, are you kidding me? ")
|
||||
data = request.data
|
||||
captcha = Captcha(request)
|
||||
if not captcha.check(data["captcha"]):
|
||||
return self.error("Invalid captcha")
|
||||
try:
|
||||
user = User.objects.get(email__iexact=data["email"])
|
||||
except User.DoesNotExist:
|
||||
return self.error("User does not exist")
|
||||
if (
|
||||
user.reset_password_token_expire_time
|
||||
and 0
|
||||
< int((user.reset_password_token_expire_time - now()).total_seconds())
|
||||
< 20 * 60
|
||||
):
|
||||
return self.error("You can only reset password once per 20 minutes")
|
||||
user.reset_password_token = rand_str()
|
||||
user.reset_password_token_expire_time = now() + timedelta(minutes=20)
|
||||
user.save()
|
||||
render_data = {
|
||||
"username": user.username,
|
||||
"website_name": SysOptions.website_name,
|
||||
"link": f"{SysOptions.website_base_url}/reset-password/{user.reset_password_token}",
|
||||
}
|
||||
email_html = render_to_string("reset_password_email.html", render_data)
|
||||
send_email_async.send(
|
||||
from_name=SysOptions.website_name_shortcut,
|
||||
to_email=user.email,
|
||||
to_name=user.username,
|
||||
subject="Reset your password",
|
||||
content=email_html,
|
||||
)
|
||||
return self.success("Succeeded")
|
||||
|
||||
|
||||
class ResetPasswordAPI(APIView):
|
||||
@validate_serializer(ResetPasswordSerializer)
|
||||
def post(self, request):
|
||||
data = request.data
|
||||
captcha = Captcha(request)
|
||||
if not captcha.check(data["captcha"]):
|
||||
return self.error("Invalid captcha")
|
||||
try:
|
||||
user = User.objects.get(reset_password_token=data["token"])
|
||||
except User.DoesNotExist:
|
||||
return self.error("Token does not exist")
|
||||
if user.reset_password_token_expire_time < now():
|
||||
return self.error("Token has expired")
|
||||
user.reset_password_token = None
|
||||
user.two_factor_auth = False
|
||||
user.set_password(data["password"])
|
||||
user.save()
|
||||
return self.success("Succeeded")
|
||||
|
||||
|
||||
class SessionManagementAPI(APIView):
|
||||
@login_required
|
||||
def get(self, request):
|
||||
engine = import_module(settings.SESSION_ENGINE)
|
||||
session_store = engine.SessionStore
|
||||
current_session = request.session.session_key
|
||||
session_keys = request.user.session_keys
|
||||
result = []
|
||||
modified = False
|
||||
for key in session_keys[:]:
|
||||
session = session_store(key)
|
||||
# session does not exist or is expiry
|
||||
if not session._session:
|
||||
session_keys.remove(key)
|
||||
modified = True
|
||||
continue
|
||||
|
||||
s = {}
|
||||
if current_session == key:
|
||||
s["current_session"] = True
|
||||
s["ip"] = session["ip"]
|
||||
s["user_agent"] = session["user_agent"]
|
||||
s["last_activity"] = datetime2str(session["last_activity"])
|
||||
s["session_key"] = key
|
||||
result.append(s)
|
||||
if modified:
|
||||
request.user.save()
|
||||
return self.success(result)
|
||||
|
||||
@login_required
|
||||
def delete(self, request):
|
||||
session_key = request.GET.get("session_key")
|
||||
if not session_key:
|
||||
return self.error("Parameter Error")
|
||||
request.session.delete(session_key)
|
||||
if session_key in request.user.session_keys:
|
||||
request.user.session_keys.remove(session_key)
|
||||
request.user.save()
|
||||
return self.success("Succeeded")
|
||||
else:
|
||||
return self.error("Invalid session_key")
|
||||
|
||||
|
||||
class UserRankAPI(APIView):
|
||||
def get(self, request):
|
||||
rule_type = request.GET.get("rule")
|
||||
username = request.GET.get("username", "")
|
||||
try:
|
||||
n = int(request.GET.get("n", "0"))
|
||||
except ValueError:
|
||||
n = 0
|
||||
|
||||
profiles = (
|
||||
UserProfile.objects.filter(
|
||||
user__admin_type__in=[AdminType.REGULAR_USER, AdminType.STUDENT_ADMIN],
|
||||
user__is_disabled=False,
|
||||
user__username__icontains=username,
|
||||
if rule_type not in ContestRuleType.choices():
|
||||
rule_type = ContestRuleType.ACM
|
||||
profiles = UserProfile.objects.filter(
|
||||
user__admin_type=AdminType.REGULAR_USER,
|
||||
user__is_disabled=False,
|
||||
user__username__icontains=username,
|
||||
).select_related("user")
|
||||
if rule_type == ContestRuleType.ACM:
|
||||
profiles = profiles.filter(accepted_number__gte=0).order_by(
|
||||
"-accepted_number", "submission_number"
|
||||
)
|
||||
.select_related("user")
|
||||
.filter(accepted_number__gte=0)
|
||||
.order_by("-accepted_number", "submission_number")
|
||||
)
|
||||
else:
|
||||
profiles = profiles.filter(total_score__gt=0).order_by("-total_score")
|
||||
if n > 0:
|
||||
profiles = profiles[:n]
|
||||
return self.success(await self.async_paginate_data(request, profiles, RankInfoSerializer))
|
||||
return self.success(self.paginate_data(request, profiles, RankInfoSerializer))
|
||||
|
||||
|
||||
class UserActivityRankAPI(AsyncAPIView):
|
||||
async def get(self, request):
|
||||
class UserActivityRankAPI(APIView):
|
||||
def get(self, request):
|
||||
start = request.GET.get("start")
|
||||
if not start:
|
||||
return self.error("start time is required")
|
||||
cache_key = f"{CacheKey.user_activity_rank}:{start}"
|
||||
cached = await async_cache_get(cache_key)
|
||||
if cached is not None:
|
||||
return self.success(cached)
|
||||
|
||||
hidden_names = User.objects.filter(Q(admin_type=AdminType.SUPER_ADMIN) | Q(is_disabled=True)).values_list("username", flat=True)
|
||||
hidden_names = User.objects.filter(
|
||||
Q(admin_type=AdminType.SUPER_ADMIN)
|
||||
| Q(admin_type=AdminType.ADMIN)
|
||||
| Q(is_disabled=True)
|
||||
).values_list("username", flat=True)
|
||||
submissions = Submission.objects.filter(
|
||||
contest_id__isnull=True,
|
||||
create_time__gte=start,
|
||||
result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED],
|
||||
).exclude(username__in=hidden_names)
|
||||
data = [row async for row in submissions.values("username").annotate(count=Count("problem_id", distinct=True)).order_by("-count")[:10]]
|
||||
await async_cache_set(cache_key, data, 600)
|
||||
return self.success(data)
|
||||
|
||||
|
||||
class UserProblemRankAPI(AsyncAPIView):
|
||||
async def get(self, request):
|
||||
problem_id = request.GET.get("problem_id")
|
||||
user = request.user
|
||||
if not user.is_authenticated:
|
||||
return self.error("User is not authenticated")
|
||||
|
||||
problem = await Problem.objects.aget(_id__iexact=problem_id, contest_id__isnull=True, visible=True)
|
||||
submissions = Submission.objects.filter(problem=problem, result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED])
|
||||
|
||||
all_ac_count = await submissions.values("user_id").distinct().acount()
|
||||
|
||||
class_name = user.class_name or ""
|
||||
class_ac_count = 0
|
||||
|
||||
if class_name:
|
||||
users = User.objects.filter(class_name=user.class_name, is_disabled=False).values_list("id", flat=True)
|
||||
user_ids = [user_id async for user_id in users]
|
||||
submissions = submissions.filter(user_id__in=user_ids)
|
||||
class_ac_count = await submissions.values("user_id").distinct().acount()
|
||||
|
||||
my_submissions = submissions.filter(user_id=user.id)
|
||||
|
||||
if not await my_submissions.aexists():
|
||||
return self.success(
|
||||
{
|
||||
"class_name": class_name,
|
||||
"rank": -1,
|
||||
"class_ac_count": class_ac_count,
|
||||
"all_ac_count": all_ac_count,
|
||||
}
|
||||
)
|
||||
|
||||
my_first_submission = await my_submissions.order_by("create_time").afirst()
|
||||
rank = await submissions.filter(create_time__lte=my_first_submission.create_time).acount()
|
||||
return self.success(
|
||||
{
|
||||
"class_name": class_name,
|
||||
"rank": rank,
|
||||
"class_ac_count": class_ac_count,
|
||||
"all_ac_count": all_ac_count,
|
||||
}
|
||||
contest_id__isnull=True, create_time__gte=start, result=JudgeStatus.ACCEPTED
|
||||
)
|
||||
counts = (
|
||||
submissions.values("username")
|
||||
.annotate(count=Count("problem_id", distinct=True))
|
||||
.order_by("-count")[: 10 + len(hidden_names)]
|
||||
)
|
||||
data = []
|
||||
for count in counts:
|
||||
if count["username"] not in hidden_names:
|
||||
data.append(count)
|
||||
return self.success(data[:10])
|
||||
|
||||
|
||||
class ProfileProblemDisplayIDRefreshAPI(AsyncAPIView):
|
||||
class ProfileProblemDisplayIDRefreshAPI(APIView):
|
||||
@login_required
|
||||
async def get(self, request):
|
||||
profile = await UserProfile.objects.aget(user=request.user)
|
||||
def get(self, request):
|
||||
profile = request.user.userprofile
|
||||
acm_problems = profile.acm_problems_status.get("problems", {})
|
||||
ids = list(acm_problems.keys())
|
||||
oi_problems = profile.oi_problems_status.get("problems", {})
|
||||
ids = list(acm_problems.keys()) + list(oi_problems.keys())
|
||||
if not ids:
|
||||
return self.success()
|
||||
display_ids = [did async for did in Problem.objects.filter(id__in=ids, visible=True).values_list("_id", flat=True)]
|
||||
display_ids = Problem.objects.filter(id__in=ids, visible=True).values_list(
|
||||
"_id", flat=True
|
||||
)
|
||||
id_map = dict(zip(ids, display_ids))
|
||||
for k, v in acm_problems.items():
|
||||
v["_id"] = id_map[k]
|
||||
await profile.asave(update_fields=["acm_problems_status"])
|
||||
for k, v in oi_problems.items():
|
||||
v["_id"] = id_map[k]
|
||||
profile.save(update_fields=["acm_problems_status", "oi_problems_status"])
|
||||
return self.success()
|
||||
|
||||
|
||||
class OpenAPIAppkeyAPI(APIView):
|
||||
@login_required
|
||||
def post(self, request):
|
||||
user = request.user
|
||||
if not user.open_api:
|
||||
return self.error("OpenAPI function is truned off for you")
|
||||
api_appkey = rand_str()
|
||||
user.open_api_appkey = api_appkey
|
||||
user.save()
|
||||
return self.success({"appkey": api_appkey})
|
||||
|
||||
|
||||
class SSOAPI(CSRFExemptAPIView):
|
||||
@login_required
|
||||
def get(self, request):
|
||||
token = rand_str()
|
||||
request.user.auth_token = token
|
||||
request.user.save()
|
||||
return self.success({"token": token})
|
||||
|
||||
@method_decorator(csrf_exempt)
|
||||
@validate_serializer(SSOSerializer)
|
||||
def post(self, request):
|
||||
try:
|
||||
user = User.objects.get(auth_token=request.data["token"])
|
||||
except User.DoesNotExist:
|
||||
return self.error("User does not exist")
|
||||
return self.success(
|
||||
{
|
||||
"username": user.username,
|
||||
"avatar": user.userprofile.avatar,
|
||||
"admin_type": user.admin_type,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -1,10 +0,0 @@
|
||||
from django.apps import AppConfig
|
||||
|
||||
|
||||
class AchievementConfig(AppConfig):
|
||||
default_auto_field = "django.db.models.BigAutoField"
|
||||
name = "achievement"
|
||||
|
||||
def ready(self):
|
||||
# 导入以触发 @metric 装饰器把指标注册进 METRIC_REGISTRY
|
||||
import achievement.metrics # noqa: F401
|
||||
@@ -1,99 +0,0 @@
|
||||
"""成就判定核心,被异步任务和管理命令共用。
|
||||
|
||||
判定刻意做成"一次查询取全部候选 + 内存比对",与 problemset 里逐条查询的
|
||||
旧写法相反:一次判题只多 3~4 条 SQL。
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from django.db import transaction
|
||||
from django.db.models import F
|
||||
|
||||
from achievement.metrics import META_METRICS, METRIC_REGISTRY, build_ctx
|
||||
from achievement.models import Achievement, Operator, UserAchievement, UserStat
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def evaluate(user, metrics, only_metrics=None):
|
||||
"""返回该用户应解锁但尚未解锁的成就列表。"""
|
||||
unlocked_ids = set(UserAchievement.objects.filter(user=user).values_list("achievement_id", flat=True))
|
||||
qs = Achievement.objects.filter(visible=True).exclude(id__in=unlocked_ids)
|
||||
if only_metrics is not None:
|
||||
qs = qs.filter(metric__in=only_metrics)
|
||||
|
||||
hits = []
|
||||
for achievement in qs:
|
||||
value = metrics.get(achievement.metric)
|
||||
# 指标从未产生有效值时 key 不存在,直接跳过:
|
||||
# 否则极小值型指标(求 min、配 lte 用的那种)会对新用户恒成立
|
||||
if value is None:
|
||||
continue
|
||||
if achievement.operator == Operator.GTE and value >= achievement.threshold:
|
||||
hits.append(achievement)
|
||||
elif achievement.operator == Operator.LTE and value <= achievement.threshold:
|
||||
hits.append(achievement)
|
||||
return hits
|
||||
|
||||
|
||||
def unlock(user, achievements, backfilled=False, notified=False):
|
||||
"""写入解锁记录并累加 unlock_count,返回实际新建的记录。
|
||||
|
||||
刻意逐条 get_or_create 而不是 bulk_create:unique_user_achievement 约束负责
|
||||
并发竞态,was_created 是"这一条确实是我新建的"的唯一可信判据。用
|
||||
bulk_create(ignore_conflicts=True) 则无法区分新建与已存在,并发判题时会把
|
||||
unlock_count 重复累加(获得率永久偏高),并对同一个奖杯重复推送通知。
|
||||
|
||||
循环次数是"本次新解锁的成就数",常态为 0,因此常态零查询。
|
||||
"""
|
||||
created = []
|
||||
for achievement in achievements:
|
||||
record, was_created = UserAchievement.objects.get_or_create(
|
||||
user=user,
|
||||
achievement=achievement,
|
||||
defaults={"backfilled": backfilled, "notified": notified},
|
||||
)
|
||||
if was_created:
|
||||
Achievement.objects.filter(id=achievement.id).update(unlock_count=F("unlock_count") + 1)
|
||||
created.append(record)
|
||||
return created
|
||||
|
||||
|
||||
def run_for_submission(user, submission):
|
||||
"""判题后的完整流程:更新指标 → 第一轮判定 → 元指标第二轮判定。"""
|
||||
ctx = build_ctx(user.id, submission)
|
||||
if ctx["skip"]:
|
||||
# 比赛提交不计入成就
|
||||
return []
|
||||
|
||||
with transaction.atomic():
|
||||
stat = UserStat.objects.select_for_update().get_or_create(user=user)[0]
|
||||
for key, m in METRIC_REGISTRY.items():
|
||||
if key in META_METRICS:
|
||||
continue
|
||||
m.on_submission(stat.metrics, submission, ctx)
|
||||
stat.save(update_fields=["metrics", "update_time"])
|
||||
|
||||
first = unlock(user, evaluate(user, stat.metrics))
|
||||
if not first:
|
||||
return []
|
||||
|
||||
# 第二轮:只重算元指标、只判定依赖元指标的成就,不再有第三轮
|
||||
meta_values = {key: METRIC_REGISTRY[key].recompute(user) for key in META_METRICS}
|
||||
|
||||
# 必须重新取锁并重新读一次 stat:上面那个 stat 对象的 metrics 是解锁前的快照,
|
||||
# 直接 save 会把整份字典写回,覆盖掉并发判题在这期间已提交的增量
|
||||
# (同一用户两次提交并发判题时会让提交数/AC 数静默倒退,且无定期重算兜底)。
|
||||
# 这里只合并元指标那几个 key。
|
||||
with transaction.atomic():
|
||||
stat = UserStat.objects.select_for_update().get(user=user)
|
||||
for key, value in meta_values.items():
|
||||
if value is None:
|
||||
stat.metrics.pop(key, None)
|
||||
else:
|
||||
stat.metrics[key] = value
|
||||
stat.save(update_fields=["metrics", "update_time"])
|
||||
|
||||
second = unlock(user, evaluate(user, stat.metrics, only_metrics=META_METRICS))
|
||||
|
||||
return first + second
|
||||
@@ -1,228 +0,0 @@
|
||||
"""成就系统部署后自检。
|
||||
|
||||
只读,不修改任何数据,可以反复跑。
|
||||
|
||||
这套系统开发全程没有数据库可用,所有后端逻辑都只经过纸面审查。本命令覆盖
|
||||
那些"错了也不报错、只是悄悄发错奖杯"的地方——每一项都对应一个已知的、
|
||||
静态审查抓不到的失败模式。
|
||||
|
||||
python manage.py check_achievement_deploy
|
||||
"""
|
||||
|
||||
from django.core.management.base import BaseCommand
|
||||
from django.db.models import IntegerField
|
||||
from django.db.models.fields.json import KeyTextTransform
|
||||
from django.db.models.functions import Cast
|
||||
|
||||
from account.models import User
|
||||
from achievement.metrics import META_METRICS, METRIC_REGISTRY
|
||||
from achievement.models import Achievement, Operator, UserAchievement, UserStat
|
||||
|
||||
# on_submission 依赖的辅助键 -> 它应该与哪个顶层指标保持一致
|
||||
STATE_PAIRS = {
|
||||
"_active_dates": "active_days",
|
||||
"_languages": "languages_used",
|
||||
}
|
||||
|
||||
|
||||
class Command(BaseCommand):
|
||||
help = "成就系统部署后自检(只读)"
|
||||
|
||||
def handle(self, *args, **options):
|
||||
self.failures = 0
|
||||
self.warnings = 0
|
||||
|
||||
self._check_migration()
|
||||
self._check_registry()
|
||||
self._check_min_metric_absent()
|
||||
self._check_jsonb_cast()
|
||||
self._check_recompute_state()
|
||||
self._check_unlock_count()
|
||||
self._check_dead_achievements()
|
||||
|
||||
self.stdout.write("")
|
||||
if self.failures:
|
||||
self.stdout.write(self.style.ERROR(f"{self.failures} 项未通过,先别配成就。"))
|
||||
elif self.warnings:
|
||||
self.stdout.write(self.style.WARNING(f"全部通过,{self.warnings} 项提醒。"))
|
||||
else:
|
||||
self.stdout.write(self.style.SUCCESS("全部通过。"))
|
||||
|
||||
# ---------- 输出 ----------
|
||||
|
||||
def _ok(self, title, detail=""):
|
||||
self.stdout.write(self.style.SUCCESS(f"[PASS] {title}") + (f" {detail}" if detail else ""))
|
||||
|
||||
def _fail(self, title, detail):
|
||||
self.failures += 1
|
||||
self.stdout.write(self.style.ERROR(f"[FAIL] {title}"))
|
||||
self.stdout.write(f" {detail}")
|
||||
|
||||
def _skip(self, title, detail):
|
||||
self.stdout.write(self.style.WARNING(f"[SKIP] {title}") + f" {detail}")
|
||||
|
||||
def _warn(self, title, detail):
|
||||
self.warnings += 1
|
||||
self.stdout.write(self.style.WARNING(f"[WARN] {title}"))
|
||||
self.stdout.write(f" {detail}")
|
||||
|
||||
# ---------- 检查项 ----------
|
||||
|
||||
def _check_migration(self):
|
||||
title = "迁移已应用"
|
||||
try:
|
||||
Achievement.objects.count()
|
||||
UserStat.objects.count()
|
||||
UserAchievement.objects.count()
|
||||
except Exception as e:
|
||||
self._fail(title, f"三张表至少有一张不存在:{e}\n 跑 python manage.py migrate achievement")
|
||||
return
|
||||
self._ok(title)
|
||||
|
||||
def _check_registry(self):
|
||||
title = "指标注册表已加载"
|
||||
count = len(METRIC_REGISTRY)
|
||||
if count == 0:
|
||||
self._fail(title, "METRIC_REGISTRY 是空的,AppConfig.ready() 没有导入 metrics")
|
||||
return
|
||||
self._ok(title, f"{count} 个指标,元指标 {sorted(META_METRICS)}")
|
||||
|
||||
def _check_min_metric_absent(self):
|
||||
"""lte 类成就用的指标,对没有 AC 记录的用户必须返回 None。
|
||||
|
||||
返回 0 的话,"最短 AC 代码 ≤ 50 字符"这类成就会白送给每一个从没做出过题的新生。
|
||||
线上目前没有 lte 成就(min_ac_code_chars 已删),本项因此会 SKIP;
|
||||
将来配了任何 lte 成就,它会自动开始检查对应的指标。
|
||||
"""
|
||||
title = "极小值指标对零 AC 用户返回 None"
|
||||
metric_keys = sorted(set(Achievement.objects.filter(operator=Operator.LTE, visible=True).values_list("metric", flat=True)))
|
||||
metrics = [(k, METRIC_REGISTRY[k]) for k in metric_keys if k in METRIC_REGISTRY]
|
||||
if not metrics:
|
||||
self._skip(title, "没有上架的 lte 类成就")
|
||||
return
|
||||
|
||||
user = User.objects.filter(is_disabled=False, userprofile__accepted_number=0).first()
|
||||
if user is None:
|
||||
self._skip(title, "找不到 accepted_number=0 的用户,无法验证")
|
||||
return
|
||||
|
||||
bad = [(key, value) for key, m in metrics if (value := m.recompute(user)) is not None]
|
||||
if bad:
|
||||
detail = "、".join(f"{key} 得到 {value!r}" for key, value in bad)
|
||||
self._fail(
|
||||
title,
|
||||
f"用户 {user.username}:{detail},都应为 None。\n 这些 lte 成就正在白送给全部零 AC 用户。",
|
||||
)
|
||||
else:
|
||||
self._ok(title, f"用户 {user.username},检查了 {len(metrics)} 个指标")
|
||||
|
||||
def _check_jsonb_cast(self):
|
||||
"""JSONB 数字必须按整数比较,不能按字符串序。
|
||||
|
||||
没有显式 Cast 时 Postgres 会认为 "9" > "50",后台调低阈值触发的补发
|
||||
会发给错误的人群,且不报任何错。
|
||||
"""
|
||||
title = "JSONB 阈值比较按整数而非字符串序"
|
||||
key = "submission_count"
|
||||
rows = list(UserStat.objects.filter(metrics__has_key=key).values_list("id", "metrics"))
|
||||
if len(rows) < 2:
|
||||
self._skip(title, f"有 {key} 的 UserStat 少于 2 条,无法验证")
|
||||
return
|
||||
|
||||
values = [(rid, m.get(key)) for rid, m in rows if isinstance(m.get(key), int)]
|
||||
if not values:
|
||||
self._skip(title, f"{key} 没有整数值")
|
||||
return
|
||||
|
||||
# 挑一个能区分整数序和字符串序的阈值:字符串比较下 "9" > "50"
|
||||
threshold = 9
|
||||
expected = {rid for rid, v in values if v >= threshold}
|
||||
|
||||
actual = set(
|
||||
UserStat.objects.filter(metrics__has_key=key).annotate(v=Cast(KeyTextTransform(key, "metrics"), IntegerField())).filter(v__gte=threshold).values_list("id", flat=True)
|
||||
)
|
||||
|
||||
if actual == expected:
|
||||
self._ok(title, f"阈值 {threshold} 命中 {len(actual)} 人,与 Python 侧一致")
|
||||
else:
|
||||
missing = len(expected - actual)
|
||||
extra = len(actual - expected)
|
||||
self._fail(
|
||||
title,
|
||||
f"数据库筛出 {len(actual)} 人,正确答案是 {len(expected)} 人(少 {missing}、多 {extra})。\n rescan_achievement 会把成就补发给错误的人群。",
|
||||
)
|
||||
|
||||
def _check_recompute_state(self):
|
||||
"""重算必须一并重建 on_submission 依赖的辅助键。
|
||||
|
||||
辅助键缺失时,用户的下一次提交会把 active_days / languages_used
|
||||
打回 1,且要等到下一次重算才恢复。
|
||||
"""
|
||||
title = "重算重建了增量辅助键"
|
||||
stat = UserStat.objects.filter(metrics__has_key="active_days").exclude(metrics__active_days=0).first()
|
||||
if stat is None:
|
||||
self._skip(title, "还没有任何用户的 active_days,先跑 recompute_achievements --user <id>")
|
||||
return
|
||||
|
||||
problems = []
|
||||
for state_key, top_key in STATE_PAIRS.items():
|
||||
top = stat.metrics.get(top_key)
|
||||
state = stat.metrics.get(state_key)
|
||||
if top is None:
|
||||
continue
|
||||
if state is None:
|
||||
problems.append(f"{state_key} 缺失({top_key}={top})")
|
||||
elif len(state) != top:
|
||||
problems.append(f"len({state_key})={len(state)} != {top_key}={top}")
|
||||
|
||||
if problems:
|
||||
self._fail(
|
||||
title,
|
||||
";".join(problems) + f"\n 用户 id={stat.user_id} 的下一次提交会让这些指标回退。",
|
||||
)
|
||||
else:
|
||||
self._ok(title, f"用户 id={stat.user_id} 的辅助键与顶层值一致")
|
||||
|
||||
def _check_unlock_count(self):
|
||||
"""unlock_count 是独立计数器,必须与实际解锁人数相等。
|
||||
|
||||
并发解锁重复累加、或删过 UserAchievement 但没重置计数器,
|
||||
都会让获得率虚高。
|
||||
"""
|
||||
title = "unlock_count 与实际解锁人数一致"
|
||||
if not Achievement.objects.exists():
|
||||
self._skip(title, "还没有配置任何成就")
|
||||
return
|
||||
|
||||
drifted = []
|
||||
for a in Achievement.objects.all():
|
||||
real = UserAchievement.objects.filter(achievement=a).count()
|
||||
if a.unlock_count != real:
|
||||
drifted.append(f"「{a.name}」计数器 {a.unlock_count} vs 实际 {real}")
|
||||
|
||||
if drifted:
|
||||
self._fail(
|
||||
title,
|
||||
";".join(drifted[:5]) + ("…" if len(drifted) > 5 else "") + "\n 获得率显示会不准。修:Achievement.objects.update(unlock_count=0) 后重跑重算。",
|
||||
)
|
||||
else:
|
||||
self._ok(title, f"{Achievement.objects.count()} 条成就全部一致")
|
||||
|
||||
def _check_dead_achievements(self):
|
||||
"""长期零解锁的成就多半是阈值配错了。"""
|
||||
title = "没有疑似配错阈值的成就"
|
||||
if not Achievement.objects.exists():
|
||||
self._skip(title, "还没有配置任何成就")
|
||||
return
|
||||
|
||||
dead = list(Achievement.objects.filter(visible=True, unlock_count=0, hidden=False).values_list("name", "metric", "operator", "threshold"))
|
||||
if dead:
|
||||
self._warn(
|
||||
title,
|
||||
f"{len(dead)} 条公开成就零解锁:"
|
||||
+ ";".join(f"「{n}」{m} {o} {t}" for n, m, o, t in dead[:5])
|
||||
+ ("…" if len(dead) > 5 else "")
|
||||
+ "\n 刚上线属正常;配置一周后仍为 0 就要怀疑阈值。",
|
||||
)
|
||||
else:
|
||||
self._ok(title)
|
||||
@@ -1,65 +0,0 @@
|
||||
"""全量重算指标并补发成就。
|
||||
|
||||
用途:
|
||||
1. 系统首次上线,给存量用户补发(用 --silent,否则学生一登录被 30 个奖杯糊脸)
|
||||
2. 新增指标或改了口径后重算(不带 --silent,新解锁会正常弹出)
|
||||
|
||||
这也是增量逻辑的安全网:怀疑 UserStat 漂移时跑一遍即可。
|
||||
"""
|
||||
|
||||
from django.core.management.base import BaseCommand
|
||||
|
||||
from account.models import User
|
||||
from achievement import checker
|
||||
from achievement.metrics import META_METRICS, METRIC_REGISTRY
|
||||
from achievement.models import UserStat
|
||||
|
||||
|
||||
class Command(BaseCommand):
|
||||
help = "重算全部用户的成就指标并补发成就"
|
||||
|
||||
def add_arguments(self, parser):
|
||||
parser.add_argument("--user", type=int, default=None, help="只处理指定用户 id")
|
||||
parser.add_argument("--silent", action="store_true", help="补发的成就标记为已通知,不弹窗")
|
||||
|
||||
def handle(self, *args, **options):
|
||||
users = User.objects.filter(is_disabled=False)
|
||||
if options["user"]:
|
||||
users = users.filter(id=options["user"])
|
||||
|
||||
silent = options["silent"]
|
||||
total_users = users.count()
|
||||
total_unlocked = 0
|
||||
|
||||
for index, user in enumerate(users.iterator(), start=1):
|
||||
stat, _ = UserStat.objects.get_or_create(user=user)
|
||||
metrics = {}
|
||||
for key, m in METRIC_REGISTRY.items():
|
||||
if key in META_METRICS:
|
||||
continue
|
||||
value = m.recompute(user)
|
||||
# None 表示该指标无有效值,key 必须缺席而不是置 0
|
||||
if value is not None:
|
||||
metrics[key] = value
|
||||
# 一并重建增量辅助键,否则重算后的第一次判题会把
|
||||
# active_days / languages_used / max_ac_in_one_day 打回 1
|
||||
metrics.update(m.recompute_state(user))
|
||||
stat.metrics = metrics
|
||||
stat.save(update_fields=["metrics", "update_time"])
|
||||
|
||||
records = checker.unlock(user, checker.evaluate(user, metrics), backfilled=True, notified=silent)
|
||||
|
||||
# 元指标第二轮
|
||||
for key in META_METRICS:
|
||||
value = METRIC_REGISTRY[key].recompute(user)
|
||||
if value is not None:
|
||||
metrics[key] = value
|
||||
stat.metrics = metrics
|
||||
stat.save(update_fields=["metrics", "update_time"])
|
||||
records += checker.unlock(user, checker.evaluate(user, metrics, only_metrics=META_METRICS), backfilled=True, notified=silent)
|
||||
|
||||
total_unlocked += len(records)
|
||||
if index % 50 == 0:
|
||||
self.stdout.write(f"processed {index}/{total_users}")
|
||||
|
||||
self.stdout.write(self.style.SUCCESS(f"完成:{total_users} 个用户,新解锁 {total_unlocked} 条"))
|
||||
@@ -1,403 +0,0 @@
|
||||
"""成就指标注册表。
|
||||
|
||||
这里定义"能测量什么",后台定义"多少算达成"。管理员在后台看到的指标下拉框
|
||||
就是 METRIC_REGISTRY 的 key 列表。加一个新维度必须改本文件并部署,
|
||||
之后在该维度上加任意多条成就都是纯配置。
|
||||
|
||||
约定:指标从未产生过有效值时,recompute 返回 None,调用方删除该 key。
|
||||
metrics 字典里 key 不存在 == 未达标,判定时直接跳过。
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from django.db.models import Count, Q
|
||||
from django.utils import timezone
|
||||
|
||||
from submission.models import JudgeStatus, Submission, is_accepted
|
||||
from utils.constants import Difficulty
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
METRIC_REGISTRY = {}
|
||||
META_METRICS = set()
|
||||
|
||||
# AC 口径与全项目一致:AST_CHECK_FAILED(代码结构检查未通过但测试点全过)也算通过。
|
||||
# 见 submission/models.py 的 is_accepted()。ORM 过滤用这个常量,标量比较用 is_accepted()。
|
||||
ACCEPTED_RESULTS = (JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED)
|
||||
|
||||
|
||||
class BaseMetric:
|
||||
key = ""
|
||||
name = ""
|
||||
help_text = ""
|
||||
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
"""判题完成后增量更新 metrics(原地修改)。"""
|
||||
raise NotImplementedError
|
||||
|
||||
def recompute(self, user):
|
||||
"""全量重算。返回 None 表示该用户此指标无有效值。"""
|
||||
raise NotImplementedError
|
||||
|
||||
def recompute_state(self, user):
|
||||
"""重算时一并重建增量所需的辅助键,返回 {key: value}。
|
||||
|
||||
默认无辅助状态。凡是 on_submission 依赖 `_` 前缀辅助键的指标都必须重写
|
||||
本方法:否则全量重算把辅助键丢掉后,下一次判题会从零重新累积,指标当场
|
||||
回退(active_days 从 45 掉回 1),且要等到下一次重算才恢复。
|
||||
"""
|
||||
return {}
|
||||
|
||||
|
||||
def metric(key, name, help_text="", meta=False):
|
||||
def deco(cls):
|
||||
cls.key = key
|
||||
cls.name = name
|
||||
cls.help_text = help_text
|
||||
METRIC_REGISTRY[key] = cls()
|
||||
if meta:
|
||||
META_METRICS.add(key)
|
||||
return cls
|
||||
|
||||
return deco
|
||||
|
||||
|
||||
def _practice_submissions(user_id):
|
||||
"""成就只统计平时练习的提交,比赛提交不计入。"""
|
||||
return Submission.objects.filter(user_id=user_id, contest_id__isnull=True)
|
||||
|
||||
|
||||
def build_ctx(user_id, sub):
|
||||
"""判题后预查一次,所有指标复用,避免每个指标各自查库。
|
||||
|
||||
比赛提交不参与成就统计,此时返回 skip=True,调用方直接跳过整轮判定。
|
||||
"""
|
||||
if sub.contest_id is not None:
|
||||
return {"skip": True}
|
||||
|
||||
prior = _practice_submissions(user_id).filter(problem_id=sub.problem_id).exclude(id=sub.id)
|
||||
prior_stats = prior.aggregate(
|
||||
total=Count("id"),
|
||||
accepted=Count("id", filter=Q(result__in=ACCEPTED_RESULTS)),
|
||||
)
|
||||
local_now = timezone.localtime(sub.create_time)
|
||||
sub_is_accepted = is_accepted(sub.result)
|
||||
is_first_ac_of_problem = sub_is_accepted and prior_stats["accepted"] == 0
|
||||
|
||||
# 难度只有首次 AC 时才用得上,其余情况不查这一次库——
|
||||
# 绝大多数提交都不是首次 AC,放在外面等于给每次判题白加一条 SQL
|
||||
difficulty = None
|
||||
if is_first_ac_of_problem:
|
||||
from problem.models import Problem
|
||||
|
||||
difficulty = Problem.objects.filter(id=sub.problem_id).values_list("difficulty", flat=True).first()
|
||||
|
||||
return {
|
||||
"skip": False,
|
||||
"is_accepted": sub_is_accepted,
|
||||
# 该题此前的提交次数与 AC 次数
|
||||
"prior_count": prior_stats["total"],
|
||||
"prior_accepted": prior_stats["accepted"],
|
||||
# 首次 AC 这道题(此前从未 AC 过)
|
||||
"is_first_ac_of_problem": is_first_ac_of_problem,
|
||||
# 一发入魂:此前无任何提交且本次 AC
|
||||
"is_first_try_ac": sub_is_accepted and prior_stats["total"] == 0,
|
||||
# 本题难度,仅首次 AC 时有值
|
||||
"problem_difficulty": difficulty,
|
||||
"local_date": local_now.date().isoformat(),
|
||||
"local_hour": local_now.hour,
|
||||
}
|
||||
|
||||
|
||||
@metric("accepted_count", "AC 题目数", "去重后通过的题目数量(不含比赛)")
|
||||
class AcceptedCount(BaseMetric):
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
if ctx["is_first_ac_of_problem"]:
|
||||
metrics["accepted_count"] = metrics.get("accepted_count", 0) + 1
|
||||
|
||||
def recompute(self, user):
|
||||
return _practice_submissions(user.id).filter(result__in=ACCEPTED_RESULTS).order_by().values("problem_id").distinct().count()
|
||||
|
||||
|
||||
class _DifficultyAcCount(BaseMetric):
|
||||
"""按难度去重统计 AC 题数。子类只需指定 difficulty。
|
||||
|
||||
增量靠 ctx["problem_difficulty"],它只在首次 AC 时才有值——与
|
||||
is_first_ac_of_problem 是同一个条件,所以两者一起判即可。
|
||||
"""
|
||||
|
||||
difficulty = ""
|
||||
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
if ctx["is_first_ac_of_problem"] and ctx["problem_difficulty"] == self.difficulty:
|
||||
metrics[self.key] = metrics.get(self.key, 0) + 1
|
||||
|
||||
def recompute(self, user):
|
||||
return _practice_submissions(user.id).filter(result__in=ACCEPTED_RESULTS, problem__difficulty=self.difficulty).order_by().values("problem_id").distinct().count()
|
||||
|
||||
|
||||
@metric("mid_ac_count", "中等题 AC 数", "去重后通过的中等难度题目数(不含比赛)")
|
||||
class MidAcCount(_DifficultyAcCount):
|
||||
difficulty = Difficulty.MID
|
||||
|
||||
|
||||
@metric("hard_ac_count", "困难题 AC 数", "去重后通过的困难题目数(不含比赛)")
|
||||
class HardAcCount(_DifficultyAcCount):
|
||||
difficulty = Difficulty.HIGH
|
||||
|
||||
|
||||
@metric("submission_count", "提交总数", "提交次数(不含比赛)")
|
||||
class SubmissionCount(BaseMetric):
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
metrics["submission_count"] = metrics.get("submission_count", 0) + 1
|
||||
|
||||
def recompute(self, user):
|
||||
return _practice_submissions(user.id).count()
|
||||
|
||||
|
||||
@metric("active_days", "活跃天数", "有过提交的累计天数")
|
||||
class ActiveDays(BaseMetric):
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
seen = metrics.get("_active_dates", [])
|
||||
if ctx["local_date"] not in seen:
|
||||
seen.append(ctx["local_date"])
|
||||
metrics["_active_dates"] = seen
|
||||
metrics["active_days"] = len(seen)
|
||||
|
||||
def recompute(self, user):
|
||||
dates = {timezone.localtime(t).date().isoformat() for t in _practice_submissions(user.id).values_list("create_time", flat=True)}
|
||||
return len(dates)
|
||||
|
||||
def recompute_state(self, user):
|
||||
dates = sorted({timezone.localtime(t).date().isoformat() for t in _practice_submissions(user.id).values_list("create_time", flat=True)})
|
||||
return {"_active_dates": dates}
|
||||
|
||||
|
||||
@metric("max_ac_streak_days", "最长连续 AC 天数", "连续每天至少 AC 一题的最长天数")
|
||||
class MaxAcStreakDays(BaseMetric):
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
if not ctx["is_accepted"]:
|
||||
return
|
||||
today = ctx["local_date"]
|
||||
last = metrics.get("_last_ac_date")
|
||||
if last == today:
|
||||
return
|
||||
current = metrics.get("_current_ac_streak", 0)
|
||||
if last and (timezone.datetime.fromisoformat(today) - timezone.datetime.fromisoformat(last)).days == 1:
|
||||
current += 1
|
||||
else:
|
||||
current = 1
|
||||
metrics["_last_ac_date"] = today
|
||||
metrics["_current_ac_streak"] = current
|
||||
metrics["max_ac_streak_days"] = max(metrics.get("max_ac_streak_days", 0), current)
|
||||
|
||||
def recompute(self, user):
|
||||
dates = sorted({timezone.localtime(t).date() for t in _practice_submissions(user.id).filter(result__in=ACCEPTED_RESULTS).values_list("create_time", flat=True)})
|
||||
if not dates:
|
||||
return None
|
||||
best = current = 1
|
||||
for prev, cur in zip(dates, dates[1:]):
|
||||
current = current + 1 if (cur - prev).days == 1 else 1
|
||||
best = max(best, current)
|
||||
return best
|
||||
|
||||
def recompute_state(self, user):
|
||||
dates = sorted({timezone.localtime(t).date() for t in _practice_submissions(user.id).filter(result__in=ACCEPTED_RESULTS).values_list("create_time", flat=True)})
|
||||
if not dates:
|
||||
return {}
|
||||
current = 1
|
||||
for prev, cur in zip(dates, dates[1:]):
|
||||
current = current + 1 if (cur - prev).days == 1 else 1
|
||||
return {"_last_ac_date": dates[-1].isoformat(), "_current_ac_streak": current}
|
||||
|
||||
|
||||
@metric("languages_used", "使用语言数", "用过多少种编程语言")
|
||||
class LanguagesUsed(BaseMetric):
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
seen = metrics.get("_languages", [])
|
||||
if sub.language not in seen:
|
||||
seen.append(sub.language)
|
||||
metrics["_languages"] = seen
|
||||
metrics["languages_used"] = len(seen)
|
||||
|
||||
def recompute(self, user):
|
||||
return _practice_submissions(user.id).order_by().values("language").distinct().count()
|
||||
|
||||
def recompute_state(self, user):
|
||||
# order_by() 不能省:Submission.Meta 有默认排序 ("-create_time",),
|
||||
# Django 会把排序字段并入 DISTINCT,于是每条提交各成一行——
|
||||
# 实测某用户返回 659 条而不是 5 种语言。
|
||||
# recompute 侥幸正确只是因为 .count() 会清掉排序,不能依赖这一点。
|
||||
return {"_languages": list(_practice_submissions(user.id).order_by().values_list("language", flat=True).distinct())}
|
||||
|
||||
|
||||
@metric("contest_joined", "参赛场次", "参加过的比赛数量(本指标是比赛维度,不受比赛提交不计入的限制)")
|
||||
class ContestJoined(BaseMetric):
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
# 比赛提交在 build_ctx 就被跳过,本指标只走 recompute
|
||||
return
|
||||
|
||||
def recompute(self, user):
|
||||
return Submission.objects.filter(user_id=user.id, contest_id__isnull=False).order_by().values("contest_id").distinct().count()
|
||||
|
||||
|
||||
@metric("badge_count", "题单奖章数", "获得的题单奖章数量")
|
||||
class BadgeCount(BaseMetric):
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
# 奖章由题单流程颁发,本指标只走 recompute,见 Task 7
|
||||
return
|
||||
|
||||
def recompute(self, user):
|
||||
from problemset.models import UserBadge
|
||||
|
||||
return UserBadge.objects.filter(user=user).count()
|
||||
|
||||
|
||||
@metric("problemset_completed", "完成题单数", "完成的题单数量")
|
||||
class ProblemSetCompleted(BaseMetric):
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
return
|
||||
|
||||
def recompute(self, user):
|
||||
from problemset.models import ProblemSetProgress
|
||||
|
||||
return ProblemSetProgress.objects.filter(user=user, complete_time__isnull=False).count()
|
||||
|
||||
|
||||
@metric("first_try_ac_count", "一发入魂次数", "首次提交即通过的次数")
|
||||
class FirstTryAcCount(BaseMetric):
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
if ctx["is_first_try_ac"]:
|
||||
metrics["first_try_ac_count"] = metrics.get("first_try_ac_count", 0) + 1
|
||||
|
||||
def recompute(self, user):
|
||||
count = 0
|
||||
seen = set()
|
||||
for s in _practice_submissions(user.id).order_by("create_time").values("problem_id", "result"):
|
||||
if s["problem_id"] in seen:
|
||||
continue
|
||||
seen.add(s["problem_id"])
|
||||
if is_accepted(s["result"]):
|
||||
count += 1
|
||||
return count
|
||||
|
||||
|
||||
@metric("midnight_submissions", "凌晨提交次数", "0:00–5:00 之间的提交次数")
|
||||
class MidnightSubmissions(BaseMetric):
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
if 0 <= ctx["local_hour"] < 5:
|
||||
metrics["midnight_submissions"] = metrics.get("midnight_submissions", 0) + 1
|
||||
|
||||
def recompute(self, user):
|
||||
return sum(1 for t in _practice_submissions(user.id).values_list("create_time", flat=True) if 0 <= timezone.localtime(t).hour < 5)
|
||||
|
||||
|
||||
@metric("early_bird_submissions", "早起提交次数", "5:00–7:00 之间的提交次数")
|
||||
class EarlyBirdSubmissions(BaseMetric):
|
||||
"""与 midnight_submissions 对称的作息维度。5 点是两者的分界,不重叠。"""
|
||||
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
if 5 <= ctx["local_hour"] < 7:
|
||||
metrics["early_bird_submissions"] = metrics.get("early_bird_submissions", 0) + 1
|
||||
|
||||
def recompute(self, user):
|
||||
return sum(1 for t in _practice_submissions(user.id).values_list("create_time", flat=True) if 5 <= timezone.localtime(t).hour < 7)
|
||||
|
||||
|
||||
@metric("compile_error_count", "编译错误次数", "累计编译错误的次数")
|
||||
class CompileErrorCount(BaseMetric):
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
if sub.result == JudgeStatus.COMPILE_ERROR:
|
||||
metrics["compile_error_count"] = metrics.get("compile_error_count", 0) + 1
|
||||
|
||||
def recompute(self, user):
|
||||
return _practice_submissions(user.id).filter(result=JudgeStatus.COMPILE_ERROR).count()
|
||||
|
||||
|
||||
@metric("max_wa_before_ac", "屡败屡战", "单题失败最多多少次后终于通过")
|
||||
class MaxWaBeforeAc(BaseMetric):
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
if ctx["is_first_ac_of_problem"]:
|
||||
metrics["max_wa_before_ac"] = max(metrics.get("max_wa_before_ac", 0), ctx["prior_count"])
|
||||
|
||||
def recompute(self, user):
|
||||
best = None
|
||||
attempts = {}
|
||||
for s in _practice_submissions(user.id).order_by("create_time").values("problem_id", "result"):
|
||||
pid = s["problem_id"]
|
||||
if pid in attempts and attempts[pid] is None:
|
||||
continue
|
||||
if is_accepted(s["result"]):
|
||||
best = max(best or 0, attempts.get(pid, 0))
|
||||
attempts[pid] = None
|
||||
else:
|
||||
attempts[pid] = attempts.get(pid, 0) + 1
|
||||
return best
|
||||
|
||||
|
||||
@metric("max_ac_in_one_day", "单日最多 AC", "一天之内最多通过多少题")
|
||||
class MaxAcInOneDay(BaseMetric):
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
if not ctx["is_first_ac_of_problem"]:
|
||||
return
|
||||
counts = metrics.get("_ac_per_day", {})
|
||||
counts[ctx["local_date"]] = counts.get(ctx["local_date"], 0) + 1
|
||||
metrics["_ac_per_day"] = counts
|
||||
metrics["max_ac_in_one_day"] = max(counts.values())
|
||||
|
||||
def recompute(self, user):
|
||||
counts = {}
|
||||
seen = set()
|
||||
for s in _practice_submissions(user.id).filter(result__in=ACCEPTED_RESULTS).order_by("create_time").values("problem_id", "create_time"):
|
||||
if s["problem_id"] in seen:
|
||||
continue
|
||||
seen.add(s["problem_id"])
|
||||
day = timezone.localtime(s["create_time"]).date().isoformat()
|
||||
counts[day] = counts.get(day, 0) + 1
|
||||
return max(counts.values()) if counts else None
|
||||
|
||||
def recompute_state(self, user):
|
||||
counts = {}
|
||||
seen = set()
|
||||
for s in _practice_submissions(user.id).filter(result__in=ACCEPTED_RESULTS).order_by("create_time").values("problem_id", "create_time"):
|
||||
if s["problem_id"] in seen:
|
||||
continue
|
||||
seen.add(s["problem_id"])
|
||||
day = timezone.localtime(s["create_time"]).date().isoformat()
|
||||
counts[day] = counts.get(day, 0) + 1
|
||||
return {"_ac_per_day": counts}
|
||||
|
||||
|
||||
# 曾经这里有 min_ac_code_chars(最短 AC 代码)。线上实测 1314 个用户的分布,
|
||||
# 最小值 8、p5=10:有道题 8 个字符就能通过,于是它测的是"谁做过那道水题"
|
||||
# 而不是"谁写得简洁",配不出有意义的成就,2026-08-05 删除。
|
||||
# 要重新引入,得先按题目难度加权,或排除掉那类水题。
|
||||
|
||||
|
||||
@metric("max_code_lines", "最长代码行数", "提交过的最长代码有多少行")
|
||||
class MaxCodeLines(BaseMetric):
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
lines = len(sub.code.splitlines())
|
||||
metrics["max_code_lines"] = max(metrics.get("max_code_lines", 0), lines)
|
||||
|
||||
def recompute(self, user):
|
||||
counts = [len(c.splitlines()) for c in _practice_submissions(user.id).values_list("code", flat=True)]
|
||||
return max(counts) if counts else None
|
||||
|
||||
|
||||
@metric("achievement_unlocked_count", "已解锁成就数", "已解锁的成就数量(不含白金档)", meta=True)
|
||||
class AchievementUnlockedCount(BaseMetric):
|
||||
"""自引用指标:解锁成就会改变它。
|
||||
|
||||
因此判定流程限定为最多两轮(见 checker.py),且口径排除白金档自身,
|
||||
避免「集齐 N 个成就」这类白金奖杯把自己算进分子。
|
||||
"""
|
||||
|
||||
def on_submission(self, metrics, sub, ctx):
|
||||
# 由 checker 在第一轮解锁后显式重算,不参与增量更新
|
||||
return
|
||||
|
||||
def recompute(self, user):
|
||||
from achievement.models import Rarity, UserAchievement
|
||||
|
||||
return UserAchievement.objects.filter(user=user).exclude(achievement__rarity=Rarity.PLATINUM).count()
|
||||
@@ -1,70 +0,0 @@
|
||||
# Generated by Django 6.0.4 on 2026-08-04 06:35
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.conf import settings
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
initial = True
|
||||
|
||||
dependencies = [
|
||||
migrations.swappable_dependency(settings.AUTH_USER_MODEL),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='Achievement',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('name', models.TextField(verbose_name='成就名称')),
|
||||
('description', models.TextField(verbose_name='成就描述')),
|
||||
('icon', models.TextField(verbose_name='图标')),
|
||||
('rarity', models.TextField(choices=[('bronze', '青铜'), ('silver', '白银'), ('gold', '黄金'), ('platinum', '白金')], default='bronze', verbose_name='稀有度')),
|
||||
('hidden', models.BooleanField(db_default=False, default=False, verbose_name='是否隐藏')),
|
||||
('metric', models.TextField(verbose_name='指标名')),
|
||||
('operator', models.TextField(choices=[('gte', '大于等于'), ('lte', '小于等于')], default='gte', verbose_name='比较符')),
|
||||
('threshold', models.IntegerField(verbose_name='阈值')),
|
||||
('visible', models.BooleanField(db_default=True, default=True, verbose_name='是否上架')),
|
||||
('unlock_count', models.IntegerField(db_default=0, default=0, verbose_name='已解锁人数')),
|
||||
('order', models.IntegerField(db_default=0, default=0, verbose_name='排序')),
|
||||
('create_time', models.DateTimeField(auto_now_add=True)),
|
||||
],
|
||||
options={
|
||||
'verbose_name': '成就',
|
||||
'verbose_name_plural': '成就',
|
||||
'db_table': 'achievement',
|
||||
'ordering': ('order', 'id'),
|
||||
},
|
||||
),
|
||||
migrations.CreateModel(
|
||||
name='UserStat',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('metrics', models.JSONField(db_default=models.Value({}, output_field=models.JSONField()), default=dict)),
|
||||
('update_time', models.DateTimeField(auto_now=True)),
|
||||
('user', models.OneToOneField(on_delete=django.db.models.deletion.CASCADE, related_name='achievement_stat', to=settings.AUTH_USER_MODEL)),
|
||||
],
|
||||
options={
|
||||
'db_table': 'user_stat',
|
||||
},
|
||||
),
|
||||
migrations.CreateModel(
|
||||
name='UserAchievement',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('unlock_time', models.DateTimeField(auto_now_add=True)),
|
||||
('backfilled', models.BooleanField(db_default=False, default=False)),
|
||||
('notified', models.BooleanField(db_default=False, default=False)),
|
||||
('achievement', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to='achievement.achievement')),
|
||||
('user', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='achievements', to=settings.AUTH_USER_MODEL)),
|
||||
],
|
||||
options={
|
||||
'db_table': 'user_achievement',
|
||||
'ordering': ('-unlock_time',),
|
||||
'indexes': [models.Index(fields=['user', '-unlock_time'], name='user_achv_time_idx'), models.Index(fields=['user', 'notified'], name='user_achv_notified_idx')],
|
||||
'constraints': [models.UniqueConstraint(fields=('user', 'achievement'), name='unique_user_achievement')],
|
||||
},
|
||||
),
|
||||
]
|
||||
@@ -1,123 +0,0 @@
|
||||
"""把 Achievement.icon 从 emoji 字符换成 iconify 图标名。
|
||||
|
||||
机房里的 Chrome 91 之类的老浏览器缺 emoji 字体,emoji 会渲染成方块。
|
||||
iconify 出来的是 SVG,跟浏览器字体无关。
|
||||
|
||||
映射到 noto(Noto Color Emoji 的 SVG 版),视觉上和原 emoji 基本一致。
|
||||
表里没有的 emoji 保持原值不动,前端会按纯文本兜底渲染,之后在管理后台手工补。
|
||||
"""
|
||||
|
||||
from django.db import migrations
|
||||
|
||||
# 带 ️(variation selector-16)的 emoji 后台可能存成带或不带两种形态,
|
||||
# 迁移时统一去掉再查表,所以这里的 key 一律不带 ️
|
||||
EMOJI_TO_ICONIFY = {
|
||||
# 计划里首批成就用到的
|
||||
"🌱": "noto:seedling",
|
||||
"📗": "noto:green-book",
|
||||
"📚": "noto:books",
|
||||
"📅": "noto:calendar",
|
||||
"🗓": "noto:spiral-calendar",
|
||||
"🌐": "noto:globe-with-meridians",
|
||||
"🎖": "noto:military-medal",
|
||||
"🎯": "noto:direct-hit",
|
||||
"🏹": "noto:bow-and-arrow",
|
||||
"🦉": "noto:owl",
|
||||
"🌙": "noto:crescent-moon",
|
||||
"💥": "noto:collision",
|
||||
"🔥": "noto:fire",
|
||||
"⚡": "noto:high-voltage",
|
||||
"✂": "noto:scissors",
|
||||
"📜": "noto:scroll",
|
||||
"💎": "noto:gem-stone",
|
||||
"🏆": "noto:trophy",
|
||||
"❓": "noto:red-question-mark",
|
||||
# 其余常用图标,后台可能已经用上了
|
||||
"❔": "noto:white-question-mark",
|
||||
"📖": "noto:open-book",
|
||||
"🥇": "noto:1st-place-medal",
|
||||
"🥈": "noto:2nd-place-medal",
|
||||
"🥉": "noto:3rd-place-medal",
|
||||
"🏅": "noto:sports-medal",
|
||||
"⭐": "noto:star",
|
||||
"🌟": "noto:glowing-star",
|
||||
"✨": "noto:sparkles",
|
||||
"🎉": "noto:party-popper",
|
||||
"🎊": "noto:confetti-ball",
|
||||
"🚀": "noto:rocket",
|
||||
"💡": "noto:light-bulb",
|
||||
"🐛": "noto:bug",
|
||||
"👑": "noto:crown",
|
||||
"🥷": "noto:ninja",
|
||||
"🧠": "noto:brain",
|
||||
"💪": "noto:flexed-biceps",
|
||||
"🎓": "noto:graduation-cap",
|
||||
"📈": "noto:chart-increasing",
|
||||
"🎁": "noto:wrapped-gift",
|
||||
"🔮": "noto:crystal-ball",
|
||||
"🦄": "noto:unicorn",
|
||||
"🐢": "noto:turtle",
|
||||
"☕": "noto:hot-beverage",
|
||||
"🌈": "noto:rainbow",
|
||||
"⏰": "noto:alarm-clock",
|
||||
"📝": "noto:memo",
|
||||
"✅": "noto:check-mark-button",
|
||||
"🔒": "noto:locked",
|
||||
"🎮": "noto:video-game",
|
||||
"🧩": "noto:puzzle-piece",
|
||||
"🔧": "noto:wrench",
|
||||
"🐍": "noto:snake",
|
||||
"💻": "noto:laptop",
|
||||
"⌨": "noto:keyboard",
|
||||
"☀": "noto:sun",
|
||||
"😴": "noto:sleeping-face",
|
||||
"🤖": "noto:robot",
|
||||
"🔍": "noto:magnifying-glass-tilted-left",
|
||||
"🎨": "noto:artist-palette",
|
||||
"⚔": "noto:crossed-swords",
|
||||
"🛡": "noto:shield",
|
||||
"🏰": "noto:castle",
|
||||
"🗿": "noto:moai",
|
||||
"🌊": "noto:water-wave",
|
||||
"🍀": "noto:four-leaf-clover",
|
||||
"🔁": "noto:repeat-button",
|
||||
"🎪": "noto:circus-tent",
|
||||
"🦅": "noto:eagle",
|
||||
"🐉": "noto:dragon",
|
||||
"💫": "noto:dizzy",
|
||||
"🧭": "noto:compass",
|
||||
"🪄": "noto:magic-wand",
|
||||
"🔔": "noto:bell",
|
||||
"📌": "noto:pushpin",
|
||||
"🏁": "noto:chequered-flag",
|
||||
"🎢": "noto:roller-coaster",
|
||||
}
|
||||
|
||||
ICONIFY_TO_EMOJI = {v: k for k, v in EMOJI_TO_ICONIFY.items()}
|
||||
|
||||
|
||||
def _convert(apps, schema_editor, table):
|
||||
Achievement = apps.get_model("achievement", "Achievement")
|
||||
updated = []
|
||||
for a in Achievement.objects.all():
|
||||
key = (a.icon or "").replace("️", "").strip()
|
||||
new = table.get(key)
|
||||
if new and new != a.icon:
|
||||
a.icon = new
|
||||
updated.append(a)
|
||||
if updated:
|
||||
Achievement.objects.bulk_update(updated, ["icon"])
|
||||
|
||||
|
||||
def emoji_to_iconify(apps, schema_editor):
|
||||
_convert(apps, schema_editor, EMOJI_TO_ICONIFY)
|
||||
|
||||
|
||||
def iconify_to_emoji(apps, schema_editor):
|
||||
_convert(apps, schema_editor, ICONIFY_TO_EMOJI)
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
dependencies = [("achievement", "0001_initial")]
|
||||
|
||||
operations = [migrations.RunPython(emoji_to_iconify, iconify_to_emoji)]
|
||||
@@ -1,75 +0,0 @@
|
||||
from django.db import models
|
||||
|
||||
from account.models import User
|
||||
from utils.models import JSONField
|
||||
|
||||
|
||||
class Rarity(models.TextChoices):
|
||||
BRONZE = "bronze", "青铜"
|
||||
SILVER = "silver", "白银"
|
||||
GOLD = "gold", "黄金"
|
||||
PLATINUM = "platinum", "白金"
|
||||
|
||||
|
||||
class Operator(models.TextChoices):
|
||||
GTE = "gte", "大于等于"
|
||||
LTE = "lte", "小于等于"
|
||||
|
||||
|
||||
class Achievement(models.Model):
|
||||
name = models.TextField(verbose_name="成就名称")
|
||||
description = models.TextField(verbose_name="成就描述")
|
||||
# iconify 图标名(如 noto:owl),不是 emoji 字符:
|
||||
# 机房的老浏览器缺 emoji 字体会渲染成方块,前端统一渲染成 SVG
|
||||
icon = models.TextField(verbose_name="图标")
|
||||
rarity = models.TextField(default=Rarity.BRONZE, choices=Rarity.choices, verbose_name="稀有度")
|
||||
hidden = models.BooleanField(default=False, db_default=False, verbose_name="是否隐藏")
|
||||
metric = models.TextField(verbose_name="指标名")
|
||||
operator = models.TextField(default=Operator.GTE, choices=Operator.choices, verbose_name="比较符")
|
||||
threshold = models.IntegerField(verbose_name="阈值")
|
||||
visible = models.BooleanField(default=True, db_default=True, verbose_name="是否上架")
|
||||
unlock_count = models.IntegerField(default=0, db_default=0, verbose_name="已解锁人数")
|
||||
order = models.IntegerField(default=0, db_default=0, verbose_name="排序")
|
||||
create_time = models.DateTimeField(auto_now_add=True)
|
||||
|
||||
class Meta:
|
||||
db_table = "achievement"
|
||||
ordering = ("order", "id")
|
||||
verbose_name = "成就"
|
||||
verbose_name_plural = "成就"
|
||||
|
||||
|
||||
class UserStat(models.Model):
|
||||
"""成就系统唯一的指标源,不复用 UserProfile 的计数器,避免两处口径漂移。
|
||||
|
||||
metrics 为 {指标名: 数值};指标从未产生过有效值时 key 不存在(而非置 0),
|
||||
否则极小值型指标(求 min、配 lte 用的那种)会对新用户恒成立。
|
||||
"""
|
||||
|
||||
user = models.OneToOneField(User, on_delete=models.CASCADE, related_name="achievement_stat")
|
||||
metrics = JSONField(default=dict, db_default=models.Value({}, output_field=models.JSONField()))
|
||||
update_time = models.DateTimeField(auto_now=True)
|
||||
|
||||
class Meta:
|
||||
db_table = "user_stat"
|
||||
|
||||
|
||||
class UserAchievement(models.Model):
|
||||
user = models.ForeignKey(User, on_delete=models.CASCADE, related_name="achievements")
|
||||
achievement = models.ForeignKey(Achievement, on_delete=models.CASCADE)
|
||||
unlock_time = models.DateTimeField(auto_now_add=True)
|
||||
# 上线补发的记录:前端不显示具体日期,只显示"已获得"
|
||||
backfilled = models.BooleanField(default=False, db_default=False)
|
||||
# 是否已向用户弹过奖杯,pending 端点据此查询
|
||||
notified = models.BooleanField(default=False, db_default=False)
|
||||
|
||||
class Meta:
|
||||
db_table = "user_achievement"
|
||||
ordering = ("-unlock_time",)
|
||||
constraints = [
|
||||
models.UniqueConstraint(fields=["user", "achievement"], name="unique_user_achievement"),
|
||||
]
|
||||
indexes = [
|
||||
models.Index(fields=["user", "-unlock_time"], name="user_achv_time_idx"),
|
||||
models.Index(fields=["user", "notified"], name="user_achv_notified_idx"),
|
||||
]
|
||||
@@ -1,59 +0,0 @@
|
||||
"""解锁通知。
|
||||
|
||||
通知走推拉结合,UserAchievement.notified 是唯一的真相来源:
|
||||
- 拉(主):前端在布局层拉 /api/achievements/pending,覆盖全部场景
|
||||
- 推(增强):WebSocket 只负责把"当场那一下"的延迟压到几百毫秒
|
||||
|
||||
必须推拉结合的原因:前端 WebSocket 不是常驻连接,useSubmissionWebSocket
|
||||
只在问题页且有提交监听时建连,纯推会丢消息。
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
from utils.websocket import push_to_user
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def notify_achievements(user_id, records):
|
||||
"""records: list[UserAchievement],已带 select_related('achievement')。"""
|
||||
if not records:
|
||||
return
|
||||
payload = [
|
||||
{
|
||||
"id": r.achievement_id,
|
||||
"name": r.achievement.name,
|
||||
"description": r.achievement.description,
|
||||
"icon": r.achievement.icon,
|
||||
"rarity": r.achievement.rarity,
|
||||
"kind": "achievement",
|
||||
}
|
||||
for r in records
|
||||
]
|
||||
_push(user_id, payload)
|
||||
|
||||
|
||||
def notify_badges(user_id, badges):
|
||||
"""badges: list[ProblemSetBadge]。题单奖章复用同一个弹窗组件。"""
|
||||
if not badges:
|
||||
return
|
||||
payload = [
|
||||
{
|
||||
"id": b.id,
|
||||
"name": b.name,
|
||||
"description": b.description,
|
||||
"icon": b.icon,
|
||||
"rarity": "bronze",
|
||||
"kind": "badge",
|
||||
}
|
||||
for b in badges
|
||||
]
|
||||
_push(user_id, payload)
|
||||
|
||||
|
||||
def _push(user_id, payload):
|
||||
# 推送失败不影响已入库的解锁记录,前端下次拉 pending 时仍会补弹
|
||||
try:
|
||||
push_to_user(user_id, "achievement_unlocked", {"achievements": payload})
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to push achievement notification: user_id={user_id}, error={e}")
|
||||
@@ -1,97 +0,0 @@
|
||||
from django.core.cache import cache
|
||||
from rest_framework import serializers
|
||||
|
||||
from account.models import User
|
||||
from achievement.models import Achievement, UserAchievement
|
||||
|
||||
|
||||
def get_active_user_count():
|
||||
"""获得率的分母:未禁用用户总数,缓存 1 小时。
|
||||
|
||||
不用"有过提交的用户数"这类动态口径:分母波动会让同一个成就的获得率
|
||||
忽高忽低,学生会当成 bug。
|
||||
"""
|
||||
count = cache.get("achievement_active_user_count")
|
||||
if count is None:
|
||||
count = User.objects.filter(is_disabled=False).count()
|
||||
cache.set("achievement_active_user_count", count, 3600)
|
||||
return count
|
||||
|
||||
|
||||
class AchievementSerializer(serializers.ModelSerializer):
|
||||
unlocked = serializers.BooleanField(read_only=True)
|
||||
unlock_time = serializers.DateTimeField(read_only=True, allow_null=True)
|
||||
backfilled = serializers.BooleanField(read_only=True)
|
||||
progress = serializers.IntegerField(read_only=True, allow_null=True)
|
||||
unlock_rate = serializers.SerializerMethodField()
|
||||
name = serializers.SerializerMethodField()
|
||||
description = serializers.SerializerMethodField()
|
||||
icon = serializers.SerializerMethodField()
|
||||
# 条件三件套也必须跟着遮掉:只遮名称和描述、却明文下发
|
||||
# metric/operator/threshold,学生打开 DevTools 就知道"凌晨提交 10 次",
|
||||
# 隐藏成就的意义全部作废
|
||||
metric = serializers.SerializerMethodField()
|
||||
operator = serializers.SerializerMethodField()
|
||||
threshold = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = Achievement
|
||||
fields = (
|
||||
"id",
|
||||
"name",
|
||||
"description",
|
||||
"icon",
|
||||
"rarity",
|
||||
"hidden",
|
||||
"metric",
|
||||
"operator",
|
||||
"threshold",
|
||||
"unlocked",
|
||||
"unlock_time",
|
||||
"backfilled",
|
||||
"progress",
|
||||
"unlock_rate",
|
||||
)
|
||||
|
||||
def _masked(self, obj):
|
||||
return obj.hidden and not getattr(obj, "unlocked", False)
|
||||
|
||||
def get_name(self, obj):
|
||||
return "???" if self._masked(obj) else obj.name
|
||||
|
||||
def get_description(self, obj):
|
||||
return "达成条件保密" if self._masked(obj) else obj.description
|
||||
|
||||
def get_icon(self, obj):
|
||||
# icon 存的是 iconify 图标名,不是 emoji 字符:机房里的老浏览器缺 emoji
|
||||
# 字体会渲染成方块,前端统一按 iconify 渲染成 SVG
|
||||
return "noto:red-question-mark" if self._masked(obj) else obj.icon
|
||||
|
||||
def get_metric(self, obj):
|
||||
return None if self._masked(obj) else obj.metric
|
||||
|
||||
def get_operator(self, obj):
|
||||
return None if self._masked(obj) else obj.operator
|
||||
|
||||
def get_threshold(self, obj):
|
||||
return None if self._masked(obj) else obj.threshold
|
||||
|
||||
def get_unlock_rate(self, obj):
|
||||
total = get_active_user_count()
|
||||
if not total:
|
||||
return 0.0
|
||||
return round(obj.unlock_count / total * 100, 1)
|
||||
|
||||
|
||||
class PendingAchievementSerializer(serializers.ModelSerializer):
|
||||
"""待弹窗的解锁记录,无需打码(已解锁)。"""
|
||||
|
||||
id = serializers.IntegerField(source="achievement_id", read_only=True)
|
||||
name = serializers.CharField(source="achievement.name", read_only=True)
|
||||
description = serializers.CharField(source="achievement.description", read_only=True)
|
||||
icon = serializers.CharField(source="achievement.icon", read_only=True)
|
||||
rarity = serializers.CharField(source="achievement.rarity", read_only=True)
|
||||
|
||||
class Meta:
|
||||
model = UserAchievement
|
||||
fields = ("id", "name", "description", "icon", "rarity")
|
||||
@@ -1,65 +0,0 @@
|
||||
import logging
|
||||
|
||||
import dramatiq
|
||||
|
||||
from account.models import User
|
||||
from achievement import checker
|
||||
from achievement.notify import notify_achievements
|
||||
from submission.models import Submission
|
||||
from utils.shortcuts import DRAMATIQ_WORKER_ARGS
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dramatiq.actor(**DRAMATIQ_WORKER_ARGS())
|
||||
def check_achievements(user_id, submission_id):
|
||||
"""判题完成后的成就判定。
|
||||
|
||||
所有异常在此吞掉:成就算错绝不能影响判题结果,这也是选异步的意义。
|
||||
"""
|
||||
try:
|
||||
user = User.objects.get(id=user_id)
|
||||
submission = Submission.objects.get(id=submission_id)
|
||||
records = checker.run_for_submission(user, submission)
|
||||
notify_achievements(user_id, records)
|
||||
except Exception as e:
|
||||
logger.exception(f"check_achievements failed: user_id={user_id}, submission_id={submission_id}, error={e}")
|
||||
|
||||
|
||||
@dramatiq.actor(**DRAMATIQ_WORKER_ARGS())
|
||||
def rescan_achievement(achievement_id):
|
||||
"""新建成就或调低阈值后,补发给已达标的存量用户。
|
||||
|
||||
判定只在判题时发生,因此后台改了阈值不会自动补发,必须显式扫一遍。
|
||||
"""
|
||||
from django.db.models import IntegerField
|
||||
from django.db.models.fields.json import KeyTextTransform
|
||||
from django.db.models.functions import Cast
|
||||
|
||||
from achievement.models import Achievement, Operator, UserAchievement, UserStat
|
||||
|
||||
try:
|
||||
achievement = Achievement.objects.get(id=achievement_id, visible=True)
|
||||
except Achievement.DoesNotExist:
|
||||
return
|
||||
|
||||
# JSONField 的默认比较是 JSON 值比较,数字会按字符串序比("9" > "50"),
|
||||
# 必须显式 cast 成整数,否则筛出来的用户是错的
|
||||
qs = UserStat.objects.filter(metrics__has_key=achievement.metric).annotate(v=Cast(KeyTextTransform(achievement.metric, "metrics"), IntegerField()))
|
||||
if achievement.operator == Operator.GTE:
|
||||
qs = qs.filter(v__gte=achievement.threshold)
|
||||
else:
|
||||
qs = qs.filter(v__lte=achievement.threshold)
|
||||
|
||||
already = set(UserAchievement.objects.filter(achievement=achievement).values_list("user_id", flat=True))
|
||||
for stat in qs.select_related("user").iterator():
|
||||
if stat.user_id in already:
|
||||
continue
|
||||
try:
|
||||
# backfilled=True:这是补发,不是学生刚刚挣到的。
|
||||
# 前端据此只显示"已获得"而不显示具体日期——否则一次补发会给几百人
|
||||
# 盖上同一个时间戳,把"最近获得"板块彻底冲垮
|
||||
records = checker.unlock(stat.user, [achievement], backfilled=True)
|
||||
notify_achievements(stat.user_id, records)
|
||||
except Exception as e:
|
||||
logger.error(f"rescan_achievement failed for user {stat.user_id}: {e}")
|
||||
@@ -1,8 +0,0 @@
|
||||
from django.urls import path
|
||||
|
||||
from achievement.views.admin import AchievementAdminAPI, AchievementMetricAdminAPI
|
||||
|
||||
urlpatterns = [
|
||||
path("achievement", AchievementAdminAPI.as_view(), name="achievement_admin_api"),
|
||||
path("achievement/metrics", AchievementMetricAdminAPI.as_view(), name="achievement_metric_admin_api"),
|
||||
]
|
||||
@@ -1,9 +0,0 @@
|
||||
from django.urls import path
|
||||
|
||||
from achievement.views.oj import AchievementListAPI, AchievementPendingAPI, AchievementSummaryAPI
|
||||
|
||||
urlpatterns = [
|
||||
path("achievements", AchievementListAPI.as_view(), name="achievement_list_api"),
|
||||
path("achievements/summary", AchievementSummaryAPI.as_view(), name="achievement_summary_api"),
|
||||
path("achievements/pending", AchievementPendingAPI.as_view(), name="achievement_pending_api"),
|
||||
]
|
||||
@@ -1,126 +0,0 @@
|
||||
from account.decorators import super_admin_required
|
||||
from achievement.metrics import METRIC_REGISTRY
|
||||
from achievement.models import Achievement, Rarity
|
||||
from achievement.tasks import rescan_achievement
|
||||
from utils.api import APIView
|
||||
from utils.shortcuts import check_is_id
|
||||
|
||||
|
||||
class AchievementAdminAPI(APIView):
|
||||
@super_admin_required
|
||||
def get(self, request):
|
||||
achievement_id = request.GET.get("id")
|
||||
if achievement_id:
|
||||
try:
|
||||
achievement = Achievement.objects.get(id=achievement_id)
|
||||
except Achievement.DoesNotExist:
|
||||
return self.error("成就不存在")
|
||||
return self.success(_serialize(achievement))
|
||||
return self.success([_serialize(a) for a in Achievement.objects.all()])
|
||||
|
||||
@super_admin_required
|
||||
def post(self, request):
|
||||
data = request.data
|
||||
error = _validate(data)
|
||||
if error:
|
||||
return self.error(error)
|
||||
achievement = Achievement.objects.create(
|
||||
name=data["name"],
|
||||
description=data["description"],
|
||||
icon=data["icon"],
|
||||
rarity=data["rarity"],
|
||||
hidden=data.get("hidden", False),
|
||||
metric=data["metric"],
|
||||
operator=data["operator"],
|
||||
threshold=data["threshold"],
|
||||
visible=data.get("visible", True),
|
||||
order=data.get("order", 0),
|
||||
)
|
||||
# 新建成就需要补发给已达标的存量用户
|
||||
rescan_achievement.send(achievement.id)
|
||||
return self.success(_serialize(achievement))
|
||||
|
||||
@super_admin_required
|
||||
def put(self, request):
|
||||
data = request.data
|
||||
if not check_is_id(data.get("id")):
|
||||
return self.error("参数错误")
|
||||
try:
|
||||
achievement = Achievement.objects.get(id=data["id"])
|
||||
except Achievement.DoesNotExist:
|
||||
return self.error("成就不存在")
|
||||
error = _validate(data)
|
||||
if error:
|
||||
return self.error(error)
|
||||
|
||||
before = (achievement.metric, achievement.operator, achievement.threshold, achievement.visible)
|
||||
for field in ("name", "description", "icon", "rarity", "hidden", "metric", "operator", "threshold", "visible", "order"):
|
||||
if field in data:
|
||||
setattr(achievement, field, data[field])
|
||||
achievement.save()
|
||||
|
||||
# 只要"谁能达成"这件事可能变了就补发,不去精细判断是否放宽。
|
||||
# 补发是幂等的后台任务(unlock 用 get_or_create),多跑一次只花一次扫描;
|
||||
# 漏跑却是学生已达标却拿不到,两个方向代价不对称。
|
||||
# 早先的 loosened 谓词只看 operator/threshold,会漏掉两种情况:
|
||||
# 换了 metric(换了维度)、以及从下架改成上架(草稿期已达标的人)。
|
||||
after = (achievement.metric, achievement.operator, achievement.threshold, achievement.visible)
|
||||
if achievement.visible and before != after:
|
||||
rescan_achievement.send(achievement.id)
|
||||
return self.success(_serialize(achievement))
|
||||
|
||||
@super_admin_required
|
||||
def delete(self, request):
|
||||
achievement_id = request.GET.get("id")
|
||||
if not check_is_id(achievement_id):
|
||||
return self.error("参数错误")
|
||||
Achievement.objects.filter(id=achievement_id).delete()
|
||||
return self.success("删除成功")
|
||||
|
||||
|
||||
class AchievementMetricAdminAPI(APIView):
|
||||
@super_admin_required
|
||||
def get(self, request):
|
||||
"""供后台指标下拉框使用。这里的列表就是代码里注册了什么。"""
|
||||
return self.success([{"key": key, "name": m.name, "help_text": m.help_text} for key, m in METRIC_REGISTRY.items()])
|
||||
|
||||
|
||||
def _validate(data):
|
||||
for field in ("name", "description", "icon", "rarity", "metric", "operator"):
|
||||
if not data.get(field):
|
||||
return f"{field} 不能为空"
|
||||
if data["metric"] not in METRIC_REGISTRY:
|
||||
return "指标不存在"
|
||||
if data["operator"] not in ("gte", "lte"):
|
||||
return "比较符不合法"
|
||||
# rarity 不校验的话,一个乱填的值会让 AchievementSummaryAPI 的四档统计
|
||||
# 对不上:它按 Rarity.choices 遍历,野值算进总数却不出现在任何一档里
|
||||
if data["rarity"] not in Rarity.values:
|
||||
return "稀有度不合法"
|
||||
if not isinstance(data.get("threshold"), int):
|
||||
return "阈值必须是整数"
|
||||
return None
|
||||
|
||||
|
||||
def _serialize(a):
|
||||
return {
|
||||
"id": a.id,
|
||||
"name": a.name,
|
||||
"description": a.description,
|
||||
"icon": a.icon,
|
||||
"rarity": a.rarity,
|
||||
"hidden": a.hidden,
|
||||
"metric": a.metric,
|
||||
"metric_name": METRIC_REGISTRY[a.metric].name if a.metric in METRIC_REGISTRY else a.metric,
|
||||
"operator": a.operator,
|
||||
"threshold": a.threshold,
|
||||
"visible": a.visible,
|
||||
# 后台列表必须显示这个:阈值配错时学生永远拿不到也永远不会来问,
|
||||
# 这个计数器是唯一的仪表盘
|
||||
"unlock_count": a.unlock_count,
|
||||
"order": a.order,
|
||||
# 必须 isoformat():这里是手写的 dict 而不是 DRF 序列化器,
|
||||
# 原始 datetime 交给 JSON 编码器会抛 TypeError,整个管理接口 500。
|
||||
# 表为空时列表接口看着正常(不进循环),一旦有数据就全挂。
|
||||
"create_time": a.create_time.isoformat() if a.create_time else None,
|
||||
}
|
||||
@@ -1,100 +0,0 @@
|
||||
from django.db.models import Count
|
||||
|
||||
from account.decorators import login_required
|
||||
from account.models import User
|
||||
from achievement.models import Achievement, Rarity, UserAchievement, UserStat
|
||||
from achievement.serializers import AchievementSerializer, PendingAchievementSerializer
|
||||
from utils.api import APIView
|
||||
|
||||
|
||||
def _resolve_user(request):
|
||||
"""?name=<username> 指定他人,不传则为自己。与 /api/profile 的约定一致。"""
|
||||
username = request.GET.get("name")
|
||||
if username:
|
||||
return User.objects.filter(username=username, is_disabled=False).first()
|
||||
return request.user
|
||||
|
||||
|
||||
def _decorate(achievements, unlocked_map, metrics):
|
||||
"""给成就对象挂上该用户的解锁状态与进度,供序列化器读取。"""
|
||||
for a in achievements:
|
||||
record = unlocked_map.get(a.id)
|
||||
a.unlocked = record is not None
|
||||
a.unlock_time = record.unlock_time if record else None
|
||||
a.backfilled = record.backfilled if record else False
|
||||
value = metrics.get(a.metric)
|
||||
# 隐藏且未解锁的成就不下发进度,否则能反推出条件
|
||||
a.progress = None if (a.hidden and not a.unlocked) else (value or 0)
|
||||
return achievements
|
||||
|
||||
|
||||
class AchievementListAPI(APIView):
|
||||
@login_required
|
||||
def get(self, request):
|
||||
user = _resolve_user(request)
|
||||
if user is None:
|
||||
return self.error("用户不存在")
|
||||
|
||||
achievements = list(Achievement.objects.filter(visible=True))
|
||||
unlocked_map = {r.achievement_id: r for r in UserAchievement.objects.filter(user=user)}
|
||||
stat = UserStat.objects.filter(user=user).first()
|
||||
metrics = stat.metrics if stat else {}
|
||||
|
||||
_decorate(achievements, unlocked_map, metrics)
|
||||
return self.success(
|
||||
{
|
||||
"username": user.username,
|
||||
"achievements": AchievementSerializer(achievements, many=True).data,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class AchievementSummaryAPI(APIView):
|
||||
@login_required
|
||||
def get(self, request):
|
||||
user = _resolve_user(request)
|
||||
if user is None:
|
||||
return self.error("用户不存在")
|
||||
|
||||
total_by_rarity = dict(Achievement.objects.filter(visible=True).values_list("rarity").annotate(c=Count("id")))
|
||||
unlocked_by_rarity = dict(UserAchievement.objects.filter(user=user, achievement__visible=True).values_list("achievement__rarity").annotate(c=Count("id")))
|
||||
total = sum(total_by_rarity.values())
|
||||
unlocked = sum(unlocked_by_rarity.values())
|
||||
|
||||
recent = list(UserAchievement.objects.filter(user=user, achievement__visible=True).select_related("achievement").order_by("-unlock_time")[:10])
|
||||
|
||||
return self.success(
|
||||
{
|
||||
"username": user.username,
|
||||
"total": total,
|
||||
"unlocked": unlocked,
|
||||
"percent": round(unlocked / total * 100, 1) if total else 0.0,
|
||||
"rarity": [
|
||||
{
|
||||
"rarity": value,
|
||||
"label": label,
|
||||
"total": total_by_rarity.get(value, 0),
|
||||
"unlocked": unlocked_by_rarity.get(value, 0),
|
||||
}
|
||||
for value, label in Rarity.choices
|
||||
],
|
||||
"recent": PendingAchievementSerializer(recent, many=True).data,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class AchievementPendingAPI(APIView):
|
||||
@login_required
|
||||
def get(self, request):
|
||||
"""返回尚未弹过的解锁记录。前端在布局层路由切换时拉取。"""
|
||||
records = UserAchievement.objects.filter(user=request.user, notified=False, achievement__visible=True).select_related("achievement").order_by("unlock_time")
|
||||
return self.success(PendingAchievementSerializer(records, many=True).data)
|
||||
|
||||
@login_required
|
||||
def post(self, request):
|
||||
"""弹完后标记已读,避免下次导航重复弹。body: {"ids": [成就 id]}"""
|
||||
ids = request.data.get("ids") or []
|
||||
if not isinstance(ids, list):
|
||||
return self.error("参数错误")
|
||||
UserAchievement.objects.filter(user=request.user, achievement_id__in=ids).update(notified=True)
|
||||
return self.success("ok")
|
||||
@@ -1,6 +0,0 @@
|
||||
from django.apps import AppConfig
|
||||
|
||||
|
||||
class AiConfig(AppConfig):
|
||||
default_auto_field = 'django.db.models.BigAutoField'
|
||||
name = 'ai'
|
||||
@@ -1,34 +0,0 @@
|
||||
# Generated by Django 5.2.3 on 2025-09-24 12:59
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.conf import settings
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
initial = True
|
||||
|
||||
dependencies = [
|
||||
migrations.swappable_dependency(settings.AUTH_USER_MODEL),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='AIAnalysis',
|
||||
fields=[
|
||||
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('provider', models.TextField(default='deepseek')),
|
||||
('data', models.JSONField()),
|
||||
('system_prompt', models.TextField()),
|
||||
('user_prompt', models.TextField()),
|
||||
('analysis', models.TextField()),
|
||||
('create_time', models.DateTimeField(auto_now_add=True)),
|
||||
('user', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to=settings.AUTH_USER_MODEL)),
|
||||
],
|
||||
options={
|
||||
'db_table': 'ai_analysis',
|
||||
'ordering': ['-create_time'],
|
||||
},
|
||||
),
|
||||
]
|
||||
@@ -1,18 +0,0 @@
|
||||
# Generated by Django 5.2.3 on 2025-09-24 13:02
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('ai', '0001_initial'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='aianalysis',
|
||||
name='model',
|
||||
field=models.TextField(default='deepseek-chat'),
|
||||
),
|
||||
]
|
||||
@@ -1,18 +0,0 @@
|
||||
# Generated by Django 6.0 on 2026-04-27 12:31
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('ai', '0002_aianalysis_model'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AlterField(
|
||||
model_name='aianalysis',
|
||||
name='model',
|
||||
field=models.TextField(default='deepseek-v4-flash'),
|
||||
),
|
||||
]
|
||||
@@ -1,18 +0,0 @@
|
||||
# Generated by Django 6.0.4 on 2026-06-04 14:39
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('ai', '0003_alter_aianalysis_model'),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddField(
|
||||
model_name='aianalysis',
|
||||
name='is_pinned',
|
||||
field=models.BooleanField(default=False),
|
||||
),
|
||||
]
|
||||
19
ai/models.py
19
ai/models.py
@@ -1,19 +0,0 @@
|
||||
from django.db import models
|
||||
|
||||
from account.models import User
|
||||
|
||||
|
||||
class AIAnalysis(models.Model):
|
||||
user = models.ForeignKey(User, on_delete=models.CASCADE)
|
||||
provider = models.TextField(default="deepseek")
|
||||
model = models.TextField(default="deepseek-v4-flash")
|
||||
data = models.JSONField()
|
||||
system_prompt = models.TextField()
|
||||
user_prompt = models.TextField()
|
||||
analysis = models.TextField()
|
||||
is_pinned = models.BooleanField(default=False)
|
||||
create_time = models.DateTimeField(auto_now_add=True)
|
||||
|
||||
class Meta:
|
||||
db_table = "ai_analysis"
|
||||
ordering = ["-create_time"]
|
||||
@@ -1,27 +0,0 @@
|
||||
from utils.api import serializers
|
||||
|
||||
from .models import AIAnalysis
|
||||
|
||||
|
||||
class AIAnalysisListSerializer(serializers.ModelSerializer):
|
||||
username = serializers.CharField(source="user.username")
|
||||
analysis_excerpt = serializers.SerializerMethodField()
|
||||
|
||||
class Meta:
|
||||
model = AIAnalysis
|
||||
fields = ["id", "create_time", "username", "analysis_excerpt", "is_pinned"]
|
||||
|
||||
def get_analysis_excerpt(self, obj):
|
||||
if not obj.analysis:
|
||||
return ""
|
||||
text = " ".join(obj.analysis.split())
|
||||
return text[:120] if len(text) <= 120 else text[:120] + "…"
|
||||
|
||||
|
||||
class AIAnalysisDetailSerializer(serializers.ModelSerializer):
|
||||
username = serializers.CharField(source="user.username")
|
||||
class_name = serializers.CharField(source="user.class_name")
|
||||
|
||||
class Meta:
|
||||
model = AIAnalysis
|
||||
fields = ["id", "create_time", "username", "class_name", "analysis"]
|
||||
@@ -1,7 +0,0 @@
|
||||
from django.urls import path
|
||||
|
||||
from ..views.admin import AIAnalysisAdminAPI
|
||||
|
||||
urlpatterns = [
|
||||
path("ai/reports", AIAnalysisAdminAPI.as_view()),
|
||||
]
|
||||
@@ -1,25 +0,0 @@
|
||||
from django.urls import path
|
||||
|
||||
from ..views.oj import (
|
||||
AIAnalysisAPI,
|
||||
AIDetailDataAPI,
|
||||
AIDurationDataAPI,
|
||||
AIHeatmapDataAPI,
|
||||
AIHintAPI,
|
||||
AILoginSummaryAPI,
|
||||
AIPinnedReportAPI,
|
||||
ClassPKAnalysisAPI,
|
||||
SingleClassAnalysisAPI,
|
||||
)
|
||||
|
||||
urlpatterns = [
|
||||
path("ai/detail", AIDetailDataAPI.as_view()),
|
||||
path("ai/duration", AIDurationDataAPI.as_view()),
|
||||
path("ai/analysis", AIAnalysisAPI.as_view()),
|
||||
path("ai/hint", AIHintAPI.as_view()),
|
||||
path("ai/heatmap", AIHeatmapDataAPI.as_view()),
|
||||
path("ai/login_summary", AILoginSummaryAPI.as_view()),
|
||||
path("ai/pinned", AIPinnedReportAPI.as_view()),
|
||||
path("ai/class_pk", ClassPKAnalysisAPI.as_view()),
|
||||
path("ai/class_single", SingleClassAnalysisAPI.as_view()),
|
||||
]
|
||||
@@ -1,45 +0,0 @@
|
||||
from account.decorators import teacher_admin_required
|
||||
from utils.api import APIView
|
||||
|
||||
from ..models import AIAnalysis
|
||||
from ..serializers import AIAnalysisDetailSerializer, AIAnalysisListSerializer
|
||||
|
||||
|
||||
class AIAnalysisAdminAPI(APIView):
|
||||
@teacher_admin_required
|
||||
def get(self, request):
|
||||
report_id = request.GET.get("id")
|
||||
if report_id:
|
||||
try:
|
||||
report = AIAnalysis.objects.select_related("user").get(id=report_id)
|
||||
except AIAnalysis.DoesNotExist:
|
||||
return self.error("AIAnalysis not found")
|
||||
return self.success(AIAnalysisDetailSerializer(report).data)
|
||||
|
||||
qs = AIAnalysis.objects.select_related("user").order_by("-create_time")
|
||||
|
||||
username = request.GET.get("username")
|
||||
if username:
|
||||
qs = qs.filter(user__username__icontains=username)
|
||||
|
||||
if request.GET.get("pinned_only") == "true":
|
||||
pinned = qs.filter(is_pinned=True)
|
||||
return self.success(AIAnalysisListSerializer(pinned, many=True).data)
|
||||
|
||||
return self.success(self.paginate_data(request, qs, AIAnalysisListSerializer))
|
||||
|
||||
@teacher_admin_required
|
||||
def post(self, request):
|
||||
report_id = request.data.get("id")
|
||||
try:
|
||||
report = AIAnalysis.objects.select_related("user").get(id=report_id)
|
||||
except AIAnalysis.DoesNotExist:
|
||||
return self.error("AIAnalysis not found")
|
||||
|
||||
if report.is_pinned:
|
||||
report.is_pinned = False
|
||||
else:
|
||||
AIAnalysis.objects.filter(user=report.user, is_pinned=True).update(is_pinned=False)
|
||||
report.is_pinned = True
|
||||
report.save(update_fields=["is_pinned"])
|
||||
return self.success({"is_pinned": report.is_pinned})
|
||||
907
ai/views/oj.py
907
ai/views/oj.py
@@ -1,907 +0,0 @@
|
||||
import calendar
|
||||
import hashlib
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
from django.core.cache import cache
|
||||
from django.db.models import Count, Min
|
||||
from django.db.models.functions import TruncDate
|
||||
from django.http import StreamingHttpResponse
|
||||
from django.utils import timezone
|
||||
from django.utils.dateparse import parse_datetime
|
||||
|
||||
from account.decorators import login_required, teacher_admin_required
|
||||
from account.models import User
|
||||
from ai.models import AIAnalysis
|
||||
from ai.serializers import AIAnalysisDetailSerializer
|
||||
from flowchart.models import FlowchartSubmission, FlowchartSubmissionStatus
|
||||
from problem.models import Problem
|
||||
from submission.models import JudgeStatus, Submission
|
||||
from utils.api import APIView
|
||||
from utils.openai import get_ai_client, get_async_ai_client
|
||||
from utils.shortcuts import datetime2str
|
||||
|
||||
CACHE_TIMEOUT = 300
|
||||
DIFFICULTY_MAP = {"Low": "简单", "Mid": "中等", "High": "困难"}
|
||||
DEFAULT_CLASS_SIZE = 45
|
||||
|
||||
# 评级阈值配置:(百分位上限, 评级)
|
||||
GRADE_THRESHOLDS = [
|
||||
(10, "S"), # 前10%: S级 - 卓越
|
||||
(35, "A"), # 前35%: A级 - 优秀
|
||||
(75, "B"), # 前75%: B级 - 良好
|
||||
(100, "C"), # 其余: C级 - 及格
|
||||
]
|
||||
|
||||
# 小规模参与惩罚配置:(最小人数, 等级降级映射)
|
||||
SMALL_SCALE_PENALTY = {
|
||||
"threshold": 10,
|
||||
"downgrade": {"S": "A", "A": "B"},
|
||||
}
|
||||
|
||||
# 等级权重映射(用于加权平均计算)
|
||||
GRADE_WEIGHTS = {"S": 4, "A": 3, "B": 2, "C": 1}
|
||||
|
||||
# 平均等级阈值:(最小权重, 等级)
|
||||
AVERAGE_GRADE_THRESHOLDS = [(3.5, "S"), (2.5, "A"), (1.5, "B")]
|
||||
|
||||
|
||||
def shift_months(dt, months):
|
||||
"""按月平移。落到不存在的日期时收缩到当月最后一天(1月31日 + 1个月 = 2月28/29日)"""
|
||||
month_index = dt.month - 1 + months
|
||||
year = dt.year + month_index // 12
|
||||
month = month_index % 12 + 1
|
||||
day = min(dt.day, calendar.monthrange(year, month)[1])
|
||||
return dt.replace(year=year, month=month, day=day)
|
||||
|
||||
|
||||
def get_cache_key(prefix, *args):
|
||||
return hashlib.md5(f"{prefix}:{'_'.join(map(str, args))}".encode()).hexdigest()
|
||||
|
||||
|
||||
def get_difficulty(difficulty):
|
||||
return DIFFICULTY_MAP.get(difficulty, "中等")
|
||||
|
||||
|
||||
def get_grade(rank, submission_count, reference_count=None):
|
||||
"""
|
||||
计算题目完成评级
|
||||
|
||||
评级标准:
|
||||
- S级:前10%,卓越水平(10%的人)
|
||||
- A级:前35%,优秀水平(25%的人)
|
||||
- B级:前75%,良好水平(40%的人)
|
||||
- C级:75%之后,及格水平(25%的人)
|
||||
|
||||
特殊规则:
|
||||
- 小规模惩罚用 reference_count(全时段人数)判断,避免同期窗口窄导致惩罚误触发
|
||||
- reference_count 未传时退化为 submission_count
|
||||
"""
|
||||
if not rank or rank <= 0 or submission_count <= 0:
|
||||
return "C"
|
||||
|
||||
percentile = (rank - 1) / submission_count * 100
|
||||
|
||||
base_grade = "C"
|
||||
for threshold, grade in GRADE_THRESHOLDS:
|
||||
if percentile < threshold:
|
||||
base_grade = grade
|
||||
break
|
||||
|
||||
penalty_count = reference_count if reference_count is not None else submission_count
|
||||
if penalty_count < SMALL_SCALE_PENALTY["threshold"]:
|
||||
base_grade = SMALL_SCALE_PENALTY["downgrade"].get(base_grade, base_grade)
|
||||
|
||||
return base_grade
|
||||
|
||||
|
||||
def calculate_average_grade(grades):
|
||||
"""根据等级列表计算加权平均等级"""
|
||||
scores = [GRADE_WEIGHTS[g] for g in grades if g in GRADE_WEIGHTS]
|
||||
if not scores:
|
||||
return ""
|
||||
avg = sum(scores) / len(scores)
|
||||
for threshold, grade in AVERAGE_GRADE_THRESHOLDS:
|
||||
if avg >= threshold:
|
||||
return grade
|
||||
return "C"
|
||||
|
||||
|
||||
def find_user_rank(ranking_list, user_id):
|
||||
"""在排名列表中找到用户的排名(1-based),未找到返回 None"""
|
||||
return next(
|
||||
(idx + 1 for idx, rec in enumerate(ranking_list) if rec["user_id"] == user_id),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def get_class_user_ids(user):
|
||||
if not user.class_name:
|
||||
return []
|
||||
|
||||
cache_key = get_cache_key("class_users", user.class_name)
|
||||
user_ids = cache.get(cache_key)
|
||||
if user_ids is None:
|
||||
user_ids = list(User.objects.filter(class_name=user.class_name).values_list("id", flat=True))
|
||||
cache.set(cache_key, user_ids, CACHE_TIMEOUT)
|
||||
return user_ids
|
||||
|
||||
|
||||
def get_user_first_ac_submissions(user_id, start, end, class_user_ids=None, use_class_scope=False, include_all_time=True):
|
||||
# 用户自己的 AC 记录按时间范围过滤
|
||||
user_first_ac = list(
|
||||
Submission.objects.filter(
|
||||
result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED],
|
||||
user_id=user_id,
|
||||
create_time__gte=start,
|
||||
create_time__lte=end,
|
||||
)
|
||||
.values("problem_id")
|
||||
.annotate(first_ac_time=Min("create_time"))
|
||||
)
|
||||
if not user_first_ac:
|
||||
return [], {}, []
|
||||
|
||||
problem_ids = [item["problem_id"] for item in user_first_ac]
|
||||
|
||||
if not include_all_time:
|
||||
return user_first_ac, {}, problem_ids
|
||||
|
||||
# 排名基于全局数据(不限时间),后注册的学生与所有人公平竞争
|
||||
rank_qs = Submission.objects.filter(
|
||||
result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED],
|
||||
problem_id__in=problem_ids,
|
||||
)
|
||||
if use_class_scope and class_user_ids:
|
||||
rank_qs = rank_qs.filter(user_id__in=class_user_ids)
|
||||
|
||||
ranked_first_ac = list(rank_qs.values("user_id", "problem_id").annotate(first_ac_time=Min("create_time")))
|
||||
|
||||
by_problem = defaultdict(list)
|
||||
for item in ranked_first_ac:
|
||||
by_problem[item["problem_id"]].append(item)
|
||||
for submissions in by_problem.values():
|
||||
submissions.sort(key=lambda x: (x["first_ac_time"], x["user_id"]))
|
||||
|
||||
return user_first_ac, by_problem, problem_ids
|
||||
|
||||
|
||||
async def stream_ai_response(client, system_prompt, user_prompt, on_complete=None):
|
||||
"""SSE 流式响应异步生成器,on_complete(full_text) 在流结束时调用。
|
||||
|
||||
必须是异步生成器:在 ASGI 下,同步生成器会被 Django 的
|
||||
StreamingHttpResponse 通过 sync_to_async(list) 一次性消费完才发送,
|
||||
导致整段内容生成完毕后才返回,失去流式效果。
|
||||
"""
|
||||
try:
|
||||
stream = await client.chat.completions.create(
|
||||
model="deepseek-v4-flash",
|
||||
messages=[
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_prompt},
|
||||
],
|
||||
stream=True,
|
||||
extra_body={"thinking": {"type": "disabled"}},
|
||||
)
|
||||
except Exception as exc:
|
||||
yield f"data: {json.dumps({'type': 'error', 'message': str(exc)})}\n\n"
|
||||
yield "event: end\n\n"
|
||||
return
|
||||
|
||||
yield "event: start\n\n"
|
||||
chunks = []
|
||||
try:
|
||||
async for chunk in stream:
|
||||
if not chunk.choices:
|
||||
continue
|
||||
choice = chunk.choices[0]
|
||||
if choice.finish_reason:
|
||||
if on_complete:
|
||||
await on_complete("".join(chunks).strip())
|
||||
yield f"data: {json.dumps({'type': 'done'})}\n\n"
|
||||
break
|
||||
content = choice.delta.content
|
||||
if content:
|
||||
chunks.append(content)
|
||||
yield f"data: {json.dumps({'type': 'delta', 'content': content})}\n\n"
|
||||
except Exception as exc:
|
||||
yield f"data: {json.dumps({'type': 'error', 'message': str(exc)})}\n\n"
|
||||
finally:
|
||||
yield "event: end\n\n"
|
||||
|
||||
|
||||
def make_sse_response(generator):
|
||||
"""创建 SSE StreamingHttpResponse"""
|
||||
response = StreamingHttpResponse(
|
||||
streaming_content=generator,
|
||||
content_type="text/event-stream",
|
||||
)
|
||||
response["Cache-Control"] = "no-cache"
|
||||
# 关闭反向代理(如 nginx)对流式响应的缓冲
|
||||
response["X-Accel-Buffering"] = "no"
|
||||
return response
|
||||
|
||||
|
||||
class AIDetailDataAPI(APIView):
|
||||
@login_required
|
||||
def get(self, request):
|
||||
start = request.GET.get("start")
|
||||
end = request.GET.get("end")
|
||||
username = request.GET.get("username")
|
||||
|
||||
if not start or not end:
|
||||
return self.error("参数 start 和 end 不能为空")
|
||||
if not parse_datetime(start):
|
||||
return self.error("start 格式无效,请使用 ISO 8601 格式")
|
||||
if not parse_datetime(end):
|
||||
return self.error("end 格式无效,请使用 ISO 8601 格式")
|
||||
|
||||
user = request.user
|
||||
if username and request.user.is_teacher_or_above():
|
||||
try:
|
||||
user = User.objects.get(username=username)
|
||||
except User.DoesNotExist:
|
||||
return self.error("User not found")
|
||||
|
||||
cache_key = get_cache_key("ai_detail", user.id, user.class_name or "", start, end)
|
||||
cached_result = cache.get(cache_key)
|
||||
if cached_result:
|
||||
return self.success(cached_result)
|
||||
|
||||
class_user_ids = get_class_user_ids(user)
|
||||
use_class_scope = bool(user.class_name) and len(class_user_ids) > 1
|
||||
user_first_ac, by_problem, problem_ids = get_user_first_ac_submissions(user.id, start, end, class_user_ids, use_class_scope)
|
||||
|
||||
# 同期排名:只统计时间窗口内解题的人
|
||||
by_problem_period = defaultdict(list)
|
||||
if problem_ids:
|
||||
period_qs = Submission.objects.filter(
|
||||
result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED],
|
||||
problem_id__in=problem_ids,
|
||||
create_time__gte=start,
|
||||
create_time__lte=end,
|
||||
)
|
||||
if use_class_scope and class_user_ids:
|
||||
period_qs = period_qs.filter(user_id__in=class_user_ids)
|
||||
for item in period_qs.values("user_id", "problem_id").annotate(first_ac_time=Min("create_time")):
|
||||
by_problem_period[item["problem_id"]].append(item)
|
||||
for lst in by_problem_period.values():
|
||||
lst.sort(key=lambda x: (x["first_ac_time"], x["user_id"]))
|
||||
|
||||
result = {
|
||||
"user": user.username,
|
||||
"class_name": user.class_name,
|
||||
"start": start,
|
||||
"end": end,
|
||||
"solved": [],
|
||||
"flowcharts": [],
|
||||
"grade": "",
|
||||
"tags": {},
|
||||
"difficulty": {},
|
||||
"contest_count": 0,
|
||||
}
|
||||
|
||||
if user_first_ac:
|
||||
problems = {p.id: p for p in Problem.objects.filter(id__in=problem_ids).select_related("contest").prefetch_related("tags")}
|
||||
solved, contest_ids = self._build_solved_records(user_first_ac, by_problem, by_problem_period, problems, user.id)
|
||||
# 查找 flowchart submissions
|
||||
flowcharts_query = FlowchartSubmission.objects.filter(
|
||||
user_id=user,
|
||||
status=FlowchartSubmissionStatus.COMPLETED,
|
||||
)
|
||||
|
||||
# 添加时间范围过滤
|
||||
if start:
|
||||
flowcharts_query = flowcharts_query.filter(create_time__gte=start)
|
||||
if end:
|
||||
flowcharts_query = flowcharts_query.filter(create_time__lte=end)
|
||||
|
||||
flowcharts = flowcharts_query.select_related("problem").only(
|
||||
"id",
|
||||
"create_time",
|
||||
"ai_score",
|
||||
"ai_grade",
|
||||
"problem___id",
|
||||
"problem__title",
|
||||
)
|
||||
|
||||
# 按problem分组
|
||||
problem_groups = defaultdict(list)
|
||||
for flowchart in flowcharts:
|
||||
problem_id = flowchart.problem._id
|
||||
problem_groups[problem_id].append(flowchart)
|
||||
|
||||
flowcharts_data = []
|
||||
for problem_id, submissions in problem_groups.items():
|
||||
if not submissions:
|
||||
continue
|
||||
|
||||
# 获取第一个提交的基本信息
|
||||
first_submission = submissions[0]
|
||||
|
||||
# 计算统计数据
|
||||
scores = [s.ai_score for s in submissions if s.ai_score is not None]
|
||||
times = [s.create_time for s in submissions]
|
||||
|
||||
# 找到最高分和对应的等级
|
||||
best_score = max(scores) if scores else 0
|
||||
best_submission = next((s for s in submissions if s.ai_score == best_score), submissions[0])
|
||||
best_grade = best_submission.ai_grade or ""
|
||||
|
||||
# 计算平均分
|
||||
avg_score = sum(scores) / len(scores) if scores else 0
|
||||
|
||||
# 最新提交时间
|
||||
latest_time = max(times) if times else first_submission.create_time
|
||||
|
||||
merged_item = {
|
||||
"problem__id": problem_id,
|
||||
"problem_title": first_submission.problem.title,
|
||||
"submission_count": len(submissions),
|
||||
"best_score": best_score,
|
||||
"best_grade": best_grade,
|
||||
"latest_submission_time": latest_time.isoformat() if latest_time else None,
|
||||
"avg_score": round(avg_score, 0),
|
||||
}
|
||||
|
||||
flowcharts_data.append(merged_item)
|
||||
|
||||
# 按最新提交时间排序
|
||||
flowcharts_data.sort(key=lambda x: x["latest_submission_time"] or "", reverse=True)
|
||||
|
||||
result.update(
|
||||
{
|
||||
"solved": solved,
|
||||
"flowcharts": flowcharts_data,
|
||||
"grade": calculate_average_grade([s["grade"] for s in solved]),
|
||||
"tags": self._calculate_top_tags(problems.values()),
|
||||
"difficulty": self._calculate_difficulty_distribution(problems.values()),
|
||||
"contest_count": len(set(contest_ids)),
|
||||
}
|
||||
)
|
||||
|
||||
cache.set(cache_key, result, CACHE_TIMEOUT)
|
||||
return self.success(result)
|
||||
|
||||
def _build_solved_records(self, user_first_ac, by_problem, by_problem_period, problems, user_id):
|
||||
solved, contest_ids = [], []
|
||||
for item in user_first_ac:
|
||||
pid = item["problem_id"]
|
||||
problem = problems.get(pid)
|
||||
if not problem:
|
||||
continue
|
||||
|
||||
ranking_list = by_problem.get(pid, [])
|
||||
rank = find_user_rank(ranking_list, user_id)
|
||||
|
||||
period_ranking_list = by_problem_period.get(pid, [])
|
||||
period_rank = find_user_rank(period_ranking_list, user_id)
|
||||
|
||||
if problem.contest_id:
|
||||
contest_ids.append(problem.contest_id)
|
||||
|
||||
solved.append(
|
||||
{
|
||||
"problem": {
|
||||
"display_id": problem._id,
|
||||
"title": problem.title,
|
||||
"contest_id": problem.contest_id,
|
||||
"contest_title": getattr(problem.contest, "title", ""),
|
||||
},
|
||||
"ac_time": timezone.localtime(item["first_ac_time"]).isoformat(),
|
||||
"rank": rank,
|
||||
"ac_count": len(ranking_list),
|
||||
"grade": get_grade(period_rank, len(period_ranking_list), reference_count=len(ranking_list)),
|
||||
"period_rank": period_rank,
|
||||
"period_ac_count": len(period_ranking_list),
|
||||
"difficulty": get_difficulty(problem.difficulty),
|
||||
}
|
||||
)
|
||||
|
||||
return sorted(solved, key=lambda x: x["ac_time"]), contest_ids
|
||||
|
||||
def _calculate_top_tags(self, problems):
|
||||
tags_counter = defaultdict(int)
|
||||
for problem in problems:
|
||||
for tag in problem.tags.all():
|
||||
if tag.name:
|
||||
tags_counter[tag.name] += 1
|
||||
return dict(sorted(tags_counter.items(), key=lambda x: x[1], reverse=True)[:5])
|
||||
|
||||
def _calculate_difficulty_distribution(self, problems):
|
||||
diff_counter = {"Low": 0, "Mid": 0, "High": 0}
|
||||
for problem in problems:
|
||||
diff_counter[problem.difficulty if problem.difficulty in diff_counter else "Mid"] += 1
|
||||
return {get_difficulty(k): v for k, v in sorted(diff_counter.items(), key=lambda x: x[1], reverse=True)}
|
||||
|
||||
|
||||
class AIDurationDataAPI(APIView):
|
||||
@login_required
|
||||
def get(self, request):
|
||||
end_iso = request.GET.get("end")
|
||||
duration = request.GET.get("duration")
|
||||
username = request.GET.get("username")
|
||||
|
||||
user = request.user
|
||||
if username and request.user.is_teacher_or_above():
|
||||
try:
|
||||
user = User.objects.get(username=username)
|
||||
except User.DoesNotExist:
|
||||
return self.error("User not found")
|
||||
|
||||
cache_key = get_cache_key("ai_duration", user.id, user.class_name or "", end_iso, duration)
|
||||
cached_result = cache.get(cache_key)
|
||||
if cached_result:
|
||||
return self.success(cached_result)
|
||||
|
||||
class_user_ids = get_class_user_ids(user)
|
||||
use_class_scope = bool(user.class_name) and len(class_user_ids) > 1
|
||||
time_config = self._parse_duration(duration)
|
||||
start = time_config["rewind"](datetime.fromisoformat(end_iso))
|
||||
|
||||
duration_data = []
|
||||
for i in range(time_config["show_count"]):
|
||||
start = time_config["advance"](start)
|
||||
period_end = time_config["advance"](start)
|
||||
|
||||
submission_count = Submission.objects.filter(user_id=user.id, create_time__gte=start, create_time__lte=period_end).count()
|
||||
|
||||
period_data = {
|
||||
"unit": time_config["show_unit"],
|
||||
"index": time_config["show_count"] - 1 - i,
|
||||
"start": start.isoformat(),
|
||||
"end": period_end.isoformat(),
|
||||
"problem_count": 0,
|
||||
"submission_count": submission_count,
|
||||
"grade": "",
|
||||
}
|
||||
|
||||
if submission_count > 0:
|
||||
user_first_ac, _, problem_ids = get_user_first_ac_submissions(
|
||||
user.id,
|
||||
start.isoformat(),
|
||||
period_end.isoformat(),
|
||||
class_user_ids,
|
||||
use_class_scope,
|
||||
include_all_time=False,
|
||||
)
|
||||
if user_first_ac:
|
||||
period_data["problem_count"] = len(problem_ids)
|
||||
by_problem_period = defaultdict(list)
|
||||
period_qs = Submission.objects.filter(
|
||||
result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED],
|
||||
problem_id__in=problem_ids,
|
||||
create_time__gte=start,
|
||||
create_time__lte=period_end,
|
||||
)
|
||||
if use_class_scope and class_user_ids:
|
||||
period_qs = period_qs.filter(user_id__in=class_user_ids)
|
||||
for row in period_qs.values("user_id", "problem_id").annotate(first_ac_time=Min("create_time")):
|
||||
by_problem_period[row["problem_id"]].append(row)
|
||||
for lst in by_problem_period.values():
|
||||
lst.sort(key=lambda x: (x["first_ac_time"], x["user_id"]))
|
||||
grades = [
|
||||
get_grade(
|
||||
find_user_rank(by_problem_period.get(item["problem_id"], []), user.id),
|
||||
len(by_problem_period.get(item["problem_id"], [])),
|
||||
)
|
||||
for item in user_first_ac
|
||||
]
|
||||
period_data["grade"] = calculate_average_grade(grades)
|
||||
|
||||
duration_data.append(period_data)
|
||||
|
||||
cache.set(cache_key, duration_data, CACHE_TIMEOUT)
|
||||
return self.success(duration_data)
|
||||
|
||||
def _parse_duration(self, duration):
|
||||
unit, count = duration.split(":")
|
||||
count = int(count)
|
||||
|
||||
# rewind 把结束时间倒推到区间起点,advance 前进一格。
|
||||
# 按月的档位不能用 timedelta(月长不固定),用 shift_months 处理
|
||||
configs = {
|
||||
("months", 2): {
|
||||
"show_count": 8,
|
||||
"show_unit": "weeks",
|
||||
"rewind": lambda dt: dt - timedelta(weeks=9),
|
||||
"advance": lambda dt: dt + timedelta(weeks=1),
|
||||
},
|
||||
("months", 6): {
|
||||
"show_count": 6,
|
||||
"show_unit": "months",
|
||||
"rewind": lambda dt: shift_months(dt, -7),
|
||||
"advance": lambda dt: shift_months(dt, 1),
|
||||
},
|
||||
("years", 1): {
|
||||
"show_count": 12,
|
||||
"show_unit": "months",
|
||||
"rewind": lambda dt: shift_months(dt, -13),
|
||||
"advance": lambda dt: shift_months(dt, 1),
|
||||
},
|
||||
}
|
||||
|
||||
return configs.get(
|
||||
(unit, count),
|
||||
{
|
||||
"show_count": 4,
|
||||
"show_unit": "weeks",
|
||||
"rewind": lambda dt: dt - timedelta(weeks=5),
|
||||
"advance": lambda dt: dt + timedelta(weeks=1),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class AILoginSummaryAPI(APIView):
|
||||
@login_required
|
||||
def get(self, request):
|
||||
user = request.user
|
||||
end_time = timezone.now()
|
||||
start_time = self._resolve_start_time(request, user, end_time)
|
||||
|
||||
problems_qs = Problem.objects.filter(
|
||||
create_time__gte=start_time,
|
||||
create_time__lte=end_time,
|
||||
contest_id__isnull=True,
|
||||
visible=True,
|
||||
)
|
||||
new_problem_count = problems_qs.count()
|
||||
|
||||
submissions_qs = Submission.objects.filter(user_id=user.id, create_time__gte=start_time, create_time__lte=end_time)
|
||||
submission_count = submissions_qs.count()
|
||||
accepted_count = submissions_qs.filter(result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED]).count()
|
||||
solved_count = submissions_qs.filter(result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED]).values("problem_id").distinct().count()
|
||||
flowchart_submission_count = FlowchartSubmission.objects.filter(user_id=user.id, create_time__gte=start_time, create_time__lte=end_time).count()
|
||||
|
||||
summary = {
|
||||
"start": datetime2str(start_time),
|
||||
"end": datetime2str(end_time),
|
||||
"new_problem_count": new_problem_count,
|
||||
"submission_count": submission_count,
|
||||
"accepted_count": accepted_count,
|
||||
"solved_count": solved_count,
|
||||
"flowchart_submission_count": flowchart_submission_count,
|
||||
}
|
||||
|
||||
analysis = ""
|
||||
analysis_error = ""
|
||||
if submission_count >= 3:
|
||||
analysis, analysis_error = self._get_ai_analysis(summary)
|
||||
|
||||
data = {"summary": summary, "analysis": analysis}
|
||||
if analysis_error:
|
||||
data["analysis_error"] = analysis_error
|
||||
return self.success(data)
|
||||
|
||||
def _resolve_start_time(self, request, user, end_time):
|
||||
start_raw = request.session.get("prev_login") or request.GET.get("start")
|
||||
start_time = parse_datetime(start_raw) if start_raw else None
|
||||
|
||||
if start_time and timezone.is_naive(start_time):
|
||||
start_time = timezone.make_aware(start_time, timezone.get_current_timezone())
|
||||
|
||||
if not start_time:
|
||||
if user.last_login and user.last_login < end_time:
|
||||
start_time = user.last_login
|
||||
elif user.create_time:
|
||||
start_time = user.create_time
|
||||
else:
|
||||
start_time = end_time - timedelta(days=7)
|
||||
|
||||
if start_time >= end_time:
|
||||
start_time = end_time - timedelta(days=1)
|
||||
|
||||
return start_time
|
||||
|
||||
def _get_ai_analysis(self, summary):
|
||||
try:
|
||||
client = get_ai_client()
|
||||
except Exception as exc:
|
||||
return "", str(exc)
|
||||
|
||||
system_prompt = "你是 OnlineJudge 的学习助教。请根据统计数据给出简短分析(1-2句),再给出一行结论,结论用“结论:”开头。"
|
||||
user_prompt = (
|
||||
f"时间范围:{summary['start']} 到 {summary['end']}\n"
|
||||
f"新题目数:{summary['new_problem_count']}\n"
|
||||
f"提交次数:{summary['submission_count']}\n"
|
||||
f"AC 次数:{summary['accepted_count']}\n"
|
||||
f"AC 题目数:{summary['solved_count']}\n"
|
||||
f"流程图提交数:{summary['flowchart_submission_count']}\n"
|
||||
)
|
||||
|
||||
try:
|
||||
completion = client.chat.completions.create(
|
||||
model="deepseek-v4-flash",
|
||||
messages=[
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": user_prompt},
|
||||
],
|
||||
extra_body={"thinking": {"type": "disabled"}},
|
||||
)
|
||||
except Exception as exc:
|
||||
return "", str(exc)
|
||||
|
||||
if not completion.choices:
|
||||
return "", ""
|
||||
|
||||
content = completion.choices[0].message.content or ""
|
||||
return content.strip(), ""
|
||||
|
||||
|
||||
class AIAnalysisAPI(APIView):
|
||||
@login_required
|
||||
def post(self, request):
|
||||
details = request.data.get("details")
|
||||
duration = request.data.get("duration")
|
||||
|
||||
client = get_async_ai_client()
|
||||
|
||||
system_prompt = (
|
||||
"你是一个风趣的编程老师,学生使用判题狗平台进行编程练习。"
|
||||
"请根据学生提供的详细数据和每周数据,给出用户的学习建议,最后写一句鼓励学生的话。"
|
||||
"请使用 markdown 格式输出,不要在代码块中输出。"
|
||||
)
|
||||
user_prompt = f"这段时间内的详细数据: {details}\n(其中部分字段含义是 flowcharts:流程图的提交,solved:代码的提交)\n每周或每月的数据: {duration}"
|
||||
|
||||
user = request.user
|
||||
|
||||
async def on_complete(full_text):
|
||||
await AIAnalysis.objects.acreate(
|
||||
user=user,
|
||||
provider="deepseek",
|
||||
model="deepseek-v4-flash",
|
||||
data={"details": details, "duration": duration},
|
||||
system_prompt=system_prompt,
|
||||
user_prompt="这段时间内的详细数据,每周或每月的数据。",
|
||||
analysis=full_text,
|
||||
)
|
||||
|
||||
return make_sse_response(stream_ai_response(client, system_prompt, user_prompt, on_complete))
|
||||
|
||||
|
||||
class ClassPKAnalysisAPI(APIView):
|
||||
@teacher_admin_required
|
||||
def post(self, request):
|
||||
comparisons = request.data.get("comparisons")
|
||||
time_range_label = request.data.get("time_range_label", "全部时间")
|
||||
|
||||
if not comparisons or len(comparisons) < 2:
|
||||
return self.error("至少需要2个班级的数据")
|
||||
|
||||
client = get_async_ai_client()
|
||||
|
||||
system_prompt = (
|
||||
"你是一位经验丰富的编程教育数据分析专家,专注于职业院校计算机编程教学效果评估。"
|
||||
"请根据在线评测系统(OJ)提供的班级对比数据,给出深入、专业的分析报告。"
|
||||
"请使用 markdown 格式输出,不要在代码块中输出。"
|
||||
)
|
||||
user_prompt = self._build_prompt(comparisons, time_range_label)
|
||||
|
||||
return make_sse_response(stream_ai_response(client, system_prompt, user_prompt))
|
||||
|
||||
def _build_prompt(self, comparisons, time_range_label):
|
||||
def fmt_class(name):
|
||||
return f"{name[:2]}计算机{name[2:]}班"
|
||||
|
||||
lines = [
|
||||
f"时间范围:{time_range_label}",
|
||||
"",
|
||||
"## 指标说明",
|
||||
"- 总AC数:班级累计通过的题目总数",
|
||||
'- 平均AC / 中位数AC:平均受尖子生影响,中位数反映"典型学生"的真实水平',
|
||||
"- 前10%/中间80%/后10%均值:梯队分层,判断是尖子班还是均衡班",
|
||||
"- 优秀率 / 及格率:达到优秀线、及格线的学生比例",
|
||||
"- 参与度:有过任意提交的学生比例(反映主动性)",
|
||||
"- AC率:提交通过率(反映代码编写质量)",
|
||||
"- 综合分:综合多项指标计算的总评分(满分100)",
|
||||
"",
|
||||
"## 班级数据",
|
||||
]
|
||||
|
||||
for i, c in enumerate(comparisons):
|
||||
class_display = fmt_class(c["class_name"])
|
||||
lines.append(f"\n### 第{i + 1}名:{class_display}(综合分 {c['composite_score']:.1f})")
|
||||
lines.append(f"- 人数:{c['user_count']}")
|
||||
lines.append(f"- 总AC数:{c['total_ac']},总提交数:{c['total_submission']},AC率:{c['ac_rate']:.1f}%")
|
||||
lines.append(f"- 平均AC:{c['avg_ac']:.2f},中位数AC:{c['median_ac']:.2f}")
|
||||
lines.append(f"- Q1:{c['q1_ac']:.2f},Q3:{c['q3_ac']:.2f},IQR(四分位距):{c['iqr']:.2f},标准差:{c['std_dev']:.2f}")
|
||||
lines.append(f"- 前10%均值:{c['top_10_avg']:.2f},中间80%均值:{c['middle_80_avg']:.2f},后10%均值:{c['bottom_10_avg']:.2f}")
|
||||
lines.append(f"- 优秀率:{c['excellent_rate']:.1f}%,及格率:{c['pass_rate']:.1f}%,参与度:{c['active_rate']:.1f}%")
|
||||
|
||||
if c.get("recent_total_ac") is not None:
|
||||
lines.append(
|
||||
f"- 时间段内新增:总AC {c['recent_total_ac']},"
|
||||
f"平均AC {c.get('recent_avg_ac', 0):.2f},"
|
||||
f"中位数AC {c.get('recent_median_ac', 0):.2f},"
|
||||
f"前10%平均 {c.get('recent_top_10_avg', 0):.2f},"
|
||||
f"活跃学生数 {c.get('recent_active_count', 0)}"
|
||||
)
|
||||
|
||||
lines += [
|
||||
"",
|
||||
"## 请从以下7个维度分析,输出中文报告:",
|
||||
"",
|
||||
"**1. 总体排名**:直接给出排名,说明各班综合分差距大小。",
|
||||
"",
|
||||
"**2. 参与积极性**:对比参与度和总提交数,谁的班学生更积极主动?",
|
||||
"",
|
||||
'**3. 典型学生水平**:重点用中位数AC数对比(而非平均值),分析谁班的"普通学生"更强。若均值明显高于中位数,说明均值被少数强者拉高,需指出。',
|
||||
"",
|
||||
'**4. 班级内部均衡性**:结合标准差、IQR、前10%与后10%差距,判断哪个班是"均衡型",哪个班是"两极型"。',
|
||||
"",
|
||||
"**5. 梯队深度对比**:对比各班前10%均值(尖子生天花板)和后10%均值(薄弱学生水平),分析各班在培养尖子生和帮扶后进生上的差异。",
|
||||
"",
|
||||
'**6. 代码提交质量**:对比AC率,是否有班级存在"凑提交次数但不思考"的问题?',
|
||||
"",
|
||||
"**7. 综合结论与建议**:用1句话明确说明胜负;对落后班级给出2~3条具体可操作的改进建议;点出领先班级1条值得借鉴的做法。",
|
||||
"",
|
||||
"分析对象是班级任课教师,语言专业但不过分学术。",
|
||||
]
|
||||
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
class SingleClassAnalysisAPI(APIView):
|
||||
@login_required
|
||||
def post(self, request):
|
||||
comparison = request.data.get("comparison")
|
||||
if not comparison:
|
||||
return self.error("缺少班级数据")
|
||||
|
||||
client = get_async_ai_client()
|
||||
|
||||
class_name = comparison.get("class_name", "")
|
||||
class_display = f"{class_name[:2]}计算机{class_name[2:]}班"
|
||||
|
||||
system_prompt = (
|
||||
"你是一位经验丰富的编程教育数据分析专家,专注于职业院校计算机编程教学效果评估。"
|
||||
"请根据在线评测系统(OJ)提供的班级数据,给出深入、专业的分析报告。"
|
||||
"请使用 markdown 格式输出,不要在代码块中输出。"
|
||||
)
|
||||
|
||||
lines = [
|
||||
f"## 班级:{class_display}",
|
||||
"",
|
||||
"## 指标说明",
|
||||
"- 总AC数:班级累计通过的题目总数",
|
||||
'- 平均AC / 中位数AC:平均受尖子生影响,中位数反映"典型学生"的真实水平',
|
||||
"- 前10%/中间80%/后10%均值:梯队分层",
|
||||
"- 优秀率 / 及格率:分别基于班级内部Q3/Q1阈值统计",
|
||||
"- 参与度:有过任意提交的学生比例",
|
||||
"- AC率:提交通过率",
|
||||
"- 综合分:综合多项指标计算的总评分(满分100)",
|
||||
"",
|
||||
"## 班级数据",
|
||||
f"- 人数:{comparison['user_count']}",
|
||||
f"- 总AC数:{comparison['total_ac']},总提交数:{comparison['total_submission']},AC率:{comparison['ac_rate']:.1f}%",
|
||||
f"- 平均AC:{comparison['avg_ac']:.2f},中位数AC:{comparison['median_ac']:.2f}",
|
||||
f"- Q1:{comparison['q1_ac']:.2f},Q3:{comparison['q3_ac']:.2f},IQR:{comparison['iqr']:.2f},标准差:{comparison['std_dev']:.2f}",
|
||||
f"- 前10%均值:{comparison['top_10_avg']:.2f},中间80%均值:{comparison['middle_80_avg']:.2f},后10%均值:{comparison['bottom_10_avg']:.2f}",
|
||||
f"- 优秀率:{comparison['excellent_rate']:.1f}%,及格率:{comparison['pass_rate']:.1f}%,参与度:{comparison['active_rate']:.1f}%",
|
||||
f"- 综合分:{comparison['composite_score']:.1f}",
|
||||
"",
|
||||
"## 请从以下5个维度分析,输出中文报告:",
|
||||
"",
|
||||
"**1. 整体水平**:基于总AC数、均值与中位数,评价班级的整体编程水平。若均值明显高于中位数,需指出被少数强者拉高的情况。",
|
||||
"",
|
||||
"**2. 参与积极性**:结合参与度和AC率,评价学生的学习主动性和代码质量。",
|
||||
"",
|
||||
'**3. 班级内部均衡性**:结合标准差、IQR、前10%与后10%差距,判断是"均衡型"还是"两极型"班级。',
|
||||
"",
|
||||
"**4. 梯队分析**:分析前10%(尖子生)、中间80%(中坚力量)、后10%(需帮扶学生)的水平差异,给出针对性建议。",
|
||||
"",
|
||||
"**5. 改进建议**:结合优秀率、及格率等数据,给出3条具体可操作的教学改进建议。最后用一句话鼓励这个班级。",
|
||||
"",
|
||||
"分析对象是班级任课教师,语言专业但不过分学术。",
|
||||
]
|
||||
|
||||
user_prompt = "\n".join(lines)
|
||||
|
||||
return make_sse_response(stream_ai_response(client, system_prompt, user_prompt))
|
||||
|
||||
|
||||
class AIPinnedReportAPI(APIView):
|
||||
@login_required
|
||||
def get(self, request):
|
||||
try:
|
||||
report = AIAnalysis.objects.get(user=request.user, is_pinned=True)
|
||||
except AIAnalysis.DoesNotExist:
|
||||
return self.success(None)
|
||||
return self.success(AIAnalysisDetailSerializer(report).data)
|
||||
|
||||
|
||||
class AIHintAPI(APIView):
|
||||
@login_required
|
||||
def post(self, request):
|
||||
submission_id = request.data.get("submission_id")
|
||||
if not submission_id:
|
||||
return self.error("submission_id is required")
|
||||
|
||||
try:
|
||||
submission = Submission.objects.get(id=submission_id, user_id=request.user.id)
|
||||
except Submission.DoesNotExist:
|
||||
return self.error("Submission not found")
|
||||
|
||||
problem = submission.problem
|
||||
client = get_async_ai_client()
|
||||
|
||||
# 获取参考答案(同语言优先,否则取第一个)
|
||||
answers = problem.answers or []
|
||||
ref_answer = next(
|
||||
(a["code"] for a in answers if a["language"] == submission.language),
|
||||
answers[0]["code"] if answers else "",
|
||||
)
|
||||
|
||||
system_prompt = (
|
||||
"你是编程助教。你知道题目的参考答案,请按照以下规则给学生提示:\n\n"
|
||||
"【核心规则】\n"
|
||||
"- 【绝对禁止】直接给出答案或核心算法代码,也不能暗示完整解法。\n"
|
||||
"- 提示要循序渐进:先指出问题所在的方向,再给出一个小的思考点,让学生自己推导。\n"
|
||||
"- 对照参考答案分析学生代码,找出最关键的一个问题重点提示,不要一次列出所有问题。\n\n"
|
||||
"【输入处理例外】\n"
|
||||
"- 如果学生的代码在【读取输入】部分有错误(例如:输入格式解析错误、"
|
||||
"未正确读取多组输入、split/scanf 使用有误等),"
|
||||
"则【直接给出正确的输入读取代码片段】,并解释为什么这样写。"
|
||||
"输入处理不属于算法核心,可以直接告诉学生。\n\n"
|
||||
"【回复格式】\n"
|
||||
"语气鼓励,使用 Markdown 格式,回复简洁(不超过6句话)。"
|
||||
)
|
||||
user_prompt = (
|
||||
f"题目:{problem.title}\n"
|
||||
f"题目描述:{problem.description[:500]}\n"
|
||||
f"参考答案(仅供你分析,不可透露给学生):\n```\n{ref_answer[:2000]}\n```\n"
|
||||
f"学生提交语言:{submission.language}\n"
|
||||
f"判题结果:{submission.result}\n"
|
||||
f"错误信息:{submission.statistic_info.get('err_info', '无')}\n"
|
||||
f"学生代码:\n```\n{submission.code[:2000]}\n```"
|
||||
)
|
||||
|
||||
return make_sse_response(stream_ai_response(client, system_prompt, user_prompt))
|
||||
|
||||
|
||||
class AIHeatmapDataAPI(APIView):
|
||||
@login_required
|
||||
def get(self, request):
|
||||
username = request.GET.get("username")
|
||||
user = request.user
|
||||
if username and request.user.is_teacher_or_above():
|
||||
try:
|
||||
user = User.objects.get(username=username)
|
||||
except User.DoesNotExist:
|
||||
return self.error("User not found")
|
||||
end = timezone.now()
|
||||
today = end.date().isoformat()
|
||||
cache_key = get_cache_key("ai_heatmap", user.id, user.class_name or "", today)
|
||||
cached_result = cache.get(cache_key)
|
||||
if cached_result:
|
||||
return self.success(cached_result)
|
||||
start = end - timedelta(days=365)
|
||||
|
||||
# 使用单次查询获取所有数据,按日期分组统计
|
||||
submission_counts = (
|
||||
Submission.objects.filter(user_id=user.id, create_time__gte=start, create_time__lte=end)
|
||||
.annotate(date=TruncDate("create_time"))
|
||||
.values("date")
|
||||
.annotate(count=Count("id"))
|
||||
.order_by("date")
|
||||
)
|
||||
|
||||
# 将查询结果转换为字典,便于快速查找
|
||||
submission_dict = {item["date"]: item["count"] for item in submission_counts}
|
||||
|
||||
# 生成365天的热力图数据
|
||||
heatmap_data = []
|
||||
current_date = start.date()
|
||||
for i in range(365):
|
||||
day_date = current_date + timedelta(days=i)
|
||||
submission_count = submission_dict.get(day_date, 0)
|
||||
heatmap_data.append(
|
||||
{
|
||||
"timestamp": int(datetime.combine(day_date, datetime.min.time()).timestamp() * 1000),
|
||||
"value": submission_count,
|
||||
}
|
||||
)
|
||||
|
||||
cache.set(cache_key, heatmap_data, CACHE_TIMEOUT)
|
||||
return self.success(heatmap_data)
|
||||
@@ -1,19 +0,0 @@
|
||||
# Generated by Django 6.0 on 2026-04-23 20:07
|
||||
|
||||
from django.conf import settings
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
dependencies = [
|
||||
('announcement', '0001_initial'),
|
||||
migrations.swappable_dependency(settings.AUTH_USER_MODEL),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.AddIndex(
|
||||
model_name='announcement',
|
||||
index=models.Index(fields=['visible', '-top', '-create_time'], name='announcement_list_idx'),
|
||||
),
|
||||
]
|
||||
@@ -17,10 +17,4 @@ class Announcement(models.Model):
|
||||
|
||||
class Meta:
|
||||
db_table = "announcement"
|
||||
ordering = (
|
||||
"-top",
|
||||
"-create_time",
|
||||
)
|
||||
indexes = [
|
||||
models.Index(fields=["visible", "-top", "-create_time"], name="announcement_list_idx"),
|
||||
]
|
||||
ordering = ("-top", "-create_time",)
|
||||
|
||||
@@ -25,7 +25,7 @@ class AnnouncementListSerializer(serializers.ModelSerializer):
|
||||
|
||||
class Meta:
|
||||
model = Announcement
|
||||
exclude = ["content"]
|
||||
exclude = ['content']
|
||||
|
||||
|
||||
class EditAnnouncementSerializer(serializers.Serializer):
|
||||
|
||||
48
announcement/tests.py
Normal file
48
announcement/tests.py
Normal file
@@ -0,0 +1,48 @@
|
||||
from utils.api.tests import APITestCase
|
||||
|
||||
from .models import Announcement
|
||||
|
||||
|
||||
class AnnouncementAdminTest(APITestCase):
|
||||
def setUp(self):
|
||||
self.user = self.create_super_admin()
|
||||
self.url = self.reverse("announcement_admin_api")
|
||||
|
||||
def test_announcement_list(self):
|
||||
response = self.client.get(self.url)
|
||||
self.assertSuccess(response)
|
||||
|
||||
def create_announcement(self):
|
||||
return self.client.post(self.url, data={"title": "test", "content": "test", "visible": True})
|
||||
|
||||
def test_create_announcement(self):
|
||||
resp = self.create_announcement()
|
||||
self.assertSuccess(resp)
|
||||
return resp
|
||||
|
||||
def test_edit_announcement(self):
|
||||
data = {"id": self.create_announcement().data["data"]["id"], "title": "ahaha", "content": "test content",
|
||||
"visible": False}
|
||||
resp = self.client.put(self.url, data=data)
|
||||
self.assertSuccess(resp)
|
||||
resp_data = resp.data["data"]
|
||||
self.assertEqual(resp_data["title"], "ahaha")
|
||||
self.assertEqual(resp_data["content"], "test content")
|
||||
self.assertEqual(resp_data["visible"], False)
|
||||
|
||||
def test_delete_announcement(self):
|
||||
id = self.test_create_announcement().data["data"]["id"]
|
||||
resp = self.client.delete(self.url + "?id=" + str(id))
|
||||
self.assertSuccess(resp)
|
||||
self.assertFalse(Announcement.objects.filter(id=id).exists())
|
||||
|
||||
|
||||
class AnnouncementAPITest(APITestCase):
|
||||
def setUp(self):
|
||||
self.user = self.create_super_admin()
|
||||
Announcement.objects.create(title="title", content="content", visible=True, created_by=self.user)
|
||||
self.url = self.reverse("announcement_api")
|
||||
|
||||
def test_get_announcement_list(self):
|
||||
resp = self.client.get(self.url)
|
||||
self.assertSuccess(resp)
|
||||
@@ -1,11 +1,12 @@
|
||||
from account.decorators import super_admin_required
|
||||
from utils.api import APIView, validate_serializer
|
||||
|
||||
from announcement.models import Announcement
|
||||
from announcement.serializers import (
|
||||
AnnouncementSerializer,
|
||||
CreateAnnouncementSerializer,
|
||||
EditAnnouncementSerializer,
|
||||
)
|
||||
from utils.api import APIView, validate_serializer
|
||||
|
||||
|
||||
class AnnouncementAdminAPI(APIView):
|
||||
@@ -52,7 +53,9 @@ class AnnouncementAdminAPI(APIView):
|
||||
announcement = Announcement.objects.all().order_by("-create_time")
|
||||
if request.GET.get("visible") == "true":
|
||||
announcement = announcement.filter(visible=True)
|
||||
return self.success(self.paginate_data(request, announcement, AnnouncementSerializer))
|
||||
return self.success(
|
||||
self.paginate_data(request, announcement, AnnouncementSerializer)
|
||||
)
|
||||
|
||||
@super_admin_required
|
||||
def delete(self, request):
|
||||
|
||||
@@ -1,19 +1,20 @@
|
||||
from utils.api import APIView
|
||||
|
||||
from announcement.models import Announcement
|
||||
from announcement.serializers import AnnouncementListSerializer, AnnouncementSerializer
|
||||
from utils.api import AsyncAPIView
|
||||
from announcement.serializers import AnnouncementSerializer, AnnouncementListSerializer
|
||||
|
||||
|
||||
class AnnouncementAPI(AsyncAPIView):
|
||||
async def get(self, request):
|
||||
class AnnouncementAPI(APIView):
|
||||
def get(self, request):
|
||||
id = request.GET.get("id")
|
||||
if id:
|
||||
try:
|
||||
announcement = await Announcement.objects.select_related("created_by").filter(id=id, visible=True).afirst()
|
||||
if announcement is None:
|
||||
raise Announcement.DoesNotExist
|
||||
return self.success(await self.async_serialize_data(AnnouncementSerializer, announcement))
|
||||
announcement = Announcement.objects.get(id=id, visible=True)
|
||||
return self.success(AnnouncementSerializer(announcement).data)
|
||||
except Announcement.DoesNotExist:
|
||||
return self.error("Announcement does not exist")
|
||||
|
||||
announcements = Announcement.objects.select_related("created_by").filter(visible=True)
|
||||
return self.success(await self.async_paginate_data(request, announcements, AnnouncementListSerializer))
|
||||
announcements = Announcement.objects.filter(visible=True)
|
||||
return self.success(
|
||||
self.paginate_data(request, announcements, AnnouncementListSerializer)
|
||||
)
|
||||
|
||||
@@ -1,35 +0,0 @@
|
||||
from tree_sitter import Parser
|
||||
|
||||
from .engines import get_engine
|
||||
from .mappings import get_language, get_mapping
|
||||
|
||||
|
||||
def check_ast(code: str, language: str, rules: list[dict]) -> tuple[bool, list[dict]]:
|
||||
if not rules:
|
||||
return True, []
|
||||
|
||||
ts_language = get_language(language)
|
||||
if ts_language is None:
|
||||
return True, []
|
||||
|
||||
mapping = get_mapping(language)
|
||||
|
||||
try:
|
||||
parser = Parser(ts_language)
|
||||
tree = parser.parse(code.encode("utf-8"))
|
||||
except Exception:
|
||||
return True, []
|
||||
|
||||
results = []
|
||||
all_passed = True
|
||||
for rule in rules:
|
||||
engine = get_engine(rule.get("engine", ""))
|
||||
if engine is None:
|
||||
continue
|
||||
errors = engine.check(tree, rule, language, mapping)
|
||||
passed = len(errors) == 0
|
||||
if not passed:
|
||||
all_passed = False
|
||||
results.append({"description": engine.describe(rule, language, mapping), "passed": passed})
|
||||
|
||||
return all_passed, results
|
||||
@@ -1,23 +0,0 @@
|
||||
from .function_call import CountFunctionCallEngine, MustCallFunctionEngine, MustNotCallFunctionEngine
|
||||
from .method_call import MustCallMethodEngine, MustNotCallMethodEngine
|
||||
from .nesting import MustHaveNestingEngine
|
||||
from .node_count import CountNodeEngine
|
||||
from .node_exists import MustExistNodeEngine, MustNotExistNodeEngine
|
||||
from .operator import MustUseOperatorEngine
|
||||
|
||||
ENGINES = {
|
||||
"must_exist_node": MustExistNodeEngine(),
|
||||
"must_not_exist_node": MustNotExistNodeEngine(),
|
||||
"count_node": CountNodeEngine(),
|
||||
"must_call_function": MustCallFunctionEngine(),
|
||||
"must_not_call_function": MustNotCallFunctionEngine(),
|
||||
"count_function_call": CountFunctionCallEngine(),
|
||||
"must_call_method": MustCallMethodEngine(),
|
||||
"must_not_call_method": MustNotCallMethodEngine(),
|
||||
"must_use_operator": MustUseOperatorEngine(),
|
||||
"must_have_nesting": MustHaveNestingEngine(),
|
||||
}
|
||||
|
||||
|
||||
def get_engine(name: str):
|
||||
return ENGINES.get(name)
|
||||
@@ -1,21 +0,0 @@
|
||||
class BaseEngine:
|
||||
@staticmethod
|
||||
def collect_nodes(node, node_type):
|
||||
results = []
|
||||
if node.type == node_type:
|
||||
results.append(node)
|
||||
for child in node.children:
|
||||
results.extend(BaseEngine.collect_nodes(child, node_type))
|
||||
return results
|
||||
|
||||
@staticmethod
|
||||
def has_node(node, node_type):
|
||||
if node.type == node_type:
|
||||
return True
|
||||
return any(BaseEngine.has_node(child, node_type) for child in node.children)
|
||||
|
||||
def check(self, tree, rule, language, mapping) -> list[str]:
|
||||
raise NotImplementedError
|
||||
|
||||
def describe(self, rule, language, mapping) -> str:
|
||||
raise NotImplementedError
|
||||
@@ -1,78 +0,0 @@
|
||||
from .base import BaseEngine
|
||||
|
||||
CALL_NODE_TYPES = {
|
||||
"Python3": "call",
|
||||
"C": "call_expression",
|
||||
}
|
||||
|
||||
|
||||
class _FunctionCallBase(BaseEngine):
|
||||
def _find_function_calls(self, root, func_name, language):
|
||||
call_type = CALL_NODE_TYPES.get(language, "call")
|
||||
calls = self.collect_nodes(root, call_type)
|
||||
matches = []
|
||||
for call in calls:
|
||||
func_node = call.child_by_field_name("function")
|
||||
if func_node and func_node.type == "identifier" and func_node.text.decode() == func_name:
|
||||
matches.append(call)
|
||||
return matches
|
||||
|
||||
|
||||
class MustCallFunctionEngine(_FunctionCallBase):
|
||||
def _message(self, rule):
|
||||
return rule.get("message") or f"必须调用 {rule['target']}()"
|
||||
|
||||
def check(self, tree, rule, language, mapping):
|
||||
if not self._find_function_calls(tree.root_node, rule["target"], language):
|
||||
return [self._message(rule)]
|
||||
return []
|
||||
|
||||
def describe(self, rule, language, mapping):
|
||||
return self._message(rule)
|
||||
|
||||
|
||||
class MustNotCallFunctionEngine(_FunctionCallBase):
|
||||
def _message(self, rule):
|
||||
return rule.get("message") or f"不能调用 {rule['target']}()"
|
||||
|
||||
def check(self, tree, rule, language, mapping):
|
||||
if self._find_function_calls(tree.root_node, rule["target"], language):
|
||||
return [self._message(rule)]
|
||||
return []
|
||||
|
||||
def describe(self, rule, language, mapping):
|
||||
return self._message(rule)
|
||||
|
||||
|
||||
class CountFunctionCallEngine(_FunctionCallBase):
|
||||
def _message(self, rule, count):
|
||||
target = rule["target"]
|
||||
exact = rule.get("exact")
|
||||
if exact is not None and count != exact:
|
||||
return rule.get("message") or f"{target}() 需要调用 {exact} 次,当前 {count} 次"
|
||||
min_count = rule.get("min")
|
||||
max_count = rule.get("max")
|
||||
if min_count is not None and count < min_count:
|
||||
return rule.get("message") or f"{target}() 至少调用 {min_count} 次,当前 {count} 次"
|
||||
if max_count is not None and count > max_count:
|
||||
return rule.get("message") or f"{target}() 至多调用 {max_count} 次,当前 {count} 次"
|
||||
return None
|
||||
|
||||
def check(self, tree, rule, language, mapping):
|
||||
count = len(self._find_function_calls(tree.root_node, rule["target"], language))
|
||||
msg = self._message(rule, count)
|
||||
return [msg] if msg else []
|
||||
|
||||
def describe(self, rule, language, mapping):
|
||||
target = rule["target"]
|
||||
if rule.get("message"):
|
||||
return rule["message"]
|
||||
exact = rule.get("exact")
|
||||
if exact is not None:
|
||||
return f"{target}() 调用 {exact} 次"
|
||||
parts = []
|
||||
if rule.get("min") is not None:
|
||||
parts.append(f"至少 {rule['min']} 次")
|
||||
if rule.get("max") is not None:
|
||||
parts.append(f"至多 {rule['max']} 次")
|
||||
return f"{target}() " + "、".join(parts)
|
||||
@@ -1,48 +0,0 @@
|
||||
from .base import BaseEngine
|
||||
|
||||
CALL_NODE_TYPES = {
|
||||
"Python3": "call",
|
||||
"C": "call_expression",
|
||||
}
|
||||
|
||||
|
||||
class _MethodCallBase(BaseEngine):
|
||||
def _find_method_calls(self, root, method_name, language):
|
||||
if language == "C":
|
||||
return []
|
||||
call_type = CALL_NODE_TYPES.get(language, "call")
|
||||
calls = self.collect_nodes(root, call_type)
|
||||
matches = []
|
||||
for call in calls:
|
||||
func_node = call.child_by_field_name("function")
|
||||
if func_node and func_node.type == "attribute":
|
||||
attr_node = func_node.child_by_field_name("attribute")
|
||||
if attr_node and attr_node.text.decode() == method_name:
|
||||
matches.append(call)
|
||||
return matches
|
||||
|
||||
|
||||
class MustCallMethodEngine(_MethodCallBase):
|
||||
def _message(self, rule):
|
||||
return rule.get("message") or f"必须调用 .{rule['target']}()"
|
||||
|
||||
def check(self, tree, rule, language, mapping):
|
||||
if not self._find_method_calls(tree.root_node, rule["target"], language):
|
||||
return [self._message(rule)]
|
||||
return []
|
||||
|
||||
def describe(self, rule, language, mapping):
|
||||
return self._message(rule)
|
||||
|
||||
|
||||
class MustNotCallMethodEngine(_MethodCallBase):
|
||||
def _message(self, rule):
|
||||
return rule.get("message") or f"不能调用 .{rule['target']}()"
|
||||
|
||||
def check(self, tree, rule, language, mapping):
|
||||
if self._find_method_calls(tree.root_node, rule["target"], language):
|
||||
return [self._message(rule)]
|
||||
return []
|
||||
|
||||
def describe(self, rule, language, mapping):
|
||||
return self._message(rule)
|
||||
@@ -1,34 +0,0 @@
|
||||
from ast_checker.labels import label
|
||||
|
||||
from .base import BaseEngine
|
||||
|
||||
|
||||
class MustHaveNestingEngine(BaseEngine):
|
||||
def _has_inner_in_subtree(self, node, inner_type):
|
||||
for child in node.children:
|
||||
if self.has_node(child, inner_type):
|
||||
return True
|
||||
return False
|
||||
|
||||
def _message(self, rule):
|
||||
if rule.get("message"):
|
||||
return rule["message"]
|
||||
outer = rule.get("outer", "")
|
||||
inner = rule.get("inner", "")
|
||||
outer_label = label(outer)
|
||||
inner_label = label(inner)
|
||||
if outer == inner:
|
||||
return f"必须使用 {outer_label} 嵌套"
|
||||
return f"必须在 {outer_label} 中嵌套使用 {inner_label}"
|
||||
|
||||
def check(self, tree, rule, language, mapping):
|
||||
outer_type = mapping.get(rule["outer"], rule["outer"])
|
||||
inner_type = mapping.get(rule["inner"], rule["inner"])
|
||||
outer_nodes = self.collect_nodes(tree.root_node, outer_type)
|
||||
for outer_node in outer_nodes:
|
||||
if self._has_inner_in_subtree(outer_node, inner_type):
|
||||
return []
|
||||
return [self._message(rule)]
|
||||
|
||||
def describe(self, rule, language, mapping):
|
||||
return self._message(rule)
|
||||
@@ -1,39 +0,0 @@
|
||||
from ast_checker.labels import label
|
||||
|
||||
from .base import BaseEngine
|
||||
|
||||
|
||||
class CountNodeEngine(BaseEngine):
|
||||
def _message(self, rule, count):
|
||||
name = rule.get("label") or label(rule["target"])
|
||||
exact = rule.get("exact")
|
||||
if exact is not None and count != exact:
|
||||
return rule.get("message") or f"{name} 需要出现 {exact} 次,当前 {count} 次"
|
||||
min_count = rule.get("min")
|
||||
max_count = rule.get("max")
|
||||
if min_count is not None and count < min_count:
|
||||
return rule.get("message") or f"{name} 至少出现 {min_count} 次,当前 {count} 次"
|
||||
if max_count is not None and count > max_count:
|
||||
return rule.get("message") or f"{name} 至多出现 {max_count} 次,当前 {count} 次"
|
||||
return None
|
||||
|
||||
def check(self, tree, rule, language, mapping):
|
||||
target = rule["target"]
|
||||
node_type = mapping.get(target, target)
|
||||
count = len(self.collect_nodes(tree.root_node, node_type))
|
||||
msg = self._message(rule, count)
|
||||
return [msg] if msg else []
|
||||
|
||||
def describe(self, rule, language, mapping):
|
||||
name = rule.get("label") or label(rule["target"])
|
||||
if rule.get("message"):
|
||||
return rule["message"]
|
||||
exact = rule.get("exact")
|
||||
if exact is not None:
|
||||
return f"{name} 出现 {exact} 次"
|
||||
parts = []
|
||||
if rule.get("min") is not None:
|
||||
parts.append(f"至少 {rule['min']} 次")
|
||||
if rule.get("max") is not None:
|
||||
parts.append(f"至多 {rule['max']} 次")
|
||||
return f"{name} " + "、".join(parts)
|
||||
@@ -1,31 +0,0 @@
|
||||
from ast_checker.labels import label
|
||||
|
||||
from .base import BaseEngine
|
||||
|
||||
|
||||
class MustExistNodeEngine(BaseEngine):
|
||||
def _message(self, rule):
|
||||
return rule.get("message") or f"必须使用 {rule.get('label') or label(rule['target'])}"
|
||||
|
||||
def check(self, tree, rule, language, mapping):
|
||||
node_type = mapping.get(rule["target"], rule["target"])
|
||||
if not self.has_node(tree.root_node, node_type):
|
||||
return [self._message(rule)]
|
||||
return []
|
||||
|
||||
def describe(self, rule, language, mapping):
|
||||
return self._message(rule)
|
||||
|
||||
|
||||
class MustNotExistNodeEngine(BaseEngine):
|
||||
def _message(self, rule):
|
||||
return rule.get("message") or f"不能使用 {rule.get('label') or label(rule['target'])}"
|
||||
|
||||
def check(self, tree, rule, language, mapping):
|
||||
node_type = mapping.get(rule["target"], rule["target"])
|
||||
if self.has_node(tree.root_node, node_type):
|
||||
return [self._message(rule)]
|
||||
return []
|
||||
|
||||
def describe(self, rule, language, mapping):
|
||||
return self._message(rule)
|
||||
@@ -1,15 +0,0 @@
|
||||
from .base import BaseEngine
|
||||
|
||||
|
||||
class MustUseOperatorEngine(BaseEngine):
|
||||
def _message(self, rule):
|
||||
return rule.get("message") or f"必须使用 {rule['target']} 运算符"
|
||||
|
||||
def check(self, tree, rule, language, mapping):
|
||||
mapped_op = mapping.get(rule["target"], rule["target"])
|
||||
if not self.has_node(tree.root_node, mapped_op):
|
||||
return [self._message(rule)]
|
||||
return []
|
||||
|
||||
def describe(self, rule, language, mapping):
|
||||
return self._message(rule)
|
||||
@@ -1,21 +0,0 @@
|
||||
TARGET_LABELS: dict[str, str] = {
|
||||
"for_loop": "for 循环",
|
||||
"while_loop": "while 循环",
|
||||
"if_statement": "if 条件",
|
||||
"else_clause": "else 子句",
|
||||
"function_definition": "函数定义",
|
||||
"return": "return 语句",
|
||||
"break": "break 语句",
|
||||
"continue": "continue 语句",
|
||||
"list_comprehension": "列表推导式",
|
||||
"list_literal": "列表",
|
||||
"dict_literal": "字典",
|
||||
"set_literal": "集合",
|
||||
"f_string": "f-string",
|
||||
"try_except": "try-except",
|
||||
"class_definition": "类定义",
|
||||
}
|
||||
|
||||
|
||||
def label(target: str) -> str:
|
||||
return TARGET_LABELS.get(target, target)
|
||||
@@ -1,37 +0,0 @@
|
||||
from tree_sitter import Language
|
||||
|
||||
from .c import C_MAPPING
|
||||
from .python import PYTHON_MAPPING
|
||||
|
||||
_MAPPINGS = {
|
||||
"Python3": PYTHON_MAPPING,
|
||||
"C": C_MAPPING,
|
||||
}
|
||||
|
||||
_LANGUAGES: dict[str, Language] = {}
|
||||
|
||||
|
||||
def _init_languages():
|
||||
try:
|
||||
import tree_sitter_python as tspython
|
||||
|
||||
_LANGUAGES["Python3"] = Language(tspython.language())
|
||||
except ImportError:
|
||||
pass
|
||||
try:
|
||||
import tree_sitter_c as tsc
|
||||
|
||||
_LANGUAGES["C"] = Language(tsc.language())
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
|
||||
_init_languages()
|
||||
|
||||
|
||||
def get_mapping(language: str) -> dict:
|
||||
return _MAPPINGS.get(language, {})
|
||||
|
||||
|
||||
def get_language(language: str) -> Language | None:
|
||||
return _LANGUAGES.get(language)
|
||||
@@ -1,39 +0,0 @@
|
||||
C_MAPPING = {
|
||||
"for_loop": "for_statement",
|
||||
"while_loop": "while_statement",
|
||||
"do_while": "do_statement",
|
||||
"if_statement": "if_statement",
|
||||
"else_clause": "else_clause",
|
||||
"break": "break_statement",
|
||||
"continue": "continue_statement",
|
||||
"function_definition": "function_definition",
|
||||
"return": "return_statement",
|
||||
"switch_statement": "switch_statement",
|
||||
"case_statement": "case_statement",
|
||||
"assignment": "assignment_expression",
|
||||
"struct": "struct_specifier",
|
||||
"include": "preproc_include",
|
||||
"+": "+",
|
||||
"-": "-",
|
||||
"*": "*",
|
||||
"/": "/",
|
||||
"%": "%",
|
||||
"+=": "+=",
|
||||
"-=": "-=",
|
||||
"*=": "*=",
|
||||
"/=": "/=",
|
||||
"%=": "%=",
|
||||
"==": "==",
|
||||
"!=": "!=",
|
||||
">": ">",
|
||||
">=": ">=",
|
||||
"<": "<",
|
||||
"<=": "<=",
|
||||
"and": "&&",
|
||||
"or": "||",
|
||||
"not": "!",
|
||||
"&": "&",
|
||||
"|": "|",
|
||||
"++": "++",
|
||||
"--": "--",
|
||||
}
|
||||
@@ -1,45 +0,0 @@
|
||||
PYTHON_MAPPING = {
|
||||
"for_loop": "for_statement",
|
||||
"while_loop": "while_statement",
|
||||
"if_statement": "if_statement",
|
||||
"else_clause": "else_clause",
|
||||
"elif_clause": "elif_clause",
|
||||
"break": "break_statement",
|
||||
"continue": "continue_statement",
|
||||
"function_definition": "function_definition",
|
||||
"return": "return_statement",
|
||||
"try_except": "try_statement",
|
||||
"with_statement": "with_statement",
|
||||
"list_comprehension": "list_comprehension",
|
||||
"list_literal": "list",
|
||||
"dict_literal": "dictionary",
|
||||
"set_literal": "set",
|
||||
"f_string": "format_string",
|
||||
"import": "import_statement",
|
||||
"import_from": "import_from_statement",
|
||||
"assignment": "assignment",
|
||||
"class_definition": "class_definition",
|
||||
"+": "+",
|
||||
"-": "-",
|
||||
"*": "*",
|
||||
"/": "/",
|
||||
"//": "//",
|
||||
"%": "%",
|
||||
"**": "**",
|
||||
"+=": "+=",
|
||||
"-=": "-=",
|
||||
"*=": "*=",
|
||||
"/=": "/=",
|
||||
"%=": "%=",
|
||||
"==": "==",
|
||||
"!=": "!=",
|
||||
">": ">",
|
||||
">=": ">=",
|
||||
"<": "<",
|
||||
"<=": "<=",
|
||||
"and": "and",
|
||||
"or": "or",
|
||||
"not": "not",
|
||||
"&": "&",
|
||||
"|": "|",
|
||||
}
|
||||
@@ -1 +0,0 @@
|
||||
# Register your models here.
|
||||
@@ -1,7 +0,0 @@
|
||||
from django.apps import AppConfig
|
||||
|
||||
|
||||
class ClassPkConfig(AppConfig):
|
||||
default_auto_field = 'django.db.models.BigAutoField'
|
||||
name = 'class_pk'
|
||||
verbose_name = '班级PK'
|
||||
@@ -1,2 +0,0 @@
|
||||
# 空文件
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
# 如果需要存储班级PK历史记录,可以在这里定义模型
|
||||
# 目前暂时不需要,因为都是实时计算
|
||||
@@ -1,3 +0,0 @@
|
||||
# 如果需要序列化器,可以在这里定义
|
||||
# 目前使用APIView的paginate_data方法,暂时不需要
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
# 空文件
|
||||
|
||||
@@ -1,9 +0,0 @@
|
||||
from django.urls import path
|
||||
|
||||
from ..views.oj import ClassPKAPI, ClassRankAPI, UserClassRankAPI
|
||||
|
||||
urlpatterns = [
|
||||
path("class_rank", ClassRankAPI.as_view()),
|
||||
path("user_class_rank", UserClassRankAPI.as_view()),
|
||||
path("class_pk", ClassPKAPI.as_view()),
|
||||
]
|
||||
@@ -1,337 +0,0 @@
|
||||
import math
|
||||
import statistics
|
||||
from datetime import datetime
|
||||
|
||||
from django.db.models import Avg, Sum
|
||||
from django.utils import timezone
|
||||
|
||||
from account.decorators import login_required
|
||||
from account.models import AdminType, User, UserProfile
|
||||
from submission.models import JudgeStatus, Submission
|
||||
from utils.api import APIView
|
||||
|
||||
|
||||
class ClassRankAPI(APIView):
|
||||
"""获取班级排名列表"""
|
||||
|
||||
def get(self, request):
|
||||
# 获取年级参数
|
||||
grade = int(request.GET.get("grade"))
|
||||
# 获取所有有用户的班级
|
||||
classes = (
|
||||
User.objects.filter(
|
||||
class_name__isnull=False,
|
||||
is_disabled=False,
|
||||
admin_type__in=[AdminType.REGULAR_USER, AdminType.STUDENT_ADMIN],
|
||||
class_name__startswith=str(grade),
|
||||
)
|
||||
.values("class_name")
|
||||
.distinct()
|
||||
)
|
||||
|
||||
class_stats = []
|
||||
for class_info in classes:
|
||||
class_name = class_info["class_name"]
|
||||
users = User.objects.filter(
|
||||
class_name=class_name,
|
||||
is_disabled=False,
|
||||
admin_type__in=[AdminType.REGULAR_USER, AdminType.STUDENT_ADMIN],
|
||||
)
|
||||
user_ids = list(users.values_list("id", flat=True))
|
||||
|
||||
profiles = UserProfile.objects.filter(user_id__in=user_ids)
|
||||
|
||||
total_ac = profiles.aggregate(total=Sum("accepted_number"))["total"] or 0
|
||||
total_submission = profiles.aggregate(total=Sum("submission_number"))["total"] or 0
|
||||
avg_ac = profiles.aggregate(avg=Avg("accepted_number"))["avg"] or 0
|
||||
|
||||
user_count = users.count()
|
||||
|
||||
class_stats.append(
|
||||
{
|
||||
"class_name": class_name,
|
||||
"user_count": user_count,
|
||||
"total_ac": int(total_ac),
|
||||
"total_submission": int(total_submission),
|
||||
"avg_ac": round(avg_ac, 2),
|
||||
"ac_rate": round(total_ac / total_submission * 100, 2) if total_submission > 0 else 0,
|
||||
}
|
||||
)
|
||||
|
||||
# 按总AC数排序
|
||||
class_stats.sort(key=lambda x: (-x["total_ac"], x["total_submission"]))
|
||||
|
||||
# 添加排名
|
||||
for i, stat in enumerate(class_stats):
|
||||
stat["rank"] = i + 1
|
||||
|
||||
return self.success(class_stats)
|
||||
|
||||
|
||||
class UserClassRankAPI(APIView):
|
||||
"""获取用户在班级中的排名"""
|
||||
|
||||
@login_required
|
||||
def get(self, request):
|
||||
user = request.user
|
||||
if not user.class_name:
|
||||
return self.error("用户没有班级信息")
|
||||
scope = request.GET.get("scope", "").lower()
|
||||
show_all = scope == "all"
|
||||
try:
|
||||
limit = int(request.GET.get("limit", "10"))
|
||||
except ValueError:
|
||||
limit = 10
|
||||
if limit <= 0 or limit > 250:
|
||||
limit = 10
|
||||
try:
|
||||
offset = int(request.GET.get("offset", "0"))
|
||||
except ValueError:
|
||||
offset = 0
|
||||
if offset < 0:
|
||||
offset = 0
|
||||
|
||||
# 获取同班所有用户
|
||||
class_users = User.objects.filter(
|
||||
class_name=user.class_name,
|
||||
is_disabled=False,
|
||||
admin_type__in=[AdminType.REGULAR_USER, AdminType.STUDENT_ADMIN],
|
||||
).select_related("userprofile")
|
||||
|
||||
user_ranks = []
|
||||
for class_user in class_users:
|
||||
profile = class_user.userprofile
|
||||
user_ranks.append(
|
||||
{
|
||||
"user_id": class_user.id,
|
||||
"username": class_user.username,
|
||||
"accepted_number": profile.accepted_number,
|
||||
"submission_number": profile.submission_number,
|
||||
}
|
||||
)
|
||||
|
||||
# 按AC数排序
|
||||
user_ranks.sort(key=lambda x: (-x["accepted_number"], x["submission_number"]))
|
||||
|
||||
# 添加排名
|
||||
my_rank = -1
|
||||
for i, rank_info in enumerate(user_ranks):
|
||||
rank_info["rank"] = i + 1
|
||||
if rank_info["user_id"] == user.id:
|
||||
my_rank = i + 1
|
||||
|
||||
trimmed_ranks = user_ranks
|
||||
if not show_all and my_rank > 0 and len(user_ranks) > 10:
|
||||
center_index = my_rank - 1
|
||||
start = max(0, center_index - 5)
|
||||
end = start + 10
|
||||
if end > len(user_ranks):
|
||||
end = len(user_ranks)
|
||||
start = max(0, end - 10)
|
||||
trimmed_ranks = user_ranks[start:end]
|
||||
elif show_all:
|
||||
trimmed_ranks = user_ranks[offset : offset + limit]
|
||||
|
||||
return self.success(
|
||||
{
|
||||
"class_name": user.class_name,
|
||||
"my_rank": my_rank,
|
||||
"total": len(user_ranks),
|
||||
"ranks": trimmed_ranks,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class ClassPKAPI(APIView):
|
||||
"""班级PK比较 - 多维度教育评价"""
|
||||
|
||||
def post(self, request):
|
||||
class_names = request.data.get("class_name", [])
|
||||
if not class_names or len(class_names) < 1:
|
||||
return self.error("至少需要选择1个班级")
|
||||
|
||||
# 获取时间段参数
|
||||
start_time = request.data.get("start_time")
|
||||
end_time = request.data.get("end_time")
|
||||
|
||||
# 将时间字符串转换为datetime对象
|
||||
# 处理空字符串、None 或 undefined 的情况
|
||||
if start_time and isinstance(start_time, str) and start_time.strip():
|
||||
try:
|
||||
start_time = datetime.fromisoformat(start_time.replace("Z", "+00:00"))
|
||||
if timezone.is_naive(start_time):
|
||||
start_time = timezone.make_aware(start_time)
|
||||
except (ValueError, AttributeError):
|
||||
start_time = None
|
||||
else:
|
||||
start_time = None
|
||||
|
||||
if end_time and isinstance(end_time, str) and end_time.strip():
|
||||
try:
|
||||
end_time = datetime.fromisoformat(end_time.replace("Z", "+00:00"))
|
||||
if timezone.is_naive(end_time):
|
||||
end_time = timezone.make_aware(end_time)
|
||||
except (ValueError, AttributeError):
|
||||
end_time = None
|
||||
else:
|
||||
end_time = None
|
||||
|
||||
class_comparisons = []
|
||||
|
||||
# 预计算全局阈值(所有参与PK班级的学生AC数合并)
|
||||
all_user_ids = list(
|
||||
User.objects.filter(
|
||||
class_name__in=class_names,
|
||||
is_disabled=False,
|
||||
admin_type__in=[AdminType.REGULAR_USER, AdminType.STUDENT_ADMIN],
|
||||
).values_list("id", flat=True)
|
||||
)
|
||||
all_ac_list = sorted(
|
||||
[p.accepted_number for p in UserProfile.objects.filter(user_id__in=all_user_ids)],
|
||||
reverse=True,
|
||||
)
|
||||
if len(all_ac_list) > 1:
|
||||
_quantiles = statistics.quantiles(all_ac_list, n=4)
|
||||
global_q1 = _quantiles[0]
|
||||
global_q3 = _quantiles[2]
|
||||
else:
|
||||
global_q1 = all_ac_list[0] if all_ac_list else 0
|
||||
global_q3 = all_ac_list[0] if all_ac_list else 0
|
||||
|
||||
for class_name in class_names:
|
||||
users = User.objects.filter(
|
||||
class_name=class_name,
|
||||
is_disabled=False,
|
||||
admin_type__in=[AdminType.REGULAR_USER, AdminType.STUDENT_ADMIN],
|
||||
)
|
||||
user_ids = list(users.values_list("id", flat=True))
|
||||
|
||||
# 获取所有学生的AC数列表(用于统计计算)
|
||||
profiles = UserProfile.objects.filter(user_id__in=user_ids)
|
||||
ac_list = sorted([p.accepted_number for p in profiles], reverse=True)
|
||||
submission_list = sorted([p.submission_number for p in profiles], reverse=True)
|
||||
|
||||
user_count = len(ac_list)
|
||||
if user_count == 0:
|
||||
continue
|
||||
|
||||
# 基础统计
|
||||
total_ac = sum(ac_list)
|
||||
total_submission = sum(submission_list)
|
||||
avg_ac = statistics.mean(ac_list) if ac_list else 0
|
||||
|
||||
# 中位数和分位数
|
||||
median_ac = statistics.median(ac_list) if ac_list else 0
|
||||
q1_ac = statistics.quantiles(ac_list, n=4)[0] if len(ac_list) > 1 else 0
|
||||
q3_ac = statistics.quantiles(ac_list, n=4)[2] if len(ac_list) > 1 else 0
|
||||
iqr = q3_ac - q1_ac
|
||||
|
||||
# 标准差
|
||||
std_dev = statistics.stdev(ac_list) if len(ac_list) > 1 else 0
|
||||
|
||||
# 前10%和后10%统计
|
||||
top_10_count = max(1, math.ceil(user_count * 0.10))
|
||||
bottom_10_count = max(1, math.ceil(user_count * 0.10))
|
||||
top_10_avg = statistics.mean(ac_list[:top_10_count]) if top_10_count > 0 else 0
|
||||
bottom_10_avg = statistics.mean(ac_list[-bottom_10_count:]) if bottom_10_count > 0 else 0
|
||||
|
||||
# 中间80%均值(截尾均值,去掉前10%和后10%)
|
||||
if top_10_count + bottom_10_count < user_count:
|
||||
middle_list = ac_list[top_10_count:-bottom_10_count]
|
||||
else:
|
||||
middle_list = ac_list
|
||||
middle_80_avg = statistics.mean(middle_list) if middle_list else avg_ac
|
||||
|
||||
# 优秀率(AC数 >= 全局Q3,即超过PK组所有学生的前25%)
|
||||
excellent_count = sum(1 for ac in ac_list if ac >= global_q3)
|
||||
excellent_rate = (excellent_count / user_count * 100) if user_count > 0 else 0
|
||||
|
||||
# 及格率(AC数 >= 全局Q1,即超过PK组所有学生的后25%)
|
||||
pass_count = sum(1 for ac in ac_list if ac >= global_q1)
|
||||
pass_rate = (pass_count / user_count * 100) if user_count > 0 else 0
|
||||
|
||||
# 参与度(有提交记录的学生比例)
|
||||
active_count = sum(1 for sub in submission_list if sub > 0)
|
||||
active_rate = (active_count / user_count * 100) if user_count > 0 else 0
|
||||
|
||||
# 时间段内的统计(如果提供了时间段)
|
||||
recent_stats = {}
|
||||
if start_time and end_time:
|
||||
submissions = Submission.objects.filter(
|
||||
user_id__in=user_ids,
|
||||
create_time__gte=start_time,
|
||||
create_time__lte=end_time,
|
||||
)
|
||||
recent_ac = submissions.filter(result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED]).values("user_id", "problem_id").distinct().count()
|
||||
recent_submission = submissions.count()
|
||||
|
||||
# 时间段内的用户AC数列表
|
||||
recent_user_ac = {}
|
||||
for user_id in user_ids:
|
||||
user_recent_ac = submissions.filter(user_id=user_id, result__in=[JudgeStatus.ACCEPTED, JudgeStatus.AST_CHECK_FAILED]).values("problem_id").distinct().count()
|
||||
recent_user_ac[user_id] = user_recent_ac
|
||||
|
||||
recent_ac_list = sorted(recent_user_ac.values(), reverse=True)
|
||||
if recent_ac_list:
|
||||
recent_stats = {
|
||||
"recent_total_ac": recent_ac,
|
||||
"recent_total_submission": recent_submission,
|
||||
"recent_avg_ac": statistics.mean(recent_ac_list),
|
||||
"recent_median_ac": statistics.median(recent_ac_list),
|
||||
"recent_top_10_avg": statistics.mean(recent_ac_list[: max(1, math.ceil(len(recent_ac_list) * 0.10))]) if recent_ac_list else 0,
|
||||
"recent_active_count": sum(1 for ac in recent_ac_list if ac > 0),
|
||||
}
|
||||
|
||||
class_comparisons.append(
|
||||
{
|
||||
"class_name": class_name,
|
||||
"user_count": user_count,
|
||||
# 基础统计
|
||||
"total_ac": int(total_ac),
|
||||
"total_submission": int(total_submission),
|
||||
"avg_ac": round(avg_ac, 2),
|
||||
# 中位数和分位数
|
||||
"median_ac": round(median_ac, 2),
|
||||
"q1_ac": round(q1_ac, 2),
|
||||
"q3_ac": round(q3_ac, 2),
|
||||
"iqr": round(iqr, 2),
|
||||
# 标准差
|
||||
"std_dev": round(std_dev, 2),
|
||||
# 分层统计
|
||||
"top_10_avg": round(top_10_avg, 2),
|
||||
"middle_80_avg": round(middle_80_avg, 2),
|
||||
"bottom_10_avg": round(bottom_10_avg, 2),
|
||||
# 比率统计
|
||||
"excellent_rate": round(excellent_rate, 2),
|
||||
"pass_rate": round(pass_rate, 2),
|
||||
"active_rate": round(active_rate, 2),
|
||||
# 正确率
|
||||
"ac_rate": round(total_ac / total_submission * 100, 2) if total_submission > 0 else 0,
|
||||
# 时间段统计(如果有)
|
||||
**recent_stats,
|
||||
}
|
||||
)
|
||||
|
||||
# 计算综合分(需要所有班级数据就绪后才能归一化)
|
||||
max_median = max((c["median_ac"] for c in class_comparisons), default=1) or 1
|
||||
max_middle = max((c["middle_80_avg"] for c in class_comparisons), default=1) or 1
|
||||
|
||||
for c in class_comparisons:
|
||||
score = (
|
||||
0.40 * (c["median_ac"] / max_median * 100)
|
||||
+ 0.15 * (c["middle_80_avg"] / max_middle * 100)
|
||||
+ 0.20 * c["active_rate"]
|
||||
+ 0.15 * c["pass_rate"]
|
||||
+ 0.10 * c["excellent_rate"]
|
||||
)
|
||||
c["composite_score"] = round(score, 1)
|
||||
|
||||
# 按综合分排序(主),中位数(次)
|
||||
class_comparisons.sort(key=lambda x: (-x["composite_score"], -x["median_ac"]))
|
||||
|
||||
return self.success(
|
||||
{
|
||||
"comparisons": class_comparisons,
|
||||
"has_time_range": bool(start_time and end_time),
|
||||
}
|
||||
)
|
||||
38
comment/migrations/0001_initial.py
Normal file
38
comment/migrations/0001_initial.py
Normal file
@@ -0,0 +1,38 @@
|
||||
# Generated by Django 5.2.3 on 2025-06-14 08:51
|
||||
|
||||
import django.db.models.deletion
|
||||
from django.conf import settings
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
|
||||
initial = True
|
||||
|
||||
dependencies = [
|
||||
('problem', '0001_initial'),
|
||||
('submission', '0001_initial'),
|
||||
migrations.swappable_dependency(settings.AUTH_USER_MODEL),
|
||||
]
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name='Comment',
|
||||
fields=[
|
||||
('id', models.AutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||
('language', models.CharField(choices=[('Python', 'Python'), ('C', 'C'), ('C++', 'C++'), ('Java', 'Java')], default='Python', max_length=10, verbose_name='解决这道题使用的语言')),
|
||||
('description_rating', models.PositiveSmallIntegerField(default=5, verbose_name='题目描述的分数')),
|
||||
('difficulty_rating', models.PositiveSmallIntegerField(default=5, verbose_name='题目难度的分数')),
|
||||
('comprehensive_rating', models.PositiveSmallIntegerField(default=5, verbose_name='综合的分数')),
|
||||
('content', models.TextField(blank=True, null=True)),
|
||||
('create_time', models.DateTimeField(auto_now_add=True)),
|
||||
('problem', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to='problem.problem')),
|
||||
('submission', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to='submission.submission')),
|
||||
('user', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, to=settings.AUTH_USER_MODEL)),
|
||||
],
|
||||
options={
|
||||
'db_table': 'comment',
|
||||
'ordering': ('-create_time',),
|
||||
},
|
||||
),
|
||||
]
|
||||
44
comment/models.py
Normal file
44
comment/models.py
Normal file
@@ -0,0 +1,44 @@
|
||||
from django.db import models
|
||||
|
||||
from account.models import User
|
||||
from problem.models import Problem
|
||||
from submission.models import Submission
|
||||
|
||||
|
||||
class Languages(models.TextChoices):
|
||||
Python = "Python", "Python"
|
||||
C = "C", "C"
|
||||
Cpp = "C++", "C++"
|
||||
Java = "Java", "Java"
|
||||
|
||||
|
||||
class Comment(models.Model):
|
||||
problem = models.ForeignKey(Problem, on_delete=models.CASCADE)
|
||||
user = models.ForeignKey(User, on_delete=models.CASCADE)
|
||||
submission = models.ForeignKey(Submission, on_delete=models.CASCADE)
|
||||
language = models.CharField(
|
||||
max_length=10,
|
||||
default=Languages.Python,
|
||||
choices=Languages.choices,
|
||||
verbose_name="解决这道题使用的语言",
|
||||
)
|
||||
description_rating = models.PositiveSmallIntegerField(
|
||||
default=5,
|
||||
verbose_name="题目描述的分数",
|
||||
)
|
||||
difficulty_rating = models.PositiveSmallIntegerField(
|
||||
default=5,
|
||||
verbose_name="题目难度的分数",
|
||||
)
|
||||
comprehensive_rating = models.PositiveSmallIntegerField(
|
||||
default=5,
|
||||
verbose_name="综合的分数",
|
||||
)
|
||||
content = models.TextField(null=True, blank=True)
|
||||
create_time = models.DateTimeField(auto_now_add=True)
|
||||
|
||||
class Meta:
|
||||
db_table = "comment"
|
||||
ordering = ("-create_time",)
|
||||
|
||||
|
||||
31
comment/serializers.py
Normal file
31
comment/serializers.py
Normal file
@@ -0,0 +1,31 @@
|
||||
from comment.models import Comment
|
||||
from utils.api import UsernameSerializer, serializers
|
||||
|
||||
|
||||
class CreateCommentSerializer(serializers.Serializer):
|
||||
problem_id = serializers.IntegerField()
|
||||
description_rating = serializers.IntegerField()
|
||||
difficulty_rating = serializers.IntegerField()
|
||||
comprehensive_rating = serializers.IntegerField()
|
||||
content = serializers.CharField(required=False, allow_blank=True)
|
||||
|
||||
|
||||
class CommentSerializer(serializers.ModelSerializer):
|
||||
class Meta:
|
||||
model = Comment
|
||||
fields = [
|
||||
"comprehensive_rating",
|
||||
"description_rating",
|
||||
"difficulty_rating",
|
||||
"content",
|
||||
"create_time",
|
||||
]
|
||||
|
||||
|
||||
class CommentListSerializer(serializers.ModelSerializer):
|
||||
problem = serializers.SlugRelatedField(read_only=True, slug_field="_id")
|
||||
user = UsernameSerializer()
|
||||
|
||||
class Meta:
|
||||
model = Comment
|
||||
fields = "__all__"
|
||||
8
comment/urls/admin.py
Normal file
8
comment/urls/admin.py
Normal file
@@ -0,0 +1,8 @@
|
||||
from django.urls import path
|
||||
|
||||
from ..views.admin import CommentAPI
|
||||
|
||||
|
||||
urlpatterns = [
|
||||
path("comment", CommentAPI.as_view()),
|
||||
]
|
||||
9
comment/urls/oj.py
Normal file
9
comment/urls/oj.py
Normal file
@@ -0,0 +1,9 @@
|
||||
from django.urls import path
|
||||
|
||||
from ..views.oj import CommentAPI, CommentStatisticsAPI
|
||||
|
||||
|
||||
urlpatterns = [
|
||||
path("comment", CommentAPI.as_view()),
|
||||
path("comment/statistics", CommentStatisticsAPI.as_view()),
|
||||
]
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user