import json import os from typing import Any, Dict, List from pathlib import Path import httpx from config import Config def _dump_json_content(data: Dict[str, Any]) -> str: return json.dumps(data, ensure_ascii=False, indent=2) 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: """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() @staticmethod def _project_root() -> Path: """返回项目根目录。""" return Path(__file__).resolve().parents[2] def _collect_sql_gen_documents(self) -> tuple[List[Dict[str, Any]], List[Dict[str, str]]]: """收集 SQL 生成文档,并跳过空文件/非法 JSON。""" prompts_dir = self._project_root() / "config" / "sql_gen_prompts" documents: List[Dict[str, Any]] = [] warnings: List[Dict[str, str]] = [] for name in os.listdir(prompts_dir): if not name.endswith(".json"): continue path = prompts_dir / name raw_text = path.read_text(encoding="utf-8") if not raw_text.strip(): warnings.append({"file": name, "reason": "empty_file"}) continue try: prompt = json.loads(raw_text) except json.JSONDecodeError as exc: warnings.append({"file": name, "reason": f"invalid_json:{exc}"}) continue table = prompt.get("table") or path.stem documents.append({"filename": f"{table}.txt", "content": _dump_json_content(prompt)}) return documents, warnings def upload_documents(self, dataset_id: str, documents: List[Dict[str, Any]]) -> Dict[str, Any]: """上传文档到指定知识库(每个文档单独上传)""" 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 {} 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}") results.append(result) 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 {"Content-Type": "application/json"} payload = {"document_ids": doc_ids} with httpx.Client(timeout=60) as client: 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,无法更新表名检索文档") tables_file = self._project_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 sync_table_retrieval(self) -> Dict[str, Any]: """兼容旧脚本:同步表名检索文档,采用覆盖更新避免旧表残留。""" return self.update_table_retrieval_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 生成文档") documents, warnings = self._collect_sql_gen_documents() if not documents: raise RuntimeError(f"SQL 生成提示词目录中没有可同步的有效 JSON 文档,warnings={warnings}") result = self.replace_documents(self._sql_gen_dataset_id, documents) result["warnings"] = warnings result["valid_document_count"] = len(documents) return result def sync_sql_gen_prompts(self) -> Dict[str, Any]: """兼容旧脚本:同步 SQL 生成提示词文档,采用覆盖更新避免旧 prompt 残留。""" return self.update_sql_gen_documents() def upload_table_retrieval(self) -> Dict[str, Any]: """上传表名检索模板文档 - 直接上传整个 JSON 文件""" if not self._table_retrieval_dataset_id: raise RuntimeError("未配置 ragflow.table_retrieval_dataset_id,无法上传表名检索模板") tables_file = self._project_root() / "config" / "table_retrieval_prompts" / "tables.json" if not os.path.exists(tables_file): raise RuntimeError(f"表名检索模板文件不存在: {tables_file}") # 读取整个 JSON 文件内容 with open(tables_file, "r", encoding="utf-8") as f: data = json.load(f) 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 生成提示词文档""" if not self._sql_gen_dataset_id: raise RuntimeError("未配置 ragflow.sql_gen_dataset_id,无法上传 SQL 生成提示词") prompts_dir = self._project_root() / "config" / "sql_gen_prompts" if not os.path.exists(prompts_dir): raise RuntimeError(f"SQL 生成提示词目录不存在: {prompts_dir}") documents, warnings = self._collect_sql_gen_documents() if not documents: raise RuntimeError(f"SQL 生成提示词目录中没有可上传的有效 JSON 文档,warnings={warnings}") result = self.upload_documents(self._sql_gen_dataset_id, documents) result["warnings"] = warnings result["valid_document_count"] = len(documents) return result