refactor(problem): 标签创建逻辑收敛到 services.resolve_tags

原本 4 处题目保存逻辑各自复制了一份 get-or-create,且大小写敏感、有竞态。

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
2026-08-05 04:15:36 -06:00
parent e698ef8bfa
commit ec7cf71ea7
2 changed files with 58 additions and 28 deletions

52
problem/services.py Normal file
View File

@@ -0,0 +1,52 @@
from django.core.cache import cache
from django.db import IntegrityError
from utils.constants import CacheKey
from .models import ProblemTag
def clear_tag_cache():
"""标签列表接口按 keyword 分键缓存,标签一变就整批清掉"""
cache.delete_pattern(f"{CacheKey.problem_tags}:*")
def find_tags(names):
"""按名字大小写不敏感查已有标签,查不到就跳过,不创建"""
tags = []
seen = set()
for raw in names:
name = (raw or "").strip()
if not name or name.lower() in seen:
continue
seen.add(name.lower())
tag = ProblemTag.objects.filter(name__iexact=name).first()
if tag is not None:
tags.append(tag)
return tags
def resolve_tags(names):
"""把前端传来的标签名解析成 ProblemTag去空格、大小写不敏感复用已有标签没有才新建"""
tags = []
seen = set()
created = False
for raw in names:
name = (raw or "").strip()
if not name or name.lower() in seen:
continue
seen.add(name.lower())
tag = ProblemTag.objects.filter(name__iexact=name).first()
if tag is None:
try:
tag = ProblemTag.objects.create(name=name)
created = True
except IntegrityError:
# 并发下另一个请求刚建好同名标签,回查复用
tag = ProblemTag.objects.filter(name__iexact=name).first()
if tag is None:
continue
tags.append(tag)
if created:
clear_tag_cache()
return tags

View File

@@ -19,7 +19,7 @@ from utils.api import APIError, APIView, CSRFExemptAPIView, validate_serializer
from utils.openai import get_ai_client from utils.openai import get_ai_client
from utils.shortcuts import natural_sort_key, rand_str from utils.shortcuts import natural_sort_key, rand_str
from ..models import Problem, ProblemTag from ..models import Problem
from ..serializers import ( from ..serializers import (
AddContestProblemSerializer, AddContestProblemSerializer,
ContestProblemMakePublicSerializer, ContestProblemMakePublicSerializer,
@@ -32,6 +32,7 @@ from ..serializers import (
SQLTestCasePreviewSerializer, SQLTestCasePreviewSerializer,
TestCaseUploadForm, TestCaseUploadForm,
) )
from ..services import resolve_tags
from ..utils import generate_sql_display from ..utils import generate_sql_display
@@ -248,12 +249,7 @@ class ProblemAPI(ProblemBase):
data["created_by"] = request.user data["created_by"] = request.user
problem = Problem.objects.create(**data) problem = Problem.objects.create(**data)
for item in tags: problem.tags.set(resolve_tags(tags))
try:
tag = ProblemTag.objects.get(name=item)
except ProblemTag.DoesNotExist:
tag = ProblemTag.objects.create(name=item)
problem.tags.add(tag)
return self.success(ProblemAdminSerializer(problem).data) return self.success(ProblemAdminSerializer(problem).data)
@problem_permission_required @problem_permission_required
@@ -310,14 +306,7 @@ class ProblemAPI(ProblemBase):
setattr(problem, k, v) setattr(problem, k, v)
problem.save() problem.save()
problem.tags.remove(*problem.tags.all()) problem.tags.set(resolve_tags(tags))
for tag in tags:
try:
tag = ProblemTag.objects.get(name=tag)
except ProblemTag.DoesNotExist:
tag = ProblemTag.objects.create(name=tag)
problem.tags.add(tag)
return self.success() return self.success()
@problem_permission_required @problem_permission_required
@@ -364,12 +353,7 @@ class ContestProblemAPI(ProblemBase):
data["created_by"] = request.user data["created_by"] = request.user
problem = Problem.objects.create(**data) problem = Problem.objects.create(**data)
for item in tags: problem.tags.set(resolve_tags(tags))
try:
tag = ProblemTag.objects.get(name=item)
except ProblemTag.DoesNotExist:
tag = ProblemTag.objects.create(name=item)
problem.tags.add(tag)
return self.success(ProblemAdminSerializer(problem).data) return self.success(ProblemAdminSerializer(problem).data)
def get(self, request): def get(self, request):
@@ -434,13 +418,7 @@ class ContestProblemAPI(ProblemBase):
setattr(problem, k, v) setattr(problem, k, v)
problem.save() problem.save()
problem.tags.remove(*problem.tags.all()) problem.tags.set(resolve_tags(tags))
for tag in tags:
try:
tag = ProblemTag.objects.get(name=tag)
except ProblemTag.DoesNotExist:
tag = ProblemTag.objects.create(name=tag)
problem.tags.add(tag)
return self.success() return self.success()
def delete(self, request): def delete(self, request):