From 81c0d2cc747c9e6bf343c6c212ec39aaed0b55a8 Mon Sep 17 00:00:00 2001 From: Wu Clan Date: Mon, 23 Oct 2023 19:30:30 +0800 Subject: [PATCH] Fix pytest interface unit tests (#233) --- README.md | 6 +++--- README.zh-CN.md | 6 +++--- backend/app/api/v1/auth/auth.py | 4 ++-- backend/app/tests/api_v1/test_auth.py | 9 +++++---- backend/app/tests/conftest.py | 7 ++++++- backend/app/tests/utils/db_mysql.py | 13 +++---------- backend/app/tests/utils/get_headers.py | 6 +++--- backend/app/utils/request_parse.py | 3 +++ 8 files changed, 28 insertions(+), 26 deletions(-) diff --git a/README.md b/README.md index fcc538d9..09e016d7 100644 --- a/README.md +++ b/README.md @@ -174,14 +174,14 @@ Click [fastapi_best_architecture_ui](https://github.com/fastapi-practices/fastap Execute unittests via pytest 1. Create the test database `fba_test`, select utf8mb4 encoding -2. Enter the app directory +2. Using `backend/sql/create_tables.sql` file to create database tables +3. Initialize the test data using the `backend/sql/init_pytest_data.sql` file +4. Enter the app directory ```shell cd backend/app/ ``` -3. Using `backend/sql/create_tables.sql` file to create database tables -4. Initialize the test data using the `backend/sql/init_pytest_data.sql` file 5. Execute the test command ```shell diff --git a/README.zh-CN.md b/README.zh-CN.md index 65e33b34..f843d27c 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -168,14 +168,14 @@ TODO: 通过 pytest 执行单元测试 1. 创建测试数据库 `fba_test`,选择 utf8mb4 编码 -2. 进入app目录 +2. 使用 `backend/sql/create_tables.sql` 文件创建数据库表 +3. 使用 `backend/sql/init_pytest_data.sql` 文件初始化测试数据 +4. 进入app目录 ```shell cd backend/app/ ``` -3. 使用 `backend/sql/create_tables.sql` 文件创建数据库表 -4. 使用 `backend/sql/init_pytest_data.sql` 文件初始化测试数据 5. 执行测试命令 ```shell diff --git a/backend/app/api/v1/auth/auth.py b/backend/app/api/v1/auth/auth.py index 2c69dc4b..1f60513a 100644 --- a/backend/app/api/v1/auth/auth.py +++ b/backend/app/api/v1/auth/auth.py @@ -19,7 +19,7 @@ router = APIRouter() @router.post('/swagger_login', summary='swagger 表单登录', description='form 格式登录,仅用于 swagger 文档调试接口') async def swagger_user_login(form_data: OAuth2PasswordRequestForm = Depends()) -> GetSwaggerToken: token, user = await AuthService().swagger_login(form_data=form_data) - return GetSwaggerToken(access_token=token, user=user) + return GetSwaggerToken(access_token=token, user=user) # type: ignore @router.post( @@ -37,7 +37,7 @@ async def user_login(request: Request, obj: AuthLogin, background_tasks: Backgro refresh_token=refresh_token, access_token_expire_time=access_expire, refresh_token_expire_time=refresh_expire, - user=user, + user=user, # type: ignore ) return await response_base.success(data=data) diff --git a/backend/app/tests/api_v1/test_auth.py b/backend/app/tests/api_v1/test_auth.py index 345a8b7a..7d48766f 100644 --- a/backend/app/tests/api_v1/test_auth.py +++ b/backend/app/tests/api_v1/test_auth.py @@ -3,16 +3,17 @@ from starlette.testclient import TestClient from backend.app.core.conf import settings +from backend.app.tests.conftest import PYTEST_USERNAME, PYTEST_PASSWORD def test_login(client: TestClient) -> None: data = { - 'username': 'admin', - 'password': '123456', + 'username': PYTEST_USERNAME, + 'password': PYTEST_PASSWORD, } - response = client.post(f'{settings.API_V1_STR}/auth/login', json=data) + response = client.post(f'{settings.API_V1_STR}/auth/swagger_login', data=data) assert response.status_code == 200 - assert response.json()['data']['access_token_type'] == 'Bearer' + assert response.json()['token_type'] == 'Bearer' def test_logout(client: TestClient, token_headers: dict[str, str]) -> None: diff --git a/backend/app/tests/conftest.py b/backend/app/tests/conftest.py index 74f646d7..ae6985fc 100644 --- a/backend/app/tests/conftest.py +++ b/backend/app/tests/conftest.py @@ -17,6 +17,11 @@ from backend.app.tests.utils.db_mysql import override_get_db app.dependency_overrides[get_db] = override_get_db +# Test user +PYTEST_USERNAME = 'admin' +PYTEST_PASSWORD = '123456' + + @pytest.fixture(scope='module') def client() -> Generator: with TestClient(app) as c: @@ -25,4 +30,4 @@ def client() -> Generator: @pytest.fixture(scope='module') def token_headers(client: TestClient) -> Dict[str, str]: - return get_token_headers(client=client, username='admin', password='123456') + return get_token_headers(client=client, username=PYTEST_USERNAME, password=PYTEST_PASSWORD) diff --git a/backend/app/tests/utils/db_mysql.py b/backend/app/tests/utils/db_mysql.py index 280d981d..90720ff6 100644 --- a/backend/app/tests/utils/db_mysql.py +++ b/backend/app/tests/utils/db_mysql.py @@ -3,22 +3,21 @@ from sqlalchemy.ext.asyncio import AsyncSession from backend.app.core.conf import settings -from backend.app.models.base import MappedBase from backend.app.database.db_mysql import create_engine_and_session TEST_DB_DATABASE = settings.DB_DATABASE + '_test' -SQLALCHEMY_DATABASE_URL = ( +TEST_SQLALCHEMY_DATABASE_URL = ( f'mysql+asyncmy://{settings.DB_USER}:{settings.DB_PASSWORD}@{settings.DB_HOST}:' f'{settings.DB_PORT}/{TEST_DB_DATABASE}?charset={settings.DB_CHARSET}' ) -async_engine, async_db_session = create_engine_and_session(SQLALCHEMY_DATABASE_URL) +test_async_engine, test_async_db_session = create_engine_and_session(TEST_SQLALCHEMY_DATABASE_URL) async def override_get_db() -> AsyncSession: """session 生成器""" - session = async_db_session() + session = test_async_db_session() try: yield session except Exception as se: @@ -26,9 +25,3 @@ async def override_get_db() -> AsyncSession: raise se finally: await session.close() - - -async def create_table(): - """创建数据库表""" - async with async_engine.begin() as coon: - await coon.run_sync(MappedBase.metadata.create_all) diff --git a/backend/app/tests/utils/get_headers.py b/backend/app/tests/utils/get_headers.py index 8ad13686..5e582cb2 100644 --- a/backend/app/tests/utils/get_headers.py +++ b/backend/app/tests/utils/get_headers.py @@ -12,8 +12,8 @@ def get_token_headers(client: TestClient, username: str, password: str) -> Dict[ 'username': username, 'password': password, } - response = client.post(f'{settings.API_V1_STR}/auth/login', json=data) - token_type = response.json()['data']['access_token_type'] - access_token = response.json()['data']['access_token'] + response = client.post(f'{settings.API_V1_STR}/auth/swagger_login', data=data) + token_type = response.json()['token_type'] + access_token = response.json()['access_token'] headers = {'Authorization': f'{token_type} {access_token}'} return headers diff --git a/backend/app/utils/request_parse.py b/backend/app/utils/request_parse.py index c77a98da..31e04694 100644 --- a/backend/app/utils/request_parse.py +++ b/backend/app/utils/request_parse.py @@ -24,6 +24,9 @@ def get_request_ip(request: Request) -> str: ip = forwarded.split(',')[0] else: ip = request.client.host + # 忽略 pytest + if ip == 'testclient': + ip = '127.0.0.1' return ip