131 lines
3.9 KiB
Python
131 lines
3.9 KiB
Python
"""
|
|
PostgreSQL AI Agent - 使用LangChain实现的PostgreSQL数据库查询代理
|
|
需要设置环境变量: POSTGRES_HOST, POSTGRES_PORT, POSTGRES_USER, POSTGRES_PASSWORD, POSTGRES_DATABASE, OPENAI_API_KEY
|
|
"""
|
|
import os
|
|
import time
|
|
import logging
|
|
from langchain_community.utilities import SQLDatabase
|
|
from langchain.agents import create_sql_agent
|
|
from langchain.agents.agent_toolkits import SQLDatabaseToolkit
|
|
from langchain_openai import ChatOpenAI
|
|
from langchain.agents.agent_types import AgentType
|
|
|
|
logging.basicConfig(level=logging.INFO)
|
|
logger = logging.getLogger(__name__)
|
|
|
|
POD_NAME = os.getenv("POD_NAME", "unknown")
|
|
TEMPLATE_TYPE = os.getenv("TEMPLATE_TYPE", "postgresql_agent")
|
|
|
|
# PostgreSQL数据库配置
|
|
POSTGRES_HOST = os.getenv("POSTGRES_HOST", "localhost")
|
|
POSTGRES_PORT = os.getenv("POSTGRES_PORT", "5432")
|
|
POSTGRES_USER = os.getenv("POSTGRES_USER", "postgres")
|
|
POSTGRES_PASSWORD = os.getenv("POSTGRES_PASSWORD", "")
|
|
POSTGRES_DATABASE = os.getenv("POSTGRES_DATABASE", "postgres")
|
|
|
|
# OpenAI配置
|
|
OPENAI_API_KEY = os.getenv("OPENAI_API_KEY", "")
|
|
|
|
|
|
def create_postgresql_agent():
|
|
"""创建PostgreSQL数据库Agent"""
|
|
|
|
# 构建数据库URI
|
|
db_uri = f"postgresql+psycopg2://{POSTGRES_USER}:{POSTGRES_PASSWORD}@{POSTGRES_HOST}:{POSTGRES_PORT}/{POSTGRES_DATABASE}"
|
|
|
|
try:
|
|
# 连接数据库
|
|
db = SQLDatabase.from_uri(db_uri)
|
|
logger.info(f"✅ 成功连接到PostgreSQL数据库: {POSTGRES_HOST}:{POSTGRES_PORT}/{POSTGRES_DATABASE}")
|
|
|
|
# 显示可用的表
|
|
tables = db.get_usable_table_names()
|
|
logger.info(f"可用的表: {tables}")
|
|
|
|
except Exception as e:
|
|
logger.error(f"❌ 数据库连接失败: {str(e)}")
|
|
return None
|
|
|
|
# 初始化LLM
|
|
if not OPENAI_API_KEY:
|
|
logger.error("❌ 未设置OPENAI_API_KEY")
|
|
return None
|
|
|
|
llm = ChatOpenAI(
|
|
temperature=0,
|
|
model="gpt-3.5-turbo",
|
|
openai_api_key=OPENAI_API_KEY
|
|
)
|
|
|
|
# 创建SQL工具包
|
|
toolkit = SQLDatabaseToolkit(db=db, llm=llm)
|
|
|
|
# 创建SQL Agent
|
|
agent_executor = create_sql_agent(
|
|
llm=llm,
|
|
toolkit=toolkit,
|
|
agent_type=AgentType.ZERO_SHOT_REACT_DESCRIPTION,
|
|
verbose=True,
|
|
handle_parsing_errors=True,
|
|
max_iterations=5
|
|
)
|
|
|
|
return agent_executor
|
|
|
|
|
|
def main():
|
|
"""主函数 - PostgreSQL Agent主循环"""
|
|
logger.info(f"PostgreSQL Agent启动: {POD_NAME} (模板: {TEMPLATE_TYPE})")
|
|
logger.info(f"数据库配置: {POSTGRES_HOST}:{POSTGRES_PORT}/{POSTGRES_DATABASE}")
|
|
|
|
# 创建Agent
|
|
agent = create_postgresql_agent()
|
|
|
|
if agent is None:
|
|
logger.error("Agent创建失败,请检查配置")
|
|
# 保持容器运行
|
|
while True:
|
|
logger.info(f"[{POD_NAME}] 等待正确的配置...")
|
|
time.sleep(30)
|
|
return
|
|
|
|
logger.info("✅ PostgreSQL Agent创建成功,开始运行...")
|
|
|
|
# 示例查询列表
|
|
sample_queries = [
|
|
"列出数据库中所有的表和视图",
|
|
"描述每个表的结构和主键",
|
|
"统计每个表的记录数",
|
|
"查询数据库的版本信息",
|
|
"显示最大的3个表",
|
|
]
|
|
|
|
query_index = 0
|
|
|
|
while True:
|
|
try:
|
|
# 每2分钟执行一次示例查询
|
|
query = sample_queries[query_index % len(sample_queries)]
|
|
logger.info(f"\n{'='*60}")
|
|
logger.info(f"🐘 执行查询: {query}")
|
|
logger.info(f"{'='*60}\n")
|
|
|
|
# 执行Agent
|
|
result = agent.invoke({"input": query})
|
|
|
|
logger.info(f"\n✅ 结果:\n{result['output']}\n")
|
|
|
|
query_index += 1
|
|
|
|
except Exception as e:
|
|
logger.error(f"❌ 查询执行失败: {str(e)}")
|
|
|
|
# 等待120秒后执行下一个查询
|
|
logger.info(f"[{POD_NAME}] 等待下一次查询...")
|
|
time.sleep(120)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|