refactor(knowledge): 降低知识库路由复杂度
拆分示例问题生成逻辑并收敛 URL 入库与上传限制,修复 uid、MIME 与结果统计等路由问题。
This commit is contained in:
parent
acf150d47d
commit
afdb5fbb02
@ -1,13 +1,10 @@
|
||||
"""知识库工具模块
|
||||
|
||||
包含知识库相关的工具函数:
|
||||
- kb_utils: 知识库通用工具函数
|
||||
- indexing: 文件处理和索引相关功能
|
||||
"""
|
||||
"""知识库工具模块。"""
|
||||
|
||||
from .kb_utils import (
|
||||
calculate_content_hash,
|
||||
is_minio_url,
|
||||
merge_processing_params,
|
||||
parse_minio_url,
|
||||
prepare_item_metadata,
|
||||
resolve_processing_params,
|
||||
sanitize_processing_params,
|
||||
@ -15,8 +12,10 @@ from .kb_utils import (
|
||||
|
||||
__all__ = [
|
||||
"calculate_content_hash",
|
||||
"is_minio_url",
|
||||
"merge_processing_params",
|
||||
"parse_minio_url",
|
||||
"prepare_item_metadata",
|
||||
"resolve_processing_params",
|
||||
"sanitize_processing_params",
|
||||
"merge_processing_params",
|
||||
]
|
||||
|
||||
@ -51,11 +51,11 @@ async def calculate_content_hash(data: bytes | bytearray) -> str:
|
||||
|
||||
async def prepare_item_metadata(item: str, content_type: str, kb_id: str, params: dict | None = None) -> dict:
|
||||
"""
|
||||
准备文件或URL的元数据,文件来源必须是 MinIO URL。
|
||||
准备 MinIO 文件元数据;URL 导入需先通过 fetch-url 预处理为 MinIO 文件。
|
||||
|
||||
Args:
|
||||
item: MinIO URL 或 URL
|
||||
content_type: 内容类型 ("file" 或 "url")
|
||||
item: MinIO URL
|
||||
content_type: 内容类型,目前仅支持 "file"
|
||||
kb_id: 数据库ID
|
||||
params: 处理参数,可选
|
||||
"""
|
||||
@ -133,16 +133,6 @@ async def prepare_item_metadata(item: str, content_type: str, kb_id: str, params
|
||||
file_size = file_sizes.get(item)
|
||||
file_id = f"file_{hashstr(str(item_path) + str(time.time()), 6)}"
|
||||
|
||||
elif content_type == "url":
|
||||
# URL 处理
|
||||
filename = item # 使用完整 URL 作为文件名
|
||||
filename_display = item
|
||||
file_type = "url"
|
||||
item_path = item
|
||||
content_hash = None # URL 没有 content_hash
|
||||
file_size = None
|
||||
file_id = f"url_{hashstr(item + str(time.time()), 6)}"
|
||||
|
||||
else:
|
||||
raise ValueError(f"Unsupported content_type: {content_type}")
|
||||
|
||||
@ -191,7 +181,6 @@ def merge_processing_params(metadata_params: dict | None, request_params: dict |
|
||||
return merged_params
|
||||
|
||||
|
||||
|
||||
def is_minio_url(file_path: str) -> bool:
|
||||
"""检测是否是本系统生成的 MinIO 存储 URL。"""
|
||||
from urllib.parse import urlparse
|
||||
|
||||
@ -4,6 +4,8 @@ import json
|
||||
import textwrap
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from yuxi import config, knowledge_base
|
||||
from yuxi.models import select_model
|
||||
from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository
|
||||
@ -63,18 +65,6 @@ MINDMAP_SYSTEM_PROMPT = """你是一个专业的知识整理助手。
|
||||
"""
|
||||
|
||||
|
||||
class MindmapNotFoundError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
class MindmapValidationError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
class MindmapGenerationError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
def build_database_file_list(files: dict[str, dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
return [
|
||||
{
|
||||
@ -136,7 +126,7 @@ def parse_mindmap_content(content: str) -> dict[str, Any]:
|
||||
async def get_mindmap_database_files(kb_id: str) -> dict[str, Any]:
|
||||
db_info = await knowledge_base.get_database_info(kb_id)
|
||||
if not db_info:
|
||||
raise MindmapNotFoundError(f"知识库 {kb_id} 不存在")
|
||||
raise HTTPException(status_code=404, detail=f"知识库 {kb_id} 不存在")
|
||||
|
||||
file_list = build_database_file_list(db_info.get("files", {}))
|
||||
return {
|
||||
@ -154,13 +144,13 @@ async def generate_database_mindmap(
|
||||
) -> dict[str, Any]:
|
||||
db_info = await knowledge_base.get_database_info(kb_id)
|
||||
if not db_info:
|
||||
raise MindmapNotFoundError(f"知识库 {kb_id} 不存在")
|
||||
raise HTTPException(status_code=404, detail=f"知识库 {kb_id} 不存在")
|
||||
|
||||
db_name = db_info.get("name", "知识库")
|
||||
all_files = db_info.get("files", {})
|
||||
selected_file_ids = list(file_ids or all_files.keys())
|
||||
if not selected_file_ids:
|
||||
raise MindmapValidationError("知识库中没有文件")
|
||||
raise HTTPException(status_code=400, detail="知识库中没有文件")
|
||||
|
||||
original_count = len(selected_file_ids)
|
||||
if len(selected_file_ids) > 20:
|
||||
@ -169,7 +159,7 @@ async def generate_database_mindmap(
|
||||
|
||||
files_info = collect_mindmap_files(all_files, selected_file_ids)
|
||||
if not files_info:
|
||||
raise MindmapValidationError("选择的文件不存在")
|
||||
raise HTTPException(status_code=400, detail="选择的文件不存在")
|
||||
|
||||
logger.info(f"开始生成思维导图,知识库: {db_name}, 文件数量: {len(files_info)}")
|
||||
|
||||
@ -185,7 +175,7 @@ async def generate_database_mindmap(
|
||||
mindmap_data = parse_mindmap_content(content)
|
||||
except ValueError as e:
|
||||
logger.error(f"AI返回的JSON解析失败: {e}, 原始内容: {content}")
|
||||
raise MindmapGenerationError(f"AI返回格式错误: {str(e)}") from e
|
||||
raise HTTPException(status_code=500, detail=f"AI返回格式错误: {str(e)}") from e
|
||||
|
||||
logger.info("思维导图生成成功")
|
||||
|
||||
@ -234,7 +224,7 @@ async def get_mindmap_databases_overview(uid: str) -> dict[str, Any]:
|
||||
async def get_database_mindmap_data(kb_id: str) -> dict[str, Any]:
|
||||
kb = await KnowledgeBaseRepository().get_by_kb_id(kb_id)
|
||||
if kb is None:
|
||||
raise MindmapNotFoundError(f"知识库 {kb_id} 不存在")
|
||||
raise HTTPException(status_code=404, detail=f"知识库 {kb_id} 不存在")
|
||||
|
||||
return {
|
||||
"message": "success",
|
||||
|
||||
142
backend/package/yuxi/knowledge/utils/sample_question_utils.py
Normal file
142
backend/package/yuxi/knowledge/utils/sample_question_utils.py
Normal file
@ -0,0 +1,142 @@
|
||||
"""知识库示例问题生成工具。"""
|
||||
|
||||
import json
|
||||
import textwrap
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
from yuxi import config, knowledge_base
|
||||
from yuxi.knowledge.factory import KnowledgeBaseFactory
|
||||
from yuxi.models import select_model
|
||||
from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository
|
||||
from yuxi.utils import logger
|
||||
|
||||
SAMPLE_QUESTIONS_SYSTEM_PROMPT = """你是一个专业的知识库问答测试专家。
|
||||
|
||||
你的任务是根据知识库中的文件列表,生成有价值的测试问题。
|
||||
|
||||
要求:
|
||||
1. 问题要具体、有针对性,基于文件名称和类型推测可能的内容
|
||||
2. 问题要涵盖不同方面和难度
|
||||
3. 问题要简洁明了,适合用于检索测试
|
||||
4. 问题要多样化,包括事实查询、概念解释、操作指导等
|
||||
5. 问题长度控制在10-30字之间
|
||||
6. 直接返回JSON数组格式,不要其他说明
|
||||
|
||||
返回格式:
|
||||
```json
|
||||
{
|
||||
"questions": [
|
||||
"问题1?",
|
||||
"问题2?",
|
||||
"问题3?"
|
||||
]
|
||||
}
|
||||
```
|
||||
"""
|
||||
|
||||
|
||||
def build_sample_question_file_list(files: dict[str, dict[str, Any]]) -> list[dict[str, str]]:
|
||||
return [
|
||||
{
|
||||
"filename": file_info.get("filename", ""),
|
||||
"type": file_info.get("type") or file_info.get("file_type", ""),
|
||||
}
|
||||
for file_info in files.values()
|
||||
]
|
||||
|
||||
|
||||
def build_sample_questions_user_message(db_name: str, files_info: list[dict[str, str]], count: int) -> str:
|
||||
files_text = "\n".join([f"- {file_info['filename']} ({file_info['type']})" for file_info in files_info[:20]])
|
||||
file_count_text = f"(共{len(files_info)}个文件)" if len(files_info) > 20 else ""
|
||||
|
||||
return textwrap.dedent(f"""请为知识库\"{db_name}\"生成{count}个测试问题。
|
||||
|
||||
知识库文件列表{file_count_text}:
|
||||
{files_text}
|
||||
|
||||
请根据这些文件的名称和类型,生成{count}个有价值的测试问题。""")
|
||||
|
||||
|
||||
def parse_sample_questions_content(content: str) -> list[str]:
|
||||
if "```json" in content:
|
||||
json_start = content.find("```json") + 7
|
||||
json_end = content.find("```", json_start)
|
||||
if json_end == -1:
|
||||
raise ValueError("AI返回的JSON代码块不完整")
|
||||
content = content[json_start:json_end].strip()
|
||||
elif "```" in content:
|
||||
json_start = content.find("```") + 3
|
||||
json_end = content.find("```", json_start)
|
||||
if json_end == -1:
|
||||
raise ValueError("AI返回的代码块不完整")
|
||||
content = content[json_start:json_end].strip()
|
||||
|
||||
questions_data = json.loads(content)
|
||||
questions = questions_data.get("questions", []) if isinstance(questions_data, dict) else []
|
||||
if not questions or not isinstance(questions, list):
|
||||
raise ValueError("AI返回的问题格式不正确")
|
||||
return questions
|
||||
|
||||
|
||||
async def generate_database_sample_questions(kb_id: str, count: int = 10) -> dict[str, Any]:
|
||||
db_info = await knowledge_base.get_database_info(kb_id)
|
||||
if not db_info:
|
||||
raise HTTPException(status_code=404, detail=f"知识库 {kb_id} 不存在")
|
||||
|
||||
kb_type = (db_info.get("kb_type") or "").lower()
|
||||
if not KnowledgeBaseFactory.get_kb_class(kb_type).supports_documents:
|
||||
raise HTTPException(status_code=400, detail=f"{db_info.get('name') or kb_type} 不支持基于文件生成测试问题")
|
||||
|
||||
db_name = db_info.get("name", "")
|
||||
all_files = db_info.get("files", {})
|
||||
if not all_files:
|
||||
raise HTTPException(status_code=400, detail="知识库中没有文件")
|
||||
|
||||
files_info = build_sample_question_file_list(all_files)
|
||||
logger.info(f"开始生成知识库问题,知识库: {db_name}, 文件数量: {len(files_info)}, 问题数量: {count}")
|
||||
|
||||
model = select_model(model_spec=config.default_model)
|
||||
messages = [
|
||||
{"role": "system", "content": SAMPLE_QUESTIONS_SYSTEM_PROMPT},
|
||||
{"role": "user", "content": build_sample_questions_user_message(db_name, files_info, count)},
|
||||
]
|
||||
response = await model.call(messages, stream=False)
|
||||
content = response.content if hasattr(response, "content") else str(response)
|
||||
|
||||
try:
|
||||
questions = parse_sample_questions_content(content)
|
||||
except (json.JSONDecodeError, ValueError) as e:
|
||||
logger.error(f"AI返回的JSON解析失败: {e}, 原始内容: {content}")
|
||||
raise HTTPException(status_code=500, detail=f"AI返回格式错误: {str(e)}") from e
|
||||
|
||||
logger.info(f"成功生成{len(questions)}个问题")
|
||||
|
||||
try:
|
||||
await KnowledgeBaseRepository().update(kb_id, {"sample_questions": questions})
|
||||
logger.info(f"成功保存 {len(questions)} 个问题到知识库 {kb_id}")
|
||||
except Exception as save_error:
|
||||
logger.error(f"保存问题失败: {save_error}")
|
||||
|
||||
return {
|
||||
"message": "success",
|
||||
"questions": questions,
|
||||
"count": len(questions),
|
||||
"kb_id": kb_id,
|
||||
"db_name": db_name,
|
||||
}
|
||||
|
||||
|
||||
async def get_database_sample_questions(kb_id: str) -> dict[str, Any]:
|
||||
kb = await KnowledgeBaseRepository().get_by_kb_id(kb_id)
|
||||
if kb is None:
|
||||
raise HTTPException(status_code=404, detail=f"知识库 {kb_id} 不存在")
|
||||
|
||||
questions = kb.sample_questions or []
|
||||
return {
|
||||
"message": "success",
|
||||
"questions": questions,
|
||||
"count": len(questions),
|
||||
"kb_id": kb_id,
|
||||
}
|
||||
@ -1,6 +1,6 @@
|
||||
import ipaddress
|
||||
import socket
|
||||
from urllib.parse import urlparse
|
||||
from urllib.parse import urljoin, urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
@ -93,19 +93,7 @@ async def fetch_url_content(url: str, max_size: int = MAX_DOWNLOAD_SIZE) -> tupl
|
||||
if not location:
|
||||
raise ValueError("Redirect response missing Location header")
|
||||
|
||||
# Handle relative redirects
|
||||
if location.startswith("/"):
|
||||
parsed_current = urlparse(current_url)
|
||||
current_url = f"{parsed_current.scheme}://{parsed_current.netloc}{location}"
|
||||
elif not location.startswith("http"):
|
||||
# Handle relative path without /? or other weird cases, or assume absolute
|
||||
parsed_current = urlparse(current_url)
|
||||
# simple join
|
||||
from urllib.parse import urljoin
|
||||
|
||||
current_url = urljoin(current_url, location)
|
||||
else:
|
||||
current_url = location
|
||||
current_url = urljoin(current_url, location)
|
||||
|
||||
# Validate the new URL
|
||||
is_valid, error_msg = validate_url(current_url)
|
||||
|
||||
@ -6,6 +6,8 @@ from fastapi import UploadFile
|
||||
|
||||
from yuxi.storage.minio import aupload_file_to_minio
|
||||
|
||||
MAX_UPLOAD_SIZE_BYTES = 100 * 1024 * 1024
|
||||
|
||||
|
||||
async def write_upload_to_buffer(
|
||||
upload: UploadFile,
|
||||
|
||||
@ -12,14 +12,14 @@ import aiofiles
|
||||
from fastapi import HTTPException, UploadFile
|
||||
from fastapi.responses import FileResponse, StreamingResponse
|
||||
from yuxi.agents.backends.sandbox.paths import _global_user_data_dir, ensure_workspace_default_files
|
||||
from yuxi.services.upload_utils import write_upload_to_buffer
|
||||
from yuxi.services.upload_utils import MAX_UPLOAD_SIZE_BYTES, write_upload_to_buffer
|
||||
from yuxi.services.viewer_filesystem_service import _detect_preview_type
|
||||
from yuxi.storage.postgres.models_business import User
|
||||
from yuxi.utils.datetime_utils import utc_isoformat_from_timestamp
|
||||
from yuxi.utils.paths import VIRTUAL_PATH_WORKSPACE, WORKSPACE_DIR_NAME
|
||||
|
||||
EDITABLE_WORKSPACE_SUFFIXES = {".md", ".markdown", ".mdx", ".txt"}
|
||||
MAX_WORKSPACE_UPLOAD_SIZE_BYTES = 100 * 1024 * 1024
|
||||
MAX_WORKSPACE_UPLOAD_SIZE_BYTES = MAX_UPLOAD_SIZE_BYTES
|
||||
|
||||
|
||||
def _workspace_root(user: User) -> Path:
|
||||
|
||||
@ -1,13 +1,11 @@
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import textwrap
|
||||
import traceback
|
||||
import time
|
||||
from urllib.parse import quote, unquote
|
||||
|
||||
import aiofiles
|
||||
from fastapi import APIRouter, Body, Depends, File, HTTPException, Query, Request, UploadFile
|
||||
from fastapi import APIRouter, Body, Depends, File, HTTPException, Query, UploadFile
|
||||
from pydantic import BaseModel
|
||||
from fastapi.responses import FileResponse
|
||||
from starlette.responses import StreamingResponse
|
||||
@ -18,18 +16,20 @@ from yuxi import config, knowledge_base
|
||||
from yuxi.knowledge.factory import KnowledgeBaseFactory
|
||||
from yuxi.knowledge.graphs.milvus_graph_service import GRAPH_TASK_TYPE, MilvusGraphService
|
||||
from yuxi.plugins.parser import Parser, SUPPORTED_FILE_EXTENSIONS, is_supported_file_extension
|
||||
from yuxi.knowledge.utils import calculate_content_hash
|
||||
from yuxi.knowledge.utils.kb_utils import is_minio_url, parse_minio_url
|
||||
from yuxi.knowledge.utils import calculate_content_hash, is_minio_url, parse_minio_url
|
||||
from yuxi.knowledge.utils.mindmap_utils import (
|
||||
MindmapGenerationError,
|
||||
MindmapNotFoundError,
|
||||
MindmapValidationError,
|
||||
generate_database_mindmap,
|
||||
get_database_mindmap_data,
|
||||
get_mindmap_database_files,
|
||||
get_mindmap_databases_overview,
|
||||
)
|
||||
from yuxi.knowledge.utils.sample_question_utils import (
|
||||
generate_database_sample_questions,
|
||||
get_database_sample_questions,
|
||||
)
|
||||
from yuxi.knowledge.utils.url_fetcher import fetch_url_content
|
||||
from yuxi.services.model_cache import model_cache
|
||||
from yuxi.services.upload_utils import MAX_UPLOAD_SIZE_BYTES, read_upload_with_limit, write_upload_to_path
|
||||
from yuxi.services.workspace_service import MAX_WORKSPACE_UPLOAD_SIZE_BYTES, resolve_workspace_file_path
|
||||
from yuxi.storage.postgres.models_business import User
|
||||
from yuxi.storage.minio.client import MinIOClient, StorageError, aupload_file_to_minio, get_minio_client
|
||||
@ -147,9 +147,9 @@ async def create_database(
|
||||
description: str = Body(...),
|
||||
embedding_model_spec: str | None = Body(None),
|
||||
kb_type: str = Body("milvus"),
|
||||
additional_params: dict = Body({}),
|
||||
additional_params: dict | None = Body(None),
|
||||
llm_model_spec: str | None = Body(None),
|
||||
share_config: dict = Body(None),
|
||||
share_config: dict | None = Body(None),
|
||||
current_user: User = Depends(get_admin_user),
|
||||
):
|
||||
"""创建知识库"""
|
||||
@ -238,24 +238,16 @@ async def get_accessible_databases(current_user: User = Depends(get_required_use
|
||||
return {"message": f"获取可访问知识库列表失败: {str(e)}", "databases": []}
|
||||
|
||||
|
||||
def _raise_mindmap_http_exception(error: Exception, operation: str) -> None:
|
||||
if isinstance(error, MindmapNotFoundError):
|
||||
raise HTTPException(status_code=404, detail=str(error))
|
||||
if isinstance(error, MindmapValidationError):
|
||||
raise HTTPException(status_code=400, detail=str(error))
|
||||
if isinstance(error, MindmapGenerationError):
|
||||
raise HTTPException(status_code=500, detail=str(error))
|
||||
logger.error(f"{operation}失败: {error}, {traceback.format_exc()}")
|
||||
raise HTTPException(status_code=500, detail=f"{operation}失败: {str(error)}")
|
||||
|
||||
|
||||
@knowledge.get("/mindmap/databases")
|
||||
async def get_mindmap_databases(current_user: User = Depends(get_admin_user)):
|
||||
"""获取所有知识库的概览信息,用于思维导图界面选择。"""
|
||||
try:
|
||||
return await get_mindmap_databases_overview(current_user.uid)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
_raise_mindmap_http_exception(e, "获取知识库列表")
|
||||
logger.error(f"获取知识库列表失败: {e}, {traceback.format_exc()}")
|
||||
raise HTTPException(status_code=500, detail=f"获取知识库列表失败: {str(e)}")
|
||||
|
||||
|
||||
@knowledge.get("/databases/{kb_id}/mindmap/files")
|
||||
@ -263,8 +255,11 @@ async def get_database_mindmap_files(kb_id: str, current_user: User = Depends(ge
|
||||
"""获取指定知识库的所有文件列表。"""
|
||||
try:
|
||||
return await get_mindmap_database_files(kb_id)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
_raise_mindmap_http_exception(e, "获取文件列表")
|
||||
logger.error(f"获取文件列表失败: {e}, {traceback.format_exc()}")
|
||||
raise HTTPException(status_code=500, detail=f"获取文件列表失败: {str(e)}")
|
||||
|
||||
|
||||
@knowledge.post("/databases/{kb_id}/mindmap/generate")
|
||||
@ -277,8 +272,11 @@ async def generate_mindmap(
|
||||
"""使用 AI 分析知识库文件,生成思维导图结构。"""
|
||||
try:
|
||||
return await generate_database_mindmap(kb_id, file_ids, user_prompt)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
_raise_mindmap_http_exception(e, "生成思维导图")
|
||||
logger.error(f"生成思维导图失败: {e}, {traceback.format_exc()}")
|
||||
raise HTTPException(status_code=500, detail=f"生成思维导图失败: {str(e)}")
|
||||
|
||||
|
||||
@knowledge.get("/databases/{kb_id}/mindmap")
|
||||
@ -286,8 +284,11 @@ async def get_database_mindmap(kb_id: str, current_user: User = Depends(get_admi
|
||||
"""获取知识库关联的思维导图。"""
|
||||
try:
|
||||
return await get_database_mindmap_data(kb_id)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
_raise_mindmap_http_exception(e, "获取知识库思维导图")
|
||||
logger.error(f"获取知识库思维导图失败: {e}, {traceback.format_exc()}")
|
||||
raise HTTPException(status_code=500, detail=f"获取知识库思维导图失败: {str(e)}")
|
||||
|
||||
|
||||
@knowledge.get("/databases/{kb_id}")
|
||||
@ -342,6 +343,8 @@ async def update_database_info(
|
||||
operator_department_id=current_user.department_id,
|
||||
)
|
||||
return {"message": "更新成功", "database": database}
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"更新数据库失败 {e}, {traceback.format_exc()}")
|
||||
raise HTTPException(status_code=400, detail=f"更新数据库失败: {e}")
|
||||
@ -401,9 +404,10 @@ async def configure_graph_build(
|
||||
@knowledge.post("/databases/{kb_id}/graph-build/index")
|
||||
async def index_graph_build(
|
||||
kb_id: str,
|
||||
data: dict = Body(default={}),
|
||||
data: dict | None = Body(default=None),
|
||||
current_user: User = Depends(get_admin_user),
|
||||
):
|
||||
data = data or {}
|
||||
try:
|
||||
if await _has_running_graph_build_task(kb_id):
|
||||
raise HTTPException(status_code=409, detail="该知识库已有正在运行的图谱构建任务")
|
||||
@ -449,9 +453,10 @@ async def index_graph_build(
|
||||
@knowledge.post("/databases/{kb_id}/graph-build/reset")
|
||||
async def reset_graph_build(
|
||||
kb_id: str,
|
||||
data: dict = Body(default={}),
|
||||
data: dict | None = Body(default=None),
|
||||
current_user: User = Depends(get_admin_user),
|
||||
):
|
||||
data = data or {}
|
||||
try:
|
||||
if await _has_running_graph_build_task(kb_id):
|
||||
raise HTTPException(status_code=409, detail="该知识库存在正在运行的图谱构建任务,无法重置")
|
||||
@ -485,9 +490,11 @@ async def export_database(
|
||||
if not os.path.exists(file_path):
|
||||
raise HTTPException(status_code=404, detail="Exported file not found.")
|
||||
|
||||
media_type = media_types.get(format, "application/octet-stream")
|
||||
media_type = media_types.get(f".{format}", "application/octet-stream")
|
||||
|
||||
return FileResponse(path=file_path, filename=os.path.basename(file_path), media_type=media_type)
|
||||
except HTTPException:
|
||||
raise
|
||||
except NotImplementedError as e:
|
||||
logger.warning(f"A disabled feature was accessed: {e}")
|
||||
raise HTTPException(status_code=501, detail=str(e))
|
||||
@ -521,148 +528,142 @@ async def add_documents(
|
||||
if isinstance(chunk_parser_config, dict):
|
||||
indexing_params["chunk_parser_config"] = chunk_parser_config
|
||||
|
||||
# URL 解析与入库(需白名单验证)
|
||||
if content_type == "url":
|
||||
raise HTTPException(status_code=400, detail="URL 处理方式已变更,请使用 fetch-url 接口先获取内容")
|
||||
if content_type != "file":
|
||||
raise HTTPException(status_code=400, detail=f"Unsupported content_type: {content_type}")
|
||||
|
||||
if content_type == "file":
|
||||
for item in items:
|
||||
if not is_minio_url(item):
|
||||
raise HTTPException(status_code=400, detail="File source must be a MinIO URL")
|
||||
for item in items:
|
||||
if not is_minio_url(item):
|
||||
raise HTTPException(status_code=400, detail="File source must be a MinIO URL")
|
||||
|
||||
async def run_ingest(context: TaskContext):
|
||||
await context.set_message("任务初始化")
|
||||
await context.set_progress(5.0, "准备处理文档")
|
||||
|
||||
total = len(items)
|
||||
processed_items = []
|
||||
|
||||
# 存储第一阶段成功添加的文件记录 {item: (file_id, file_meta)}
|
||||
added_files = {}
|
||||
processed_items: list[dict | None] = [None] * total
|
||||
added_files: list[dict] = []
|
||||
|
||||
try:
|
||||
# ========== 第一阶段:批量添加文件记录 ==========
|
||||
await context.set_message("第一阶段:添加文件记录")
|
||||
for idx, item in enumerate(items, 1):
|
||||
await context.raise_if_cancelled()
|
||||
|
||||
# 第一阶段进度:5% ~ 30%
|
||||
progress = 5.0 + (idx / total) * 25.0
|
||||
await context.set_progress(progress, f"[1/2] 添加记录 {idx}/{total}")
|
||||
await context.set_progress(progress, f"[1/3] 添加记录 {idx}/{total}")
|
||||
|
||||
try:
|
||||
# 1. Add file record (UPLOADED)
|
||||
file_meta = await knowledge_base.add_file_record(
|
||||
kb_id, item, params=params, operator_id=current_user.uid
|
||||
)
|
||||
file_id = file_meta["file_id"]
|
||||
added_files[item] = (file_id, file_meta)
|
||||
added_files.append(
|
||||
{
|
||||
"index": idx - 1,
|
||||
"item": item,
|
||||
"file_id": file_meta["file_id"],
|
||||
"file_meta": file_meta,
|
||||
}
|
||||
)
|
||||
except Exception as add_error:
|
||||
logger.error(f"添加文件记录失败 {item}: {add_error}")
|
||||
error_type = "timeout" if isinstance(add_error, TimeoutError) else "add_failed"
|
||||
error_msg = "添加超时" if isinstance(add_error, TimeoutError) else "添加记录失败"
|
||||
processed_items.append(
|
||||
{
|
||||
"item": item,
|
||||
"status": "failed",
|
||||
"error": f"{error_msg}: {str(add_error)}",
|
||||
"error_type": error_type,
|
||||
}
|
||||
)
|
||||
processed_items[idx - 1] = {
|
||||
"item": item,
|
||||
"status": "failed",
|
||||
"error": f"{error_msg}: {str(add_error)}",
|
||||
"error_type": error_type,
|
||||
}
|
||||
|
||||
# ========== 第二阶段:批量解析文件 ==========
|
||||
await context.set_message("第二阶段:解析文件")
|
||||
parse_success_count = 0
|
||||
# 计算解析阶段的进度范围
|
||||
parse_progress_range = 30.0 if not auto_index else 25.0
|
||||
|
||||
for idx, (item, (file_id, add_file_meta)) in enumerate(added_files.items(), 1):
|
||||
parse_end = 60.0 if auto_index else 95.0
|
||||
parse_total = len(added_files)
|
||||
for idx, record in enumerate(added_files, 1):
|
||||
await context.raise_if_cancelled()
|
||||
|
||||
# 第二阶段进度:25%~55% 或 30%~60%
|
||||
progress = parse_progress_range + (idx / len(added_files)) * 30.0
|
||||
await context.set_progress(progress, f"[2/2] 解析文件 {idx}/{len(added_files)}")
|
||||
progress = 30.0 + (idx / parse_total) * (parse_end - 30.0)
|
||||
await context.set_progress(progress, f"[2/3] 解析文件 {idx}/{parse_total}")
|
||||
|
||||
item = record["item"]
|
||||
file_id = record["file_id"]
|
||||
try:
|
||||
# 2. Parse file (PARSING -> PARSED)
|
||||
file_meta = await knowledge_base.parse_file(kb_id, file_id, operator_id=current_user.uid)
|
||||
added_files[item] = (file_id, file_meta)
|
||||
processed_items.append(file_meta)
|
||||
parse_success_count += 1
|
||||
record["file_meta"] = file_meta
|
||||
if not auto_index or file_meta.get("status") != "parsed":
|
||||
processed_items[record["index"]] = file_meta
|
||||
except Exception as parse_error:
|
||||
logger.error(f"解析文件失败 {item} (file_id={file_id}): {parse_error}")
|
||||
error_type = "timeout" if isinstance(parse_error, TimeoutError) else "parse_failed"
|
||||
error_msg = "解析超时" if isinstance(parse_error, TimeoutError) else "解析失败"
|
||||
processed_items.append(
|
||||
{
|
||||
"item": item,
|
||||
"status": "failed",
|
||||
"error": f"{error_msg}: {str(parse_error)}",
|
||||
"error_type": error_type,
|
||||
}
|
||||
)
|
||||
processed_items[record["index"]] = {
|
||||
"item": item,
|
||||
"status": "failed",
|
||||
"error": f"{error_msg}: {str(parse_error)}",
|
||||
"error_type": error_type,
|
||||
}
|
||||
|
||||
# ========== 第三阶段:自动入库 ==========
|
||||
if auto_index:
|
||||
await context.set_message("第三阶段:自动入库")
|
||||
parsed_files = [(item, data) for item, data in added_files.items() if data[1].get("status") == "parsed"]
|
||||
parsed_files = [record for record in added_files if record["file_meta"].get("status") == "parsed"]
|
||||
total_parsed = len(parsed_files)
|
||||
|
||||
for idx, (item, (file_id, file_meta)) in enumerate(parsed_files, 1):
|
||||
for idx, record in enumerate(parsed_files, 1):
|
||||
await context.raise_if_cancelled()
|
||||
|
||||
# 第三阶段进度:55%~95% 或 60%~95%
|
||||
progress = 55.0 + (idx / total_parsed) * 40.0
|
||||
progress = 60.0 + (idx / total_parsed) * 35.0
|
||||
await context.set_progress(progress, f"[3/3] 入库文件 {idx}/{total_parsed}")
|
||||
|
||||
item = record["item"]
|
||||
file_id = record["file_id"]
|
||||
try:
|
||||
# 1. 更新入库参数
|
||||
await knowledge_base.update_file_params(
|
||||
kb_id, file_id, indexing_params, operator_id=current_user.uid
|
||||
)
|
||||
# 2. 执行入库(传入 indexing_params 确保使用的参数与用户设置一致)
|
||||
result = await knowledge_base.index_file(
|
||||
kb_id, file_id, operator_id=current_user.uid, params=indexing_params
|
||||
)
|
||||
processed_items.append(result)
|
||||
processed_items[record["index"]] = result
|
||||
except Exception as index_error:
|
||||
logger.error(f"自动入库失败 {item} (file_id={file_id}): {index_error}")
|
||||
processed_items.append(
|
||||
{
|
||||
"item": item,
|
||||
"status": "failed",
|
||||
"error": f"入库失败: {str(index_error)}",
|
||||
"error_type": "index_failed",
|
||||
}
|
||||
)
|
||||
processed_items[record["index"]] = {
|
||||
"item": item,
|
||||
"status": "failed",
|
||||
"error": f"入库失败: {str(index_error)}",
|
||||
"error_type": "index_failed",
|
||||
}
|
||||
|
||||
except asyncio.CancelledError:
|
||||
await context.set_progress(100.0, "任务已取消")
|
||||
raise
|
||||
except Exception as task_error:
|
||||
# 处理整体任务的其他异常(如内存不足、网络错误等)
|
||||
logger.exception(f"Task processing failed: {task_error}")
|
||||
await context.set_progress(100.0, f"任务处理失败: {str(task_error)}")
|
||||
# 注意:不需要手动标记未处理的文件为失败,因为:
|
||||
# 1. 内层异常处理已记录所有处理过的文件(成功/失败)
|
||||
# 2. 未处理的文件没有进入 processed_items,前端会正确显示
|
||||
# 3. 用户可以重新提交未处理的文件
|
||||
raise
|
||||
|
||||
item_type = "URL" if content_type == "url" else "文件"
|
||||
# Check for failed status (including ERROR_PARSING)
|
||||
failed_count = len([_p for _p in processed_items if "error" in _p or _p.get("status") == "failed"])
|
||||
final_items = [
|
||||
item
|
||||
if item is not None
|
||||
else {
|
||||
"item": items[index],
|
||||
"status": "failed",
|
||||
"error": "文件未处理",
|
||||
"error_type": "not_processed",
|
||||
}
|
||||
for index, item in enumerate(processed_items)
|
||||
]
|
||||
failed_count = len([item for item in final_items if "error" in item or item.get("status") == "failed"])
|
||||
|
||||
summary = {
|
||||
"kb_id": kb_id,
|
||||
"item_type": item_type,
|
||||
"submitted": len(processed_items),
|
||||
"item_type": "文件",
|
||||
"submitted": total,
|
||||
"failed": failed_count,
|
||||
}
|
||||
message = f"{item_type}处理完成,失败 {failed_count} 个" if failed_count else f"{item_type}处理完成"
|
||||
await context.set_result(summary | {"items": processed_items})
|
||||
message = f"文件处理完成,失败 {failed_count} 个" if failed_count else "文件处理完成"
|
||||
await context.set_result(summary | {"items": final_items})
|
||||
await context.set_progress(100.0, message)
|
||||
return summary | {"items": processed_items}
|
||||
return summary | {"items": final_items}
|
||||
|
||||
try:
|
||||
database = await knowledge_base.get_database_info(kb_id)
|
||||
@ -740,15 +741,15 @@ async def parse_documents(kb_id: str, file_ids: list[str] = Body(...), current_u
|
||||
async def index_documents(
|
||||
kb_id: str,
|
||||
file_ids: list[str] = Body(...),
|
||||
params: dict = Body({}),
|
||||
params: dict | None = Body(None),
|
||||
current_user: User = Depends(get_admin_user),
|
||||
):
|
||||
"""手动触发文档入库(Indexing),支持更新参数"""
|
||||
params = params or {}
|
||||
logger.debug(f"Index documents for kb_id {kb_id}: {file_ids} {params=}")
|
||||
await _ensure_database_supports_documents(kb_id, "文档入库")
|
||||
|
||||
# extract operator_id safely before background task
|
||||
operator_id = current_user.id
|
||||
operator_id = current_user.uid
|
||||
|
||||
async def run_index(context: TaskContext):
|
||||
await context.set_message("任务初始化")
|
||||
@ -928,7 +929,7 @@ async def delete_document(kb_id: str, doc_id: str, current_user: User = Depends(
|
||||
|
||||
|
||||
@knowledge.get("/databases/{kb_id}/documents/{doc_id}/download")
|
||||
async def download_document(kb_id: str, doc_id: str, request: Request, current_user: User = Depends(get_admin_user)):
|
||||
async def download_document(kb_id: str, doc_id: str, current_user: User = Depends(get_admin_user)):
|
||||
"""下载原始文件"""
|
||||
logger.debug(f"Download document {doc_id} from {kb_id}")
|
||||
await _ensure_database_supports_documents(kb_id, "文档下载")
|
||||
@ -1080,6 +1081,8 @@ async def update_knowledge_base_query_params(
|
||||
|
||||
return {"message": "success", "data": params}
|
||||
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"更新知识库查询参数失败: {e}")
|
||||
raise HTTPException(status_code=500, detail=f"更新查询参数失败: {str(e)}")
|
||||
@ -1121,148 +1124,16 @@ def _merge_saved_options(params: dict, saved_options: dict) -> dict:
|
||||
# =============================================================================
|
||||
|
||||
|
||||
SAMPLE_QUESTIONS_SYSTEM_PROMPT = """你是一个专业的知识库问答测试专家。
|
||||
|
||||
你的任务是根据知识库中的文件列表,生成有价值的测试问题。
|
||||
|
||||
要求:
|
||||
1. 问题要具体、有针对性,基于文件名称和类型推测可能的内容
|
||||
2. 问题要涵盖不同方面和难度
|
||||
3. 问题要简洁明了,适合用于检索测试
|
||||
4. 问题要多样化,包括事实查询、概念解释、操作指导等
|
||||
5. 问题长度控制在10-30字之间
|
||||
6. 直接返回JSON数组格式,不要其他说明
|
||||
|
||||
返回格式:
|
||||
```json
|
||||
{
|
||||
"questions": [
|
||||
"问题1?",
|
||||
"问题2?",
|
||||
"问题3?"
|
||||
]
|
||||
}
|
||||
```
|
||||
"""
|
||||
|
||||
|
||||
@knowledge.post("/databases/{kb_id}/sample-questions")
|
||||
async def generate_sample_questions(
|
||||
kb_id: str,
|
||||
request_body: dict = Body(...),
|
||||
current_user: User = Depends(get_admin_user),
|
||||
):
|
||||
"""
|
||||
AI生成针对知识库的测试问题
|
||||
|
||||
Args:
|
||||
kb_id: 知识库ID
|
||||
request_body: 请求体,包含 count 字段
|
||||
|
||||
Returns:
|
||||
生成的问题列表
|
||||
"""
|
||||
"""AI生成针对知识库的测试问题。"""
|
||||
try:
|
||||
db_info = await knowledge_base.get_database_info(kb_id)
|
||||
if not db_info:
|
||||
raise HTTPException(status_code=404, detail=f"知识库 {kb_id} 不存在")
|
||||
kb_type = (db_info.get("kb_type") or "").lower()
|
||||
if not KnowledgeBaseFactory.get_kb_class(kb_type).supports_documents:
|
||||
raise HTTPException(status_code=400, detail=f"{db_info.get('name') or kb_type} 不支持基于文件生成测试问题")
|
||||
|
||||
from yuxi.models import select_model
|
||||
|
||||
# 从请求体中提取参数
|
||||
count = request_body.get("count", 10)
|
||||
|
||||
db_name = db_info.get("name", "")
|
||||
all_files = db_info.get("files", {})
|
||||
|
||||
if not all_files:
|
||||
raise HTTPException(status_code=400, detail="知识库中没有文件")
|
||||
|
||||
# 收集文件信息
|
||||
files_info = []
|
||||
for file_id, file_info in all_files.items():
|
||||
files_info.append(
|
||||
{
|
||||
"filename": file_info.get("filename", ""),
|
||||
"type": file_info.get("type", ""),
|
||||
}
|
||||
)
|
||||
|
||||
# 构建AI提示词
|
||||
system_prompt = SAMPLE_QUESTIONS_SYSTEM_PROMPT
|
||||
|
||||
# 构建用户消息
|
||||
files_text = "\n".join(
|
||||
[
|
||||
f"- {f['filename']} ({f['type']})"
|
||||
for f in files_info[:20] # 最多列举20个文件
|
||||
]
|
||||
)
|
||||
|
||||
file_count_text = f"(共{len(files_info)}个文件)" if len(files_info) > 20 else ""
|
||||
|
||||
user_message = textwrap.dedent(f"""请为知识库"{db_name}"生成{count}个测试问题。
|
||||
|
||||
知识库文件列表{file_count_text}:
|
||||
{files_text}
|
||||
|
||||
请根据这些文件的名称和类型,生成{count}个有价值的测试问题。""")
|
||||
|
||||
# 调用AI生成
|
||||
logger.info(f"开始生成知识库问题,知识库: {db_name}, 文件数量: {len(files_info)}, 问题数量: {count}")
|
||||
|
||||
# 选择模型并调用
|
||||
model = select_model(model_spec=config.default_model)
|
||||
messages = [{"role": "system", "content": system_prompt}, {"role": "user", "content": user_message}]
|
||||
response = await model.call(messages, stream=False)
|
||||
|
||||
# 解析AI返回的JSON
|
||||
try:
|
||||
# 提取JSON内容
|
||||
content = response.content if hasattr(response, "content") else str(response)
|
||||
|
||||
# 尝试从markdown代码块中提取JSON
|
||||
if "```json" in content:
|
||||
json_start = content.find("```json") + 7
|
||||
json_end = content.find("```", json_start)
|
||||
content = content[json_start:json_end].strip()
|
||||
elif "```" in content:
|
||||
json_start = content.find("```") + 3
|
||||
json_end = content.find("```", json_start)
|
||||
content = content[json_start:json_end].strip()
|
||||
|
||||
questions_data = json.loads(content)
|
||||
questions = questions_data.get("questions", [])
|
||||
|
||||
if not questions or not isinstance(questions, list):
|
||||
raise ValueError("AI返回的问题格式不正确")
|
||||
|
||||
logger.info(f"成功生成{len(questions)}个问题")
|
||||
|
||||
# 保存问题到知识库元数据
|
||||
try:
|
||||
from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository
|
||||
|
||||
await KnowledgeBaseRepository().update(kb_id, {"sample_questions": questions})
|
||||
logger.info(f"成功保存 {len(questions)} 个问题到知识库 {kb_id}")
|
||||
except Exception as save_error:
|
||||
logger.error(f"保存问题失败: {save_error}")
|
||||
|
||||
return {
|
||||
"message": "success",
|
||||
"questions": questions,
|
||||
"count": len(questions),
|
||||
"kb_id": kb_id,
|
||||
"db_name": db_name,
|
||||
}
|
||||
|
||||
except json.JSONDecodeError as e:
|
||||
logger.error(f"AI返回的JSON解析失败: {e}, 原始内容: {content}")
|
||||
raise HTTPException(status_code=500, detail=f"AI返回格式错误: {str(e)}")
|
||||
|
||||
return await generate_database_sample_questions(kb_id, count=count)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
@ -1272,33 +1143,9 @@ async def generate_sample_questions(
|
||||
|
||||
@knowledge.get("/databases/{kb_id}/sample-questions")
|
||||
async def get_sample_questions(kb_id: str, current_user: User = Depends(get_admin_user)):
|
||||
"""
|
||||
获取知识库的测试问题
|
||||
|
||||
Args:
|
||||
kb_id: 知识库ID
|
||||
|
||||
Returns:
|
||||
问题列表
|
||||
"""
|
||||
"""获取知识库的测试问题。"""
|
||||
try:
|
||||
from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository
|
||||
|
||||
kb_repo = KnowledgeBaseRepository()
|
||||
kb = await kb_repo.get_by_kb_id(kb_id)
|
||||
|
||||
if kb is None:
|
||||
raise HTTPException(status_code=404, detail=f"知识库 {kb_id} 不存在")
|
||||
|
||||
questions = kb.sample_questions or []
|
||||
|
||||
return {
|
||||
"message": "success",
|
||||
"questions": questions,
|
||||
"count": len(questions),
|
||||
"kb_id": kb_id,
|
||||
}
|
||||
|
||||
return await get_database_sample_questions(kb_id)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
@ -1322,6 +1169,8 @@ async def create_folder(
|
||||
try:
|
||||
await _ensure_database_supports_documents(kb_id, "文件夹创建")
|
||||
return await knowledge_base.create_folder(kb_id, folder_name, parent_id)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"创建文件夹失败 {e}, {traceback.format_exc()}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
@ -1339,6 +1188,8 @@ async def move_document(
|
||||
try:
|
||||
await _ensure_database_supports_documents(kb_id, "文件移动")
|
||||
return await knowledge_base.move_file(kb_id, doc_id, new_parent_id)
|
||||
except HTTPException:
|
||||
raise
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
except Exception as e:
|
||||
@ -1357,10 +1208,6 @@ async def fetch_url(
|
||||
"""
|
||||
logger.debug(f"Fetching URL: {url} for kb_id: {kb_id}")
|
||||
try:
|
||||
from yuxi.knowledge.utils.url_fetcher import fetch_url_content
|
||||
from yuxi.storage.minio import get_minio_client
|
||||
from yuxi.knowledge.utils import calculate_content_hash
|
||||
|
||||
# 1. 下载内容 (包含白名单校验、大小限制、类型检查)
|
||||
content_bytes, final_url = await fetch_url_content(url)
|
||||
|
||||
@ -1511,7 +1358,14 @@ async def upload_file(
|
||||
# 直接使用原始文件名(小写)
|
||||
filename = f"{basename}{ext}".lower()
|
||||
|
||||
file_bytes = await file.read()
|
||||
try:
|
||||
file_bytes = await read_upload_with_limit(
|
||||
file,
|
||||
max_size_bytes=MAX_UPLOAD_SIZE_BYTES,
|
||||
too_large_message="文件过大,当前仅支持 100 MB 以内的文件",
|
||||
)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
|
||||
content_hash = await calculate_content_hash(file_bytes)
|
||||
|
||||
@ -1523,8 +1377,6 @@ async def upload_file(
|
||||
)
|
||||
|
||||
# 直接上传到MinIO,添加时间戳区分版本
|
||||
import time
|
||||
|
||||
timestamp = int(time.time() * 1000)
|
||||
minio_filename = f"{basename}_{timestamp}{ext}"
|
||||
|
||||
@ -1574,16 +1426,20 @@ async def mark_it_down(file: UploadFile = File(...), current_user: User = Depend
|
||||
temp_path = None
|
||||
|
||||
try:
|
||||
content = await file.read()
|
||||
|
||||
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as temp_file:
|
||||
temp_path = temp_file.name
|
||||
|
||||
async with aiofiles.open(temp_path, "wb") as temp_buffer:
|
||||
await temp_buffer.write(content)
|
||||
await write_upload_to_path(
|
||||
file,
|
||||
temp_path,
|
||||
max_size_bytes=MAX_UPLOAD_SIZE_BYTES,
|
||||
too_large_message="文件过大,当前仅支持 100 MB 以内的文件",
|
||||
)
|
||||
|
||||
markdown_content = await Parser.aparse(temp_path)
|
||||
return {"markdown_content": markdown_content, "message": "success"}
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
except Exception as e:
|
||||
logger.error(f"文件解析失败 {e}, {traceback.format_exc()}")
|
||||
return {"message": f"文件解析失败 {e}", "markdown_content": ""}
|
||||
@ -1631,7 +1487,7 @@ async def get_knowledge_base_statistics(current_user: User = Depends(get_admin_u
|
||||
async def generate_description(
|
||||
name: str = Body(..., description="知识库名称"),
|
||||
current_description: str = Body("", description="当前描述(可选,用于优化)"),
|
||||
file_list: list[str] = Body([], description="文件列表"),
|
||||
file_list: list[str] | None = Body(None, description="文件列表"),
|
||||
current_user: User = Depends(get_admin_user),
|
||||
):
|
||||
"""使用 LLM 生成或优化知识库描述
|
||||
@ -1640,6 +1496,7 @@ async def generate_description(
|
||||
"""
|
||||
from yuxi.models import select_model
|
||||
|
||||
file_list = file_list or []
|
||||
logger.debug(f"Generating description for knowledge base: {name}, files: {len(file_list)}")
|
||||
|
||||
# 构建文件列表文本
|
||||
|
||||
@ -709,7 +709,7 @@ async def test_sample_questions_endpoints(test_client, admin_headers, knowledge_
|
||||
get_response = await test_client.get(f"/api/knowledge/databases/{slug}/sample-questions", headers=admin_headers)
|
||||
assert get_response.status_code == 200, get_response.text
|
||||
get_payload = get_response.json()
|
||||
assert get_payload["slug"] == slug
|
||||
assert get_payload["kb_id"] == slug
|
||||
assert "questions" in get_payload
|
||||
assert get_payload["count"] == 0 # 空知识库没有问题
|
||||
|
||||
|
||||
@ -1,3 +1,5 @@
|
||||
import pytest
|
||||
|
||||
from yuxi.knowledge.utils.kb_utils import prepare_item_metadata
|
||||
|
||||
|
||||
@ -31,3 +33,8 @@ async def test_prepare_item_metadata_preserves_preprocessed_file_size():
|
||||
|
||||
assert metadata["size"] == 5678
|
||||
assert "_preprocessed_map" not in (metadata.get("processing_params") or {})
|
||||
|
||||
|
||||
async def test_prepare_item_metadata_rejects_direct_url_content_type():
|
||||
with pytest.raises(ValueError, match="Unsupported content_type"):
|
||||
await prepare_item_metadata("https://example.com", "url", "db")
|
||||
|
||||
108
backend/test/unit/knowledge/test_sample_question_utils.py
Normal file
108
backend/test/unit/knowledge/test_sample_question_utils.py
Normal file
@ -0,0 +1,108 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from yuxi.knowledge.utils import sample_question_utils as sq
|
||||
|
||||
|
||||
def test_parse_sample_questions_content_strips_json_fence():
|
||||
questions = sq.parse_sample_questions_content('```json\n{"questions": ["什么是测试?"]}\n```')
|
||||
|
||||
assert questions == ["什么是测试?"]
|
||||
|
||||
|
||||
def test_parse_sample_questions_content_rejects_invalid_payload():
|
||||
with pytest.raises(ValueError, match="问题格式"):
|
||||
sq.parse_sample_questions_content('{"items": []}')
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_database_sample_questions_rejects_empty_files(monkeypatch):
|
||||
class FakeKnowledgeBase:
|
||||
async def get_database_info(self, kb_id: str) -> dict:
|
||||
return {"name": "空知识库", "kb_type": "milvus", "files": {}}
|
||||
|
||||
monkeypatch.setattr(sq, "knowledge_base", FakeKnowledgeBase())
|
||||
monkeypatch.setattr(
|
||||
sq.KnowledgeBaseFactory,
|
||||
"get_kb_class",
|
||||
lambda _kb_type: SimpleNamespace(supports_documents=True),
|
||||
)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await sq.generate_database_sample_questions("kb_1")
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "没有文件" in exc_info.value.detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_database_sample_questions_saves_and_returns_questions(monkeypatch):
|
||||
saved: dict = {}
|
||||
|
||||
class FakeKnowledgeBase:
|
||||
async def get_database_info(self, kb_id: str) -> dict:
|
||||
return {
|
||||
"name": "测试知识库",
|
||||
"kb_type": "milvus",
|
||||
"files": {"file_1": {"filename": "demo.md", "file_type": "md"}},
|
||||
}
|
||||
|
||||
class FakeModel:
|
||||
async def call(self, messages, stream: bool = False):
|
||||
assert messages[0]["role"] == "system"
|
||||
assert "demo.md" in messages[1]["content"]
|
||||
return SimpleNamespace(content='{"questions": ["如何使用 demo?"]}')
|
||||
|
||||
class FakeRepository:
|
||||
async def update(self, kb_id: str, data: dict) -> None:
|
||||
saved[kb_id] = data["sample_questions"]
|
||||
|
||||
async def get_by_kb_id(self, kb_id: str):
|
||||
return SimpleNamespace(name="测试知识库", sample_questions=saved.get(kb_id))
|
||||
|
||||
monkeypatch.setattr(sq, "knowledge_base", FakeKnowledgeBase())
|
||||
monkeypatch.setattr(
|
||||
sq.KnowledgeBaseFactory,
|
||||
"get_kb_class",
|
||||
lambda _kb_type: SimpleNamespace(supports_documents=True),
|
||||
)
|
||||
monkeypatch.setattr(sq, "select_model", lambda model_spec: FakeModel())
|
||||
monkeypatch.setattr(sq, "KnowledgeBaseRepository", lambda: FakeRepository())
|
||||
|
||||
generated = await sq.generate_database_sample_questions("kb_1", count=1)
|
||||
stored = await sq.get_database_sample_questions("kb_1")
|
||||
|
||||
assert generated["questions"] == ["如何使用 demo?"]
|
||||
assert generated["count"] == 1
|
||||
assert stored["questions"] == ["如何使用 demo?"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_generate_database_sample_questions_maps_invalid_json(monkeypatch):
|
||||
class FakeKnowledgeBase:
|
||||
async def get_database_info(self, kb_id: str) -> dict:
|
||||
return {
|
||||
"name": "测试知识库",
|
||||
"kb_type": "milvus",
|
||||
"files": {"file_1": {"filename": "demo.md", "file_type": "md"}},
|
||||
}
|
||||
|
||||
class FakeModel:
|
||||
async def call(self, messages, stream: bool = False):
|
||||
return SimpleNamespace(content="not json")
|
||||
|
||||
monkeypatch.setattr(sq, "knowledge_base", FakeKnowledgeBase())
|
||||
monkeypatch.setattr(
|
||||
sq.KnowledgeBaseFactory,
|
||||
"get_kb_class",
|
||||
lambda _kb_type: SimpleNamespace(supports_documents=True),
|
||||
)
|
||||
monkeypatch.setattr(sq, "select_model", lambda model_spec: FakeModel())
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await sq.generate_database_sample_questions("kb_1")
|
||||
|
||||
assert exc_info.value.status_code == 500
|
||||
assert "AI返回格式错误" in exc_info.value.detail
|
||||
136
backend/test/unit/routers/test_knowledge_router_cleanup.py
Normal file
136
backend/test/unit/routers/test_knowledge_router_cleanup.py
Normal file
@ -0,0 +1,136 @@
|
||||
from io import BytesIO
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException, UploadFile
|
||||
|
||||
from server.routers import knowledge_router
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
class FakeTaskContext:
|
||||
def __init__(self):
|
||||
self.result = None
|
||||
|
||||
async def set_message(self, message: str) -> None:
|
||||
return None
|
||||
|
||||
async def set_progress(self, progress: float, message: str | None = None) -> None:
|
||||
return None
|
||||
|
||||
async def set_result(self, result: dict) -> None:
|
||||
self.result = result
|
||||
|
||||
async def raise_if_cancelled(self) -> None:
|
||||
return None
|
||||
|
||||
|
||||
async def test_upload_file_rejects_oversized_file(monkeypatch):
|
||||
monkeypatch.setattr(knowledge_router, "MAX_UPLOAD_SIZE_BYTES", 5)
|
||||
upload = UploadFile(filename="demo.txt", file=BytesIO(b"123456"))
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await knowledge_router.upload_file(upload, kb_id="kb_1", current_user=SimpleNamespace(uid="user_1"))
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "100 MB" in exc_info.value.detail
|
||||
|
||||
|
||||
async def test_markdown_endpoint_rejects_oversized_file(monkeypatch):
|
||||
monkeypatch.setattr(knowledge_router, "MAX_UPLOAD_SIZE_BYTES", 5)
|
||||
upload = UploadFile(filename="demo.txt", file=BytesIO(b"123456"))
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await knowledge_router.mark_it_down(upload, current_user=SimpleNamespace(uid="user_1"))
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "100 MB" in exc_info.value.detail
|
||||
|
||||
|
||||
async def test_index_documents_uses_uid_for_operator(monkeypatch):
|
||||
captured = {}
|
||||
|
||||
async def fake_ensure_database_supports_documents(kb_id: str, operation: str) -> None:
|
||||
return None
|
||||
|
||||
async def fake_get_database_info(kb_id: str) -> dict:
|
||||
return {"name": "测试知识库"}
|
||||
|
||||
async def fake_index_file(kb_id: str, file_id: str, operator_id: str | None = None, params: dict | None = None):
|
||||
captured["operator_id"] = operator_id
|
||||
return {"file_id": file_id, "status": "indexed"}
|
||||
|
||||
async def fake_enqueue(name: str, task_type: str, payload: dict, coroutine):
|
||||
await coroutine(FakeTaskContext())
|
||||
return SimpleNamespace(id="task_1")
|
||||
|
||||
monkeypatch.setattr(
|
||||
knowledge_router,
|
||||
"_ensure_database_supports_documents",
|
||||
fake_ensure_database_supports_documents,
|
||||
)
|
||||
monkeypatch.setattr(knowledge_router.knowledge_base, "get_database_info", fake_get_database_info)
|
||||
monkeypatch.setattr(knowledge_router.knowledge_base, "index_file", fake_index_file)
|
||||
monkeypatch.setattr(knowledge_router.tasker, "enqueue", fake_enqueue)
|
||||
|
||||
result = await knowledge_router.index_documents(
|
||||
"kb_1",
|
||||
["file_1"],
|
||||
params={},
|
||||
current_user=SimpleNamespace(id="numeric-id", uid="uid-user"),
|
||||
)
|
||||
|
||||
assert result["status"] == "queued"
|
||||
assert captured["operator_id"] == "uid-user"
|
||||
|
||||
|
||||
async def test_add_documents_auto_index_returns_one_final_result_per_item(monkeypatch):
|
||||
context = FakeTaskContext()
|
||||
item = "minio://knowledgebases/kb_1/upload/demo.txt"
|
||||
|
||||
async def fake_ensure_database_supports_documents(kb_id: str, operation: str) -> None:
|
||||
return None
|
||||
|
||||
async def fake_get_database_info(kb_id: str) -> dict:
|
||||
return {"name": "测试知识库"}
|
||||
|
||||
async def fake_add_file_record(kb_id: str, item_path: str, params: dict, operator_id: str | None = None):
|
||||
return {"file_id": "file_1", "status": "indexing"}
|
||||
|
||||
async def fake_parse_file(kb_id: str, file_id: str, operator_id: str | None = None):
|
||||
return {"file_id": file_id, "status": "parsed"}
|
||||
|
||||
async def fake_update_file_params(kb_id: str, file_id: str, params: dict, operator_id: str | None = None):
|
||||
return None
|
||||
|
||||
async def fake_index_file(kb_id: str, file_id: str, operator_id: str | None = None, params: dict | None = None):
|
||||
return {"file_id": file_id, "status": "indexed"}
|
||||
|
||||
async def fake_enqueue(name: str, task_type: str, payload: dict, coroutine):
|
||||
await coroutine(context)
|
||||
return SimpleNamespace(id="task_1")
|
||||
|
||||
monkeypatch.setattr(
|
||||
knowledge_router,
|
||||
"_ensure_database_supports_documents",
|
||||
fake_ensure_database_supports_documents,
|
||||
)
|
||||
monkeypatch.setattr(knowledge_router.knowledge_base, "get_database_info", fake_get_database_info)
|
||||
monkeypatch.setattr(knowledge_router.knowledge_base, "add_file_record", fake_add_file_record)
|
||||
monkeypatch.setattr(knowledge_router.knowledge_base, "parse_file", fake_parse_file)
|
||||
monkeypatch.setattr(knowledge_router.knowledge_base, "update_file_params", fake_update_file_params)
|
||||
monkeypatch.setattr(knowledge_router.knowledge_base, "index_file", fake_index_file)
|
||||
monkeypatch.setattr(knowledge_router.tasker, "enqueue", fake_enqueue)
|
||||
|
||||
result = await knowledge_router.add_documents(
|
||||
"kb_1",
|
||||
[item],
|
||||
params={"content_type": "file", "auto_index": True},
|
||||
current_user=SimpleNamespace(uid="uid-user"),
|
||||
)
|
||||
|
||||
assert result["status"] == "queued"
|
||||
assert context.result["submitted"] == 1
|
||||
assert context.result["failed"] == 0
|
||||
assert context.result["items"] == [{"file_id": "file_1", "status": "indexed"}]
|
||||
@ -45,7 +45,7 @@ async def test_import_workspace_files_uploads_workspace_file_to_minio(tmp_path,
|
||||
monkeypatch.setattr(knowledge_router, "aupload_file_to_minio", fake_upload)
|
||||
|
||||
result = await knowledge_router.import_workspace_files(
|
||||
knowledge_router.WorkspaceImportRequest(slug="db_1", paths=["/note.md"]),
|
||||
knowledge_router.WorkspaceImportRequest(kb_id="db_1", paths=["/note.md"]),
|
||||
current_user=SimpleNamespace(id="user_1"),
|
||||
)
|
||||
|
||||
@ -77,7 +77,7 @@ async def test_import_workspace_files_rejects_directory(tmp_path, monkeypatch):
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await knowledge_router.import_workspace_files(
|
||||
knowledge_router.WorkspaceImportRequest(slug="db_1", paths=["/folder"]),
|
||||
knowledge_router.WorkspaceImportRequest(kb_id="db_1", paths=["/folder"]),
|
||||
current_user=SimpleNamespace(id="user_1"),
|
||||
)
|
||||
|
||||
|
||||
@ -38,6 +38,7 @@
|
||||
### 0.7.0 开发记录
|
||||
|
||||
<!-- 0.7.0 的内容请放在这里 -->
|
||||
- 降低知识库路由与工具模块复杂度:示例问题生成迁移到知识库 utils,文件上传统一 100 MB 限制,URL 预处理入库路径与旧 `content_type=url` 行为收敛,并修复 uid、导出 MIME 与异常透传等路由问题。
|
||||
- 重构智能体配置语义:用户可见的 `AgentConfig` 收敛为数据库持久化的一级 `Agent`,内置 Python Agent 改为智能体后端;新增 `/api/agent` 管理与运行接口,聊天、运行任务、恢复审批和文件预览均从线程绑定的 Agent 解析运行时上下文,前端只提交 `agent_id`,并在模型配置页新增“智能体”管理页签。
|
||||
- 删除 Upload 与 LightRAG 图谱/知识库能力:知识库类型收敛为 Milvus 与 Dify,只保留 Milvus 知识库内图谱构建/展示/检索,移除独立 `/graph` 页面和默认上传图谱工具。
|
||||
- 收敛只读知识源连接器:新增 `ReadOnlyConnectors` 基类,Dify 改为声明自身创建参数与校验规则,新增 Notion Data Source 只读知识库并支持 Search/Find/Open;知识库类型接口返回创建参数 schema,前端新建表单按类型动态渲染非 Milvus 配置并统一保存到 `additional_params`。
|
||||
|
||||
Loading…
Reference in New Issue
Block a user