Files
ss-tools/backend/tests/test_app_middleware.py
busya 488a8f349b test(backend): raise coverage to 95%+ statements and branches (97.8%/95.0%)
- ~60 new/extended test files across api, core, plugins, services, schemas:
  routes, superset clients, task_manager, lineage, git, translate,
  dashboard-testing, load-testing, migration, llm_analysis, scheduler, ssl
- .coveragerc: enable branch coverage; exclude src/__tests__ (test files)
  and src/scripts (CLI/ops tools) from the denominator
- bug fixes found while testing:
  * settings: PUT /settings/reports registered under duplicated prefix
  * schemas/lineage: FleetReportDTO missing run_status (route always 500)
  * dashboard_testing/baseline_inheritance: visual entry read wrong field
  * superset_client/_databases: logger extra name shadowed LogRecord attr
  * routes/datasets: _yaml_string_paths recursion without yield from
  * translate/sql_generator: restore explicit-type timestamp contract
  * baseline_catalog: remove unreachable dashboard_id fallback
- conftest fixes: pytest_plugins to rootdir conftest (pytest 9), test
  filename collision, TMPDIR-safe integration fixtures
2026-08-19 17:14:32 +03:00

368 lines
16 KiB
Python

# #region Test.AppModule.Middleware [C:3] [TYPE Module] [SEMANTICS test,app,middleware,hsts,cors,logging]
# @BRIEF Tests for app.py — middleware: HSTS, log_requests, CORS, Session, TraceContext, router registration.
# @RELATION BINDS_TO -> [App.AppModule]
# @TEST_EDGE: hsts_enabled -> header set when FORCE_HTTPS=true
# @TEST_EDGE: hsts_disabled -> no header when FORCE_HTTPS=false
# @TEST_EDGE: polling_endpoint_no_log -> polling /api/tasks GET does not log
# @TEST_INVARIANT: hsts_gate -> VERIFIED_BY: test_hsts_middleware_enabled, test_hsts_middleware_disabled
from pathlib import Path
import sys
import os
sys.path.insert(0, str(Path(__file__).parent.parent / "src"))
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from fastapi import Request, HTTPException
# #region _make_mock_request [C:1] [TYPE Function]
def _make_mock_request(method="GET", path="/api/test", client_host="127.0.0.1"):
req = MagicMock(spec=Request)
req.method = method
req.url = MagicMock()
req.url.path = path
req.client = MagicMock()
req.client.host = client_host
req.query_params = {}
return req
# #endregion _make_mock_request
class TestHSTSMiddleware:
"""HSTSMiddleware — Strict-Transport-Security header."""
# #region Test.AppModule.TestHstsEnabled [C:2] [TYPE Function]
@patch.dict(os.environ, {"FORCE_HTTPS": "true"}, clear=False)
@pytest.mark.asyncio
async def test_hsts_enabled(self):
from src.app import HSTSMiddleware
async def call_next(r):
resp = MagicMock(); resp.headers = {}; return resp
mw = HSTSMiddleware(MagicMock())
resp = await mw.dispatch(_make_mock_request(), call_next)
assert resp.headers.get("Strict-Transport-Security") == "max-age=31536000; includeSubDomains"
# #endregion Test.AppModule.TestHstsEnabled
# #region Test.AppModule.TestHstsDisabled [C:2] [TYPE Function]
@patch.dict(os.environ, {"FORCE_HTTPS": "false"}, clear=False)
@pytest.mark.asyncio
async def test_hsts_disabled(self):
from src.app import HSTSMiddleware
async def call_next(r):
resp = MagicMock(); resp.headers = {}; return resp
mw = HSTSMiddleware(MagicMock())
assert mw._enabled is False
resp = await mw.dispatch(_make_mock_request(), call_next)
assert "Strict-Transport-Security" not in resp.headers
# #endregion Test.AppModule.TestHstsDisabled
# #region Test.AppModule.TestHstsEmptyEnv [C:2] [TYPE Function]
@patch.dict(os.environ, {}, clear=False)
def test_hsts_empty_env(self):
os.environ.pop("FORCE_HTTPS", None)
from src.app import HSTSMiddleware
mw = HSTSMiddleware(MagicMock())
assert mw._enabled is False
# #endregion Test.AppModule.TestHstsEmptyEnv
class TestLogRequests:
"""log_requests — HTTP request/response logging."""
# #region Test.AppModule.TestLogNonPolling [C:2] [TYPE Function]
@pytest.mark.asyncio
async def test_log_non_polling(self):
from src.app import log_requests
async def call_next(r):
resp = MagicMock(); resp.status_code = 200; return resp
resp = await log_requests(_make_mock_request(path="/api/dashboards"), call_next)
assert resp.status_code == 200
# #endregion Test.AppModule.TestLogNonPolling
# #region Test.AppModule.TestLogPollingSkipped [C:2] [TYPE Function]
@pytest.mark.asyncio
async def test_log_polling_skipped(self):
from src.app import log_requests
async def call_next(r):
resp = MagicMock(); resp.status_code = 200; return resp
resp = await log_requests(_make_mock_request(method="GET", path="/api/tasks"), call_next)
assert resp.status_code == 200
# #endregion Test.AppModule.TestLogPollingSkipped
# #region Test.AppModule.TestLogPostPollingNotSkipped [C:2] [TYPE Function]
@pytest.mark.asyncio
async def test_log_post_polling_not_skipped(self):
from src.app import log_requests
async def call_next(r):
resp = MagicMock(); resp.status_code = 201; return resp
resp = await log_requests(_make_mock_request(method="POST", path="/api/tasks"), call_next)
assert resp.status_code == 201
# #endregion Test.AppModule.TestLogPostPollingNotSkipped
# #region Test.AppModule.TestLogTaskIdPollingSkipped [C:2] [TYPE Function]
# @TEST_EDGE: polling_task_id_no_log -> GET /api/tasks/{id} (progress polling) is suppressed
@pytest.mark.asyncio
async def test_log_task_id_polling_skipped(self):
from src.app import log_requests
async def call_next(_r):
resp = MagicMock(); resp.status_code = 200; return resp
resp = await log_requests(
_make_mock_request(method="GET", path="/api/tasks/0364fcbe-3af4-4cc2-a98c-599e6cbb93ef"),
call_next,
)
assert resp.status_code == 200
# #endregion Test.AppModule.TestLogTaskIdPollingSkipped
# #region Test.AppModule.TestLogAgentLlmConfigPollingSkipped [C:2] [TYPE Function]
# @TEST_EDGE: llm_config_polling_no_log -> /api/agent/llm-config requests are suppressed
@pytest.mark.asyncio
async def test_log_agent_llm_config_polling_skipped(self):
from src.app import log_requests
async def call_next(_r):
resp = MagicMock(); resp.status_code = 401; return resp
resp = await log_requests(
_make_mock_request(method="GET", path="/api/agent/llm-config"),
call_next,
)
assert resp.status_code == 401
# #endregion Test.AppModule.TestLogAgentLlmConfigPollingSkipped
# #region Test.AppModule.TestLogSessionActivityPollingSkipped [C:2] [TYPE Function]
# @TEST_EDGE: session_activity_no_log -> POST /api/auth/session/activity is suppressed
@pytest.mark.asyncio
async def test_log_session_activity_polling_skipped(self):
from src.app import log_requests
async def call_next(_r):
resp = MagicMock(); resp.status_code = 200; return resp
resp = await log_requests(
_make_mock_request(method="POST", path="/api/auth/session/activity"),
call_next,
)
assert resp.status_code == 200
# #endregion Test.AppModule.TestLogSessionActivityPollingSkipped
# #region Test.AppModule.TestLogSettingsConsolidatedPollingSkipped [C:2] [TYPE Function]
# @TEST_EDGE: settings_consolidated_no_log -> GET /api/settings/consolidated is suppressed
@pytest.mark.asyncio
async def test_log_settings_consolidated_polling_skipped(self):
from src.app import log_requests
async def call_next(_r):
resp = MagicMock(); resp.status_code = 200; return resp
resp = await log_requests(
_make_mock_request(method="GET", path="/api/settings/consolidated"),
call_next,
)
assert resp.status_code == 200
# #endregion Test.AppModule.TestLogSettingsConsolidatedPollingSkipped
# #region Test.AppModule.TestLogNetworkError [C:2] [TYPE Function]
@pytest.mark.asyncio
async def test_log_network_error(self):
from src.app import log_requests
from src.core.utils.network import NetworkError
async def call_next(r):
raise NetworkError("timeout")
with pytest.raises(HTTPException) as e:
await log_requests(_make_mock_request(path="/api/dashboards"), call_next)
assert e.value.status_code == 503
# #endregion Test.AppModule.TestLogNetworkError
# #region Test.AppModule.TestLogExistingTraceSkipsSeed [C:2] [TYPE Function]
@pytest.mark.asyncio
async def test_log_existing_trace_skips_seed(self):
"""When a trace_id is already set, log_requests must not re-seed (484->486)."""
from src.app import log_requests
from ss_tools.shared.cot_logger import seed_trace_id
seed_trace_id()
async def call_next(r):
resp = MagicMock(); resp.status_code = 200; return resp
with patch("src.app.seed_trace_id") as mseed:
resp = await log_requests(_make_mock_request(path="/api/dashboards"), call_next)
assert resp.status_code == 200
mseed.assert_not_called()
# #endregion Test.AppModule.TestLogExistingTraceSkipsSeed
class TestLogRequestsStderrFallback:
"""log_requests — stderr JSON dump when logger level > INFO."""
# #region Test.AppModule.TestStderrFallbackReq [C:2] [TYPE Function]
@pytest.mark.asyncio
async def test_stderr_fallback_req(self):
from src.app import log_requests
async def call_next(r):
resp = MagicMock(); resp.status_code = 200; return resp
with patch("src.app.logger.isEnabledFor", return_value=False):
resp = await log_requests(_make_mock_request(path="/api/test-stderr"), call_next)
assert resp.status_code == 200
# #endregion Test.AppModule.TestStderrFallbackReq
# #region Test.AppModule.TestStderrFallbackResp [C:2] [TYPE Function]
@pytest.mark.asyncio
async def test_stderr_fallback_resp(self):
from src.app import log_requests
async def call_next(r):
resp = MagicMock(); resp.status_code = 500; return resp
with patch("src.app.logger.isEnabledFor", return_value=False):
resp = await log_requests(_make_mock_request(path="/api/test-stderr-resp"), call_next)
assert resp.status_code == 500
# #endregion Test.AppModule.TestStderrFallbackResp
class TestMiddlewareConfig:
"""Middleware registration — Session, CORS, HSTS, TraceContext."""
def test_session_middleware(self):
from src.app import app
names = [m.cls.__name__ for m in app.user_middleware]
assert "SessionMiddleware" in names
def test_cors_middleware(self):
from src.app import app
names = [m.cls.__name__ for m in app.user_middleware]
assert "CORSMiddleware" in names
def test_hsts_middleware(self):
from src.app import app
names = [m.cls.__name__ for m in app.user_middleware]
assert "HSTSMiddleware" in names
def test_trace_context_middleware(self):
from src.app import app
names = [m.cls.__name__ for m in app.user_middleware]
assert "TraceContextMiddleware" in names
class TestCorsConfig:
"""CORS and Session configuration edge cases."""
def test_cors_empty_allowed_origins(self):
from src.app import app
names = [m.cls.__name__ for m in app.user_middleware]
assert "CORSMiddleware" in names
def test_session_secret_fallback(self):
from src.app import app
names = [m.cls.__name__ for m in app.user_middleware]
assert "SessionMiddleware" in names
class TestCorsAllowedOriginsEnv:
"""CORS allow_origins populated from ALLOWED_ORIGINS at import time (line 364)."""
def test_cors_allow_origins_from_env(self):
"""ALLOWED_ORIGINS set before import -> CORSMiddleware allow_origins list."""
import json
import subprocess
import sys
import textwrap
backend = str(Path(__file__).resolve().parent.parent)
script = textwrap.dedent(f"""
import os, sys, json
os.environ["ALLOWED_ORIGINS"] = "https://one.example, https://two.example, "
sys.path.insert(0, {backend!r})
import pathlib as _p
_orig_mkdir = _p.Path.mkdir
def _safe_mkdir(self, mode=0o777, parents=False, exist_ok=False):
if str(self).startswith("/app"):
return
return _orig_mkdir(self, mode, parents=parents, exist_ok=exist_ok)
_p.Path.mkdir = _safe_mkdir
_orig_makedirs = os.makedirs
def _safe_makedirs(path, mode=0o777, exist_ok=False):
if str(path).startswith("/app"):
return
return _orig_makedirs(path, mode, exist_ok=exist_ok)
os.makedirs = _safe_makedirs
from src.app import app
cors = [m for m in app.user_middleware if m.cls.__name__ == "CORSMiddleware"]
print("CORS_RESULT=" + json.dumps(cors[0].kwargs.get("allow_origins")))
""")
env = {**os.environ, "ALLOWED_ORIGINS": "https://one.example, https://two.example, "}
proc = subprocess.run(
[sys.executable, "-c", script],
capture_output=True, text=True, cwd=backend, env=env, timeout=180,
)
assert proc.returncode == 0, proc.stderr
cors_line = next(l for l in proc.stdout.splitlines() if l.startswith("CORS_RESULT="))
assert json.loads(cors_line.split("=", 1)[1]) == ["https://one.example", "https://two.example"]
def test_cors_allow_origins_unset_empty(self):
"""ALLOWED_ORIGINS unset -> CORSMiddleware allow_origins is empty."""
import json
import subprocess
import sys
import textwrap
backend = str(Path(__file__).resolve().parent.parent)
script = textwrap.dedent(f"""
import os, sys, json
os.environ.pop("ALLOWED_ORIGINS", None)
sys.path.insert(0, {backend!r})
import pathlib as _p
_orig_mkdir = _p.Path.mkdir
def _safe_mkdir(self, mode=0o777, parents=False, exist_ok=False):
if str(self).startswith("/app"):
return
return _orig_mkdir(self, mode, parents=parents, exist_ok=exist_ok)
_p.Path.mkdir = _safe_mkdir
_orig_makedirs = os.makedirs
def _safe_makedirs(path, mode=0o777, exist_ok=False):
if str(path).startswith("/app"):
return
return _orig_makedirs(path, mode, exist_ok=exist_ok)
os.makedirs = _safe_makedirs
from src.app import app
cors = [m for m in app.user_middleware if m.cls.__name__ == "CORSMiddleware"]
print("CORS_RESULT=" + json.dumps(cors[0].kwargs.get("allow_origins")))
""")
env = {k: v for k, v in os.environ.items() if k != "ALLOWED_ORIGINS"}
proc = subprocess.run(
[sys.executable, "-c", script],
capture_output=True, text=True, cwd=backend, env=env, timeout=180,
)
assert proc.returncode == 0, proc.stderr
cors_line = next(l for l in proc.stdout.splitlines() if l.startswith("CORS_RESULT="))
assert json.loads(cors_line.split("=", 1)[1]) == []
class TestRouterRegistration:
"""Verify all expected route groups are registered."""
def test_expected_routers_registered(self):
from src.app import app
routes = "\n".join(r.path for r in app.routes)
assert "/api/dashboards" in routes or "/api/plugins" in routes
assert "/api/tasks" in routes
assert "/api/settings" in routes
assert "/api/git" in routes or "/api/mappings" in routes
assert "/ws/logs/" in routes
assert "/ws/task-events" in routes
class TestAlembicMigrations:
"""run_alembic_migrations — test happy path and error path."""
def test_run_alembic_migrations_success(self):
"""Alembic migrations apply successfully (lines 205-213)."""
from src.app import run_alembic_migrations
with patch("alembic.command.upgrade") as mock_upgrade, \
patch("alembic.config.Config") as mock_cfg_cls:
mock_cfg = MagicMock()
mock_cfg_cls.return_value = mock_cfg
run_alembic_migrations()
mock_upgrade.assert_called_once_with(mock_cfg, "head")
def test_run_alembic_migrations_failure(self):
"""Alembic migration failure raises (lines 214-219)."""
from src.app import run_alembic_migrations
with patch("alembic.command.upgrade", side_effect=Exception("Migration failed")) as mock_upgrade, \
patch("alembic.config.Config") as mock_cfg_cls:
mock_cfg = MagicMock()
mock_cfg_cls.return_value = mock_cfg
with pytest.raises(Exception, match="Migration failed"):
run_alembic_migrations()
mock_upgrade.assert_called_once_with(mock_cfg, "head")
# #endregion Test.AppModule.Middleware