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