"""
流式敏感词过滤器 ------ 完整可运行示例
功能:
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()