Restore 0.1.5 version from stash
This commit is contained in:
@@ -48,6 +48,16 @@ class Settings(BaseSettings):
|
||||
openai_api_key: str = ""
|
||||
lightrag_db_url: str
|
||||
lightrag_collection: str = "wolai-docs"
|
||||
lightrag_llm_model: str = "qwen3:32b"
|
||||
lightrag_embedding_model: str = "qwen3-embedding:8b"
|
||||
lightrag_rerank_model: str = "dengcao/Qwen3-Reranker-4B:Q5_K_M"
|
||||
ollama_base_url: str = "http://127.0.0.1:11434"
|
||||
lightrag_embedding_dim: int = 4096
|
||||
deepseek_api_key: str = ""
|
||||
deepseek_base_url: str = "https://api.deepseek.com"
|
||||
deepseek_model: str = "deepseek-chat"
|
||||
searxng_base_url: Optional[str] = None
|
||||
searxng_api_token: Optional[str] = None
|
||||
|
||||
|
||||
@lru_cache
|
||||
|
||||
@@ -1,55 +1,149 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from typing import AsyncIterator, Optional
|
||||
|
||||
from fastapi import APIRouter, HTTPException, status
|
||||
from fastapi.responses import StreamingResponse
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.deps import AuthDep
|
||||
from app.services.lightrag_service import lightrag_service
|
||||
from app.services.simple_ai_service import simple_ai_service
|
||||
from app.services.searxng_client import searxng_client
|
||||
from app.services.supabase_rest import supabase_rest
|
||||
|
||||
router = APIRouter(prefix="/chat")
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@router.get("")
|
||||
async def chat(
|
||||
class ChatRequest(BaseModel):
|
||||
query: str
|
||||
document_id: Optional[str] = None
|
||||
workspace_id: Optional[str] = None
|
||||
model: Optional[str] = None
|
||||
use_web_search: bool = False
|
||||
|
||||
|
||||
async def _build_chat_response(
|
||||
query: str,
|
||||
auth: AuthDep,
|
||||
document_id: Optional[str] = None,
|
||||
document_id: Optional[str],
|
||||
workspace_id: Optional[str],
|
||||
model: Optional[str],
|
||||
use_web_search: bool,
|
||||
) -> StreamingResponse:
|
||||
"""
|
||||
基于 LightRAG 的 SSE 流式回答。
|
||||
query: 必填问题
|
||||
document_id: 可选,指定所属 workspace
|
||||
"""
|
||||
"""统一封装 GET/POST 的对话逻辑,便于同时支持长文本 POST。"""
|
||||
if not query or not query.strip():
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="query 不能为空")
|
||||
|
||||
workspace_id: Optional[str] = None
|
||||
resolved_workspace: Optional[str] = None
|
||||
if workspace_id:
|
||||
membership = supabase_rest.select_one(
|
||||
"workspace_members", {"workspace_id": workspace_id, "user_id": auth.user_id}
|
||||
)
|
||||
if not membership:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN, detail="无权访问该 workspace"
|
||||
)
|
||||
resolved_workspace = workspace_id
|
||||
|
||||
if document_id:
|
||||
document = supabase_rest.select_one(
|
||||
"documents", {"id": document_id, "user_id": auth.user_id}
|
||||
)
|
||||
if not document:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Document not found")
|
||||
workspace_id = (
|
||||
str(document.get("workspace_id")) if document.get("workspace_id") else None
|
||||
)
|
||||
doc_workspace = str(document.get("workspace_id")) if document.get("workspace_id") else None
|
||||
if resolved_workspace and doc_workspace and resolved_workspace != doc_workspace:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail="workspace 与文档不匹配"
|
||||
)
|
||||
resolved_workspace = resolved_workspace or doc_workspace
|
||||
|
||||
health = await lightrag_service.run_healthcheck()
|
||||
if not health.get("ok"):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail=health.get("error", "LightRAG 未就绪"),
|
||||
)
|
||||
fallback_result: Optional[dict[str, object]] = None
|
||||
lightrag_result: Optional[dict[str, object]] = None
|
||||
web_search_result: Optional[dict[str, object]] = None
|
||||
|
||||
lightrag_result = await lightrag_service.stream_answer(
|
||||
query_text=query, user_id=auth.user_id, workspace_id=workspace_id
|
||||
)
|
||||
fallback_reason: Optional[str] = None
|
||||
# 联网搜索模式:直接用搜索上下文 + DeepSeek/Ollama 生成
|
||||
if use_web_search:
|
||||
search_results = searxng_client.search(query)
|
||||
context = "\n\n".join(
|
||||
[f"[{idx+1}] {item['title']}\n{item.get('snippet','')}\n{item['url']}" for idx, item in enumerate(search_results)]
|
||||
)
|
||||
web_search_result = await lightrag_service.answer_with_context(
|
||||
query_text=query,
|
||||
context=context or "未获取到搜索结果",
|
||||
prefer_deepseek=True,
|
||||
)
|
||||
# 将搜索结果作为引用
|
||||
if web_search_result is not None:
|
||||
web_search_result["references"] = search_results
|
||||
else:
|
||||
health = await lightrag_service.run_healthcheck()
|
||||
if not health.get("ok"):
|
||||
fallback_reason = health.get("error", "LightRAG 未就绪")
|
||||
else:
|
||||
try:
|
||||
lightrag_result = await lightrag_service.stream_answer(
|
||||
query_text=query,
|
||||
user_id=auth.user_id,
|
||||
workspace_id=resolved_workspace,
|
||||
model_choice=model,
|
||||
document_id=document_id,
|
||||
)
|
||||
except Exception as exc: # pragma: no cover - 运行时保护
|
||||
fallback_reason = str(exc)
|
||||
logger.exception("LightRAG stream_answer failed, fallback to simple summary: %s", fallback_reason)
|
||||
|
||||
if fallback_reason:
|
||||
fallback_result = simple_ai_service.summarize_document(
|
||||
query=query,
|
||||
user_id=auth.user_id,
|
||||
document_id=document_id,
|
||||
workspace_id=resolved_workspace,
|
||||
)
|
||||
logger.warning("LightRAG unavailable, use simple summary. reason=%s", fallback_reason)
|
||||
|
||||
async def event_stream() -> AsyncIterator[str]:
|
||||
if web_search_result is not None:
|
||||
if web_search_result.get("is_streaming") and web_search_result.get("iterator"):
|
||||
iterator = web_search_result["iterator"]
|
||||
async for chunk in iterator:
|
||||
payload = json.dumps({"type": "chunk", "content": chunk})
|
||||
yield f"data: {payload}\n\n"
|
||||
else:
|
||||
payload = json.dumps(
|
||||
{"type": "chunk", "content": web_search_result.get("content", "")}
|
||||
)
|
||||
yield f"data: {payload}\n\n"
|
||||
references = web_search_result.get("references", []) if isinstance(web_search_result, dict) else []
|
||||
yield f"data: {json.dumps({'type': 'references', 'data': references})}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
return
|
||||
|
||||
if fallback_result is not None:
|
||||
payload = json.dumps(
|
||||
{
|
||||
"type": "chunk",
|
||||
"content": fallback_result.get("content", ""),
|
||||
}
|
||||
)
|
||||
yield f"data: {payload}\n\n"
|
||||
references = fallback_result.get("references", [])
|
||||
yield f"data: {json.dumps({'type': 'references', 'data': references})}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
return
|
||||
|
||||
if lightrag_result is None:
|
||||
payload = json.dumps({"type": "chunk", "content": "LightRAG 暂不可用"})
|
||||
yield f"data: {payload}\n\n"
|
||||
yield f"data: {json.dumps({'type': 'references', 'data': []})}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
return
|
||||
|
||||
if lightrag_result.get("is_streaming") and lightrag_result.get("iterator"):
|
||||
iterator = lightrag_result["iterator"]
|
||||
async for chunk in iterator:
|
||||
@@ -65,3 +159,41 @@ async def chat(
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
return StreamingResponse(event_stream(), media_type="text/event-stream")
|
||||
|
||||
|
||||
@router.get("")
|
||||
async def chat(
|
||||
query: str,
|
||||
auth: AuthDep,
|
||||
document_id: Optional[str] = None,
|
||||
workspace_id: Optional[str] = None,
|
||||
model: Optional[str] = None,
|
||||
use_web_search: bool = False,
|
||||
) -> StreamingResponse:
|
||||
"""
|
||||
基于 LightRAG 的 SSE 流式回答。
|
||||
query: 必填问题
|
||||
document_id: 可选,指定所属 workspace
|
||||
workspace_id: 可选,直接指定 workspace,优先级高于 document_id
|
||||
"""
|
||||
return await _build_chat_response(
|
||||
query=query,
|
||||
auth=auth,
|
||||
document_id=document_id,
|
||||
workspace_id=workspace_id,
|
||||
model=model,
|
||||
use_web_search=use_web_search,
|
||||
)
|
||||
|
||||
|
||||
@router.post("")
|
||||
async def chat_post(payload: ChatRequest, auth: AuthDep) -> StreamingResponse:
|
||||
"""POST 版本,适配长问题与 JSON 传参。"""
|
||||
return await _build_chat_response(
|
||||
query=payload.query,
|
||||
auth=auth,
|
||||
document_id=payload.document_id,
|
||||
workspace_id=payload.workspace_id,
|
||||
model=payload.model,
|
||||
use_web_search=payload.use_web_search,
|
||||
)
|
||||
|
||||
+4
@@ -0,0 +1,4 @@
|
||||
<?xml version='1.0' encoding='utf-8'?>
|
||||
<graphml xmlns="http://graphml.graphdrawing.org/xmlns" xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance" xsi:schemaLocation="http://graphml.graphdrawing.org/xmlns http://graphml.graphdrawing.org/xmlns/1.0/graphml.xsd">
|
||||
<graph edgedefault="undirected" />
|
||||
</graphml>
|
||||
@@ -7,8 +7,12 @@ import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
from contextvars import ContextVar
|
||||
from typing import Any, AsyncIterator, Dict, List, Optional
|
||||
from urllib.parse import urlparse
|
||||
import httpx
|
||||
import numpy as np
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
|
||||
def _ensure_lightrag_available() -> None:
|
||||
@@ -23,10 +27,18 @@ def _ensure_lightrag_available() -> None:
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 优先将本地 LightRAG 代码加入 sys.path,避免导入失败走占位实现
|
||||
_ensure_lightrag_available()
|
||||
|
||||
_EMBEDDING_MODEL_OVERRIDE: ContextVar[Optional[str]] = ContextVar(
|
||||
"lightrag_embedding_model_override",
|
||||
default=None,
|
||||
)
|
||||
|
||||
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.llm.ollama import ollama_model_complete, ollama_embed
|
||||
from lightrag.utils import logger as lightrag_logger
|
||||
_LIGHTRAG_AVAILABLE = True
|
||||
_IMPORT_ERROR: Optional[Exception] = None
|
||||
@@ -34,6 +46,12 @@ except Exception as exc: # pragma: no cover - 本地缺失或版本不兼容时
|
||||
_LIGHTRAG_AVAILABLE = False
|
||||
_IMPORT_ERROR = exc
|
||||
lightrag_logger = logging.getLogger("lightrag_stub")
|
||||
# 占位的 ollama 方法,避免引用错误
|
||||
async def ollama_model_complete(*_: Any, **__: Any) -> str: # type: ignore[override]
|
||||
return ""
|
||||
|
||||
async def ollama_embed(*_: Any, **__: Any) -> list[list[float]]: # type: ignore[override]
|
||||
return []
|
||||
|
||||
class QueryParam: # type: ignore[override]
|
||||
def __init__(self, mode: str = "mix") -> None:
|
||||
@@ -70,6 +88,7 @@ except Exception as exc: # pragma: no cover - 本地缺失或版本不兼容时
|
||||
_ensure_lightrag_available()
|
||||
|
||||
from app.config import settings
|
||||
from app.services.supabase_rest import supabase_rest
|
||||
|
||||
|
||||
class LightRAGService:
|
||||
@@ -92,6 +111,19 @@ class LightRAGService:
|
||||
)
|
||||
self._working_dir.mkdir(parents=True, exist_ok=True)
|
||||
self._availability_error: Optional[str] = None
|
||||
self._ollama_host = settings.ollama_base_url
|
||||
self._use_deepseek = bool(getattr(settings, "deepseek_api_key", ""))
|
||||
self._deepseek_client: Optional[AsyncOpenAI] = None
|
||||
if self._use_deepseek:
|
||||
self._deepseek_client = AsyncOpenAI(
|
||||
api_key=settings.deepseek_api_key,
|
||||
base_url=getattr(settings, "deepseek_base_url", None),
|
||||
)
|
||||
self._llm_model_name = (
|
||||
getattr(settings, "deepseek_model", "deepseek-chat")
|
||||
if self._use_deepseek
|
||||
else settings.lightrag_llm_model
|
||||
)
|
||||
|
||||
def _configure_pg_env(self) -> None:
|
||||
"""根据配置将 pgvector 连接信息注入 LightRAG 需要的环境变量。"""
|
||||
@@ -110,18 +142,19 @@ class LightRAGService:
|
||||
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_KEY,LightRAG 将使用占位实现"
|
||||
# Ollama 走本地 HTTP,无需 OpenAI key
|
||||
# 允许通过环境变量关闭 KG 抽取,避免大模型深拷贝异常
|
||||
os.environ.setdefault("LIGHTRAG_DISABLE_ENTITY_RELATION", "true")
|
||||
# 为大维度向量配置 IVFFlat,避免 HNSW 2000 维限制
|
||||
os.environ.setdefault("POSTGRES_VECTOR_INDEX_TYPE", "IVFFlat")
|
||||
os.environ.setdefault("EMBEDDING_DIM", str(settings.lightrag_embedding_dim if hasattr(settings, "lightrag_embedding_dim") else 4096))
|
||||
|
||||
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 _namespace(self, workspace_id: Optional[str], user_id: str, *, prefer_deepseek: bool = False) -> str:
|
||||
"""生成 LightRAG workspace 名称,优先 workspace,其次 user;根据模型标记区分实例。"""
|
||||
base = f"workspace_{workspace_id}" if workspace_id else f"user_{user_id}"
|
||||
if prefer_deepseek:
|
||||
return f"{base}_deepseek"
|
||||
return base
|
||||
|
||||
def _is_available(self) -> tuple[bool, Optional[str]]:
|
||||
"""
|
||||
@@ -135,11 +168,117 @@ class LightRAGService:
|
||||
return False, self._availability_error
|
||||
return True, None
|
||||
|
||||
async def _wait_with_timeout(self, coro: Any, *, timeout: float = 6.0) -> Any:
|
||||
async def _wait_with_timeout(self, coro: Any, *, timeout: float = 300.0) -> Any:
|
||||
"""为外部调用包一层超时,避免卡住 worker / healthcheck。"""
|
||||
return await asyncio.wait_for(coro, timeout=timeout)
|
||||
|
||||
async def _get_instance(self, workspace: str) -> LightRAG:
|
||||
async def _ollama_llm(self, *args: Any, **kwargs: Any) -> Any:
|
||||
"""
|
||||
固定使用配置好的 Ollama LLM。
|
||||
LightRAG 会传入 prompt/system_prompt/history_messages。
|
||||
"""
|
||||
prompt = kwargs.pop("prompt", None)
|
||||
if prompt is None and args:
|
||||
prompt = args[0]
|
||||
if prompt is None:
|
||||
raise ValueError("缺少 prompt,无法调用 Ollama LLM")
|
||||
stream_flag = bool(kwargs.pop("stream", False))
|
||||
timeout = kwargs.pop("timeout", None)
|
||||
return await ollama_model_complete(
|
||||
prompt=prompt,
|
||||
host=self._ollama_host,
|
||||
timeout=timeout,
|
||||
stream=stream_flag,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
async def _deepseek_llm(self, *args: Any, **kwargs: Any) -> Any:
|
||||
"""
|
||||
使用 DeepSeek 在线模型,符合 LightRAG 的 llm_model_func 接口。
|
||||
"""
|
||||
if not self._deepseek_client:
|
||||
raise RuntimeError("DeepSeek 客户端未初始化")
|
||||
prompt = kwargs.pop("prompt", None)
|
||||
if prompt is None and args:
|
||||
prompt = args[0]
|
||||
if prompt is None:
|
||||
raise ValueError("缺少 prompt,无法调用 DeepSeek LLM")
|
||||
system_prompt = kwargs.pop("system_prompt", None)
|
||||
history = kwargs.pop("history_messages", []) or []
|
||||
stream_flag = bool(kwargs.pop("stream", False))
|
||||
temperature = kwargs.pop("temperature", 0.2)
|
||||
max_tokens = kwargs.pop("max_tokens", 512)
|
||||
# 构建消息
|
||||
messages = []
|
||||
if system_prompt:
|
||||
messages.append({"role": "system", "content": system_prompt})
|
||||
messages.extend(history)
|
||||
messages.append({"role": "user", "content": prompt})
|
||||
client = self._deepseek_client
|
||||
if stream_flag:
|
||||
response = await client.chat.completions.create(
|
||||
model=self._llm_model_name,
|
||||
messages=messages,
|
||||
stream=True,
|
||||
temperature=temperature,
|
||||
max_tokens=max_tokens,
|
||||
)
|
||||
|
||||
async def _aiter():
|
||||
async for chunk in response:
|
||||
delta = chunk.choices[0].delta.content or ""
|
||||
if delta:
|
||||
yield delta
|
||||
|
||||
return _aiter()
|
||||
else:
|
||||
response = await client.chat.completions.create(
|
||||
model=self._llm_model_name,
|
||||
messages=messages,
|
||||
stream=False,
|
||||
temperature=temperature,
|
||||
max_tokens=max_tokens,
|
||||
)
|
||||
return response.choices[0].message.content
|
||||
|
||||
async def _ollama_embed(self, texts: list[str], **kwargs: Any) -> Any:
|
||||
"""固定使用配置好的 Ollama embedding 模型。"""
|
||||
timeout = kwargs.pop("timeout", None)
|
||||
embed_model_override = kwargs.pop("embed_model", None) or _EMBEDDING_MODEL_OVERRIDE.get()
|
||||
embed_model = embed_model_override or settings.lightrag_embedding_model
|
||||
return await ollama_embed(
|
||||
texts,
|
||||
embed_model=embed_model,
|
||||
host=self._ollama_host,
|
||||
timeout=timeout,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
async def _embedding_rerank(self, query: str, documents: List[str], **_: Any) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
简易基于向量的 rerank:使用指定 embedding 模型计算 query/doc 向量并按余弦相似度排序。
|
||||
若调用失败则返回原序。
|
||||
"""
|
||||
if not documents:
|
||||
return []
|
||||
try:
|
||||
query_vec = await self._ollama_embed([query])
|
||||
doc_vecs = await self._ollama_embed(documents)
|
||||
q = np.array(query_vec[0], dtype=float)
|
||||
d = np.array(doc_vecs, dtype=float)
|
||||
q_norm = np.linalg.norm(q) + 1e-8
|
||||
d_norm = np.linalg.norm(d, axis=1) + 1e-8
|
||||
scores = (d @ q) / (d_norm * q_norm)
|
||||
order = np.argsort(-scores)
|
||||
return [
|
||||
{"index": int(idx), "relevance_score": float(scores[idx])}
|
||||
for idx in order
|
||||
]
|
||||
except Exception as exc:
|
||||
logger.warning("rerank 失败,回退原序:%s", exc)
|
||||
return [{"index": i, "relevance_score": 0.0} for i in range(len(documents))]
|
||||
|
||||
async def _get_instance(self, workspace: str, *, prefer_deepseek: bool = False) -> LightRAG:
|
||||
if workspace in self._instances:
|
||||
return self._instances[workspace]
|
||||
|
||||
@@ -148,15 +287,25 @@ class LightRAGService:
|
||||
if workspace in self._instances:
|
||||
return self._instances[workspace]
|
||||
|
||||
llm_func = self._deepseek_llm if (self._use_deepseek and prefer_deepseek) else self._ollama_llm
|
||||
rag = LightRAG(
|
||||
working_dir=str(self._working_dir),
|
||||
workspace=workspace,
|
||||
kv_storage="PGKVStorage",
|
||||
vector_storage="PGVectorStorage",
|
||||
graph_storage="PGGraphStorage",
|
||||
graph_storage="NetworkXStorage",
|
||||
doc_status_storage="PGDocStatusStorage",
|
||||
llm_model_func=gpt_4o_mini_complete,
|
||||
embedding_func=openai_embed,
|
||||
llm_model_func=llm_func,
|
||||
llm_model_name=self._llm_model_name if (self._use_deepseek and prefer_deepseek) else settings.lightrag_llm_model,
|
||||
llm_model_kwargs={
|
||||
"options": {
|
||||
# 限制生成长度,避免本地大模型回答过慢
|
||||
"num_predict": 512,
|
||||
"temperature": 0.2,
|
||||
}
|
||||
},
|
||||
embedding_func=self._ollama_embed,
|
||||
rerank_model_func=self._embedding_rerank,
|
||||
)
|
||||
await self._wait_with_timeout(rag.initialize_storages())
|
||||
if not self._pipeline_ready:
|
||||
@@ -217,13 +366,43 @@ class LightRAGService:
|
||||
user_id: str,
|
||||
workspace_id: Optional[str],
|
||||
stream: bool = True,
|
||||
mode: str = "mix",
|
||||
mode: str = "naive",
|
||||
model_choice: Optional[str] = None,
|
||||
rag_settings: Optional[Dict[str, Any]] = None,
|
||||
) -> Dict[str, Any]:
|
||||
workspace = self._namespace(workspace_id, user_id)
|
||||
rag = await self._get_instance(workspace)
|
||||
param = QueryParam(mode=mode)
|
||||
prefer_deepseek = model_choice == "deepseek"
|
||||
workspace = self._namespace(workspace_id, user_id, prefer_deepseek=prefer_deepseek)
|
||||
rag = await self._get_instance(workspace, prefer_deepseek=prefer_deepseek)
|
||||
param = QueryParam(mode=mode, top_k=6, chunk_top_k=6, max_total_tokens=4096)
|
||||
if rag_settings:
|
||||
try:
|
||||
if rag_settings.get("mode"):
|
||||
param.mode = str(rag_settings["mode"])
|
||||
if isinstance(rag_settings.get("top_k"), (int, float)):
|
||||
param.top_k = int(rag_settings["top_k"])
|
||||
if isinstance(rag_settings.get("chunk_top_k"), (int, float)):
|
||||
param.chunk_top_k = int(rag_settings["chunk_top_k"])
|
||||
if isinstance(rag_settings.get("max_entity_tokens"), (int, float)):
|
||||
param.max_entity_tokens = int(rag_settings["max_entity_tokens"])
|
||||
if isinstance(rag_settings.get("max_relation_tokens"), (int, float)):
|
||||
param.max_relation_tokens = int(rag_settings["max_relation_tokens"])
|
||||
if isinstance(rag_settings.get("max_total_tokens"), (int, float)):
|
||||
param.max_total_tokens = int(rag_settings["max_total_tokens"])
|
||||
if "enable_rerank" in rag_settings:
|
||||
param.enable_rerank = bool(rag_settings["enable_rerank"])
|
||||
if rag_settings.get("user_prompt"):
|
||||
param.user_prompt = str(rag_settings["user_prompt"])
|
||||
except Exception as exc:
|
||||
logger.warning("解析 RAG 参数失败,继续使用默认值: %s", exc)
|
||||
param.stream = stream
|
||||
result = await self._wait_with_timeout(rag.aquery_llm(query_text, param))
|
||||
embedding_token = None
|
||||
if rag_settings and isinstance(rag_settings.get("embedding_model"), str):
|
||||
embedding_token = _EMBEDDING_MODEL_OVERRIDE.set(str(rag_settings["embedding_model"]))
|
||||
try:
|
||||
result = await self._wait_with_timeout(rag.aquery_llm(query_text, param))
|
||||
finally:
|
||||
if embedding_token is not None:
|
||||
_EMBEDDING_MODEL_OVERRIDE.reset(embedding_token)
|
||||
return result
|
||||
|
||||
async def stream_answer(
|
||||
@@ -232,6 +411,8 @@ class LightRAGService:
|
||||
query_text: str,
|
||||
user_id: str,
|
||||
workspace_id: Optional[str],
|
||||
model_choice: Optional[str] = None,
|
||||
document_id: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""统一返回结构,包含 SSE 需要的 iterator 与引用。"""
|
||||
ok, reason = self._is_available()
|
||||
@@ -243,11 +424,16 @@ class LightRAGService:
|
||||
"is_streaming": False,
|
||||
"metadata": {"skipped": True, "reason": reason},
|
||||
}
|
||||
rag_settings = self._load_rag_settings(document_id)
|
||||
rag_mode = (rag_settings or {}).get("mode")
|
||||
result = await self.query_async(
|
||||
query_text=query_text,
|
||||
user_id=user_id,
|
||||
workspace_id=workspace_id,
|
||||
stream=True,
|
||||
model_choice=model_choice,
|
||||
mode=str(rag_mode) if isinstance(rag_mode, str) else "naive",
|
||||
rag_settings=rag_settings,
|
||||
)
|
||||
llm_resp = result.get("llm_response", {})
|
||||
references: List[Dict[str, Any]] = (
|
||||
@@ -261,6 +447,59 @@ class LightRAGService:
|
||||
"metadata": result.get("metadata", {}),
|
||||
}
|
||||
|
||||
def _load_rag_settings(self, document_id: Optional[str]) -> Optional[Dict[str, Any]]:
|
||||
if not document_id:
|
||||
return None
|
||||
try:
|
||||
doc = supabase_rest.select_one("documents", {"id": document_id})
|
||||
except Exception as exc:
|
||||
logger.warning("读取文档 %s 的 RAG 配置失败:%s", document_id, exc)
|
||||
return None
|
||||
value = doc.get("rag_settings") if doc else None
|
||||
if isinstance(value, dict):
|
||||
return value
|
||||
return None
|
||||
|
||||
async def answer_with_context(
|
||||
self,
|
||||
*,
|
||||
query_text: str,
|
||||
context: str,
|
||||
prefer_deepseek: bool = False,
|
||||
) -> Dict[str, Any]:
|
||||
"""基于外部上下文的简单问答,优先走 DeepSeek,无则回退 Ollama。"""
|
||||
prompt = (
|
||||
"你是检索结果总结助手。根据以下搜索摘要与来源回答用户问题,"
|
||||
"答案需简洁且引用要点。保持中文输出,并在内容后附上引用编号。\n\n"
|
||||
f"【搜索摘要】\n{context}\n\n【用户问题】{query_text}"
|
||||
)
|
||||
prefer_deepseek = prefer_deepseek and self._use_deepseek
|
||||
stream_flag = True
|
||||
iterator: Optional[AsyncIterator[str]] = None
|
||||
content: Optional[str] = None
|
||||
reason: Optional[str] = None
|
||||
try:
|
||||
if prefer_deepseek and self._deepseek_client:
|
||||
resp = await self._deepseek_llm(prompt, stream=stream_flag)
|
||||
iterator = resp if hasattr(resp, "__aiter__") else None
|
||||
content = None if iterator else str(resp)
|
||||
else:
|
||||
resp = await self._ollama_llm(prompt, stream=stream_flag)
|
||||
iterator = resp if hasattr(resp, "__aiter__") else None
|
||||
content = None if iterator else str(resp)
|
||||
except Exception as exc: # pragma: no cover - 运行时保护
|
||||
iterator = None
|
||||
content = f"生成失败:{exc}"
|
||||
reason = str(exc)
|
||||
|
||||
return {
|
||||
"references": [],
|
||||
"iterator": iterator,
|
||||
"content": content,
|
||||
"is_streaming": iterator is not None,
|
||||
"metadata": {"reason": reason} if reason else {},
|
||||
}
|
||||
|
||||
async def run_healthcheck(self) -> Dict[str, Any]:
|
||||
workspace = "__healthcheck__"
|
||||
ok, reason = self._is_available()
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import List, Dict, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from app.config import settings
|
||||
|
||||
|
||||
class SearxngClient:
|
||||
"""封装 SearxNG 简单搜索接口。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.base_url: Optional[str] = getattr(settings, "searxng_base_url", None)
|
||||
self.token: Optional[str] = getattr(settings, "searxng_api_token", None)
|
||||
|
||||
def search(self, query: str, *, limit: int = 5, categories: str = "general") -> List[Dict[str, str]]:
|
||||
"""返回精简的搜索结果列表。"""
|
||||
if not self.base_url:
|
||||
return []
|
||||
params = {
|
||||
"q": query,
|
||||
"format": "json",
|
||||
"engines": "",
|
||||
"language": "zh-CN",
|
||||
"categories": categories,
|
||||
"limit": limit,
|
||||
}
|
||||
headers = {}
|
||||
if self.token:
|
||||
headers["Authorization"] = f"Token {self.token}"
|
||||
try:
|
||||
resp = httpx.get(
|
||||
self.base_url.rstrip("/") + "/search",
|
||||
params=params,
|
||||
headers=headers,
|
||||
timeout=10.0,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
except Exception:
|
||||
return []
|
||||
results = []
|
||||
for item in data.get("results", [])[:limit]:
|
||||
title = item.get("title") or ""
|
||||
content = item.get("content") or ""
|
||||
url = item.get("url") or ""
|
||||
if not url:
|
||||
continue
|
||||
results.append(
|
||||
{
|
||||
"title": title,
|
||||
"snippet": content,
|
||||
"url": url,
|
||||
}
|
||||
)
|
||||
return results
|
||||
|
||||
|
||||
searxng_client = SearxngClient()
|
||||
@@ -0,0 +1,91 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Dict, Optional
|
||||
|
||||
from app.services.supabase_rest import supabase_rest
|
||||
|
||||
|
||||
class SimpleAIService:
|
||||
"""提供在 LightRAG 不可用时的兜底回答。"""
|
||||
|
||||
def summarize_document(
|
||||
self,
|
||||
*,
|
||||
query: str,
|
||||
user_id: str,
|
||||
document_id: Optional[str],
|
||||
workspace_id: Optional[str],
|
||||
) -> Dict[str, object]:
|
||||
document = None
|
||||
if document_id:
|
||||
document = supabase_rest.select_one(
|
||||
"documents",
|
||||
{
|
||||
"id": document_id,
|
||||
"user_id": user_id,
|
||||
},
|
||||
)
|
||||
|
||||
title = (document or {}).get("title") or "当前页面"
|
||||
raw_text = (document or {}).get("raw_text") or ""
|
||||
if not raw_text and document and document.get("content"):
|
||||
raw_text = self._blocks_to_text(document["content"])
|
||||
|
||||
snippet = raw_text.strip().replace("\n", " ")
|
||||
if len(snippet) > 400:
|
||||
snippet = snippet[:400].rstrip() + "…"
|
||||
|
||||
if snippet:
|
||||
summary = f"《{title}》当前摘要:{snippet}"
|
||||
else:
|
||||
summary = f"《{title}》暂未填写正文内容,可直接编辑后再次提问。"
|
||||
|
||||
answer = "\n\n".join(
|
||||
[
|
||||
summary,
|
||||
f"你的问题:{query}",
|
||||
"(提示:LightRAG 暂未就绪,已使用本地摘要兜底回答)",
|
||||
]
|
||||
)
|
||||
|
||||
references = []
|
||||
if document_id:
|
||||
ref_path = f"doc://{document_id}"
|
||||
if title:
|
||||
ref_path = f"{ref_path}?title={title}"
|
||||
references.append(
|
||||
{
|
||||
"file_path": ref_path,
|
||||
"workspace": workspace_id or "",
|
||||
}
|
||||
)
|
||||
|
||||
return {"content": answer, "references": references}
|
||||
|
||||
def _blocks_to_text(self, content: object) -> str:
|
||||
if isinstance(content, str):
|
||||
return content
|
||||
if isinstance(content, list):
|
||||
parts = []
|
||||
for item in content:
|
||||
text = self._blocks_to_text(item)
|
||||
if text:
|
||||
parts.append(text)
|
||||
return " ".join(parts)
|
||||
if isinstance(content, dict):
|
||||
if "text" in content and isinstance(content["text"], str):
|
||||
return content["text"]
|
||||
parts = []
|
||||
for value in content.values():
|
||||
text = self._blocks_to_text(value)
|
||||
if text:
|
||||
parts.append(text)
|
||||
return " ".join(parts)
|
||||
try:
|
||||
return json.dumps(content, ensure_ascii=False)
|
||||
except TypeError:
|
||||
return ""
|
||||
|
||||
|
||||
simple_ai_service = SimpleAIService()
|
||||
Reference in New Issue
Block a user