Files
mnote/wolai-backend/app/services/lightrag_service.py
T
2025-12-06 16:47:17 +08:00

278 lines
10 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""LightRAG 集成。负责管理实例、增量索引与问答。"""
from __future__ import annotations
import asyncio
import logging
import os
from pathlib import Path
import sys
from typing import Any, AsyncIterator, Dict, List, Optional
from urllib.parse import urlparse
def _ensure_lightrag_available() -> None:
"""确保本地 LightRAG 代码可被 Python 找到."""
repo_root = Path(__file__).resolve().parents[3]
local_pkg = repo_root / "LightRAG"
if local_pkg.exists():
path_str = str(local_pkg)
if path_str not in sys.path:
sys.path.insert(0, path_str)
logger = logging.getLogger(__name__)
try:
from lightrag import LightRAG, QueryParam
from lightrag.kg.shared_storage import initialize_pipeline_status
from lightrag.llm.openai import gpt_4o_mini_complete, openai_embed
from lightrag.utils import logger as lightrag_logger
_LIGHTRAG_AVAILABLE = True
_IMPORT_ERROR: Optional[Exception] = None
except Exception as exc: # pragma: no cover - 本地缺失或版本不兼容时使用占位实现
_LIGHTRAG_AVAILABLE = False
_IMPORT_ERROR = exc
lightrag_logger = logging.getLogger("lightrag_stub")
class QueryParam: # type: ignore[override]
def __init__(self, mode: str = "mix") -> None:
self.mode = mode
self.stream = True
class LightRAG: # type: ignore[override]
async def initialize_storages(self) -> None:
return None
async def ainsert(self, *_: Any, **__: Any) -> str:
return "lightrag-skipped"
async def aquery_llm(self, *_: Any, **__: Any) -> Dict[str, Any]:
return {
"llm_response": {
"response_iterator": iter(()),
"content": "LightRAG 已跳过",
"is_streaming": False,
},
"data": {"references": []},
"metadata": {"skipped": True},
}
async def initialize_pipeline_status() -> None: # type: ignore[override]
return None
def gpt_4o_mini_complete(*_: Any, **__: Any) -> str: # type: ignore[override]
return ""
def openai_embed(*_: Any, **__: Any) -> List[float]: # type: ignore[override]
return []
_ensure_lightrag_available()
from app.config import settings
class LightRAGService:
"""
管理 LightRAG 单例、工作空间隔离与增量索引。
- 每个 workspace 映射为一个 LightRAG 实例(共享 Postgres
- 支持在同步/异步上下文中调用
"""
def __init__(self) -> None:
self.collection = settings.lightrag_collection
self._instances: Dict[str, LightRAG] = {}
self._locks: Dict[str, asyncio.Lock] = {}
self._pipeline_ready = False
self._configure_pg_env()
if not _LIGHTRAG_AVAILABLE and _IMPORT_ERROR:
logger.warning("LightRAG 不可用,使用占位实现:%s", _IMPORT_ERROR)
self._working_dir = (
Path(__file__).resolve().parent.parent / "runtime" / "lightrag_cache"
)
self._working_dir.mkdir(parents=True, exist_ok=True)
self._availability_error: Optional[str] = None
def _configure_pg_env(self) -> None:
"""根据配置将 pgvector 连接信息注入 LightRAG 需要的环境变量。"""
parsed = urlparse(settings.lightrag_db_url)
if parsed.scheme not in {"postgresql", "postgres"}:
self._availability_error = "LIGHTRAG_DB_URL 必须是 Postgres 连接串"
return
if parsed.username:
os.environ.setdefault("POSTGRES_USER", parsed.username)
if parsed.password:
os.environ.setdefault("POSTGRES_PASSWORD", parsed.password)
if parsed.hostname:
os.environ.setdefault("POSTGRES_HOST", parsed.hostname)
if parsed.port:
os.environ.setdefault("POSTGRES_PORT", str(parsed.port))
if parsed.path and len(parsed.path) > 1:
os.environ.setdefault("POSTGRES_DATABASE", parsed.path.lstrip("/"))
# OpenAI key 也交给环境变量,缺失时依旧由外部控制
if settings.openai_api_key:
os.environ.setdefault("OPENAI_API_KEY", settings.openai_api_key)
elif self._availability_error is None:
# 允许无 key 但标记提示,便于健康检查返回可诊断信息
self._availability_error = "缺少 OPENAI_API_KEYLightRAG 将使用占位实现"
def _namespace(self, workspace_id: Optional[str], user_id: str) -> str:
"""生成 LightRAG workspace 名称,优先 workspace,其次 user。"""
if workspace_id:
return f"workspace_{workspace_id}"
return f"user_{user_id}"
def _is_available(self) -> tuple[bool, Optional[str]]:
"""
判断 LightRAG 是否具备运行条件。
- import 失败或配置错误时返回 False 和原因
"""
if not _LIGHTRAG_AVAILABLE:
return False, f"LightRAG 导入失败: {self._availability_error or _IMPORT_ERROR}"
if self._availability_error:
# 缺少关键配置时也视为不可用
return False, self._availability_error
return True, None
async def _wait_with_timeout(self, coro: Any, *, timeout: float = 6.0) -> Any:
"""为外部调用包一层超时,避免卡住 worker / healthcheck。"""
return await asyncio.wait_for(coro, timeout=timeout)
async def _get_instance(self, workspace: str) -> LightRAG:
if workspace in self._instances:
return self._instances[workspace]
lock = self._locks.setdefault(workspace, asyncio.Lock())
async with lock:
if workspace in self._instances:
return self._instances[workspace]
rag = LightRAG(
working_dir=str(self._working_dir),
workspace=workspace,
kv_storage="PGKVStorage",
vector_storage="PGVectorStorage",
graph_storage="PGGraphStorage",
doc_status_storage="PGDocStatusStorage",
llm_model_func=gpt_4o_mini_complete,
embedding_func=openai_embed,
)
await self._wait_with_timeout(rag.initialize_storages())
if not self._pipeline_ready:
await self._wait_with_timeout(initialize_pipeline_status())
self._pipeline_ready = True
self._instances[workspace] = rag
lightrag_logger.info("LightRAG workspace %s ready", workspace)
return rag
async def index_document_async(
self,
*,
document_id: str,
user_id: str,
workspace_id: Optional[str],
text: str,
title: Optional[str] = None,
) -> str:
workspace = self._namespace(workspace_id, user_id)
rag = await self._get_instance(workspace)
clean_text = text.strip()
if not clean_text:
raise ValueError("空文本无法建立 LightRAG 索引")
file_path = f"doc://{document_id}"
if title:
file_path = f"{file_path}?title={title}"
track_id = await rag.ainsert(
clean_text,
ids=[document_id],
file_paths=[file_path],
)
return track_id
def index_document(
self,
*,
document_id: str,
user_id: str,
workspace_id: Optional[str],
text: str,
title: Optional[str] = None,
) -> str:
"""同步环境(如 Celery)调用的封装。"""
return asyncio.run(
self.index_document_async(
document_id=document_id,
user_id=user_id,
workspace_id=workspace_id,
text=text,
title=title,
)
)
async def query_async(
self,
*,
query_text: str,
user_id: str,
workspace_id: Optional[str],
stream: bool = True,
mode: str = "mix",
) -> Dict[str, Any]:
workspace = self._namespace(workspace_id, user_id)
rag = await self._get_instance(workspace)
param = QueryParam(mode=mode)
param.stream = stream
result = await self._wait_with_timeout(rag.aquery_llm(query_text, param))
return result
async def stream_answer(
self,
*,
query_text: str,
user_id: str,
workspace_id: Optional[str],
) -> Dict[str, Any]:
"""统一返回结构,包含 SSE 需要的 iterator 与引用。"""
ok, reason = self._is_available()
if not ok:
return {
"references": [],
"iterator": iter(()),
"content": f"LightRAG 未就绪:{reason}",
"is_streaming": False,
"metadata": {"skipped": True, "reason": reason},
}
result = await self.query_async(
query_text=query_text,
user_id=user_id,
workspace_id=workspace_id,
stream=True,
)
llm_resp = result.get("llm_response", {})
references: List[Dict[str, Any]] = (
result.get("data", {}).get("references", []) or []
)
return {
"references": references,
"iterator": llm_resp.get("response_iterator"),
"content": llm_resp.get("content"),
"is_streaming": llm_resp.get("is_streaming", False),
"metadata": result.get("metadata", {}),
}
async def run_healthcheck(self) -> Dict[str, Any]:
workspace = "__healthcheck__"
ok, reason = self._is_available()
if not ok:
return {"ok": False, "workspace": workspace, "error": reason}
try:
rag = await self._get_instance(workspace)
await self._wait_with_timeout(rag.doc_status.initialize(), timeout=5.0)
return {"ok": True, "workspace": workspace}
except Exception as exc: # pragma: no cover - 调试辅助
return {"ok": False, "workspace": workspace, "error": str(exc)}
lightrag_service = LightRAGService()