feat(server): add project RBAC and guarded workflows

This commit is contained in:
2026-07-30 16:45:09 +08:00
parent 3fbb17bb30
commit ae1a657554
29 changed files with 1431 additions and 412 deletions
+4 -2
View File
@@ -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:
+39
View File
@@ -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),
)
+1
View File
@@ -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(
+3
View File
@@ -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,
)
+57 -53
View File
@@ -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
+297
View File
@@ -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": "管网更新成功"})
+11 -63
View File
@@ -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="文件名"),
+6 -11
View File
@@ -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,
+3 -45
View File
@@ -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)
+7 -2
View File
@@ -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
View File
@@ -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],
)
+162
View File
@@ -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)),
)
+81 -89
View File
@@ -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
+13
View File
@@ -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]
+2 -2
View File
@@ -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)
);
+32
View File
@@ -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'));
+77
View File
@@ -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"]
+17 -12
View File
@@ -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,
+16 -2
View File
@@ -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",
}
+7 -1
View File
@@ -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:
+1 -1
View File
@@ -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)
+100
View File
@@ -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()
+19 -14
View File
@@ -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):
-21
View File
@@ -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 = []
+159
View File
@@ -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
+14
View File
@@ -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
+71
View File
@@ -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"]