diff --git a/REFACTOR.md b/REFACTOR.md index 226b5a2c..b764c828 100644 --- a/REFACTOR.md +++ b/REFACTOR.md @@ -27,6 +27,6 @@ - [ ] Config spacy model 的 load - [x] 在工作区的文件编辑的时候,保存和取消的按钮应该是悬浮在编辑框的右上角,而不是在 header 上面 - [ ] default enable all build in tools / kbs / skills / mcps / subagents -- [ ] 链接 Notion 和 feishu 目前来看,都是支持的 -- [ ] 知识库的权限调整,修改为三个等级,全局共享、部门共享(选择多个部门,默认是自己部门且必须包含自己部门)、指定人可访问(选择多个用户,默认是仅自己,可以添加其他人)。UI 上也需要调整,三个卡片不再是等宽,而是选中的会宽一点,并展示描述以及选择按钮,未选中的则是默认宽度,仅显示标题。对于选中的卡片除了展示描述、按钮之外,还包括“X 个部门可访问”、“X 个用户可访问”的信息展示。全局的就看是所有用户可访问。所以等级的字段配置也要重新设计,不需要考虑兼容,所有知识库都会重新构建。 +- [x] 链接 Notion 和 feishu 目前来看,都是支持的 +- [x] 知识库的权限调整,修改为三个等级,全局共享、部门共享(选择多个部门,默认是自己部门且必须包含自己部门)、指定人可访问(选择多个用户,默认是仅自己,可以添加其他人)。UI 上也需要调整,三个卡片不再是等宽,而是选中的会宽一点,并展示描述以及选择按钮,未选中的则是默认宽度,仅显示标题。对于选中的卡片除了展示描述、按钮之外,还包括“X 个部门可访问”、“X 个用户可访问”的信息展示。全局的就看是所有用户可访问。所以等级的字段配置也要重新设计,不需要考虑兼容,所有知识库都会重新构建。 - [ ] databaseinfo 的重构,在左侧展示那个 tab 标签吧,将文件管理(filetable)以及右侧的那些图谱、检索、检索配置、评估之类的,都列为不同的 tab,进入之后默认激活的是 filetable。这样页面布局就好的多。作恶侧边栏除了这些 tab 之外,顶部是和 Skill Detail 那里的 header 一样, \ No newline at end of file diff --git a/backend/package/yuxi/knowledge/manager.py b/backend/package/yuxi/knowledge/manager.py index 99d6e839..7a23bceb 100644 --- a/backend/package/yuxi/knowledge/manager.py +++ b/backend/package/yuxi/knowledge/manager.py @@ -9,6 +9,10 @@ from yuxi.utils import logger from yuxi.utils.datetime_utils import utc_isoformat +DEFAULT_SHARE_CONFIG = {"access_level": "global", "department_ids": [], "user_uids": []} +ACCESS_LEVELS = {"global", "department", "user"} + + class KnowledgeBaseManager: """ 知识库管理器 @@ -150,6 +154,47 @@ class KnowledgeBaseManager: """ return await self._get_kb_for_database(db_id) + def _normalize_share_config( + self, + share_config: dict | None, + *, + user_uid: str | None = None, + department_id: int | str | None = None, + ) -> dict: + config = share_config or {} + access_level = config.get("access_level") or "global" + if access_level not in ACCESS_LEVELS: + raise ValueError("无效的知识库权限等级") + + if access_level == "global": + return DEFAULT_SHARE_CONFIG.copy() + + if access_level == "department": + department_ids = self._normalize_department_ids(config.get("department_ids")) + if department_id is not None: + department_ids.append(int(department_id)) + department_ids = sorted(set(department_ids)) + if not department_ids: + raise ValueError("部门共享至少需要选择一个部门") + return {"access_level": "department", "department_ids": department_ids, "user_uids": []} + + user_uids = self._normalize_user_uids(config.get("user_uids")) + if user_uid: + user_uids.append(str(user_uid)) + user_uids = sorted(set(user_uids)) + if not user_uids: + raise ValueError("指定人可访问至少需要选择一个用户") + return {"access_level": "user", "department_ids": [], "user_uids": user_uids} + + def _normalize_department_ids(self, department_ids: list | None) -> list[int]: + normalized = [] + for department_id in department_ids or []: + normalized.append(int(department_id)) + return normalized + + def _normalize_user_uids(self, user_uids: list | None) -> list[str]: + return [uid for uid in (str(uid).strip() for uid in user_uids or []) if uid] + async def get_databases(self) -> dict: """获取所有数据库信息""" from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository @@ -178,7 +223,7 @@ class KnowledgeBaseManager: continue # 补充 share_config 和 additional_params - db_info["share_config"] = row.share_config or {"is_shared": True, "accessible_departments": []} + db_info["share_config"] = row.share_config or DEFAULT_SHARE_CONFIG.copy() db_info["additional_params"] = kb_instance.normalize_additional_params(row.additional_params) db_info["created_by"] = row.created_by all_databases.append(db_info) @@ -205,28 +250,30 @@ class KnowledgeBaseManager: if kb is None: return False - share_config = kb.share_config or {} - is_shared = share_config.get("is_shared", True) - - # 如果是全员共享,则有权限 - if is_shared: + user_uid = str(user.get("uid") or "") + if user_uid and kb.created_by == user_uid: return True - # 检查部门权限 - user_department_id = user.get("department_id") - accessible_departments = share_config.get("accessible_departments", []) + share_config = kb.share_config or DEFAULT_SHARE_CONFIG.copy() + access_level = share_config.get("access_level") - if user_department_id is None: - return False + if access_level == "global": + return True - # 转换为整数进行比较(前端可能传递字符串,后端存储为整数) - try: - user_department_id = int(user_department_id) - accessible_departments = [int(d) for d in accessible_departments] - except (ValueError, TypeError): - return False + if access_level == "department": + user_department_id = user.get("department_id") + if user_department_id is None: + return False + try: + department_ids = [int(dept_id) for dept_id in share_config.get("department_ids") or []] + return int(user_department_id) in department_ids + except (ValueError, TypeError): + return False - return user_department_id in accessible_departments + if access_level == "user": + return bool(user_uid and user_uid in (share_config.get("user_uids") or [])) + + return False async def get_databases_by_uid(self, uid: str) -> dict: """根据 uid 获取知识库列表""" @@ -248,6 +295,7 @@ class KnowledgeBaseManager: user_info = user else: user_info = { + "uid": user.uid, "role": user.role, "department_id": user.department_id, } @@ -304,6 +352,7 @@ class KnowledgeBaseManager: llm_model_spec: str | None = None, share_config: dict | None = None, created_by: str | None = None, + created_by_department_id: int | str | None = None, **kwargs, ) -> dict: """ @@ -317,6 +366,7 @@ class KnowledgeBaseManager: llm_model_spec: LLM 模型 spec share_config: 共享配置 created_by: 创建者 uid + created_by_department_id: 创建者部门 ID **kwargs: 其他配置参数 Returns: @@ -330,9 +380,11 @@ class KnowledgeBaseManager: if await self.database_name_exists(database_name): raise ValueError(f"知识库名称 '{database_name}' 已存在,请使用其他名称") - # 默认共享配置 - if share_config is None: - share_config = {"is_shared": True, "accessible_departments": []} + share_config = self._normalize_share_config( + share_config, + user_uid=created_by, + department_id=created_by_department_id, + ) kb_instance = self._get_or_create_kb_instance(kb_type) kwargs = kb_instance.normalize_additional_params(kwargs) @@ -444,7 +496,7 @@ class KnowledgeBaseManager: # 添加数据库中的附加字段 db_info["additional_params"] = kb_instance.normalize_additional_params(kb.additional_params) - db_info["share_config"] = kb.share_config or {"is_shared": True, "accessible_departments": []} + db_info["share_config"] = kb.share_config or DEFAULT_SHARE_CONFIG.copy() db_info["mindmap"] = kb.mindmap db_info["sample_questions"] = kb.sample_questions or [] db_info["query_params"] = kb.query_params @@ -622,6 +674,8 @@ class KnowledgeBaseManager: update_llm_model_spec: bool = False, additional_params: dict | None = None, share_config: dict | None = None, + operator_uid: str | None = None, + operator_department_id: int | str | None = None, ) -> dict: """更新数据库""" from yuxi.repositories.knowledge_base_repository import KnowledgeBaseRepository @@ -655,7 +709,11 @@ class KnowledgeBaseManager: kb_instance.databases_meta[db_id]["metadata"] = merged_additional_params if share_config is not None: - update_data["share_config"] = share_config + update_data["share_config"] = self._normalize_share_config( + share_config, + user_uid=operator_uid, + department_id=operator_department_id, + ) # 保存到数据库 await kb_repo.update(db_id, update_data) diff --git a/backend/server/routers/auth_router.py b/backend/server/routers/auth_router.py index d7b882bc..4fae2877 100644 --- a/backend/server/routers/auth_router.py +++ b/backend/server/routers/auth_router.py @@ -87,6 +87,14 @@ class UserResponse(BaseModel): last_login: str | None = None +class UserAccessOption(BaseModel): + uid: str + username: str + role: str + department_id: int | None = None + department_name: str | None = None + + class InitializeAdmin(BaseModel): uid: str # 直接输入用户ID password: str @@ -477,6 +485,11 @@ async def create_user( else: # 普通管理员创建用户时,自动继承该管理员的部门 department_id = current_user.department_id + if department_id is None: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="管理员必须属于部门才能创建用户", + ) # 非超级管理员不能指定部门 if user_data.department_id is not None: raise HTTPException( @@ -528,6 +541,41 @@ async def read_users( return users +def _ensure_user_in_current_department(current_user: User, target_user: User) -> None: + if current_user.role == "superadmin": + return + if target_user.department_id != current_user.department_id: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail="只能管理本部门用户", + ) + + +@auth.get("/users/access-options", response_model=list[UserAccessOption]) +async def read_user_access_options( + skip: int = 0, + limit: int = 1000, + current_user: User = Depends(get_admin_user), +): + user_repo = UserRepository() + if current_user.role == "superadmin": + users_with_dept = await user_repo.list_with_department(skip=skip, limit=limit) + else: + users_with_dept = await user_repo.list_with_department( + skip=skip, limit=limit, department_id=current_user.department_id + ) + return [ + { + "uid": user.uid, + "username": user.username, + "role": user.role, + "department_id": user.department_id, + "department_name": dept_name, + } + for user, dept_name in users_with_dept + ] + + # 路由:获取特定用户信息(管理员权限) @auth.get("/users/{user_id}", response_model=UserResponse) async def read_user(user_id: int, current_user: User = Depends(get_admin_user), db: AsyncSession = Depends(get_db)): @@ -538,6 +586,7 @@ async def read_user(user_id: int, current_user: User = Depends(get_admin_user), status_code=status.HTTP_404_NOT_FOUND, detail="用户不存在", ) + _ensure_user_in_current_department(current_user, user) return user.to_dict() @@ -570,6 +619,8 @@ async def update_user( detail="用户不存在", ) + _ensure_user_in_current_department(current_user, user) + # 检查权限 if user.role == "superadmin" and current_user.role != "superadmin": raise HTTPException( @@ -584,6 +635,18 @@ async def update_user( detail="不能降级超级管理员账户", ) + if current_user.role == "admin": + if user.role != "user": + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail="管理员只能修改普通用户账户", + ) + if user_data.role is not None and user_data.role != "user": + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail="管理员只能将用户角色设置为普通用户", + ) + # 更新信息 update_details = [] @@ -664,6 +727,8 @@ async def delete_user( detail="用户不存在", ) + _ensure_user_in_current_department(current_user, user) + # 不能删除超级管理员账户 if user.role == "superadmin": raise HTTPException( @@ -671,6 +736,12 @@ async def delete_user( detail="不能删除超级管理员账户", ) + if current_user.role == "admin" and user.role != "user": + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail="管理员只能删除普通用户账户", + ) + # 检查是否是部门的唯一管理员 if user.role == "admin" and current_user.role != "superadmin": result = await db.execute( @@ -849,7 +920,7 @@ async def impersonate_user( return { "access_token": access_token, "token_type": "bearer", - "uid": target_user.id, + "user_id": target_user.id, "username": target_user.username, "uid": target_user.uid, "phone_number": target_user.phone_number, diff --git a/backend/server/routers/knowledge_router.py b/backend/server/routers/knowledge_router.py index ed2ff0a5..f5ce4902 100644 --- a/backend/server/routers/knowledge_router.py +++ b/backend/server/routers/knowledge_router.py @@ -190,6 +190,7 @@ async def create_database( llm_model_spec=llm_model_spec, share_config=share_config, created_by=current_user.uid, + created_by_department_id=current_user.department_id, **additional_params, ) @@ -276,6 +277,8 @@ async def update_database_info( update_llm_model_spec=update_llm_model_spec, additional_params=additional_params, share_config=data.share_config, + operator_uid=current_user.uid, + operator_department_id=current_user.department_id, ) return {"message": "更新成功", "database": database} except Exception as e: diff --git a/backend/server/routers/workspace_router.py b/backend/server/routers/workspace_router.py index 59abef8e..2e6f9d83 100644 --- a/backend/server/routers/workspace_router.py +++ b/backend/server/routers/workspace_router.py @@ -36,6 +36,7 @@ class UpdateWorkspaceFileContentRequest(BaseModel): async def _ensure_knowledge_read_access(current_user: User, db_id: str) -> None: allowed = await knowledge_base.check_accessible( { + "uid": current_user.uid, "role": current_user.role, "department_id": current_user.department_id, }, diff --git a/backend/test/integration/api/test_auth_router.py b/backend/test/integration/api/test_auth_router.py index fcef09f2..a1666864 100644 --- a/backend/test/integration/api/test_auth_router.py +++ b/backend/test/integration/api/test_auth_router.py @@ -11,6 +11,68 @@ import pytest pytestmark = [pytest.mark.asyncio, pytest.mark.integration] +async def _require_superadmin(test_client, headers): + response = await test_client.get("/api/auth/me", headers=headers) + assert response.status_code == 200, response.text + if response.json()["role"] != "superadmin": + pytest.fail("This test requires TEST_USERNAME to be a superadmin account.") + + +async def _create_department_with_admin(test_client, headers, label: str) -> dict: + suffix = uuid.uuid4().hex[:8] + admin_uid = f"adm{label}_{suffix}" + admin_password = f"Pw!{suffix}" + response = await test_client.post( + "/api/departments", + json={ + "name": f"pytest_{label}_{suffix}", + "description": "pytest managed department", + "admin_uid": admin_uid, + "admin_password": admin_password, + }, + headers=headers, + ) + assert response.status_code == 201, response.text + + login_response = await test_client.post( + "/api/auth/token", + data={"username": admin_uid, "password": admin_password}, + ) + assert login_response.status_code == 200, login_response.text + + login_payload = login_response.json() + return { + "department": response.json(), + "admin_id": login_payload["user_id"], + "admin_headers": {"Authorization": f"Bearer {login_payload['access_token']}"}, + } + + +async def _create_user(test_client, headers, label: str, role: str = "user", department_id: int | None = None) -> dict: + suffix = uuid.uuid4().hex[:8] + payload = { + "username": f"u{label}_{suffix}", + "password": f"Pw!{suffix}", + "role": role, + } + if department_id is not None: + payload["department_id"] = department_id + + response = await test_client.post("/api/auth/users", json=payload, headers=headers) + assert response.status_code == 200, response.text + return response.json() + + +async def _cleanup_user(test_client, headers, user_id: int) -> None: + response = await test_client.delete(f"/api/auth/users/{user_id}", headers=headers) + assert response.status_code in {200, 404}, response.text + + +async def _cleanup_department(test_client, headers, department_id: int) -> None: + response = await test_client.delete(f"/api/departments/{department_id}", headers=headers) + assert response.status_code in {200, 404}, response.text + + async def test_login_with_invalid_credentials(test_client): response = await test_client.post("/api/auth/token", data={"username": "invalid", "password": "invalid"}) assert response.status_code == 401 @@ -25,9 +87,7 @@ async def test_user_is_locked_after_repeated_failed_logins(test_client, standard assert response.status_code == 401, response.text assert response.json()["detail"] == "用户名或密码错误" - locked_response = await test_client.post( - "/api/auth/token", data={"username": uid, "password": "wrong-password"} - ) + locked_response = await test_client.post("/api/auth/token", data={"username": uid, "password": "wrong-password"}) assert locked_response.status_code == 423, locked_response.text assert "X-Lock-Remaining" in locked_response.headers assert "账户已被锁定" in locked_response.json()["detail"] @@ -47,7 +107,7 @@ async def test_admin_can_login_and_fetch_profile(test_client, admin_headers): data = profile_response.json() assert data["role"] in {"admin", "superadmin"} assert data["username"] - assert data["user_id"] + assert data["id"] async def test_admin_can_create_and_delete_user(test_client, admin_headers): @@ -71,6 +131,96 @@ async def test_admin_can_create_and_delete_user(test_client, admin_headers): assert delete_payload["message"] == "用户已删除" +async def test_department_admin_is_limited_to_own_department_users(test_client, admin_headers): + await _require_superadmin(test_client, admin_headers) + + user_ids: list[int] = [] + admin_ids: list[int] = [] + department_ids: list[int] = [] + + try: + dept_a = await _create_department_with_admin(test_client, admin_headers, "a") + dept_b = await _create_department_with_admin(test_client, admin_headers, "b") + department_a = dept_a["department"] + department_b = dept_b["department"] + department_ids.extend([department_a["id"], department_b["id"]]) + admin_ids.extend([dept_a["admin_id"], dept_b["admin_id"]]) + + user_a = await _create_user(test_client, dept_a["admin_headers"], "a") + user_b = await _create_user(test_client, dept_b["admin_headers"], "b") + superadmin_created_user = await _create_user(test_client, admin_headers, "s", department_id=department_b["id"]) + user_ids.extend([user_a["id"], user_b["id"], superadmin_created_user["id"]]) + + assert user_a["department_id"] == department_a["id"] + assert superadmin_created_user["department_id"] == department_b["id"] + + forbidden_create = await test_client.post( + "/api/auth/users", + json={ + "username": f"ux_{uuid.uuid4().hex[:8]}", + "password": "routerTest123!", + "role": "user", + "department_id": department_b["id"], + }, + headers=dept_a["admin_headers"], + ) + assert forbidden_create.status_code == 403, forbidden_create.text + + list_response = await test_client.get("/api/auth/users", headers=dept_a["admin_headers"]) + assert list_response.status_code == 200, list_response.text + listed_users = list_response.json() + listed_user_ids = {user["id"] for user in listed_users} + assert user_a["id"] in listed_user_ids + assert user_b["id"] not in listed_user_ids + assert all(user["department_id"] == department_a["id"] for user in listed_users) + + options_response = await test_client.get("/api/auth/users/access-options", headers=dept_a["admin_headers"]) + assert options_response.status_code == 200, options_response.text + access_options = options_response.json() + option_uids = {user["uid"] for user in access_options} + assert user_a["uid"] in option_uids + assert user_b["uid"] not in option_uids + assert all(user["department_id"] == department_a["id"] for user in access_options) + + superadmin_list_response = await test_client.get("/api/auth/users", headers=admin_headers) + assert superadmin_list_response.status_code == 200, superadmin_list_response.text + superadmin_user_ids = {user["id"] for user in superadmin_list_response.json()} + assert user_a["id"] in superadmin_user_ids + assert user_b["id"] in superadmin_user_ids + + own_read = await test_client.get(f"/api/auth/users/{user_a['id']}", headers=dept_a["admin_headers"]) + assert own_read.status_code == 200, own_read.text + + cross_read = await test_client.get(f"/api/auth/users/{user_b['id']}", headers=dept_a["admin_headers"]) + assert cross_read.status_code == 403, cross_read.text + + cross_update = await test_client.put( + f"/api/auth/users/{user_b['id']}", + json={"username": f"ub_{uuid.uuid4().hex[:8]}"}, + headers=dept_a["admin_headers"], + ) + assert cross_update.status_code == 403, cross_update.text + + role_escalation = await test_client.put( + f"/api/auth/users/{user_a['id']}", json={"role": "admin"}, headers=dept_a["admin_headers"] + ) + assert role_escalation.status_code == 403, role_escalation.text + + cross_delete = await test_client.delete(f"/api/auth/users/{user_b['id']}", headers=dept_a["admin_headers"]) + assert cross_delete.status_code == 403, cross_delete.text + + own_delete = await test_client.delete(f"/api/auth/users/{user_a['id']}", headers=dept_a["admin_headers"]) + assert own_delete.status_code == 200, own_delete.text + user_ids.remove(user_a["id"]) + finally: + for user_id in user_ids: + await _cleanup_user(test_client, admin_headers, user_id) + for admin_id in admin_ids: + await _cleanup_user(test_client, admin_headers, admin_id) + for department_id in department_ids: + await _cleanup_department(test_client, admin_headers, department_id) + + async def test_invalid_token_is_rejected(test_client): headers = {"Authorization": "Bearer not-a-real-token"} response = await test_client.get("/api/auth/me", headers=headers) diff --git a/backend/test/integration/api/test_knowledge_router.py b/backend/test/integration/api/test_knowledge_router.py index f45ec452..b6c61469 100644 --- a/backend/test/integration/api/test_knowledge_router.py +++ b/backend/test/integration/api/test_knowledge_router.py @@ -20,6 +20,94 @@ def _assert_forbidden_response(response): assert isinstance(payload["detail"], str) +async def _create_test_department(test_client, admin_headers, prefix="pytest_dept"): + suffix = uuid.uuid4().hex[:8] + admin_uid = f"deptadmin_{suffix}" + response = await test_client.post( + "/api/departments", + json={ + "name": f"{prefix}_{suffix}", + "description": "pytest department", + "admin_uid": admin_uid, + "admin_password": f"Pw!{suffix}", + }, + headers=admin_headers, + ) + assert response.status_code == 201, response.text + payload = response.json() + payload["admin_uid"] = admin_uid + return payload + + +async def _create_test_user(test_client, admin_headers, department_id): + suffix = uuid.uuid4().hex[:8] + password = f"Pw!{suffix}" + response = await test_client.post( + "/api/auth/users", + json={ + "username": f"pytest_user_{suffix}", + "password": password, + "role": "user", + "department_id": department_id, + }, + headers=admin_headers, + ) + assert response.status_code == 200, response.text + user = response.json() + + login_response = await test_client.post( + "/api/auth/token", + data={"username": user["uid"], "password": password}, + ) + assert login_response.status_code == 200, login_response.text + return {"user": user, "headers": {"Authorization": f"Bearer {login_response.json()['access_token']}"}} + + +async def _delete_user_by_id(test_client, admin_headers, user_id): + response = await test_client.delete(f"/api/auth/users/{user_id}", headers=admin_headers) + assert response.status_code in (200, 404), response.text + + +async def _find_user_id_by_uid(test_client, admin_headers, uid): + response = await test_client.get("/api/auth/users", headers=admin_headers) + assert response.status_code == 200, response.text + for user in response.json(): + if user["uid"] == uid: + return user["id"] + return None + + +async def _delete_department_with_admin(test_client, admin_headers, department): + admin_user_id = await _find_user_id_by_uid(test_client, admin_headers, department["admin_uid"]) + if admin_user_id: + await _delete_user_by_id(test_client, admin_headers, admin_user_id) + response = await test_client.delete(f"/api/departments/{department['id']}", headers=admin_headers) + assert response.status_code in (200, 404), response.text + + +async def _create_test_database(test_client, admin_headers, share_config=None): + response = await test_client.post( + "/api/knowledge/databases", + json={ + "database_name": f"pytest_acl_{uuid.uuid4().hex[:8]}", + "description": "Knowledge permission test", + "embedding_model_spec": "siliconflow-cn:Pro/BAAI/bge-m3", + "kb_type": "milvus", + "additional_params": {}, + "share_config": share_config, + }, + headers=admin_headers, + ) + assert response.status_code == 200, response.text + return response.json() + + +async def _accessible_db_ids(test_client, headers): + response = await test_client.get("/api/knowledge/databases/accessible", headers=headers) + assert response.status_code == 200, response.text + return {item["db_id"] for item in response.json().get("databases", [])} + + async def test_admin_can_manage_knowledge_databases(test_client, admin_headers, knowledge_database): db_id = knowledge_database["db_id"] @@ -395,6 +483,96 @@ async def test_get_accessible_databases(test_client, admin_headers, knowledge_da assert knowledge_database["db_id"] in db_ids +async def test_create_database_defaults_to_global_share_config(test_client, admin_headers): + database = await _create_test_database(test_client, admin_headers) + db_id = database["db_id"] + try: + assert database["share_config"] == {"access_level": "global", "department_ids": [], "user_uids": []} + finally: + await test_client.delete(f"/api/knowledge/databases/{db_id}", headers=admin_headers) + + +async def test_department_share_config_filters_accessible_databases(test_client, admin_headers): + department_a = await _create_test_department(test_client, admin_headers, "pytest_dept_a") + department_b = await _create_test_department(test_client, admin_headers, "pytest_dept_b") + user_a = user_b = None + database = None + + try: + user_a = await _create_test_user(test_client, admin_headers, department_a["id"]) + user_b = await _create_test_user(test_client, admin_headers, department_b["id"]) + database = await _create_test_database( + test_client, + admin_headers, + {"access_level": "department", "department_ids": [department_a["id"]], "user_uids": []}, + ) + + saved_config = database["share_config"] + assert saved_config["access_level"] == "department" + assert department_a["id"] in saved_config["department_ids"] + + assert database["db_id"] in await _accessible_db_ids(test_client, user_a["headers"]) + assert database["db_id"] not in await _accessible_db_ids(test_client, user_b["headers"]) + finally: + if database: + await test_client.delete(f"/api/knowledge/databases/{database['db_id']}", headers=admin_headers) + if user_a: + await _delete_user_by_id(test_client, admin_headers, user_a["user"]["id"]) + if user_b: + await _delete_user_by_id(test_client, admin_headers, user_b["user"]["id"]) + await _delete_department_with_admin(test_client, admin_headers, department_a) + await _delete_department_with_admin(test_client, admin_headers, department_b) + + +async def test_user_share_config_filters_accessible_databases(test_client, admin_headers): + department_a = await _create_test_department(test_client, admin_headers, "pytest_dept_a") + department_b = await _create_test_department(test_client, admin_headers, "pytest_dept_b") + user_a = user_b = None + database = None + + try: + user_a = await _create_test_user(test_client, admin_headers, department_a["id"]) + user_b = await _create_test_user(test_client, admin_headers, department_b["id"]) + database = await _create_test_database( + test_client, + admin_headers, + {"access_level": "user", "department_ids": [], "user_uids": [user_a["user"]["uid"]]}, + ) + + saved_config = database["share_config"] + assert saved_config["access_level"] == "user" + assert user_a["user"]["uid"] in saved_config["user_uids"] + + assert database["db_id"] in await _accessible_db_ids(test_client, user_a["headers"]) + assert database["db_id"] not in await _accessible_db_ids(test_client, user_b["headers"]) + finally: + if database: + await test_client.delete(f"/api/knowledge/databases/{database['db_id']}", headers=admin_headers) + if user_a: + await _delete_user_by_id(test_client, admin_headers, user_a["user"]["id"]) + if user_b: + await _delete_user_by_id(test_client, admin_headers, user_b["user"]["id"]) + await _delete_department_with_admin(test_client, admin_headers, department_a) + await _delete_department_with_admin(test_client, admin_headers, department_b) + + +async def test_user_access_options_include_all_departments_for_admin(test_client, admin_headers): + department = await _create_test_department(test_client, admin_headers, "pytest_access_options") + user = None + + try: + user = await _create_test_user(test_client, admin_headers, department["id"]) + response = await test_client.get("/api/auth/users/access-options", headers=admin_headers) + assert response.status_code == 200, response.text + uids = {item["uid"] for item in response.json()} + assert user["user"]["uid"] in uids + assert department["admin_uid"] in uids + finally: + if user: + await _delete_user_by_id(test_client, admin_headers, user["user"]["id"]) + await _delete_department_with_admin(test_client, admin_headers, department) + + async def test_get_knowledge_base_types(test_client, admin_headers): """测试获取支持的知识库类型""" response = await test_client.get("/api/knowledge/types", headers=admin_headers) diff --git a/web/src/apis/auth_api.js b/web/src/apis/auth_api.js index 6142c4bd..faaab76d 100644 --- a/web/src/apis/auth_api.js +++ b/web/src/apis/auth_api.js @@ -2,6 +2,8 @@ * 认证相关 API */ +import { apiAdminGet } from './base' + async function parseErrorDetail(response, fallbackMessage) { const contentType = response.headers.get('content-type') || '' @@ -57,6 +59,10 @@ async function getOIDCLoginUrl(redirectPath = '/') { * department_name: string | null * }>} */ +async function getUserAccessOptions() { + return apiAdminGet('/api/auth/users/access-options') +} + async function exchangeOIDCCode(code) { const response = await fetch('/api/auth/oidc/exchange-code', { method: 'POST', @@ -77,5 +83,6 @@ async function exchangeOIDCCode(code) { export const authApi = { getOIDCConfig, getOIDCLoginUrl, + getUserAccessOptions, exchangeOIDCCode } diff --git a/web/src/components/KnowledgeBaseCard.vue b/web/src/components/KnowledgeBaseCard.vue index 1d4aa4c1..95119701 100644 --- a/web/src/components/KnowledgeBaseCard.vue +++ b/web/src/components/KnowledgeBaseCard.vue @@ -39,7 +39,7 @@ - +