Files
TJWaterServerBinary/tests/api/test_openapi_contract.py

448 lines
14 KiB
Python

from __future__ import annotations
import inspect
from datetime import datetime, timezone
from pathlib import Path
from uuid import uuid4
import pytest
from fastapi import APIRouter, FastAPI, Query
from fastapi.routing import APIRoute
from fastapi.testclient import TestClient
from app.api.v1.endpoints import schemes as schemes_endpoint
from app.api.v1.endpoints import simulation as simulation_endpoint
from app.api.pagination import PaginatedList
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.auth.metadata_dependencies import get_current_metadata_user
from app.auth.project_dependencies import (
ProjectContext,
get_project_business_routing,
get_project_context,
get_project_simulation_routing,
)
from app.infra.db.project_routing import (
ActiveProjectRouting,
get_project_pgconn_string,
get_project_timescale_pgconn_string,
)
from scripts.check_openapi import current_contract_bytes, validate
def _override_project_routing(
app: FastAPI,
project_context: ProjectContext,
) -> None:
app.dependency_overrides[get_project_context] = lambda: project_context
business = ActiveProjectRouting(
project_code=project_context.project_code,
business_dsn=f"postgresql://user:password@biz/{project_context.project_code}",
)
simulation = ActiveProjectRouting(
project_code=project_context.project_code,
business_dsn=business.business_dsn,
timescale_dsn=(
f"postgresql://user:password@timescale/{project_context.project_code}"
),
)
app.dependency_overrides[get_project_business_routing] = lambda: business
app.dependency_overrides[get_project_simulation_routing] = lambda: simulation
def test_rest_router_preserves_every_distinct_source_operation() -> None:
skipped_names = {"fastapi_get_json", "fastapi_test_dict"}
source_names = {
route.name
for route in source_api_router.routes
if isinstance(route, APIRoute) and route.name not in skipped_names
}
rest_names = {
route.name for route in api_router.routes if isinstance(route, APIRoute)
}
assert rest_names == source_names
def test_rest_router_has_unique_method_path_pairs() -> None:
pairs: list[tuple[str, str]] = []
for route in api_router.routes:
if not isinstance(route, APIRoute):
continue
pairs.extend((method, route.path) for method in route.methods or set())
assert len(pairs) == len(set(pairs))
def test_removed_redis_management_routes_are_not_published() -> None:
published_paths = {
route.path for route in api_router.routes if isinstance(route, APIRoute)
}
assert published_paths.isdisjoint(
{"/redis-keys/detail", "/redis-keys", "/all-redis", "/redis"}
)
def test_rest_router_rejects_duplicate_method_path_pairs() -> None:
first = APIRoute(
"/duplicate",
lambda: None,
methods={"POST"},
name="first_endpoint",
)
second = APIRoute(
"/duplicate",
lambda: None,
methods={"POST"},
name="second_endpoint",
)
with pytest.raises(RuntimeError, match="REST route collision"):
build_rest_router([first, second])
def test_handler_router_defines_only_the_public_rest_operations() -> None:
source_operations = {
(method, route.path)
for route in source_api_router.routes
if isinstance(route, APIRoute)
for method in route.methods or set()
}
public_operations = {
(method, route.path)
for route in api_router.routes
if isinstance(route, APIRoute)
for method in route.methods or set()
}
assert source_operations == public_operations
assert ("POST", "/burst-analysis") not in source_operations
assert ("GET", "/getpipeproperties/") not in source_operations
def test_legacy_project_user_operations_are_not_exposed() -> None:
source_paths = {
route.path
for route in source_api_router.routes
if isinstance(route, APIRoute)
}
assert {
"/network-schemas/user",
"/users",
"/users/detail",
}.isdisjoint(source_paths)
def test_rest_openapi_satisfies_contract_invariants() -> None:
from fastapi import FastAPI
app = FastAPI(redirect_slashes=False)
app.include_router(api_router, prefix="/api/v1")
errors = validate(app.openapi())
assert errors == []
def test_openapi_snapshot_matches_current_application() -> None:
contract = Path(__file__).resolve().parents[2] / "contracts/server-v1.openapi.json"
assert contract.read_bytes() == current_contract_bytes()
def test_rest_contract_uses_header_project_context() -> None:
from fastapi import FastAPI
app = FastAPI(redirect_slashes=False)
app.include_router(api_router, prefix="/api/v1")
document = app.openapi()
for path_item in document["paths"].values():
for operation in path_item.values():
if not isinstance(operation, dict):
continue
query_names = {
parameter["name"]
for parameter in operation.get("parameters", [])
if parameter.get("in") == "query"
}
assert "network" not in query_names
assert "network_name" not in query_names
for name, schema in document["components"]["schemas"].items():
if name.endswith("Rest"):
assert "network" 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/getpipeproperties/" not in document["paths"]
pipes_collection = document["paths"]["/api/v1/pipes"]["get"]
query_names = {
parameter["name"]
for parameter in pipes_collection["parameters"]
if parameter["in"] == "query"
}
assert {"limit", "offset"} <= query_names
assert "204" in document["paths"]["/api/v1/pipes"]["delete"]["responses"]
def test_side_effecting_analysis_routes_are_post() -> None:
methods_by_path = {
route.path: route.methods
for route in api_router.routes
if isinstance(route, APIRoute)
}
assert methods_by_path["/burst-analyses"] == {"POST"}
assert methods_by_path["/flushing-analyses"] == {"POST"}
assert methods_by_path["/contaminant-simulations"] == {"POST"}
def test_valve_isolation_route_uses_the_isolation_handler() -> None:
route = next(
route
for route in api_router.routes
if isinstance(route, APIRoute)
and route.path == "/valve-isolation-analyses"
and route.methods == {"POST"}
)
assert route.name == "valve_isolation_endpoint"
def test_open_project_route_requires_business_and_timescale_routing() -> None:
route = next(
route
for route in api_router.routes
if isinstance(route, APIRoute)
and route.path == "/projects/current"
and route.methods == {"POST"}
)
routing_parameter = inspect.signature(route.endpoint).parameters[
"_rest_project_routing"
]
assert routing_parameter.default.dependency is get_project_simulation_routing
def test_valve_isolation_runtime_accepts_frontend_query(monkeypatch) -> None:
captured: dict[str, object] = {}
def fake_analyze_valve_isolation(network, accident_element, disabled_valves):
captured.update(
network=network,
accident_element=accident_element,
disabled_valves=disabled_valves,
business_dsn=get_project_pgconn_string(network),
timescale_dsn=get_project_timescale_pgconn_string(network),
)
return {"isolatable": True, "must_close_valves": ["V-1"]}
monkeypatch.setattr(
simulation_endpoint,
"analyze_valve_isolation",
fake_analyze_valve_isolation,
)
app = FastAPI(redirect_slashes=False)
app.include_router(api_router, prefix="/api/v1")
project_context = ProjectContext(
project_id=uuid4(),
project_code="fengyang",
user_id=uuid4(),
project_role="member",
)
_override_project_routing(app, project_context)
response = TestClient(app, raise_server_exceptions=False).post(
"/api/v1/valve-isolation-analyses",
params=[
("accident_element", "P-1"),
("accident_element", "P-2"),
("disabled_valves", "V-9"),
],
)
assert response.status_code == 200
assert response.json() == {"isolatable": True, "must_close_valves": ["V-1"]}
assert captured == {
"network": "fengyang",
"accident_element": ["P-1", "P-2"],
"disabled_valves": ["V-9"],
"business_dsn": "postgresql://user:password@biz/fengyang",
"timescale_dsn": "postgresql://user:password@timescale/fengyang",
}
def test_scada_cleaning_runs_are_post() -> None:
methods_by_path = {
route.path: route.methods
for route in api_router.routes
if isinstance(route, APIRoute)
}
assert methods_by_path["/timeseries/scada-cleaning-runs"] == {"POST"}
def test_sensor_placement_excel_export_is_post() -> None:
methods_by_path = {
route.path: route.methods
for route in api_router.routes
if isinstance(route, APIRoute)
}
assert methods_by_path[
"/sensor-placement-schemes/{scheme_id}/exports/excel"
] == {"POST"}
def test_rest_runtime_consumes_injected_project_context(monkeypatch) -> None:
captured: dict[str, object] = {}
def fake_get_all_schemes(network, scheme_type=None, query_date=None):
captured.update(
network=network,
scheme_type=scheme_type,
query_date=query_date,
business_dsn=get_project_pgconn_string(network),
)
return [{"scheme_name": "burst_case", "scheme_type": scheme_type}]
monkeypatch.setattr(
schemes_endpoint,
"get_all_schemes",
fake_get_all_schemes,
)
app = FastAPI(redirect_slashes=False)
app.include_router(api_router, prefix="/api/v1")
project_context = ProjectContext(
project_id=uuid4(),
project_code="fengyang",
user_id=uuid4(),
project_role="viewer",
)
_override_project_routing(app, project_context)
response = TestClient(app, raise_server_exceptions=False).get(
"/api/v1/schemes",
params={"scheme_type": "burst_analysis"},
)
assert response.status_code == 200
assert captured == {
"network": "fengyang",
"scheme_type": "burst_analysis",
"query_date": None,
"business_dsn": "postgresql://user:password@biz/fengyang",
}
assert response.json()["items"] == [
{"scheme_name": "burst_case", "scheme_type": "burst_analysis"}
]
def test_rest_runtime_wraps_handler_paginated_list() -> None:
source_router = APIRouter()
@source_router.get("/records", response_model=list[int])
async def list_records(
skip: int = Query(0, ge=0),
limit: int = Query(2, ge=1, le=10),
) -> list[int]:
records = [10, 20, 30, 40]
return PaginatedList(records[skip : skip + limit], total=len(records))
app = FastAPI(redirect_slashes=False)
app.include_router(build_rest_router(source_router.routes), prefix="/api/v1")
response = TestClient(app, raise_server_exceptions=False).get(
"/api/v1/records",
params={"skip": 1, "limit": 2},
)
assert response.status_code == 200
assert response.json() == {
"items": [20, 30],
"total": 4,
"limit": 2,
"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")
project_context = ProjectContext(
project_id=uuid4(),
project_code="project_a",
user_id=uuid4(),
project_role="member",
)
_override_project_routing(app, project_context)
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_rest_runtime_json_encodes_untyped_datetime_response() -> None:
source_router = APIRouter()
@source_router.get("/simulation-result")
async def simulation_result():
return {
"result": [
{
"time": datetime(2026, 7, 30, 4, tzinfo=timezone.utc),
"id": "4277",
}
]
}
app = FastAPI(redirect_slashes=False)
app.include_router(build_rest_router(source_router.routes), prefix="/api/v1")
response = TestClient(app, raise_server_exceptions=False).get(
"/api/v1/simulation-result"
)
assert response.status_code == 200
assert response.json() == {
"result": [
{
"time": "2026-07-30T04:00:00+00:00",
"id": "4277",
}
]
}