This commit is contained in:
2026-02-26 18:06:17 +08:00
parent 68200cdfe6
commit ad0082d811
33 changed files with 356 additions and 249 deletions
+12 -6
View File
@@ -35,8 +35,7 @@ more_dots/
├── config/ # 配置层 ├── config/ # 配置层
│ └── settings.py # 配置读取 │ └── settings.py # 配置读取
│ └── prompts.yaml # 提示词配置 │ └── prompts.yaml # 提示词配置
│ └── ragflow_templates/ # RAGFlow 模板(表名 -> 模板列表) │ └── table_retrieval_prompts/ # 表名检索提示词(表名 -> 模板列表)
│ └── table_metadata_prompts/ # 表模型元数据提示词
├── tools/ # 工具模块 ├── tools/ # 工具模块
│ ├── calculator.py # 计算器工具 │ ├── calculator.py # 计算器工具
│ └── web_search.py # 网络搜索工具(占位符) │ └── web_search.py # 网络搜索工具(占位符)
@@ -148,14 +147,21 @@ heartbeat_interval = 5
### RAGFlow 模板同步 ### RAGFlow 模板同步
模板文件位于 `config/ragflow_templates`,每个 JSON 对应一个表名与模板列表。 模板文件位于 `config/table_retrieval_prompts/tables.json`,单文件包含多个表名与模板列表。
同步脚本: 同步脚本:
```bash ```bash
python scripts/sync_ragflow_templates.py python scripts/sync_ragflow_templates.py
``` ```
请在 `config/config.ini` 中配置 `ragflow.upload` 上传接口。 请在 `config/config.ini` 中配置 `ragflow.upload` 上传接口,并分别设置:
`table_retrieval_dataset_id` 与 `sql_gen_dataset_id`。
默认使用覆盖更新模式(`ragflow.upload_mode = overwrite`)。
热更新接口:
- `POST /api/ragflow/table-retrieval/reload`
- `POST /api/ragflow/sql-gen/reload`
# 使用对话工作流 # 使用对话工作流
result = manager.execute_workflow( result = manager.execute_workflow(
@@ -191,7 +197,7 @@ class CustomTool(BaseTool):
创建新的代理类型: 创建新的代理类型:
```python ```python
from agents.base_agent import BaseAgent from agent.graph import BaseAgent
class CustomAgent(BaseAgent): class CustomAgent(BaseAgent):
def _build_graph(self): def _build_graph(self):
@@ -215,7 +221,7 @@ class CustomAgent(BaseAgent):
### 添加新功能 ### 添加新功能
1. **新工具**: 在 `tools/` 目录下创建新的工具类 1. **新工具**: 在 `tools/` 目录下创建新的工具类
2. **新代理**: 在 `agents/` 目录下继承 `BaseAgent` 类 2. **新代理**: 在 `agent/` 目录下继承 `BaseAgent` 类
3. **新工作流**: 在 `workflows/` 目录下扩展工作流管理器 3. **新工作流**: 在 `workflows/` 目录下扩展工作流管理器
### 测试 ### 测试
View File
+8 -1
View File
@@ -19,11 +19,13 @@ class ConversationAgent(BaseAgent):
workflow.add_node("analyze_intent", self._analyze_intent) workflow.add_node("analyze_intent", self._analyze_intent)
workflow.add_node("normalize_input", self._normalize_input) workflow.add_node("normalize_input", self._normalize_input)
workflow.add_node("generate_sql", self._generate_sql)
workflow.add_node("generate_response", self._generate_response) workflow.add_node("generate_response", self._generate_response)
workflow.add_node("update_context", self._update_context) workflow.add_node("update_context", self._update_context)
workflow.add_edge("analyze_intent", "normalize_input") workflow.add_edge("analyze_intent", "normalize_input")
workflow.add_edge("normalize_input", "generate_response") workflow.add_edge("normalize_input", "generate_sql")
workflow.add_edge("generate_sql", "generate_response")
workflow.add_edge("generate_response", "update_context") workflow.add_edge("generate_response", "update_context")
workflow.add_edge("update_context", END) workflow.add_edge("update_context", END)
@@ -61,6 +63,11 @@ class ConversationAgent(BaseAgent):
state.current_step = "response_generated" state.current_step = "response_generated"
return state return state
def _generate_sql(self, state: AgentState) -> AgentState:
"""生成 SQL"""
from . import nodes
return nodes.generate_sql(state, self.model)
def _update_context(self, state: AgentState) -> AgentState: def _update_context(self, state: AgentState) -> AgentState:
"""更新对话上下文与历史""" """更新对话上下文与历史"""
for message in state.messages: for message in state.messages:
+7 -1
View File
@@ -20,10 +20,12 @@ class BaseAgent:
workflow.add_node("process_input", nodes.process_input) workflow.add_node("process_input", nodes.process_input)
workflow.add_node("normalize_input", self._normalize_input) workflow.add_node("normalize_input", self._normalize_input)
workflow.add_node("generate_sql", self._generate_sql)
workflow.add_node("generate_response", self._generate_response) workflow.add_node("generate_response", self._generate_response)
workflow.add_edge("process_input", "normalize_input") workflow.add_edge("process_input", "normalize_input")
workflow.add_edge("normalize_input", "generate_response") workflow.add_edge("normalize_input", "generate_sql")
workflow.add_edge("generate_sql", "generate_response")
workflow.add_edge("generate_response", END) workflow.add_edge("generate_response", END)
workflow.set_entry_point("process_input") workflow.set_entry_point("process_input")
@@ -38,6 +40,10 @@ class BaseAgent:
"""规范化用户输入""" """规范化用户输入"""
return nodes.normalize_input(state, self.model) return nodes.normalize_input(state, self.model)
def _generate_sql(self, state: AgentState) -> AgentState:
"""生成 SQL"""
return nodes.generate_sql(state, self.model)
def run(self, user_input: str, **kwargs) -> Dict[str, Any]: def run(self, user_input: str, **kwargs) -> Dict[str, Any]:
"""运行代理并处理用户输入""" """运行代理并处理用户输入"""
initial_state = AgentState( initial_state = AgentState(
+44 -5
View File
@@ -1,7 +1,11 @@
from langchain_core.messages import BaseMessage, HumanMessage, SystemMessage import json
from langchain_core.messages import BaseMessage, HumanMessage, SystemMessage, AIMessage
from .state import AgentState from .state import AgentState
from services.prompt_manager import PromptManager from services.prompt_manager import get_prompt_manager
from services.template_matcher import TemplateMatcher from services.template_matcher import get_template_matcher
from services.sql_prompt_manager import SqlPromptManager
from tools.sr_api_tool import SrApiQueryTool
def process_input(state: AgentState) -> AgentState: def process_input(state: AgentState) -> AgentState:
@@ -12,6 +16,14 @@ def process_input(state: AgentState) -> AgentState:
def generate_response(state: AgentState, model) -> AgentState: def generate_response(state: AgentState, model) -> AgentState:
"""使用 LLM 生成回复""" """使用 LLM 生成回复"""
sr_api_result = state.context.get("sr_api_result")
if sr_api_result:
state.messages.append(AIMessage(content=str(sr_api_result)))
return state
final_sql = state.context.get("final_sql")
if final_sql:
state.messages.append(AIMessage(content=final_sql))
return state
if state.messages: if state.messages:
response = model.invoke(state.messages) response = model.invoke(state.messages)
state.messages.append(response) state.messages.append(response)
@@ -27,7 +39,7 @@ def normalize_input(state: AgentState, model) -> AgentState:
if not isinstance(last_message, HumanMessage): if not isinstance(last_message, HumanMessage):
return state return state
prompt_manager = PromptManager() prompt_manager = get_prompt_manager()
system_prompt = SystemMessage( system_prompt = SystemMessage(
content=prompt_manager.get("system", "english_normalizer") content=prompt_manager.get("system", "english_normalizer")
) )
@@ -38,6 +50,33 @@ def normalize_input(state: AgentState, model) -> AgentState:
state.context["original_input"] = last_message.content state.context["original_input"] = last_message.content
state.context["normalized_input"] = normalized state.context["normalized_input"] = normalized
matcher = TemplateMatcher() matcher = get_template_matcher()
state.context["table_match"] = matcher.match(normalized) state.context["table_match"] = matcher.match(normalized)
return state return state
def generate_sql(state: AgentState, model) -> AgentState:
"""根据表名与提示词生成 SQL"""
table_match = state.context.get("table_match") or {}
table_name = table_match.get("table_name")
normalized = state.context.get("normalized_input")
if not table_name or not normalized:
return state
prompt_manager = SqlPromptManager()
prompt_data = prompt_manager.get_prompt(table_name)
if not prompt_data:
return state
prompt_text = json.dumps(prompt_data, ensure_ascii=False, indent=2)
system_template = get_prompt_manager().get("system", "sql_mysql_select_only")
system_content = system_template.format(table_prompt_json=prompt_text)
user_content = f"User question (normalized English): {normalized}"
response = model.invoke([SystemMessage(content=system_content), HumanMessage(content=user_content)])
sql_text = response.content if hasattr(response, "content") else str(response)
state.context["final_sql"] = sql_text
tool = SrApiQueryTool()
state.context["sr_api_result"] = tool.run(json.dumps({"sql": sql_text}, ensure_ascii=False))
return state
+8 -1
View File
@@ -28,10 +28,12 @@ class ToolAgent(BaseAgent):
workflow = StateGraph(AgentState) workflow = StateGraph(AgentState)
workflow.add_node("normalize_input", self._normalize_input) workflow.add_node("normalize_input", self._normalize_input)
workflow.add_node("generate_sql", self._generate_sql)
workflow.add_node("agent", self._agent_node) workflow.add_node("agent", self._agent_node)
workflow.add_node("tools", self.tool_node) workflow.add_node("tools", self.tool_node)
workflow.add_edge("normalize_input", "agent") workflow.add_edge("normalize_input", "generate_sql")
workflow.add_edge("generate_sql", "agent")
workflow.add_edge("tools", "agent") workflow.add_edge("tools", "agent")
workflow.add_conditional_edges( workflow.add_conditional_edges(
@@ -65,6 +67,11 @@ class ToolAgent(BaseAgent):
return state return state
def _generate_sql(self, state: AgentState) -> AgentState:
"""生成 SQL"""
from . import nodes
return nodes.generate_sql(state, self.model)
def _should_use_tools(self, state: AgentState) -> str: def _should_use_tools(self, state: AgentState) -> str:
"""判断是否需要使用工具""" """判断是否需要使用工具"""
last_message = state.messages[-1] last_message = state.messages[-1]
-6
View File
@@ -1,6 +0,0 @@
"""兼容导出:请优先使用 agent 包"""
from agent.graph import BaseAgent
from agent.state import AgentState
__all__ = ["BaseAgent", "AgentState"]
-53
View File
@@ -1,53 +0,0 @@
from typing import Dict, Any, List, Optional
from langchain_core.messages import BaseMessage, HumanMessage, AIMessage
from langgraph.graph import StateGraph, END
from .base_agent import BaseAgent, AgentState
class ConversationAgent(BaseAgent):
"""处理多轮对话的代理"""
def __init__(self, model_section: Optional[str] = None):
super().__init__(model_section)
self.conversation_history: List[BaseMessage] = []
def _build_graph(self) -> StateGraph:
"""构建对话专用图"""
workflow = StateGraph(AgentState)
# 添加节点
workflow.add_node("analyze_intent", self._analyze_intent)
workflow.add_node("generate_response", self._generate_response)
workflow.add_node("update_context", self._update_context)
# 定义边
workflow.add_edge("analyze_intent", "generate_response")
workflow.add_edge("generate_response", "update_context")
workflow.add_edge("update_context", END)
# 设置入口节点
workflow.set_entry_point("analyze_intent")
return workflow.compile()
def _analyze_intent(self, state: AgentState) -> AgentState:
"""分析用户意图与对话上下文"""
# 简单意图分析,可用更复杂逻辑增强
user_message = state.messages[-1] if state.messages else None
if user_message and isinstance(user_message, HumanMessage):
content = user_message.content.lower()
# 基础意图识别
if any(word in content for word in ["hello", "hi", "hey", "greetings"]):
state.context["intent"] = "greeting"
elif any(word in content for word in ["help", "assist", "support"]):
state.context["intent"] = "help"
elif "?" in content:
state.context["intent"] = "question"
else:
"""兼容导出:请优先使用 agent 包"""
from agent.conversation import ConversationAgent
__all__ = ["ConversationAgent"]
-5
View File
@@ -1,5 +0,0 @@
"""兼容导出:请优先使用 agent 包"""
from agent.tool import ToolAgent
__all__ = ["ToolAgent"]
+4
View File
@@ -15,3 +15,7 @@ def get_service_config(request: Request):
def get_tool_router(request: Request): def get_tool_router(request: Request):
return request.app.state.tool_router return request.app.state.tool_router
def get_prompt_manager(request: Request):
return request.app.state.prompt_manager
+45 -4
View File
@@ -6,7 +6,8 @@ from schemas.agent_output import AgentOutput
from schemas.tool_input import ToolInput from schemas.tool_input import ToolInput
from schemas.tool_output import ToolOutput from schemas.tool_output import ToolOutput
from workflows.workflow_manager import WorkflowType from workflows.workflow_manager import WorkflowType
from api.dependencies import get_workflow_manager, get_nacos_manager, get_service_config, get_tool_router from api.dependencies import get_workflow_manager, get_nacos_manager, get_service_config, get_tool_router, get_prompt_manager
from services.ragflow_sync import RagflowSync
router = APIRouter() router = APIRouter()
@@ -63,12 +64,32 @@ def run_workflow_stream(payload: AgentInput, workflow_manager=Depends(get_workfl
if workflow_type != WorkflowType.CONVERSATION: if workflow_type != WorkflowType.CONVERSATION:
raise HTTPException(status_code=400, detail="仅支持对话工作流的流式输出") raise HTTPException(status_code=400, detail="仅支持对话工作流的流式输出")
agent = workflow_manager.get_workflow(workflow_type) def _extract_output_text(result: dict) -> str:
context = (result.get("context") or {}) if isinstance(result, dict) else {}
if "sr_api_result" in context:
return str(context.get("sr_api_result") or "")
messages = result.get("messages") if isinstance(result, dict) else None
if messages:
last = messages[-1]
if hasattr(last, "content"):
return str(last.content or "")
return ""
def event_stream(): def event_stream():
try: try:
for token in agent.stream_run(payload.input): result = workflow_manager.execute_workflow(
yield f"data: {token}\n\n" workflow_type=workflow_type,
user_input=payload.input,
session_id=payload.session_id,
)
text = _extract_output_text(result.get("result") or {})
if not text:
yield "event: end\ndata: [DONE]\n\n"
return
chunk_size = 512
for i in range(0, len(text), chunk_size):
chunk = text[i : i + chunk_size]
yield f"data: {chunk}\n\n"
yield "event: end\ndata: [DONE]\n\n" yield "event: end\ndata: [DONE]\n\n"
except Exception as e: except Exception as e:
yield f"event: error\ndata: {str(e)}\n\n" yield f"event: error\ndata: {str(e)}\n\n"
@@ -80,3 +101,23 @@ def run_workflow_stream(payload: AgentInput, workflow_manager=Depends(get_workfl
def run_tool(payload: ToolInput, tool_router=Depends(get_tool_router)): def run_tool(payload: ToolInput, tool_router=Depends(get_tool_router)):
result = tool_router.call(payload.tool_name, payload.payload) result = tool_router.call(payload.tool_name, payload.payload)
return ToolOutput(**result) return ToolOutput(**result)
@router.post("/api/prompts/reload")
def reload_prompts(prompt_manager=Depends(get_prompt_manager)):
prompt_manager.reload()
return {"ok": True}
@router.post("/api/ragflow/table-retrieval/reload")
def reload_table_retrieval():
syncer = RagflowSync()
result = syncer.sync_table_retrieval()
return {"ok": True, "result": result}
@router.post("/api/ragflow/sql-gen/reload")
def reload_sql_gen():
syncer = RagflowSync()
result = syncer.sync_sql_gen_prompts()
return {"ok": True, "result": result}
+4 -8
View File
@@ -32,17 +32,13 @@ model_section = gpt-4o
[ragflow] [ragflow]
url = http://10.122.176.97:21020 url = http://10.122.176.97:21020
dataset_ids = f5b8b854d63a11f083230242c0a8e006
document_ids = 819cf100f52611f0a7fa0242c0a8e006
api_key = ragflow-xxxxx api_key = ragflow-xxxxx
retrieval = /api/v1/retrieval retrieval = /api/v1/retrieval
upload = /api/v1/documents upload = /api/v1/documents
cache_ttl = 600 # 上传模式:overwrite(覆盖更新)或 append(追加)
upload_mode = overwrite
[redis] table_retrieval_dataset_id =
enabled = false sql_gen_dataset_id =
url = redis://localhost:6379/0
db = 0
[nacos] [nacos]
# 是否启用 Nacos 注册 # 是否启用 Nacos 注册
+7
View File
@@ -4,6 +4,13 @@ system:
sql_generator: | sql_generator: |
You are an assistant that converts user intent into SQL. You are an assistant that converts user intent into SQL.
Ensure the SQL is correct, safe, and syntactically valid. Ensure the SQL is correct, safe, and syntactically valid.
sql_mysql_select_only: |
You are an expert SQL generator.
Only output a single MySQL SELECT statement.
Do not use LIMIT.
Do not output any other text.
Use the following table prompt JSON to generate SQL:
{table_prompt_json}
user: user:
default: | default: |
-4
View File
@@ -1,4 +0,0 @@
本目录存放用于 RAGFlow 匹配的模板配置(JSON)。
约定:每个 JSON 文件对应一个表名(key),包含模板列表(value)。
示例文件:orders.json、customers.json。
@@ -1,8 +0,0 @@
{
"table": "example_table",
"templates": [
"example_table created in {date}",
"count of records in example_table",
"example_table where status = {status}"
]
}
-4
View File
@@ -1,4 +0,0 @@
本目录存放业务表模型元数据提示词(JSON 格式)。
约定:每个 JSON 文件对应一个数据库表模型,描述表与字段元数据。
示例文件:order_metadata.json、customer_metadata.json 等。
@@ -1,20 +0,0 @@
{
"table": "example_table",
"description": "示例表模型元数据提示词",
"fields": [
{
"name": "id",
"type": "int",
"description": "主键"
},
{
"name": "created_at",
"type": "datetime",
"description": "创建时间"
}
],
"relationships": [],
"notes": [
"字段含义与业务规则可在此补充"
]
}
+4
View File
@@ -0,0 +1,4 @@
本目录存放用于 RAGFlow 匹配的表名检索提示词(JSON)。
约定:使用单一 JSON 文件维护多个表名及其模板列表。
示例文件:tables.json。
@@ -0,0 +1,23 @@
{
"tables": {
"apbo_eta_ful": [
"SO ETA INFO",
"BO list",
"CC",
"Commodity Code",
"BO QTY",
"open order",
"BO status",
"BO report",
"model",
"MTM SN",
"backlog",
"Premier",
"SN machine_sn",
"work order information",
"recovery ETA",
"history order",
"Warranty type"
]
}
}
-1
View File
@@ -8,4 +8,3 @@ uvicorn>=0.30.0
nacos-sdk-python>=2.0.9 nacos-sdk-python>=2.0.9
httpx>=0.27.0 httpx>=0.27.0
pyyaml>=6.0.1 pyyaml>=6.0.1
redis>=5.0.0
+4 -58
View File
@@ -1,62 +1,8 @@
import json
import os
from typing import Dict, List
import httpx
from config import Config
def load_templates(dir_path: str) -> List[Dict[str, any]]:
items = []
for name in os.listdir(dir_path):
if not name.endswith(".json"):
continue
with open(os.path.join(dir_path, name), "r", encoding="utf-8") as f:
items.append(json.load(f))
return items
def build_document(item: Dict[str, any]) -> str:
table = item.get("table", "")
templates = item.get("templates", [])
lines = [f"table: {table}"]
for t in templates:
lines.append(f"- {t}")
return "\n".join(lines)
def main(): def main():
cfg = Config.get_section("ragflow") from services.ragflow_sync import RagflowSync
base_url = cfg.get("url", "").rstrip("/") syncer = RagflowSync()
api_key = cfg.get("api_key", "") syncer.sync_table_retrieval()
dataset_ids = cfg.get("dataset_ids", "") print("表名检索模板同步完成")
upload_path = cfg.get("upload", "")
if not upload_path:
raise RuntimeError("未配置 ragflow.upload 上传接口,请在 config/config.ini 中设置")
url = base_url + "/" + upload_path.lstrip("/")
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
root = os.path.dirname(os.path.dirname(__file__))
templates_dir = os.path.join(root, "config", "ragflow_templates")
items = load_templates(templates_dir)
payload = []
for item in items:
payload.append(
{
"dataset_ids": dataset_ids,
"content": build_document(item),
"metadata": {"table": item.get("table")},
}
)
with httpx.Client(timeout=60) as client:
response = client.post(url, json={"documents": payload}, headers=headers)
response.raise_for_status()
print("同步完成")
if __name__ == "__main__": if __name__ == "__main__":
+11
View File
@@ -0,0 +1,11 @@
from services.ragflow_sync import RagflowSync
def main():
syncer = RagflowSync()
syncer.sync_sql_gen_prompts()
print("SQL 生成提示词同步完成")
if __name__ == "__main__":
main()
+3
View File
@@ -8,6 +8,7 @@ from config import Config
from workflows.workflow_manager import WorkflowManager from workflows.workflow_manager import WorkflowManager
from services.nacos_service import load_nacos_config, load_service_config, NacosManager from services.nacos_service import load_nacos_config, load_service_config, NacosManager
from services.tool_router import ToolRouter from services.tool_router import ToolRouter
from services.prompt_manager import get_prompt_manager
from api import endpoints from api import endpoints
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -22,6 +23,7 @@ def create_app() -> FastAPI:
workflow_manager = WorkflowManager(default_model_section=default_model_section) workflow_manager = WorkflowManager(default_model_section=default_model_section)
nacos_manager = NacosManager(nacos_config=nacos_config, service_config=service_config) nacos_manager = NacosManager(nacos_config=nacos_config, service_config=service_config)
tool_router = ToolRouter() tool_router = ToolRouter()
prompt_manager = get_prompt_manager()
@asynccontextmanager @asynccontextmanager
async def lifespan(app: FastAPI): async def lifespan(app: FastAPI):
@@ -31,6 +33,7 @@ def create_app() -> FastAPI:
app.state.nacos_manager = nacos_manager app.state.nacos_manager = nacos_manager
app.state.service_config = service_config app.state.service_config = service_config
app.state.tool_router = tool_router app.state.tool_router = tool_router
app.state.prompt_manager = prompt_manager
await nacos_manager.start() await nacos_manager.start()
logger.info("✅ 服务准备就绪: %s on %s:%s", service_config.service_name, service_config.ip, service_config.port) logger.info("✅ 服务准备就绪: %s on %s:%s", service_config.service_name, service_config.ip, service_config.port)
-18
View File
@@ -2,11 +2,6 @@ from __future__ import annotations
from typing import Optional from typing import Optional
try:
import redis
except Exception:
redis = None
class CacheBase: class CacheBase:
"""缓存接口""" """缓存接口"""
@@ -28,16 +23,3 @@ class NoopCache(CacheBase):
return None return None
class RedisCache(CacheBase):
"""Redis 缓存"""
def __init__(self, url: str, db: int = 0):
if redis is None:
raise ImportError("未安装 redis 依赖")
self._client = redis.Redis.from_url(url, db=db, decode_responses=True)
def get(self, key: str) -> Optional[str]:
return self._client.get(key)
def set(self, key: str, value: str, ttl: int) -> None:
self._client.set(key, value, ex=ttl)
+1 -6
View File
@@ -3,12 +3,7 @@ import logging
import socket import socket
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Dict, Optional from typing import Any, Dict, Optional
import nacos
try:
import nacos
except Exception:
nacos = None
from config import Config from config import Config
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
+11
View File
@@ -29,3 +29,14 @@ class PromptManager:
def list_prompts(self, group: str) -> list[str]: def list_prompts(self, group: str) -> list[str]:
"""列出分组内提示词""" """列出分组内提示词"""
return list(self._data.get(group, {}).keys()) return list(self._data.get(group, {}).keys())
_GLOBAL_PROMPT_MANAGER: Optional[PromptManager] = None
def get_prompt_manager(config_path: Optional[str] = None) -> PromptManager:
"""获取全局 PromptManager(单例)"""
global _GLOBAL_PROMPT_MANAGER
if _GLOBAL_PROMPT_MANAGER is None:
_GLOBAL_PROMPT_MANAGER = PromptManager(config_path=config_path)
return _GLOBAL_PROMPT_MANAGER
+8 -3
View File
@@ -14,22 +14,27 @@ class RagflowClient:
self._base_url = cfg.get("url", "") self._base_url = cfg.get("url", "")
self._api_key = cfg.get("api_key", "") self._api_key = cfg.get("api_key", "")
self._retrieval_path = cfg.get("retrieval", "/api/v1/retrieval") self._retrieval_path = cfg.get("retrieval", "/api/v1/retrieval")
self._dataset_ids = cfg.get("dataset_ids", "") self._table_retrieval_dataset_id = cfg.get("table_retrieval_dataset_id", "")
self._sql_gen_dataset_id = cfg.get("sql_gen_dataset_id", "")
def _build_url(self) -> str: def _build_url(self) -> str:
return self._base_url.rstrip("/") + "/" + self._retrieval_path.lstrip("/") return self._base_url.rstrip("/") + "/" + self._retrieval_path.lstrip("/")
def retrieve(self, query: str, top_k: int = 3) -> Dict[str, Any]: def retrieve(self, query: str, top_k: int = 3, dataset_id: Optional[str] = None, document_ids: Optional[str] = None) -> Dict[str, Any]:
"""检索匹配文档""" """检索匹配文档"""
if not self._base_url or not self._retrieval_path: if not self._base_url or not self._retrieval_path:
raise RuntimeError("未配置 ragflow.url 或 ragflow.retrieval") raise RuntimeError("未配置 ragflow.url 或 ragflow.retrieval")
if not dataset_id:
raise ValueError("未提供 ragflow.dataset_id,无法进行检索")
url = self._build_url() url = self._build_url()
headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {} headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {}
payload = { payload = {
"dataset_ids": self._dataset_ids, "dataset_ids": dataset_id or "",
"query": query, "query": query,
"top_k": top_k, "top_k": top_k,
} }
if document_ids:
payload["document_ids"] = document_ids
with httpx.Client(timeout=30) as client: with httpx.Client(timeout=30) as client:
response = client.post(url, json=payload, headers=headers) response = client.post(url, json=payload, headers=headers)
+106
View File
@@ -0,0 +1,106 @@
import json
import os
from typing import Dict, List, Optional
import httpx
from config import Config
def _build_document_for_table(table: str, templates: List[str]) -> str:
lines = [f"table: {table}"]
for t in templates:
lines.append(f"- {t}")
return "\n".join(lines)
def _build_sql_gen_document(table: str, prompt: Dict[str, any]) -> str:
system_prompt = prompt.get("system_prompt", "")
business_prompt = prompt.get("business_prompt", "")
constraints = prompt.get("constraints", [])
lines = [f"table: {table}"]
if system_prompt:
lines.append("[system] " + system_prompt)
if business_prompt:
lines.append("[business] " + business_prompt)
if constraints:
lines.append("[constraints]")
for c in constraints:
lines.append(f"- {c}")
return "\n".join(lines)
class RagflowSync:
"""RAGFlow 同步工具"""
def __init__(self):
cfg = Config.get_section("ragflow")
self._base_url = cfg.get("url", "").rstrip("/")
self._api_key = cfg.get("api_key", "")
self._table_retrieval_dataset_id = (cfg.get("table_retrieval_dataset_id") or "").strip()
self._sql_gen_dataset_id = (cfg.get("sql_gen_dataset_id") or "").strip()
self._upload_path = (cfg.get("upload") or "").strip()
self._upload_mode = (cfg.get("upload_mode") or "overwrite").strip().lower()
def _validate_common(self) -> None:
if not self._base_url:
raise RuntimeError("未配置 ragflow.url")
if not self._upload_path:
raise RuntimeError("未配置 ragflow.upload 上传接口,请在 config/config.ini 中设置")
if self._upload_mode not in ("overwrite", "append"):
raise RuntimeError("ragflow.upload_mode 仅支持 overwrite 或 append")
def _post(self, documents: List[Dict[str, any]]):
self._validate_common()
url = self._base_url + "/" + self._upload_path.lstrip("/")
headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {}
payload: Dict[str, any] = {"documents": documents}
if self._upload_mode == "overwrite":
payload["mode"] = "overwrite"
with httpx.Client(timeout=60) as client:
response = client.post(url, json=payload, headers=headers)
response.raise_for_status()
return response.json()
def sync_table_retrieval(self) -> Dict[str, any]:
if not self._table_retrieval_dataset_id:
raise RuntimeError("未配置 ragflow.table_retrieval_dataset_id,无法同步表名检索模板")
root = os.path.dirname(os.path.dirname(__file__))
tables_file = os.path.join(root, "config", "table_retrieval_prompts", "tables.json")
with open(tables_file, "r", encoding="utf-8") as f:
data = json.load(f)
tables = data.get("tables", {})
documents = []
for table, templates in tables.items():
doc = {
"dataset_ids": self._table_retrieval_dataset_id,
"content": _build_document_for_table(table, templates),
"metadata": {"table": table},
}
documents.append(doc)
return self._post(documents)
def sync_sql_gen_prompts(self) -> Dict[str, any]:
if not self._sql_gen_dataset_id:
raise RuntimeError("未配置 ragflow.sql_gen_dataset_id,无法同步 SQL 生成提示词")
root = os.path.dirname(os.path.dirname(__file__))
prompts_dir = os.path.join(root, "config", "sql_gen_prompts")
documents = []
for name in os.listdir(prompts_dir):
if not name.endswith(".json"):
continue
path = os.path.join(prompts_dir, name)
with open(path, "r", encoding="utf-8") as f:
prompt = json.load(f)
table = prompt.get("table") or os.path.splitext(name)[0]
doc = {
"dataset_ids": self._sql_gen_dataset_id,
"content": _build_sql_gen_document(table, prompt),
"metadata": {"table": table},
}
documents.append(doc)
return self._post(documents)
+26
View File
@@ -0,0 +1,26 @@
import json
import os
from typing import Any, Dict, Optional
class SqlPromptManager:
"""按表名读取 SQL 提示词"""
def __init__(self, base_dir: Optional[str] = None):
root_dir = os.path.dirname(os.path.dirname(__file__))
self._base_dir = base_dir or os.path.join(root_dir, "config", "sql_gen_prompts")
@staticmethod
def _safe_filename(name: str) -> str:
return name.replace("..", "").replace("/", "_").replace("\\", "_")
def get_prompt(self, table_name: str) -> Optional[Dict[str, Any]]:
"""读取指定表的提示词 JSON"""
if not table_name:
return None
filename = self._safe_filename(table_name) + ".json"
path = os.path.join(self._base_dir, filename)
if not os.path.exists(path):
return None
with open(path, "r", encoding="utf-8") as f:
return json.load(f)
+20 -37
View File
@@ -1,53 +1,28 @@
import hashlib from typing import Any, Dict
import json
from typing import Any, Dict, Optional
from config import Config from config import Config
from services.cache import NoopCache, RedisCache
from services.ragflow_client import RagflowClient, extract_table_name from services.ragflow_client import RagflowClient, extract_table_name
class TemplateMatcher: class TemplateMatcher:
"""模板匹配器:RAGFlow + Redis 缓存""" """模板匹配器:RAGFlow 检索"""
def __init__(self): def __init__(self):
self._ragflow = RagflowClient() self._ragflow = RagflowClient()
self._cache = self._init_cache()
cfg = Config.get_section("ragflow") cfg = Config.get_section("ragflow")
self._cache_ttl = int(cfg.get("cache_ttl", 600)) self._dataset_id = (cfg.get("table_retrieval_dataset_id") or "").strip()
def _init_cache(self): def _validate(self) -> None:
cfg = Config.get_section("redis") if not self._dataset_id:
enabled = str(cfg.get("enabled", "false")).lower() in ("1", "true", "yes") raise RuntimeError("未配置 ragflow.table_retrieval_dataset_id,无法进行表名检索")
if not enabled:
return NoopCache()
url = cfg.get("url")
db = int(cfg.get("db", 0))
if not url:
return NoopCache()
try:
return RedisCache(url=url, db=db)
except Exception:
return NoopCache()
@staticmethod
def _cache_key(text: str) -> str:
return "ragflow:table:" + hashlib.sha256(text.encode("utf-8")).hexdigest()
def match(self, normalized_text: str) -> Dict[str, Any]: def match(self, normalized_text: str) -> Dict[str, Any]:
"""返回匹配的表名与原始响应""" """返回匹配的表名与原始响应"""
key = self._cache_key(normalized_text) self._validate()
cached = self._cache.get(key)
if cached:
return json.loads(cached)
try: try:
response = self._ragflow.retrieve(normalized_text, top_k=3) response = self._ragflow.retrieve(normalized_text, top_k=3, dataset_id=self._dataset_id)
except Exception as e: except Exception as e:
result = {"table_name": None, "raw": {"error": str(e)}} return {"table_name": None, "raw": {"error": str(e)}}
self._cache.set(key, json.dumps(result, ensure_ascii=False), self._cache_ttl)
return result
candidates = [] candidates = []
data = response.get("data") if isinstance(response, dict) else None data = response.get("data") if isinstance(response, dict) else None
@@ -58,7 +33,15 @@ class TemplateMatcher:
candidates.append(table_name) candidates.append(table_name)
matched = candidates[0] if candidates else None matched = candidates[0] if candidates else None
result = {"table_name": matched, "raw": response} return {"table_name": matched, "raw": response}
self._cache.set(key, json.dumps(result, ensure_ascii=False), self._cache_ttl)
return result _GLOBAL_TEMPLATE_MATCHER: TemplateMatcher | None = None
def get_template_matcher() -> TemplateMatcher:
"""获取全局 TemplateMatcher(单例)"""
global _GLOBAL_TEMPLATE_MATCHER
if _GLOBAL_TEMPLATE_MATCHER is None:
_GLOBAL_TEMPLATE_MATCHER = TemplateMatcher()
return _GLOBAL_TEMPLATE_MATCHER