feat(auth): add opaque bearer sessions
This commit is contained in:
53
server/tests/auth_helpers.py
Normal file
53
server/tests/auth_helpers.py
Normal 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(),
|
||||
)
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
148
server/tests/test_auth_session_endpoints.py
Normal file
148
server/tests/test_auth_session_endpoints.py
Normal 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
|
||||
@@ -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()
|
||||
|
||||
65
server/tests/test_bootstrap_security.py
Normal file
65
server/tests/test_bootstrap_security.py
Normal 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
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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"},
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user