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",
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from django.contrib import admin
|
||||
|
||||
from .models import SensitiveWord
|
||||
|
||||
|
||||
@admin.register(SensitiveWord)
|
||||
class SensitiveWordAdmin(admin.ModelAdmin):
|
||||
list_display = (
|
||||
"word",
|
||||
"normalized_word",
|
||||
"category",
|
||||
"action",
|
||||
"is_active",
|
||||
"updated_at",
|
||||
)
|
||||
list_filter = ("category", "action", "is_active")
|
||||
search_fields = ("word", "normalized_word", "category")
|
||||
readonly_fields = ("normalized_word", "created_at", "updated_at")
|
||||
@@ -0,0 +1,10 @@
|
||||
from django.apps import AppConfig
|
||||
|
||||
|
||||
class ModerationConfig(AppConfig):
|
||||
default_auto_field = "django.db.models.BigAutoField"
|
||||
name = "apps.moderation"
|
||||
verbose_name = "内容安全"
|
||||
|
||||
def ready(self) -> None:
|
||||
from . import signals # noqa: F401
|
||||
@@ -0,0 +1,46 @@
|
||||
# Generated by Codex for T-604.
|
||||
|
||||
from django.db import migrations, models
|
||||
|
||||
|
||||
class Migration(migrations.Migration):
|
||||
initial = True
|
||||
|
||||
dependencies = []
|
||||
|
||||
operations = [
|
||||
migrations.CreateModel(
|
||||
name="SensitiveWord",
|
||||
fields=[
|
||||
("id", models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name="ID")),
|
||||
("word", models.CharField(max_length=255, verbose_name="敏感词")),
|
||||
("normalized_word", models.CharField(editable=False, max_length=255, verbose_name="归一化敏感词")),
|
||||
("category", models.CharField(blank=True, default="custom", max_length=64, verbose_name="分类")),
|
||||
("action", models.CharField(choices=[("block", "拦截")], default="block", max_length=16, verbose_name="动作")),
|
||||
("is_active", models.BooleanField(default=True, verbose_name="启用")),
|
||||
("created_at", models.DateTimeField(auto_now_add=True, verbose_name="创建时间")),
|
||||
("updated_at", models.DateTimeField(auto_now=True, verbose_name="更新时间")),
|
||||
],
|
||||
options={
|
||||
"verbose_name": "敏感词",
|
||||
"verbose_name_plural": "敏感词",
|
||||
"db_table": "sensitive_word",
|
||||
"ordering": ("category", "word"),
|
||||
},
|
||||
),
|
||||
migrations.AddIndex(
|
||||
model_name="sensitiveword",
|
||||
index=models.Index(fields=["is_active", "category"], name="sw_active_category_idx"),
|
||||
),
|
||||
migrations.AddIndex(
|
||||
model_name="sensitiveword",
|
||||
index=models.Index(fields=["normalized_word"], name="sw_normalized_word_idx"),
|
||||
),
|
||||
migrations.AddConstraint(
|
||||
model_name="sensitiveword",
|
||||
constraint=models.UniqueConstraint(
|
||||
fields=("category", "normalized_word"),
|
||||
name="unique_sensitive_word_per_category",
|
||||
),
|
||||
),
|
||||
]
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from django.core.exceptions import ValidationError
|
||||
from django.db import models
|
||||
|
||||
from .normalization import normalize_text
|
||||
|
||||
|
||||
class SensitiveWord(models.Model):
|
||||
class Action(models.TextChoices):
|
||||
BLOCK = "block", "拦截"
|
||||
|
||||
word = models.CharField("敏感词", max_length=255)
|
||||
normalized_word = models.CharField("归一化敏感词", max_length=255, editable=False)
|
||||
category = models.CharField("分类", max_length=64, default="custom", blank=True)
|
||||
action = models.CharField("动作", max_length=16, choices=Action.choices, default=Action.BLOCK)
|
||||
is_active = models.BooleanField("启用", default=True)
|
||||
created_at = models.DateTimeField("创建时间", auto_now_add=True)
|
||||
updated_at = models.DateTimeField("更新时间", auto_now=True)
|
||||
|
||||
class Meta:
|
||||
db_table = "sensitive_word"
|
||||
verbose_name = "敏感词"
|
||||
verbose_name_plural = "敏感词"
|
||||
ordering = ("category", "word")
|
||||
constraints = [
|
||||
models.UniqueConstraint(
|
||||
fields=("category", "normalized_word"),
|
||||
name="unique_sensitive_word_per_category",
|
||||
),
|
||||
]
|
||||
indexes = [
|
||||
models.Index(fields=("is_active", "category"), name="sw_active_category_idx"),
|
||||
models.Index(fields=("normalized_word",), name="sw_normalized_word_idx"),
|
||||
]
|
||||
|
||||
def __str__(self) -> str:
|
||||
return self.word
|
||||
|
||||
def clean(self) -> None:
|
||||
self.word = (self.word or "").strip()
|
||||
self.category = (self.category or "custom").strip() or "custom"
|
||||
self.normalized_word = normalize_text(self.word)
|
||||
if not self.normalized_word:
|
||||
raise ValidationError({"word": "敏感词归一化后不能为空"})
|
||||
if self.action != self.Action.BLOCK:
|
||||
raise ValidationError({"action": "MVP 只支持 block 动作"})
|
||||
|
||||
def save(self, *args, **kwargs) -> None:
|
||||
self.full_clean()
|
||||
super().save(*args, **kwargs)
|
||||
@@ -0,0 +1,54 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import unicodedata
|
||||
from functools import lru_cache
|
||||
|
||||
_ZERO_WIDTH_CHARS = {
|
||||
"\u200b",
|
||||
"\u200c",
|
||||
"\u200d",
|
||||
"\ufeff",
|
||||
"\u2060",
|
||||
}
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _opencc_converter():
|
||||
try:
|
||||
from opencc import OpenCC
|
||||
except Exception:
|
||||
return None
|
||||
try:
|
||||
return OpenCC("t2s")
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _to_simplified(value: str) -> str:
|
||||
converter = _opencc_converter()
|
||||
if converter is None:
|
||||
return value
|
||||
return converter.convert(value)
|
||||
|
||||
|
||||
def normalize_text(value: str | None) -> str:
|
||||
"""Normalize text before keyword matching.
|
||||
|
||||
The local keyword MVP removes obvious bypass characters only; it does not
|
||||
attempt semantic rewriting or fuzzy matching.
|
||||
"""
|
||||
if not value:
|
||||
return ""
|
||||
|
||||
normalized = unicodedata.normalize("NFKC", value).lower()
|
||||
normalized = _to_simplified(normalized)
|
||||
|
||||
chars: list[str] = []
|
||||
for char in normalized:
|
||||
if char in _ZERO_WIDTH_CHARS:
|
||||
continue
|
||||
category = unicodedata.category(char)
|
||||
if category[0] in {"P", "Z"}:
|
||||
continue
|
||||
chars.append(char)
|
||||
return "".join(chars)
|
||||
@@ -0,0 +1,59 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from django.conf import settings
|
||||
|
||||
from .providers import ModerationConfigError, ModerationProvider, Verdict
|
||||
|
||||
|
||||
def moderation_enabled() -> bool:
|
||||
return bool(getattr(settings, "MODERATION_ENABLED", False))
|
||||
|
||||
|
||||
def fail_closed() -> bool:
|
||||
return bool(getattr(settings, "MODERATION_FAIL_CLOSED", True))
|
||||
|
||||
|
||||
def block_on_review() -> bool:
|
||||
return bool(getattr(settings, "MODERATION_BLOCK_ON_REVIEW", False))
|
||||
|
||||
|
||||
def verdict_is_blocking(verdict: Verdict) -> bool:
|
||||
if verdict == Verdict.BLOCK:
|
||||
return True
|
||||
if verdict == Verdict.REVIEW:
|
||||
return block_on_review()
|
||||
return False
|
||||
|
||||
|
||||
def selected_provider_name() -> str:
|
||||
return str(getattr(settings, "MODERATION_PROVIDER", "") or "").strip().lower()
|
||||
|
||||
|
||||
def _build_keyword() -> ModerationProvider:
|
||||
from .providers.keyword import KeywordModerationProvider
|
||||
|
||||
return KeywordModerationProvider()
|
||||
|
||||
|
||||
_PROVIDER_BUILDERS = {
|
||||
"keyword": _build_keyword,
|
||||
}
|
||||
|
||||
|
||||
def biz_type(stage: str) -> str:
|
||||
return stage
|
||||
|
||||
|
||||
def _build_provider() -> ModerationProvider:
|
||||
"""按 MODERATION_PROVIDER 选具体厂商。未选 / 未知 → ModerationConfigError(上层按 fail-closed 处理)。"""
|
||||
name = selected_provider_name()
|
||||
if not name:
|
||||
raise ModerationConfigError("MODERATION_PROVIDER 未设置(内容安全厂商未选定)")
|
||||
builder = _PROVIDER_BUILDERS.get(name)
|
||||
if builder is None:
|
||||
raise ModerationConfigError(f"未知内容安全厂商: {name}")
|
||||
return builder()
|
||||
|
||||
|
||||
def get_text_provider() -> ModerationProvider:
|
||||
return _build_provider()
|
||||
@@ -0,0 +1,17 @@
|
||||
from .base import (
|
||||
ModerationConfigError,
|
||||
ModerationError,
|
||||
ModerationProvider,
|
||||
ModerationResult,
|
||||
ModerationUnavailableError,
|
||||
Verdict,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"ModerationConfigError",
|
||||
"ModerationError",
|
||||
"ModerationProvider",
|
||||
"ModerationResult",
|
||||
"ModerationUnavailableError",
|
||||
"Verdict",
|
||||
]
|
||||
@@ -0,0 +1,49 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Any, Protocol
|
||||
|
||||
|
||||
class ModerationError(RuntimeError):
|
||||
"""Base error for content-moderation failures."""
|
||||
|
||||
|
||||
class ModerationUnavailableError(ModerationError):
|
||||
"""Provider could not return a verdict (down / timeout / not implemented).
|
||||
|
||||
The orchestration layer converts this into a fail-closed BLOCK or a
|
||||
fail-open PASS depending on ``MODERATION_FAIL_CLOSED`` policy.
|
||||
"""
|
||||
|
||||
|
||||
class ModerationConfigError(ModerationError):
|
||||
"""Provider is misconfigured (missing credentials / biz type / endpoint)."""
|
||||
|
||||
|
||||
class Verdict(str, Enum):
|
||||
PASS = "pass" # 放行
|
||||
REVIEW = "review" # 建议人工复审(是否等同拦截由 policy 决定)
|
||||
BLOCK = "block" # 拦截
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModerationResult:
|
||||
"""Normalized moderation verdict, decoupled from any vendor's response shape.
|
||||
|
||||
NOTE: ``raw`` is only for debugging in-process; never persist it verbatim
|
||||
(may contain the moderated content / vendor internals). Compliance logging
|
||||
stores a summary (verdict + labels + request_id), not ``raw``.
|
||||
"""
|
||||
|
||||
verdict: Verdict
|
||||
labels: tuple[str, ...] = ()
|
||||
score: int | None = None
|
||||
keywords: tuple[str, ...] = ()
|
||||
request_id: str = ""
|
||||
raw: Any = field(default=None, repr=False, compare=False)
|
||||
|
||||
|
||||
class ModerationProvider(Protocol):
|
||||
def moderate_text(self, text: str, *, biz_type: str = "") -> ModerationResult:
|
||||
...
|
||||
@@ -0,0 +1,117 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from dataclasses import dataclass
|
||||
|
||||
try:
|
||||
from ahocorapy.keywordtree import KeywordTree
|
||||
except ImportError: # pragma: no cover - exercised only when dependency is missing.
|
||||
KeywordTree = None
|
||||
|
||||
from apps.moderation.models import SensitiveWord
|
||||
from apps.moderation.normalization import normalize_text
|
||||
from apps.moderation.versioning import get_sensitive_words_version
|
||||
|
||||
from .base import ModerationResult, ModerationUnavailableError, Verdict
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SensitiveWordMatch:
|
||||
id: int
|
||||
word: str
|
||||
normalized_word: str
|
||||
category: str
|
||||
|
||||
|
||||
class KeywordMatcher:
|
||||
def __init__(self, records: list[SensitiveWordMatch]) -> None:
|
||||
self._records_by_keyword: dict[str, list[SensitiveWordMatch]] = {}
|
||||
if KeywordTree is None:
|
||||
raise ModerationUnavailableError("ahocorapy is not installed")
|
||||
self._tree = KeywordTree(case_insensitive=False)
|
||||
for record in records:
|
||||
self._records_by_keyword.setdefault(record.normalized_word, []).append(record)
|
||||
for keyword in self._records_by_keyword:
|
||||
self._tree.add(keyword)
|
||||
self._tree.finalize()
|
||||
|
||||
def search(self, text: str) -> list[SensitiveWordMatch]:
|
||||
normalized = normalize_text(text)
|
||||
if not normalized:
|
||||
return []
|
||||
|
||||
matches: list[SensitiveWordMatch] = []
|
||||
seen: set[int] = set()
|
||||
for keyword, _index in self._tree.search_all(normalized):
|
||||
for record in self._records_by_keyword.get(keyword, ()):
|
||||
if record.id in seen:
|
||||
continue
|
||||
seen.add(record.id)
|
||||
matches.append(record)
|
||||
return matches
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _MatcherState:
|
||||
version: str
|
||||
matcher: KeywordMatcher
|
||||
|
||||
|
||||
_matcher_lock = threading.Lock()
|
||||
_matcher_state: _MatcherState | None = None
|
||||
|
||||
|
||||
def reset_keyword_matcher_cache() -> None:
|
||||
global _matcher_state
|
||||
with _matcher_lock:
|
||||
_matcher_state = None
|
||||
|
||||
|
||||
def _load_active_records() -> list[SensitiveWordMatch]:
|
||||
rows = SensitiveWord.objects.filter(
|
||||
is_active=True,
|
||||
action=SensitiveWord.Action.BLOCK,
|
||||
).only("id", "word", "normalized_word", "category")
|
||||
return [
|
||||
SensitiveWordMatch(
|
||||
id=row.id,
|
||||
word=row.word,
|
||||
normalized_word=row.normalized_word,
|
||||
category=row.category,
|
||||
)
|
||||
for row in rows.order_by("id")
|
||||
if row.normalized_word
|
||||
]
|
||||
|
||||
|
||||
def get_keyword_matcher() -> KeywordMatcher:
|
||||
global _matcher_state
|
||||
version = get_sensitive_words_version()
|
||||
state = _matcher_state
|
||||
if state is not None and state.version == version:
|
||||
return state.matcher
|
||||
|
||||
with _matcher_lock:
|
||||
state = _matcher_state
|
||||
if state is not None and state.version == version:
|
||||
return state.matcher
|
||||
matcher = KeywordMatcher(_load_active_records())
|
||||
_matcher_state = _MatcherState(version=version, matcher=matcher)
|
||||
return matcher
|
||||
|
||||
|
||||
class KeywordModerationProvider:
|
||||
def moderate_text(self, text: str, *, biz_type: str = "") -> ModerationResult:
|
||||
matches = get_keyword_matcher().search(text)
|
||||
if not matches:
|
||||
return ModerationResult(verdict=Verdict.PASS)
|
||||
|
||||
labels = tuple(dict.fromkeys(match.category for match in matches))
|
||||
keywords = tuple(dict.fromkeys(match.word for match in matches))
|
||||
word_ids = tuple(match.id for match in matches)
|
||||
return ModerationResult(
|
||||
verdict=Verdict.BLOCK,
|
||||
labels=labels,
|
||||
keywords=keywords,
|
||||
raw={"word_ids": word_ids},
|
||||
)
|
||||
@@ -0,0 +1,94 @@
|
||||
"""Prompt moderation orchestration.
|
||||
|
||||
T-604 只做本地关键词 prompt 审核:在图片下载、定价、预扣点和上游调用之前执行。
|
||||
MVP 不做输出审核、不做图片审核,也不保存用户 prompt 原文。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from . import policy
|
||||
from .providers import ModerationError, ModerationResult, Verdict
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModerationOutcome:
|
||||
stage: str
|
||||
verdict: Verdict
|
||||
blocked: bool
|
||||
labels: tuple[str, ...] = ()
|
||||
keywords: tuple[str, ...] = ()
|
||||
matched_word_ids: tuple[int, ...] = ()
|
||||
request_id: str = ""
|
||||
reason: str = ""
|
||||
result: ModerationResult | None = field(default=None, repr=False)
|
||||
|
||||
|
||||
def _passed(stage: str) -> ModerationOutcome:
|
||||
return ModerationOutcome(stage=stage, verdict=Verdict.PASS, blocked=False)
|
||||
|
||||
|
||||
def _on_provider_failure(stage: str, exc: Exception) -> ModerationOutcome:
|
||||
"""provider 不可用时的兜底:fail-closed → BLOCK,fail-open → PASS。"""
|
||||
if policy.fail_closed():
|
||||
logger.warning("moderation fail-closed at %s: %s", stage, exc.__class__.__name__)
|
||||
return ModerationOutcome(
|
||||
stage=stage,
|
||||
verdict=Verdict.BLOCK,
|
||||
blocked=True,
|
||||
reason=f"fail-closed: {exc.__class__.__name__}",
|
||||
)
|
||||
logger.warning("moderation fail-open at %s: %s", stage, exc.__class__.__name__)
|
||||
return ModerationOutcome(stage=stage, verdict=Verdict.PASS, blocked=False, reason="fail-open")
|
||||
|
||||
|
||||
def _evaluate(stage: str, call) -> ModerationOutcome:
|
||||
if not policy.moderation_enabled():
|
||||
return _passed(stage)
|
||||
try:
|
||||
result = call()
|
||||
except ModerationError as exc:
|
||||
outcome = _on_provider_failure(stage, exc)
|
||||
_record(outcome)
|
||||
return outcome
|
||||
|
||||
outcome = ModerationOutcome(
|
||||
stage=stage,
|
||||
verdict=result.verdict,
|
||||
blocked=policy.verdict_is_blocking(result.verdict),
|
||||
labels=result.labels,
|
||||
keywords=result.keywords,
|
||||
matched_word_ids=tuple(result.raw.get("word_ids", ())) if isinstance(result.raw, dict) else (),
|
||||
request_id=result.request_id,
|
||||
result=result,
|
||||
)
|
||||
_record(outcome)
|
||||
return outcome
|
||||
|
||||
|
||||
def _record(outcome: ModerationOutcome) -> None:
|
||||
logger.info(
|
||||
"moderation stage=%s verdict=%s blocked=%s labels=%s word_ids=%s request_id=%s",
|
||||
outcome.stage,
|
||||
outcome.verdict.value,
|
||||
outcome.blocked,
|
||||
",".join(outcome.labels),
|
||||
",".join(str(word_id) for word_id in outcome.matched_word_ids),
|
||||
outcome.request_id,
|
||||
)
|
||||
|
||||
|
||||
def moderate_prompt(*, user=None, api_key=None, prompt: str) -> ModerationOutcome:
|
||||
return _evaluate(
|
||||
"prompt",
|
||||
lambda: policy.get_text_provider().moderate_text(
|
||||
prompt, biz_type=policy.biz_type("input")
|
||||
),
|
||||
)
|
||||
|
||||
def moderate_input(*, user=None, prompt: str, **kwargs) -> ModerationOutcome:
|
||||
return moderate_prompt(user=user, prompt=prompt)
|
||||
@@ -0,0 +1,18 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from django.db.models.signals import post_delete, post_save
|
||||
from django.dispatch import receiver
|
||||
|
||||
from .models import SensitiveWord
|
||||
from .versioning import bump_sensitive_words_version
|
||||
|
||||
|
||||
@receiver(post_save, sender=SensitiveWord)
|
||||
@receiver(post_delete, sender=SensitiveWord)
|
||||
def invalidate_sensitive_words_cache(**kwargs) -> None:
|
||||
bump_sensitive_words_version()
|
||||
try:
|
||||
from .providers.keyword import reset_keyword_matcher_cache
|
||||
except Exception:
|
||||
return
|
||||
reset_keyword_matcher_cache()
|
||||
@@ -0,0 +1,106 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from django.core.cache import cache
|
||||
from django.core.exceptions import ValidationError
|
||||
from django.test import TestCase, override_settings
|
||||
|
||||
from .models import SensitiveWord
|
||||
from .normalization import normalize_text
|
||||
from .providers.keyword import reset_keyword_matcher_cache
|
||||
from .services import moderate_prompt
|
||||
from .versioning import get_sensitive_words_version
|
||||
|
||||
|
||||
@override_settings(MODERATION_CACHE_VERSION_KEY="test:moderation:sensitive_words:version")
|
||||
class SensitiveWordNormalizationTests(TestCase):
|
||||
def setUp(self):
|
||||
cache.clear()
|
||||
reset_keyword_matcher_cache()
|
||||
|
||||
def tearDown(self):
|
||||
reset_keyword_matcher_cache()
|
||||
cache.clear()
|
||||
|
||||
def test_normalize_text_removes_separators_zero_width_and_case(self):
|
||||
self.assertEqual(normalize_text("Bad-\u200b Word"), "badword")
|
||||
self.assertEqual(normalize_text("敏 感-词"), normalize_text("敏感词"))
|
||||
|
||||
def test_sensitive_word_saves_normalized_word_and_rejects_empty_normalized_word(self):
|
||||
word = SensitiveWord.objects.create(word=" 敏 感-词 ", category="")
|
||||
|
||||
self.assertEqual(word.category, "custom")
|
||||
self.assertEqual(word.normalized_word, "敏感词")
|
||||
|
||||
with self.assertRaises(ValidationError):
|
||||
SensitiveWord.objects.create(word=" - \u200b ")
|
||||
|
||||
def test_word_save_bumps_shared_cache_version(self):
|
||||
before = get_sensitive_words_version()
|
||||
|
||||
SensitiveWord.objects.create(word="敏感词")
|
||||
|
||||
self.assertNotEqual(get_sensitive_words_version(), before)
|
||||
|
||||
|
||||
@override_settings(
|
||||
MODERATION_ENABLED=True,
|
||||
MODERATION_PROVIDER="keyword",
|
||||
MODERATION_CACHE_VERSION_KEY="test:moderation:sensitive_words:version",
|
||||
)
|
||||
class KeywordModerationTests(TestCase):
|
||||
def setUp(self):
|
||||
cache.clear()
|
||||
reset_keyword_matcher_cache()
|
||||
|
||||
def tearDown(self):
|
||||
reset_keyword_matcher_cache()
|
||||
cache.clear()
|
||||
|
||||
def test_prompt_blocks_normalized_keyword_without_storing_prompt(self):
|
||||
word = SensitiveWord.objects.create(word="敏感词", category="policy")
|
||||
prompt = "请处理敏-感\u200b 词内容"
|
||||
|
||||
with self.assertLogs("apps.moderation.services", level="INFO") as logs:
|
||||
outcome = moderate_prompt(prompt=prompt)
|
||||
|
||||
self.assertTrue(outcome.blocked)
|
||||
self.assertEqual(outcome.labels, ("policy",))
|
||||
self.assertEqual(outcome.keywords, ("敏感词",))
|
||||
self.assertEqual(outcome.matched_word_ids, (word.id,))
|
||||
self.assertNotIn(prompt, "\n".join(logs.output))
|
||||
|
||||
def test_non_matching_prompt_passes(self):
|
||||
SensitiveWord.objects.create(word="敏感词")
|
||||
|
||||
outcome = moderate_prompt(prompt="正常业务标题")
|
||||
|
||||
self.assertFalse(outcome.blocked)
|
||||
|
||||
def test_matcher_rebuilds_after_shared_version_changes(self):
|
||||
self.assertFalse(moderate_prompt(prompt="新增词").blocked)
|
||||
|
||||
SensitiveWord.objects.create(word="新增词")
|
||||
|
||||
self.assertTrue(moderate_prompt(prompt="这个提示包含新增词").blocked)
|
||||
|
||||
|
||||
@override_settings(
|
||||
MODERATION_ENABLED=False,
|
||||
MODERATION_PROVIDER="keyword",
|
||||
MODERATION_CACHE_VERSION_KEY="test:moderation:sensitive_words:version",
|
||||
)
|
||||
class DisabledModerationTests(TestCase):
|
||||
def setUp(self):
|
||||
cache.clear()
|
||||
reset_keyword_matcher_cache()
|
||||
|
||||
def tearDown(self):
|
||||
reset_keyword_matcher_cache()
|
||||
cache.clear()
|
||||
|
||||
def test_disabled_moderation_is_noop(self):
|
||||
SensitiveWord.objects.create(word="敏感词")
|
||||
|
||||
outcome = moderate_prompt(prompt="敏感词")
|
||||
|
||||
self.assertFalse(outcome.blocked)
|
||||
@@ -0,0 +1,27 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from uuid import uuid4
|
||||
|
||||
from django.conf import settings
|
||||
from django.core.cache import cache
|
||||
|
||||
|
||||
def sensitive_words_version_key() -> str:
|
||||
return str(
|
||||
getattr(
|
||||
settings,
|
||||
"MODERATION_CACHE_VERSION_KEY",
|
||||
"moderation:sensitive_words:version",
|
||||
)
|
||||
or "moderation:sensitive_words:version"
|
||||
)
|
||||
|
||||
|
||||
def get_sensitive_words_version() -> str:
|
||||
return str(cache.get(sensitive_words_version_key()) or "0")
|
||||
|
||||
|
||||
def bump_sensitive_words_version() -> str:
|
||||
version = uuid4().hex
|
||||
cache.set(sensitive_words_version_key(), version, timeout=None)
|
||||
return version
|
||||
Reference in New Issue
Block a user