diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..87ea4f3 --- /dev/null +++ b/.dockerignore @@ -0,0 +1,13 @@ +__pycache__/ +*.pyc +*.pyo +*.pyd +*.log +.pytest_cache/ +.mypy_cache/ +.git/ +.gitignore +.vscode/ +.idea/ +_trial_temp/ +*.ipynb diff --git a/.gitignore b/.gitignore index f7d85eb..025bede 100644 --- a/.gitignore +++ b/.gitignore @@ -16,5 +16,3 @@ __pycache__/ .DS_Store *.log -# 本地配置 -config/config.ini diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..2f1b57f --- /dev/null +++ b/Dockerfile @@ -0,0 +1,22 @@ +FROM python:3.12-bookworm + +ENV PYTHONDONTWRITEBYTECODE=1 \ + PYTHONUNBUFFERED=1 \ + PIP_NO_CACHE_DIR=1 \ + PIP_PROGRESS_BAR=off \ + PIP_NO_COLOR=1 \ + PIP_QUIET=1 + +WORKDIR /app + +RUN pip install --no-input --no-cache-dir --upgrade pip setuptools wheel + +COPY requirements.txt /app/requirements.txt +RUN pip install --no-input --no-cache-dir \ + -r /app/requirements.txt \ + -i https://pypi.tuna.tsinghua.edu.cn/simple/ + +COPY . /app + +EXPOSE 8000 +CMD ["python", "server.py"] \ No newline at end of file diff --git a/Jenkinsfile b/Jenkinsfile new file mode 100644 index 0000000..8bf316d --- /dev/null +++ b/Jenkinsfile @@ -0,0 +1,64 @@ +stage('Deploy to k3s') { + steps { + dir('more_dots') { + script { + def cluster = params.TARGET_CLUSTER + def fullImageWithTag = "${FULL_IMAGE_NAME}:${env.IMAGE_TAG}" + + if (cluster == 'cluster1' || cluster == 'both') { + withCredentials([ + file(credentialsId: 'k3s-cluster1-config', variable: 'KUBECONFIG_CLUSTER1'), + usernamePassword(credentialsId: REGISTRY_CREDENTIALS_ID, usernameVariable: 'REGISTRY_USER', passwordVariable: 'REGISTRY_PASS') + ]) { + sh """ + set -eux + + # 查找kubectl路径 + KUBECTL_PATH=\$(command -v kubectl 2>/dev/null || true) + if [ -z "\$KUBECTL_PATH" ]; then + for p in /usr/local/bin/kubectl /usr/bin/kubectl /bin/kubectl; do + if [ -x "\$p" ]; then + KUBECTL_PATH="\$p" + break + fi + done + fi + + echo "使用kubectl路径: \$KUBECTL_PATH" + + # 定义kubectl函数 + k() { + sudo \$KUBECTL_PATH --kubeconfig=${KUBECONFIG_CLUSTER1} "\$@" + } + + # 检查命名空间 + k get namespace ${params.DEPLOY_ENV} || k create namespace ${params.DEPLOY_ENV} + + # 创建imagePullSecret + k create secret docker-registry regcred-130 \\ + --docker-server=${REGISTRY_URL} \\ + --docker-username=${REGISTRY_USER} \\ + --docker-password=${REGISTRY_PASS} \\ + --namespace=${params.DEPLOY_ENV} \\ + --dry-run=client -o yaml | k apply -f - + + # 替换镜像并部署 + sed "s|image:.*more_dots.*|image: ${fullImageWithTag}|g" k8s/deployment.yaml > /tmp/deployment-${params.DEPLOY_ENV}.yaml + k apply -f /tmp/deployment-${params.DEPLOY_ENV}.yaml -n ${params.DEPLOY_ENV} + + # 确保使用imagePullSecret(注意 deployment 名称是 more-dots) + k patch deployment more-dots -n ${params.DEPLOY_ENV} \\ + -p '{"spec":{"template":{"spec":{"imagePullSecrets":[{"name":"regcred-130"}]}}}}' || true + + # 重启并等待 + k rollout restart deployment/more-dots -n ${params.DEPLOY_ENV} 2>/dev/null || true + k rollout status deployment/more-dots -n ${params.DEPLOY_ENV} --timeout=300s + """ + } + } + + // cluster2 部分同样修改... + } + } + } +} \ No newline at end of file diff --git a/agent/conversation.py b/agent/conversation.py index e1039ac..534e337 100644 --- a/agent/conversation.py +++ b/agent/conversation.py @@ -53,13 +53,9 @@ class ConversationAgent(BaseAgent): return state def _generate_response(self, state: AgentState) -> AgentState: - """结合对话历史生成回复""" - all_messages = self.conversation_history + state.messages - - if all_messages: - response = self.model.invoke(all_messages) - state.messages.append(response) - + """优先返回 SQL 执行结果,其次返回生成 SQL,再回退到模型回复""" + from . import nodes + state = nodes.generate_response(state, self.model) state.current_step = "response_generated" return state diff --git a/agent/nodes.py b/agent/nodes.py index 7006a83..e812fff 100644 --- a/agent/nodes.py +++ b/agent/nodes.py @@ -4,29 +4,40 @@ from langchain_core.messages import BaseMessage, HumanMessage, SystemMessage, AI from .state import AgentState from services.prompt_manager import get_prompt_manager from services.template_matcher import get_template_matcher -from services.sql_prompt_manager import SqlPromptManager +from services.sql_prompt_manager import get_sql_prompt_manager from tools.sr_api_tool import SrApiQueryTool +def _short(value, max_len: int = 500) -> str: + text = str(value) + return text if len(text) <= max_len else text[:max_len] + "..." + + def process_input(state: AgentState) -> AgentState: """处理用户输入""" + print("[process_input][in] messages=", _short(state.messages)) state.current_step = "processed" + print("[process_input][out] current_step=", state.current_step) return state def generate_response(state: AgentState, model) -> AgentState: """使用 LLM 生成回复""" + print("[generate_response][in] context_keys=", list((state.context or {}).keys())) sr_api_result = state.context.get("sr_api_result") if sr_api_result: state.messages.append(AIMessage(content=str(sr_api_result))) + print("[generate_response][out] source=sr_api_result") return state final_sql = state.context.get("final_sql") if final_sql: state.messages.append(AIMessage(content=final_sql)) + print("[generate_response][out] source=final_sql") return state if state.messages: response = model.invoke(state.messages) state.messages.append(response) + print("[generate_response][out] source=model_invoke") return state @@ -39,19 +50,25 @@ def normalize_input(state: AgentState, model) -> AgentState: if not isinstance(last_message, HumanMessage): return state + print("[normalize_input][in] user_input=", _short(last_message.content)) + prompt_manager = get_prompt_manager() - system_prompt = SystemMessage( - content=prompt_manager.get("system", "english_normalizer") + normalizer_prompt = ( + prompt_manager.get("system", "english_normalizer") + or prompt_manager.get("user", "english_normalizer") ) + system_prompt = SystemMessage(content=normalizer_prompt) response = model.invoke([system_prompt, HumanMessage(content=last_message.content)]) normalized = response.content if hasattr(response, "content") else str(response) + print("[normalize_input][out] normalized=", _short(normalized)) state.context["original_input"] = last_message.content state.context["normalized_input"] = normalized matcher = get_template_matcher() state.context["table_match"] = matcher.match(normalized) + print("[normalize_input][out] table_match=", _short(state.context.get("table_match"))) return state @@ -62,11 +79,16 @@ def generate_sql(state: AgentState, model) -> AgentState: normalized = state.context.get("normalized_input") if not table_name or not normalized: + print("[generate_sql][skip] missing table_name or normalized") return state - prompt_manager = SqlPromptManager() + print("[generate_sql][in] table_name=", table_name) + print("[generate_sql][in] normalized=", _short(normalized)) + + prompt_manager = get_sql_prompt_manager() prompt_data = prompt_manager.get_prompt(table_name) if not prompt_data: + print("[generate_sql][skip] prompt not found for table=", table_name) return state prompt_text = json.dumps(prompt_data, ensure_ascii=False, indent=2) @@ -75,8 +97,11 @@ def generate_sql(state: AgentState, model) -> AgentState: 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) + print("[generate_sql][out] sql=", _short(sql_text)) state.context["final_sql"] = sql_text - tool = SrApiQueryTool() - state.context["sr_api_result"] = tool.run(json.dumps({"sql": sql_text}, ensure_ascii=False)) + if not state.context.get("skip_sr_api"): + tool = SrApiQueryTool() + state.context["sr_api_result"] = tool.run(json.dumps({"sql": sql_text}, ensure_ascii=False)) + print("[generate_sql][out] sr_api_result=", _short(state.context.get("sr_api_result"))) return state diff --git a/api/endpoints.py b/api/endpoints.py index 2945229..ecf0c38 100644 --- a/api/endpoints.py +++ b/api/endpoints.py @@ -1,13 +1,23 @@ +import asyncio +import json +import time +import uuid + from fastapi import APIRouter, HTTPException, Depends from fastapi.responses import StreamingResponse +from config import Config from schemas.agent_input import AgentInput from schemas.agent_output import AgentOutput from schemas.tool_input import ToolInput from schemas.tool_output import ToolOutput +from schemas.chat_message_response import ChatMessageResponseDTO from workflows.workflow_manager import WorkflowType from api.dependencies import get_workflow_manager, get_nacos_manager, get_service_config, get_tool_router, get_prompt_manager +from services.app_errors import AppError, ErrorCode from services.ragflow_sync import RagflowSync +from services.structured_logger import get_structured_logger +from tools.sr_api_tool import SrApiQueryTool router = APIRouter() @@ -17,7 +27,17 @@ def _resolve_workflow_type(value: str) -> WorkflowType: try: return WorkflowType(value) except Exception as e: - raise ValueError(f"不支持的工作流类型: {value}") from e + raise AppError( + code=ErrorCode.INVALID_WORKFLOW_TYPE, + message=f"不支持的工作流类型: {value}", + status_code=400, + ) from e + + +def _to_http_error(e: Exception) -> HTTPException: + if isinstance(e, AppError): + return HTTPException(status_code=e.status_code, detail=e.to_dict()) + return HTTPException(status_code=500, detail={"code": ErrorCode.INTERNAL_ERROR.value, "message": str(e)}) @router.get("/health") @@ -36,16 +56,25 @@ def nacos_status(nacos_manager=Depends(get_nacos_manager)): @router.post("/api/workflows", response_model=AgentOutput) def run_workflow(payload: AgentInput, workflow_manager=Depends(get_workflow_manager)): + trace_id = uuid.uuid4().hex + slog = get_structured_logger() + slog.log("INFO", "run_workflow.start", trace_id, {"workflow_type": payload.workflow_type}) try: workflow_type = _resolve_workflow_type(payload.workflow_type) - except ValueError as e: - raise HTTPException(status_code=400, detail=str(e)) + except Exception as e: + slog.log("ERROR", "run_workflow.invalid_type", trace_id, error_code=ErrorCode.INVALID_WORKFLOW_TYPE.value, payload={"workflow_type": payload.workflow_type}) + raise _to_http_error(e) - result = workflow_manager.execute_workflow( - workflow_type=workflow_type, - user_input=payload.input, - session_id=payload.session_id, - ) + try: + result = workflow_manager.execute_workflow( + workflow_type=workflow_type, + user_input=payload.input, + session_id=payload.session_id, + ) + slog.log("INFO", "run_workflow.success", trace_id, {"session_id": result.get("session_id")}) + except Exception as e: + slog.log("ERROR", "run_workflow.failed", trace_id, error_code=ErrorCode.INTERNAL_ERROR.value, payload={"error": str(e)}) + raise _to_http_error(e) return AgentOutput( session_id=result["session_id"], @@ -54,45 +83,112 @@ def run_workflow(payload: AgentInput, workflow_manager=Depends(get_workflow_mana ) -@router.post("/api/workflows/stream") -def run_workflow_stream(payload: AgentInput, workflow_manager=Depends(get_workflow_manager)): +@router.post("/api/sql/generate") +def generate_sql(payload: AgentInput, workflow_manager=Depends(get_workflow_manager)): + """仅生成 SQL,不调用 SR API""" + trace_id = uuid.uuid4().hex + slog = get_structured_logger() try: workflow_type = _resolve_workflow_type(payload.workflow_type) - except ValueError as e: - raise HTTPException(status_code=400, detail=str(e)) + except Exception as e: + slog.log("ERROR", "generate_sql.invalid_type", trace_id, error_code=ErrorCode.INVALID_WORKFLOW_TYPE.value) + raise _to_http_error(e) + + result = workflow_manager.execute_workflow( + workflow_type=workflow_type, + user_input=payload.input, + session_id=payload.session_id, + skip_sr_api=True, + ) + + context = (result.get("result") or {}).get("context") or {} + sql_text = context.get("final_sql") + if not sql_text: + e = AppError(code=ErrorCode.SQL_GENERATION_FAILED, message="SQL 生成失败") + slog.log("ERROR", "generate_sql.failed", trace_id, error_code=e.code.value, payload={"context_keys": list(context.keys())}) + raise _to_http_error(e) + + slog.log("INFO", "generate_sql.success", trace_id, {"sql_len": len(sql_text)}) + + return { + "session_id": result.get("session_id"), + "workflow_type": result.get("workflow_type"), + "sql": sql_text, + } + + +@router.post("/api/workflows/stream") +def run_workflow_stream(payload: AgentInput, workflow_manager=Depends(get_workflow_manager)): + trace_id = uuid.uuid4().hex + slog = get_structured_logger() + try: + workflow_type = _resolve_workflow_type(payload.workflow_type) + except Exception as e: + raise _to_http_error(e) if workflow_type != WorkflowType.CONVERSATION: - raise HTTPException(status_code=400, detail="仅支持对话工作流的流式输出") + raise _to_http_error(AppError(code=ErrorCode.INVALID_WORKFLOW_TYPE, message="仅支持对话工作流的流式输出", status_code=400)) - 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 "" + stream_cfg = Config.get_section("stream") + progress_interval = float(stream_cfg.get("progress_interval", 0.3)) + task_id = uuid.uuid4().hex - def event_stream(): + def _build_message(conversation_id: str, answer: str) -> str: + dto = ChatMessageResponseDTO( + id=uuid.uuid4().hex, + event="message", + task_id=task_id, + message_id=uuid.uuid4().hex, + conversation_id=conversation_id, + answer=answer, + created_at=int(time.time()), + ) + return f"event: message\ndata: {json.dumps(dto.model_dump(), ensure_ascii=False)}\n\n" + + async def event_stream(): try: - result = workflow_manager.execute_workflow( - workflow_type=workflow_type, - user_input=payload.input, - session_id=payload.session_id, + slog.log("INFO", "stream.start", trace_id, {"workflow_type": payload.workflow_type}) + # 1) 先仅生成 SQL(不执行 SR API) + result = await asyncio.to_thread( + workflow_manager.execute_workflow, + workflow_type, + payload.input, + payload.session_id, + skip_sr_api=True, ) - text = _extract_output_text(result.get("result") or {}) - if not text: + conversation_id = str(result.get("session_id") or payload.session_id or task_id) + context = (result.get("result") or {}).get("context") or {} + sql_text = str(context.get("final_sql") or "") + + if not sql_text: + reason = "SQL 生成失败,可能是表未匹配或对应 SQL 提示词不存在" + slog.log("ERROR", "stream.sql_generation_failed", trace_id, error_code=ErrorCode.SQL_GENERATION_FAILED.value, payload={"conversation_id": conversation_id}) + yield _build_message(conversation_id, reason) 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" + + # 2) 先流式返回 SQL + yield _build_message(conversation_id, sql_text) + + # 3) 异步执行 SQL,并及时流式返回执行结果 + tool = SrApiQueryTool() + task = asyncio.create_task( + asyncio.to_thread(tool.run, json.dumps({"sql": sql_text}, ensure_ascii=False)) + ) + + while not task.done(): + yield _build_message(conversation_id, "executing_sql") + await asyncio.sleep(progress_interval) + + sql_result = await task + slog.log("INFO", "stream.sql_executed", trace_id, {"result_len": len(str(sql_result))}) + yield _build_message(conversation_id, str(sql_result)) yield "event: end\ndata: [DONE]\n\n" except Exception as e: - yield f"event: error\ndata: {str(e)}\n\n" + slog.log("ERROR", "stream.failed", trace_id, error_code=ErrorCode.INTERNAL_ERROR.value, payload={"error": str(e)}) + conversation_id = str(payload.session_id or task_id) + yield _build_message(conversation_id, str(e)) + yield "event: end\ndata: [DONE]\n\n" return StreamingResponse(event_stream(), media_type="text/event-stream") @@ -128,11 +224,11 @@ def upload_table_retrieval(): @router.put("/api/ragflow/table-retrieval/update") -def update_table_retrieval(config: dict): - """更新表名检索知识库配置""" +def update_table_retrieval(): + """更新表名检索文档(仅文档内容)""" syncer = RagflowSync() try: - result = syncer.update_dataset(syncer._table_retrieval_dataset_id, config) + result = syncer.update_table_retrieval_documents() return {"ok": True, "result": result} except Exception as e: raise HTTPException(status_code=500, detail=str(e)) @@ -150,11 +246,11 @@ def upload_sql_gen(): @router.put("/api/ragflow/sql-gen/update") -def update_sql_gen(config: dict): - """更新 SQL 生成知识库配置""" +def update_sql_gen(): + """更新 SQL 生成文档(仅文档内容)""" syncer = RagflowSync() try: - result = syncer.update_dataset(syncer._sql_gen_dataset_id, config) + result = syncer.update_sql_gen_documents() return {"ok": True, "result": result} except Exception as e: raise HTTPException(status_code=500, detail=str(e)) diff --git a/config/config.ini b/config/config.ini new file mode 100644 index 0000000..3541dc6 --- /dev/null +++ b/config/config.ini @@ -0,0 +1,73 @@ +[General] +# 默认使用的模型配置区域 +DEFAULT_MODEL_SECTION = qwen-80b +# 最大重试次数 +MAX_RETRIES = 3 +# 请求超时时间 +TIMEOUT = 30 + + +[qwen-32b] +# 模型名称 +MODEL_NAME = qwen-32b +# API Key +OPENAI_API_KEY = 123ffa18-3baf-4f39-9698-fdf680d15f36 +# URL +URL = http://led-gateway.lenovo.com:30089/intranet/qwen3-coder-30b/v1 + +[qwen-80b] +MODEL_NAME = Qwen3-Next-80B-A3B-Instruct +OPENAI_API_KEY = 123ffa18-3baf-4f39-9698-fdf680d15f36 +URL = http://led-gateway.lenovo.com:30089/intranet/qwen3-next-80b-a3b-instruct/v1 + +[sr_api] +url = http://ipc.lenovo.com/gateway/bgs-ai/Report/fetchData +llzAppkey = e0578e6a-f045-4f3e-93c0-5957c8b1915e +llzSercret = BFhItTCY5FqgtHAdn0eaxAGix/25oR8YfxYKfSPgrJb8zxYG06kBeoip0DEGjekfh43223atdBDXBwUmw18NXp/2piZamnUwlFWYFAGfpBCwy3K+921KI5ZaWRWcMCCcCKOYF/2feMg72owUhno2JXqSUcb8HBhang== + + +[ragflow] +url = http://led-ai-ragflow.lenovo.com +api_key = ragflow-gOOQ93FvWdyhP3i1J1xXFKVxWucgtF0bwbthuDJk5vs +retrieval = /api/v1/retrieval +retrieval_top_k = 3 +cache_ttl = 600 +table_retrieval_dataset_id = 9945baf512ea11f18ccb6a681b3130b2 +sql_gen_dataset_id = ee68f53a12ec11f18e436a681b3130b2 + +[redis] +enabled = true +host = led-redis.lenovo.com +port = 30398 +password = bgs123456 +database = 0 +sql_prompt_ttl = 600 + +[stream] +progress_interval = 0.3 + +[logging_mysql] +enabled = false +host = 127.0.0.1 +port = 3306 +user = root +password = +database = more_dots +table = structured_logs +connect_timeout = 5 + + +[app] +service_name = local-model-streaming-api +host = 0.0.0.0 +port = 8000 +version = 1.0.0 + +[nacos] +enabled = false +server = 10.122.132.204:8848 +namespace = prod +group_name = BGS +username = nacos +password = bgs20250901 + diff --git a/config/config.ini.example b/config/config.ini.example index b89b59f..5a6c60e 100644 --- a/config/config.ini.example +++ b/config/config.ini.example @@ -34,12 +34,33 @@ model_section = gpt-4o url = http://10.122.176.97:21020 api_key = ragflow-xxxxx retrieval = /api/v1/retrieval -upload = /api/v1/documents -# 上传模式:overwrite(覆盖更新)或 append(追加) -upload_mode = overwrite +retrieval_top_k = 3 table_retrieval_dataset_id = sql_gen_dataset_id = +[redis] +# 是否启用 Redis 缓存(用于 sql_gen_prompts) +enabled = false +url = redis://localhost:6379/0 +db = 0 +# SQL 提示词缓存过期秒数 +sql_prompt_ttl = 600 + +[stream] +# /api/workflows/stream 进度事件间隔(秒) +progress_interval = 0.3 + +[logging_mysql] +# 是否启用结构化日志写入 MySQL +enabled = false +host = 127.0.0.1 +port = 3306 +user = root +password = +database = more_dots +table = structured_logs +connect_timeout = 5 + [nacos] # 是否启用 Nacos 注册 enabled = false diff --git a/config/sql_gen_prompts/apbo_hic_ssoc_consumption.json b/config/sql_gen_prompts/apbo_hic_ssoc_consumption.json new file mode 100644 index 0000000..616fda3 --- /dev/null +++ b/config/sql_gen_prompts/apbo_hic_ssoc_consumption.json @@ -0,0 +1,214 @@ +{ + "meta": { + "domain": "订单消耗与积压统计", + "data_source": "dwd_ai.apbo_hic_ssoc_consumption", + "description": "此模型用于分析HIC和SSOC系统的订单消耗情况、积压订单统计及订单状态追踪,通过订单创建日期进行时间维度分析。", + "fields_list": ["so", "soid", "original_pn", "ship_pn", "service_delivery_type", "warranty", "actual_wh", "item_status", "item_creation_date", "allocation_datetime", "backlog_status", "system_flag", "machine_type"] + }, + "data_model_specification": { + "core_status_fields": { + "backlog_status": { + "type": "varchar(64)", + "meaning": "积压状态标识,'1'表示积压订单,其他值为非积压订单", + "business_rule": "backlog_status = '1' 标识为积压订单" + }, + "system_flag": { + "type": "varchar(64)", + "meaning": "系统来源标识", + "values": ["HIC", "SSOC"], + "business_rule": "区分订单来自HIC系统还是SSOC系统" + }, + "machine_type": { + "type": "varchar(64)", + "meaning": "客户购买成品的系列", + "business_rule": "用于按产品系列进行过滤和分析,例如:machine_type = '20LN'" + } + }, + "mandatory_display_fields": { + "rule": "以下三个核心维度字段建议出现在SELECT子句中,特别是进行分组统计时。否则应展示fields_list中的全部字段", + "actual_fields": ["item_creation_date", "system_flag", "backlog_status"], + "user_alias": ["创建日期", "系统来源", "积压状态"], + "exception_condition": "如果用户明确指定字段数≤2,且均为具体业务字段,则可以不包含全部建议字段。" + }, + "other_fields": { + "order_info_fields": ["so", "soid", "original_pn", "ship_pn"], + "service_fields": ["service_delivery_type", "warranty"], + "warehouse_fields": ["actual_wh"], + "status_fields": ["item_status"], + "time_fields": ["allocation_datetime"], + "product_series_fields": ["machine_type"] + }, + "date_field": { + "primary_date": "item_creation_date", + "format": "DATE类型", + "usage": "主要的时间维度字段,用于按日期统计" + } + }, + "business_logic_rules": { + "alias_handling": "将用户提到的业务术语别名转换为完整字段名后再生成SQL", + "field_selection_logic": { + "priority_order": [ + { + "level": 1, + "name": "用户指定字段", + "rule": "首先添加用户明确提到的所有字段。" + }, + { + "level": 2, + "name": "必要补充字段", + "rule": "如果用户查询涉及时间、系统或状态分析,建议补充相关维度字段。", + "threshold": "用户查询涉及分组统计时建议补充item_creation_date和system_flag" + }, + { + "level": 3, + "name": "智能推断字段", + "rules": [ + { + "trigger": ["积压", "backlog", "积压订单", "待处理"], + "add_fields": ["backlog_status"], + "where_condition": "backlog_status = '1'" + }, + { + "trigger": ["系统", "来源", "HIC", "SSOC"], + "add_fields": ["system_flag"] + }, + { + "trigger": ["日期", "时间", "创建时间", "item_creation"], + "add_fields": ["item_creation_date"] + }, + { + "trigger": ["订单", "so", "订单号"], + "add_fields": ["so", "soid"] + }, + { + "trigger": ["物料", "PN", "料号", "零件"], + "add_fields": ["original_pn", "ship_pn"] + }, + { + "trigger": ["仓库", "actual_wh", "出货仓库"], + "add_fields": ["actual_wh"] + }, + { + "trigger": ["状态", "item_status", "订单状态"], + "add_fields": ["item_status"] + }, + { + "trigger": ["系列", "成品系列", "machine_type", "20LN", "产品类型"], + "add_fields": ["machine_type"] + } + ] + }, + { + "level": 4, + "name": "时间维度处理", + "rules": [ + { + "condition": "用户查询涉及统计、趋势、每日等时间概念", + "action": "必须包含item_creation_date字段" + }, + { + "condition": "用户指定具体日期范围", + "action": "在WHERE条件中添加日期范围过滤" + } + ] + } + ], + "backlog_logic": { + "definition": "backlog_status = '1' 表示积压订单,这是核心业务标识", + "inclusion_rule": "当用户查询涉及积压相关分析时,默认WHERE条件包含backlog_status = '1'", + "important_note": "统计总订单数时需要同时考虑积压和非积压状态" + } + }, + "aggregation_mode": { + "trigger_keywords": ["统计", "汇总", "总数", "有多少", "数量", "按...分组", "每个", "每日", "count", "sum", "total", "group by", "趋势", "分布"], + "column_naming_rules": { + "aggregate_functions": { + "COUNT": "COUNT({column}) AS {column}_count", + "COUNT_DISTINCT": "COUNT(DISTINCT {column}) AS {column}_distinct_count", + "SUM_CASE": "SUM(CASE WHEN {condition} THEN 1 ELSE 0 END) AS {alias}_count" + }, + "group_by_fields": "保持原字段名", + "examples": [ + "用户说'按日期统计积压订单数' → SELECT item_creation_date, COUNT(*) AS backlog_count WHERE backlog_status = '1' GROUP BY item_creation_date", + "用户说'统计每个系统的订单数' → SELECT system_flag, COUNT(*) AS order_count GROUP BY system_flag", + "用户说'按产品系列统计订单数' → SELECT machine_type, COUNT(*) AS order_count GROUP BY machine_type" + ] + }, + "allowed_group_by_fields": ["item_creation_date", "system_flag", "backlog_status", "actual_wh", "item_status", "service_delivery_type", "warranty", "machine_type"] + }, + "default_behavior": { + "sorting": { + "date_queries": "item_creation_date DESC", + "count_queries": "按统计字段降序", + "default": "item_creation_date DESC" + }, + "limit": "禁止使用LIMIT", + "backlog_filter": "当用户明确查询'积压订单'时,WHERE条件包含backlog_status = '1'" + }, + "special_scenarios": { + "积压率计算": "积压订单数 / 总订单数", + "系统对比": "HIC vs SSOC 系统订单分布", + "时间趋势": "按日/月统计订单创建趋势", + "物料分析": "按original_pn或ship_pn统计热门物料", + "产品系列分析": "按machine_type统计不同产品系列的订单分布" + } + }, + "field_mapping_reference": { + "key_fields": { + "backlog_status": "积压状态标识('1'=积压订单)", + "system_flag": "系统来源(HIC/SSOC)", + "item_creation_date": "订单创建日期(主要时间维度)", + "so": "订单号", + "soid": "订单ID", + "original_pn": "原始物料号", + "ship_pn": "出货物料号", + "machine_type": "客户购买成品的系列" + }, + "alias_mapping": { + "创建日期": "item_creation_date", + "日期": "item_creation_date", + "系统": "system_flag", + "来源": "system_flag", + "积压": "backlog_status", + "订单号": "so", + "订单": "so", + "物料": "original_pn", + "料号": "original_pn", + "PN": "original_pn", + "出货PN": "ship_pn", + "仓库": "actual_wh", + "状态": "item_status", + "系列": "machine_type", + "成品系列": "machine_type", + "产品类型": "machine_type", + "机器类型": "machine_type" + }, + "mapping_rule": "生成SQL时必须使用右侧的完整字段名,用户别名仅用于理解意图" + }, + "examples": { + "统计查询": { + "user": "统计每日的订单总数和积压订单数", + "sql": "SELECT item_creation_date, COUNT(*) AS total_orders, SUM(CASE WHEN backlog_status = '1' THEN 1 ELSE 0 END) AS backlog_orders FROM dwd_ai.apbo_hic_ssoc_consumption GROUP BY item_creation_date ORDER BY item_creation_date DESC" + }, + "物料分析": { + "user": "查看积压最多的物料Top 10", + "sql": "SELECT original_pn, COUNT(*) AS backlog_count FROM dwd_ai.apbo_hic_ssoc_consumption WHERE backlog_status = '1' GROUP BY original_pn ORDER BY backlog_count DESC" + }, + "时间范围查询": { + "user": "查询2024年1月的所有订单", + "sql": "SELECT so, soid, original_pn, system_flag, backlog_status, item_creation_date FROM dwd_ai.apbo_hic_ssoc_consumption WHERE item_creation_date >= '2024-01-01' AND item_creation_date <= '2024-01-31' ORDER BY item_creation_date DESC" + }, + "复合条件查询": { + "user": "查询SSOC系统中实际仓库为3001的非积压订单", + "sql": "SELECT so, soid, original_pn, item_creation_date, item_status FROM dwd_ai.apbo_hic_ssoc_consumption WHERE system_flag = 'SSOC' AND actual_wh = '3001' AND backlog_status != '1' ORDER BY item_creation_date DESC" + }, + "按日期和产品系列过滤查询": { + "user": "查询20260115, 20LN系列的订单", + "sql": "SELECT so, soid, machine_type, original_pn, ship_pn, service_delivery_type, warranty, actual_wh, item_status, item_creation_date, allocation_datetime, backlog_status, system_flag FROM dwd_ai.apbo_hic_ssoc_consumption WHERE machine_type = '20LN' AND item_creation_date = '2026-01-15' ORDER BY item_creation_date DESC" + }, + "产品系列统计": { + "user": "按产品系列统计订单总数", + "sql": "SELECT machine_type, COUNT(*) AS order_count FROM dwd_ai.apbo_hic_ssoc_consumption GROUP BY machine_type ORDER BY order_count DESC" + } + } +} \ No newline at end of file diff --git a/config/sql_gen_prompts/apbo_milestone_info.json b/config/sql_gen_prompts/apbo_milestone_info.json new file mode 100644 index 0000000..5e72ad3 --- /dev/null +++ b/config/sql_gen_prompts/apbo_milestone_info.json @@ -0,0 +1,97 @@ +{ + "meta": { + "domain": "订单物流节点(Milestone)追踪与查询", + "keywords": ["milestone", "物流节点", "运输节点", "logistics node", "shipment tracking", "节点查询", "节点状态", "节点时间", "节点跟踪", "里程碑", "节点明细"], + "description": "此模型用于查询订单在物流运输全链路中的关键节点(milestone)信息,追踪从订单创建、提货、运输到签收的完整状态流。每个订单可对应多个节点记录以反映运输进度。" + }, + "data_model_specification": { + "data_source": "dwd_ai.apbo_milestone_info", + "fields_list": ["service_order_id", "soid", "part_number", "PO", "po_creation_date", "prid", "DN", "dn_date", "gi_date", "BOL", "hawb", "pickup_time", "etd", "atd", "eta", "ata", "pod", "gr_date", "status", "status_date", "service_order_creation_date", "so_eta", "topmost_pn", "commodity_code", "ship_to_country", "region", "dc_plant", "mtm", "machine_sn","machine_type", "whether_premier", "stm_planner", "category", "lenovo_ref_no"], + "mandatory_display_fields": { + "rule1": "字段默认全部展示。日期字段需格式化为指定字符串格式。", + "rule2": "select时,所有NULL值使用空字符串''代替", + "rule3": "部分字段select时的顺序如下:po,po_creation_date,prid,dn,dn_date,gi_date,bol,hawb,pickup_time,etd,atd,eta,ata,pod,gr_date", + "rule4": "禁止使用limit" + }, + + "optional_fields": { + "material_fields": ["topmost_pn", "commodity_code"], + "location_fields": ["ship_to_country", "region", "dc_plant"], + "machine_fields": ["mtm", "machine_sn","machine_type"], + "service_fields": ["whether_premier", "stm_planner", "lenovo_ref_no"], + "Field_combination": ["category", "status", "status_date"] + } + }, + + "business_logic_rules": { + "query_recognition_rules": { + "milestone_keywords": ["物流节点", "运输节点", "节点", "里程碑", "milestone", "logistics node", "节点时间", "节点状态", "节点查询", "节点跟踪"], + "filter_based_patterns": ["查询{单据号}的物流节点", "追踪{订单号}的节点信息", "显示{节点状态}的订单", "查找{节点时间}的物流记录", "获取{节点类型}的详细信息"] + }, + "default_behavior": { + "sorting": "status_date DESC, service_order_creation_date DESC" + } + }, + + "field_mapping_reference": { + "key_identifiers": { + "service_order_id": {"type": "varchar(20)", "desc": "主订单号,唯一标识", "example": "4020438779", "required": true, "alias": ["SO", "订单号", "服务订单", "so", "order_no", "Service Order"]}, + "soid": {"type": "varchar(22)", "desc": "服务订单明细ID=service_order_id+两位序号", "example": "402043877920", "required": true, "alias": ["SOID", "订单明细", "服务订单明细", "soid", "Service Order Item Detail"]}, + "part_number": {"type": "varchar(30)", "desc": "原始申请物料号", "example": "5CB1L57599", "required": true, "alias": ["pn", "物料号", "零件号", "Part Number", "材料号"]} + }, + + "milestone_document_fields": { + "bol": {"type": "varchar(50)", "desc": "提单号(Bill of Lading)", "example": "BOL20231027001", "alias": ["提单", "海运提单", "提单号", "Bill of Lading", "B/L"]}, + "dn": {"type": "varchar(50)", "desc": "发货单号(Delivery Note)", "example": "DN20231027001", "alias": ["发货单", "送货单", "DN单", "Delivery Note", "发货单据"]}, + "po": {"type": "varchar(50)", "desc": "采购单号(Purchase Order)", "example": "PO20231027001", "alias": ["采购单", "PO单", "采购订单", "Purchase Order", "采购订单号"]}, + "hawb": {"type": "varchar(50)", "desc": "空运主单号(House Air Waybill)", "example": "HAWB12345678", "alias": ["空运单", "航空运单", "HAWB", "Air Waybill", "空运提单"]}, + "prid": {"type": "varchar(50)", "desc": "运输参考ID(PRID)", "example": "PRID123456", "alias": ["运输参考号", "PRID号", "运输ID", "运输参考ID", "Shipping Reference ID"]} + }, + + "milestone_date_fields": { + "service_order_creation_date": {"type": "date", "desc": "服务订单创建日期", "format": "YYYY-MM-DD HH:MM:SS", "query_format": "CAST(field_name AS VARCHAR(2048)) AS field_name 作为字符串展示", "alias": ["订单创建时间", "SO创建时间", "服务订单创建时间", "Order Creation Date"]}, + "po_creation_date": {"type": "date", "desc": "采购单创建日期", "format": "YYYY-MM-DD", "query_format": "CAST(field_name AS VARCHAR(2048)) AS field_name 作为字符串展示", "alias": ["采购单创建时间", "PO创建时间", "采购订单创建日期", "Purchase Order Creation Date"]}, + "dn_date": {"type": "date", "desc": "发货单创建日期", "format": "YYYY-MM-DD HH:MM:SS", "query_format": "CAST(field_name AS VARCHAR(2048)) AS field_name 作为字符串展示", "alias": ["发货单日期", "DN日期", "发货时间", "Delivery Note Date"]}, + "gi_date": {"type": "date", "desc": "货物发出日期(Goods Issue)", "format": "YYYY-MM-DD HH:MM:SS", "query_format": "CAST(field_name AS VARCHAR(2048)) AS field_name 作为字符串展示", "alias": ["货物发出时间", "发货日期", "出库时间", "Goods Issue Date", "出库日期"]}, + "pickup_time": {"type": "datetime", "desc": "提货时间", "format": "YYYY-MM-DD HH:MM:SS", "query_format": "CAST(field_name AS VARCHAR(2048)) AS field_name 作为字符串展示", "alias": ["提货日期", "取货时间", "提货时间点", "Pickup Time", "提取时间"]}, + "etd": {"type": "date", "desc": "预计出发时间(Estimated Time of Departure)", "format": "YYYY-MM-DD HH:MM:SS", "query_format": "CAST(field_name AS VARCHAR(2048)) AS field_name 作为字符串展示", "alias": ["预计出发", "计划出发时间", "ETD", "预计离港时间"]}, + "atd": {"type": "date", "desc": "实际出发时间(Actual Time of Departure)", "format": "YYYY-MM-DD HH:MM:SS", "query_format": "CAST(field_name AS VARCHAR(2048)) AS field_name 作为字符串展示", "alias": ["实际出发", "实际离港时间", "ATD", "实际出发时间"]}, + "eta": {"type": "date", "desc": "预计到达时间(Estimated Time of Arrival)", "format": "YYYY-MM-DD HH:MM:SS", "query_format": "CAST(field_name AS VARCHAR(2048)) AS field_name 作为字符串展示", "alias": ["预计到达", "计划到达时间", "ETA", "预计到港时间"]}, + "ata": {"type": "date", "desc": "实际到达时间(Actual Time of Arrival)", "format": "YYYY-MM-DD HH:MM:SS", "query_format": "CAST(field_name AS VARCHAR(2048)) AS field_name 作为字符串展示", "alias": ["实际到达", "实际到港时间", "ATA", "实际到达时间"]}, + "pod": {"type": "date", "desc": "签收单收到日期(Proof of Delivery)", "format": "YYYY-MM-DD HH:MM:SS", "query_format": "CAST(field_name AS VARCHAR(2048)) AS field_name 作为字符串展示", "alias": ["签收时间", "POD时间", "签收日期", "Proof of Delivery", "签收证明时间"]}, + "gr_date": {"type": "date", "desc": "收货日期(Goods Receipt)", "format": "YYYY-MM-DD HH:MM:SS", "query_format": "CAST(field_name AS VARCHAR(2048)) AS field_name 作为字符串展示", "alias": ["收货时间", "入库时间", "GR时间", "Goods Receipt Date", "收货日期"]}, + "status_date": {"type": "date", "desc": "状态最后更新时间", "format": "YYYY-MM-DD HH:MM:SS", "query_format": "CAST(field_name AS VARCHAR(2048)) AS field_name 作为字符串展示", "alias": ["状态更新时间", "最后更新", "状态日期", "Status Update Date"]}, + "so_eta": {"type": "date", "desc": "订单预计到达时间(SO ETA)", "format": "YYYY-MM-DD HH:MM:SS", "query_format": "CAST(field_name AS VARCHAR(2048)) AS field_name 作为字符串展示", "alias": ["订单最终到达时间", "订单ETA", "服务订单预计到达", "Service Order ETA"]} + }, + + "additional_fields": { + "topmost_pn": {"type": "varchar(30)", "desc": "最优替换物料号", "example": "5CB1L57599", "alias": ["tp", "最优物料", "Topmost Part Number", "首选物料号", "最优替换零件"]}, + "commodity_code": {"type": "varchar(2)", "desc": "商品分类代码", "example": "PL", "alias": ["cc", "商品代码", "物料分类", "Commodity Code", "商品类别"]}, + "ship_to_country": {"type": "varchar(2)", "desc": "目的地国家ISO代码", "example": "PH", "alias": ["目的地国家", "收货国家", "Ship to Country", "目标国家", "送达国家"]}, + "region": {"type": "varchar(10)", "desc": "区域代码", "example": "CAP", "alias": ["区域", "地区", "Region Code", "大区", "地域代码"]}, + "dc_plant": {"type": "varchar(10)", "desc": "发货配送中心", "example": "VN01", "alias": ["配送中心", "DC", "发货中心", "Distribution Center", "发货工厂"]}, + "mtm": {"type": "varchar(20)", "desc": "机器型号,请注意与mt区分,这不是mt,这是mtm,这是mtm,这是mtm", "example": "20QUS0SQ00", "alias": ["机型", "型号", "Machine Type Model", "设备型号", "机器型号代码"]}, + "machine_sn": {"type": "varchar(50)", "desc": "机器序列号", "example": "PW01554B", "alias": ["序列号", "SN", "Machine Serial Number", "设备序列号", "机器SN"]}, + "machine_type": {"type": "varchar(2048)","desc": "机器类型, 请注意与mtm区分,这不是mtm,这是mt,这是mt,这是mt","example": "20Q1","alias": ["MT","machine type", "机器类型"]}, + "whether_premier": {"type": "tinyint", "desc": "是否Premier服务", "values": {"0": "非Premier", "1": "Premier"}, "alias": ["是否Premier", "Premier服务", "是否优先", "Whether Premier", "优先服务标识"]}, + "stm_planner": {"type": "varchar(50)", "desc": "STM计划员姓名", "alias": ["计划员", "STM计划员", "STM Planner", "服务计划员", "运输计划员"]}, + "category": {"type": "varchar(20)", "desc": "订单类别", "values": ["Standard", "Emergency", "Critical"], "alias": ["类别", "订单类型", "Category", "服务类别", "订单分类"]}, + "status": {"type": "varchar(4096)", "desc": "状态", "example": "3001:02HL023*1(Need geo verify demand:2025-10-30)", "alias": ["状态"]}, + "lenovo_ref_no": {"type": "varchar(50)", "desc": "联想内部参考号", "alias": ["联想参考号", "内部参考号", "Lenovo Reference", "联想内部编号", "参考编号"]}, + "case_number": {"type": "string", "desc": "事件编号", "example": "CS12345678", "alias": ["case", "事件号", "案例号", "Case Number", "投诉编号", "事件编号"]} + } + }, + + "examples": { + "基础查询": { + "user": "查询订单的物流节点信息", + "sql": "SELECT service_order_id, soid, part_number, bol, dn, po, hawb, prid, CAST(service_order_creation_date AS VARCHAR(2048)) as service_order_creation_date, CAST(po_creation_date AS VARCHAR(2048)) as po_creation_date, CAST(dn_date AS VARCHAR(2048)) as dn_date, CAST(gi_date AS VARCHAR(2048)) as gi_date, CAST(pickup_time AS VARCHAR(2048)) as pickup_time, CAST(etd AS VARCHAR(2048)) as etd, CAST(atd AS VARCHAR(2048)) as atd, CAST(eta AS VARCHAR(2048)) as eta, CAST(ata AS VARCHAR(2048)) as ata, CAST(pod AS VARCHAR(2048)) as pod, CAST(gr_date AS VARCHAR(2048)) as gr_date, CAST(status_date AS VARCHAR(2048)) as status_date, CAST(so_eta AS VARCHAR(2048)) as so_eta FROM dwd_ai.apbo_milestone_info ORDER BY status_date DESC, service_order_creation_date DESC", + "field_selection_reason": "基础查询显示所有必填的milestone字段,包括新增的service_order_creation_date, so_eta,日期字段已格式化为指定字符串格式" + }, + "按主订单号查询": { + "user": "SO为4019630464的milestone信息", + "sql": "SELECT service_order_id, soid, part_number, bol, dn, po, hawb, prid, CAST(service_order_creation_date AS VARCHAR(2048)) as service_order_creation_date, CAST(po_creation_date AS VARCHAR(2048)) as po_creation_date, CAST(dn_date AS VARCHAR(2048)) as dn_date, CAST(gi_date AS VARCHAR(2048)) as gi_date, CAST(pickup_time AS VARCHAR(2048)) as pickup_time, CAST(etd AS VARCHAR(2048)) as etd, CAST(atd AS VARCHAR(2048)) as atd, CAST(eta AS VARCHAR(2048)) as eta, CAST(ata AS VARCHAR(2048)) as ata, CAST(pod AS VARCHAR(2048)) as pod, CAST(gr_date AS VARCHAR(2048)) as gr_date, CAST(status_date AS VARCHAR(2048)) as status_date, CAST(so_eta AS VARCHAR(2048)) as so_eta FROM dwd_ai.apbo_milestone_info WHERE service_order_id = '4019630464' ORDER BY status_date DESC, service_order_creation_date DESC", + "field_selection_reason": "用户提到'SO',根据映射规则应查询service_order_id字段,包含所有必填字段,日期字段已格式化为字符串" + } + } +} \ No newline at end of file diff --git a/config/sql_gen_prompts/apbo_region_usage_ib_report.json b/config/sql_gen_prompts/apbo_region_usage_ib_report.json new file mode 100644 index 0000000..89bae43 --- /dev/null +++ b/config/sql_gen_prompts/apbo_region_usage_ib_report.json @@ -0,0 +1,431 @@ +{ + "meta": { + "domain": "亚太区物料库存与用量分析", + "keywords": ["库存", "用量", "物料分析", "亚太区", "国家用量", "历史使用量", "近期用量", "IB", "库存数量"], + "description": "此模型用于分析亚太地区各物料在不同区域,配送中心及目的地国家的库存情况,近期用量及历史使用量。特别注意:国家字段(AU, VN, JP等)是数值型用量字段,表示该物料在该国家的历史用量,不是国家代码。", + "data_source": "dwd_ai.apbo_region_usage_ib_report" + }, + "data_model_specification": { + "fields_list": ["topmost_pn", "region", "commodity_code", "dc_plant", "country_total_cnt", "AU", "NZ", "LK", "VN", "JP", "HK", "SG", "TH", "PH", "IN", "BN", "NP", "BD", "KR", "ID", "FJ", "MY", "TW", "location_name", "ib", "usage_qty_8_week", "usage_qty_52_week", "total_history_usage"], + + "core_country_fields": { + "list": ["AU", "NZ", "LK", "VN", "JP", "HK", "SG", "TH", "PH", "IN", "BN", "NP", "BD", "KR", "ID", "FJ", "MY", "TW"], + "meaning": "各国家过去某几周的总使用量,字段类型为BIGINT.这些字段表示该物料在该国家的历史用量值.", + "critical_notes": "这些字段是数值型用量字段,不是国家代码。例如:VN字段表示物料在越南的历史使用量数值,不是字符串'VN'。" + }, + + "mandatory_display_fields": { + "rule1": "默认查询所有字段(SELECT *)", + "rule2": "当用户明确指定某些字段时,只查询这些字段,并保持用户提到的顺序", + "rule3": "WHERE条件中使用的所有字段(除了国家字段的数值比较外),必须在SELECT子句中展示", + "rule4": "国家字段是数值型,直接展示数值,不需要特殊处理", + "rule5": "select的所有字段必须用``包裹,如select `IN`,防止与SQL关键字冲突", + "actual_fields_order": ["topmost_pn", "region", "commodity_code", "dc_plant", "country_total_cnt", "AU", "NZ", "LK", "VN", "JP", "HK", "SG", "TH", "PH", "IN", "BN", "NP", "BD", "KR", "ID", "FJ", "MY", "TW", "location_name", "ib", "usage_qty_8_week", "usage_qty_52_week", "total_history_usage"] + } + }, + "business_logic_rules": { + "query_recognition_rules": { + "material_keywords": ["物料", "topmost_pn", "物料号", "零件号", "Part Number", "PN"], + "region_keywords": ["region", "区域", "大区", "地区", "Region Code", "CAP", "EMEA", "AMER"], + "commodity_keywords": ["commodity_code", "商品代码", "编码", "物料分类", "Commodity Code", "CC"], + "inventory_keywords": ["库存", "ib", "库存数量", "库存量", "Inventory Balance"], + "usage_keywords": ["用量", "使用量", "历史用量", "近期用量", "usage", "使用数量"], + "location_keywords": ["dc_plant", "plant", "配送中心", "工厂", "location", "地点"], + "country_usage_keywords": ["国家用量", "国家使用量", "country usage", "国家用量分析"], + "specific_country_keywords": { + "AU": ["澳大利亚", "澳洲", "AU", "Australia"], + "VN": ["越南", "VN", "Vietnam"], + "JP": ["日本", "JP", "Japan"], + "HK": ["香港", "HK", "Hong Kong"], + "SG": ["新加坡", "SG", "Singapore"], + "TH": ["泰国", "TH", "Thailand"], + "IN": ["印度", "IN", "India"], + "KR": ["韩国", "KR", "Korea"], + "ID": ["印度尼西亚", "印尼", "ID", "Indonesia"], + "MY": ["马来西亚", "MY", "Malaysia"], + "TW": ["台湾", "TW", "Taiwan"], + "NZ": ["新西兰", "NZ", "New Zealand"], + "PH": ["菲律宾", "PH", "Philippines"] + } + }, + + "country_field_interpretation": { + "fundamental_rule": "国家字段(AU, VN, JP等)是数值型用量字段(BIGINT),表示该物料在该国家的历史用量,不是国家代码", + "correct_usage": { + "in_select": "直接展示数值,如:SELECT AU, VN, JP", + "in_where": "不需要添加任何默认过滤条件,除非用户明确要求", + "as_numeric": "作为数值字段处理,支持比较运算符:>, <, >=, <=, =" + }, + "incorrect_usage": [ + "WHERE VN = 'VN' (错误:VN是数值字段,不是字符串)", + "WHERE VN LIKE '%VN%' (错误:VN是数值字段,不支持LIKE)", + "WHERE VN IN ('VN', 'AU') (错误:VN是数值字段,不是枚举值)" + ], + "user_intent_interpretation": { + "当用户说'country为VN'": "用户只是提到VN字段,但不一定要求VN>0,不需要添加过滤条件", + "当用户说'VN国家'": "用户指的是VN字段,按字段处理", + "当用户说'有VN用量的物料'": "需要添加WHERE VN > 0", + "当用户说'VN用量超过100'": "需要添加WHERE VN >= 100", + "当用户说'VN用量为0'": "需要添加WHERE VN = 0", + "当用户说'没有VN用量的'": "需要添加WHERE VN = 0" + } + }, + + "field_selection_logic": { + "priority_order": [ + { + "level": 1, + "name": "用户明确指定的字段", + "rule": "精确添加用户提到的字段,按用户问题中出现的顺序排列,不添加未提及的字段" + }, + { + "level": 2, + "name": "默认字段选择", + "rule": "如果用户没有明确指定字段,默认查询所有字段(SELECT *)" + } + ], + + "where_condition_fields_must_in_select": { + "rule": "WHERE条件中使用的所有字段(除了国家字段的数值比较外),必须在SELECT子句中展示", + "purpose": "确保查询结果的完整性,用户能看到过滤依据", + "examples": [ + "用户说'region为CAP的数据' → SELECT * ... WHERE region = 'CAP'", + "用户说'ib大于500的物料号' → SELECT topmost_pn, ib ... WHERE ib > 500", + "用户说'region为CAP且commodity_code为LF的物料号' → SELECT topmost_pn, region, commodity_code ... WHERE region = 'CAP' AND commodity_code = 'LF'" + ] + } + }, + + "where_condition_generation": { + "strict_rule": "只生成用户明确要求的过滤条件,不添加任何默认或假设的条件", + "condition_types": { + "exact_match": { + "pattern": "{字段}为{值}", + "sql": "{field} = '{value}'", + "examples": ["region为CAP → region = 'CAP'", "commodity_code为LF → commodity_code = 'LF'", "dc_plant为HKGDC → dc_plant = 'HKGDC'"] + }, + "country_field_handling": { + "important_note": "当用户提到'country为VN'时,不需要在WHERE中添加VN > 0,除非用户明确要求过滤用量", + "correct_interpretation": "用户只是提到VN字段,但没有要求过滤VN的值", + "only_add_filter_when": [ + "用户明确说'有VN用量的' → WHERE VN > 0", + "用户明确说'VN用量大于0的' → WHERE VN > 0", + "用户明确说'VN超过100的' → WHERE VN >= 100", + "用户明确说'VN用量为0的' → WHERE VN = 0", + "用户明确说'没有VN用量的' → WHERE VN = 0" + ], + "incorrect_interpretation": [ + "用户说'country为VN' → 错误:添加WHERE VN > 0", + "用户说'VN国家' → 错误:添加WHERE VN > 0", + "用户说'查看VN' → 错误:添加WHERE VN > 0" + ] + }, + "range_filter": { + "pattern": ["大于{值}", "超过{值}", "少于{值}", "小于{值}", "不低于{值}", "不超过{值}"], + "sql_mapping": { + "大于{值}": "{field} > {value}", + "超过{值}": "{field} >= {value}", + "少于{值}": "{field} < {value}", + "小于{值}": "{field} < {value}", + "不低于{值}": "{field} >= {value}", + "不超过{值}": "{field} <= {value}" + }, + "examples": [ + "库存大于500 → ib > 500", + "VN用量超过100 → VN >= 100", + "近期用量少于50 → usage_qty_8_week < 50" + ] + }, + "multiple_conditions": { + "pattern": ["且", "并且", "和", "同时", ","], + "sql": "AND", + "examples": [ + "region为CAP且commodity_code为LF → region = 'CAP' AND commodity_code = 'LF'", + "库存大于500且有VN用量 → ib > 500 AND VN > 0" + ] + } + }, + + "no_default_conditions": { + "rule": "绝对不添加任何用户没有明确要求的WHERE条件", + "examples_of_what_not_to_add": [ + "不要添加WHERE region = 'CAP'(除非用户明确要求)", + "不要添加WHERE ib > 0(除非用户明确要求)", + "不要添加WHERE commodity_code IS NOT NULL(除非用户明确要求)", + "不要添加任何假设性的过滤条件" + ] + } + }, + + "default_behavior": { + "sorting": "默认按ib(库存数量)降序排列:ORDER BY ib DESC", + "limit": "禁止使用limit", + "select_all": "默认查询所有字段:SELECT *", + "field_order": "当用户指定字段时,按用户提到的顺序排列字段", + "date_handling": "本表无日期字段,不需要日期转换" + } + }, + "field_mapping_reference": { + "critical_note": "此表包含亚太地区各物料的库存、近期用量及历史使用量的详细数据。特别注意:国家字段(AU, VN, JP等)是数值型用量字段,表示该物料在该国家的历史用量,不是国家代码。这些字段是BIGINT类型,支持数值比较操作。", + + "material_identifier_fields": { + "topmost_pn": { + "type": "string", + "desc": "顶级物料号,物料的唯一标识", + "example": "5CB1L57599", + "query_pattern": "物料为{值} → topmost_pn = '{value}'", + "alias": ["物料号", "零件号", "Part Number", "PN", "topmost", "物料编码"] + }, + "commodity_code": { + "type": "string", + "desc": "商品代码,物料分类标识", + "example": "LF", + "query_pattern": "commodity_code为{值} → commodity_code = '{value}'", + "alias": ["商品编码", "物料分类", "Commodity Code", "CC", "编码", "商品类别"] + } + }, + + "geographic_dimension_fields": { + "region": { + "type": "string", + "desc": "区域划分,如CAP(亚太区)、EMEA(欧洲中东非洲)等", + "example": "CAP", + "values": ["CAP", "EMEA", "AMER"], + "query_pattern": "region为{值} → region = '{value}'", + "alias": ["区域", "大区", "Region", "地区", "地理区域"] + }, + "dc_plant": { + "type": "string", + "desc": "配送中心/工厂代码", + "example": "HKGDC", + "query_pattern": "dc_plant为{值} → dc_plant = '{value}'", + "alias": ["工厂", "配送中心", "plant", "DC", "发货中心", "Distribution Center"] + }, + "location_name": { + "type": "string", + "desc": "地点名称", + "example": "Hong Kong Distribution Center", + "alias": ["地点名称", "位置名称", "location", "地点", "场所名称"] + } + }, + + "country_usage_fields_section": { + "important_note": "以下所有字段都是数值型(BIGINT),表示该物料在该国家的历史使用量,不是国家代码。这些字段支持数值比较操作(>, <, >=, <=, =, !=)。", + "critical_warning": "绝对不要将这些字段作为字符串处理,不要使用单引号,不要使用LIKE操作符。", + + "country_total_cnt": { + "type": "bigint", + "desc": "所有国家的总使用量计数", + "calculation_note": "可能是各国家字段的汇总或其他计算逻辑", + "alias": ["国家总计数", "总使用量", "country total", "总计"] + }, + + "country_usage_fields": { + "AU": { + "type": "bigint", + "desc": "澳大利亚的历史使用量", + "alias": ["澳大利亚", "澳洲", "AU用量", "Australia用量"] + }, + "NZ": { + "type": "bigint", + "desc": "新西兰的历史使用量", + "alias": ["新西兰", "NZ用量", "New Zealand用量"] + }, + "LK": { + "type": "bigint", + "desc": "斯里兰卡的历史使用量", + "alias": ["斯里兰卡", "LK用量", "Sri Lanka用量"] + }, + "VN": { + "type": "bigint", + "desc": "越南的历史使用量", + "alias": ["越南", "VN用量", "Vietnam用量"] + }, + "JP": { + "type": "bigint", + "desc": "日本的历史使用量", + "alias": ["日本", "JP用量", "Japan用量"] + }, + "HK": { + "type": "bigint", + "desc": "香港的历史使用量", + "alias": ["香港", "HK用量", "Hong Kong用量"] + }, + "SG": { + "type": "bigint", + "desc": "新加坡的历史使用量", + "alias": ["新加坡", "SG用量", "Singapore用量"] + }, + "TH": { + "type": "bigint", + "desc": "泰国的历史使用量", + "alias": ["泰国", "TH用量", "Thailand用量"] + }, + "PH": { + "type": "bigint", + "desc": "菲律宾的历史使用量", + "alias": ["菲律宾", "PH用量", "Philippines用量"] + }, + "IN": { + "type": "bigint", + "desc": "印度的历史使用量", + "alias": ["印度", "IN用量", "India用量"] + }, + "BN": { + "type": "bigint", + "desc": "文莱的历史使用量", + "alias": ["文莱", "BN用量", "Brunei用量"] + }, + "NP": { + "type": "bigint", + "desc": "尼泊尔的历史使用量", + "alias": ["尼泊尔", "NP用量", "Nepal用量"] + }, + "BD": { + "type": "bigint", + "desc": "孟加拉国的历史使用量", + "alias": ["孟加拉国", "BD用量", "Bangladesh用量"] + }, + "KR": { + "type": "bigint", + "desc": "韩国的历史使用量", + "alias": ["韩国", "KR用量", "Korea用量"] + }, + "ID": { + "type": "bigint", + "desc": "印度尼西亚的历史使用量", + "alias": ["印度尼西亚", "印尼", "ID用量", "Indonesia用量"] + }, + "FJ": { + "type": "bigint", + "desc": "斐济的历史使用量", + "alias": ["斐济", "FJ用量", "Fiji用量"] + }, + "MY": { + "type": "bigint", + "desc": "马来西亚的历史使用量", + "alias": ["马来西亚", "MY用量", "Malaysia用量"] + }, + "TW": { + "type": "bigint", + "desc": "台湾的历史使用量", + "alias": ["台湾", "TW用量", "Taiwan用量"] + } + }, + + "query_examples": { + "correct": [ + "WHERE VN > 0 (查询有越南用量的物料)", + "WHERE JP >= 100 (查询日本用量超过100的物料)", + "WHERE AU = 0 (查询没有澳大利亚用量的物料)", + "WHERE SG < 50 (查询新加坡用量少于50的物料)" + ], + "incorrect": [ + "WHERE VN = 'VN' (错误:VN是数值,不是字符串)", + "WHERE JP LIKE '%JP%' (错误:JP是数值,不支持LIKE)", + "WHERE AU IN ('AU', 'VN') (错误:AU是数值,不是枚举)" + ] + } + }, + + "inventory_metrics": { + "ib": { + "type": "bigint", + "desc": "库存数量(Inventory Balance)", + "query_pattern": "库存超过{值} → ib > {value}, 库存大于{值} → ib > {value}, 库存少于{值} → ib < {value}", + "alias": ["库存", "库存量", "库存数量", "Inventory", "库存余额", "IB"] + } + }, + + "usage_time_series_fields": { + "usage_qty_8_week": { + "type": "bigint", + "desc": "近8周的使用量", + "alias": ["近8周用量", "近期用量", "短期用量", "8周用量", "近期使用量"] + }, + "usage_qty_52_week": { + "type": "bigint", + "desc": "近52周的使用量", + "alias": ["近52周用量", "年度用量", "长期用量", "52周用量", "年度使用量"] + }, + "total_history_usage": { + "type": "bigint", + "desc": "历史总使用量", + "alias": ["历史总用量", "总使用量", "累计用量", "历史累计", "total usage"] + } + } + }, + "examples": { + "简单查询-全部字段": { + "user": "region为CAP的数据", + "sql": "SELECT * FROM dwd_ai.apbo_region_usage_ib_report WHERE region = 'CAP' ORDER BY ib DESC", + "field_selection_reason": "用户没有指定字段,默认查询所有字段。WHERE条件中使用了region字段。" + }, + + "简单查询-基础过滤": { + "user": "region为CAP,commodity_code为LF的数据", + "sql": "SELECT * FROM dwd_ai.apbo_region_usage_ib_report WHERE region = 'CAP' AND commodity_code = 'LF' ORDER BY ib DESC", + "field_selection_reason": "用户没有指定字段,默认查询所有字段。WHERE条件中使用了region和commodity_code字段。" + }, + + "简单查询-带国家字段但不过滤": { + "user": "region为CAP,commodity_code为LF的,country为VN的数据", + "sql": "SELECT * FROM dwd_ai.apbo_region_usage_ib_report WHERE region = 'CAP' AND commodity_code = 'LF' ORDER BY ib DESC", + "field_selection_reason": "用户提到'country为VN'但没有要求过滤VN用量,所以不添加VN > 0条件。默认查询所有字段。" + }, + + "简单查询-带国家用量过滤": { + "user": "region为CAP,commodity_code为LF的,有VN用量的数据", + "sql": "SELECT * FROM dwd_ai.apbo_region_usage_ib_report WHERE region = 'CAP' AND commodity_code = 'LF' AND VN > 0 ORDER BY ib DESC", + "field_selection_reason": "用户明确要求'有VN用量的',所以添加VN > 0条件。默认查询所有字段。" + }, + + "简单查询-指定字段": { + "user": "查看物料号和库存数量", + "sql": "SELECT topmost_pn, ib FROM dwd_ai.apbo_region_usage_ib_report ORDER BY ib DESC", + "field_selection_reason": "用户明确指定了topmost_pn和ib字段,只查询这两个字段,按用户提到的顺序排列。" + }, + + "简单查询-多条件组合": { + "user": "region为CAP,commodity_code为LF,库存大于500,有VN用量的数据", + "sql": "SELECT * FROM dwd_ai.apbo_region_usage_ib_report WHERE region = 'CAP' AND commodity_code = 'LF' AND ib > 500 AND VN > 0 ORDER BY ib DESC", + "field_selection_reason": "用户没有指定字段,默认查询所有字段。WHERE条件包含region、commodity_code、ib和VN字段的过滤。" + }, + + "指定字段且包含WHERE字段": { + "user": "region为CAP的物料号和库存", + "sql": "SELECT topmost_pn, ib, region FROM dwd_ai.apbo_region_usage_ib_report WHERE region = 'CAP' ORDER BY ib DESC", + "field_selection_reason": "用户指定了topmost_pn和ib字段,但WHERE条件中使用了region字段,所以必须包含region字段在SELECT中。" + }, + + "国家用量范围查询": { + "user": "VN用量超过100且库存大于200的物料", + "sql": "SELECT * FROM dwd_ai.apbo_region_usage_ib_report WHERE VN >= 100 AND ib > 200 ORDER BY ib DESC", + "field_selection_reason": "用户没有指定字段,默认查询所有字段。WHERE条件包含VN和ib字段的数值范围过滤。" + }, + + "多国家用量查询": { + "user": "有VN用量且有AU用量的物料号", + "sql": "SELECT topmost_pn, VN, AU FROM dwd_ai.apbo_region_usage_ib_report WHERE VN > 0 AND AU > 0 ORDER BY ib DESC", + "field_selection_reason": "用户指定了topmost_pn字段,并提到了VN和AU用量,所以包含这些字段。WHERE条件包含VN>0和AU>0。" + }, + + "混合条件复杂查询": { + "user": "region为CAP,commodity_code为LF,库存大于500,有VN用量且超过50,近期用量少于100的物料信息", + "sql": "SELECT * FROM dwd_ai.apbo_region_usage_ib_report WHERE region = 'CAP' AND commodity_code = 'LF' AND ib > 500 AND VN > 50 AND usage_qty_8_week < 100 ORDER BY ib DESC", + "field_selection_reason": "用户没有指定字段,默认查询所有字段。WHERE条件包含多个字段的复杂过滤。" + }, + + "国家用量为零查询": { + "user": "没有VN用量的物料", + "sql": "SELECT * FROM dwd_ai.apbo_region_usage_ib_report WHERE VN = 0 ORDER BY ib DESC", + "field_selection_reason": "用户明确要求'没有VN用量的',所以添加VN = 0条件。默认查询所有字段。" + }, + + "仅查看特定国家用量": { + "user": "查看VN和JP的用量", + "sql": "SELECT VN, JP FROM dwd_ai.apbo_region_usage_ib_report ORDER BY ib DESC", + "field_selection_reason": "用户明确指定了VN和JP字段,只查询这两个字段。没有WHERE条件。" + } + } +} \ No newline at end of file diff --git a/config/sql_gen_prompts/apbo_tp_multiple_impact.json b/config/sql_gen_prompts/apbo_tp_multiple_impact.json new file mode 100644 index 0000000..9545596 --- /dev/null +++ b/config/sql_gen_prompts/apbo_tp_multiple_impact.json @@ -0,0 +1,182 @@ +{ + "meta": { + "domain": "TP物料多重影响分析", + "keywords": ["multiple impact", "tp物料", "物料影响", "pal_2h", "达标状态", "不达标状态", "物料分析", "TP分析", "影响分析"], + "description": "此模型用于分析TP物料基于pal_2h状态的详细记录和统计信息。pal_2h='N'表示不达标,pal_2h='Y'表示达标。支持明细查询和统计分析两种模式,根据不同查询意图自动切换模式。", + "data_source": "dwd_ai.apbo_tp_multiple_impact" + }, + + "data_model_specification": { + "fields_list": ["topmost_pn", "service_order_id", "soid", "pal_2h"], + "mandatory_display_fields": { + "rule1": "所有核心字段必须出现在SELECT子句中,除非用户明确指定排除。", + "rule2": "pal_2h字段在两种模式下都必须显示", + "rule3": "SELECT子句中字段顺序建议为:pal_2h, topmost_pn, service_order_id, soid", + "action": "detail_mode下显示所有字段,statistical_mode下显示分组字段和统计结果" + }, + + "optional_fields": { + "key_fields": ["topmost_pn", "service_order_id", "soid"], + "status_fields": ["pal_2h"], + "grouping_fields": ["topmost_pn", "pal_2h"] + } + }, + + "business_logic_rules": { + "query_recognition_rules": { + "detail_mode_keywords": ["有哪些", "查看", "列出", "显示", "查询", "搜索", "找出", "记录", "明细", "详情", "具体"], + "statistical_mode_keywords": ["统计", "汇总", "总数", "有多少", "数量", "count", "条数", "计数", "分组", "分布", "占比", "比例"], + "tp_material_keywords": ["tp", "物料", "topmost_pn", "零件", "零件号", "物料号", "TP物料"], + "order_keywords": ["订单", "so", "service_order_id", "SOID", "soid"], + "status_keywords": ["状态", "pal_2h", "达标", "不达标", "Y", "N", "status"] + }, + + "default_behavior": { + "detail_mode": { + "sorting": "默认按topmost_pn, service_order_id排序", + "pal_2h_display": "pal_2h字段必须显示在SELECT结果中", + "pal_2h_filter": "用户未指定状态条件时,禁止使用pal_2h过滤" + }, + "statistical_mode": { + "sorting": "COUNT(*) DESC", + "limit": "禁止使用limit", + "pal_2h_filter": "用户未指定状态条件时,禁止使用pal_2h过滤" + }, + "alias_handling": "将用户提到的别名转换为完整字段名后再生成SQL" + } + }, + + "field_mapping_reference": { + "critical_note": "注意区分detail_mode(明细查询)和statistical_mode(统计分析)两种模式,根据用户query中的关键词自动判断模式。detail_mode禁止使用聚合函数,statistical_mode必须使用聚合函数。", + + "key_identifiers": { + "topmost_pn": { + "type": "varchar", + "desc": "TP物料编号", + "example": "02HK965", + "required": true, + "alias": ["tp", "物料", "物料号", "零件号", "TP物料", "零件编号", "物料编码"] + }, + "service_order_id": { + "type": "varchar", + "desc": "服务订单ID", + "example": "4020438779", + "required": true, + "alias": ["so", "订单", "订单号", "service_order", "订单ID", "SO", "服务订单"] + }, + "soid": { + "type": "varchar", + "desc": "SOID(服务订单明细ID)", + "example": "402043877920", + "required": true, + "alias": ["SOID", "子单号", "明细ID", "订单明细", "服务订单明细"] + } + }, + + "status_fields": { + "pal_2h": { + "type": "varchar(1)", + "desc": "状态字段:'Y'表示达标,'N'表示不达标", + "values": { + "Y": "达标", + "N": "不达标", + "NULL": "无状态" + }, + "business_rule": "用户未指定状态条件时,默认查询不达标记录(pal_2h = 'N')", + "alias": ["状态", "达标状态", "不达标状态", "pal_2h状态", "status", "达标标识"] + } + }, + + "condition_mapping": { + "达标": "pal_2h = 'Y'", + "不达标": "pal_2h = 'N'", + "无状态": "pal_2h IS NULL", + "所有状态": "pal_2h IN ('Y', 'N')", + "有状态的": "pal_2h IN ('Y', 'N')", + "状态完整": "pal_2h IN ('Y', 'N')" + } + }, + + "examples": { + "detail_mode_examples": { + "example1": { + "user": "tp为02HK965的multiple impact有哪些", + "mode": "detail_mode", + "reason": "包含'有哪些'关键词,表示查看具体记录;未指定状态,默认查不达标", + "sql": "SELECT pal_2h, topmost_pn, service_order_id, soid FROM dwd_ai.apbo_tp_multiple_impact WHERE topmost_pn = '02HK965' ORDER BY topmost_pn, service_order_id" + }, + "example2": { + "user": "查看达标的记录", + "mode": "detail_mode", + "reason": "包含'查看'关键词,表示查看具体记录;指定了达标状态", + "sql": "SELECT pal_2h, topmost_pn, service_order_id, soid FROM dwd_ai.apbo_tp_multiple_impact WHERE ORDER BY topmost_pn, service_order_id" + }, + "example3": { + "user": "列出topmost_pn为02HK965的记录", + "mode": "detail_mode", + "reason": "包含'列出'关键词,表示查看具体记录;未指定状态,默认查不达标", + "sql": "SELECT pal_2h, topmost_pn, service_order_id, soid FROM dwd_ai.apbo_tp_multiple_impact WHERE topmost_pn = '02HK965' ORDER BY topmost_pn, service_order_id" + }, + "example4": { + "user": "查询订单4020438779的TP物料影响", + "mode": "detail_mode", + "reason": "包含'查询'关键词,表示查看具体记录;未指定状态,默认查不达标", + "sql": "SELECT pal_2h, topmost_pn, service_order_id, soid FROM dwd_ai.apbo_tp_multiple_impact WHERE service_order_id = '4020438779' ORDER BY topmost_pn, soid" + } + }, + + "statistical_mode_examples": { + "example1": { + "user": "统计不同状态的物料数量", + "mode": "statistical_mode", + "reason": "包含'统计'关键词,表示统计数量;统计不同状态,需要显示所有状态", + "sql": "SELECT pal_2h, COUNT(topmost_pn) AS topmost_pn_count FROM dwd_ai.apbo_tp_multiple_impact GROUP BY pal_2h ORDER BY topmost_pn_count DESC" + }, + "example2": { + "user": "按topmost_pn分组统计每个物料的记录数", + "mode": "statistical_mode", + "reason": "包含'统计'和'分组'关键词;未指定状态,默认只统计不达标", + "sql": "SELECT topmost_pn, pal_2h, COUNT(*) AS count FROM dwd_ai.apbo_tp_multiple_impact GROUP BY topmost_pn, pal_2h ORDER BY count DESC" + }, + "example3": { + "user": "汇总各TP物料的影响分布", + "mode": "statistical_mode", + "reason": "包含'汇总'关键词,表示统计分布;未指定状态,默认只统计不达标", + "sql": "SELECT topmost_pn, COUNT(*) AS record_count FROM dwd_ai.apbo_tp_multiple_impact GROUP BY topmost_pn ORDER BY record_count DESC" + } + }, + + "edge_cases": { + "example1": { + "user": "有多少条tp为02HK965的记录", + "mode": "statistical_mode", + "reason": "包含'有多少'关键词,虽然指定了具体物料,但目的是获取数量;未指定状态,默认只统计不达标", + "sql": "SELECT COUNT(*) AS count FROM dwd_ai.apbo_tp_multiple_impact WHERE topmost_pn = '02HK965'" + }, + "example2": { + "user": "显示所有状态为达标和不达标的记录", + "mode": "detail_mode", + "reason": "包含'显示'关键词,表示查看具体记录;明确要求查看两种状态,不使用默认值", + "sql": "SELECT pal_2h, topmost_pn, service_order_id, soid FROM dwd_ai.apbo_tp_multiple_impact ORDER BY topmost_pn, service_order_id" + }, + "example3": { + "user": "统计状态为NULL的记录数量", + "mode": "statistical_mode", + "reason": "包含'统计'关键词,表示统计数量;明确指定NULL状态,不使用默认值", + "sql": "SELECT COUNT(*) AS count FROM dwd_ai.apbo_tp_multiple_impact WHERE pal_2h IS NULL" + }, + "example4": { + "user": "查看所有状态的记录详情", + "mode": "detail_mode", + "reason": "包含'查看'和'详情'关键词,表示查看具体记录;要求所有状态,不使用默认值", + "sql": "SELECT pal_2h, topmost_pn, service_order_id, soid FROM dwd_ai.apbo_tp_multiple_impact OR pal_2h IS NULL ORDER BY pal_2h, topmost_pn, service_order_id" + }, + "example5": { + "user": "统计每个TP物料的不同状态数量", + "mode": "statistical_mode", + "reason": "包含'统计'关键词,表示统计分析;统计每个物料的各状态分布", + "sql": "SELECT topmost_pn, pal_2h, COUNT(*) AS count FROM dwd_ai.apbo_tp_multiple_impact GROUP BY topmost_pn, pal_2h ORDER BY topmost_pn, pal_2h" + } + } + } +} \ No newline at end of file diff --git a/config/sql_gen_prompts/example_table.json b/config/sql_gen_prompts/example_table.json deleted file mode 100644 index 5386308..0000000 --- a/config/sql_gen_prompts/example_table.json +++ /dev/null @@ -1,10 +0,0 @@ -{ - "table": "example_table", - "description": "示例表模型提示词", - "system_prompt": "You are an expert SQL generator.", - "business_prompt": "Generate SQL for example_table based on the user's intent.", - "constraints": [ - "Use only fields defined in this table.", - "Return only SQL without explanations." - ] -} diff --git a/config/table_retrieval_prompts/tables.json b/config/table_retrieval_prompts/tables.json index 4223632..02dd22d 100644 --- a/config/table_retrieval_prompts/tables.json +++ b/config/table_retrieval_prompts/tables.json @@ -16,9 +16,31 @@ "work order information", "recovery ETA", "history order", - "Warranty type" + "Warranty type", + "category" ], - "x_table_name": [ - "xx" + "apbo_milestone_info": [ + "GR or POD", + "POD", + "GR", + "milestone status", + "milestone", + "物流节点", + "shipment tracking", + "里程碑", + "节点明细" + ], + "apbo_hic_ssoc_consumption": [ + "consumption消耗记录", + "consumption order" + ], + "apbo_tp_multiple_impact": [ + "multiple impact orders", + "不达标的multiple impact", + "multiple impact status" + ], + "apbo_region_usage_ib_report": [ + "ib", + "usage" ] } \ No newline at end of file diff --git a/k8s/deployment.yaml b/k8s/deployment.yaml new file mode 100644 index 0000000..8b1aba6 --- /dev/null +++ b/k8s/deployment.yaml @@ -0,0 +1,35 @@ +apiVersion: apps/v1 +kind: Deployment +metadata: + name: more_dots + namespace: default +spec: + replicas: 2 + selector: + matchLabels: + app: more_dots + template: + metadata: + labels: + app: more_dots + spec: + containers: + - name: more_dots + image: harbor.yourdomain.com/library/more_dots:${IMAGE_TAG} + ports: + - containerPort: 8000 + env: + - name: ENVIRONMENT + value: ${DEPLOY_ENV} +--- +apiVersion: v1 +kind: Service +metadata: + name: more_dots-service +spec: + selector: + app: more_dots + ports: + - port: 80 + targetPort: 8000 + type: ClusterIP \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index 491d7c1..d1158fa 100644 --- a/requirements.txt +++ b/requirements.txt @@ -5,6 +5,8 @@ langchain-openai>=1.1.6 pydantic>=2.0.0 fastapi>=0.110.0 uvicorn>=0.30.0 -nacos-sdk-python>=2.0.9 +nacos-sdk-python==2.0.9 httpx>=0.27.0 pyyaml>=6.0.1 +redis>=5.0.0 +pymysql>=1.1.1 diff --git a/schemas/chat_message_response.py b/schemas/chat_message_response.py new file mode 100644 index 0000000..50a43af --- /dev/null +++ b/schemas/chat_message_response.py @@ -0,0 +1,13 @@ +from pydantic import BaseModel + + +class ChatMessageResponseDTO(BaseModel): + """流式消息响应模型""" + + id: str + event: str = "message" + task_id: str + message_id: str + conversation_id: str + answer: str + created_at: int diff --git a/scripts/test_endpoints.py b/scripts/test_endpoints.py new file mode 100644 index 0000000..0dbed19 --- /dev/null +++ b/scripts/test_endpoints.py @@ -0,0 +1,139 @@ +import json +import sys +from typing import Any, Dict, List + +import httpx + +from config import Config + + +def _base_url() -> str: + app = Config.get_section("app") + host = app.get("host", "127.0.0.1") + port = app.get("port", "8000") + if host in ("0.0.0.0", "::"): + host = "127.0.0.1" + return f"http://{host}:{port}" + + +def _post(client: httpx.Client, url: str, payload: Dict[str, Any]) -> None: + resp = client.post(url, json=payload) + print(f"POST {url} -> {resp.status_code}") + print(resp.text) + + +def _put(client: httpx.Client, url: str, payload: Dict[str, Any]) -> None: + resp = client.put(url, json=payload) + print(f"PUT {url} -> {resp.status_code}") + print(resp.text) + + +def _get(client: httpx.Client, url: str) -> None: + resp = client.get(url) + print(f"GET {url} -> {resp.status_code}") + print(resp.text) + + +def _stream_sse(client: httpx.Client, url: str, payload: Dict[str, Any]) -> None: + with client.stream("POST", url, json=payload) as resp: + print(f"POST {url} -> {resp.status_code}") + current_event = "message" + for raw in resp.iter_lines(): + if raw is None: + continue + line = raw.strip() + if not line: + continue + if line.startswith("event:"): + current_event = line.split(":", 1)[1].strip() or "message" + continue + if line.startswith("data:"): + data = line.split(":", 1)[1].strip() + print(f"[{current_event}] {data}") + + +def main() -> None: + base = _base_url() + menu: List[str] = [ + "1) GET /health", + "2) GET /nacos/status", + "3) POST /api/workflows (conversation)", + "4) POST /api/workflows/stream (conversation)", + "5) POST /api/sql/generate", + "6) POST /api/tools/execute", + "7) POST /api/prompts/reload", + "8) POST /api/ragflow/table-retrieval/upload", + "9) PUT /api/ragflow/table-retrieval/update", + "10) POST /api/ragflow/sql-gen/upload", + "11) PUT /api/ragflow/sql-gen/update", + "0) Exit", + ] + + with httpx.Client(timeout=60) as client: + while True: + print("\n可用接口:") + for line in menu: + print(line) + + choice = input("\n请选择编号: ").strip() + if choice == "0": + break + + if choice == "1": + _get(client, f"{base}/health") + elif choice == "2": + _get(client, f"{base}/nacos/status") + elif choice == "3": + payload = { + "input": "查询 SO 4020438779 的 eta 信息", + "session_id": None, + "workflow_type": "conversation", + } + _post(client, f"{base}/api/workflows", payload) + elif choice == "4": + payload = { + "input": "查询 SO 4016769041 的 eta 信息", + "session_id": None, + "workflow_type": "conversation", + } + _stream_sse(client, f"{base}/api/workflows/stream", payload) + elif choice == "5": + payload = { + "input": "查询 SO 4020438779 的 eta 信息", + "session_id": None, + "workflow_type": "conversation", + } + _post(client, f"{base}/api/sql/generate", payload) + elif choice == "6": + payload = { + "tool_name": "sr_api_query", + "payload": { + "sql": "SELECT 1", + "page": 1, + "rows": 1, + "orderBySelect": True, + "timeout": 30, + }, + } + _post(client, f"{base}/api/tools/execute", payload) + elif choice == "7": + _post(client, f"{base}/api/prompts/reload", {}) + elif choice == "8": + _post(client, f"{base}/api/ragflow/table-retrieval/upload", {}) + elif choice == "9": + cfg = {"name": "table_retrieval_dataset"} + _put(client, f"{base}/api/ragflow/table-retrieval/update", cfg) + elif choice == "10": + _post(client, f"{base}/api/ragflow/sql-gen/upload", {}) + elif choice == "11": + cfg = {"name": "sql_gen_dataset"} + _put(client, f"{base}/api/ragflow/sql-gen/update", cfg) + else: + print("无效选择") + + +if __name__ == "__main__": + try: + main() + except KeyboardInterrupt: + sys.exit(0) diff --git a/services/app_errors.py b/services/app_errors.py new file mode 100644 index 0000000..f4b3f2f --- /dev/null +++ b/services/app_errors.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +from dataclasses import dataclass +from enum import Enum +from typing import Any, Dict, Optional + + +class ErrorCode(str, Enum): + INVALID_WORKFLOW_TYPE = "INVALID_WORKFLOW_TYPE" + SQL_GENERATION_FAILED = "SQL_GENERATION_FAILED" + TABLE_MATCH_FAILED = "TABLE_MATCH_FAILED" + SQL_EXECUTION_FAILED = "SQL_EXECUTION_FAILED" + RAGFLOW_RETRIEVE_FAILED = "RAGFLOW_RETRIEVE_FAILED" + CONFIG_INVALID = "CONFIG_INVALID" + INTERNAL_ERROR = "INTERNAL_ERROR" + + +@dataclass +class AppError(Exception): + code: ErrorCode + message: str + status_code: int = 500 + detail: Optional[Dict[str, Any]] = None + + def to_dict(self) -> Dict[str, Any]: + return { + "code": self.code.value, + "message": self.message, + "detail": self.detail or {}, + } diff --git a/services/cache.py b/services/cache.py index ce2fddb..27f9515 100644 --- a/services/cache.py +++ b/services/cache.py @@ -2,6 +2,11 @@ from __future__ import annotations from typing import Optional +try: + import redis +except Exception: + redis = None + class CacheBase: """缓存接口""" @@ -23,3 +28,18 @@ class NoopCache(CacheBase): 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) + + diff --git a/services/ragflow_client.py b/services/ragflow_client.py index f875a90..51d4bd9 100644 --- a/services/ragflow_client.py +++ b/services/ragflow_client.py @@ -20,6 +20,15 @@ class RagflowClient: def _build_url(self) -> str: return self._base_url.rstrip("/") + "/" + self._retrieval_path.lstrip("/") + @staticmethod + def _normalize_dataset_ids(dataset_id: Optional[str]) -> list[str]: + """将配置值规范化为 RAGFlow 需要的 list[string]""" + if not dataset_id: + return [] + # 兼容逗号分隔配置 + parts = [p.strip() for p in str(dataset_id).split(",") if p.strip()] + return parts + 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: @@ -29,10 +38,12 @@ class RagflowClient: url = self._build_url() headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {} payload = { - "dataset_ids": dataset_id or "", - "query": query, + "dataset_ids": self._normalize_dataset_ids(dataset_id), + "question": query, "top_k": top_k, } + # 兼容部分版本字段 + payload["query"] = query if document_ids: payload["document_ids"] = document_ids @@ -57,6 +68,17 @@ def extract_table_name(record: Dict[str, Any]) -> Optional[str]: return record.get(key) content = record.get("content") or record.get("text") or "" + + # 兼容 content 为 JSON 字符串:{"table":"xxx", ...} + try: + parsed = json.loads(str(content)) + if isinstance(parsed, dict): + for key in ("table", "table_name"): + if parsed.get(key): + return str(parsed.get(key)) + except Exception: + pass + for line in str(content).splitlines(): if line.lower().startswith("table:"): return line.split(":", 1)[1].strip() diff --git a/services/ragflow_sync.py b/services/ragflow_sync.py index 3a62c91..4322ed7 100644 --- a/services/ragflow_sync.py +++ b/services/ragflow_sync.py @@ -7,45 +7,18 @@ import httpx from config import Config -def _build_document_for_table(table: str, templates: List[str]) -> str: - """构建表名检索文档 - 使用更标准的格式""" - lines = [ - f"# 表名检索模板: {table}", - "", - "## 可用模板:", - "" - ] - for i, t in enumerate(templates, 1): - lines.append(f"{i}. {t}") - lines.extend(["", f"表名: {table}", "类型: 表名检索模板"]) - return "\n".join(lines) +def _dump_json_content(data: Dict[str, Any]) -> str: + return json.dumps(data, ensure_ascii=False, indent=2) -def _build_sql_gen_document(table: str, prompt: Dict[str, any]) -> str: - """构建 SQL 生成文档 - 使用更标准的格式""" - system_prompt = prompt.get("system_prompt", "") - business_prompt = prompt.get("business_prompt", "") - constraints = prompt.get("constraints", []) - - lines = [ - f"# SQL 生成提示词: {table}", - "", - "## 系统提示词:", - system_prompt, - "", - "## 业务提示词:", - business_prompt, - "" - ] - - if constraints: - lines.extend(["## 约束条件:", ""]) - for i, c in enumerate(constraints, 1): - lines.append(f"{i}. {c}") - lines.append("") - - lines.extend([f"表名: {table}", "类型: SQL 生成提示词"]) - return "\n".join(lines) +def _extract_tables_map(data: Dict[str, Any]) -> Dict[str, Any]: + """兼容两种结构:{"tables": {...}} 或直接 {...}""" + tables = data.get("tables") if isinstance(data, dict) else None + if isinstance(tables, dict): + return tables + if isinstance(data, dict): + return data + return {} class RagflowSync: @@ -58,143 +31,271 @@ class RagflowSync: 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() - 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 "{dataset_id}" not in self._upload_path: - raise RuntimeError("上传接口路径必须包含 {dataset_id} 占位符") - if self._upload_mode not in ("overwrite", "append"): - raise RuntimeError("ragflow.upload_mode 仅支持 overwrite 或 append") - - def _post(self, documents: List[Dict[str, Any]], dataset_id: str): - """上传文档到指定知识库 - 使用 multipart/form-data 格式""" - if not self._base_url: - raise RuntimeError("未配置 ragflow.url") - - # 构建正确的 URL - upload_path = self._upload_path.replace("{dataset_id}", dataset_id) - url = self._base_url + "/" + upload_path.lstrip("/") - - headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {} - - # 由于接口使用 multipart/form-data,我们需要创建临时文件 - import tempfile - - # 创建临时文件并写入文档内容 - with tempfile.NamedTemporaryFile(mode='w', suffix='.txt', delete=False, encoding='utf-8') as f: - # 将文档内容写入文件 - for doc in documents: - content = doc.get('content', '') - f.write(content + '\n\n') - temp_file_path = f.name - - try: - # 使用 multipart/form-data 上传文件 - files = {'file': open(temp_file_path, 'rb')} - - print(f"请求 URL: {url}") # 调试信息 - print(f"上传文件: {temp_file_path}") # 调试信息 - - with httpx.Client(timeout=60) as client: - response = client.post(url, files=files, headers=headers) - response.raise_for_status() - result = response.json() - print(f"RAGFlow 上传响应: {result}") # 调试信息 - return result - finally: - # 清理临时文件 - import os - if os.path.exists(temp_file_path): - os.unlink(temp_file_path) - def upload_documents(self, dataset_id: str, documents: List[Dict[str, Any]]) -> Dict[str, Any]: - """上传文档到指定知识库 - 使用 multipart/form-data 格式 - - 根据官方文档: POST /api/v1/datasets/{dataset_id}/documents - """ + """上传文档到指定知识库(每个文档单独上传)""" if not self._base_url: raise RuntimeError("未配置 ragflow.url") + if not dataset_id: + raise RuntimeError("dataset_id 为空,无法上传文档") + if not documents: + raise RuntimeError("没有可上传的文档内容") # 构建正确的 URL url = f"{self._base_url}/api/v1/datasets/{dataset_id}/documents" headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {} - - # 由于接口使用 multipart/form-data,我们需要创建临时文件 - import tempfile - import os - - # 创建临时文件并写入文档内容 - with tempfile.NamedTemporaryFile(mode='w', suffix='.txt', delete=False, encoding='utf-8') as f: - # 将文档内容写入文件 - for doc in documents: - content = doc.get('content', '') - f.write(content + '\n\n') - temp_file_path = f.name - - try: - # 使用 multipart/form-data 上传文件 - # 确保文件在 with 块内打开和关闭 - with open(temp_file_path, 'rb') as file_obj: - files = {'file': file_obj} - - print(f"上传文档 URL: {url}") # 调试信息 - print(f"上传文件: {temp_file_path}") # 调试信息 - - with httpx.Client(timeout=60) as client: - response = client.post(url, files=files, headers=headers) + + results: List[Dict[str, Any]] = [] + with httpx.Client(timeout=60) as client: + for idx, doc in enumerate(documents, start=1): + content = str(doc.get("content", "")) + filename = str(doc.get("filename") or f"doc_{idx}.txt") + files = {"file": (filename, content.encode("utf-8"), "text/plain")} + + print(f"上传文档 URL: {url}") + print(f"上传文件名: {filename}") + + response = client.post(url, files=files, headers=headers) response.raise_for_status() result = response.json() - print(f"RAGFlow 上传响应: {result}") # 调试信息 - - # 检查文档处理状态 - if result.get('code') == 0 and result.get('data'): - doc_id = result['data'][0].get('id') - if doc_id: - print(f"文档已上传,ID: {doc_id}") - print("注意: 文档处理需要时间,请等待 RAGFlow 完成分块处理") - print("可以在 RAGFlow 界面查看处理进度") - - return result - finally: - # 清理临时文件 - if os.path.exists(temp_file_path): - try: - os.unlink(temp_file_path) - except PermissionError: - # 如果文件被占用,等待一下再重试 - import time - time.sleep(0.1) - try: - os.unlink(temp_file_path) - except PermissionError: - print(f"警告: 无法删除临时文件 {temp_file_path}") + print(f"RAGFlow 上传响应: {result}") + results.append(result) - def update_dataset(self, dataset_id: str, config: Dict[str, Any]) -> Dict[str, Any]: - """更新知识库配置 - - 根据官方文档: PUT /api/v1/datasets/{dataset_id} - """ - if not self._base_url: - raise RuntimeError("未配置 ragflow.url") - + dataset_detail = self._get_dataset_detail(dataset_id) + chunk_method = self._extract_chunk_method(dataset_detail) + if chunk_method is None: + chunk_method = self._extract_chunk_method_from_upload_results(results) + # 参考 Java 实现:查询知识库文档 ID 后统一调用 chunks 解析 + doc_ids = self._list_document_ids(dataset_id) + parse_results = self._auto_parse_documents(dataset_id, doc_ids) + upload_status = self._build_parse_status_from_upload_results(results) + + return { + "ok": True, + "count": len(results), + "results": results, + "chunk_method": chunk_method, + "upload_status": upload_status, + "parse": parse_results, + } + + def _get_dataset_detail(self, dataset_id: str) -> Dict[str, Any]: + """查询知识库详情(用于读取 chunk_method)""" url = f"{self._base_url}/api/v1/datasets/{dataset_id}" + headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {} + with httpx.Client(timeout=30) as client: + resp = client.get(url, headers=headers) + resp.raise_for_status() + return resp.json() + + @staticmethod + def _extract_chunk_method(dataset_detail: Dict[str, Any]) -> Any: + """从知识库详情提取 chunk_method""" + data = dataset_detail.get("data") + if isinstance(data, dict): + if "chunk_method" in data: + return data.get("chunk_method") + parser_cfg = data.get("parser_config") or {} + if isinstance(parser_cfg, dict): + return parser_cfg.get("chunk_method") + return None + + @staticmethod + def _extract_chunk_method_from_upload_results(upload_results: List[Dict[str, Any]]) -> Any: + """从上传响应中提取 chunk_method(兼容不同版本返回结构)""" + for item in upload_results: + data = item.get("data") + records = data if isinstance(data, list) else [data] if isinstance(data, dict) else [] + for rec in records: + if not isinstance(rec, dict): + continue + if rec.get("chunk_method"): + return rec.get("chunk_method") + parser_cfg = rec.get("parser_config") or {} + if isinstance(parser_cfg, dict) and parser_cfg.get("chunk_method"): + return parser_cfg.get("chunk_method") + return None + + @staticmethod + def _extract_uploaded_doc_ids(upload_results: List[Dict[str, Any]]) -> List[str]: + """从上传结果中提取文档 ID""" + ids: List[str] = [] + for item in upload_results: + data = item.get("data") + if isinstance(data, list): + for d in data: + if isinstance(d, dict) and d.get("id"): + ids.append(str(d.get("id"))) + elif isinstance(data, dict) and data.get("id"): + ids.append(str(data.get("id"))) + return ids + + @staticmethod + def _build_parse_status_from_upload_results(upload_results: List[Dict[str, Any]]) -> Dict[str, Any]: + """根据上传返回构造解析状态(上传接口已触发解析,无需额外 parse API)""" + details: List[Dict[str, Any]] = [] + for item in upload_results: + data = item.get("data") + records = data if isinstance(data, list) else [data] if isinstance(data, dict) else [] + for rec in records: + if not isinstance(rec, dict): + continue + details.append( + { + "doc_id": rec.get("id"), + "name": rec.get("name") or rec.get("location"), + "run": rec.get("run"), + "chunk_method": rec.get("chunk_method") + or (rec.get("parser_config") or {}).get("chunk_method"), + } + ) + return { + "ok": True, + "trigger": "upload_endpoint", + "message": "文档上传接口已触发解析流程,无需单独调用 parse API", + "count": len(details), + "details": details, + } + + def _auto_parse_documents(self, dataset_id: str, doc_ids: List[str]) -> Dict[str, Any]: + """调用官方 chunks 接口触发解析""" + if not doc_ids: + return {"ok": False, "message": "未提取到文档ID,无法触发解析", "count": 0, "details": []} + + url = f"{self._base_url}/api/v1/datasets/{dataset_id}/chunks" headers = { "Content-Type": "application/json", - "Authorization": f"Bearer {self._api_key}" if self._api_key else "" - } - - print(f"更新知识库 URL: {url}") # 调试信息 - print(f"更新配置: {config}") # 调试信息 - + "Authorization": f"Bearer {self._api_key}", + } if self._api_key else {"Content-Type": "application/json"} + payload = {"document_ids": doc_ids} + with httpx.Client(timeout=60) as client: - response = client.put(url, json=config, headers=headers) - response.raise_for_status() - result = response.json() - print(f"RAGFlow 更新响应: {result}") # 调试信息 - return result + resp = client.post(url, headers=headers, json=payload) + + if resp.status_code >= 400: + return { + "ok": False, + "trigger": "chunks_api", + "status": resp.status_code, + "message": resp.text, + "count": len(doc_ids), + "details": [{"doc_id": d} for d in doc_ids], + } + + body: Any + try: + body = resp.json() + except Exception: + body = resp.text + + return { + "ok": True, + "trigger": "chunks_api", + "count": len(doc_ids), + "details": [{"doc_id": d} for d in doc_ids], + "response": body, + } + + def _list_document_ids(self, dataset_id: str) -> List[str]: + """获取知识库中的全部文档 ID(用于覆盖更新)""" + if not self._base_url: + raise RuntimeError("未配置 ragflow.url") + if not dataset_id: + raise RuntimeError("dataset_id 为空,无法查询文档") + + url = f"{self._base_url}/api/v1/datasets/{dataset_id}/documents" + headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {} + + ids: List[str] = [] + page = 1 + page_size = 100 + + with httpx.Client(timeout=60) as client: + while True: + resp = client.get(url, headers=headers, params={"page": page, "page_size": page_size}) + resp.raise_for_status() + body = resp.json() + data = body.get("data") + if isinstance(data, dict): + docs = data.get("docs") or data.get("list") or [] + elif isinstance(data, list): + docs = data + else: + docs = [] + + if not docs: + break + + for item in docs: + if isinstance(item, dict) and item.get("id"): + ids.append(str(item.get("id"))) + + if len(docs) < page_size: + break + page += 1 + + return ids + + def _delete_documents(self, dataset_id: str, doc_ids: List[str]) -> Dict[str, Any]: + """按 ID 删除文档""" + if not doc_ids: + return {"ok": True, "deleted": 0} + + url = f"{self._base_url}/api/v1/datasets/{dataset_id}/documents" + headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {} + payload = {"ids": doc_ids} + + with httpx.Client(timeout=60) as client: + resp = client.request("DELETE", url, headers=headers, json=payload) + resp.raise_for_status() + return resp.json() + + def replace_documents(self, dataset_id: str, documents: List[Dict[str, Any]]) -> Dict[str, Any]: + """覆盖更新:先删后传,避免“update 变新增”""" + ids = self._list_document_ids(dataset_id) + if ids: + self._delete_documents(dataset_id, ids) + return self.upload_documents(dataset_id, documents) + + def update_table_retrieval_documents(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 = _extract_tables_map(data) + documents = [ + { + "filename": f"{k}.txt", + "content": _dump_json_content({"table": k, "templates": v}), + } + for k, v in tables.items() + ] + return self.replace_documents(self._table_retrieval_dataset_id, documents) + + def update_sql_gen_documents(self) -> Dict[str, Any]: + """更新 SQL 生成文档(仅文档内容)""" + 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: List[Dict[str, Any]] = [] + + 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] + documents.append({"filename": f"{table}.txt", "content": _dump_json_content(prompt)}) + + return self.replace_documents(self._sql_gen_dataset_id, documents) def upload_table_retrieval(self) -> Dict[str, Any]: """上传表名检索模板文档 - 直接上传整个 JSON 文件""" @@ -210,22 +311,17 @@ class RagflowSync: # 读取整个 JSON 文件内容 with open(tables_file, "r", encoding="utf-8") as f: data = json.load(f) - - # 将 JSON 内容转换为字符串 - json_content = json.dumps(data, ensure_ascii=False, indent=2) - - # 构建文档 - doc = { - "content": f"# 表名检索模板库\n\n以下是所有表名检索模板的 JSON 数据:\n\n```json\n{json_content}\n```\n\n包含的表:{list(data.keys())}", - "metadata": {"type": "table_retrieval_templates", "format": "json"}, - "title": "表名检索模板库", - "type": "table_template_library" - } - - print(f"生成的表名检索文档: {doc}") - - # 上传整个 JSON 文件内容 - return self.upload_documents(self._table_retrieval_dataset_id, [doc]) + + tables = _extract_tables_map(data) + # 每个 key 一个文档,配合 One 解析时每个表单独成块 + documents = [ + { + "filename": f"{k}.txt", + "content": _dump_json_content({"table": k, "templates": v}), + } + for k, v in tables.items() + ] + return self.upload_documents(self._table_retrieval_dataset_id, documents) def upload_sql_gen(self) -> Dict[str, Any]: """上传 SQL 生成提示词文档""" @@ -247,13 +343,7 @@ class RagflowSync: with open(path, "r", encoding="utf-8") as f: prompt = json.load(f) table = prompt.get("table") or os.path.splitext(name)[0] - doc = { - "content": _build_sql_gen_document(table, prompt), - "metadata": {"table": table}, - "title": f"SQL Prompt: {table}", - "type": "sql_prompt" - } - documents.append(doc) - print(f"生成的 SQL 提示词文档: {doc}") + json_content = _dump_json_content(prompt) + documents.append({"filename": f"{table}.txt", "content": json_content}) return self.upload_documents(self._sql_gen_dataset_id, documents) diff --git a/services/sql_prompt_manager.py b/services/sql_prompt_manager.py index 3fb09af..62dbda1 100644 --- a/services/sql_prompt_manager.py +++ b/services/sql_prompt_manager.py @@ -2,6 +2,9 @@ import json import os from typing import Any, Dict, Optional +from config import Config +from services.cache import NoopCache, RedisCache + class SqlPromptManager: """按表名读取 SQL 提示词""" @@ -9,18 +12,83 @@ class SqlPromptManager: 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") + self._cache = self._init_cache() + self._cache_ttl = self._get_cache_ttl() + + @staticmethod + def _get_cache_ttl() -> int: + redis_cfg = Config.get_section("redis") + try: + return int(redis_cfg.get("sql_prompt_ttl", 600)) + except Exception: + return 600 + + @staticmethod + def _init_cache(): + redis_cfg = Config.get_section("redis") + enabled = str(redis_cfg.get("enabled", "false")).lower() in ("1", "true", "yes") + if not enabled: + return NoopCache() + + # 优先使用完整 URL;否则使用 host/port/password/database 拼接 + url = redis_cfg.get("url") + db = int(redis_cfg.get("db", redis_cfg.get("database", 0))) + if not url: + host = redis_cfg.get("host") + port = redis_cfg.get("port", "6379") + password = redis_cfg.get("password", "") + database = redis_cfg.get("database", str(db)) + if host: + auth = f":{password}@" if password else "" + url = f"redis://{auth}{host}:{port}/{database}" + + if not url: + return NoopCache() + try: + return RedisCache(url=url, db=db) + except Exception: + return NoopCache() @staticmethod def _safe_filename(name: str) -> str: return name.replace("..", "").replace("/", "_").replace("\\", "_") + @staticmethod + def _cache_key(table_name: str, mtime: float) -> str: + return f"sql_prompt:{table_name}:{int(mtime)}" + 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" + safe_name = self._safe_filename(table_name) + filename = safe_name + ".json" path = os.path.join(self._base_dir, filename) if not os.path.exists(path): return None + + mtime = os.path.getmtime(path) + key = self._cache_key(safe_name, mtime) + cached = self._cache.get(key) + if cached: + try: + return json.loads(cached) + except Exception: + pass + with open(path, "r", encoding="utf-8") as f: - return json.load(f) + prompt = json.load(f) + + self._cache.set(key, json.dumps(prompt, ensure_ascii=False), self._cache_ttl) + return prompt + + +_GLOBAL_SQL_PROMPT_MANAGER: Optional[SqlPromptManager] = None + + +def get_sql_prompt_manager(base_dir: Optional[str] = None) -> SqlPromptManager: + """获取全局 SqlPromptManager(单例)""" + global _GLOBAL_SQL_PROMPT_MANAGER + if _GLOBAL_SQL_PROMPT_MANAGER is None: + _GLOBAL_SQL_PROMPT_MANAGER = SqlPromptManager(base_dir=base_dir) + return _GLOBAL_SQL_PROMPT_MANAGER diff --git a/services/structured_logger.py b/services/structured_logger.py new file mode 100644 index 0000000..96ba37e --- /dev/null +++ b/services/structured_logger.py @@ -0,0 +1,104 @@ +from __future__ import annotations + +import json +from datetime import datetime +from typing import Any, Dict, Optional + +import pymysql + +from config import Config + + +class StructuredLogger: + def __init__(self): + cfg = Config.get_section("logging_mysql") + self.enabled = str(cfg.get("enabled", "false")).lower() in ("1", "true", "yes") + self.host = cfg.get("host", "127.0.0.1") + self.port = int(cfg.get("port", 3306)) + self.user = cfg.get("user", "root") + self.password = cfg.get("password", "") + self.database = cfg.get("database", "more_dots") + self.table = cfg.get("table", "structured_logs") + self.connect_timeout = int(cfg.get("connect_timeout", 5)) + self._inited = False + + def _get_conn(self): + return pymysql.connect( + host=self.host, + port=self.port, + user=self.user, + password=self.password, + database=self.database, + charset="utf8mb4", + autocommit=True, + connect_timeout=self.connect_timeout, + ) + + def _ensure_table(self) -> None: + if self._inited or not self.enabled: + return + sql = f""" + CREATE TABLE IF NOT EXISTS {self.table} ( + id BIGINT PRIMARY KEY AUTO_INCREMENT, + trace_id VARCHAR(64) NOT NULL, + level VARCHAR(16) NOT NULL, + event VARCHAR(128) NOT NULL, + error_code VARCHAR(64) NULL, + payload JSON NULL, + created_at DATETIME NOT NULL + ) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4; + """ + try: + with self._get_conn() as conn: + with conn.cursor() as cur: + cur.execute(sql) + self._inited = True + except Exception: + # 开发阶段容错,避免日志失败影响主流程 + self.enabled = False + + def log(self, level: str, event: str, trace_id: str, payload: Optional[Dict[str, Any]] = None, error_code: Optional[str] = None) -> None: + print(json.dumps({ + "trace_id": trace_id, + "level": level, + "event": event, + "error_code": error_code, + "payload": payload or {}, + "created_at": datetime.now().isoformat(), + }, ensure_ascii=False)) + + if not self.enabled: + return + + self._ensure_table() + if not self.enabled: + return + + insert_sql = f"INSERT INTO {self.table}(trace_id, level, event, error_code, payload, created_at) VALUES(%s,%s,%s,%s,%s,%s)" + try: + with self._get_conn() as conn: + with conn.cursor() as cur: + cur.execute( + insert_sql, + ( + trace_id, + level, + event, + error_code, + json.dumps(payload or {}, ensure_ascii=False), + datetime.now(), + ), + ) + except Exception: + # 开发阶段容错,避免日志失败影响主流程 + return + + +_GLOBAL_STRUCTURED_LOGGER: Optional[StructuredLogger] = None + + +def get_structured_logger() -> StructuredLogger: + global _GLOBAL_STRUCTURED_LOGGER + if _GLOBAL_STRUCTURED_LOGGER is None: + _GLOBAL_STRUCTURED_LOGGER = StructuredLogger() + return _GLOBAL_STRUCTURED_LOGGER diff --git a/services/template_matcher.py b/services/template_matcher.py index 4054752..7361ef7 100644 --- a/services/template_matcher.py +++ b/services/template_matcher.py @@ -11,6 +11,7 @@ class TemplateMatcher: self._ragflow = RagflowClient() cfg = Config.get_section("ragflow") self._dataset_id = (cfg.get("table_retrieval_dataset_id") or "").strip() + self._top_k = int(cfg.get("retrieval_top_k", 3)) def _validate(self) -> None: if not self._dataset_id: @@ -20,17 +21,24 @@ class TemplateMatcher: """返回匹配的表名与原始响应""" self._validate() try: - response = self._ragflow.retrieve(normalized_text, top_k=3, dataset_id=self._dataset_id) + response = self._ragflow.retrieve(normalized_text, top_k=self._top_k, dataset_id=self._dataset_id) except Exception as e: return {"table_name": None, "raw": {"error": str(e)}} candidates = [] data = response.get("data") if isinstance(response, dict) else None + records = [] if isinstance(data, list): - for item in data: - table_name = extract_table_name(item) - if table_name: - candidates.append(table_name) + records = data + elif isinstance(data, dict): + chunks = data.get("chunks") + if isinstance(chunks, list): + records = chunks + + for item in records: + table_name = extract_table_name(item) + if table_name: + candidates.append(table_name) matched = candidates[0] if candidates else None return {"table_name": matched, "raw": response}