fix(security): close backend merge blockers

This commit is contained in:
2026-08-18 17:51:29 +08:00
parent 2581631b51
commit 8853877fcd
15 changed files with 216 additions and 50 deletions
+14
View File
@@ -0,0 +1,14 @@
from __future__ import annotations
from collections.abc import Iterable
from typing import Generic, TypeVar
T = TypeVar("T")
class PaginatedList(list[T], Generic[T]):
"""A page of items carrying the total count from its data source."""
def __init__(self, items: Iterable[T], *, total: int) -> None:
super().__init__(items)
self.total = total
+19 -2
View File
@@ -10,6 +10,7 @@ from app.auth.metadata_dependencies import (
get_current_metadata_admin, get_current_metadata_admin,
get_current_metadata_user, get_current_metadata_user,
) )
from app.api.pagination import PaginatedList
from app.core.audit import AuditAction, log_audit_event from app.core.audit import AuditAction, log_audit_event
from app.domain.schemas.audit import AuditLogResponse from app.domain.schemas.audit import AuditLogResponse
from app.infra.db.metadb.database import get_metadata_session from app.infra.db.metadb.database import get_metadata_session
@@ -46,7 +47,7 @@ async def get_audit_logs(
_current_user=Depends(get_current_metadata_admin), _current_user=Depends(get_current_metadata_admin),
audit_repo: AuditRepository = Depends(get_audit_repository), audit_repo: AuditRepository = Depends(get_audit_repository),
) -> list[AuditLogResponse]: ) -> list[AuditLogResponse]:
return await audit_repo.get_logs( items = await audit_repo.get_logs(
user_id=user_id, user_id=user_id,
project_id=project_id, project_id=project_id,
action=action, action=action,
@@ -56,6 +57,15 @@ async def get_audit_logs(
skip=skip, skip=skip,
limit=limit, limit=limit,
) )
total = await audit_repo.get_log_count(
user_id=user_id,
project_id=project_id,
action=action,
resource_type=resource_type,
start_time=start_time,
end_time=end_time,
)
return PaginatedList(items, total=total)
@router.get( @router.get(
@@ -119,7 +129,7 @@ async def get_my_audit_logs(
current_user=Depends(get_current_metadata_user), current_user=Depends(get_current_metadata_user),
audit_repo: AuditRepository = Depends(get_audit_repository), audit_repo: AuditRepository = Depends(get_audit_repository),
) -> list[AuditLogResponse]: ) -> list[AuditLogResponse]:
return await audit_repo.get_logs( items = await audit_repo.get_logs(
user_id=current_user.id, user_id=current_user.id,
action=action, action=action,
start_time=start_time, start_time=start_time,
@@ -127,3 +137,10 @@ async def get_my_audit_logs(
skip=skip, skip=skip,
limit=limit, limit=limit,
) )
total = await audit_repo.get_log_count(
user_id=current_user.id,
action=action,
start_time=start_time,
end_time=end_time,
)
return PaginatedList(items, total=total)
+9 -3
View File
@@ -14,6 +14,7 @@ from pydantic import BaseModel, JsonValue, create_model
from starlette.responses import Response from starlette.responses import Response
from app.api.problem_details import ProblemDetails from app.api.problem_details import ProblemDetails
from app.api.pagination import PaginatedList
from app.api.v1.router import api_router as handler_api_router from app.api.v1.router import api_router as handler_api_router
from app.auth.metadata_dependencies import get_current_metadata_user from app.auth.metadata_dependencies import get_current_metadata_user
from app.auth.project_dependencies import ProjectContext, get_project_context from app.auth.project_dependencies import ProjectContext, get_project_context
@@ -41,8 +42,8 @@ _PUBLIC_PARAMETER_RENAMES = {
"burst_ID": "burst_id", "burst_ID": "burst_id",
"drainage_node_ID": "drainage_node_id", "drainage_node_ID": "drainage_node_id",
} }
_MODEL_NAME_IS_NETWORK = {"RunSimulationManuallyByDate"} _MODEL_NAME_IS_NETWORK = {"RunSimulationManuallyByDate", "PressureSensorPlacement"}
_MODEL_USERNAME_FROM_AUTH: set[str] = set() _MODEL_USERNAME_FROM_AUTH = {"PressureSensorPlacement"}
def _clean_name(name: str) -> str: def _clean_name(name: str) -> str:
@@ -230,9 +231,14 @@ def _with_pagination(endpoint):
if not isinstance(result, list): if not isinstance(result, list):
return result return result
if handler_handles_pagination: if handler_handles_pagination:
if not isinstance(result, PaginatedList):
raise RuntimeError(
f"Paginated handler {endpoint.__name__!r} must return "
"PaginatedList with the real total"
)
return Page( return Page(
items=result, items=result,
total=offset + len(result), total=result.total,
limit=limit or len(result), limit=limit or len(result),
offset=offset, offset=offset,
) )
+3 -1
View File
@@ -53,6 +53,7 @@ from app.api.v1.endpoints.timeseries import (
) )
from app.auth.permissions import ( from app.auth.permissions import (
BURST_RUN, BURST_RUN,
ENVIRONMENT_MANAGE,
OPTIMIZATION_RUN, OPTIMIZATION_RUN,
RISK_RUN, RISK_RUN,
SCADA_CLEAN, SCADA_CLEAN,
@@ -88,6 +89,7 @@ simulation_access = Depends(
webgis_view_access = Depends(require_permission(WEBGIS_VIEW)) webgis_view_access = Depends(require_permission(WEBGIS_VIEW))
simulation_run_access = Depends(require_permission(SIMULATION_RUN)) simulation_run_access = Depends(require_permission(SIMULATION_RUN))
environment_manage_access = Depends(require_permission(ENVIRONMENT_MANAGE))
burst_run_access = Depends(require_permission(BURST_RUN)) burst_run_access = Depends(require_permission(BURST_RUN))
risk_run_access = Depends(require_permission(RISK_RUN)) risk_run_access = Depends(require_permission(RISK_RUN))
optimization_run_access = Depends(require_permission(OPTIMIZATION_RUN)) optimization_run_access = Depends(require_permission(OPTIMIZATION_RUN))
@@ -169,7 +171,7 @@ api_router.include_router(
api_router.include_router( api_router.include_router(
cache.router, cache.router,
tags=["Cache"], tags=["Cache"],
dependencies=[simulation_run_access], dependencies=[environment_manage_access],
) )
api_router.include_router( api_router.include_router(
web_search.router, web_search.router,
+1
View File
@@ -130,6 +130,7 @@ def sanitize_sensitive_data(data: dict) -> dict:
"token", "token",
"api_key", "api_key",
"apikey", "apikey",
"dsn",
"credit_card", "credit_card",
"ssn", "ssn",
"social_security", "social_security",
+5 -6
View File
@@ -1,6 +1,5 @@
# influxdb数据库连接信息 from app.core.config import settings
url = "http://127.0.0.1:8086" # 替换为你的InfluxDB实例地址
token = "kMPX2V5HsbzPpUT2B9HPBu1sTG1Emf-lPlT2UjxYnGAuocpXq_f_0lK4HHs-TbbKyjsZpICkMsyXG_V2D7P7yQ==" # 替换为你的InfluxDB Token url = settings.INFLUXDB_URL
# _ENCODED_TOKEN = "eEdETTVSWnFSSkF1ekFHUy1vdFhVZEMyTkZkWTc1cUpBalJMcUFCNHA1V2NJSUFsSVVwT3BUOF95QTE2QU9IbUpXZXJ3UV8wOGd3Yjg0c3k0MmpuWlE9PQ==" token = settings.INFLUXDB_TOKEN
# token = base64.b64decode(_ENCODED_TOKEN).decode("utf-8") org = settings.INFLUXDB_ORG
org = "TJWATERORG" # 替换为你的Organization名称
+18 -6
View File
@@ -1,3 +1,4 @@
from collections.abc import Mapping, Sequence
from typing import Any from typing import Any
from psycopg.rows import dict_row, Row from psycopg.rows import dict_row, Row
from .connection import project_connection from .connection import project_connection
@@ -82,27 +83,38 @@ class DbChangeSet:
return DbChangeSet(redo_sql, undo_sql, redo_cs_s, undo_cs_s) return DbChangeSet(redo_sql, undo_sql, redo_cs_s, undo_cs_s)
def read(name: str, sql: str) -> Row: QueryParams = Sequence[Any] | Mapping[str, Any]
def _execute(cur, sql: str, params: QueryParams | None = None):
return cur.execute(sql, params) if params is not None else cur.execute(sql)
def read(name: str, sql: str, params: QueryParams | None = None) -> Row:
with project_connection(name) as conn: with project_connection(name) as conn:
with conn.cursor(row_factory=dict_row) as cur: with conn.cursor(row_factory=dict_row) as cur:
cur.execute(sql) _execute(cur, sql, params)
row = cur.fetchone() row = cur.fetchone()
if row == None: if row == None:
raise Exception(sql) raise Exception(sql)
return row return row
def read_all(name: str, sql: str) -> list[Row]: def read_all(
name: str, sql: str, params: QueryParams | None = None
) -> list[Row]:
with project_connection(name) as conn: with project_connection(name) as conn:
with conn.cursor(row_factory=dict_row) as cur: with conn.cursor(row_factory=dict_row) as cur:
cur.execute(sql) _execute(cur, sql, params)
return cur.fetchall() return cur.fetchall()
def try_read(name: str, sql: str) -> Row | None: def try_read(
name: str, sql: str, params: QueryParams | None = None
) -> Row | None:
with project_connection(name) as conn: with project_connection(name) as conn:
with conn.cursor(row_factory=dict_row) as cur: with conn.cursor(row_factory=dict_row) as cur:
cur.execute(sql) _execute(cur, sql, params)
return cur.fetchone() return cur.fetchone()
+20 -9
View File
@@ -1,3 +1,4 @@
from psycopg import sql
from psycopg.rows import dict_row, Row from psycopg.rows import dict_row, Row
from .connection import project_connection from .connection import project_connection
from .database import read from .database import read
@@ -49,7 +50,12 @@ ELEMENT_TYPES : dict[str, int] = {
def _get_from(name: str, id: str, base_type: str) -> Row | None: def _get_from(name: str, id: str, base_type: str) -> Row | None:
with project_connection(name) as conn: with project_connection(name) as conn:
with conn.cursor(row_factory=dict_row) as cur: with conn.cursor(row_factory=dict_row) as cur:
cur.execute(f"select * from {base_type} where id = '{id}'") cur.execute(
sql.SQL("select * from {} where id = %s").format(
sql.Identifier(base_type)
),
(id,),
)
return cur.fetchone() return cur.fetchone()
@@ -243,11 +249,17 @@ def get_node_links(name: str, id: str) -> list[str]:
with project_connection(name) as conn: with project_connection(name) as conn:
with conn.cursor(row_factory=dict_row) as cur: with conn.cursor(row_factory=dict_row) as cur:
links: list[str] = [] links: list[str] = []
for p in cur.execute(f"select id from pipes where node1 = '{id}' or node2 = '{id}'").fetchall(): for p in cur.execute(
"select id from pipes where node1 = %s or node2 = %s", (id, id)
).fetchall():
links.append(p['id']) links.append(p['id'])
for p in cur.execute(f"select id from pumps where node1 = '{id}' or node2 = '{id}'").fetchall(): for p in cur.execute(
"select id from pumps where node1 = %s or node2 = %s", (id, id)
).fetchall():
links.append(p['id']) links.append(p['id'])
for p in cur.execute(f"select id from valves where node1 = '{id}' or node2 = '{id}'").fetchall(): for p in cur.execute(
"select id from valves where node1 = %s or node2 = %s", (id, id)
).fetchall():
links.append(p['id']) links.append(p['id'])
return links return links
@@ -255,16 +267,15 @@ def get_node_links(name: str, id: str) -> list[str]:
def get_link_nodes(name: str, id: str) -> list[str]: def get_link_nodes(name: str, id: str) -> list[str]:
row = {} row = {}
if is_pipe(name, id): if is_pipe(name, id):
row = read(name, f"select node1, node2 from pipes where id = '{id}'") row = read(name, "select node1, node2 from pipes where id = %s", (id,))
elif is_pump(name, id): elif is_pump(name, id):
row = read(name, f"select node1, node2 from pumps where id = '{id}'") row = read(name, "select node1, node2 from pumps where id = %s", (id,))
elif is_valve(name, id): elif is_valve(name, id):
row = read(name, f"select node1, node2 from valves where id = '{id}'") row = read(name, "select node1, node2 from valves where id = %s", (id,))
return [str(row['node1']), str(row['node2'])] return [str(row['node1']), str(row['node2'])]
def get_region_type(name: str, id: str)->str: def get_region_type(name: str, id: str)->str:
if(is_region(name,id)): if(is_region(name,id)):
type = read(name, f"select type from _region where id = '{id}'") type = read(name, "select type from _region where id = %s", (id,))
return type return type
+8 -2
View File
@@ -23,7 +23,11 @@ def from_postgis_point(coord: str) -> dict[str, float]:
def get_node_coord(name: str, node: str) -> dict[str, float]: def get_node_coord(name: str, node: str) -> dict[str, float]:
row = try_read(name, f"select st_astext(coord) as coord_geom from coordinates where node = '{node}'") row = try_read(
name,
"select st_astext(coord) as coord_geom from coordinates where node = %s",
(node,),
)
if row == None: if row == None:
write(name, sql_insert_coord(node, 0.0, 0.0)) write(name, sql_insert_coord(node, 0.0, 0.0))
return {'x': 0.0, 'y': 0.0} return {'x': 0.0, 'y': 0.0}
@@ -66,7 +70,9 @@ def get_links_in_extent(name: str, x1: float, y1: float, x2: float, y2: float) -
def node_has_coord(name: str, node: str) -> bool: def node_has_coord(name: str, node: str) -> bool:
return try_read(name, f"select node from coordinates where node = '{node}'") != None return try_read(
name, "select node from coordinates where node = %s", (node,)
) != None
#-------------------------------------------------------------- #--------------------------------------------------------------
+1 -1
View File
@@ -12,7 +12,7 @@ def get_junction_schema(name: str) -> dict[str, dict[str, Any]]:
def get_junction(name: str, id: str) -> dict[str, Any]: def get_junction(name: str, id: str) -> dict[str, Any]:
j = try_read(name, f"select * from junctions where id = '{id}'") j = try_read(name, "select * from junctions where id = %s", (id,))
if j == None: if j == None:
return {} return {}
xy = get_node_coord(name, id) xy = get_node_coord(name, id)
+1 -1
View File
@@ -3,7 +3,7 @@
"contracts": { "contracts": {
"server": { "server": {
"file": "server-v1.openapi.json", "file": "server-v1.openapi.json",
"sha256": "9ad12d3cd789fd42c341faec5b859d76bceeabc74399e3991e62ace1129e69a2" "sha256": "df7ae927dcf5ae32c3c1ad9be3245b1b78b984ce1902dd91e6313770860e0d48"
} }
} }
} }
+5 -17
View File
@@ -1859,7 +1859,7 @@
"title": "PressureRegulationRest", "title": "PressureRegulationRest",
"type": "object" "type": "object"
}, },
"PressureSensorPlacement": { "PressureSensorPlacementRest": {
"properties": { "properties": {
"min_diameter": { "min_diameter": {
"default": 0, "default": 0,
@@ -1867,11 +1867,6 @@
"title": "Min Diameter", "title": "Min Diameter",
"type": "integer" "type": "integer"
}, },
"name": {
"description": "管网名称(或数据库名称)",
"title": "Name",
"type": "string"
},
"scheme_name": { "scheme_name": {
"description": "方案名称", "description": "方案名称",
"title": "Scheme Name", "title": "Scheme Name",
@@ -1881,20 +1876,13 @@
"description": "传感器数量", "description": "传感器数量",
"title": "Sensor Number", "title": "Sensor Number",
"type": "integer" "type": "integer"
},
"username": {
"description": "用户名",
"title": "Username",
"type": "string"
} }
}, },
"required": [ "required": [
"name",
"scheme_name", "scheme_name",
"sensor_number", "sensor_number"
"username"
], ],
"title": "PressureSensorPlacement", "title": "PressureSensorPlacementRest",
"type": "object" "type": "object"
}, },
"ProblemDetails": { "ProblemDetails": {
@@ -24894,7 +24882,7 @@
"content": { "content": {
"application/json": { "application/json": {
"schema": { "schema": {
"$ref": "#/components/schemas/PressureSensorPlacement", "$ref": "#/components/schemas/PressureSensorPlacementRest",
"description": "传感器放置分析参数" "description": "传感器放置分析参数"
} }
} }
@@ -25134,7 +25122,7 @@
"content": { "content": {
"application/json": { "application/json": {
"schema": { "schema": {
"$ref": "#/components/schemas/PressureSensorPlacement", "$ref": "#/components/schemas/PressureSensorPlacementRest",
"description": "传感器放置分析参数" "description": "传感器放置分析参数"
} }
} }
+15
View File
@@ -7,6 +7,21 @@ from fastapi.testclient import TestClient
from app.infra.audit import middleware as audit_middleware from app.infra.audit import middleware as audit_middleware
from app.infra.audit.middleware import AuditMiddleware from app.infra.audit.middleware import AuditMiddleware
from app.core.audit import sanitize_sensitive_data
def test_sanitize_sensitive_data_redacts_database_dsn() -> None:
raw_dsn = "postgresql://alice:supersecret@db.internal/project"
sanitized = sanitize_sensitive_data(
{"dsn": raw_dsn, "database": {"readonly_dsn": raw_dsn}}
)
assert sanitized == {
"dsn": "***REDACTED***",
"database": {"readonly_dsn": "***REDACTED***"},
}
assert raw_dsn not in str(sanitized)
def test_post_streaming_response_survives_audit_body_capture(monkeypatch): def test_post_streaming_response_survives_audit_body_capture(monkeypatch):
+76 -2
View File
@@ -2,6 +2,7 @@ from __future__ import annotations
from datetime import datetime, timezone from datetime import datetime, timezone
from pathlib import Path from pathlib import Path
from unittest.mock import Mock
from uuid import uuid4 from uuid import uuid4
import pytest import pytest
@@ -11,8 +12,11 @@ from fastapi.testclient import TestClient
from app.api.v1.endpoints import schemes as schemes_endpoint from app.api.v1.endpoints import schemes as schemes_endpoint
from app.api.v1.endpoints import simulation as simulation_endpoint from app.api.v1.endpoints import simulation as simulation_endpoint
from app.api.v1.endpoints import cache as cache_endpoint
from app.api.pagination import PaginatedList
from app.api.v1.rest_router import api_router, build_rest_router from app.api.v1.rest_router import api_router, build_rest_router
from app.api.v1.router import api_router as source_api_router from app.api.v1.router import api_router as source_api_router
from app.auth.metadata_dependencies import get_current_metadata_user
from app.auth.project_dependencies import ProjectContext, get_project_context from app.auth.project_dependencies import ProjectContext, get_project_context
from scripts.check_openapi import current_contract_bytes, validate from scripts.check_openapi import current_contract_bytes, validate
@@ -130,6 +134,12 @@ def test_rest_contract_uses_header_project_context() -> None:
assert "network" not in schema.get("properties", {}) assert "network" not in schema.get("properties", {})
assert "network_name" not in schema.get("properties", {}) assert "network_name" not in schema.get("properties", {})
placement_schema = document["components"]["schemas"][
"PressureSensorPlacementRest"
]
assert "name" not in placement_schema["properties"]
assert "username" not in placement_schema["properties"]
assert "/api/v1/burst-analysis" not in document["paths"] assert "/api/v1/burst-analysis" not in document["paths"]
assert "/api/v1/getpipeproperties/" not in document["paths"] assert "/api/v1/getpipeproperties/" not in document["paths"]
@@ -280,7 +290,7 @@ def test_rest_runtime_wraps_handler_paginated_list() -> None:
limit: int = Query(2, ge=1, le=10), limit: int = Query(2, ge=1, le=10),
) -> list[int]: ) -> list[int]:
records = [10, 20, 30, 40] records = [10, 20, 30, 40]
return records[skip : skip + limit] return PaginatedList(records[skip : skip + limit], total=len(records))
app = FastAPI(redirect_slashes=False) app = FastAPI(redirect_slashes=False)
app.include_router(build_rest_router(source_router.routes), prefix="/api/v1") app.include_router(build_rest_router(source_router.routes), prefix="/api/v1")
@@ -293,12 +303,76 @@ def test_rest_runtime_wraps_handler_paginated_list() -> None:
assert response.status_code == 200 assert response.status_code == 200
assert response.json() == { assert response.json() == {
"items": [20, 30], "items": [20, 30],
"total": 3, "total": 4,
"limit": 2, "limit": 2,
"offset": 1, "offset": 1,
} }
def test_sensor_placement_body_uses_authenticated_project_and_user(
monkeypatch,
) -> None:
captured: dict[str, object] = {}
def fake_pressure_sensor_placement_kmeans(**kwargs):
captured.update(kwargs)
monkeypatch.setattr(
simulation_endpoint,
"pressure_sensor_placement_kmeans",
fake_pressure_sensor_placement_kmeans,
)
app = FastAPI(redirect_slashes=False)
app.include_router(api_router, prefix="/api/v1")
app.dependency_overrides[get_project_context] = lambda: ProjectContext(
project_id=uuid4(),
project_code="project_a",
user_id=uuid4(),
project_role="member",
)
app.dependency_overrides[get_current_metadata_user] = lambda: type(
"User", (), {"username": "alice"}
)()
response = TestClient(app, raise_server_exceptions=False).post(
"/api/v1/pressure-sensor-placement-kmeans",
json={
"scheme_name": "placement_01",
"sensor_number": 5,
"min_diameter": 100,
},
)
assert response.status_code == 200
assert captured == {
"name": "project_a",
"scheme_name": "placement_01",
"sensor_number": 5,
"min_diameter": 100,
"username": "alice",
}
def test_cache_management_requires_environment_permission(monkeypatch) -> None:
flushdb = Mock(return_value=True)
monkeypatch.setattr(cache_endpoint.redis_client, "flushdb", flushdb)
app = FastAPI(redirect_slashes=False)
app.include_router(api_router, prefix="/api/v1")
app.dependency_overrides[get_project_context] = lambda: ProjectContext(
project_id=uuid4(),
project_code="project_a",
user_id=uuid4(),
project_role="member",
)
response = TestClient(app, raise_server_exceptions=False).delete(
"/api/v1/all-redis"
)
assert response.status_code == 403
flushdb.assert_not_called()
def test_rest_runtime_json_encodes_untyped_datetime_response() -> None: def test_rest_runtime_json_encodes_untyped_datetime_response() -> None:
source_router = APIRouter() source_router = APIRouter()
+21
View File
@@ -0,0 +1,21 @@
from app.native.wndb import s2_junctions
def test_get_junction_binds_untrusted_identifier(monkeypatch) -> None:
calls: list[tuple[str, str, tuple[str, ...]]] = []
malicious_id = "J-1'; DELETE FROM junctions; --"
def fake_try_read(name, statement, params):
calls.append((name, statement, params))
return None
monkeypatch.setattr(s2_junctions, "try_read", fake_try_read)
assert s2_junctions.get_junction("project_a", malicious_id) == {}
assert calls == [
(
"project_a",
"select * from junctions where id = %s",
(malicious_id,),
)
]