feat: implement t-604 keyword prompt moderation

This commit is contained in:
QiuSW
2026-07-06 14:57:25 +08:00
parent c89cb4cf28
commit ca0ffb8dbe
24 changed files with 787 additions and 10 deletions
+10
View File
@@ -29,6 +29,8 @@ from apps.billing.services import (
mark_call_success,
)
from apps.moderation.services import moderate_prompt
from .errors import api_error
from .storage import save_generated_image
@@ -60,6 +62,7 @@ def generate_title_response(*, user, api_key, request_data: Mapping[str, Any]) -
alias = request_data.get("model") or None
resolution = normalize_resolution(request_data.get("resolution") or "1K") or "1K"
parameters = dict(request_data.get("parameters") or {})
moderate_prompt_or_raise(user=user, api_key=api_key, prompt=prompt)
image_input = load_image_input(request_data)
model_alias = resolve_model_alias_or_raise(CallRecord.OperationType.TITLE, alias)
@@ -134,6 +137,7 @@ def generate_image_response(*, user, api_key, request, request_data: Mapping[str
resolution = normalize_resolution(request_data.get("resolution") or "1K") or "1K"
aspect_ratio = request_data.get("aspect_ratio") or "1:1"
parameters = dict(request_data.get("parameters") or {})
moderate_prompt_or_raise(user=user, api_key=api_key, prompt=prompt)
image_input = load_image_input(request_data)
model_alias = resolve_model_alias_or_raise(CallRecord.OperationType.IMAGE, alias)
@@ -202,6 +206,12 @@ def generate_image_response(*, user, api_key, request, request_data: Mapping[str
}
def moderate_prompt_or_raise(*, user, api_key, prompt: str) -> None:
outcome = moderate_prompt(user=user, api_key=api_key, prompt=prompt)
if outcome.blocked:
raise ApiRequestError("content_blocked", "输入内容未通过安全审核", status.HTTP_400_BAD_REQUEST)
def resolve_model_alias_or_raise(operation_type: str, alias: str | None):
try:
return resolve_model_alias(operation_type, alias)
+51
View File
@@ -39,6 +39,8 @@ from apps.billing.payment_gateways import (
build_mock_body_signature,
)
from apps.billing.services import RechargePayment
from apps.moderation.models import SensitiveWord
from apps.moderation.providers.keyword import reset_keyword_matcher_cache
from apps.users.models import ApiKey
from apps.users.models import UserWallet
@@ -824,9 +826,11 @@ def dns_result(address: str):
class GenerateApiTests(TestCase):
def setUp(self):
cache.clear()
reset_keyword_matcher_cache()
suffix = uuid.uuid4().hex[:8]
self.media_dir = tempfile.TemporaryDirectory()
self.addCleanup(self.media_dir.cleanup)
self.addCleanup(reset_keyword_matcher_cache)
media_override = override_settings(
MEDIA_ROOT=self.media_dir.name,
MEDIA_URL="/media/",
@@ -905,6 +909,53 @@ class GenerateApiTests(TestCase):
self.assertFalse(CallRecord.objects.filter(user=self.user).exists())
self.assertFalse(PointsLedger.objects.filter(user=self.user).exists())
@override_settings(
MODERATION_ENABLED=True,
MODERATION_PROVIDER="keyword",
MODERATION_CACHE_VERSION_KEY="test:api:moderation:sensitive_words:version",
)
def test_blocked_prompt_returns_content_blocked_before_image_download_or_charge(self):
SensitiveWord.objects.create(word="敏感词", category="policy")
with (
patch("apps.api.generation.socket.getaddrinfo") as dns_lookup,
patch("apps.api.generation.requests.Session.get") as image_get,
):
response = self.post_with_provider(
"/api/v1/generate/image",
{
"prompt": "请生成敏-感\u200b 词图片",
"model": self.image_alias,
"image_url": "https://safe.example.com/input.jpg",
"resolution": "1K",
"aspect_ratio": "1:1",
},
)
self.assertEqual(response.status_code, 400)
self.assertEqual(response.data["error"]["code"], "content_blocked")
dns_lookup.assert_not_called()
image_get.assert_not_called()
self.assertEqual(self.provider.image_calls, [])
self.assert_generation_not_charged()
@override_settings(
MODERATION_ENABLED=False,
MODERATION_PROVIDER="keyword",
MODERATION_CACHE_VERSION_KEY="test:api:moderation:sensitive_words:version",
)
def test_disabled_moderation_does_not_block_matching_prompt(self):
SensitiveWord.objects.create(word="敏感词")
response = self.post_with_provider(
"/api/v1/generate/title",
{"prompt": "敏感词", "model": self.title_alias},
)
self.assertEqual(response.status_code, 200)
self.assertEqual(response.data["points_cost"], 2)
self.assertEqual(len(self.provider.text_calls), 1)
def test_generate_title_uses_default_alias_charges_points_and_writes_call_record(self):
response = self.post_with_provider(
"/api/v1/generate/title",