feat(auth): add opaque bearer sessions

This commit is contained in:
caoxiaozhu
2026-07-13 14:45:36 +08:00
parent 661990b27b
commit 653eda0596
59 changed files with 1408 additions and 408 deletions

View File

@@ -0,0 +1,53 @@
from __future__ import annotations
from typing import Annotated
from fastapi import FastAPI, Header
from app.api.deps import CurrentUserContext, get_current_user
def install_legacy_header_auth_override(app: FastAPI) -> None:
"""让领域接口测试专注业务断言,不把旧请求头带回生产认证链路。"""
app.dependency_overrides[get_current_user] = _read_test_user_headers
def _read_test_user_headers(
username: Annotated[str | None, Header(alias="X-Auth-Username")] = None,
name: Annotated[str | None, Header(alias="X-Auth-Name")] = None,
role_codes: Annotated[str | None, Header(alias="X-Auth-Role-Codes")] = None,
is_admin: Annotated[str | None, Header(alias="X-Auth-Is-Admin")] = None,
department: Annotated[str | None, Header(alias="X-Auth-Department")] = None,
cost_center: Annotated[str | None, Header(alias="X-Auth-Cost-Center")] = None,
position: Annotated[str | None, Header(alias="X-Auth-Position")] = None,
grade: Annotated[str | None, Header(alias="X-Auth-Grade")] = None,
employee_no: Annotated[str | None, Header(alias="X-Auth-Employee-No")] = None,
manager_name: Annotated[str | None, Header(alias="X-Auth-Manager-Name")] = None,
) -> CurrentUserContext:
normalized_username = str(username or "").strip()
normalized_name = str(name or normalized_username).strip()
if not normalized_username and not normalized_name:
normalized_username = "test-admin"
normalized_name = "Test Admin"
is_admin = "true"
normalized_roles = [
normalized
for item in str(role_codes or "").split(",")
if (normalized := item.strip().lower())
]
admin_flag = str(is_admin or "").strip().lower() in {"1", "true", "yes", "on"}
admin_flag = admin_flag or normalized_username.lower() in {"admin", "superadmin"}
admin_flag = admin_flag or bool(set(normalized_roles) & {"admin", "superadmin"})
return CurrentUserContext(
username=normalized_username or normalized_name,
name=normalized_name or normalized_username,
role_codes=normalized_roles,
is_admin=admin_flag,
department_name=str(department or "").strip(),
cost_center=str(cost_center or "").strip(),
position=str(position or "").strip(),
grade=str(grade or "").strip(),
employee_no=str(employee_no or "").strip(),
manager_name=str(manager_name or "").strip(),
)

View File

@@ -2,6 +2,7 @@ from __future__ import annotations
from collections.abc import Generator
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
@@ -23,7 +24,8 @@ def build_client() -> tuple[TestClient, sessionmaker[Session]]:
Base.metadata.create_all(bind=engine)
session_factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
app = create_app()
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()

View File

@@ -3,6 +3,7 @@ from __future__ import annotations
from collections.abc import Generator
from datetime import UTC, datetime, timedelta
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
@@ -37,6 +38,7 @@ def build_client() -> tuple[TestClient, sessionmaker[Session]]:
Base.metadata.create_all(bind=engine)
session_factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()

View File

@@ -5,6 +5,7 @@ from collections.abc import Generator
from datetime import UTC, date, datetime
from decimal import Decimal
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import select
from sqlalchemy.orm import Session, selectinload
@@ -17,8 +18,8 @@ from app.models.employee import Employee
from app.models.financial_record import ExpenseClaim, ExpenseClaimItem
from app.schemas.ocr import OcrRecognizeBatchRead, OcrRecognizeDocumentRead, OcrRecognizeFieldRead
from app.services.attachment_association_jobs import clear_attachment_association_jobs_for_tests
from app.services.expense_claims import ExpenseClaimService
from app.services.expense_claim_attachment_storage import ExpenseClaimAttachmentStorage
from app.services.expense_claims import ExpenseClaimService
from app.services.ocr import OcrService
from app.services.receipt_folder import ReceiptFolderService
from app.test_helpers.db import build_in_memory_session_factory
@@ -27,6 +28,7 @@ from app.test_helpers.db import build_in_memory_session_factory
def build_client(monkeypatch) -> tuple[TestClient, object]:
session_factory = build_in_memory_session_factory()
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()

View File

@@ -1,13 +1,17 @@
from __future__ import annotations
from sqlalchemy import create_engine
from __future__ import annotations
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.pool import StaticPool
from app.db.base import Base
from app.db.base import Base
from app.models.auth_session import AuthSession
from app.models.user_session_metric import UserSessionMetric
from app.schemas.auth import LoginRequest
from app.schemas.settings import SettingsWrite
from app.services.auth import AuthService, AuthenticatedUser
from app.services.auth import AuthenticatedUser, AuthService
from app.services.auth_sessions import AuthSessionService
from app.services.employee import EmployeeService
from app.services.settings import SettingsService
@@ -37,6 +41,11 @@ def test_employee_can_login_with_seed_default_password() -> None:
assert result.user.grade == employee.grade
assert result.user.roleCodes
assert result.user.isAdmin is False
assert result.accessToken
assert result.tokenType == "Bearer"
stored_session = db.query(AuthSession).one()
assert stored_session.token_hash == AuthSessionService.hash_token(result.accessToken)
assert stored_session.token_hash != result.accessToken
def test_current_user_snapshot_refreshes_employee_position() -> None:
@@ -113,8 +122,15 @@ def test_employee_login_skips_directory_bootstrap_when_employee_exists(monkeypat
calls.append("ensure_directory_ready")
raise AssertionError("existing employee login should not run directory bootstrap")
monkeypatch.setattr(AuthService, "_find_employee_by_email", lambda self, _: ExistingEmployee())
monkeypatch.setattr("app.services.auth.verify_password", lambda password, password_hash: True)
monkeypatch.setattr(
AuthService,
"_find_employee_by_email",
lambda self, _: ExistingEmployee(),
)
monkeypatch.setattr(
"app.services.auth.verify_password",
lambda password, password_hash: True,
)
monkeypatch.setattr(
AuthService,
"_build_employee_user",
@@ -143,3 +159,35 @@ def test_employee_login_skips_directory_bootstrap_when_employee_exists(monkeypat
assert user is not None
assert user.username == "demo@example.com"
assert calls == []
def test_login_session_write_rolls_back_metric_when_token_issue_fails(monkeypatch) -> None:
with build_session() as db:
user = AuthenticatedUser(
username="rollback@example.com",
name="Rollback User",
role="使用者",
department="",
position="",
grade="",
employee_no="",
manager_name="",
location="",
cost_center="",
finance_owner_name="",
risk_profile={},
role_codes=["user"],
email="rollback@example.com",
avatar="R",
)
def fail_issue(*args, **kwargs):
raise RuntimeError("token issue failed")
monkeypatch.setattr(AuthSessionService, "issue", fail_issue)
with pytest.raises(RuntimeError, match="token issue failed"):
AuthService(db)._build_login_response(user)
assert db.query(AuthSession).count() == 0
assert db.query(UserSessionMetric).count() == 0

View File

@@ -0,0 +1,148 @@
from __future__ import annotations
from collections.abc import Generator
from datetime import UTC, datetime, timedelta
import pytest
from fastapi import HTTPException
from fastapi.testclient import TestClient
from sqlalchemy import create_engine, select
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.pool import StaticPool
from app.api.deps import CurrentUserContext, get_db, require_platform_admin_user
from app.db.base import Base
from app.main import create_app
from app.models.auth_session import AuthSession
from app.models.user_session_metric import UserSessionMetric
from app.schemas.settings import SettingsWrite
from app.services.auth_sessions import AuthSessionService
from app.services.settings import SettingsService
def build_client() -> tuple[TestClient, sessionmaker[Session]]:
engine = create_engine(
"sqlite+pysqlite:///:memory:",
connect_args={"check_same_thread": False},
poolclass=StaticPool,
)
Base.metadata.create_all(bind=engine)
session_factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
with session_factory() as db:
settings_service = SettingsService(db)
payload = settings_service.get_settings_snapshot().model_dump()
payload["adminForm"]["adminAccount"] = "auth-admin"
payload["adminForm"]["newPassword"] = "safe-admin-password"
payload["adminForm"]["confirmPassword"] = "safe-admin-password"
settings_service.save_settings_snapshot(SettingsWrite(**payload))
app = create_app()
def override_db() -> Generator[Session, None, None]:
with session_factory() as db:
yield db
app.dependency_overrides[get_db] = override_db
return TestClient(app), session_factory
def login_admin(client: TestClient) -> dict:
response = client.post(
"/api/v1/auth/login",
json={"username": "auth-admin", "password": "safe-admin-password"},
)
assert response.status_code == 200
return response.json()
def test_login_issues_opaque_bearer_token_and_me_uses_server_session() -> None:
client, session_factory = build_client()
payload = login_admin(client)
token = payload["accessToken"]
assert payload["tokenType"] == "Bearer"
assert payload["expiresAt"]
with session_factory() as db:
auth_session = db.scalars(select(AuthSession)).one()
assert auth_session.token_hash == AuthSessionService.hash_token(token)
assert auth_session.token_hash != token
response = client.get(
"/api/v1/auth/me",
headers={"Authorization": f"Bearer {token}"},
)
assert response.status_code == 200
assert response.json()["username"] == "auth-admin"
assert response.json()["isAdmin"] is True
def test_forged_identity_headers_no_longer_authenticate() -> None:
client, _ = build_client()
response = client.get(
"/api/v1/auth/me",
headers={
"X-Auth-Username": "superadmin",
"X-Auth-Role-Codes": "manager",
"X-Auth-Is-Admin": "true",
},
)
assert response.status_code == 401
assert response.headers["www-authenticate"] == "Bearer"
def test_expired_and_revoked_tokens_are_rejected() -> None:
client, session_factory = build_client()
first_payload = login_admin(client)
first_token = first_payload["accessToken"]
with session_factory() as db:
auth_session = db.scalars(
select(AuthSession).where(
AuthSession.token_hash == AuthSessionService.hash_token(first_token)
)
).one()
auth_session.expires_at = datetime.now(UTC) - timedelta(seconds=1)
db.commit()
expired_response = client.get(
"/api/v1/auth/me",
headers={"Authorization": f"Bearer {first_token}"},
)
assert expired_response.status_code == 401
second_payload = login_admin(client)
second_token = second_payload["accessToken"]
logout_response = client.post(
"/api/v1/auth/logout",
headers={"Authorization": f"Bearer {second_token}"},
json={"sessionId": second_payload["sessionId"], "reason": "manual"},
)
assert logout_response.status_code == 200
with session_factory() as db:
metric_session = db.scalars(
select(UserSessionMetric).where(
UserSessionMetric.session_id == second_payload["sessionId"]
)
).one()
assert metric_session.status == "closed"
revoked_response = client.get(
"/api/v1/auth/me",
headers={"Authorization": f"Bearer {second_token}"},
)
assert revoked_response.status_code == 401
def test_manager_role_is_not_platform_admin() -> None:
manager = CurrentUserContext(
username="manager@example.com",
name="Manager",
role_codes=["manager"],
is_admin=False,
)
with pytest.raises(HTTPException) as exc_info:
require_platform_admin_user(manager)
assert exc_info.value.status_code == 403

View File

@@ -4,6 +4,7 @@ from collections.abc import Generator
from datetime import UTC, datetime
from decimal import Decimal
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
@@ -25,6 +26,7 @@ def build_client() -> tuple[TestClient, sessionmaker[Session]]:
Base.metadata.create_all(bind=engine)
session_factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()
@@ -112,7 +114,10 @@ def test_expense_claims_support_page_envelope_and_keep_legacy_list() -> None:
def test_employee_directory_supports_backend_pagination() -> None:
client, _ = build_client()
response = client.get("/api/v1/employees?page=2&page_size=10")
response = client.get(
"/api/v1/employees?page=2&page_size=10",
headers={"x-auth-username": "admin"},
)
assert response.status_code == 200
payload = response.json()

View File

@@ -0,0 +1,65 @@
from __future__ import annotations
from types import SimpleNamespace
import pytest
from fastapi import HTTPException
from app.api.v1.endpoints import bootstrap as bootstrap_endpoint
from app.schemas.bootstrap import BootstrapSetupPayload
def completed_settings() -> SimpleNamespace:
return SimpleNamespace(
setup_completed=True,
company_name="X-Financial",
company_code="XF-001",
admin_email="admin@example.com",
web_host="0.0.0.0",
web_port=5273,
app_host="0.0.0.0",
app_port=8000,
postgres_host="postgres.internal",
postgres_port=5432,
postgres_db="x_financial",
postgres_user="postgres-admin",
postgres_password="secret",
redis_url="redis://redis.internal:6379/0",
)
def setup_payload() -> BootstrapSetupPayload:
return BootstrapSetupPayload(
company_name="X-Financial",
company_code="XF-001",
admin_email="admin@example.com",
postgres_host="postgres.internal",
postgres_port=5432,
postgres_db="x_financial",
postgres_user="postgres-admin",
postgres_password="secret",
redis_url="redis://redis.internal:6379/0",
)
def test_completed_bootstrap_state_redacts_infrastructure(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(bootstrap_endpoint, "get_settings", completed_settings)
state = bootstrap_endpoint.get_bootstrap_state()
assert state.initialized is True
assert state.database.host == ""
assert state.database.username == ""
assert state.database.password_configured is True
assert state.redis.url == ""
def test_completed_bootstrap_rejects_anonymous_reconfiguration(
monkeypatch: pytest.MonkeyPatch,
) -> None:
monkeypatch.setattr(bootstrap_endpoint, "get_settings", completed_settings)
with pytest.raises(HTTPException) as exc_info:
bootstrap_endpoint.initialize_bootstrap(setup_payload(), None)
assert exc_info.value.status_code == 403

View File

@@ -4,6 +4,7 @@ from collections.abc import Generator
from datetime import UTC, datetime
from decimal import Decimal
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
@@ -32,6 +33,7 @@ def build_session_factory() -> sessionmaker[Session]:
def build_client() -> tuple[TestClient, sessionmaker[Session]]:
session_factory = build_session_factory()
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()

View File

@@ -4,6 +4,7 @@ from collections.abc import Generator
from datetime import UTC, date, datetime, timedelta
from decimal import Decimal
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
@@ -249,6 +250,7 @@ def test_latest_profile_endpoint_returns_approval_payload() -> None:
seed_profile_data(db)
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()
@@ -303,6 +305,7 @@ def test_current_employee_profile_endpoint_resolves_login_user() -> None:
)
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()
@@ -370,6 +373,7 @@ def test_current_admin_profile_endpoint_returns_account_usage_profile() -> None:
db.commit()
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()
@@ -414,6 +418,7 @@ def test_current_admin_profile_endpoint_uses_online_session_without_agent_runs()
)
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()
@@ -464,6 +469,7 @@ def test_finish_session_endpoint_closes_active_session() -> None:
db.commit()
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()

View File

@@ -137,7 +137,12 @@ def test_legacy_bootstrap_excludes_migration_owned_tables(
table_names = set(inspect(engine).get_table_names())
assert "employees" in table_names
assert {"expense_cases", "expense_case_links", "business_events"}.isdisjoint(table_names)
assert {
"auth_sessions",
"expense_cases",
"expense_case_links",
"business_events",
}.isdisjoint(table_names)
def test_event_write_is_idempotent_for_same_business_operation() -> None:

View File

@@ -3,6 +3,7 @@ from __future__ import annotations
from collections.abc import Generator
from datetime import UTC, datetime
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy.orm import Session
@@ -12,7 +13,9 @@ from app.main import create_app
from app.models.employee import Employee
from app.models.financial_record import ExpenseClaim
from app.schemas.orchestrator import OrchestratorResponse, OrchestratorTraceSummary
from app.services.linked_reimbursement_draft_jobs import clear_linked_reimbursement_draft_jobs_for_tests
from app.services.linked_reimbursement_draft_jobs import (
clear_linked_reimbursement_draft_jobs_for_tests,
)
from app.services.orchestrator import OrchestratorService
from app.test_helpers.db import build_in_memory_session_factory
@@ -53,6 +56,7 @@ def seed_employee_and_application(db: Session) -> None:
def build_client(monkeypatch) -> tuple[TestClient, object]:
session_factory = build_in_memory_session_factory()
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()

View File

@@ -2,6 +2,7 @@ from __future__ import annotations
from collections.abc import Generator
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
@@ -34,6 +35,7 @@ def build_client() -> TestClient:
Base.metadata.create_all(bind=engine)
session_factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()

View File

@@ -2,6 +2,7 @@ from __future__ import annotations
from collections.abc import Generator
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
@@ -11,7 +12,12 @@ from app.api.deps import get_db
from app.core.config import get_settings
from app.db.base import Base
from app.main import create_app
from app.schemas.ocr import OcrRecognizeBatchRead, OcrRecognizeDocumentRead, OcrRecognizeFieldRead, OcrRecognizeLineRead
from app.schemas.ocr import (
OcrRecognizeBatchRead,
OcrRecognizeDocumentRead,
OcrRecognizeFieldRead,
OcrRecognizeLineRead,
)
from app.services.ocr import OcrService
@@ -24,6 +30,7 @@ def build_client() -> TestClient:
Base.metadata.create_all(bind=engine)
session_factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()

View File

@@ -3,13 +3,14 @@ from __future__ import annotations
from collections.abc import Generator
import pytest
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.pool import StaticPool
from app.core.agent_enums import AgentName, AgentRunSource, AgentRunStatus
from app.api.deps import get_db
from app.core.agent_enums import AgentName, AgentRunSource, AgentRunStatus
from app.db.base import Base
from app.schemas.ontology import OntologyParseRequest
from app.services.ontology import LlmOntologyParseResult, SemanticOntologyService
@@ -31,7 +32,8 @@ def build_client() -> tuple[TestClient, sessionmaker[Session]]:
session_factory = build_session_factory()
from app.main import create_app
app = create_app()
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()

View File

@@ -6,6 +6,7 @@ from collections.abc import Generator
from datetime import UTC, date, datetime
from decimal import Decimal
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
@@ -38,6 +39,7 @@ def build_session_factory() -> sessionmaker[Session]:
def build_client() -> tuple[TestClient, sessionmaker[Session]]:
session_factory = build_session_factory()
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()

View File

@@ -5,12 +5,14 @@ from datetime import UTC, datetime
from decimal import Decimal
import pytest
from auth_helpers import install_legacy_header_auth_override
from fastapi import FastAPI
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
from sqlalchemy.pool import StaticPool
from app.algorithem.risk_graph.replay import AlgorithmReplaySetBuilder
from app.api.deps import get_db
from app.api.v1.endpoints.risk_observations import router as risk_observations_router
from app.db.base import Base
@@ -18,7 +20,6 @@ from app.models.employee import Employee
from app.models.financial_record import ExpenseClaim
from app.models.risk_observation import RiskObservation
from app.schemas.risk_observation import RiskObservationFeedbackCreate
from app.algorithem.risk_graph.replay import AlgorithmReplaySetBuilder
from app.services.risk_observations import RiskObservationService
@@ -266,6 +267,7 @@ def _build_client() -> tuple[TestClient, sessionmaker[Session]]:
session_factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
app = FastAPI()
app.include_router(risk_observations_router, prefix="/api/v1")
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()

View File

@@ -2,6 +2,7 @@ from __future__ import annotations
from collections.abc import Generator
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import create_engine, select
from sqlalchemy.orm import Session, sessionmaker
@@ -45,6 +46,7 @@ def build_client() -> tuple[TestClient, sessionmaker[Session]]:
Base.metadata.create_all(bind=engine)
session_factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()

View File

@@ -2,6 +2,7 @@ from __future__ import annotations
from collections.abc import Generator
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
@@ -13,8 +14,8 @@ from app.db.base import Base
from app.main import create_app
from app.models.agent_asset import AgentAsset
from app.schemas.agent_asset import AgentAssetRiskRuleGenerateRequest
from app.services.agent_asset_rule_library import AgentAssetRuleLibraryManager
from app.services.agent_asset_risk_rule_regeneration import AgentAssetRiskRuleRegenerationService
from app.services.agent_asset_rule_library import AgentAssetRuleLibraryManager
from app.services.agent_assets import AgentAssetService
from app.services.risk_rule_generation import RiskRuleGenerationService
@@ -33,6 +34,7 @@ def build_client() -> tuple[TestClient, sessionmaker[Session]]:
Base.metadata.create_all(bind=engine)
session_factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()

View File

@@ -1,5 +1,6 @@
from __future__ import annotations
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from app.main import create_app
@@ -49,11 +50,13 @@ def test_risk_rule_template_catalog_groups_and_dsl_examples() -> None:
def test_risk_rule_template_endpoint_requires_login_and_returns_groups() -> None:
client = TestClient(create_app())
app = create_app()
client = TestClient(app)
unauthorized = client.get("/api/v1/agent-assets/risk-rules/templates")
assert unauthorized.status_code == 401
install_legacy_header_auth_override(app)
response = client.get(
"/api/v1/agent-assets/risk-rules/templates",
headers={"x-auth-username": "finance", "x-auth-role-codes": "finance"},

View File

@@ -3,6 +3,7 @@ from __future__ import annotations
from collections.abc import Generator
from datetime import UTC, datetime
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import create_engine, select
from sqlalchemy.orm import Session, sessionmaker
@@ -30,6 +31,7 @@ def build_session_factory() -> sessionmaker[Session]:
def build_client() -> tuple[TestClient, sessionmaker[Session]]:
session_factory = build_session_factory()
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()

View File

@@ -2,6 +2,7 @@ from __future__ import annotations
from collections.abc import Generator
from auth_helpers import install_legacy_header_auth_override
from fastapi.testclient import TestClient
from sqlalchemy import create_engine
from sqlalchemy.orm import Session, sessionmaker
@@ -23,6 +24,7 @@ def build_client() -> TestClient:
Base.metadata.create_all(bind=engine)
session_factory = sessionmaker(bind=engine, autoflush=False, autocommit=False)
app = create_app()
install_legacy_header_auth_override(app)
def override_db() -> Generator[Session, None, None]:
db = session_factory()