feat(server): add project RBAC and guarded workflows
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
@@ -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(
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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": "管网更新成功"})
|
||||
@@ -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="文件名"),
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
"""
|
||||
与服务器同步
|
||||
|
||||
|
||||
+212
-92
@@ -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],
|
||||
)
|
||||
|
||||
@@ -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)),
|
||||
)
|
||||
@@ -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
|
||||
|
||||
@@ -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]
|
||||
@@ -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"]
|
||||
|
||||
|
||||
@@ -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()
|
||||
]
|
||||
|
||||
@@ -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)
|
||||
);
|
||||
|
||||
|
||||
@@ -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'));
|
||||
@@ -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"]
|
||||
@@ -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,
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
@@ -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):
|
||||
|
||||
@@ -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 = []
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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"]
|
||||
Reference in New Issue
Block a user