95 lines
3.0 KiB
Python
95 lines
3.0 KiB
Python
"""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)
|