mirror of
https://github.com/fastapi-practices/fastapi-best-architecture.git
synced 2026-09-21 13:12:24 +00:00
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
This commit is contained in:
@@ -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=['任务管理'])
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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'),
|
||||
)
|
||||
|
||||
@@ -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'
|
||||
|
||||
@@ -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)
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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='创建时间')
|
||||
@@ -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
|
||||
@@ -49,7 +49,7 @@ class GetUserInfoNoRelation(_UserInfoBase):
|
||||
orm_mode = True
|
||||
|
||||
|
||||
class GetUserInfo(GetUserInfoNoRelation):
|
||||
class GetAllUserInfo(GetUserInfoNoRelation):
|
||||
dept: GetAllDept | None = None
|
||||
roles: list[GetAllRole]
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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:
|
||||
|
||||
@@ -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 '未知'
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user