From 4e4c6fbe95cb3320ffcd948efd4dac7649736e37 Mon Sep 17 00:00:00 2001 From: Wu Clan Date: Sat, 27 May 2023 22:55:25 +0800 Subject: [PATCH] add login logs (#76) * simplify crud method naming * update get_user_list to get_select * add sign in logs * Perform pre-commit fix * Encapsulated request ip address resolution * Delete login log records for uncertain exceptions * Add login log deletion interface * Add login logging to background tasks * update the user agent parse --- backend/app/api/routers.py | 4 +- backend/app/api/v1/auth/auth.py | 11 ++- backend/app/api/v1/login_log.py | 38 ++++++++ backend/app/api/v1/user.py | 12 +-- .../app/common/exception/exception_handler.py | 2 +- backend/app/core/conf.py | 3 + backend/app/crud/crud_login_log.py | 30 +++++++ backend/app/crud/crud_user.py | 9 +- backend/app/models/__init__.py | 1 + backend/app/models/sys_login_log.py | 26 ++++++ backend/app/schemas/login_log.py | 33 +++++++ backend/app/schemas/user.py | 2 +- backend/app/services/login_log_service.py | 59 ++++++++++++ backend/app/services/user_service.py | 90 ++++++++++++------- backend/app/utils/request_parse.py | 37 ++++++++ backend/app/utils/serializers.py | 13 ++- requirements.txt | 1 + 17 files changed, 318 insertions(+), 53 deletions(-) create mode 100644 backend/app/api/v1/login_log.py create mode 100644 backend/app/crud/crud_login_log.py create mode 100644 backend/app/models/sys_login_log.py create mode 100644 backend/app/schemas/login_log.py create mode 100644 backend/app/services/login_log_service.py create mode 100644 backend/app/utils/request_parse.py diff --git a/backend/app/api/routers.py b/backend/app/api/routers.py index 9e74eb30..f5e3ff63 100644 --- a/backend/app/api/routers.py +++ b/backend/app/api/routers.py @@ -9,8 +9,9 @@ from backend.app.api.v1.dept import router as dept_router from backend.app.api.v1.role import router as role_router from backend.app.api.v1.menu import router as menu_router from backend.app.api.v1.api import router as api_router -from backend.app.api.v1.task_demo import router as task_demo_router from backend.app.api.v1.config import router as config_router +from backend.app.api.v1.login_log import router as login_log_router +from backend.app.api.v1.task_demo import router as task_demo_router v1 = APIRouter(prefix='/v1') @@ -22,4 +23,5 @@ v1.include_router(role_router, prefix='/roles', tags=['角色管理']) v1.include_router(menu_router, prefix='/menus', tags=['菜单管理']) v1.include_router(api_router, prefix='/apis', tags=['API管理']) v1.include_router(config_router, prefix='/configs', tags=['系统配置']) +v1.include_router(login_log_router, prefix='/login_logs', tags=['登录日志管理']) v1.include_router(task_demo_router, prefix='/tasks', tags=['任务管理']) diff --git a/backend/app/api/v1/auth/auth.py b/backend/app/api/v1/auth/auth.py index df945978..1ddda9b9 100644 --- a/backend/app/api/v1/auth/auth.py +++ b/backend/app/api/v1/auth/auth.py @@ -3,6 +3,7 @@ from fastapi import APIRouter, Depends, Request from fastapi.security import OAuth2PasswordRequestForm from fastapi_limiter.depends import RateLimiter +from starlette.background import BackgroundTasks from backend.app.common.jwt import DependsUser, get_token, jwt_decode, CurrentJwtAuth from backend.app.common.response.response_schema import response_base @@ -15,7 +16,7 @@ router = APIRouter() @router.post('/swagger_login', summary='swagger 表单登录', description='form 格式登录,仅用于 swagger 文档调试接口') async def swagger_user_login(form_data: OAuth2PasswordRequestForm = Depends()) -> SwaggerToken: - token, user = await UserService.swagger_login(form_data) + token, user = await UserService().swagger_login(form_data) return SwaggerToken(access_token=token, user=user) @@ -25,8 +26,10 @@ async def swagger_user_login(form_data: OAuth2PasswordRequestForm = Depends()) - description='json 格式登录, 仅支持在第三方api工具调试接口, 例如: postman', dependencies=[Depends(RateLimiter(times=5, minutes=15))], ) -async def user_login(obj: Auth): - access_token, refresh_token, access_expire, refresh_expire, user = await UserService.login(obj) +async def user_login(request: Request, obj: Auth, background_tasks: BackgroundTasks): + access_token, refresh_token, access_expire, refresh_expire, user = await UserService().login( + request=request, obj=obj, background_tasks=background_tasks + ) data = LoginToken( access_token=access_token, refresh_token=refresh_token, @@ -41,7 +44,7 @@ async def user_login(obj: Auth): async def get_refresh_token(request: Request, custom_time: RefreshTokenTime): token = get_token(request) user_id, _ = jwt_decode(token) - refresh_token, refresh_expire = await UserService.refresh_token(user_id, custom_time) + refresh_token, refresh_expire = await UserService.refresh_token(user_id=user_id, custom_time=custom_time) data = RefreshToken(refresh_token=refresh_token, refresh_token_expire_time=refresh_expire) return response_base.success(data=data) diff --git a/backend/app/api/v1/login_log.py b/backend/app/api/v1/login_log.py new file mode 100644 index 00000000..5b44dc57 --- /dev/null +++ b/backend/app/api/v1/login_log.py @@ -0,0 +1,38 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +from typing import Annotated + +from fastapi import APIRouter, Query + +from backend.app.common.casbin_rbac import DependsRBAC +from backend.app.common.jwt import DependsUser +from backend.app.common.pagination import paging_data, PageDepends +from backend.app.common.response.response_schema import response_base +from backend.app.database.db_mysql import CurrentSession +from backend.app.schemas.login_log import GetAllLoginLog +from backend.app.services.login_log_service import LoginLogService + +router = APIRouter() + + +@router.get('', summary='获取所有登录日志', dependencies=[DependsUser, PageDepends]) +async def get_all_login_logs(db: CurrentSession): + log_select = await LoginLogService.get_select() + page_data = await paging_data(db, log_select, GetAllLoginLog) + return response_base.success(data=page_data) + + +@router.delete('', summary='(批量)删除登录日志', dependencies=[DependsRBAC]) +async def delete_login_log(pk: Annotated[list[int], Query(...)]): + count = await LoginLogService.delete(pk) + if count > 0: + return response_base.success() + return response_base.fail() + + +@router.delete('/all', summary='清空登录日志', dependencies=[DependsRBAC]) +async def delete_all_login_logs(): + count = await LoginLogService.delete_all() + if count > 0: + return response_base.success() + return response_base.fail() diff --git a/backend/app/api/v1/user.py b/backend/app/api/v1/user.py index b08033b9..f202a01f 100644 --- a/backend/app/api/v1/user.py +++ b/backend/app/api/v1/user.py @@ -6,7 +6,7 @@ from backend.app.common.jwt import DependsUser, CurrentUser, DependsSuperUser from backend.app.common.pagination import paging_data, PageDepends from backend.app.common.response.response_schema import response_base from backend.app.database.db_mysql import CurrentSession -from backend.app.schemas.user import CreateUser, GetUserInfo, ResetPassword, UpdateUser, Avatar +from backend.app.schemas.user import CreateUser, GetAllUserInfo, ResetPassword, UpdateUser, Avatar from backend.app.services.user_service import UserService from backend.app.utils.serializers import select_to_json @@ -21,14 +21,16 @@ async def user_register(obj: CreateUser): @router.post('/password/reset', summary='密码重置') async def password_reset(obj: ResetPassword): - await UserService.pwd_reset(obj) - return response_base.success() + count = await UserService.pwd_reset(obj) + if count > 0: + return response_base.success() + return response_base.fail() @router.get('/{username}', summary='查看用户信息', dependencies=[DependsUser]) async def userinfo(username: str): current_user = await UserService.get_userinfo(username) - data = GetUserInfo(**select_to_json(current_user)) + data = GetAllUserInfo(**select_to_json(current_user)) return response_base.success(data=data) @@ -51,7 +53,7 @@ async def update_avatar(username: str, avatar: Avatar, current_user: CurrentUser @router.get('', summary='获取所有用户', dependencies=[DependsUser, PageDepends]) async def get_all_users(db: CurrentSession): user_select = await UserService.get_select() - page_data = await paging_data(db, user_select, GetUserInfo) + page_data = await paging_data(db, user_select, GetAllUserInfo) return response_base.success(data=page_data) diff --git a/backend/app/common/exception/exception_handler.py b/backend/app/common/exception/exception_handler.py index 3c47e2c3..5c0fa774 100644 --- a/backend/app/common/exception/exception_handler.py +++ b/backend/app/common/exception/exception_handler.py @@ -109,6 +109,6 @@ def register_exception(app: FastAPI): return JSONResponse( status_code=500, content=response_base.fail(code=500, msg=str(exc)) - if settings.UVICORN_RELOAD + if settings.ENVIRONMENT != 'pro' else response_base.fail(code=500, msg='Internal Server Error'), ) diff --git a/backend/app/core/conf.py b/backend/app/core/conf.py index 91ad89bd..4fe894fc 100644 --- a/backend/app/core/conf.py +++ b/backend/app/core/conf.py @@ -54,6 +54,9 @@ class Settings(BaseSettings): # Static Server STATIC_FILES: bool = False + # Location Parse + LOCATION_PARSE: bool = True # 将会导致登录延时,建议关闭,有条件自行使用第三方离线数据库 + # MySQL DB_ECHO: bool = False DB_DATABASE: str = 'fba' diff --git a/backend/app/crud/crud_login_log.py b/backend/app/crud/crud_login_log.py new file mode 100644 index 00000000..b257dbae --- /dev/null +++ b/backend/app/crud/crud_login_log.py @@ -0,0 +1,30 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +from typing import NoReturn + +from sqlalchemy import Select, select, desc, delete +from sqlalchemy.ext.asyncio import AsyncSession + +from backend.app.crud.base import CRUDBase +from backend.app.models import LoginLog +from backend.app.schemas.login_log import CreateLoginLog, UpdateLoginLog + + +class CRUDLoginLog(CRUDBase[LoginLog, CreateLoginLog, UpdateLoginLog]): + async def get_all(self) -> Select: + return select(self.model).order_by(desc(self.model.create_time)) + + async def create(self, db: AsyncSession, obj_in: CreateLoginLog) -> NoReturn: + await self.create_(db, obj_in) + await db.commit() + + async def delete(self, db: AsyncSession, pk: list[int]) -> int: + logs = await db.execute(delete(self.model).where(self.model.id.in_(pk))) + return logs.rowcount + + async def delete_all(self, db: AsyncSession) -> int: + logs = await db.execute(delete(self.model)) + return logs.rowcount + + +LoginLogDao: CRUDLoginLog = CRUDLoginLog(LoginLog) diff --git a/backend/app/crud/crud_user.py b/backend/app/crud/crud_user.py index ee381c89..6f03bdb0 100644 --- a/backend/app/crud/crud_user.py +++ b/backend/app/crud/crud_user.py @@ -1,8 +1,9 @@ #!/usr/bin/env python3 # -*- coding: utf-8 -*- +from datetime import datetime from typing import NoReturn -from sqlalchemy import func, select, update, desc +from sqlalchemy import select, update, desc from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import selectinload from sqlalchemy.sql import Select @@ -21,8 +22,8 @@ class CRUDUser(CRUDBase[User, CreateUser, UpdateUser]): user = await db.execute(select(self.model).where(self.model.username == username)) return user.scalars().first() - async def update_login_time(self, db: AsyncSession, username: str) -> int: - user = await db.execute(update(self.model).where(self.model.username == username).values(last_login=func.now())) + async def update_login_time(self, db: AsyncSession, username: str, login_time: datetime) -> int: + user = await db.execute(update(self.model).where(self.model.username == username).values(last_login=login_time)) return user.rowcount async def create(self, db: AsyncSession, create: CreateUser) -> NoReturn: @@ -53,7 +54,7 @@ class CRUDUser(CRUDBase[User, CreateUser, UpdateUser]): return user.rowcount async def delete(self, db: AsyncSession, user_id: int) -> int: - return await super().delete_(db, user_id) + return await self.delete_(db, user_id) async def check_email(self, db: AsyncSession, email: str) -> User | None: mail = await db.execute(select(self.model).where(self.model.email == email)) diff --git a/backend/app/models/__init__.py b/backend/app/models/__init__.py index c63e1c42..9776c56e 100644 --- a/backend/app/models/__init__.py +++ b/backend/app/models/__init__.py @@ -11,3 +11,4 @@ from backend.app.models.sys_dept import Dept from backend.app.models.sys_menu import Menu from backend.app.models.sys_role import Role from backend.app.models.sys_user import User +from backend.app.models.sys_login_log import LoginLog diff --git a/backend/app/models/sys_login_log.py b/backend/app/models/sys_login_log.py new file mode 100644 index 00000000..b58119c8 --- /dev/null +++ b/backend/app/models/sys_login_log.py @@ -0,0 +1,26 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +from datetime import datetime + +from sqlalchemy import String, func +from sqlalchemy.orm import Mapped, mapped_column + +from backend.app.database.base_class import DataClassBase, id_key + + +class LoginLog(DataClassBase): + """登录日志表""" + + __tablename__ = 'sys_login_log' + + id: Mapped[id_key] = mapped_column(init=False) + user_uuid: Mapped[str] = mapped_column(String(50), nullable=False, comment='用户UUID') + username: Mapped[str] = mapped_column(String(20), nullable=False, comment='用户名') + status: Mapped[int] = mapped_column(insert_default=0, comment='登录状态(0失败 1成功)') + ipaddr: Mapped[str] = mapped_column(String(50), nullable=False, comment='登录IP地址') + location: Mapped[str] = mapped_column(String(255), nullable=False, comment='归属地') + browser: Mapped[str] = mapped_column(String(255), nullable=False, comment='浏览器') + os: Mapped[str] = mapped_column(String(255), nullable=False, comment='操作系统') + msg: Mapped[str] = mapped_column(String(255), nullable=False, comment='提示消息') + login_time: Mapped[datetime] = mapped_column(nullable=False, comment='登录时间') + create_time: Mapped[datetime] = mapped_column(init=False, default=func.now(), comment='创建时间') diff --git a/backend/app/schemas/login_log.py b/backend/app/schemas/login_log.py new file mode 100644 index 00000000..9e7f2946 --- /dev/null +++ b/backend/app/schemas/login_log.py @@ -0,0 +1,33 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +from datetime import datetime + +from pydantic import BaseModel + + +class LoginLogBase(BaseModel): + user_uuid: str + username: str + status: int + ipaddr: str + location: str + browser: str + os: str + msg: str + login_time: datetime + + +class CreateLoginLog(LoginLogBase): + pass + + +class UpdateLoginLog(LoginLogBase): + pass + + +class GetAllLoginLog(LoginLogBase): + id: int + create_time: datetime + + class Config: + orm_mode = True diff --git a/backend/app/schemas/user.py b/backend/app/schemas/user.py index 28a6bfe2..22b581f1 100644 --- a/backend/app/schemas/user.py +++ b/backend/app/schemas/user.py @@ -49,7 +49,7 @@ class GetUserInfoNoRelation(_UserInfoBase): orm_mode = True -class GetUserInfo(GetUserInfoNoRelation): +class GetAllUserInfo(GetUserInfoNoRelation): dept: GetAllDept | None = None roles: list[GetAllRole] diff --git a/backend/app/services/login_log_service.py b/backend/app/services/login_log_service.py new file mode 100644 index 00000000..04697fce --- /dev/null +++ b/backend/app/services/login_log_service.py @@ -0,0 +1,59 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +from datetime import datetime +from typing import NoReturn + +from fastapi import Request +from sqlalchemy import Select +from sqlalchemy.ext.asyncio import AsyncSession +from user_agents import parse + +from backend.app.common.log import log +from backend.app.core.conf import settings +from backend.app.crud.crud_login_log import LoginLogDao +from backend.app.database.db_mysql import async_db_session +from backend.app.models import User +from backend.app.schemas.login_log import CreateLoginLog +from backend.app.utils import request_parse + + +class LoginLogService: + @staticmethod + async def get_select() -> Select: + return await LoginLogDao.get_all() + + @staticmethod + async def create( + *, db: AsyncSession, request: Request, user: User, login_time: datetime, status: int, msg: str + ) -> NoReturn: + try: + ip = await request_parse.get_request_ip(request) + user_agent = request.headers.get('User-Agent') + _, os_info, browser = str(parse(user_agent)).replace(' ', '').split('/') + location = await request_parse.get_location(ip, user_agent) if settings.LOCATION_PARSE else '未知' + obj_in = CreateLoginLog( + user_uuid=user.user_uuid, + username=user.username, + status=status, + ipaddr=ip, + location=location, + browser=browser, + os=os_info, + msg=msg, + login_time=login_time, + ) + await LoginLogDao.create(db, obj_in) + except Exception as e: + log.error(f'登录日志创建失败: {e}') + + @staticmethod + async def delete(pk: list[int]) -> int: + async with async_db_session.begin() as db: + count = await LoginLogDao.delete(db, pk) + return count + + @staticmethod + async def delete_all() -> int: + async with async_db_session.begin() as db: + count = await LoginLogDao.delete_all(db) + return count diff --git a/backend/app/services/user_service.py b/backend/app/services/user_service.py index 0de075dd..ff6fcb1f 100644 --- a/backend/app/services/user_service.py +++ b/backend/app/services/user_service.py @@ -1,7 +1,14 @@ #!/usr/bin/env python3 # -*- coding: utf-8 -*- +from datetime import datetime +from typing import NoReturn + from email_validator import validate_email, EmailNotValidError +from fastapi import Request from fastapi.security import OAuth2PasswordRequestForm +from pydantic.datetime_parse import parse_datetime +from sqlalchemy import Select +from starlette.background import BackgroundTasks from backend.app.common import jwt from backend.app.common.exception import errors @@ -14,12 +21,14 @@ from backend.app.database.db_mysql import async_db_session from backend.app.models import User from backend.app.schemas.token import RefreshTokenTime from backend.app.schemas.user import CreateUser, ResetPassword, UpdateUser, Avatar, Auth +from backend.app.services.login_log_service import LoginLogService from backend.app.utils import re_verify class UserService: - @staticmethod - async def swagger_login(form_data: OAuth2PasswordRequestForm): + login_time = parse_datetime(datetime.utcnow()) + + async def swagger_login(self, form_data: OAuth2PasswordRequestForm): async with async_db_session() as db: current_user = await UserDao.get_by_username(db, form_data.username) if not current_user: @@ -29,7 +38,7 @@ class UserService: elif not current_user.is_active: raise errors.AuthorizationError(msg='用户已锁定, 登陆失败') # 更新登陆时间 - await UserDao.update_login_time(db, form_data.username) + await UserDao.update_login_time(db, form_data.username, self.login_time) # 查询用户角色 user_role_ids = await UserDao.get_role_ids(db, current_user.id) # 获取最新用户信息 @@ -38,27 +47,42 @@ class UserService: access_token, _ = await jwt.create_access_token(str(user.id), role_ids=user_role_ids) return access_token, user - @staticmethod - async def login(obj: Auth): + async def login(self, *, request: Request, obj: Auth, background_tasks: BackgroundTasks): async with async_db_session() as db: - current_user = await UserDao.get_by_username(db, obj.username) - if not current_user: - raise errors.NotFoundError(msg='用户不存在') - elif not jwt.password_verify(obj.password, current_user.password): - raise errors.AuthorizationError(msg='密码错误') - elif not current_user.is_active: - raise errors.AuthorizationError(msg='用户已锁定, 登陆失败') - await UserDao.update_login_time(db, obj.username) - user_role_ids = await UserDao.get_role_ids(db, current_user.id) - user = await UserDao.get(db, current_user.id) - access_token, access_token_expire_time = await jwt.create_access_token(str(user.id), role_ids=user_role_ids) - refresh_token, refresh_token_expire_time = await jwt.create_refresh_token( - str(user.id), access_token_expire_time, role_ids=user_role_ids - ) - return access_token, refresh_token, access_token_expire_time, refresh_token_expire_time, user + try: + current_user = await UserDao.get_by_username(db, obj.username) + if not current_user: + raise errors.NotFoundError(msg='用户不存在') + elif not jwt.password_verify(obj.password, current_user.password): + raise errors.AuthorizationError(msg='密码错误') + elif not current_user.is_active: + raise errors.AuthorizationError(msg='用户已锁定, 登陆失败') + await UserDao.update_login_time(db, obj.username, self.login_time) + user_role_ids = await UserDao.get_role_ids(db, current_user.id) + user = await UserDao.get(db, current_user.id) + access_token, access_token_expire_time = await jwt.create_access_token( + str(user.id), role_ids=user_role_ids + ) + refresh_token, refresh_token_expire_time = await jwt.create_refresh_token( + str(user.id), access_token_expire_time, role_ids=user_role_ids + ) + login_logs_params = dict( + db=db, request=request, user=user, login_time=self.login_time, status=1, msg='登录成功' + ) + except errors.NotFoundError as e: + raise errors.NotFoundError(msg=e.msg) + except errors.AuthorizationError as e: + login_logs_params.update({'status': 0, 'msg': e.msg}) + background_tasks.add_task(LoginLogService.create, **login_logs_params) + raise errors.AuthorizationError(msg=e.msg) + except Exception as e: + raise e + else: + background_tasks.add_task(LoginLogService.create, **login_logs_params) + return access_token, refresh_token, access_token_expire_time, refresh_token_expire_time, user @staticmethod - async def refresh_token(user_id: int, custom_time: RefreshTokenTime): + async def refresh_token(*, user_id: int, custom_time: RefreshTokenTime) -> tuple[str, datetime]: async with async_db_session() as db: current_user = await UserDao.get(db, user_id) if not current_user: @@ -72,13 +96,12 @@ class UserService: return refresh_token, refresh_token_expire_time @staticmethod - async def logout(user_id: int): + async def logout(user_id: int) -> NoReturn: key = f'{settings.TOKEN_REDIS_PREFIX}:{user_id}:' await redis_client.delete_prefix(key) - return @staticmethod - async def register(obj: CreateUser): + async def register(obj: CreateUser) -> NoReturn: async with async_db_session.begin() as db: username = await UserDao.get_by_username(db, obj.username) if username: @@ -100,16 +123,17 @@ class UserService: await UserDao.create(db, obj) @staticmethod - async def pwd_reset(obj: ResetPassword): + async def pwd_reset(obj: ResetPassword) -> int: async with async_db_session.begin() as db: pwd1 = obj.password1 pwd2 = obj.password2 if pwd1 != pwd2: raise errors.ForbiddenError(msg='两次密码输入不一致') - await UserDao.reset_password(db, obj.id, obj.password2) + count = await UserDao.reset_password(db, obj.id, obj.password2) + return count @staticmethod - async def get_userinfo(username: str): + async def get_userinfo(username: str) -> User: async with async_db_session() as db: user = await UserDao.get_with_relation(db, username=username) if not user: @@ -117,7 +141,7 @@ class UserService: return user @staticmethod - async def update(*, username: str, current_user: User, obj: UpdateUser): + async def update(*, username: str, current_user: User, obj: UpdateUser) -> int: async with async_db_session.begin() as db: if not current_user.is_superuser: if not username == current_user.username: @@ -151,7 +175,7 @@ class UserService: return count @staticmethod - async def update_avatar(*, username: str, current_user: User, avatar: Avatar): + async def update_avatar(*, username: str, current_user: User, avatar: Avatar) -> int: async with async_db_session.begin() as db: if not current_user.is_superuser: if not username == current_user.username: @@ -163,11 +187,11 @@ class UserService: return count @staticmethod - async def get_user_list(): + async def get_select() -> Select: return UserDao.get_all() @staticmethod - async def update_permission(pk: int): + async def update_permission(pk: int) -> int: async with async_db_session.begin() as db: if await UserDao.get(db, pk): count = await UserDao.set_super(db, pk) @@ -176,7 +200,7 @@ class UserService: raise errors.NotFoundError(msg='用户不存在') @staticmethod - async def update_active(pk: int): + async def update_active(pk: int) -> int: async with async_db_session.begin() as db: if await UserDao.get(db, pk): count = await UserDao.set_active(db, pk) @@ -185,7 +209,7 @@ class UserService: raise errors.NotFoundError(msg='用户不存在') @staticmethod - async def delete(*, username: str, current_user: User): + async def delete(*, username: str, current_user: User) -> int: async with async_db_session.begin() as db: if not current_user.is_superuser: if not username == current_user.username: diff --git a/backend/app/utils/request_parse.py b/backend/app/utils/request_parse.py new file mode 100644 index 00000000..6d2a96c0 --- /dev/null +++ b/backend/app/utils/request_parse.py @@ -0,0 +1,37 @@ +#!/usr/bin/env python3 +# -*- coding: utf-8 -*- +import httpx +from httpx import HTTPError +from fastapi import Request + + +async def get_request_ip(request: Request) -> str: + """获取请求的ip地址""" + real = request.headers.get('X-Real-IP') + if real: + ip = real + else: + forwarded = request.headers.get('X-Forwarded-For') + if forwarded: + ip = forwarded.split(',')[0] + else: + ip = request.client.host + return ip + + +async def get_location(ipaddr: str, user_agent: str) -> str: + """获取ip地址归属地(临时)""" + async with httpx.AsyncClient(timeout=3) as client: + ip_api_url = f'http://ip-api.com/json/{ipaddr}?lang=zh-CN' + whois_url = f'http://whois.pconline.com.cn/ipJson.jsp?ip={ipaddr}&json=true' + headers = {'User-Agent': user_agent} + try: + resp1 = await client.get(ip_api_url, headers=headers) + city = resp1.json()['city'] + except (HTTPError, KeyError): + try: + resp2 = await client.get(whois_url, headers=headers) + city = resp2.json()['city'] + except (HTTPError, KeyError): + city = None + return city or '未知' if city != '' else '未知' diff --git a/backend/app/utils/serializers.py b/backend/app/utils/serializers.py index c601ad8c..28341c2f 100644 --- a/backend/app/utils/serializers.py +++ b/backend/app/utils/serializers.py @@ -1,11 +1,16 @@ #!/usr/bin/env python3 # -*- coding: utf-8 -*- from decimal import Decimal +from typing import Any, TypeVar -from sqlalchemy.sql import Select +from sqlalchemy import Row, RowMapping + +RowData = Row | RowMapping | Any + +R = TypeVar('R', bound=RowData) -def select_to_dict(obj: Select) -> dict: +def select_to_dict(obj: R) -> dict: """ Serialize SQLAlchemy Select to dict @@ -21,7 +26,7 @@ def select_to_dict(obj: Select) -> dict: return obj_dict -def select_to_list(obj: list) -> list: +def select_to_list(obj: list[R]) -> list: """ Serialize SQLAlchemy Select to list @@ -35,7 +40,7 @@ def select_to_list(obj: list) -> list: return ret_list -def select_to_json(obj: Select) -> dict: +def select_to_json(obj: R) -> dict: """ Serialize SQLAlchemy Select to json diff --git a/requirements.txt b/requirements.txt index 389e031a..ed812cd0 100644 --- a/requirements.txt +++ b/requirements.txt @@ -32,5 +32,6 @@ SQLAlchemy==2.0.8 starlette==0.27.0 supervisor==4.2.5 tzlocal==2.1 +user_agents==2.2.0 uvicorn[standard]==0.13.4 wait-for-it==2.2.1