2be0778ec7
- 后端:auth 路由(register/login/refresh/me)+ JWT + bcrypt 密码哈希 - 依赖注入:get_current_user + require_role 角色权限校验 - 跨数据库兼容:JSONBType(PG 用 JSONB,SQLite 用 JSON) - 测试:11 个认证测试 + 4 个健康检查测试 = 15 passed
203 lines
6.5 KiB
Python
203 lines
6.5 KiB
Python
"""认证流程测试。
|
|
|
|
测试注册、登录、获取当前用户、刷新 token 的完整流程。
|
|
"""
|
|
|
|
import pytest
|
|
from fastapi.testclient import TestClient
|
|
from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine, async_sessionmaker
|
|
from app.core.database import Base, get_db
|
|
from app.main import app
|
|
|
|
# 使用 SQLite 内存数据库做测试
|
|
TEST_DATABASE_URL = "sqlite+aiosqlite:///file::memory:?cache=shared&uri=true"
|
|
|
|
test_engine = create_async_engine(TEST_DATABASE_URL, echo=False)
|
|
test_session_factory = async_sessionmaker(test_engine, class_=AsyncSession, expire_on_commit=False)
|
|
|
|
|
|
async def override_get_db():
|
|
"""测试用数据库 session。"""
|
|
async with test_session_factory() as session:
|
|
try:
|
|
yield session
|
|
await session.commit()
|
|
except Exception:
|
|
await session.rollback()
|
|
raise
|
|
|
|
|
|
app.dependency_overrides[get_db] = override_get_db
|
|
|
|
|
|
@pytest.fixture(scope="module", autouse=True)
|
|
async def setup_db():
|
|
"""创建测试数据库表。"""
|
|
async with test_engine.begin() as conn:
|
|
await conn.run_sync(Base.metadata.create_all)
|
|
yield
|
|
async with test_engine.begin() as conn:
|
|
await conn.run_sync(Base.metadata.drop_all)
|
|
|
|
|
|
@pytest.fixture
|
|
def client():
|
|
"""创建测试客户端。"""
|
|
return TestClient(app)
|
|
|
|
|
|
class TestRegister:
|
|
"""注册测试。"""
|
|
|
|
def test_register_success(self, client: TestClient):
|
|
"""RED→GREEN: 正常注册应返回 token。"""
|
|
response = client.post(
|
|
"/api/v1/auth/register",
|
|
json={
|
|
"email": "test@example.com",
|
|
"password": "password123",
|
|
"name": "测试用户",
|
|
"tenant_name": "测试投资机构",
|
|
"role": "founder",
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["code"] == 0
|
|
assert data["data"]["access_token"] is not None
|
|
assert data["data"]["refresh_token"] is not None
|
|
assert data["data"]["token_type"] == "bearer"
|
|
|
|
def test_register_duplicate_email(self, client: TestClient):
|
|
"""RED→GREEN: 重复邮箱注册应返回 409。"""
|
|
response = client.post(
|
|
"/api/v1/auth/register",
|
|
json={
|
|
"email": "test@example.com",
|
|
"password": "password123",
|
|
"name": "重复用户",
|
|
"tenant_name": "另一个机构",
|
|
"role": "founder",
|
|
},
|
|
)
|
|
assert response.status_code == 409
|
|
|
|
def test_register_short_password(self, client: TestClient):
|
|
"""RED→GREEN: 密码太短应返回 422。"""
|
|
response = client.post(
|
|
"/api/v1/auth/register",
|
|
json={
|
|
"email": "short@example.com",
|
|
"password": "123",
|
|
"name": "短密码",
|
|
"tenant_name": "测试机构",
|
|
"role": "founder",
|
|
},
|
|
)
|
|
assert response.status_code == 422
|
|
|
|
|
|
class TestLogin:
|
|
"""登录测试。"""
|
|
|
|
def test_login_success(self, client: TestClient):
|
|
"""RED→GREEN: 正确邮箱密码登录成功。"""
|
|
response = client.post(
|
|
"/api/v1/auth/login",
|
|
json={
|
|
"email": "test@example.com",
|
|
"password": "password123",
|
|
},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["code"] == 0
|
|
assert data["data"]["access_token"] is not None
|
|
|
|
def test_login_wrong_password(self, client: TestClient):
|
|
"""RED→GREEN: 错误密码应返回 401。"""
|
|
response = client.post(
|
|
"/api/v1/auth/login",
|
|
json={
|
|
"email": "test@example.com",
|
|
"password": "wrongpassword",
|
|
},
|
|
)
|
|
assert response.status_code == 401
|
|
|
|
def test_login_nonexistent_email(self, client: TestClient):
|
|
"""RED→GREEN: 不存在的邮箱应返回 401。"""
|
|
response = client.post(
|
|
"/api/v1/auth/login",
|
|
json={
|
|
"email": "nonexistent@example.com",
|
|
"password": "password123",
|
|
},
|
|
)
|
|
assert response.status_code == 401
|
|
|
|
|
|
class TestGetMe:
|
|
"""获取当前用户信息测试。"""
|
|
|
|
def test_get_me_with_valid_token(self, client: TestClient):
|
|
"""RED→GREEN: 有效 token 应返回用户信息。"""
|
|
# 先登录获取 token
|
|
login_resp = client.post(
|
|
"/api/v1/auth/login",
|
|
json={"email": "test@example.com", "password": "password123"},
|
|
)
|
|
token = login_resp.json()["data"]["access_token"]
|
|
|
|
response = client.get(
|
|
"/api/v1/auth/me",
|
|
headers={"Authorization": f"Bearer {token}"},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["code"] == 0
|
|
assert data["data"]["email"] == "test@example.com"
|
|
assert data["data"]["name"] == "测试用户"
|
|
|
|
def test_get_me_without_token(self, client: TestClient):
|
|
"""RED→GREEN: 无 token 应返回 401。"""
|
|
response = client.get("/api/v1/auth/me")
|
|
assert response.status_code == 401
|
|
|
|
def test_get_me_with_invalid_token(self, client: TestClient):
|
|
"""RED→GREEN: 无效 token 应返回 401。"""
|
|
response = client.get(
|
|
"/api/v1/auth/me",
|
|
headers={"Authorization": "Bearer invalid-token-string"},
|
|
)
|
|
assert response.status_code == 401
|
|
|
|
|
|
class TestRefreshToken:
|
|
"""刷新 token 测试。"""
|
|
|
|
def test_refresh_success(self, client: TestClient):
|
|
"""RED→GREEN: 有效 refresh token 应返回新 token。"""
|
|
login_resp = client.post(
|
|
"/api/v1/auth/login",
|
|
json={"email": "test@example.com", "password": "password123"},
|
|
)
|
|
refresh_token = login_resp.json()["data"]["refresh_token"]
|
|
|
|
response = client.post(
|
|
"/api/v1/auth/refresh",
|
|
json={"refresh_token": refresh_token},
|
|
)
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["code"] == 0
|
|
assert data["data"]["access_token"] is not None
|
|
|
|
def test_refresh_with_invalid_token(self, client: TestClient):
|
|
"""RED→GREEN: 无效 refresh token 应返回 401。"""
|
|
response = client.post(
|
|
"/api/v1/auth/refresh",
|
|
json={"refresh_token": "invalid-token"},
|
|
)
|
|
assert response.status_code == 401
|