-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathagent.py
More file actions
159 lines (127 loc) · 4.75 KB
/
Copy pathagent.py
File metadata and controls
159 lines (127 loc) · 4.75 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
"""
AgentCore Runtime 入口文件
这是部署到 AWS AgentCore 的主入口文件。
"""
import os
import logging
from typing import Dict, Any
from bedrock_agentcore.runtime import BedrockAgentCoreApp
from nl2sql_agent import NL2SQLAgent, DatabaseConfig, setup_logging
# ============================================================================
# 日志配置
# ============================================================================
# 检测是否在 AgentCore 环境中运行
is_agentcore = os.getenv("AWS_EXECUTION_ENV", "").startswith("AWS_Lambda") or \
os.getenv("AGENTCORE_RUNTIME", "") == "true"
# 配置日志:AgentCore 环境使用 JSON 格式,本地环境使用标准格式
log_level = os.getenv("LOG_LEVEL", "INFO")
setup_logging(level=log_level, format_json=is_agentcore)
logger = logging.getLogger(__name__)
# ============================================================================
# 初始化 AgentCore 应用
# ============================================================================
app = BedrockAgentCoreApp()
# ============================================================================
# 初始化 NL2SQL Agent
# ============================================================================
# 从环境变量读取数据库配置
db_config = DatabaseConfig(
host=os.getenv("DB_ENDPOINT", "localhost"),
port=int(os.getenv("DB_PORT", "3306")),
database=os.getenv("DB_NAME", "demodb"),
user=os.getenv("DB_USER", "root"),
password=os.getenv("DB_PASSWORD", ""),
region=os.getenv("AWS_REGION", "us-west-2")
)
# 初始化 Agent
agent = NL2SQLAgent(
db_config=db_config,
model_id=os.getenv("MODEL_ID", "us.anthropic.claude-sonnet-4-20250514-v1:0"),
memory_id=os.getenv("MEMORY_ID"),
region=os.getenv("AWS_REGION", "us-west-2"),
max_retries=int(os.getenv("MAX_RETRIES", "3"))
)
logger.info(
f"NL2SQL Agent 已初始化 - "
f"数据库: {db_config.host}:{db_config.port}/{db_config.database}, "
f"Memory: {os.getenv('MEMORY_ID', '未配置')}"
)
# ============================================================================
# AgentCore 入口函数
# ============================================================================
@app.entrypoint
def invoke(payload: Dict[str, Any], context: Any) -> str:
"""
AgentCore Runtime 入口函数
参数:
payload: 包含用户查询的字典,格式:
{
"prompt": "用户的自然语言查询"
}
context: AgentCore 上下文对象,包含 session_id 等信息
返回:
分析结果字符串
"""
# 获取会话 ID
session_id = getattr(context, 'session_id', 'default')
# 获取用户查询
user_query = payload.get("prompt", "")
if not user_query:
return "错误: 请提供查询内容"
logger.info(f"收到查询请求 - Session: {session_id}, Query: {user_query[:100]}...")
try:
# 处理查询
response = agent.process_query(
user_query=user_query,
session_id=session_id
)
# 返回分析结果
if response.success:
logger.info(
f"查询成功 - Session: {session_id}, "
f"结果数: {response.result_count}, "
f"重试次数: {response.retry_count}"
)
return response.analysis
else:
logger.error(
f"查询失败 - Session: {session_id}, "
f"错误: {response.error_message}"
)
return f"查询失败: {response.error_message}"
except Exception as e:
logger.error(
f"处理查询时发生异常 - Session: {session_id}, 错误: {str(e)}",
exc_info=True
)
return f"系统错误: {str(e)}"
# ============================================================================
# 本地测试模式
# ============================================================================
if __name__ == "__main__":
"""
本地测试模式
运行此文件可以在本地环境测试 Agent 功能。
"""
print("=" * 60)
print("NL2SQL Agent - 本地测试模式")
print("=" * 60)
# 模拟 AgentCore 上下文
class MockContext:
def __init__(self):
self.session_id = "local-test-session"
# 测试查询
test_query = "查询所有客户的数量"
print(f"\n测试查询: {test_query}\n")
# 构建 payload
payload = {
"prompt": test_query
}
# 调用入口函数
try:
result = invoke(payload, MockContext())
print("结果:")
print(result)
except Exception as e:
print(f"错误: {str(e)}")
print("\n" + "=" * 60)