132 lines
4.6 KiB
Python
132 lines
4.6 KiB
Python
|
|
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
|