feat: implement t-604 keyword prompt moderation
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user