init
This commit is contained in:
@@ -0,0 +1,24 @@
|
||||
"""外部集成模块"""
|
||||
|
||||
from .ragflow_client import RagflowClient, extract_table_name
|
||||
from .ragflow_sync import RagflowSync
|
||||
|
||||
# 延迟导入 nacos(可选依赖)
|
||||
try:
|
||||
from .nacos_service import NacosManager, NacosConfig, ServiceConfig, load_nacos_config, load_service_config
|
||||
__all__ = [
|
||||
"RagflowClient",
|
||||
"extract_table_name",
|
||||
"RagflowSync",
|
||||
"NacosManager",
|
||||
"NacosConfig",
|
||||
"ServiceConfig",
|
||||
"load_nacos_config",
|
||||
"load_service_config",
|
||||
]
|
||||
except ImportError:
|
||||
__all__ = [
|
||||
"RagflowClient",
|
||||
"extract_table_name",
|
||||
"RagflowSync",
|
||||
]
|
||||
@@ -0,0 +1,279 @@
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import socket
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
try:
|
||||
import nacos # type: ignore
|
||||
except ImportError:
|
||||
nacos = None
|
||||
|
||||
from config import Config
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class NacosConfig:
|
||||
"""Nacos 配置"""
|
||||
enabled: bool
|
||||
server_addresses: str
|
||||
namespace: str
|
||||
group_name: str
|
||||
cluster_name: str
|
||||
username: Optional[str]
|
||||
password: Optional[str]
|
||||
heartbeat_interval: int
|
||||
weight: float
|
||||
ephemeral: bool
|
||||
register_port: Optional[int]
|
||||
|
||||
|
||||
@dataclass
|
||||
class ServiceConfig:
|
||||
"""服务配置"""
|
||||
service_name: str
|
||||
host: str
|
||||
port: int
|
||||
ip: str
|
||||
metadata: Dict[str, Any]
|
||||
|
||||
|
||||
def _get_local_ip() -> str:
|
||||
"""获取本地 IP 地址"""
|
||||
for env_name in ("POD_IP", "HOST_IP"):
|
||||
env_ip = os.getenv(env_name)
|
||||
if env_ip:
|
||||
return env_ip
|
||||
|
||||
try:
|
||||
s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
|
||||
s.connect(("8.8.8.8", 80))
|
||||
ip = s.getsockname()[0]
|
||||
s.close()
|
||||
return ip
|
||||
except Exception:
|
||||
try:
|
||||
host_ip = socket.gethostbyname(socket.gethostname())
|
||||
if host_ip and host_ip != "127.0.0.1":
|
||||
return host_ip
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
logger.warning("获取本地 IP 失败,使用 127.0.0.1")
|
||||
return "127.0.0.1"
|
||||
|
||||
|
||||
def load_nacos_config() -> NacosConfig:
|
||||
"""从 config.ini 读取 Nacos 配置"""
|
||||
section = "nacos"
|
||||
enabled = Config._config.getboolean(section, "enabled", fallback=False)
|
||||
register_port = Config._config.getint(section, "register_port", fallback=0)
|
||||
return NacosConfig(
|
||||
enabled=enabled,
|
||||
server_addresses=Config._config.get(section, "server", fallback="localhost:8848"),
|
||||
namespace=Config._config.get(section, "namespace", fallback="public"),
|
||||
group_name=Config._config.get(section, "group_name", fallback="DEFAULT_GROUP"),
|
||||
cluster_name=Config._config.get(section, "cluster_name", fallback="DEFAULT"),
|
||||
username=Config._config.get(section, "username", fallback="") or None,
|
||||
password=Config._config.get(section, "password", fallback="") or None,
|
||||
heartbeat_interval=Config._config.getint(section, "heartbeat_interval", fallback=5),
|
||||
weight=Config._config.getfloat(section, "weight", fallback=1.0),
|
||||
ephemeral=Config._config.getboolean(section, "ephemeral", fallback=True),
|
||||
register_port=register_port if register_port > 0 else None,
|
||||
)
|
||||
|
||||
|
||||
def load_service_config() -> ServiceConfig:
|
||||
"""从 config.ini 读取服务配置"""
|
||||
section = "app"
|
||||
service_name = Config._config.get(section, "service_name", fallback="apbo-boat-agent")
|
||||
host = Config._config.get(section, "host", fallback="0.0.0.0")
|
||||
port = Config._config.getint(section, "port", fallback=8000)
|
||||
ip = host if host not in ("0.0.0.0", "::") else _get_local_ip()
|
||||
|
||||
metadata = {
|
||||
"version": Config._config.get(section, "version", fallback="1.0.0"),
|
||||
"service_type": "fastapi",
|
||||
"api_paths": "/health,/api/workflows,/api/workflows/stream,/nacos/status",
|
||||
"streaming": "true",
|
||||
"model_section": Config._config.get(section, "model_section", fallback=Config.DEFAULT_MODEL_SECTION),
|
||||
}
|
||||
|
||||
extra_meta = Config.get_section("metadata")
|
||||
if extra_meta:
|
||||
metadata.update(extra_meta)
|
||||
|
||||
return ServiceConfig(
|
||||
service_name=service_name,
|
||||
host=host,
|
||||
port=port,
|
||||
ip=ip,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
|
||||
class NacosManager:
|
||||
"""Nacos 服务注册与心跳管理"""
|
||||
|
||||
def __init__(self, nacos_config: NacosConfig, service_config: ServiceConfig):
|
||||
self.nacos_config = nacos_config
|
||||
self.service_config = service_config
|
||||
self.client = None
|
||||
self._heartbeat_task: Optional[asyncio.Task] = None
|
||||
self._stop_event = asyncio.Event()
|
||||
self.is_registered = False
|
||||
|
||||
def _registration_port(self) -> int:
|
||||
"""Nacos 注册端口:优先使用 nacos.register_port,未配置时回退 app.port。"""
|
||||
return int(self.nacos_config.register_port or self.service_config.port)
|
||||
|
||||
def _init_client(self) -> bool:
|
||||
"""初始化 Nacos 客户端"""
|
||||
if nacos is None:
|
||||
raise ImportError("未安装 nacos-sdk-python,请先安装依赖")
|
||||
|
||||
try:
|
||||
self.client = nacos.NacosClient(
|
||||
server_addresses=self.nacos_config.server_addresses,
|
||||
namespace=self.nacos_config.namespace,
|
||||
username=self.nacos_config.username,
|
||||
password=self.nacos_config.password,
|
||||
)
|
||||
logger.info(f"Nacos 客户端初始化成功: {self.nacos_config.server_addresses}")
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"Nacos 客户端初始化失败: {e}")
|
||||
return False
|
||||
|
||||
def register_service(self) -> bool:
|
||||
"""注册服务到 Nacos"""
|
||||
if not self.client and not self._init_client():
|
||||
return False
|
||||
|
||||
try:
|
||||
self.client.add_naming_instance(
|
||||
service_name=self.service_config.service_name,
|
||||
ip=self.service_config.ip,
|
||||
port=self._registration_port(),
|
||||
cluster_name=self.nacos_config.cluster_name,
|
||||
group_name=self.nacos_config.group_name,
|
||||
weight=self.nacos_config.weight,
|
||||
metadata=self.service_config.metadata,
|
||||
ephemeral=self.nacos_config.ephemeral,
|
||||
)
|
||||
self.is_registered = True
|
||||
logger.info(
|
||||
"✅ 服务注册成功: %s (%s:%s)",
|
||||
self.service_config.service_name,
|
||||
self.service_config.ip,
|
||||
self._registration_port(),
|
||||
)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"❌ 服务注册失败: {e}")
|
||||
self.is_registered = False
|
||||
return False
|
||||
|
||||
def deregister_service(self) -> bool:
|
||||
"""从 Nacos 注销服务"""
|
||||
if not self.client or not self.is_registered:
|
||||
return True
|
||||
|
||||
try:
|
||||
self.client.remove_naming_instance(
|
||||
service_name=self.service_config.service_name,
|
||||
ip=self.service_config.ip,
|
||||
port=self._registration_port(),
|
||||
cluster_name=self.nacos_config.cluster_name,
|
||||
group_name=self.nacos_config.group_name,
|
||||
)
|
||||
self.is_registered = False
|
||||
logger.info("✅ 服务注销成功: %s", self.service_config.service_name)
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.error(f"❌ 服务注销失败: {e}")
|
||||
return False
|
||||
|
||||
def _send_heartbeat(self) -> None:
|
||||
"""发送心跳"""
|
||||
if not self.client or not self.is_registered:
|
||||
return
|
||||
|
||||
self.client.send_heartbeat(
|
||||
service_name=self.service_config.service_name,
|
||||
ip=self.service_config.ip,
|
||||
port=self._registration_port(),
|
||||
cluster_name=self.nacos_config.cluster_name,
|
||||
group_name=self.nacos_config.group_name,
|
||||
)
|
||||
|
||||
async def _heartbeat_loop(self) -> None:
|
||||
"""心跳循环"""
|
||||
interval = max(1, self.nacos_config.heartbeat_interval)
|
||||
self._stop_event.clear()
|
||||
|
||||
while not self._stop_event.is_set():
|
||||
try:
|
||||
if not self.is_registered:
|
||||
if self.register_service():
|
||||
logger.info("✅ Nacos 重试注册成功: %s", self.service_config.service_name)
|
||||
else:
|
||||
logger.warning("⚠️ Nacos 注册重试失败: %s", self.service_config.service_name)
|
||||
else:
|
||||
self._send_heartbeat()
|
||||
logger.debug("心跳发送成功: %s", self.service_config.service_name)
|
||||
except Exception as e:
|
||||
logger.warning(f"心跳发送失败: {e}")
|
||||
# 尝试重新注册
|
||||
try:
|
||||
self.register_service()
|
||||
except Exception as re:
|
||||
logger.error(f"重新注册失败: {re}")
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(self._stop_event.wait(), timeout=interval)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
|
||||
async def start(self) -> None:
|
||||
"""启动注册与心跳"""
|
||||
if not self.nacos_config.enabled:
|
||||
logger.info("Nacos 未启用,跳过注册")
|
||||
return
|
||||
|
||||
if not self.register_service():
|
||||
logger.warning("⚠️ Nacos 注册失败,服务继续运行")
|
||||
|
||||
# 无论首次注册是否成功,都启动循环以便持续重试注册
|
||||
self._heartbeat_task = asyncio.create_task(self._heartbeat_loop())
|
||||
logger.info("✅ Nacos 心跳任务已启动")
|
||||
|
||||
async def stop(self) -> None:
|
||||
"""停止心跳并注销"""
|
||||
self._stop_event.set()
|
||||
|
||||
if self._heartbeat_task:
|
||||
self._heartbeat_task.cancel()
|
||||
try:
|
||||
await self._heartbeat_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
self.deregister_service()
|
||||
|
||||
def status(self) -> Dict[str, Any]:
|
||||
"""获取当前状态"""
|
||||
return {
|
||||
"service_name": self.service_config.service_name,
|
||||
"ip": self.service_config.ip,
|
||||
"port": self.service_config.port,
|
||||
"register_port": self._registration_port(),
|
||||
"namespace": self.nacos_config.namespace,
|
||||
"group": self.nacos_config.group_name,
|
||||
"cluster": self.nacos_config.cluster_name,
|
||||
"registered": self.is_registered,
|
||||
"heartbeat_running": self._heartbeat_task is not None and not self._heartbeat_task.done(),
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
import json
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from config import Config
|
||||
|
||||
|
||||
TABLE_NAME_ALIASES = {
|
||||
"apbo_tp_multiple_impact": "apbo_eta_multiple_impact",
|
||||
}
|
||||
|
||||
|
||||
class RagflowClient:
|
||||
"""RAGFlow 客户端(仅检索)"""
|
||||
|
||||
def __init__(self):
|
||||
cfg = Config.get_section("ragflow")
|
||||
self._base_url = cfg.get("url", "")
|
||||
self._api_key = cfg.get("api_key", "")
|
||||
self._retrieval_path = cfg.get("retrieval", "/api/v1/retrieval")
|
||||
self._table_retrieval_dataset_id = cfg.get("table_retrieval_dataset_id", "")
|
||||
self._sql_gen_dataset_id = cfg.get("sql_gen_dataset_id", "")
|
||||
|
||||
def _build_url(self) -> str:
|
||||
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:
|
||||
raise RuntimeError("未配置 ragflow.url 或 ragflow.retrieval")
|
||||
if not dataset_id:
|
||||
raise ValueError("未提供 ragflow.dataset_id,无法进行检索")
|
||||
url = self._build_url()
|
||||
headers = {"Authorization": f"Bearer {self._api_key}"} if self._api_key else {}
|
||||
payload = {
|
||||
"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
|
||||
|
||||
with httpx.Client(timeout=30) as client:
|
||||
response = client.post(url, json=payload, headers=headers)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
|
||||
def extract_table_name(record: Dict[str, Any]) -> Optional[str]:
|
||||
"""从检索结果中提取表名"""
|
||||
def _canonicalize(table_name: Any) -> Optional[str]:
|
||||
normalized = str(table_name or "").strip()
|
||||
if not normalized:
|
||||
return None
|
||||
return TABLE_NAME_ALIASES.get(normalized, normalized)
|
||||
|
||||
if not record:
|
||||
return None
|
||||
|
||||
metadata = record.get("metadata") or {}
|
||||
for key in ("table", "table_name"):
|
||||
if key in metadata:
|
||||
return _canonicalize(metadata.get(key))
|
||||
|
||||
for key in ("table", "table_name"):
|
||||
if key in record:
|
||||
return _canonicalize(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 _canonicalize(parsed.get(key))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
for line in str(content).splitlines():
|
||||
if line.lower().startswith("table:"):
|
||||
return _canonicalize(line.split(":", 1)[1].strip())
|
||||
|
||||
return None
|
||||
@@ -0,0 +1,376 @@
|
||||
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
|
||||
Reference in New Issue
Block a user