Files
TJWaterServerBinary/tests/api/test_openapi_contract.py
T

257 lines
8.1 KiB
Python

from __future__ import annotations
from pathlib import Path
from uuid import uuid4
import pytest
from fastapi import FastAPI
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.v1.rest_router import api_router, build_rest_router
from app.api.v1.router import api_router as source_api_router
from app.auth.project_dependencies import ProjectContext, get_project_context
from scripts.check_openapi import current_contract_bytes, validate
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_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_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", {})
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_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,
)
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")
app.dependency_overrides[get_project_context] = lambda: ProjectContext(
project_id=uuid4(),
project_code="fengyang",
user_id=uuid4(),
project_role="member",
)
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"],
}
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,
)
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",
)
app.dependency_overrides[get_project_context] = lambda: 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,
}
assert response.json()["items"] == [
{"scheme_name": "burst_case", "scheme_type": "burst_analysis"}
]