| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593 |
- """
- 数据分析 Agent:集成 SQL 查询和 Python 代码执行能力
- 让 LLM 自动决定使用哪些工具来回答用户问题
- 安全特性:
- - SQL 注入防护:关键词黑名单 + 只读权限检查
- - 恶意代码防护:危险模块黑名单 + 安全沙箱
- - 数据泄露防护:查询行数限制 + 敏感字段脱敏
- - 速率限制:防止恶意刷接口
- - 日志审计:记录所有操作到文件
- """
- import os
- import re
- import json
- import logging
- import time
- from datetime import datetime
- from io import StringIO
- from contextlib import redirect_stdout
- from typing import Optional, Dict, Any
- from collections import defaultdict
- from dotenv import load_dotenv
- from langchain_openai import ChatOpenAI
- from langchain_core.prompts import ChatPromptTemplate
- from langchain_classic.agents import create_tool_calling_agent, AgentExecutor
- from langchain_core.tools import tool
- from langchain_community.utilities import SQLDatabase
- # ============================================================
- # 配置日志审计系统
- # ============================================================
- logging.basicConfig(
- level=logging.INFO,
- format='%(asctime)s - %(levelname)s - %(name)s - %(message)s',
- handlers=[
- logging.StreamHandler(),
- logging.FileHandler('agent_audit.log', encoding='utf-8')
- ]
- )
- log = logging.getLogger(__name__)
- # 加载环境变量配置文件
- load_dotenv()
- # ============================================================
- # 安全配置常量
- # ============================================================
- # 最大查询返回行数
- MAX_QUERY_ROWS = 1000
- # 查询超时时间(秒)
- QUERY_TIMEOUT = 30
- # 速率限制配置
- RATE_LIMIT_WINDOW = 60 # 时间窗口(秒)
- RATE_LIMIT_MAX_CALLS = 20 # 窗口内最大调用次数
- # SQL 危险关键词
- SQL_FORBIDDEN_KEYWORDS = [
- 'INSERT', 'UPDATE', 'DELETE', 'DROP', 'CREATE', 'ALTER',
- 'TRUNCATE', 'GRANT', 'REVOKE', 'EXEC', 'EXECUTE'
- ]
- # Python 危险模块
- PYTHON_FORBIDDEN_MODULES = [
- 'os', 'subprocess', 'sys', 'shutil', 'socket',
- 'pickle', 'marshal', 'ctypes', 'multiprocessing'
- ]
- # ============================================================
- # 速率限制器
- # ============================================================
- class RateLimiter:
- """简单的速率限制器,基于内存存储"""
-
- def __init__(self, window_seconds: int = 60, max_calls: int = 20):
- self.window_seconds = window_seconds
- self.max_calls = max_calls
- self.calls: Dict[str, list] = defaultdict(list)
-
- def is_allowed(self, client_id: str = 'default') -> tuple:
- """
- 检查是否允许调用
-
- 返回:
- (is_allowed: bool, remaining: int, reset_time: int)
- """
- now = time.time()
- calls = self.calls[client_id]
-
- # 清理过期的调用记录
- calls[:] = [t for t in calls if now - t < self.window_seconds]
-
- if len(calls) >= self.max_calls:
- reset_time = int(self.window_seconds - (now - calls[0]))
- return False, 0, reset_time
-
- calls.append(now)
- remaining = self.max_calls - len(calls)
- return True, remaining, self.window_seconds
- # 全局速率限制器实例
- rate_limiter = RateLimiter(RATE_LIMIT_WINDOW, RATE_LIMIT_MAX_CALLS)
- # ============================================================
- # 审计日志记录器
- # ============================================================
- class AuditLogger:
- """审计日志记录器,记录所有 Agent 操作"""
-
- @staticmethod
- def log_operation(
- operation_type: str,
- input_data: str,
- output_data: str = None,
- success: bool = True,
- error: str = None,
- user_id: str = 'anonymous',
- **extra
- ):
- """记录操作日志"""
- log_entry = {
- 'timestamp': datetime.now().isoformat(),
- 'operation': operation_type,
- 'user_id': user_id,
- 'input': input_data[:500] if input_data else None,
- 'output': output_data[:500] if output_data else None,
- 'success': success,
- 'error': error,
- **extra
- }
-
- if success:
- log.info(f"[AUDIT] {json.dumps(log_entry, ensure_ascii=False)}")
- else:
- log.error(f"[AUDIT] {json.dumps(log_entry, ensure_ascii=False)}")
- audit_logger = AuditLogger()
- # ============================================================
- # 数据库连接初始化
- # ============================================================
- db_uri = os.getenv("DATABASE_URI")
- if not db_uri:
- raise ValueError("DATABASE_URI 环境变量未设置,请检查 .env 配置文件")
- db = SQLDatabase.from_uri(db_uri)
- # ============================================================
- # 工具 1: SQL 查询工具
- # ============================================================
- @tool
- def run_sql_query(query: str) -> str:
- """
- 执行只读 SQL 查询,用于从数据库查询数据。
-
- 该工具只能执行 SELECT 查询,禁止执行 INSERT、UPDATE、DELETE 等修改操作。
-
- 安全措施:
- - 禁止危险 SQL 关键词
- - 强制限制返回行数
- - 查询超时控制
-
- 参数:
- query: SQL 查询语句(仅限 SELECT 语句)
-
- 返回:
- 查询结果的字符串格式
- """
- start_time = time.time()
-
- # 安全检查 1: SQL 关键词黑名单检测(使用单词边界匹配)
- query_upper = query.strip().upper()
- for keyword in SQL_FORBIDDEN_KEYWORDS:
- if re.search(rf'\b{keyword}\b', query_upper):
- error_msg = f"安全限制:禁止执行 {keyword} 操作,仅允许 SELECT 查询"
- audit_logger.log_operation(
- operation_type='SQL_QUERY',
- input_data=query,
- success=False,
- error=error_msg
- )
- return f"❌ {error_msg}"
-
- # 安全检查 2: 强制添加 LIMIT(如果用户没有指定)
- if 'LIMIT' not in query_upper:
- query = query.rstrip(';').strip()
- query = f"{query} LIMIT {MAX_QUERY_ROWS}"
-
- # 安全检查 3: 验证 LIMIT 值不超过最大值
- limit_match = re.search(r'LIMIT\s+(\d+)', query, re.IGNORECASE)
- if limit_match:
- limit_value = int(limit_match.group(1))
- if limit_value > MAX_QUERY_ROWS:
- query = re.sub(r'LIMIT\s+\d+', f'LIMIT {MAX_QUERY_ROWS}', query, flags=re.IGNORECASE)
-
- try:
- result = db.run(query)
- execution_time = time.time() - start_time
-
- # 记录审计日志
- audit_logger.log_operation(
- operation_type='SQL_QUERY',
- input_data=query,
- output_data=result[:200] if result else None,
- success=True,
- execution_time=f"{execution_time:.3f}s"
- )
-
- log.info(f"SQL 查询执行成功 | 耗时: {execution_time:.3f}s")
-
- # 检查是否超时
- if execution_time > QUERY_TIMEOUT:
- log.warning(f"查询执行时间超过阈值: {execution_time:.3f}s > {QUERY_TIMEOUT}s")
-
- return result
-
- except Exception as e:
- error_msg = str(e)
- audit_logger.log_operation(
- operation_type='SQL_QUERY',
- input_data=query,
- success=False,
- error=error_msg
- )
- log.error(f"SQL 查询执行失败: {error_msg}")
- return f"❌ SQL 查询执行失败: {error_msg}"
- # ============================================================
- # 工具 2: Python 代码执行工具
- # ============================================================
- @tool
- def execute_python_code(code: str) -> str:
- """
- 执行 Python 代码进行数据分析、可视化或复杂计算。
-
- 【重要】已预导入以下模块,直接使用变量名即可,无需 import:
- - plt: matplotlib.pyplot(绑图库)
- - pd: pandas(数据分析库)
- - np: numpy(数值计算库)
- - json: json(JSON处理)
-
- 示例用法:
- # 直接使用 plt,不要写 import matplotlib.pyplot as plt
- plt.figure(figsize=(10, 6))
- plt.bar(['A', 'B'], [10, 20])
- plt.savefig('chart.png')
-
- # 直接使用 pd 和 np
- df = pd.DataFrame({'name': ['张三', '李四'], 'age': [25, 30]})
- arr = np.array([1, 2, 3, 4, 5])
-
- 安全限制:
- - 禁止使用 import 语句(已预导入常用库)
- - 禁止使用 eval()、exec()、open() 等危险函数
-
- 参数:
- code: Python 代码字符串(直接使用预导入的 plt/pd/np/json)
-
- 返回:
- 代码执行的输出结果
- """
- start_time = time.time()
-
- # 安全检查 1: 禁止危险模块
- for module in PYTHON_FORBIDDEN_MODULES:
- if re.search(rf'\bimport\s+{module}\b', code) or re.search(rf'\bfrom\s+{module}\b', code):
- error_msg = f"安全限制:禁止导入 {module} 模块"
- audit_logger.log_operation(
- operation_type='PYTHON_CODE',
- input_data=code[:200],
- success=False,
- error=error_msg
- )
- return f"❌ {error_msg}"
-
- # 安全检查 2: 禁止危险内置函数
- dangerous_functions = ['eval', 'exec', 'compile', 'open', '__import__']
- for func in dangerous_functions:
- if re.search(rf'\b{func}\s*\(', code):
- error_msg = f"安全限制:禁止使用 {func}() 函数"
- audit_logger.log_operation(
- operation_type='PYTHON_CODE',
- input_data=code[:200],
- success=False,
- error=error_msg
- )
- return f"❌ {error_msg}"
-
- # 配置 matplotlib 中文字体
- import matplotlib
- import matplotlib.pyplot as plt
-
- chinese_fonts = ['SimHei', 'Microsoft YaHei', 'STSong', 'SimSun', 'KaiTi', 'FangSong']
- available_fonts = [f.name for f in matplotlib.font_manager.fontManager.ttflist]
-
- for font in chinese_fonts:
- if font in available_fonts:
- plt.rcParams['font.sans-serif'] = [font]
- break
-
- plt.rcParams['axes.unicode_minus'] = False
-
- # 创建受限的执行环境(白名单模式)
- safe_builtins = {
- # 基础函数
- 'print': print,
- 'len': len,
- 'range': range,
- 'enumerate': enumerate,
- 'zip': zip,
- 'map': map,
- 'filter': filter,
- 'sorted': sorted,
- 'reversed': reversed,
- # 类型
- 'list': list,
- 'dict': dict,
- 'tuple': tuple,
- 'set': set,
- 'str': str,
- 'int': int,
- 'float': float,
- 'bool': bool,
- 'bytes': bytes,
- # 数学函数
- 'sum': sum,
- 'min': min,
- 'max': max,
- 'abs': abs,
- 'round': round,
- 'pow': pow,
- 'divmod': divmod,
- # 类型检查
- 'isinstance': isinstance,
- 'type': type,
- 'hasattr': hasattr,
- 'getattr': getattr,
- 'callable': callable,
- # 异常
- 'Exception': Exception,
- 'ValueError': ValueError,
- 'TypeError': TypeError,
- 'KeyError': KeyError,
- 'IndexError': IndexError,
- 'AttributeError': AttributeError,
- 'RuntimeError': RuntimeError,
- 'StopIteration': StopIteration,
- # 其他安全函数
- 'repr': repr,
- 'hash': hash,
- 'id': id,
- }
-
- # 允许的安全模块白名单(用于受限的 __import__)
- ALLOWED_MODULES = {'pandas', 'numpy', 'matplotlib', 'matplotlib.pyplot', 'json'}
-
- def safe_import(name, *args, **kwargs):
- """安全的 import 函数,只允许导入白名单中的模块"""
- if name.split('.')[0] not in {'pandas', 'numpy', 'matplotlib', 'json'}:
- raise ImportError(f"安全限制:禁止导入模块 '{name}',请直接使用预导入的 plt/pd/np/json")
- return __import__(name, *args, **kwargs)
-
- safe_builtins['__import__'] = safe_import
-
- local_vars = {
- 'pd': __import__('pandas'),
- 'np': __import__('numpy'),
- 'plt': plt,
- 'json': json,
- }
-
- # 捕获标准输出
- output_buffer = StringIO()
-
- try:
- with redirect_stdout(output_buffer):
- exec(code, {"__builtins__": safe_builtins}, local_vars)
-
- output = output_buffer.getvalue()
- execution_time = time.time() - start_time
-
- if not output:
- result_vars = {k: v for k, v in local_vars.items()
- if not k.startswith('_') and k not in ['pd', 'np', 'plt', 'json']}
- if result_vars:
- output = "执行成功,生成的变量:\n"
- for var_name, var_value in result_vars.items():
- output += f" {var_name} = {repr(var_value)[:200]}\n"
- else:
- output = "✅ 代码执行成功(无输出)"
-
- # 记录审计日志
- audit_logger.log_operation(
- operation_type='PYTHON_CODE',
- input_data=code[:200],
- output_data=output[:200],
- success=True,
- execution_time=f"{execution_time:.3f}s"
- )
-
- log.info(f"Python 代码执行成功 | 耗时: {execution_time:.3f}s")
- return output
-
- except Exception as e:
- error_msg = str(e)
- audit_logger.log_operation(
- operation_type='PYTHON_CODE',
- input_data=code[:200],
- success=False,
- error=error_msg
- )
- log.error(f"Python 代码执行失败: {error_msg}")
- return f"❌ Python 代码执行失败: {error_msg}"
- # ============================================================
- # Agent 创建
- # ============================================================
- def create_data_analysis_agent(model_name: str = None, temperature: float = 0):
- """
- 创建数据分析 Agent,集成 SQL 查询和 Python 代码执行能力
-
- 参数:
- model_name: 使用的模型名称,默认从环境变量读取
- temperature: 模型温度,默认为 0(更确定性的输出)
- """
- # 从环境变量读取配置
- api_key = os.getenv("OPENAI_API_KEY")
- api_base = os.getenv("OPENAI_API_BASE")
- model_name = os.getenv("MODEL_NAME")
-
- if not api_key:
- raise ValueError("OPENAI_API_KEY 环境变量未设置,请检查 .env 配置文件")
-
- # 初始化 LLM
- llm_config = {
- "model": model_name,
- "temperature": temperature,
- }
-
- if api_base:
- llm_config["openai_api_base"] = api_base
- log.info(f"使用自定义 API 端点: {api_base}")
-
- llm = ChatOpenAI(**llm_config)
- log.info(f"初始化模型: {model_name}")
-
- # 定义工具列表
- tools = [run_sql_query, execute_python_code]
-
- # 创建 Prompt 模板
- prompt = ChatPromptTemplate.from_messages([
- ("system", """你是一个专业的数据分析助手,可以访问公司数据库并执行 Python 代码进行数据分析。
- 数据库包含以下表:
- 1. employees (员工表)
- - id: 员工ID
- - name: 姓名
- - department: 所属部门
- - salary: 月薪
- - hire_date: 入职日期
- 2. products (产品表)
- - id: 产品ID
- - product_name: 商品名称
- - category: 商品分类
- - price: 单价
- - stock: 当前库存量
- 3. orders (订单表)
- - id: 订单ID
- - employee_id: 下单员工ID
- - product_id: 购买商品ID
- - quantity: 购买数量
- - order_date: 下单日期
- 工作流程:
- 1. 理解用户的问题
- 2. 决定需要哪些数据
- 3. 使用 run_sql_query 工具查询数据库
- 4. 如果需要复杂分析或可视化,使用 execute_python_code 工具
- 5. 用自然语言总结结果
- 注意事项:
- - SQL 查询仅支持 SELECT 语句
- - 使用 Python 时,已预导入 pandas as pd、numpy as np、matplotlib.pyplot as plt
- - 如果查询结果需要进一步分析,可以在 Python 代码中使用 SQL 查询的结果
- - 始终用中文回答问题
- - 提供清晰的分析结论和建议
- - 注意数据安全,不要泄露敏感信息
- """),
- ("human", "{input}"),
- ("placeholder", "{agent_scratchpad}"),
- ])
-
- # 创建 Agent
- agent = create_tool_calling_agent(llm, tools, prompt)
-
- # 创建 AgentExecutor
- agent_executor = AgentExecutor(
- agent=agent,
- tools=tools,
- verbose=True,
- handle_parsing_errors=True,
- max_iterations=10,
- )
-
- log.info("Agent 创建成功")
- return agent_executor
- # ============================================================
- # 主函数:统一 Agent 入口(带速率限制)
- # ============================================================
- def run_data_analysis(user_query: str, model_name: str = None, client_id: str = 'default') -> str:
- """
- 统一数据分析 Agent 入口
-
- 参数:
- user_query: 用户的自然语言查询问题
- model_name: 可选,指定使用的模型名称
- client_id: 客户端标识,用于速率限制
-
- 返回:
- Agent 的分析结果
- """
- # 速率限制检查
- is_allowed, remaining, reset_time = rate_limiter.is_allowed(client_id)
-
- if not is_allowed:
- error_msg = f"请求过于频繁,请在 {reset_time} 秒后重试"
- log.warning(f"速率限制触发: client_id={client_id}")
- return f"❌ {error_msg}"
-
- # 记录用户查询
- audit_logger.log_operation(
- operation_type='USER_QUERY',
- input_data=user_query,
- client_id=client_id
- )
-
- try:
- agent = create_data_analysis_agent(model_name=model_name)
- result = agent.invoke({"input": user_query})
-
- log.info(f"查询处理完成: {user_query[:50]}... | 剩余调用次数: {remaining}")
- return result["output"]
-
- except Exception as e:
- error_msg = str(e)
- audit_logger.log_operation(
- operation_type='AGENT_ERROR',
- input_data=user_query,
- success=False,
- error=error_msg,
- client_id=client_id
- )
- log.error(f"Agent 执行失败: {error_msg}")
- return f"❌ Agent 执行失败: {error_msg}"
- # ============================================================
- # 测试入口
- # ============================================================
- if __name__ == "__main__":
- print("=" * 60)
- print("数据分析 Agent 测试(安全增强版)")
- print("=" * 60)
-
- # 测试: 绘制图表
- # print("\n【测试】绘制图表 - 各部门薪资分布柱状图:")
- # result = run_data_analysis(
- # "查询各部门的平均薪资,并用 matplotlib 绘制柱状图,"
- # "图表标题为'各部门平均薪资分布',x轴为部门名称,y轴为平均薪资,"
- # "保存图表到 d:/agentlearning/lqq-agent-study/salary_chart.png"
- # )
- # 测试: 绘制图表
- print("\n【直接回答版:")
- result = run_data_analysis(
- "今天上海天气怎么样,会来台风吗"
- )
- print(result)
|