大模型流式输出过滤敏感词

python 复制代码
"""
流式敏感词过滤器 ------ 完整可运行示例

功能:
  1. 支持跨 chunk 检测(敏感词被拆到多个 delta 里也能拦截)
  2. 命中时先输出安全前缀,再输出拦截提示,最后终止流
  3. 使用 Aho-Corasick 风格的最长词保留策略,保证不漏检、不误切
  4. 附带 mock 流演示 + 真实 OpenAI 调用入口

运行:
  python stream_guard_demo.py          # 用 mock 数据演示(无需 API Key)
  python stream_guard_demo.py --real   # 调用真实 OpenAI(需设置 OPENAI_API_KEY)
"""

import os
import re
import sys
import asyncio
import unicodedata
from typing import AsyncIterator, Iterable


# =============================================================
# 一、核心过滤器:StreamingSensitiveFilter
# =============================================================
class StreamingSensitiveFilter:
    """
    流式敏感词过滤器。

    设计要点:
      - 维护一个字符串缓冲区 buffer,所有进入的文本先追加到 buffer。
      - 每次追加后,先检测 buffer 是否命中敏感词:
          * 命中 → 释放命中位置之前的安全前缀,然后抛拦截信号并终止。
          * 未命中 → 只释放"确认安全"的前缀,末尾保留 KEEP_TAIL 个字符。
      - KEEP_TAIL = 最长敏感词长度 - 1,保证任何跨 chunk 的敏感词在补全的
        那一刻一定完整落在 buffer 中,从而被检出。
    """

    def __init__(
        self,
        sensitive_words: Iterable[str],
        *,
        block_message: str = "[系统提示:回答内容触发安全机制,已拦截]",
        normalize: bool = False,
    ):
        # 去重 + 去空,避免构造出空正则分支
        words = [w for w in dict.fromkeys(sensitive_words) if w]
        if not words:
            raise ValueError("敏感词列表不能为空")

        self.sensitive_words = words
        # re.escape 防止敏感词含 . * + ? ( ) 等正则元字符
        # 用 | 拼接成"任意一个词命中即算命中"的模式
        self._pattern = re.compile("|".join(re.escape(w) for w in words))

        # 最长词长度决定尾部保留量:
        # 只要留下 max_len - 1 个字符,跨 chunk 的词补全后一定完整出现
        self._max_len = max(len(w) for w in words)
        self._keep_tail = max(self._max_len - 1, 0)

        self.block_message = block_message
        # 是否对文本做归一化(去全半角/大小写差异),默认关闭以保持位置映射简单
        self.normalize = normalize

        # 运行期状态
        self._buffer = ""
        self._blocked = False

    # ---------- 内部工具 ----------
    def _normalize(self, text: str) -> str:
        """可选的归一化:NFKC 会把全角转半角、合并兼容字符。"""
        if not self.normalize:
            return text
        return unicodedata.normalize("NFKC", text).lower()

    def _find(self, text: str):
        """在(可能归一化后的)文本里查找敏感词,返回 match 或 None。"""
        return self._pattern.search(self._normalize(text))

    # ---------- 对外 API ----------
    def feed(self, token: str) -> list[str]:
        """
        输入一个流式 token,返回本次可以安全输出的文本片段列表。

        返回值的语义:
          - 正常情况:返回 0 或 1 个可安全输出的字符串(可能是空列表)
          - 命中敏感词:返回 [安全前缀, 拦截提示],并把 blocked 置为 True
          - 命中后再次调用:返回 [](幂等,防止误输出)
        """
        if self._blocked:
            return []

        self._buffer += token

        # ---- 1. 在完整 buffer 上检测(天然覆盖跨 chunk)----
        m = self._find(self._buffer)
        if m:
            # 命中位置以"归一化后文本"计算;未归一化时位置可直接用于原 buffer
            # 若开启了 normalize,位置可能偏移,这里退化为不输出前缀以保证安全
            if self.normalize:
                safe_prefix = ""
            else:
                safe_prefix = self._buffer[: m.start()]

            self._blocked = True
            out = []
            if safe_prefix:
                out.append(safe_prefix)
            out.append(self.block_message)
            return out

        # ---- 2. 未命中:只释放确认安全的前缀 ----
        if len(self._buffer) > self._keep_tail:
            release_len = len(self._buffer) - self._keep_tail
            safe_part = self._buffer[:release_len]
            self._buffer = self._buffer[release_len:]
            return [safe_part]

        # 缓冲区还没超过保留量,什么都不输出
        return []

    def flush(self) -> list[str]:
        """
        流结束时调用,把缓冲区剩余内容吐出去。
        若流已经因命中被终止,则什么都不返回。
        """
        if self._blocked:
            return []
        if not self._buffer:
            return []
        tail = self._buffer
        self._buffer = ""
        return [tail]

    @property
    def blocked(self) -> bool:
        """是否已经命中过敏感词(可用于外层提前终止)。"""
        return self._blocked


# =============================================================
# 二、把过滤器接到任意异步 token 流上
# =============================================================
async def guard_stream(
    token_iter: AsyncIterator[str],
    sensitive_words: Iterable[str],
    *,
    block_message: str = "[系统提示:回答内容触发安全机制,已拦截]",
) -> AsyncIterator[str]:
    """
    装饰器式用法:吃进原始 token 流,吐出经过安全过滤的文本片段。

    参数:
        token_iter      ------ 原始异步 token 流(每个元素是一小段文本)
        sensitive_words ------ 敏感词列表
        block_message   ------ 命中后输出的提示语
    """
    filt = StreamingSensitiveFilter(sensitive_words, block_message=block_message)

    async for token in token_iter:
        for piece in filt.feed(token):
            yield piece
        if filt.blocked:
            # 命中即终止:不再消费后续 token(真实场景可在此 cancel 上游请求)
            return

    # 正常结束,吐出残留尾巴
    for piece in filt.flush():
        yield piece


# =============================================================
# 三、Mock 流:模拟"敏感词被拆到多个 chunk"
# =============================================================
async def mock_token_stream() -> AsyncIterator[str]:
    """
    模拟 OpenAI 的 delta 流。
    注意 "极端词汇" 被故意拆成 "极" / "端" / "词" / "汇" 四个 chunk,
    用来验证跨 chunk 检测是否生效。
    """
    tokens = [
        "你好,",
        "这是",
        "一段",
        "正常",
        "的",
        "回答",
        "。",       # 到这里都应正常输出
        "但是",
        "这里",
        "出现",
        "极",       # ← 敏感词开始
        "端",
        "词",
        "汇",       # ← 敏感词结束,此处应被拦截
        "后面",
        "的内容",
        "不应",
        "出现",
    ]
    for t in tokens:
        await asyncio.sleep(0.05)  # 模拟网络延迟
        yield t


# =============================================================
# 四、真实 OpenAI 流(需要 OPENAI_API_KEY)
# =============================================================
async def openai_token_stream(user_query: str) -> AsyncIterator[str]:
    """
    调用真实 OpenAI 流式接口,逐个 yield delta 文本。
    需要环境变量 OPENAI_API_KEY。
    """
    from openai import AsyncOpenAI  # 延迟导入,避免 mock 模式依赖

    client = AsyncOpenAI()
    stream = await client.chat.completions.create(
        model="gpt-4o",
        messages=[{"role": "user", "content": user_query}],
        stream=True,
    )
    async for chunk in stream:
        delta = chunk.choices[0].delta.content
        if delta:
            yield delta


# =============================================================
# 五、演示 & 测试
# =============================================================
SENSITIVE_WORDS = ["暴力", "恐怖", "极端词汇", "竞品公司名字"]


async def demo_mock():
    """用 mock 流演示跨 chunk 拦截效果。"""
    print("=" * 60)
    print("【演示】Mock 流 ------ 敏感词 '极端词汇' 被拆到 4 个 chunk")
    print("=" * 60)

    collected = []
    print("过滤后输出:", end="", flush=True)
    async for piece in guard_stream(mock_token_stream(), SENSITIVE_WORDS):
        print(piece, end="", flush=True)
        collected.append(piece)
    print("\n")

    full_text = "".join(collected)
    assert "极端词汇" not in full_text, "❌ 敏感词泄露!"
    assert "[系统提示:回答内容触发安全机制,已拦截]" in full_text, "❌ 未输出拦截提示"
    print("✅ 跨 chunk 拦截成功:敏感词未泄露,且拦截提示已输出")


async def demo_edge_cases():
    """边界用例:短流、无敏感词、敏感词在开头/结尾。"""
    print("=" * 60)
    print("【测试】边界场景")
    print("=" * 60)

    cases = [
        ("无敏感词", ["hello ", "world", "!"]),
        ("敏感词在开头", ["恐怖", "袭击"]),
        ("敏感词在结尾", ["这是", "暴力"]),
        ("敏感词跨 3 个 chunk", ["极", "端", "词", "汇"]),
        ("单个 chunk 整词", ["这里出现恐怖袭击"]),
        ("空 token 混入", ["正常", "", "文本", ""]),
    ]

    for name, tokens in cases:
        async def gen(ts=tokens):
            for t in ts:
                yield t

        out = []
        async for piece in guard_stream(gen(), SENSITIVE_WORDS):
            out.append(piece)
        text = "".join(out)

        # 校验:输出中绝不能含任何完整敏感词
        leaked = [w for w in SENSITIVE_WORDS if w in text]
        status = "✅" if not leaked else "❌"
        print(f"{status} {name:16s} -> {text!r}")
        assert not leaked, f"{name} 泄露了敏感词:{leaked}"

    print()


async def demo_real(user_query: str):
    """真实 OpenAI 调用演示。"""
    print("=" * 60)
    print(f"【真实调用】问题:{user_query}")
    print("=" * 60)
    print("输出:", end="", flush=True)
    async for piece in guard_stream(openai_token_stream(user_query), SENSITIVE_WORDS):
        print(piece, end="", flush=True)
    print()


def main():
    if "--real" in sys.argv:
        if not os.getenv("OPENAI_API_KEY"):
            print("请先设置环境变量 OPENAI_API_KEY")
            sys.exit(1)
        asyncio.run(demo_real("请用一句话介绍你自己"))
    else:
        asyncio.run(demo_mock())
        asyncio.run(demo_edge_cases())


if __name__ == "__main__":
    main()
相关推荐
小溪学编程1 小时前
C语言篇:语言结构
c语言·开发语言·算法
老王爱玩车1 小时前
字符串和字符串函数
c语言·开发语言·数据结构·学习
2601_962885721 小时前
如何用 Python 统计 A 股连续涨停(几连板)?
java·linux·python
Carl_奕然1 小时前
【智能体】Loop 的四种设计模式之:Event-Driven Loop(2026 最新版)
人工智能·驱动开发·python·设计模式
mlidongfeng1 小时前
【学习笔记】【NV】GIN
开发语言·php
Carl_奕然1 小时前
【智能体】Loop 的四种设计模式之:Hill Climbing Loop(2026 最新版)
人工智能·python·设计模式
Omics Pro1 小时前
~30,000+引用!理论2005,R包2008,多组学集成AI增强
开发语言·数据库·人工智能·算法·机器学习·自然语言处理·r语言
专业程序开发源2 小时前
springboot全民健身和饮食健康管理系统29158-计算机课程设计、毕业设计
java·spring boot·后端·python·django·php·课程设计
小道士写程序2 小时前
Windows Rust 环境安装指导
开发语言·windows·rust