server.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349
  1. #!/usr/bin/env python3
  2. """
  3. Embedding Bridge - 本地 Embedding 与 Milvus 写入服务
  4. 供 Java Spring Boot 调用,完成:
  5. - 三级分块
  6. - 本地 HuggingFace 稠密向量
  7. - Milvus 2.5+ 原生 BM25 稀疏向量
  8. - Leaf-only 向量化存储
  9. 启动方式:
  10. python server.py --port 18732
  11. """
  12. import argparse
  13. import faulthandler
  14. import logging
  15. import os
  16. import sys
  17. import threading
  18. import time
  19. from pathlib import Path
  20. from typing import List, Optional
  21. # 限制 PyTorch / OpenMP 线程数,避免在 Windows + uvicorn 多线程环境下出现 segfault
  22. os.environ.setdefault("OMP_NUM_THREADS", "1")
  23. os.environ.setdefault("MKL_NUM_THREADS", "1")
  24. os.environ.setdefault("OPENBLAS_NUM_THREADS", "1")
  25. os.environ.setdefault("NUMEXPR_NUM_THREADS", "1")
  26. os.environ.setdefault("KMP_DUPLICATE_LIB_OK", "TRUE")
  27. faulthandler.enable()
  28. # ---------------------------------------------------------------------------
  29. # 父进程守护:Java 后端被 kill 后,子进程自动退出,释放端口与数据库连接
  30. # ---------------------------------------------------------------------------
  31. def _start_parent_watcher():
  32. """启动守护线程,当父进程退出时自杀。"""
  33. try:
  34. import psutil
  35. except ImportError:
  36. logging.warning("未安装 psutil,无法监听父进程状态;Java 退出后子进程可能残留")
  37. return
  38. try:
  39. parent = psutil.Process(os.getppid())
  40. except Exception:
  41. return
  42. def _watch():
  43. while True:
  44. time.sleep(2)
  45. try:
  46. if not parent.is_running() or parent.status() == psutil.STATUS_ZOMBIE:
  47. logging.info("父进程已退出,Embedding Bridge 自动终止")
  48. os._exit(0)
  49. except Exception:
  50. # 获取不到父进程信息时也退出,避免成为孤儿进程
  51. logging.info("父进程状态不可获取,Embedding Bridge 自动终止")
  52. os._exit(0)
  53. watcher = threading.Thread(target=_watch, daemon=True, name="parent-watcher")
  54. watcher.start()
  55. _start_parent_watcher()
  56. # ---------------------------------------------------------------------------
  57. # 强制 stdout/stderr 使用 UTF-8
  58. # ---------------------------------------------------------------------------
  59. for _stream in (sys.stdout, sys.stderr):
  60. try:
  61. _stream.reconfigure(encoding="utf-8", errors="replace")
  62. except Exception:
  63. pass
  64. # ---------------------------------------------------------------------------
  65. # 项目路径注入
  66. # ---------------------------------------------------------------------------
  67. BRIDGE_DIR = Path(__file__).resolve().parent
  68. if str(BRIDGE_DIR) not in sys.path:
  69. sys.path.insert(0, str(BRIDGE_DIR))
  70. # ---------------------------------------------------------------------------
  71. # 日志
  72. # ---------------------------------------------------------------------------
  73. logging.basicConfig(
  74. level=logging.INFO,
  75. format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
  76. stream=sys.stderr,
  77. )
  78. logger = logging.getLogger("embedding-bridge")
  79. # ---------------------------------------------------------------------------
  80. # FastAPI app
  81. # ---------------------------------------------------------------------------
  82. try:
  83. from fastapi import FastAPI, HTTPException
  84. from pydantic import BaseModel, Field
  85. except ImportError:
  86. logger.error("请先安装依赖: pip install -r requirements.txt")
  87. sys.exit(1)
  88. app = FastAPI(title="Embedding Bridge", version="1.0.0")
  89. # ---------------------------------------------------------------------------
  90. # 共享密钥(X-Bridge-Token 头校验)
  91. # ---------------------------------------------------------------------------
  92. _BRIDGE_AUTH_TOKEN = os.environ.get("EMBEDDING_BRIDGE_AUTH_TOKEN", "").strip()
  93. @app.middleware("http")
  94. async def _verify_token(request, call_next):
  95. """校验 X-Bridge-Token 头(/health 不校验)"""
  96. if _BRIDGE_AUTH_TOKEN and request.url.path != "/health":
  97. token = request.headers.get("X-Bridge-Token", "")
  98. if token != _BRIDGE_AUTH_TOKEN:
  99. from fastapi.responses import JSONResponse
  100. return JSONResponse(status_code=401, content={"detail": "invalid or missing X-Bridge-Token"})
  101. return await call_next(request)
  102. # ---------------------------------------------------------------------------
  103. # 模型
  104. # ---------------------------------------------------------------------------
  105. class IndexRequest(BaseModel):
  106. document_id: int
  107. text: str
  108. filename: str = ""
  109. file_type: str = ""
  110. file_path: str = ""
  111. page_number: int = 0
  112. chunk_size: int = 800
  113. chunk_overlap: int = 100
  114. category_id: Optional[int] = None
  115. class VectorizeRequest(BaseModel):
  116. document_id: int
  117. filename: str = ""
  118. file_type: str = ""
  119. file_path: str = ""
  120. chunks: List[dict]
  121. class DeleteByDocumentRequest(BaseModel):
  122. document_id: int
  123. class DeleteByIdsRequest(BaseModel):
  124. vector_ids: List[str]
  125. class RetrieveRequest(BaseModel):
  126. query: str = Field(min_length=1)
  127. top_k: int = Field(default=5, ge=1, le=100)
  128. mode: str = "hybrid"
  129. filter_expr: str = ""
  130. class EmbedRequest(BaseModel):
  131. texts: List[str] = Field(min_length=1, max_length=500)
  132. # ---------------------------------------------------------------------------
  133. # 预导入含 C 扩展的依赖(必须在主线程完成,避免在 uvicorn 工作线程中首次加载触发 segfault)
  134. # ---------------------------------------------------------------------------
  135. import pyarrow # noqa: F401
  136. import pandas # noqa: F401
  137. from backend.indexing.milvus_writer import MilvusWriter
  138. from backend.indexing.text_splitter import HierarchicalTextSplitter
  139. # ---------------------------------------------------------------------------
  140. # 延迟初始化(避免启动时立即加载大模型)
  141. # ---------------------------------------------------------------------------
  142. _splitter: Optional[HierarchicalTextSplitter] = None
  143. _writer: Optional[MilvusWriter] = None
  144. _lock = threading.Lock()
  145. def _get_splitter(chunk_size: int, chunk_overlap: int):
  146. return HierarchicalTextSplitter(chunk_size=chunk_size, chunk_overlap=chunk_overlap)
  147. def _get_writer():
  148. global _writer
  149. if _writer is None:
  150. _writer = MilvusWriter()
  151. return _writer
  152. # ---------------------------------------------------------------------------
  153. # 健康检查
  154. # ---------------------------------------------------------------------------
  155. @app.get("/health")
  156. def health():
  157. return {"status": "ok"}
  158. @app.post("/embed")
  159. def embed(req: EmbedRequest):
  160. try:
  161. from backend.indexing.embedding import embedding_service
  162. return {"vectors": embedding_service.get_embeddings(req.texts)}
  163. except Exception as exc:
  164. raise HTTPException(status_code=500, detail=f"embedding failed: {exc}") from exc
  165. # ---------------------------------------------------------------------------
  166. # 文档分块 + 向量化
  167. # ---------------------------------------------------------------------------
  168. @app.post("/index")
  169. def index_document(req: IndexRequest):
  170. start = time.time()
  171. try:
  172. splitter = _get_splitter(req.chunk_size, req.chunk_overlap)
  173. chunks = splitter.split_text(
  174. text=req.text,
  175. document_id=req.document_id,
  176. filename=req.filename,
  177. file_type=req.file_type,
  178. file_path=req.file_path,
  179. page_number=req.page_number,
  180. )
  181. if not chunks:
  182. return {"chunks": [], "collection": os.getenv("MILVUS_COLLECTION", "kb_documents"), "vector_count": 0}
  183. # 为所有 chunk 注入 document_id
  184. for c in chunks:
  185. c["document_id"] = req.document_id
  186. c["category_id"] = req.category_id
  187. writer = _get_writer()
  188. vector_ids = writer.write_chunks(chunks)
  189. leaf_index = 0
  190. for c in chunks:
  191. if c.get("chunk_level") == 3:
  192. if leaf_index < len(vector_ids):
  193. c["vector_id"] = vector_ids[leaf_index]
  194. leaf_index += 1
  195. # L1/L2 不设置 vector_id,避免 JSON 中出现 null
  196. cost = round((time.time() - start) * 1000)
  197. logger.info(
  198. "[index] document_id=%s chunks=%s leaf=%s cost=%sms",
  199. req.document_id,
  200. len(chunks),
  201. leaf_index,
  202. cost,
  203. )
  204. return {
  205. "chunks": chunks,
  206. "collection": os.getenv("MILVUS_COLLECTION", "kb_documents"),
  207. "vector_count": leaf_index,
  208. }
  209. except Exception as e:
  210. logger.error("[index] document_id=%s 失败: %s", req.document_id, e, exc_info=True)
  211. raise HTTPException(status_code=500, detail=f"向量化失败: {e}")
  212. # ---------------------------------------------------------------------------
  213. # 批量 chunk 向量化(用于单个 chunk 重试)
  214. # ---------------------------------------------------------------------------
  215. @app.post("/vectorize")
  216. def vectorize_chunks(req: VectorizeRequest):
  217. try:
  218. for c in req.chunks:
  219. c["document_id"] = req.document_id
  220. c.setdefault("filename", req.filename)
  221. c.setdefault("file_type", req.file_type)
  222. c.setdefault("file_path", req.file_path)
  223. c.setdefault("page_number", 0)
  224. writer = _get_writer()
  225. vector_ids = writer.write_chunks(req.chunks)
  226. return {"vector_ids": vector_ids}
  227. except Exception as e:
  228. logger.error("[vectorize] document_id=%s 失败: %s", req.document_id, e, exc_info=True)
  229. raise HTTPException(status_code=500, detail=f"向量化失败: {e}")
  230. # ---------------------------------------------------------------------------
  231. # 按文档删除向量
  232. # ---------------------------------------------------------------------------
  233. @app.post("/delete_by_document")
  234. def delete_by_document(req: DeleteByDocumentRequest):
  235. try:
  236. from backend.indexing.milvus_client import get_milvus_store
  237. store = get_milvus_store()
  238. count = store.delete_by_document(req.document_id)
  239. return {"deleted": count}
  240. except Exception as e:
  241. logger.error("[delete_by_document] document_id=%s 失败: %s", req.document_id, e, exc_info=True)
  242. raise HTTPException(status_code=500, detail=f"删除失败: {e}")
  243. # ---------------------------------------------------------------------------
  244. # 按 vector_id 删除向量
  245. # ---------------------------------------------------------------------------
  246. @app.post("/delete_by_vector_ids")
  247. def delete_by_vector_ids(req: DeleteByIdsRequest):
  248. try:
  249. from backend.indexing.milvus_client import get_milvus_store
  250. store = get_milvus_store()
  251. count = store.delete_by_ids(req.vector_ids)
  252. return {"deleted": count}
  253. except Exception as e:
  254. logger.error("[delete_by_vector_ids] 失败: %s", e, exc_info=True)
  255. raise HTTPException(status_code=500, detail=f"删除失败: {e}")
  256. # ---------------------------------------------------------------------------
  257. # 主入口
  258. # ---------------------------------------------------------------------------
  259. @app.post("/retrieve")
  260. def retrieve(req: RetrieveRequest):
  261. try:
  262. from backend.indexing.embedding import embedding_service
  263. from backend.indexing.milvus_client import get_milvus_store
  264. mode = req.mode.strip().lower()
  265. if mode not in ("dense", "hybrid"):
  266. raise HTTPException(status_code=400, detail="mode must be dense or hybrid")
  267. dense = embedding_service.get_embeddings([req.query])[0]
  268. store = get_milvus_store()
  269. if mode == "dense":
  270. hits = store.dense_retrieve(dense, req.top_k, req.filter_expr)
  271. else:
  272. hits = store.hybrid_retrieve(dense, req.query, req.top_k, filter_expr=req.filter_expr)
  273. return {"mode": mode, "top_k": req.top_k, "hits": hits}
  274. except HTTPException:
  275. raise
  276. except Exception as e:
  277. logger.error("[retrieve] failed: %s", e, exc_info=True)
  278. raise HTTPException(status_code=500, detail=f"retrieval failed: {e}")
  279. if __name__ == "__main__":
  280. parser = argparse.ArgumentParser()
  281. parser.add_argument("--port", type=int, default=int(os.getenv("EMBEDDING_BRIDGE_PORT", "18732")))
  282. parser.add_argument("--host", type=str, default=os.getenv("EMBEDDING_BRIDGE_HOST", "127.0.0.1"))
  283. args = parser.parse_args()
  284. import uvicorn
  285. logger.info("Embedding Bridge 启动: %s:%s", args.host, args.port)
  286. uvicorn.run(app, host=args.host, port=args.port, log_level="info")