|
1 | 1 | #!/usr/bin/env python3 |
2 | 2 | # -*- coding: utf-8 -*- |
3 | | -from typing import NoReturn |
4 | | - |
5 | 3 | from fastapi import Request |
6 | 4 | from sqlalchemy import Select |
7 | 5 |
|
|
15 | 13 | from backend.app.crud.crud_user import UserDao |
16 | 14 | from backend.app.database.db_mysql import async_db_session |
17 | 15 | from backend.app.models import User |
18 | | -from backend.app.schemas.user import CreateUser, ResetPassword, UpdateUser, Avatar, UpdateUserRole |
| 16 | +from backend.app.schemas.user import RegisterUser, ResetPassword, UpdateUser, Avatar, UpdateUserRole, AddUser |
19 | 17 |
|
20 | 18 |
|
21 | 19 | class UserService: |
22 | 20 | @staticmethod |
23 | | - async def register(*, obj: CreateUser) -> NoReturn: |
| 21 | + async def register(*, obj: RegisterUser) -> None: |
24 | 22 | async with async_db_session.begin() as db: |
25 | 23 | username = await UserDao.get_by_username(db, obj.username) |
26 | 24 | if username: |
27 | 25 | raise errors.ForbiddenError(msg='该用户名已注册') |
| 26 | + nickname = await UserDao.get_by_nickname(db, obj.nickname) |
| 27 | + if nickname: |
| 28 | + raise errors.ForbiddenError(msg='该昵称已注册') |
28 | 29 | email = await UserDao.check_email(db, obj.email) |
29 | 30 | if email: |
30 | 31 | raise errors.ForbiddenError(msg='该邮箱已注册') |
| 32 | + await UserDao.create(db, obj) |
| 33 | + |
| 34 | + @staticmethod |
| 35 | + async def add(*, obj: AddUser) -> None: |
| 36 | + async with async_db_session.begin() as db: |
| 37 | + username = await UserDao.get_by_username(db, obj.username) |
| 38 | + if username: |
| 39 | + raise errors.ForbiddenError(msg='该用户名已注册') |
| 40 | + nickname = await UserDao.get_by_nickname(db, obj.nickname) |
| 41 | + if nickname: |
| 42 | + raise errors.ForbiddenError(msg='该昵称已注册') |
31 | 43 | dept = await DeptDao.get(db, obj.dept_id) |
32 | 44 | if not dept: |
33 | 45 | raise errors.NotFoundError(msg='部门不存在') |
34 | 46 | for role_id in obj.roles: |
35 | 47 | role = await RoleDao.get(db, role_id) |
36 | 48 | if not role: |
37 | 49 | raise errors.NotFoundError(msg='角色不存在') |
38 | | - await UserDao.create(db, obj) |
| 50 | + email = await UserDao.check_email(db, obj.email) |
| 51 | + if email: |
| 52 | + raise errors.ForbiddenError(msg='该邮箱已注册') |
| 53 | + await UserDao.add(db, obj) |
39 | 54 |
|
40 | 55 | @staticmethod |
41 | 56 | async def pwd_reset(*, request: Request, obj: ResetPassword) -> int: |
@@ -72,16 +87,17 @@ async def update(*, request: Request, username: str, obj: UpdateUser) -> int: |
72 | 87 | if not input_user: |
73 | 88 | raise errors.NotFoundError(msg='用户不存在') |
74 | 89 | if input_user.username != obj.username: |
75 | | - username = await UserDao.get_by_username(db, obj.username) |
76 | | - if username: |
| 90 | + _username = await UserDao.get_by_username(db, obj.username) |
| 91 | + if _username: |
77 | 92 | raise errors.ForbiddenError(msg='该用户名已存在') |
| 93 | + if input_user.nickname != obj.nickname: |
| 94 | + nickname = await UserDao.get_by_nickname(db, obj.nickname) |
| 95 | + if nickname: |
| 96 | + raise errors.ForbiddenError(msg='改昵称已存在') |
78 | 97 | if input_user.email != obj.email: |
79 | 98 | email = await UserDao.check_email(db, obj.email) |
80 | 99 | if email: |
81 | 100 | raise errors.ForbiddenError(msg='该邮箱已注册') |
82 | | - dept = await DeptDao.get(db, obj.dept_id) |
83 | | - if not dept: |
84 | | - raise errors.NotFoundError(msg='部门不存在') |
85 | 101 | count = await UserDao.update_userinfo(db, input_user, obj) |
86 | 102 | return count |
87 | 103 |
|
|
0 commit comments