diff --git a/backend/server/routers/auth_dept_router.py b/backend/server/routers/auth_dept_router.py index 67ba2dac..fa84e809 100644 --- a/backend/server/routers/auth_dept_router.py +++ b/backend/server/routers/auth_dept_router.py @@ -10,7 +10,7 @@ from pydantic import BaseModel from sqlalchemy import delete as sqlalchemy_delete, select, func from sqlalchemy.ext.asyncio import AsyncSession -from yuxi.storage.postgres.models_business import APIKey, AgentConfig, Department, User +from yuxi.storage.postgres.models_business import APIKey, Department, User from yuxi.repositories.department_repository import DepartmentRepository from yuxi.repositories.user_repository import UserRepository from server.utils.auth_middleware import get_superadmin_user, get_admin_user, get_db @@ -230,7 +230,6 @@ async def delete_department( for user in department_users: user.department_id = 1 # 将被删除部门的用户移至默认部门 - await db.execute(sqlalchemy_delete(AgentConfig).where(AgentConfig.department_id == department_id)) await db.execute(sqlalchemy_delete(APIKey).where(APIKey.department_id == department_id)) await db.delete(department) await db.commit() diff --git a/backend/server/routers/chat_router.py b/backend/server/routers/chat_router.py index bab8de6b..86f4e4d5 100644 --- a/backend/server/routers/chat_router.py +++ b/backend/server/routers/chat_router.py @@ -10,8 +10,9 @@ from sqlalchemy.ext.asyncio import AsyncSession from yuxi.storage.postgres.models_business import User from server.routers.auth_router import get_admin_user -from server.utils.auth_middleware import get_db, get_required_user +from server.utils.auth_middleware import get_current_user, get_db, get_required_user from yuxi import config as conf +from yuxi.agents.context import filter_config_by_role from yuxi.agents.buildin import agent_manager from yuxi.models import select_model from yuxi.services.chat_service import agent_chat, get_agent_state_view, stream_agent_chat, stream_agent_resume @@ -99,6 +100,24 @@ class AgentChatRequest(BaseModel): chat = APIRouter(prefix="/chat", tags=["chat"]) + +async def get_config_user(user: User | None = Depends(get_current_user)) -> User: + if user is None: + raise HTTPException(status_code=401, detail="请登录后再访问", headers={"WWW-Authenticate": "Bearer"}) + return user + + +def _filter_agent_config_json(agent_id: str, config_json: dict | None, role: str | None) -> dict: + agent = agent_manager.get_agent(agent_id) + context_schema = agent.context_schema if agent else None + return filter_config_by_role(config_json or {}, role, context_schema=context_schema) + + +def _serialize_agent_config(item, role: str | None) -> dict: + data = item.to_dict() + data["config_json"] = _filter_agent_config_json(item.agent_id, data.get("config_json"), role) + return data + # ============================================================================= # > === 智能体管理分组 === # ============================================================================= @@ -174,7 +193,7 @@ async def get_agent(current_user: User = Depends(get_required_user)): @chat.get("/agent/{agent_id}") -async def get_single_agent(agent_id: str, current_user: User = Depends(get_required_user)): +async def get_single_agent(agent_id: str, current_user: User = Depends(get_config_user)): """获取指定智能体的完整信息(包含配置选项)(需要登录)""" try: # 检查智能体是否存在 @@ -182,7 +201,7 @@ async def get_single_agent(agent_id: str, current_user: User = Depends(get_requi raise HTTPException(status_code=404, detail=f"智能体 {agent_id} 不存在") # 获取智能体的完整信息(包含 configurable_items) - agent_info = await agent.get_info() + agent_info = await agent.get_info(user_role=current_user.role) return agent_info @@ -196,21 +215,22 @@ async def get_single_agent(agent_id: str, current_user: User = Depends(get_requi @chat.get("/agent/{agent_id}/configs") async def list_agent_configs( agent_id: str, - current_user: User = Depends(get_required_user), + current_user: User = Depends(get_config_user), db: AsyncSession = Depends(get_db), ): if not agent_manager.get_agent(agent_id): raise HTTPException(status_code=404, detail=f"智能体 {agent_id} 不存在") repo = AgentConfigRepository(db) - items = await repo.list_by_department_agent(department_id=current_user.department_id, agent_id=agent_id) + uid = str(current_user.uid) + items = await repo.list_by_user_agent(uid=uid, agent_id=agent_id) if not items: await repo.get_or_create_default( - department_id=current_user.department_id, + uid=uid, agent_id=agent_id, - created_by=str(current_user.id), + created_by=uid, ) - items = await repo.list_by_department_agent(department_id=current_user.department_id, agent_id=agent_id) + items = await repo.list_by_user_agent(uid=uid, agent_id=agent_id) configs = [ { @@ -231,7 +251,7 @@ async def list_agent_configs( async def get_agent_config_profile( agent_id: str, config_id: int, - current_user: User = Depends(get_required_user), + current_user: User = Depends(get_config_user), db: AsyncSession = Depends(get_db), ): if not agent_manager.get_agent(agent_id): @@ -239,39 +259,38 @@ async def get_agent_config_profile( repo = AgentConfigRepository(db) item = await repo.get_by_id(config_id) - if not item or item.agent_id != agent_id or item.department_id != current_user.department_id: + if not item or item.agent_id != agent_id or item.uid != str(current_user.uid): raise HTTPException(status_code=404, detail="配置不存在") - return {"config": item.to_dict()} + return {"config": _serialize_agent_config(item, current_user.role)} @chat.post("/agent/{agent_id}/configs") async def create_agent_config_profile( agent_id: str, payload: AgentConfigCreate, - current_user: User = Depends(get_admin_user), + current_user: User = Depends(get_config_user), db: AsyncSession = Depends(get_db), ): if not agent_manager.get_agent(agent_id): raise HTTPException(status_code=404, detail=f"智能体 {agent_id} 不存在") repo = AgentConfigRepository(db) + uid = str(current_user.uid) item = await repo.create( - department_id=current_user.department_id, + uid=uid, agent_id=agent_id, name=payload.name, description=payload.description, icon=payload.icon, pics=payload.pics, examples=payload.examples, - config_json=payload.config_json, + config_json=_filter_agent_config_json(agent_id, payload.config_json, current_user.role), is_default=payload.set_default, - created_by=str(current_user.id), + created_by=uid, ) - if payload.set_default: - item = await repo.set_default(config=item, updated_by=str(current_user.id)) - return {"config": item.to_dict()} + return {"config": _serialize_agent_config(item, current_user.role)} @chat.put("/agent/{agent_id}/configs/{config_id}") @@ -279,7 +298,7 @@ async def update_agent_config_profile( agent_id: str, config_id: int, payload: AgentConfigUpdate, - current_user: User = Depends(get_admin_user), + current_user: User = Depends(get_config_user), db: AsyncSession = Depends(get_db), ): if not agent_manager.get_agent(agent_id): @@ -287,7 +306,7 @@ async def update_agent_config_profile( repo = AgentConfigRepository(db) item = await repo.get_by_id(config_id) - if not item or item.agent_id != agent_id or item.department_id != current_user.department_id: + if not item or item.agent_id != agent_id or item.uid != str(current_user.uid): raise HTTPException(status_code=404, detail="配置不存在") updated = await repo.update( @@ -297,17 +316,19 @@ async def update_agent_config_profile( icon=payload.icon, pics=payload.pics, examples=payload.examples, - config_json=payload.config_json, - updated_by=str(current_user.id), + config_json=_filter_agent_config_json(agent_id, payload.config_json, current_user.role) + if payload.config_json is not None + else None, + updated_by=str(current_user.uid), ) - return {"config": updated.to_dict()} + return {"config": _serialize_agent_config(updated, current_user.role)} @chat.post("/agent/{agent_id}/configs/{config_id}/set_default") async def set_agent_config_default( agent_id: str, config_id: int, - current_user: User = Depends(get_admin_user), + current_user: User = Depends(get_config_user), db: AsyncSession = Depends(get_db), ): if not agent_manager.get_agent(agent_id): @@ -315,18 +336,18 @@ async def set_agent_config_default( repo = AgentConfigRepository(db) item = await repo.get_by_id(config_id) - if not item or item.agent_id != agent_id or item.department_id != current_user.department_id: + if not item or item.agent_id != agent_id or item.uid != str(current_user.uid): raise HTTPException(status_code=404, detail="配置不存在") - updated = await repo.set_default(config=item, updated_by=str(current_user.id)) - return {"config": updated.to_dict()} + updated = await repo.set_default(config=item, updated_by=str(current_user.uid)) + return {"config": _serialize_agent_config(updated, current_user.role)} @chat.delete("/agent/{agent_id}/configs/{config_id}") async def delete_agent_config_profile( agent_id: str, config_id: int, - current_user: User = Depends(get_admin_user), + current_user: User = Depends(get_config_user), db: AsyncSession = Depends(get_db), ): if not agent_manager.get_agent(agent_id): @@ -334,10 +355,10 @@ async def delete_agent_config_profile( repo = AgentConfigRepository(db) item = await repo.get_by_id(config_id) - if not item or item.agent_id != agent_id or item.department_id != current_user.department_id: + if not item or item.agent_id != agent_id or item.uid != str(current_user.uid): raise HTTPException(status_code=404, detail="配置不存在") - await repo.delete(config=item, updated_by=str(current_user.id)) + await repo.delete(config=item, updated_by=str(current_user.uid)) return {"success": True} diff --git a/backend/server/routers/knowledge_router.py b/backend/server/routers/knowledge_router.py index 063d3a9c..ed2ff0a5 100644 --- a/backend/server/routers/knowledge_router.py +++ b/backend/server/routers/knowledge_router.py @@ -15,7 +15,7 @@ from starlette.responses import StreamingResponse from yuxi.services.task_service import TaskContext, tasker from server.utils.auth_middleware import get_admin_user, get_required_user from yuxi import config, knowledge_base -from yuxi.knowledge.chunking.ragflow_like.presets import ensure_chunk_defaults_in_additional_params +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 @@ -28,7 +28,6 @@ from yuxi.utils import logger knowledge = APIRouter(prefix="/knowledge", tags=["knowledge"]) -DIFY_REQUIRED_PARAMS = ("dify_api_url", "dify_token", "dify_dataset_id") ACTIVE_GRAPH_BUILD_STATUSES = {"pending", "running"} @@ -97,24 +96,14 @@ async def _delete_document_storage_objects(db_id: str, doc_id: str, file_path: s logger.warning(f"从MinIO删除解析结果失败: {minio_error}") -def _validate_dify_additional_params(additional_params: dict | None) -> dict: - params = dict(additional_params or {}) - missing_fields = [field for field in DIFY_REQUIRED_PARAMS if not str(params.get(field) or "").strip()] - if missing_fields: - raise HTTPException(status_code=400, detail=f"Dify 参数缺失: {', '.join(missing_fields)}") - - api_url = str(params.get("dify_api_url") or "").strip() - if not api_url.endswith("/v1"): - raise HTTPException(status_code=400, detail="Dify api_url 必须以 /v1 结尾") - return params - - -async def _ensure_database_not_dify(db_id: str, operation: str) -> None: +async def _ensure_database_supports_documents(db_id: str, operation: str) -> None: db_info = await knowledge_base.get_database_info(db_id) if not db_info: raise HTTPException(status_code=404, detail=f"知识库 {db_id} 不存在") - if (db_info.get("kb_type") or "").lower() == "dify": - raise HTTPException(status_code=400, detail=f"Dify 知识库只支持检索,不支持{operation}") + kb_type = (db_info.get("kb_type") or "").lower() + kb_class = KnowledgeBaseFactory.get_kb_class(kb_type) + if not kb_class.supports_documents: + raise HTTPException(status_code=400, detail=f"{db_info.get('name') or kb_type} 只支持检索,不支持{operation}") async def _has_running_graph_build_task(db_id: str) -> bool: @@ -168,6 +157,11 @@ async def create_database( detail=f"知识库名称 '{database_name}' 已存在,请使用其他名称", ) + if not KnowledgeBaseFactory.is_type_supported(kb_type): + raise HTTPException(status_code=400, detail=f"Unsupported knowledge base type: {kb_type}") + + kb_class = KnowledgeBaseFactory.get_kb_class(kb_type) + additional_params = {**(additional_params or {})} additional_params["auto_generate_questions"] = False # 默认不生成问题 @@ -176,17 +170,17 @@ async def create_database( status_code=400, detail="reranker_config 已移除,请在查询参数中使用 reranker_model spec", ) - additional_params = ensure_chunk_defaults_in_additional_params(additional_params) + additional_params = kb_class.normalize_additional_params(additional_params) - if kb_type == "dify": - additional_params = _validate_dify_additional_params(additional_params) - else: + if kb_class.requires_embedding_model: if not embedding_model_spec: raise HTTPException(status_code=400, detail="embedding_model_spec 不能为空") info = model_cache.get_model_info(embedding_model_spec) if not info or info.model_type != "embedding": raise HTTPException(status_code=400, detail=f"不支持的 embedding 模型: {embedding_model_spec}") + else: + embedding_model_spec = None database_info = await knowledge_base.create_database( database_name, @@ -195,6 +189,7 @@ async def create_database( embedding_model_spec=embedding_model_spec, llm_model_spec=llm_model_spec, share_config=share_config, + created_by=current_user.uid, **additional_params, ) @@ -222,6 +217,7 @@ async def get_accessible_databases(current_user: User = Depends(get_required_use "name": db.get("name", ""), "db_id": db.get("db_id"), "description": db.get("description", ""), + "created_by": db.get("created_by"), } for db in databases.get("databases", []) ] @@ -257,17 +253,20 @@ async def update_database_info( additional_params = data.additional_params if additional_params is not None: - additional_params = ensure_chunk_defaults_in_additional_params(additional_params) - db_info = await knowledge_base.get_database_info(db_id) if not db_info: raise HTTPException(status_code=404, detail=f"知识库 {db_id} 不存在") kb_type = (db_info.get("kb_type") or "").lower() - if kb_type == "dify": - merged_params = dict(db_info.get("additional_params") or {}) - merged_params.update(additional_params) - _validate_dify_additional_params(merged_params) + kb_class = KnowledgeBaseFactory.get_kb_class(kb_type) + merged_params = dict(db_info.get("additional_params") or {}) + merged_params.update(additional_params) + kb_class.normalize_additional_params(merged_params) + additional_params = ( + kb_class.normalize_additional_params(additional_params) + if kb_class.apply_chunk_defaults + else kb_class.normalize_additional_params(merged_params) + ) database = await knowledge_base.update_database( db_id, @@ -444,7 +443,7 @@ async def add_documents( ): """添加文档到知识库(上传 -> 解析 -> 可选入库)""" logger.debug(f"Add documents for db_id {db_id}: {items} {params=}") - await _ensure_database_not_dify(db_id, "文档添加/解析/入库") + await _ensure_database_supports_documents(db_id, "文档添加/解析/入库") content_type = params.get("content_type", "file") # 自动入库参数 @@ -628,7 +627,7 @@ async def add_documents( async def parse_documents(db_id: str, file_ids: list[str] = Body(...), current_user: User = Depends(get_admin_user)): """手动触发文档解析""" logger.debug(f"Parse documents for db_id {db_id}: {file_ids}") - await _ensure_database_not_dify(db_id, "文档解析") + await _ensure_database_supports_documents(db_id, "文档解析") async def run_parse(context: TaskContext): await context.set_message("任务初始化") @@ -682,7 +681,7 @@ async def index_documents( ): """手动触发文档入库(Indexing),支持更新参数""" logger.debug(f"Index documents for db_id {db_id}: {file_ids} {params=}") - await _ensure_database_not_dify(db_id, "文档入库") + await _ensure_database_supports_documents(db_id, "文档入库") # extract operator_id safely before background task operator_id = current_user.id @@ -755,7 +754,7 @@ async def index_documents( async def get_document_info(db_id: str, doc_id: str, current_user: User = Depends(get_admin_user)): """获取文档详细信息(包含基本信息和内容信息)""" logger.debug(f"GET document {doc_id} info in {db_id}") - await _ensure_database_not_dify(db_id, "文档查看") + await _ensure_database_supports_documents(db_id, "文档查看") try: info = await knowledge_base.get_file_info(db_id, doc_id) @@ -769,7 +768,7 @@ async def get_document_info(db_id: str, doc_id: str, current_user: User = Depend async def get_document_basic_info(db_id: str, doc_id: str, current_user: User = Depends(get_admin_user)): """获取文档基本信息(仅元数据)""" logger.debug(f"GET document {doc_id} basic info in {db_id}") - await _ensure_database_not_dify(db_id, "文档查看") + await _ensure_database_supports_documents(db_id, "文档查看") try: info = await knowledge_base.get_file_basic_info(db_id, doc_id) @@ -783,7 +782,7 @@ async def get_document_basic_info(db_id: str, doc_id: str, current_user: User = async def get_document_content(db_id: str, doc_id: str, current_user: User = Depends(get_admin_user)): """获取文档内容信息(chunks和lines)""" logger.debug(f"GET document {doc_id} content in {db_id}") - await _ensure_database_not_dify(db_id, "文档查看") + await _ensure_database_supports_documents(db_id, "文档查看") try: info = await knowledge_base.get_file_content(db_id, doc_id) @@ -799,7 +798,7 @@ async def batch_delete_documents( ): """批量删除文档或文件夹""" logger.debug(f"BATCH DELETE documents {file_ids} in {db_id}") - await _ensure_database_not_dify(db_id, "批量文档删除") + await _ensure_database_supports_documents(db_id, "批量文档删除") deleted_count = 0 failed_items = [] @@ -842,7 +841,7 @@ async def batch_delete_documents( async def delete_document(db_id: str, doc_id: str, current_user: User = Depends(get_admin_user)): """删除文档或文件夹""" logger.debug(f"DELETE document {doc_id} info in {db_id}") - await _ensure_database_not_dify(db_id, "文档删除") + await _ensure_database_supports_documents(db_id, "文档删除") try: file_meta_info = await knowledge_base.get_file_basic_info(db_id, doc_id) @@ -868,7 +867,7 @@ async def delete_document(db_id: str, doc_id: str, current_user: User = Depends( async def download_document(db_id: str, doc_id: str, request: Request, current_user: User = Depends(get_admin_user)): """下载原始文件""" logger.debug(f"Download document {doc_id} from {db_id}") - await _ensure_database_not_dify(db_id, "文档下载") + await _ensure_database_supports_documents(db_id, "文档下载") try: file_info = await knowledge_base.get_file_basic_info(db_id, doc_id) file_meta = file_info.get("meta", {}) @@ -1103,8 +1102,9 @@ async def generate_sample_questions( db_info = await knowledge_base.get_database_info(db_id) if not db_info: raise HTTPException(status_code=404, detail=f"知识库 {db_id} 不存在") - if (db_info.get("kb_type") or "").lower() == "dify": - raise HTTPException(status_code=400, detail="Dify 知识库不支持基于文件生成测试问题") + 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 @@ -1256,7 +1256,7 @@ async def create_folder( ): """创建文件夹""" try: - await _ensure_database_not_dify(db_id, "文件夹创建") + await _ensure_database_supports_documents(db_id, "文件夹创建") return await knowledge_base.create_folder(db_id, folder_name, parent_id) except Exception as e: logger.error(f"创建文件夹失败 {e}, {traceback.format_exc()}") @@ -1273,7 +1273,7 @@ async def move_document( """移动文件或文件夹""" logger.debug(f"Move document {doc_id} to {new_parent_id} in {db_id}") try: - await _ensure_database_not_dify(db_id, "文件移动") + await _ensure_database_supports_documents(db_id, "文件移动") return await knowledge_base.move_file(db_id, doc_id, new_parent_id) except ValueError as e: raise HTTPException(status_code=400, detail=str(e)) @@ -1369,7 +1369,7 @@ async def import_workspace_files( if not paths: raise HTTPException(status_code=400, detail="请选择至少一个工作区文件") - await _ensure_database_not_dify(db_id, "文档添加/解析/入库") + await _ensure_database_supports_documents(db_id, "文档添加/解析/入库") bucket_name = MinIOClient.KB_BUCKETS["documents"] results = []