Files

95 lines
3.0 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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)