from __future__ import annotations from collections.abc import Generator from fastapi.testclient import TestClient from sqlalchemy import create_engine, func, select from sqlalchemy.orm import Session, sessionmaker from sqlalchemy.pool import StaticPool from app.api.deps import get_db from app.db.base import Base from app.main import create_app from app.models.employee import Employee from app.models.organization import OrganizationUnit from app.schemas.settings import SettingsWrite from app.services.settings import SettingsService from app.services.tenant_registry import TenantRegistryService 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: TenantRegistryService(db).ensure_builtin() settings_service = SettingsService(db) payload = settings_service.get_settings_snapshot().model_dump() payload["adminForm"]["adminAccount"] = "admin" payload["adminForm"]["newPassword"] = "admin" payload["adminForm"]["confirmPassword"] = "admin" 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 test_admin_selected_default_tenant_reads_default_employee_directory() -> None: client, session_factory = build_client() login_response = client.post( "/api/v1/auth/login", json={"username": "admin", "password": "admin", "tenantId": "default"}, ) assert login_response.status_code == 200 login_payload = login_response.json() assert login_payload["user"]["tenantId"] == "default" headers = {"Authorization": f"Bearer {login_payload['accessToken']}"} meta_response = client.get("/api/v1/employees/meta", headers=headers) list_response = client.get("/api/v1/employees", headers=headers) assert meta_response.status_code == 200 assert list_response.status_code == 200 assert meta_response.json()["totalEmployees"] == 30 assert len(list_response.json()) == 30 with session_factory() as db: default_employee_count = db.scalar( select(func.count()) .select_from(Employee) .where(Employee.tenant_id == "default") ) platform_employee_count = db.scalar( select(func.count()) .select_from(Employee) .where(Employee.tenant_id == "platform") ) default_unit_count = db.scalar( select(func.count()) .select_from(OrganizationUnit) .where(OrganizationUnit.tenant_id == "default") ) platform_unit_count = db.scalar( select(func.count()) .select_from(OrganizationUnit) .where(OrganizationUnit.tenant_id == "platform") ) assert default_employee_count == 30 assert platform_employee_count == 0 assert default_unit_count == 7 assert platform_unit_count == 0 def test_platform_admin_cannot_seed_platform_employee_directory() -> None: client, session_factory = build_client() login_response = client.post( "/api/v1/auth/login", json={"username": "admin", "password": "admin"}, ) assert login_response.status_code == 200 login_payload = login_response.json() assert login_payload["user"]["tenantId"] == "platform" headers = {"Authorization": f"Bearer {login_payload['accessToken']}"} meta_response = client.get("/api/v1/employees/meta", headers=headers) list_response = client.get("/api/v1/employees", headers=headers) assert meta_response.status_code == 403 assert list_response.status_code == 403 assert meta_response.json()["detail"] == "员工目录只能在企业工作域内使用。" assert list_response.json()["detail"] == "员工目录只能在企业工作域内使用。" with session_factory() as db: platform_employee_count = db.scalar( select(func.count()) .select_from(Employee) .where(Employee.tenant_id == "platform") ) platform_unit_count = db.scalar( select(func.count()) .select_from(OrganizationUnit) .where(OrganizationUnit.tenant_id == "platform") ) assert platform_employee_count == 0 assert platform_unit_count == 0