diff --git a/app/algorithms/isolation/valve.py b/app/algorithms/isolation/valve.py index 57cd1e1..c53c0f4 100644 --- a/app/algorithms/isolation/valve.py +++ b/app/algorithms/isolation/valve.py @@ -149,14 +149,16 @@ def valve_isolation_analysis( must_close_valves.sort() optional_valves.sort() + isolatable = bool(must_close_valves) result = { "accident_elements": target_elements, "disabled_valves": disabled_valves, - "affected_nodes": sorted(affected_nodes), + "affected_nodes": sorted(affected_nodes) if isolatable else [], + "affected_node_count": len(affected_nodes), "must_close_valves": must_close_valves, "optional_valves": optional_valves, - "isolatable": len(must_close_valves) > 0, + "isolatable": isolatable, } if len(target_elements) == 1: diff --git a/app/api/v1/endpoints/access.py b/app/api/v1/endpoints/access.py new file mode 100644 index 0000000..fcf6a56 --- /dev/null +++ b/app/api/v1/endpoints/access.py @@ -0,0 +1,39 @@ +from fastapi import APIRouter, Depends, Header + +from app.auth.metadata_dependencies import ( + get_current_metadata_user, + get_metadata_repository, +) +from app.auth.permissions import resolve_permissions +from app.auth.project_dependencies import resolve_project_context +from app.domain.schemas.access import AccessContextResponse +from app.infra.db.metadb.repositories.metadata_repository import MetadataRepository + +router = APIRouter() + + +@router.get("/access/context", response_model=AccessContextResponse) +async def get_access_context( + x_project_id: str | None = Header(default=None, alias="X-Project-Id"), + current_user=Depends(get_current_metadata_user), + metadata_repo: MetadataRepository = Depends(get_metadata_repository), +) -> AccessContextResponse: + project_context = ( + await resolve_project_context(x_project_id, current_user, metadata_repo) + if x_project_id + else None + ) + permissions = resolve_permissions( + project_role=project_context.project_role if project_context else None, + system_role=current_user.role, + is_superuser=current_user.is_superuser, + ) + return AccessContextResponse( + user_id=current_user.id, + username=current_user.username, + system_role=current_user.role, + is_system_admin=current_user.is_superuser or current_user.role == "admin", + project_id=project_context.project_id if project_context else None, + project_role=project_context.project_role if project_context else None, + permissions=sorted(permissions), + ) diff --git a/app/api/v1/endpoints/admin_metadata.py b/app/api/v1/endpoints/admin_metadata.py index 9ecd0ee..930a3d8 100644 --- a/app/api/v1/endpoints/admin_metadata.py +++ b/app/api/v1/endpoints/admin_metadata.py @@ -266,6 +266,7 @@ async def create_admin_project( gs_workspace=payload.gs_workspace, map_extent=payload.map_extent, status=payload.status, + creator_user_id=current_user.id, ) except IntegrityError as exc: raise HTTPException( diff --git a/app/api/v1/endpoints/agent_auth.py b/app/api/v1/endpoints/agent_auth.py index c2637de..3aed53b 100644 --- a/app/api/v1/endpoints/agent_auth.py +++ b/app/api/v1/endpoints/agent_auth.py @@ -9,6 +9,7 @@ from app.auth.project_dependencies import ( ProjectContext, get_project_context, ) +from app.auth.permissions import permissions_for_context router = APIRouter() @@ -22,6 +23,7 @@ class AgentAuthContextResponse(BaseModel): project_id: str network: str project_role: str + permissions: list[str] token_expires_at: str | None = None @@ -46,5 +48,6 @@ async def get_agent_auth_context( project_id=str(ctx.project_id), network=ctx.project_code, project_role=ctx.project_role, + permissions=sorted(permissions_for_context(ctx)), token_expires_at=token_expires_at, ) diff --git a/app/api/v1/endpoints/audit.py b/app/api/v1/endpoints/audit.py index dd0d7d6..871ff68 100644 --- a/app/api/v1/endpoints/audit.py +++ b/app/api/v1/endpoints/audit.py @@ -1,29 +1,30 @@ -""" -审计日志 API 接口 - -仅管理员可访问 -""" - -from typing import List, Optional -from uuid import UUID from datetime import datetime -from fastapi import APIRouter, Depends, Query, Path -from app.domain.schemas.audit import AuditLogResponse -from app.infra.db.metadb.repositories.audit_repository import AuditRepository +from typing import Literal +from uuid import UUID + +from fastapi import APIRouter, Depends, Query, Request, status +from pydantic import BaseModel +from sqlalchemy.ext.asyncio import AsyncSession + from app.auth.metadata_dependencies import ( get_current_metadata_admin, get_current_metadata_user, ) +from app.core.audit import AuditAction, log_audit_event +from app.domain.schemas.audit import AuditLogResponse from app.infra.db.metadb.database import get_metadata_session -from sqlalchemy.ext.asyncio import AsyncSession +from app.infra.db.metadb.repositories.audit_repository import AuditRepository router = APIRouter() +class SessionAuditEventRequest(BaseModel): + event: Literal["login", "logout"] + + async def get_audit_repository( session: AsyncSession = Depends(get_metadata_session), ) -> AuditRepository: - """获取审计日志仓储""" return AuditRepository(session) @@ -31,26 +32,21 @@ async def get_audit_repository( "/logs", summary="查询审计日志", description="查询审计日志(仅管理员)", - response_model=List[AuditLogResponse], + response_model=list[AuditLogResponse], ) async def get_audit_logs( - user_id: Optional[UUID] = Query(None, description="按用户ID过滤"), - project_id: Optional[UUID] = Query(None, description="按项目ID过滤"), - action: Optional[str] = Query(None, description="按操作类型过滤"), - resource_type: Optional[str] = Query(None, description="按资源类型过滤"), - start_time: Optional[datetime] = Query(None, description="开始时间"), - end_time: Optional[datetime] = Query(None, description="结束时间"), + user_id: UUID | None = Query(None, description="按用户ID过滤"), + project_id: UUID | None = Query(None, description="按项目ID过滤"), + action: str | None = Query(None, description="按操作类型过滤"), + resource_type: str | None = Query(None, description="按资源类型过滤"), + start_time: datetime | None = Query(None, description="开始时间"), + end_time: datetime | None = Query(None, description="结束时间"), skip: int = Query(0, ge=0, description="跳过记录数"), limit: int = Query(100, ge=1, le=1000, description="限制记录数"), - current_user=Depends(get_current_metadata_admin), + _current_user=Depends(get_current_metadata_admin), audit_repo: AuditRepository = Depends(get_audit_repository), -) -> List[AuditLogResponse]: - """ - 查询审计日志 - - 支持按用户、时间、操作类型等条件过滤,仅管理员可访问 - """ - logs = await audit_repo.get_logs( +) -> list[AuditLogResponse]: + return await audit_repo.get_logs( user_id=user_id, project_id=project_id, action=action, @@ -60,7 +56,6 @@ async def get_audit_logs( skip=skip, limit=limit, ) - return logs @router.get( @@ -69,20 +64,15 @@ async def get_audit_logs( description="获取审计日志总数(仅管理员)", ) async def get_audit_logs_count( - user_id: Optional[UUID] = Query(None, description="按用户ID过滤"), - project_id: Optional[UUID] = Query(None, description="按项目ID过滤"), - action: Optional[str] = Query(None, description="按操作类型过滤"), - resource_type: Optional[str] = Query(None, description="按资源类型过滤"), - start_time: Optional[datetime] = Query(None, description="开始时间"), - end_time: Optional[datetime] = Query(None, description="结束时间"), - current_user=Depends(get_current_metadata_admin), + user_id: UUID | None = Query(None, description="按用户ID过滤"), + project_id: UUID | None = Query(None, description="按项目ID过滤"), + action: str | None = Query(None, description="按操作类型过滤"), + resource_type: str | None = Query(None, description="按资源类型过滤"), + start_time: datetime | None = Query(None, description="开始时间"), + end_time: datetime | None = Query(None, description="结束时间"), + _current_user=Depends(get_current_metadata_admin), audit_repo: AuditRepository = Depends(get_audit_repository), ) -> dict: - """ - 获取审计日志总数 - - 获取符合条件的审计日志的总数,仅管理员可访问 - """ count = await audit_repo.get_log_count( user_id=user_id, project_id=project_id, @@ -94,27 +84,42 @@ async def get_audit_logs_count( return {"count": count} +@router.post("/session-events", status_code=status.HTTP_204_NO_CONTENT) +async def record_session_event( + payload: SessionAuditEventRequest, + request: Request, + current_user=Depends(get_current_metadata_user), + session: AsyncSession = Depends(get_metadata_session), +) -> None: + await log_audit_event( + action=AuditAction.LOGIN if payload.event == "login" else AuditAction.LOGOUT, + user_id=current_user.id, + resource_type="session", + resource_id=str(current_user.keycloak_id), + ip_address=request.client.host if request.client else None, + request_method=request.method, + request_path=request.url.path, + response_status=status.HTTP_204_NO_CONTENT, + session=session, + ) + + @router.get( "/logs/my", summary="查询我的审计日志", description="查询当前用户的审计日志", - response_model=List[AuditLogResponse], + response_model=list[AuditLogResponse], ) async def get_my_audit_logs( - action: Optional[str] = Query(None, description="按操作类型过滤"), - start_time: Optional[datetime] = Query(None, description="开始时间"), - end_time: Optional[datetime] = Query(None, description="结束时间"), + action: str | None = Query(None, description="按操作类型过滤"), + start_time: datetime | None = Query(None, description="开始时间"), + end_time: datetime | None = Query(None, description="结束时间"), skip: int = Query(0, ge=0, description="跳过记录数"), limit: int = Query(100, ge=1, le=1000, description="限制记录数"), current_user=Depends(get_current_metadata_user), audit_repo: AuditRepository = Depends(get_audit_repository), -) -> List[AuditLogResponse]: - """ - 查询当前用户的审计日志 - - 普通用户只能查看自己的操作记录 - """ - logs = await audit_repo.get_logs( +) -> list[AuditLogResponse]: + return await audit_repo.get_logs( user_id=current_user.id, action=action, start_time=start_time, @@ -122,4 +127,3 @@ async def get_my_audit_logs( skip=skip, limit=limit, ) - return logs diff --git a/app/api/v1/endpoints/model_import.py b/app/api/v1/endpoints/model_import.py new file mode 100644 index 0000000..35fc9fc --- /dev/null +++ b/app/api/v1/endpoints/model_import.py @@ -0,0 +1,297 @@ +import json +from pathlib import Path +from tempfile import NamedTemporaryFile +from uuid import UUID, uuid4 + +from fastapi import ( + APIRouter, + Body, + Depends, + File, + Header, + HTTPException, + Path as ApiPath, + Query, + Request, + UploadFile, + status, +) + +from app.auth.metadata_dependencies import ( + get_current_metadata_admin, + get_metadata_repository, +) +from app.core.audit import AuditAction, log_audit_event +from app.infra.db.metadb.repositories.metadata_repository import MetadataRepository +from app.services.network_import import network_update +from app.services.tjnetwork import ChangeSet, import_inp, run_inp + +router = APIRouter() + +MAX_INP_FILE_BYTES = 50 * 1024 * 1024 +INP_SECTIONS = ("[TITLE]", "[JUNCTIONS]", "[RESERVOIRS]", "[TANKS]", "[PIPES]") + + +async def _get_active_project(project_id: UUID, metadata_repo: MetadataRepository): + project = await metadata_repo.get_project_by_id(project_id) + if project is None: + raise HTTPException( + status_code=status.HTTP_404_NOT_FOUND, + detail="Project not found", + ) + if project.status != "active": + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail="Project is not active", + ) + return project + + +def _validate_inp_bytes(content: bytes, filename: str) -> str: + if Path(filename).suffix.lower() != ".inp": + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="Only .inp model files are accepted", + ) + if not content: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="INP file is empty", + ) + if len(content) > MAX_INP_FILE_BYTES: + raise HTTPException( + status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE, + detail="INP file exceeds the 50 MiB limit", + ) + for encoding in ("utf-8-sig", "gb18030"): + try: + text = content.decode(encoding) + break + except UnicodeDecodeError: + continue + else: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="INP file encoding is not supported", + ) + upper_text = text.upper() + if not any(section in upper_text for section in INP_SECTIONS): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="Invalid INP file structure", + ) + return text + + +async def _read_upload(file: UploadFile) -> tuple[bytes, str]: + filename = Path(file.filename or "").name + content = await file.read(MAX_INP_FILE_BYTES + 1) + _validate_inp_bytes(content, filename) + return content, filename + + +async def _audit_model_change( + *, + request: Request, + current_user, + metadata_repo: MetadataRepository, + project_id: UUID, + action: str, +) -> None: + await log_audit_event( + action=AuditAction.UPDATE, + user_id=current_user.id, + project_id=project_id, + resource_type="hydraulic_model", + resource_id=action, + request_data={"operation": action}, + ip_address=request.client.host if request.client else None, + request_method=request.method, + request_path=request.url.path, + response_status=status.HTTP_200_OK, + session=metadata_repo.session, + ) + + +async def _run_uploaded_inp(content: bytes) -> str: + target_dir = Path("inp") + target_dir.mkdir(parents=True, exist_ok=True) + model_name = f"admin_model_{uuid4().hex}" + target_path = target_dir / f"{model_name}.inp" + target_path.write_bytes(content) + return run_inp(model_name) + + +async def _update_from_inp(content: bytes) -> None: + temp_path: Path | None = None + try: + with NamedTemporaryFile(suffix=".inp", delete=False) as temp_file: + temp_file.write(content) + temp_path = Path(temp_file.name) + network_update(str(temp_path)) + finally: + if temp_path is not None: + temp_path.unlink(missing_ok=True) + + +async def _apply_model_update(content: bytes) -> None: + try: + await _update_from_inp(content) + except Exception as exc: + raise HTTPException( + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, + detail=f"数据库操作失败: {exc}", + ) from exc + + +@router.post( + "/admin/projects/{project_id}/model/import", + summary="导入桌面端水力模型", +) +async def import_project_model( + request: Request, + project_id: UUID = ApiPath(...), + file: UploadFile = File(..., description="桌面端导出的 INP 模型文件"), + current_user=Depends(get_current_metadata_admin), + metadata_repo: MetadataRepository = Depends(get_metadata_repository), +) -> dict: + project = await _get_active_project(project_id, metadata_repo) + content, filename = await _read_upload(file) + result = await _run_uploaded_inp(content) + await _audit_model_change( + request=request, + current_user=current_user, + metadata_repo=metadata_repo, + project_id=project.id, + action="import", + ) + return {"project_id": str(project.id), "filename": filename, "result": result} + + +@router.post( + "/admin/projects/{project_id}/model/update", + summary="更新桌面端水力模型", +) +async def update_project_model( + request: Request, + project_id: UUID = ApiPath(...), + file: UploadFile = File(..., description="桌面端导出的 INP 模型文件"), + current_user=Depends(get_current_metadata_admin), + metadata_repo: MetadataRepository = Depends(get_metadata_repository), +) -> dict: + project = await _get_active_project(project_id, metadata_repo) + content, filename = await _read_upload(file) + await _apply_model_update(content) + await _audit_model_change( + request=request, + current_user=current_user, + metadata_repo=metadata_repo, + project_id=project.id, + action="update", + ) + return {"project_id": str(project.id), "filename": filename, "updated": True} + + +@router.post("/importinp/", deprecated=True) +async def legacy_import_inp( + request: Request, + network: str = Query(...), + x_project_id: UUID = Header(..., alias="X-Project-Id"), + current_user=Depends(get_current_metadata_admin), + metadata_repo: MetadataRepository = Depends(get_metadata_repository), +): + project = await _get_active_project(x_project_id, metadata_repo) + if network != project.code: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail="Project scope denied", + ) + payload = await request.json() + inp_text = payload.get("inp") if isinstance(payload, dict) else None + if not isinstance(inp_text, str): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="Missing INP content", + ) + _validate_inp_bytes(inp_text.encode("utf-8"), "model.inp") + result = import_inp(network, ChangeSet({"inp": inp_text})) + await _audit_model_change( + request=request, + current_user=current_user, + metadata_repo=metadata_repo, + project_id=project.id, + action="import", + ) + return result + + +@router.post("/uploadinp/", deprecated=True) +async def legacy_upload_inp( + request: Request, + content: bytes = Body(...), + name: str = Query(...), + x_project_id: UUID = Header(..., alias="X-Project-Id"), + current_user=Depends(get_current_metadata_admin), + metadata_repo: MetadataRepository = Depends(get_metadata_repository), +) -> bool: + project = await _get_active_project(x_project_id, metadata_repo) + safe_name = Path(name).name + if safe_name != name: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail="Invalid INP file name", + ) + _validate_inp_bytes(content, safe_name) + target_dir = Path("data") + target_dir.mkdir(parents=True, exist_ok=True) + (target_dir / safe_name).write_bytes(content) + await _audit_model_change( + request=request, + current_user=current_user, + metadata_repo=metadata_repo, + project_id=project.id, + action="upload", + ) + return True + + +@router.post("/network_project/", deprecated=True) +async def legacy_network_project( + request: Request, + file: UploadFile = File(...), + x_project_id: UUID = Header(..., alias="X-Project-Id"), + current_user=Depends(get_current_metadata_admin), + metadata_repo: MetadataRepository = Depends(get_metadata_repository), +): + project = await _get_active_project(x_project_id, metadata_repo) + content, _ = await _read_upload(file) + result = await _run_uploaded_inp(content) + await _audit_model_change( + request=request, + current_user=current_user, + metadata_repo=metadata_repo, + project_id=project.id, + action="import", + ) + return result + + +@router.post("/network_update/", deprecated=True) +async def legacy_network_update( + request: Request, + file: UploadFile = File(...), + x_project_id: UUID = Header(..., alias="X-Project-Id"), + current_user=Depends(get_current_metadata_admin), + metadata_repo: MetadataRepository = Depends(get_metadata_repository), +) -> str: + project = await _get_active_project(x_project_id, metadata_repo) + content, _ = await _read_upload(file) + await _apply_model_update(content) + await _audit_model_change( + request=request, + current_user=current_user, + metadata_repo=metadata_repo, + project_id=project.id, + action="update", + ) + return json.dumps({"message": "管网更新成功"}) diff --git a/app/api/v1/endpoints/project.py b/app/api/v1/endpoints/project.py index c4424ba..b0afe84 100644 --- a/app/api/v1/endpoints/project.py +++ b/app/api/v1/endpoints/project.py @@ -1,9 +1,13 @@ import json -from fastapi import APIRouter, Request, HTTPException, Query, Path, Body, Depends +from fastapi import APIRouter, Request, HTTPException, Query, Path, Depends from fastapi.responses import PlainTextResponse from typing import Any, Dict, List from app.infra.db.metadb.repositories.metadata_repository import MetadataRepository from app.auth.project_dependencies import get_metadata_repository +from app.auth.permissions import ( + ENVIRONMENT_MANAGE, + require_permission, +) from app.domain.schemas.metadata import ProjectMetaResponse import app.services.project_info as project_info from app.infra.db.postgresql.database import get_database_instance as get_pg_db @@ -18,7 +22,6 @@ from app.services.tjnetwork import ( open_project, close_project, copy_project, - import_inp, export_inp, read_inp, dump_inp, @@ -89,7 +92,8 @@ async def have_project_endpoint( @router.post("/createproject/", summary="创建新项目", description="创建一个新的供水管网项目。如果项目已存在,可能会覆盖或报错(取决于底层实现)。") async def create_project_endpoint( - network: str = Query(..., description="管网名称(或数据库名称)") + network: str = Query(..., description="管网名称(或数据库名称)"), + _=Depends(require_permission(ENVIRONMENT_MANAGE)), ): """ 创建新项目 @@ -101,7 +105,8 @@ async def create_project_endpoint( @router.post("/deleteproject/", summary="删除项目", description="永久删除指定的供水管网项目。此操作不可恢复。") async def delete_project_endpoint( - network: str = Query(..., description="管网名称(或数据库名称)") + network: str = Query(..., description="管网名称(或数据库名称)"), + _=Depends(require_permission(ENVIRONMENT_MANAGE)), ): """ 删除项目 @@ -172,7 +177,8 @@ async def close_project_endpoint( @router.post("/copyproject/", summary="复制项目", description="将现有项目复制为新项目。") async def copy_project_endpoint( source: str = Query(..., description="管网名称(或数据库名称)"), - target: str = Query(..., description="管网名称(或数据库名称)") + target: str = Query(..., description="管网名称(或数据库名称)"), + _=Depends(require_permission(ENVIRONMENT_MANAGE)), ): """ 复制项目 @@ -183,24 +189,6 @@ async def copy_project_endpoint( copy_project(source, target) return True -@router.post("/importinp/", summary="导入 INP 文件内容", description="将 INP 格式的文本内容导入到指定项目中。") -async def import_inp_endpoint( - req: Request, - network: str = Query(..., description="管网名称(或数据库名称)") -): - """ - 导入 INP 文件内容 - - - **network**: 管网名称(或数据库名称) - - **req**: 请求体,需包含 `{"inp": "..."}` 结构 - """ - jo_root = await req.json() - inp_text = jo_root["inp"] - ps = {"inp": inp_text} - ret = import_inp(network, ChangeSet(ps)) - print(ret) - return ret - @router.get("/exportinp/", response_model=None, summary="导出项目为 ChangeSet", description="导出项目的变更集 (ChangeSet),包含顶点、SCADA 元素、DMA、SA、VD 等信息。") async def export_inp_endpoint( network: str = Query(..., description="管网名称(或数据库名称)"), @@ -331,26 +319,6 @@ def unlock_project_endpoint( return False -# inp file operations -@router.post("/uploadinp/", status_code=status.HTTP_200_OK, summary="上传 INP 文件", description="上传 INP 文件到服务器数据目录。") -async def fastapi_upload_inp( - afile: bytes = Body(..., description="文件二进制内容"), - name: str = Query(..., description="保存的文件名") -): - """ - 上传 INP 文件 - - - **afile**: 文件内容 - - **name**: 文件名 - """ - if not os.path.exists(inpDir): - os.makedirs(inpDir, exist_ok=True) - - filePath = inpDir + str(name) - with open(filePath, "wb") as f: - f.write(afile) - return True - @router.get("/downloadinp/", status_code=status.HTTP_200_OK, summary="下载 INP 文件", description="从服务器数据目录下载指定的 INP 文件。") async def fastapi_download_inp( name: str = Query(..., description="文件名"), @@ -502,26 +470,6 @@ def unlock_project_endpoint( return False -# inp file operations -@router.post("/uploadinp/", status_code=status.HTTP_200_OK, summary="上传 INP 文件", description="上传 INP 文件到服务器数据目录。") -async def fastapi_upload_inp( - afile: bytes = Body(..., description="文件二进制内容"), - name: str = Query(..., description="保存的文件名") -): - """ - 上传 INP 文件 - - - **afile**: 文件内容 - - **name**: 文件名 - """ - if not os.path.exists(inpDir): - os.makedirs(inpDir, exist_ok=True) - - filePath = inpDir + str(name) - with open(filePath, "wb") as f: - f.write(afile) - return True - @router.get("/downloadinp/", status_code=status.HTTP_200_OK, summary="下载 INP 文件", description="从服务器数据目录下载指定的 INP 文件。") async def fastapi_download_inp( name: str = Query(..., description="文件名"), diff --git a/app/api/v1/endpoints/sensor_placement.py b/app/api/v1/endpoints/sensor_placement.py index 0f0afc3..6fd87c6 100644 --- a/app/api/v1/endpoints/sensor_placement.py +++ b/app/api/v1/endpoints/sensor_placement.py @@ -41,19 +41,14 @@ def _project_network(network: str, project_context: ProjectContext) -> str: return project_context.project_code -def _can_modify_project(project_context: ProjectContext, current_user: Any) -> bool: - return bool( - project_context.project_role in {"owner", "admin", "member"} - or getattr(current_user, "role", None) == "admin" - or getattr(current_user, "is_superuser", False) - ) +def _can_modify_project(project_context: ProjectContext) -> bool: + return project_context.project_role == "member" def _require_project_write( project_context: ProjectContext, - current_user: Any, ) -> None: - if not _can_modify_project(project_context, current_user): + if not _can_modify_project(project_context): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail="当前项目角色为只读,不能修改监测点方案", @@ -88,7 +83,7 @@ def _get_scheme_response( return { **scheme, "can_edit": ( - _can_modify_project(project_context, current_user) + _can_modify_project(project_context) and can_edit_sensor_placement(current_user, scheme) ), } @@ -110,7 +105,7 @@ async def optimize_sensor_placement_scheme( current_user=Depends(get_current_metadata_user), ) -> dict[str, Any]: network = _project_network(payload.network, project_context) - _require_project_write(project_context, current_user) + _require_project_write(project_context) optimizer = ( pressure_sensor_placement_sensitivity if payload.method == "sensitivity" @@ -173,7 +168,7 @@ async def overwrite_sensor_placement_scheme( current_user=Depends(get_current_metadata_user), ) -> dict[str, Any]: network = _project_network(network, project_context) - _require_project_write(project_context, current_user) + _require_project_write(project_context) scheme = _get_scheme_response( network, scheme_id, diff --git a/app/api/v1/endpoints/simulation.py b/app/api/v1/endpoints/simulation.py index 2295aa8..b4133ea 100644 --- a/app/api/v1/endpoints/simulation.py +++ b/app/api/v1/endpoints/simulation.py @@ -1,10 +1,8 @@ from typing import Any, List, Optional from datetime import datetime, timedelta import json -import os -import shutil import threading -from fastapi import APIRouter, Depends, HTTPException, File, UploadFile, Query, Path, Body +from fastapi import APIRouter, Depends, HTTPException, Query, Path, Body from fastapi.responses import PlainTextResponse from app.auth.keycloak_dependencies import get_current_keycloak_username import app.services.simulation as simulation @@ -29,7 +27,6 @@ from app.algorithms.sensor import ( pressure_sensor_placement_kmeans, ) -from app.services.network_import import network_update from app.services.simulation_ops import ( project_management, scheduling_simulation, @@ -282,7 +279,8 @@ async def valve_isolation_endpoint( 返回隔离方案,包括: - must_close_valves: 必须关闭的阀门列表 - optional_valves: 可选关闭的阀门列表 - - affected_nodes: 受影响的节点列表 + - affected_nodes: 受影响的节点列表;不可隔离时为空列表 + - affected_node_count: 受影响的节点总数 - isolatable: 是否可以有效隔离 """ # result = { @@ -549,46 +547,6 @@ async def fastapi_daily_scheduling_analysis(data: DailySchedulingAnalysis = Body ) -@router.post("/network_project/", summary="导入网络项目", description="通过上传INP格式的管网文件导入新的网络项目。系统将自动处理文件并执行模拟。") -async def fastapi_network_project(file: UploadFile = File(..., description="INP格式的管网文件")) -> str: - """ - 导入网络项目 - - - **file**: 上传的INP格式管网文件 - - 系统将上传的文件保存到inp文件夹并执行模拟。 - """ - temp_file_dir = "./inp/" - if not os.path.exists(temp_file_dir): - os.mkdir(temp_file_dir) - temp_file_name = f'network_project_{datetime.now().strftime("%Y%m%d")}' - temp_file_path = f"{temp_file_dir}{temp_file_name}.inp" - with open(temp_file_path, "wb") as buffer: - shutil.copyfileobj(file.file, buffer) - return run_inp(temp_file_name) - - -@router.post("/network_update/", summary="管网更新(高级)", description="通过上传更新文件对管网进行高级的更新操作。系统将处理更新文件并应用到数据库。") -async def fastapi_network_update(file: UploadFile = File(..., description="包含管网更新信息的文件")) -> str: - """ - 管网更新(高级版本) - - - **file**: 包含管网更新信息的文件 - - 系统将处理上传的文件并应用管网更新。 - """ - default_folder = "./" - temp_file_name = f'network_update_{datetime.now().strftime("%Y%m%d")}' - temp_file_path = os.path.join(default_folder, temp_file_name) - try: - with open(temp_file_path, "wb") as buffer: - shutil.copyfileobj(file.file, buffer) - network_update(temp_file_path) - return json.dumps({"message": "管网更新成功"}) - except Exception as exc: - raise HTTPException(status_code=500, detail=f"数据库操作失败: {exc}") - - # @router.get("/pumpfailure/") # async def pump_failure_endpoint(network: str, pump_id: str, time: str): # return pump_failure(network, pump_id, time) diff --git a/app/api/v1/endpoints/snapshots.py b/app/api/v1/endpoints/snapshots.py index 210f58e..e3690be 100644 --- a/app/api/v1/endpoints/snapshots.py +++ b/app/api/v1/endpoints/snapshots.py @@ -1,4 +1,5 @@ -from fastapi import APIRouter, Request, Query +from fastapi import APIRouter, Depends, Request, Query +from app.auth.permissions import SIMULATION_RUN, require_permission from app.services.tjnetwork import ( ChangeSet, get_current_operation, @@ -149,7 +150,11 @@ async def pick_operation_endpoint( return pick_operation(network, operation, discard) @router.get("/syncwithserver/", summary="与服务器同步", description="将网络与服务器同步到指定操作", response_model=None) -async def sync_with_server_endpoint(network: str = Query(..., description="管网名称(或数据库名称)"), operation: int = Query(..., description="目标操作ID")) -> ChangeSet: +async def sync_with_server_endpoint( + network: str = Query(..., description="管网名称(或数据库名称)"), + operation: int = Query(..., description="目标操作ID"), + _=Depends(require_permission(SIMULATION_RUN)), +) -> ChangeSet: """ 与服务器同步 diff --git a/app/api/v1/router.py b/app/api/v1/router.py index df73930..887be17 100644 --- a/app/api/v1/router.py +++ b/app/api/v1/router.py @@ -1,120 +1,240 @@ -from fastapi import APIRouter +from fastapi import APIRouter, Depends + from app.api.v1.endpoints import ( + access, admin_metadata, agent_auth, - project, - simulation, - scada, - sensor_placement, - extension, - snapshots, - # data_query, - users, - schemes, - misc, - risk, - cache, - leakage, + audit, burst_detection, burst_location, - audit, # 新增:审计日志 - meta, - web_search, + cache, + extension, geocoding, -) -from app.api.v1.endpoints.network import ( - general, - junctions, - reservoirs, - tanks, - pipes, - pumps, - valves, - tags, - demands, - geometry, - regions, + leakage, + meta, + misc, + model_import, + project, + project_data, + risk, + scada, + schemes, + sensor_placement, + simulation, + snapshots, + users, + web_search, ) from app.api.v1.endpoints.components import ( - curves, - patterns, controls, + curves, options, + patterns, quality, visuals, ) - -from app.api.v1.endpoints import project_data +from app.api.v1.endpoints.network import ( + demands, + general, + geometry, + junctions, + pipes, + pumps, + regions, + reservoirs, + tags, + tanks, + valves, +) from app.api.v1.endpoints.timeseries import ( - realtime as ts_realtime, - scheme as ts_scheme, - scada as ts_scada, composite as ts_composite, + realtime as ts_realtime, + scada as ts_scada, + scheme as ts_scheme, +) +from app.auth.permissions import ( + BURST_RUN, + OPTIMIZATION_RUN, + RISK_RUN, + SCADA_CLEAN, + SCADA_VIEW, + SIMULATION_RUN, + SIMULATION_VIEW, + WEBGIS_EDIT, + WEBGIS_VIEW, + require_method_permission, + require_permission, ) api_router = APIRouter() -# Core Services +webgis_access = Depends( + require_method_permission( + read_permission=WEBGIS_VIEW, + write_permission=WEBGIS_EDIT, + ) +) +scada_access = Depends( + require_method_permission( + read_permission=SCADA_VIEW, + write_permission=SCADA_CLEAN, + ) +) +simulation_access = Depends( + require_method_permission( + read_permission=SIMULATION_VIEW, + write_permission=SIMULATION_RUN, + ) +) + +webgis_view_access = Depends(require_permission(WEBGIS_VIEW)) +simulation_run_access = Depends(require_permission(SIMULATION_RUN)) +burst_run_access = Depends(require_permission(BURST_RUN)) +risk_run_access = Depends(require_permission(RISK_RUN)) +optimization_run_access = Depends(require_permission(OPTIMIZATION_RUN)) + +# Core services +api_router.include_router(access.router, tags=["Access Control"]) api_router.include_router(agent_auth.router, tags=["Agent Auth"]) api_router.include_router( - admin_metadata.router, prefix="/admin", tags=["Metadata Admin"] + admin_metadata.router, + prefix="/admin", + tags=["Metadata Admin"], ) -api_router.include_router(audit.router, prefix="/audit", tags=["Audit Logs"]) # 新增 +api_router.include_router(model_import.router, tags=["Model Administration"]) +api_router.include_router(audit.router, prefix="/audit", tags=["Audit Logs"]) api_router.include_router(meta.router, tags=["Metadata"]) -api_router.include_router(project.router, tags=["Project"]) - -# Network Elements (Node/Link Types) -api_router.include_router(general.router, tags=["Network General"]) -api_router.include_router(junctions.router, tags=["Junctions"]) -api_router.include_router(reservoirs.router, tags=["Reservoirs"]) -api_router.include_router(tanks.router, tags=["Tanks"]) -api_router.include_router(pipes.router, tags=["Pipes"]) -api_router.include_router(pumps.router, tags=["Pumps"]) -api_router.include_router(valves.router, tags=["Valves"]) - -# Network Features -api_router.include_router(tags.router, tags=["Tags"]) -api_router.include_router(demands.router, tags=["Demands"]) -api_router.include_router(geometry.router, tags=["Geometry & Coordinates"]) -api_router.include_router(regions.router, tags=["Regions & DMAs"]) - -# Components & Controls -api_router.include_router(curves.router, tags=["Curves"]) -api_router.include_router(patterns.router, tags=["Patterns"]) -api_router.include_router(controls.router, tags=["Controls & Rules"]) -api_router.include_router(options.router, tags=["Options"]) -api_router.include_router(quality.router, tags=["Quality"]) -api_router.include_router(visuals.router, tags=["Visuals"]) - -# Simulation & Data -api_router.include_router(simulation.router, tags=["Simulation Control"]) -# api_router.include_router(data_query.router, tags=["Data Query & InfluxDB"]) -api_router.include_router(scada.router) -api_router.include_router(sensor_placement.router, tags=["Sensor Placement"]) -api_router.include_router(snapshots.router, tags=["Snapshots"]) -api_router.include_router(users.router, tags=["Users"]) -api_router.include_router(schemes.router, tags=["Schemes"]) -api_router.include_router(misc.router, tags=["Misc"]) -api_router.include_router(risk.router, tags=["Risk"]) -api_router.include_router(cache.router, tags=["Cache"]) -api_router.include_router(web_search.router, tags=["Web Search"]) -api_router.include_router(geocoding.router, tags=["Geocoding"]) -api_router.include_router(leakage.router, prefix="/leakage", tags=["Leakage"]) api_router.include_router( - burst_detection.router, prefix="/burst-detection", tags=["Burst Detection"] -) -api_router.include_router( - burst_location.router, prefix="/burst-location", tags=["Burst Location"] + project.router, + tags=["Project"], + dependencies=[webgis_access], ) -# TimescaleDB Data Access -api_router.include_router(ts_realtime.router, tags=["TimescaleDB - Realtime"]) -api_router.include_router(ts_scheme.router, tags=["TimescaleDB - Scheme"]) -api_router.include_router(ts_scada.router, tags=["TimescaleDB - SCADA"]) -api_router.include_router(ts_composite.router, tags=["TimescaleDB - Composite"]) +# WebGIS data +for endpoint_router, tag in ( + (general.router, "Network General"), + (junctions.router, "Junctions"), + (reservoirs.router, "Reservoirs"), + (tanks.router, "Tanks"), + (pipes.router, "Pipes"), + (pumps.router, "Pumps"), + (valves.router, "Valves"), + (tags.router, "Tags"), + (demands.router, "Demands"), + (geometry.router, "Geometry & Coordinates"), + (regions.router, "Regions & DMAs"), + (curves.router, "Curves"), + (patterns.router, "Patterns"), + (controls.router, "Controls & Rules"), + (options.router, "Options"), + (quality.router, "Quality"), + (visuals.router, "Visuals"), +): + api_router.include_router( + endpoint_router, + tags=[tag], + dependencies=[webgis_access], + ) -# Project Data (PostgreSQL) -api_router.include_router(project_data.router, tags=["Project Data"]) +# Simulation and analysis +api_router.include_router( + simulation.router, + tags=["Simulation Control"], + dependencies=[simulation_run_access], +) +api_router.include_router(scada.router, dependencies=[scada_access]) +api_router.include_router( + sensor_placement.router, + tags=["Sensor Placement"], + dependencies=[optimization_run_access], +) +api_router.include_router( + snapshots.router, + tags=["Snapshots"], + dependencies=[simulation_access], +) +api_router.include_router( + users.router, + tags=["Users"], + dependencies=[webgis_view_access], +) +api_router.include_router( + schemes.router, + tags=["Schemes"], + dependencies=[simulation_access], +) +api_router.include_router( + misc.router, + tags=["Misc"], + dependencies=[webgis_view_access], +) +api_router.include_router( + risk.router, + tags=["Risk"], + dependencies=[risk_run_access], +) +api_router.include_router( + cache.router, + tags=["Cache"], + dependencies=[simulation_run_access], +) +api_router.include_router( + web_search.router, + tags=["Web Search"], + dependencies=[webgis_view_access], +) +api_router.include_router( + geocoding.router, + tags=["Geocoding"], + dependencies=[webgis_view_access], +) +api_router.include_router( + leakage.router, + prefix="/leakage", + tags=["Leakage"], + dependencies=[burst_run_access], +) +api_router.include_router( + burst_detection.router, + prefix="/burst-detection", + tags=["Burst Detection"], + dependencies=[burst_run_access], +) +api_router.include_router( + burst_location.router, + prefix="/burst-location", + tags=["Burst Location"], + dependencies=[burst_run_access], +) -# Extension -api_router.include_router(extension.router, tags=["Extension"]) +# TimescaleDB data +for endpoint_router, tag in ( + (ts_realtime.router, "TimescaleDB - Realtime"), + (ts_scheme.router, "TimescaleDB - Scheme"), +): + api_router.include_router( + endpoint_router, + tags=[tag], + dependencies=[simulation_access], + ) + +for endpoint_router, tag in ( + (ts_scada.router, "TimescaleDB - SCADA"), + (ts_composite.router, "TimescaleDB - Composite"), +): + api_router.include_router( + endpoint_router, + tags=[tag], + dependencies=[scada_access], + ) + +api_router.include_router( + project_data.router, + tags=["Project Data"], + dependencies=[webgis_view_access], +) +api_router.include_router( + extension.router, + tags=["Extension"], + dependencies=[webgis_access], +) diff --git a/app/auth/permissions.py b/app/auth/permissions.py new file mode 100644 index 0000000..380ea5e --- /dev/null +++ b/app/auth/permissions.py @@ -0,0 +1,162 @@ +from collections.abc import Awaitable, Callable +from typing import Any + +from fastapi import Depends, HTTPException, Request, status + +from app.auth.project_dependencies import ProjectContext, get_project_context + +WEBGIS_VIEW = "webgis.view" +WEBGIS_EDIT = "webgis.edit" +SCADA_VIEW = "scada.view" +SCADA_CLEAN = "scada.clean" +SIMULATION_VIEW = "simulation.view" +SIMULATION_RUN = "simulation.run" +BURST_VIEW = "burst.view" +BURST_RUN = "burst.run" +RISK_VIEW = "risk.view" +RISK_RUN = "risk.run" +OPTIMIZATION_VIEW = "optimization.view" +OPTIMIZATION_RUN = "optimization.run" +MODEL_IMPORT = "model.import" +AUDIT_VIEW = "audit.view" +ENVIRONMENT_MANAGE = "environment.manage" +MEMBERSHIP_MANAGE = "membership.manage" + +PROJECT_MEMBER_PERMISSIONS = frozenset( + { + WEBGIS_VIEW, + WEBGIS_EDIT, + SCADA_VIEW, + SCADA_CLEAN, + SIMULATION_VIEW, + SIMULATION_RUN, + BURST_VIEW, + BURST_RUN, + RISK_VIEW, + RISK_RUN, + OPTIMIZATION_VIEW, + OPTIMIZATION_RUN, + } +) + +PROJECT_VIEWER_PERMISSIONS = frozenset( + { + WEBGIS_VIEW, + SCADA_VIEW, + SIMULATION_VIEW, + } +) + +SYSTEM_ADMIN_PERMISSIONS = frozenset( + { + MODEL_IMPORT, + AUDIT_VIEW, + ENVIRONMENT_MANAGE, + MEMBERSHIP_MANAGE, + } +) + +PROJECT_ROLE_PERMISSIONS: dict[str, frozenset[str]] = { + "member": PROJECT_MEMBER_PERMISSIONS, + "viewer": PROJECT_VIEWER_PERMISSIONS, +} + + +def resolve_permissions( + *, + project_role: str | None, + system_role: str, + is_superuser: bool, +) -> frozenset[str]: + permissions = set(PROJECT_ROLE_PERMISSIONS.get(project_role or "", frozenset())) + if is_superuser or system_role == "admin": + permissions.update(SYSTEM_ADMIN_PERMISSIONS) + return frozenset(permissions) + + +def permissions_for_context(ctx: ProjectContext) -> frozenset[str]: + return resolve_permissions( + project_role=ctx.project_role, + system_role=ctx.system_role, + is_superuser=ctx.is_superuser, + ) + + +def _permission_denied(permission: str) -> HTTPException: + return HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "code": "permission_denied", + "permission": permission, + }, + ) + + +async def _enforce_project_scope(request: Request, ctx: ProjectContext) -> None: + requested_network = ( + request.path_params.get("network") + or request.query_params.get("network") + ) + if not requested_network: + content_type = request.headers.get("content-type", "") + if content_type.startswith("application/json"): + try: + payload = await request.json() + except (ValueError, RuntimeError): + payload = None + if isinstance(payload, dict): + requested_network = payload.get("network") + + if requested_network and str(requested_network) != ctx.project_code: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "code": "project_scope_denied", + "project_id": str(ctx.project_id), + }, + ) + + +def require_permission( + permission: str, +) -> Callable[..., Awaitable[ProjectContext]]: + async def dependency( + request: Request, + ctx: ProjectContext = Depends(get_project_context), + ) -> ProjectContext: + if permission not in permissions_for_context(ctx): + raise _permission_denied(permission) + await _enforce_project_scope(request, ctx) + return ctx + + return dependency + + +def require_method_permission( + *, + read_permission: str, + write_permission: str, +) -> Callable[..., Awaitable[ProjectContext]]: + async def dependency( + request: Request, + ctx: ProjectContext = Depends(get_project_context), + ) -> ProjectContext: + permission = ( + read_permission + if request.method.upper() in {"GET", "HEAD", "OPTIONS"} + else write_permission + ) + if permission not in permissions_for_context(ctx): + raise _permission_denied(permission) + await _enforce_project_scope(request, ctx) + return ctx + + return dependency + + +def has_permission(user: Any, project_role: str | None, permission: str) -> bool: + return permission in resolve_permissions( + project_role=project_role, + system_role=str(getattr(user, "role", "user")), + is_superuser=bool(getattr(user, "is_superuser", False)), + ) diff --git a/app/auth/project_dependencies.py b/app/auth/project_dependencies.py index 0bcaa43..362a186 100644 --- a/app/auth/project_dependencies.py +++ b/app/auth/project_dependencies.py @@ -1,18 +1,21 @@ +import logging +from collections.abc import AsyncGenerator from dataclasses import dataclass -from typing import AsyncGenerator from uuid import UUID -import logging from fastapi import Depends, Header, HTTPException, status from psycopg import AsyncConnection from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession -from app.auth.keycloak_dependencies import get_current_keycloak_sub +from app.auth.metadata_dependencies import get_current_metadata_user from app.core.config import settings from app.infra.db.dynamic_manager import project_connection_manager from app.infra.db.metadb.database import get_metadata_session -from app.infra.db.metadb.repositories.metadata_repository import MetadataRepository +from app.infra.db.metadb.repositories.metadata_repository import ( + MetadataRepository, + ProjectDbRouting, +) DB_ROLE_BIZ_DATA = "biz_data" DB_ROLE_IOT_DATA = "iot_data" @@ -28,6 +31,8 @@ class ProjectContext: project_code: str user_id: UUID project_role: str + system_role: str = "user" + is_superuser: bool = False async def get_metadata_repository( @@ -36,10 +41,10 @@ async def get_metadata_repository( return MetadataRepository(session) -async def get_project_context( - x_project_id: str = Header(..., alias="X-Project-Id"), - keycloak_sub: UUID = Depends(get_current_keycloak_sub), - metadata_repo: MetadataRepository = Depends(get_metadata_repository), +async def resolve_project_context( + x_project_id: str, + current_user, + metadata_repo: MetadataRepository, ) -> ProjectContext: try: project_uuid = UUID(x_project_id) @@ -59,17 +64,9 @@ async def get_project_context( status_code=status.HTTP_403_FORBIDDEN, detail="Project is not active" ) - user = await metadata_repo.get_user_by_keycloak_id(keycloak_sub) - if not user: - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, detail="User not registered" - ) - if not user.is_active: - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, detail="Inactive user" - ) - - membership_role = await metadata_repo.get_membership_role(project_uuid, user.id) + membership_role = await metadata_repo.get_membership_role( + project_uuid, current_user.id + ) if not membership_role: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail="No access to project" @@ -87,38 +84,65 @@ async def get_project_context( return ProjectContext( project_id=project.id, project_code=project.code, - user_id=user.id, + user_id=current_user.id, project_role=membership_role, + system_role=current_user.role, + is_superuser=current_user.is_superuser, ) +async def get_project_context( + x_project_id: str = Header(..., alias="X-Project-Id"), + current_user=Depends(get_current_metadata_user), + metadata_repo: MetadataRepository = Depends(get_metadata_repository), +) -> ProjectContext: + return await resolve_project_context(x_project_id, current_user, metadata_repo) + + +async def _get_project_routing( + metadata_repo: MetadataRepository, + project_id: UUID, + db_role: str, + expected_db_type: str, + database_label: str, +) -> ProjectDbRouting: + try: + routing = await metadata_repo.get_project_db_routing(project_id, db_role) + except ValueError as exc: + logger.error( + "Invalid project %s routing DSN configuration", + database_label, + exc_info=True, + ) + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail=f"Project {database_label} routing DSN is invalid: {exc}", + ) from exc + + if not routing: + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail=f"Project {database_label} not configured", + ) + if routing.db_type != expected_db_type: + raise HTTPException( + status_code=status.HTTP_503_SERVICE_UNAVAILABLE, + detail=f"Project {database_label} type mismatch", + ) + return routing + + async def get_project_pg_session( ctx: ProjectContext = Depends(get_project_context), metadata_repo: MetadataRepository = Depends(get_metadata_repository), ) -> AsyncGenerator[AsyncSession, None]: - try: - routing = await metadata_repo.get_project_db_routing( - ctx.project_id, DB_ROLE_BIZ_DATA - ) - except ValueError as exc: - logger.error( - "Invalid project PostgreSQL routing DSN configuration", - exc_info=True, - ) - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail=f"Project PostgreSQL routing DSN is invalid: {exc}", - ) from exc - if not routing: - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail="Project PostgreSQL not configured", - ) - if routing.db_type != DB_TYPE_POSTGRES: - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail="Project PostgreSQL type mismatch", - ) + routing = await _get_project_routing( + metadata_repo, + ctx.project_id, + DB_ROLE_BIZ_DATA, + DB_TYPE_POSTGRES, + "PostgreSQL", + ) pool_min_size = routing.pool_min_size or settings.PROJECT_PG_POOL_SIZE pool_max_size = routing.pool_max_size or settings.PROJECT_PG_POOL_SIZE @@ -137,29 +161,13 @@ async def get_project_pg_connection( ctx: ProjectContext = Depends(get_project_context), metadata_repo: MetadataRepository = Depends(get_metadata_repository), ) -> AsyncGenerator[AsyncConnection, None]: - try: - routing = await metadata_repo.get_project_db_routing( - ctx.project_id, DB_ROLE_BIZ_DATA - ) - except ValueError as exc: - logger.error( - "Invalid project PostgreSQL routing DSN configuration", - exc_info=True, - ) - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail=f"Project PostgreSQL routing DSN is invalid: {exc}", - ) from exc - if not routing: - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail="Project PostgreSQL not configured", - ) - if routing.db_type != DB_TYPE_POSTGRES: - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail="Project PostgreSQL type mismatch", - ) + routing = await _get_project_routing( + metadata_repo, + ctx.project_id, + DB_ROLE_BIZ_DATA, + DB_TYPE_POSTGRES, + "PostgreSQL", + ) pool_min_size = routing.pool_min_size or settings.PROJECT_PG_POOL_SIZE pool_max_size = routing.pool_max_size or settings.PROJECT_PG_POOL_SIZE @@ -178,29 +186,13 @@ async def get_project_timescale_connection( ctx: ProjectContext = Depends(get_project_context), metadata_repo: MetadataRepository = Depends(get_metadata_repository), ) -> AsyncGenerator[AsyncConnection, None]: - try: - routing = await metadata_repo.get_project_db_routing( - ctx.project_id, DB_ROLE_IOT_DATA - ) - except ValueError as exc: - logger.error( - "Invalid project TimescaleDB routing DSN configuration", - exc_info=True, - ) - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail=f"Project TimescaleDB routing DSN is invalid: {exc}", - ) from exc - if not routing: - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail="Project TimescaleDB not configured", - ) - if routing.db_type != DB_TYPE_TIMESCALE: - raise HTTPException( - status_code=status.HTTP_503_SERVICE_UNAVAILABLE, - detail="Project TimescaleDB type mismatch", - ) + routing = await _get_project_routing( + metadata_repo, + ctx.project_id, + DB_ROLE_IOT_DATA, + DB_TYPE_TIMESCALE, + "TimescaleDB", + ) pool_min_size = routing.pool_min_size or settings.PROJECT_TS_POOL_MIN_SIZE pool_max_size = routing.pool_max_size or settings.PROJECT_TS_POOL_MAX_SIZE diff --git a/app/domain/schemas/access.py b/app/domain/schemas/access.py new file mode 100644 index 0000000..c34f22c --- /dev/null +++ b/app/domain/schemas/access.py @@ -0,0 +1,13 @@ +from uuid import UUID + +from pydantic import BaseModel + + +class AccessContextResponse(BaseModel): + user_id: UUID + username: str + system_role: str + is_system_admin: bool + project_id: UUID | None = None + project_role: str | None = None + permissions: list[str] diff --git a/app/domain/schemas/admin_metadata.py b/app/domain/schemas/admin_metadata.py index 98d2d57..6ebfc25 100644 --- a/app/domain/schemas/admin_metadata.py +++ b/app/domain/schemas/admin_metadata.py @@ -5,8 +5,8 @@ from uuid import UUID from pydantic import BaseModel, ConfigDict, Field, model_validator -BusinessRole = Literal["admin", "user", "operator", "viewer"] -ProjectRole = Literal["owner", "admin", "member", "viewer"] +BusinessRole = Literal["admin", "user"] +ProjectRole = Literal["member", "viewer"] ProjectStatus = Literal["active", "inactive", "archived"] ProjectDbRole = Literal["biz_data", "iot_data"] diff --git a/app/infra/db/metadb/repositories/metadata_repository.py b/app/infra/db/metadb/repositories/metadata_repository.py index f35ef71..b620ba5 100644 --- a/app/infra/db/metadb/repositories/metadata_repository.py +++ b/app/infra/db/metadb/repositories/metadata_repository.py @@ -211,6 +211,7 @@ class MetadataRepository: gs_workspace: str, map_extent: dict | None, status: str, + creator_user_id: UUID | None = None, ) -> models.Project: project = models.Project( id=uuid4(), @@ -224,6 +225,15 @@ class MetadataRepository: updated_at=_utcnow(), ) self.session.add(project) + if creator_user_id is not None: + self.session.add( + models.UserProjectMembership( + id=uuid4(), + user_id=creator_user_id, + project_id=project.id, + project_role="member", + ) + ) await self.session.commit() await self.session.refresh(project) return project @@ -483,7 +493,7 @@ class MetadataRepository: gs_workspace=project.gs_workspace, map_extent=project.map_extent, status=project.status, - project_role="owner", + project_role="member", ) for project in result.scalars().all() ] diff --git a/resources/sql/004_metadata_auth_management.sql b/resources/sql/004_metadata_auth_management.sql index e3d682b..3e58a88 100644 --- a/resources/sql/004_metadata_auth_management.sql +++ b/resources/sql/004_metadata_auth_management.sql @@ -42,6 +42,12 @@ ALTER TABLE users ALTER TABLE users ALTER COLUMN role SET DEFAULT 'user'; +ALTER TABLE users + DROP CONSTRAINT IF EXISTS users_role_check; +ALTER TABLE users + ADD CONSTRAINT users_role_check + CHECK (role IN ('admin', 'user')); + CREATE UNIQUE INDEX IF NOT EXISTS idx_users_keycloak_id ON users(keycloak_id); CREATE INDEX IF NOT EXISTS idx_users_role ON users(role); CREATE INDEX IF NOT EXISTS idx_users_is_active ON users(is_active); @@ -52,7 +58,9 @@ CREATE TABLE IF NOT EXISTS user_project_membership ( project_id UUID NOT NULL, project_role VARCHAR(20) DEFAULT 'viewer' NOT NULL, CONSTRAINT user_project_membership_role_check - CHECK (project_role IN ('owner', 'admin', 'member', 'viewer')), + CHECK ( + project_role IN ('member', 'viewer') + ), CONSTRAINT user_project_membership_unique UNIQUE (user_id, project_id) ); diff --git a/resources/sql/006_metadata_rbac_roles.sql b/resources/sql/006_metadata_rbac_roles.sql new file mode 100644 index 0000000..947e256 --- /dev/null +++ b/resources/sql/006_metadata_rbac_roles.sql @@ -0,0 +1,32 @@ +-- Normalize existing roles to the Web authorization model. +-- This migration is intentionally re-runnable. + +ALTER TABLE users + DROP CONSTRAINT IF EXISTS users_role_check; + +UPDATE users +SET role = 'user' +WHERE role NOT IN ('admin', 'user'); + +ALTER TABLE users + ADD CONSTRAINT users_role_check + CHECK (role IN ('admin', 'user')); + +ALTER TABLE user_project_membership + DROP CONSTRAINT IF EXISTS user_project_membership_role_check; + +UPDATE user_project_membership +SET project_role = CASE + WHEN project_role IN ( + 'owner', + 'admin', + 'modeler', + 'dispatcher' + ) THEN 'member' + ELSE 'viewer' +END +WHERE project_role NOT IN ('member', 'viewer'); + +ALTER TABLE user_project_membership + ADD CONSTRAINT user_project_membership_role_check + CHECK (project_role IN ('member', 'viewer')); diff --git a/tests/api/test_access_endpoints.py b/tests/api/test_access_endpoints.py new file mode 100644 index 0000000..71c9604 --- /dev/null +++ b/tests/api/test_access_endpoints.py @@ -0,0 +1,77 @@ +from types import SimpleNamespace +from uuid import uuid4 + +from fastapi.testclient import TestClient + +from app.api.v1.endpoints import access as access_endpoint +from app.auth.metadata_dependencies import ( + get_current_metadata_user, + get_metadata_repository, +) +from tests.conftest import build_test_app + + +def _user(**overrides): + data = { + "id": uuid4(), + "username": "alice", + "role": "user", + "is_superuser": False, + } + data.update(overrides) + return SimpleNamespace(**data) + + +def _build_client(user, repo) -> TestClient: + app = build_test_app(access_endpoint.router, "/api/v1") + app.dependency_overrides[get_current_metadata_user] = lambda: user + app.dependency_overrides[get_metadata_repository] = lambda: repo + return TestClient(app) + + +def test_access_context_returns_global_admin_permissions_without_project(): + user = _user(role="admin") + repo = SimpleNamespace() + client = _build_client(user, repo) + + response = client.get("/api/v1/access/context") + + assert response.status_code == 200 + payload = response.json() + assert payload["is_system_admin"] is True + assert payload["project_id"] is None + assert "environment.manage" in payload["permissions"] + assert "webgis.view" not in payload["permissions"] + + +def test_access_context_returns_project_member_permissions(): + project_id = uuid4() + user = _user() + + async def get_project_by_id(value): + assert value == project_id + return SimpleNamespace(id=project_id, code="demo", status="active") + + async def get_membership_role(value, user_id): + assert value == project_id + assert user_id == user.id + return "member" + + repo = SimpleNamespace( + get_project_by_id=get_project_by_id, + get_membership_role=get_membership_role, + ) + client = _build_client(user, repo) + + response = client.get( + "/api/v1/access/context", + headers={"X-Project-Id": str(project_id)}, + ) + + assert response.status_code == 200 + payload = response.json() + assert payload["project_id"] == str(project_id) + assert payload["project_role"] == "member" + assert "scada.clean" in payload["permissions"] + assert "optimization.run" in payload["permissions"] + assert "model.import" not in payload["permissions"] diff --git a/tests/api/test_admin_metadata_endpoints.py b/tests/api/test_admin_metadata_endpoints.py index 579ede4..70161dc 100644 --- a/tests/api/test_admin_metadata_endpoints.py +++ b/tests/api/test_admin_metadata_endpoints.py @@ -138,7 +138,7 @@ async def test_batch_sync_metadata_users_returns_per_user_results(monkeypatch): keycloak_id=users[1].keycloak_id, username="bob", email="bob@example.com", - role="viewer", + role="user", is_active=True, ), ] @@ -156,7 +156,7 @@ async def test_batch_sync_metadata_users_returns_per_user_results(monkeypatch): @pytest.mark.anyio async def test_update_metadata_user_updates_role_and_active_status(monkeypatch): user_id = uuid4() - updated = _user(id=user_id, role="operator", is_active=False) + updated = _user(id=user_id, role="user", is_active=False) repo = SimpleNamespace( session=object(), update_user_admin=AsyncMock(return_value=updated), @@ -165,7 +165,7 @@ async def test_update_metadata_user_updates_role_and_active_status(monkeypatch): response = await admin_metadata.update_metadata_user( MetadataUserUpdateRequest( - role="operator", + role="user", is_active=False, ), user_id=user_id, @@ -175,9 +175,9 @@ async def test_update_metadata_user_updates_role_and_active_status(monkeypatch): repo.update_user_admin.assert_awaited_once_with( user_id, - updates={"role": "operator", "is_active": False}, + updates={"role": "user", "is_active": False}, ) - assert response.role == "operator" + assert response.role == "user" admin_metadata.log_audit_event.assert_awaited_once() @@ -192,7 +192,7 @@ async def test_update_metadata_user_rejects_self_update(monkeypatch): with pytest.raises(HTTPException) as exc: await admin_metadata.update_metadata_user( - MetadataUserUpdateRequest(role="viewer"), + MetadataUserUpdateRequest(role="user"), user_id=current_user.id, current_user=current_user, metadata_repo=repo, @@ -221,6 +221,7 @@ async def test_create_project_audits_metadata_admin_change(monkeypatch): create_project=AsyncMock(return_value=project), ) monkeypatch.setattr(admin_metadata, "log_audit_event", AsyncMock()) + current_user = _user(role="admin", is_superuser=True) response = await admin_metadata.create_admin_project( AdminProjectCreateRequest( @@ -231,12 +232,16 @@ async def test_create_project_audits_metadata_admin_change(monkeypatch): map_extent={"bbox": [1, 2, 3, 4]}, status="active", ), - current_user=_user(role="admin", is_superuser=True), + current_user=current_user, metadata_repo=repo, ) assert response.project_id == project.id repo.create_project.assert_awaited_once() + assert ( + repo.create_project.await_args.kwargs["creator_user_id"] + == current_user.id + ) admin_metadata.log_audit_event.assert_awaited_once() @@ -482,7 +487,7 @@ async def test_update_project_member_role_audits_change(monkeypatch): membership = _membership( user_id=user_id, project_id=project_id, - project_role="admin", + project_role="member", ) repo = SimpleNamespace( session=object(), @@ -492,16 +497,16 @@ async def test_update_project_member_role_audits_change(monkeypatch): monkeypatch.setattr(admin_metadata, "log_audit_event", AsyncMock()) response = await admin_metadata.update_project_member( - ProjectMemberUpdateRequest(project_role="admin"), + ProjectMemberUpdateRequest(project_role="member"), project_id=project_id, user_id=user_id, current_user=_user(role="admin", is_superuser=True), metadata_repo=repo, ) - assert response.project_role == "admin" + assert response.project_role == "member" repo.update_project_member_role.assert_awaited_once_with( - project_id, user_id, "admin" + project_id, user_id, "member" ) admin_metadata.log_audit_event.assert_awaited_once() @@ -519,7 +524,7 @@ async def test_update_project_member_rejects_self_membership_change(monkeypatch) with pytest.raises(HTTPException) as exc: await admin_metadata.update_project_member( - ProjectMemberUpdateRequest(project_role="admin"), + ProjectMemberUpdateRequest(project_role="member"), project_id=project_id, user_id=current_user.id, current_user=current_user, diff --git a/tests/api/test_agent_auth_endpoints.py b/tests/api/test_agent_auth_endpoints.py index 38fbdb2..214fd93 100644 --- a/tests/api/test_agent_auth_endpoints.py +++ b/tests/api/test_agent_auth_endpoints.py @@ -30,7 +30,7 @@ def test_agent_auth_context_returns_metadata_user_and_project_context(): project_id=project_id, project_code="fengyang", user_id=user_id, - project_role="editor", + project_role="member", ), current_user=SimpleNamespace( id=user_id, @@ -52,7 +52,21 @@ def test_agent_auth_context_returns_metadata_user_and_project_context(): "is_superuser": False, "project_id": str(project_id), "network": "fengyang", - "project_role": "editor", + "project_role": "member", + "permissions": [ + "burst.run", + "burst.view", + "optimization.run", + "optimization.view", + "risk.run", + "risk.view", + "scada.clean", + "scada.view", + "simulation.run", + "simulation.view", + "webgis.edit", + "webgis.view", + ], "token_expires_at": "2026-06-11T13:10:00+00:00", } diff --git a/tests/api/test_audit_endpoints.py b/tests/api/test_audit_endpoints.py index d043800..5a7c598 100644 --- a/tests/api/test_audit_endpoints.py +++ b/tests/api/test_audit_endpoints.py @@ -1,4 +1,5 @@ from unittest.mock import AsyncMock +from uuid import uuid4 from fastapi.testclient import TestClient @@ -10,7 +11,12 @@ from app.auth.metadata_dependencies import ( from tests.conftest import build_test_app, make_audit_log -def _build_client(repo, *, metadata_admin=None, metadata_user=None) -> TestClient: +def _build_client( + repo, + *, + metadata_admin=None, + metadata_user=None, +) -> TestClient: app = build_test_app(audit_endpoint.router, "/audit") app.dependency_overrides[audit_endpoint.get_audit_repository] = lambda: repo if metadata_admin is not None: diff --git a/tests/api/test_meta_endpoints.py b/tests/api/test_meta_endpoints.py index 899c49e..6934612 100644 --- a/tests/api/test_meta_endpoints.py +++ b/tests/api/test_meta_endpoints.py @@ -54,7 +54,7 @@ def test_meta_project_returns_map_extent(monkeypatch): app = build_test_app(module.router, "/api/v1") app.dependency_overrides[module.get_project_context] = lambda: SimpleNamespace( project_id=project_id, - project_role="editor", + project_role="member", ) app.dependency_overrides[module.get_metadata_repository] = lambda: repo client = TestClient(app) diff --git a/tests/api/test_model_import_endpoints.py b/tests/api/test_model_import_endpoints.py new file mode 100644 index 0000000..52639ae --- /dev/null +++ b/tests/api/test_model_import_endpoints.py @@ -0,0 +1,100 @@ +from types import SimpleNamespace +from unittest.mock import AsyncMock +from uuid import uuid4 + +from fastapi import HTTPException +from fastapi.testclient import TestClient + +from app.api.v1.endpoints import model_import +from app.auth.metadata_dependencies import ( + get_current_metadata_admin, + get_metadata_repository, +) +from tests.conftest import build_test_app + + +VALID_INP = b"[TITLE]\nDesktop model\n[JUNCTIONS]\n;ID Elev Demand\n" + + +def _client(*, admin=None, repo=None) -> TestClient: + app = build_test_app(model_import.router, "/api/v1") + if admin is not None: + app.dependency_overrides[get_current_metadata_admin] = lambda: admin + if repo is not None: + app.dependency_overrides[get_metadata_repository] = lambda: repo + return TestClient(app) + + +def test_system_admin_can_import_model_without_project_membership( + monkeypatch, +): + project_id = uuid4() + project = SimpleNamespace(id=project_id, code="demo", status="active") + repo = SimpleNamespace( + session=object(), + get_project_by_id=AsyncMock(return_value=project), + ) + admin = SimpleNamespace(id=uuid4(), role="admin", is_superuser=False) + monkeypatch.setattr( + model_import, + "_run_uploaded_inp", + AsyncMock(return_value="imported"), + ) + monkeypatch.setattr(model_import, "log_audit_event", AsyncMock()) + client = _client(admin=admin, repo=repo) + + response = client.post( + f"/api/v1/admin/projects/{project_id}/model/import", + files={"file": ("desktop-model.inp", VALID_INP)}, + ) + + assert response.status_code == 200 + assert response.json()["project_id"] == str(project_id) + assert response.json()["result"] == "imported" + repo.get_project_by_id.assert_awaited_once_with(project_id) + model_import.log_audit_event.assert_awaited_once() + + +def test_non_admin_is_denied_model_import(): + def deny_admin(): + raise HTTPException(status_code=403, detail="Admin access required") + + app = build_test_app(model_import.router, "/api/v1") + app.dependency_overrides[get_current_metadata_admin] = deny_admin + client = TestClient(app) + + response = client.post( + f"/api/v1/admin/projects/{uuid4()}/model/import", + files={"file": ("desktop-model.inp", VALID_INP)}, + ) + + assert response.status_code == 403 + assert response.json()["detail"] == "Admin access required" + + +def test_model_import_rejects_non_inp_file(monkeypatch): + project_id = uuid4() + repo = SimpleNamespace( + session=object(), + get_project_by_id=AsyncMock( + return_value=SimpleNamespace( + id=project_id, + code="demo", + status="active", + ) + ), + ) + monkeypatch.setattr(model_import, "log_audit_event", AsyncMock()) + client = _client( + admin=SimpleNamespace(id=uuid4(), role="admin", is_superuser=False), + repo=repo, + ) + + response = client.post( + f"/api/v1/admin/projects/{project_id}/model/import", + files={"file": ("desktop-model.txt", VALID_INP)}, + ) + + assert response.status_code == 400 + assert response.json()["detail"] == "Only .inp model files are accepted" + model_import.log_audit_event.assert_not_awaited() diff --git a/tests/api/test_sensor_placement_endpoints.py b/tests/api/test_sensor_placement_endpoints.py index 34756da..a743693 100644 --- a/tests/api/test_sensor_placement_endpoints.py +++ b/tests/api/test_sensor_placement_endpoints.py @@ -2,6 +2,7 @@ from datetime import datetime, timezone from io import BytesIO from types import SimpleNamespace +import pytest from fastapi.testclient import TestClient from tests.conftest import build_test_app, install_stub, load_module_from_path @@ -254,21 +255,25 @@ def test_optimize_rejects_viewer_project_role(monkeypatch): assert response.status_code == 403 -def test_project_owner_and_admin_can_optimize(monkeypatch): +@pytest.mark.parametrize( + "project_role", + ["owner", "admin", "modeler", "dispatcher", "auditor"], +) +def test_legacy_project_roles_cannot_optimize(monkeypatch, project_role): module = _load_module(monkeypatch) - for project_role in ("owner", "admin"): - response = _client(module, project_role=project_role).post( - "/api/v1/sensor-placement-schemes/optimize", - json={ - "network": "tjwater", - "scheme_name": f"{project_role}方案", - "sensor_type": "pressure", - "method": "kmeans", - "sensor_count": 2, - "min_diameter": 300, - }, - ) - assert response.status_code == 200 + response = _client(module, project_role=project_role).post( + "/api/v1/sensor-placement-schemes/optimize", + json={ + "network": "tjwater", + "scheme_name": f"{project_role}方案", + "sensor_type": "pressure", + "method": "kmeans", + "sensor_count": 2, + "min_diameter": 300, + }, + ) + + assert response.status_code == 403 def test_optimize_maps_running_project_job_to_409(monkeypatch): diff --git a/tests/api/test_simulation_endpoints.py b/tests/api/test_simulation_endpoints.py index c5c53b3..cee6eeb 100644 --- a/tests/api/test_simulation_endpoints.py +++ b/tests/api/test_simulation_endpoints.py @@ -1,4 +1,3 @@ -from pathlib import Path from datetime import datetime, timezone from fastapi.testclient import TestClient @@ -199,26 +198,6 @@ def test_project_management_maps_named_arguments(monkeypatch): } -def test_network_update_surfaces_service_error(monkeypatch, tmp_path): - module = _load_simulation_module(monkeypatch) - monkeypatch.chdir(tmp_path) - - def boom(_path): - raise RuntimeError("write failed") - - monkeypatch.setattr(module, "network_update", boom) - client = TestClient(build_test_app(module.router, "/api/v1")) - - response = client.post( - "/api/v1/network_update/", - files={"file": ("update.txt", b"payload")}, - ) - - assert response.status_code == 500 - assert "数据库操作失败: write failed" in response.json()["detail"] - assert list(Path(tmp_path).glob("network_update_*")) - - def test_run_simulation_manually_by_date_uses_utc_aware_timestamps(monkeypatch): module = _load_simulation_module(monkeypatch) captured_calls = [] diff --git a/tests/auth/test_permissions.py b/tests/auth/test_permissions.py new file mode 100644 index 0000000..13d96f0 --- /dev/null +++ b/tests/auth/test_permissions.py @@ -0,0 +1,159 @@ +from uuid import uuid4 + +import pytest +from fastapi import HTTPException + +from app.auth.permissions import ( + AUDIT_VIEW, + ENVIRONMENT_MANAGE, + MODEL_IMPORT, + OPTIMIZATION_RUN, + SCADA_CLEAN, + SIMULATION_RUN, + SIMULATION_VIEW, + WEBGIS_EDIT, + WEBGIS_VIEW, + permissions_for_context, + require_method_permission, + require_permission, + resolve_permissions, +) +from app.auth.project_dependencies import ProjectContext + + +@pytest.fixture +def anyio_backend(): + return "asyncio" + + +def _context(project_role: str) -> ProjectContext: + return ProjectContext( + project_id=uuid4(), + project_code="demo", + user_id=uuid4(), + project_role=project_role, + ) + + +def test_project_role_permission_matrix(): + member = resolve_permissions( + project_role="member", + system_role="user", + is_superuser=False, + ) + viewer = resolve_permissions( + project_role="viewer", + system_role="user", + is_superuser=False, + ) + + assert WEBGIS_EDIT in member + assert SCADA_CLEAN in member + assert SIMULATION_RUN in member + assert OPTIMIZATION_RUN in member + assert MODEL_IMPORT not in member + assert WEBGIS_VIEW in viewer + assert SIMULATION_VIEW in viewer + assert WEBGIS_EDIT not in viewer + assert SCADA_CLEAN not in viewer + assert SIMULATION_RUN not in viewer + + +def test_system_admin_permissions_do_not_grant_project_business_access(): + permissions = resolve_permissions( + project_role=None, + system_role="admin", + is_superuser=False, + ) + + assert ENVIRONMENT_MANAGE in permissions + assert AUDIT_VIEW in permissions + assert MODEL_IMPORT in permissions + assert WEBGIS_VIEW not in permissions + + +@pytest.mark.anyio +async def test_permission_dependency_returns_context_when_allowed(): + ctx = _context("member") + dependency = require_permission(WEBGIS_EDIT) + + request = type( + "Request", + (), + { + "path_params": {}, + "query_params": {}, + "headers": {}, + }, + )() + + assert await dependency(request, ctx) is ctx + + +@pytest.mark.anyio +async def test_permission_dependency_returns_structured_403_when_denied(): + ctx = _context("viewer") + dependency = require_permission(WEBGIS_EDIT) + + with pytest.raises(HTTPException) as exc: + await dependency(None, ctx) + + assert exc.value.status_code == 403 + assert exc.value.detail == { + "code": "permission_denied", + "permission": WEBGIS_EDIT, + } + + +@pytest.mark.anyio +async def test_permission_dependency_rejects_cross_project_network(): + ctx = _context("member") + dependency = require_permission(WEBGIS_VIEW) + request = type( + "Request", + (), + { + "path_params": {}, + "query_params": {"network": "other-project"}, + "headers": {}, + }, + )() + + with pytest.raises(HTTPException) as exc: + await dependency(request, ctx) + + assert exc.value.status_code == 403 + assert exc.value.detail["code"] == "project_scope_denied" + + +def test_member_keeps_full_web_business_access(): + permissions = permissions_for_context(_context("member")) + + assert SCADA_CLEAN in permissions + assert SIMULATION_RUN in permissions + + +@pytest.mark.anyio +async def test_viewer_can_read_but_cannot_write_or_run(): + ctx = _context("viewer") + dependency = require_method_permission( + read_permission=SIMULATION_VIEW, + write_permission=SIMULATION_RUN, + ) + read_request = type( + "Request", + (), + {"method": "GET", "path_params": {}, "query_params": {}, "headers": {}}, + )() + write_request = type( + "Request", + (), + {"method": "POST", "path_params": {}, "query_params": {}, "headers": {}}, + )() + + assert await dependency(read_request, ctx) is ctx + with pytest.raises(HTTPException) as exc: + await dependency(write_request, ctx) + + assert exc.value.status_code == 403 + assert exc.value.detail["permission"] == SIMULATION_RUN diff --git a/tests/auth/test_rbac_migration.py b/tests/auth/test_rbac_migration.py new file mode 100644 index 0000000..8111be4 --- /dev/null +++ b/tests/auth/test_rbac_migration.py @@ -0,0 +1,14 @@ +from pathlib import Path + + +def test_rbac_migration_normalizes_legacy_roles(): + sql = Path("resources/sql/006_metadata_rbac_roles.sql").read_text( + encoding="utf-8" + ) + + assert "role IN ('admin', 'user')" in sql + assert "project_role IN ('member', 'viewer')" in sql + assert "'modeler'," in sql + assert "'dispatcher'" in sql + assert "THEN 'member'" in sql + assert "ELSE 'viewer'" in sql diff --git a/tests/unit/test_valve_isolation.py b/tests/unit/test_valve_isolation.py new file mode 100644 index 0000000..f628297 --- /dev/null +++ b/tests/unit/test_valve_isolation.py @@ -0,0 +1,71 @@ +from collections import defaultdict + +from app.algorithms.isolation import valve + + +def test_non_isolatable_omits_affected_node_ids_but_keeps_count(monkeypatch): + pipe_adj = defaultdict( + set, + { + "A": {"B"}, + "B": {"A", "C"}, + "C": {"B"}, + }, + ) + topology = ( + pipe_adj, + {"V-optional": ("A", "C")}, + {"P-1": ("A", "B", "pipe")}, + {"A", "B", "C"}, + ) + monkeypatch.setattr(valve, "_get_network_topology", lambda _network: topology) + + result = valve.valve_isolation_analysis("demo", "P-1") + + assert result["isolatable"] is False + assert result["affected_node_count"] == 3 + assert result["affected_nodes"] == [] + assert result["optional_valves"] == ["V-optional"] + + +def test_isolatable_keeps_affected_node_ids_and_count(monkeypatch): + pipe_adj = defaultdict(set, {"A": {"B"}, "B": {"A"}}) + topology = ( + pipe_adj, + {"V-close": ("B", "C")}, + {"P-1": ("A", "B", "pipe")}, + {"A", "B", "C"}, + ) + monkeypatch.setattr(valve, "_get_network_topology", lambda _network: topology) + + result = valve.valve_isolation_analysis("demo", "P-1") + + assert result["isolatable"] is True + assert result["affected_node_count"] == 2 + assert result["affected_nodes"] == ["A", "B"] + assert result["must_close_valves"] == ["V-close"] + + +def test_disabled_valve_expands_affected_area_before_counting(monkeypatch): + pipe_adj = defaultdict(set, {"A": {"B"}, "B": {"A"}}) + topology = ( + pipe_adj, + { + "V-disabled": ("B", "C"), + "V-close": ("C", "D"), + }, + {"P-1": ("A", "B", "pipe")}, + {"A", "B", "C", "D"}, + ) + monkeypatch.setattr(valve, "_get_network_topology", lambda _network: topology) + + result = valve.valve_isolation_analysis( + "demo", + "P-1", + disabled_valves=["V-disabled"], + ) + + assert result["isolatable"] is True + assert result["affected_node_count"] == 3 + assert result["affected_nodes"] == ["A", "B", "C"] + assert result["must_close_valves"] == ["V-close"]