Files
rag/core/llm_utils.py
lacerate551 90b915232a fix(security): 代码审查安全加固 — 三批次修复(6H/13M/12L)
第一批(快速修复):
- H6: main.py --debug 默认值 True→False,防止 Werkzeug RCE
- M2+M3: /search 增加 validate_query + top_k 范围限制(1-50)
- M4: context_count 范围限制(0-10) + 异常捕获
- L3: assert → raise RuntimeError(生产环境 API Key 检查)
- H1: SSE 错误事件移除 traceback 字段

第二批(安全加固):
- H2+H3: 文档接口路径遍历 realpath 校验 + 文件类型/大小限制
- H4+H5: 批量上传文件大小检查
- M6: LIKE 查询通配符转义
- M1: 37 处 str(e) 异常信息统一脱敏(6 文件)
- M5: CORS 生产环境限制来源
- M7: SESSION_MANAGER None 保护(503)
- M11: subprocess 参数注入防护(白名单 + -- 分隔符)

第三批(架构改进):
- M8+M9: 提取 JSON 解析共享工具(extract_json_object/list)
- M10: Prompt 注入检测防御(prompt_guard.py)
- M12: 解析器文件大小限制(Excel 50MB/TXT 20MB/PDF 100MB)
- M13: 全局单例竞态条件双重检查锁定(engine/bm25/intent_analyzer)
2026-06-05 15:26:32 +08:00

361 lines
9.8 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.
"""
LLM 调用工具函数
统一封装 LLM 调用模式,减少代码重复。
"""
import json
import re
import logging
from typing import List, Optional, Union, Iterator, Callable
logger = logging.getLogger(__name__)
def call_llm(
client,
prompt: str,
model: str,
temperature: float = 0.3,
max_tokens: int = 1000,
messages: List[dict] = None,
stream: bool = False,
**kwargs
) -> Union[str, Iterator, None]:
"""
统一的 LLM 调用封装
Args:
client: OpenAI 客户端实例
prompt: 用户提示(当 messages 为 None 时使用)
model: 模型名称
temperature: 温度参数 (0-1)
max_tokens: 最大 token 数
messages: 完整消息列表(优先于 prompt
stream: 是否启用流式输出
**kwargs: 其他参数(如 response_format, tools 等)
Returns:
流式模式返回 stream 迭代器,否则返回响应内容字符串
调用失败返回 None
Example:
# 简单调用
response = call_llm(client, "你好", model, temperature=0.3)
# 使用消息列表
response = call_llm(client, "", model, messages=[
{"role": "system", "content": "你是助手"},
{"role": "user", "content": "你好"}
])
# 流式调用
for chunk in call_llm(client, prompt, model, stream=True):
if chunk.choices and chunk.choices[0].delta.content:
yield chunk.choices[0].delta.content
"""
if messages is None:
messages = [{"role": "user", "content": prompt}]
try:
response = client.chat.completions.create(
model=model,
messages=messages,
temperature=temperature,
max_tokens=max_tokens,
stream=stream,
**kwargs
)
if stream:
return response
content = response.choices[0].message.content
# 推理模型兼容content 为空时尝试从 reasoning_content 提取
if not content or not content.strip():
reasoning = getattr(response.choices[0].message, 'reasoning_content', None)
if reasoning and reasoning.strip():
# 从思维链中提取 JSON 块作为内容
json_match = re.search(r'\{[\s\S]*\}', reasoning)
if json_match:
logger.info("LLM: content为空从reasoning_content提取JSON")
return json_match.group().strip()
logger.warning("LLM 返回空 content可能需要增大 max_tokens")
return None
return content.strip()
except Exception as e:
logger.warning(f"LLM 调用失败: {e}")
return None
def call_llm_stream(
client,
prompt: str,
model: str,
temperature: float = 0.3,
max_tokens: int = 1000,
messages: List[dict] = None,
error_prefix: str = "[错误]",
**kwargs
) -> Iterator[str]:
"""
流式 LLM 调用(生成器封装)
自动处理流式响应,逐块 yield 文本内容。
Args:
client: OpenAI 客户端实例
prompt: 用户提示
model: 模型名称
temperature: 温度参数
max_tokens: 最大 token 数
messages: 完整消息列表
error_prefix: 错误时的前缀
**kwargs: 其他参数
Yields:
响应文本片段
Example:
for text in call_llm_stream(client, "讲个故事", model):
print(text, end="", flush=True)
"""
if messages is None:
messages = [{"role": "user", "content": prompt}]
try:
stream = client.chat.completions.create(
model=model,
messages=messages,
temperature=temperature,
max_tokens=max_tokens,
stream=True,
**kwargs
)
for chunk in stream:
if chunk.choices and chunk.choices[0].delta.content:
yield chunk.choices[0].delta.content
except Exception as e:
logger.error(f"LLM 流式调用失败: {e}")
yield f"{error_prefix} 调用大模型失败: {str(e)}"
def call_llm_with_retry(
client,
prompt: str,
model: str,
max_retries: int = 3,
retry_delay: float = 1.0,
**kwargs
) -> Optional[str]:
"""
带重试的 LLM 调用
Args:
client: OpenAI 客户端实例
prompt: 用户提示
model: 模型名称
max_retries: 最大重试次数
retry_delay: 重试间隔(秒)
**kwargs: 传递给 call_llm 的其他参数
Returns:
响应内容,全部失败返回 None
"""
import time
for attempt in range(max_retries):
result = call_llm(client, prompt, model, **kwargs)
if result is not None:
return result
if attempt < max_retries - 1:
logger.debug(f"LLM 调用重试 {attempt + 1}/{max_retries}")
time.sleep(retry_delay)
logger.warning(f"LLM 调用失败,已重试 {max_retries}")
return None
def parse_json_from_response(content: str) -> Optional[dict]:
"""
从 LLM 响应中解析 JSON
自动处理 ```json 或 ``` 代码块格式
Args:
content: LLM 返回的原始内容
Returns:
解析后的字典,失败返回 None
"""
if not content:
return None
# 提取 JSON 代码块(支持 ```json 或 ```
json_match = re.search(r'```(?:json)?\s*([\s\S]*?)\s*```', content)
json_str = json_match.group(1) if json_match else content
try:
return json.loads(json_str.strip())
except (json.JSONDecodeError, TypeError, ValueError):
return None
def parse_json_list_from_response(content: str) -> Optional[List[dict]]:
"""
从 LLM 响应中解析 JSON 数组
Args:
content: LLM 返回的原始内容
Returns:
解析后的列表,失败返回 None
"""
if not content:
return None
# 提取 JSON 代码块
json_match = re.search(r'```(?:json)?\s*([\s\S]*?)\s*```', content)
json_str = json_match.group(1) if json_match else content
try:
result = json.loads(json_str.strip())
if isinstance(result, list):
return result
# 如果是 {"items": [...]} 格式,提取 items
if isinstance(result, dict):
for key in ['items', 'data', 'results', 'questions']:
if key in result and isinstance(result[key], list):
return result[key]
return None
except (json.JSONDecodeError, TypeError, ValueError):
return None
def extract_json_object(content: str) -> Optional[dict]:
"""
多策略从 LLM 响应中提取 JSON 对象(增强版)
在 parse_json_from_response 基础上增加 fallback 策略:
1. 先调用 parse_json_from_responsemarkdown 代码块 → 直接解析)
2. 失败后 fallback 到正则匹配最外层 {...} 块
Args:
content: LLM 返回的原始内容
Returns:
解析后的字典,全部策略失败返回 None
"""
if not content:
return None
# 策略1+2markdown 代码块提取 + 直接 json.loads
result = parse_json_from_response(content)
if result is not None and isinstance(result, dict):
return result
# 策略3fallback正则匹配最外层 JSON 对象 {...}
brace_match = re.search(r'\{[\s\S]*\}', content)
if brace_match:
try:
parsed = json.loads(brace_match.group(0))
if isinstance(parsed, dict):
return parsed
except (json.JSONDecodeError, TypeError, ValueError):
pass
return None
def extract_json_list(content: str) -> Optional[list]:
"""
多策略从 LLM 响应中提取 JSON 数组(增强版)
在 parse_json_list_from_response 基础上增加 fallback 策略:
1. 先调用 parse_json_list_from_responsemarkdown 代码块 → 直接解析 → 嵌套提取)
2. 失败后 fallback 到正则匹配最外层 [...] 块
Args:
content: LLM 返回的原始内容
Returns:
解析后的列表,全部策略失败返回 None
"""
if not content:
return None
# 策略1+2markdown 代码块提取 + 直接 json.loads + 嵌套 key 提取
result = parse_json_list_from_response(content)
if result is not None:
return result
# 策略3fallback正则匹配最外层 JSON 数组 [...]
bracket_match = re.search(r'\[[\s\S]*\]', content)
if bracket_match:
try:
parsed = json.loads(bracket_match.group(0))
if isinstance(parsed, list):
return parsed
except (json.JSONDecodeError, TypeError, ValueError):
pass
return None
# ==================== 便捷函数 ====================
def quick_ask(
client,
prompt: str,
model: str,
temperature: float = 0.3
) -> str:
"""
快速提问(简化版)
Args:
client: OpenAI 客户端
prompt: 问题
model: 模型
temperature: 温度
Returns:
回答内容,失败返回空字符串
"""
result = call_llm(client, prompt, model, temperature=temperature)
return result or ""
def quick_yes_no(
client,
prompt: str,
model: str,
keywords: List[str] = None
) -> bool:
"""
快速是/否判断
Args:
client: OpenAI 客户端
prompt: 问题(需包含判断标准)
model: 模型
keywords: 判断为 True 的关键词(默认 ["", "需要", "yes", "true"]
Returns:
布尔判断结果
"""
if keywords is None:
keywords = ["", "需要", "yes", "true"]
result = call_llm(client, prompt, model, temperature=0, max_tokens=10)
if result is None:
return False
result_lower = result.lower()
return any(kw.lower() in result_lower for kw in keywords)