Skip to content

Commit dc44ef0

Browse files
committed
feat(auth, account, oauth): 添加账号设置和授权认证处理器,支持密码登录、更新信息及第三方OAuth集成
1 parent 666e6c7 commit dc44ef0

14 files changed

Lines changed: 572 additions & 1 deletion

File tree

Lines changed: 66 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,66 @@
1+
from dataclasses import dataclass
2+
3+
from flask_login import login_required, current_user
4+
from injector import inject
5+
6+
from internal.schema.account_schema import (
7+
GetCurrentUserResp,
8+
UpdatePasswordReq,
9+
UpdateNameReq,
10+
UpdateAvatarReq,
11+
)
12+
from internal.service import AccountService
13+
from pkg.response import success_json, validate_error_json, success_message
14+
15+
16+
@inject
17+
@dataclass
18+
class AccountHandler:
19+
"""账号设置处理器"""
20+
21+
account_service: AccountService
22+
23+
@login_required
24+
def get_current_user(self):
25+
"""获取当前登录账号信息"""
26+
resp = GetCurrentUserResp()
27+
return success_json(resp.dump(current_user))
28+
29+
@login_required
30+
def update_password(self):
31+
"""更新当前登录账号密码"""
32+
# 1.提取请求数据并校验
33+
req = UpdatePasswordReq()
34+
if not req.validate():
35+
return validate_error_json(req.errors)
36+
37+
# 2.调用服务更新账号密码
38+
self.account_service.update_password(req.password.data, current_user)
39+
40+
return success_message("更新账号密码成功")
41+
42+
@login_required
43+
def update_name(self):
44+
"""更新当前登录账号名称"""
45+
# 1.提取请求数据并校验
46+
req = UpdateNameReq()
47+
if not req.validate():
48+
return validate_error_json(req.errors)
49+
50+
# 2.调用服务更新账号名称
51+
self.account_service.update_account(current_user, name=req.name.data)
52+
53+
return success_message("更新账号名称成功")
54+
55+
@login_required
56+
def update_avatar(self):
57+
"""更新当前账号头像信息"""
58+
# 1.提取请求数据并校验
59+
req = UpdateAvatarReq()
60+
if not req.validate():
61+
return validate_error_json(req.errors)
62+
63+
# 2.调用服务更新账号名称
64+
self.account_service.update_account(current_user, avatar=req.avatar.data)
65+
66+
return success_message("更新账号头像成功")

internal/handler/auth_handler.py

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,39 @@
1+
from dataclasses import dataclass
2+
3+
from flask_login import logout_user, login_required
4+
from injector import inject
5+
6+
from internal.schema.auth_schema import PasswordLoginReq, PasswordLoginResp
7+
from internal.service import AccountService
8+
from pkg.response import success_message, validate_error_json, success_json
9+
10+
11+
@inject
12+
@dataclass
13+
class AuthHandler:
14+
"""LLMOps平台自有授权认证处理器"""
15+
16+
account_service: AccountService
17+
18+
def password_login(self):
19+
"""账号密码登录"""
20+
# 1.提取请求并校验数据
21+
req = PasswordLoginReq()
22+
if not req.validate():
23+
return validate_error_json(req.errors)
24+
25+
# 2.调用服务登录账号
26+
credential = self.account_service.password_login(
27+
req.email.data, req.password.data
28+
)
29+
30+
# 3.创建响应结构并返回
31+
resp = PasswordLoginResp()
32+
33+
return success_json(resp.dump(credential))
34+
35+
@login_required
36+
def logout(self):
37+
"""退出登录,用于提示前端清除授权凭证"""
38+
logout_user()
39+
return success_message("退出登陆成功")

internal/handler/oauth_handler.py

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,37 @@
1+
from dataclasses import dataclass
2+
3+
from injector import inject
4+
5+
from internal.schema.oauth_schema import AuthorizeReq, AuthorizeResp
6+
from internal.service import OAuthService
7+
from pkg.response import success_json, validate_error_json
8+
9+
10+
@inject
11+
@dataclass
12+
class OAuthHandler:
13+
"""第三方授权认证处理器"""
14+
15+
oauth_service: OAuthService
16+
17+
def provider(self, provider_name: str):
18+
"""根据传递的提供商名字获取授权认证重定向地址"""
19+
# 1.根据provider_name获取授权服务提供商
20+
oauth = self.oauth_service.get_oauth_by_provider_name(provider_name)
21+
22+
# 2.调用函数获取授权地址
23+
redirect_url = oauth.get_authorization_url()
24+
25+
return success_json({"redirect_url": redirect_url})
26+
27+
def authorize(self, provider_name: str):
28+
"""根据传递的提供商名字+code获取第三方授权信息"""
29+
# 1.提取请求数据并校验
30+
req = AuthorizeReq()
31+
if not req.validate():
32+
return validate_error_json(req.errors)
33+
34+
# 2.调用服务登录账号
35+
credential = self.oauth_service.oauth_login(provider_name, req.code.data)
36+
37+
return success_json(AuthorizeResp().dump(credential))

internal/router/router.py

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,9 @@
88
from internal.handler.document_handler import DocumentHandler
99
from internal.handler.segment_handler import SegmentHandler
1010
from internal.handler.upload_file_handler import UploadFileHandler
11+
from internal.handler.oauth_handler import OAuthHandler
12+
from internal.handler.account_handler import AccountHandler
13+
from internal.handler.auth_handler import AuthHandler
1114

1215

1316
@inject
@@ -22,6 +25,9 @@ class Router:
2225
dataset_handler: DatasetHandler
2326
document_handler: DocumentHandler
2427
segment_handler: SegmentHandler
28+
oauth_handler: OAuthHandler
29+
account_handler: AccountHandler
30+
auth_handler: AuthHandler
2531

2632
def register_router(self, app: Flask):
2733
"""注册路由"""
@@ -215,5 +221,44 @@ def register_router(self, app: Flask):
215221
view_func=self.segment_handler.update_segment,
216222
)
217223

224+
# 授权认证模块
225+
bp.add_url_rule(
226+
"/oauth/<string:provider_name>",
227+
view_func=self.oauth_handler.provider,
228+
)
229+
bp.add_url_rule(
230+
"/oauth/authorize/<string:provider_name>",
231+
methods=["POST"],
232+
view_func=self.oauth_handler.authorize,
233+
)
234+
bp.add_url_rule(
235+
"/auth/password-login",
236+
methods=["POST"],
237+
view_func=self.auth_handler.password_login,
238+
)
239+
bp.add_url_rule(
240+
"/auth/logout",
241+
methods=["POST"],
242+
view_func=self.auth_handler.logout,
243+
)
244+
245+
# 账号设置模块
246+
bp.add_url_rule("/account", view_func=self.account_handler.get_current_user)
247+
bp.add_url_rule(
248+
"/account/password",
249+
methods=["POST"],
250+
view_func=self.account_handler.update_password,
251+
)
252+
bp.add_url_rule(
253+
"/account/name",
254+
methods=["POST"],
255+
view_func=self.account_handler.update_name,
256+
)
257+
bp.add_url_rule(
258+
"/account/avatar",
259+
methods=["POST"],
260+
view_func=self.account_handler.update_avatar,
261+
)
262+
218263
# 在应用上去注册蓝图
219264
app.register_blueprint(bp)

internal/schema/account_schema.py

Lines changed: 72 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,72 @@
1+
#!/usr/bin/env python
2+
from flask_wtf import FlaskForm
3+
from marshmallow import Schema, fields, pre_dump
4+
from wtforms import StringField
5+
from wtforms.validators import DataRequired, regexp, Length, URL
6+
7+
from internal.lib.helper import datetime_to_timestamp
8+
from internal.model import Account
9+
from pkg.password import password_pattern
10+
11+
12+
class GetCurrentUserResp(Schema):
13+
"""获取当前登录账号信息响应"""
14+
15+
id = fields.UUID(dump_default="")
16+
name = fields.String(dump_default="")
17+
email = fields.String(dump_default="")
18+
avatar = fields.String(dump_default="")
19+
last_login_at = fields.Integer(dump_default=0)
20+
last_login_ip = fields.String(dump_default="")
21+
created_at = fields.Integer(dump_default=0)
22+
23+
@pre_dump
24+
def process_data(self, data: Account, **kwargs):
25+
return {
26+
"id": data.id,
27+
"name": data.name,
28+
"email": data.email,
29+
"avatar": data.avatar,
30+
"last_login_at": datetime_to_timestamp(data.last_login_at),
31+
"last_login_ip": data.last_login_ip,
32+
"created_at": datetime_to_timestamp(data.created_at),
33+
}
34+
35+
36+
class UpdatePasswordReq(FlaskForm):
37+
"""更新账号密码请求"""
38+
39+
password = StringField(
40+
"password",
41+
validators=[
42+
DataRequired("登录密码不能为空"),
43+
regexp(
44+
regex=password_pattern,
45+
message="密码最少包含一个字母、一个数字,并且长度是8-16",
46+
),
47+
],
48+
)
49+
50+
51+
class UpdateNameReq(FlaskForm):
52+
"""更新账号名称请求"""
53+
54+
name = StringField(
55+
"name",
56+
validators=[
57+
DataRequired("账号名字不能为空"),
58+
Length(min=3, max=30, message="账号名称长度在3-30位"),
59+
],
60+
)
61+
62+
63+
class UpdateAvatarReq(FlaskForm):
64+
"""更新账号头像请求"""
65+
66+
avatar = StringField(
67+
"avatar",
68+
validators=[
69+
DataRequired("账号头像不能为空"),
70+
URL("账号头像必须是URL图片地址"),
71+
],
72+
)

internal/schema/auth_schema.py

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,36 @@
1+
from flask_wtf import FlaskForm
2+
from marshmallow import Schema, fields
3+
from wtforms import StringField
4+
from wtforms.validators import DataRequired, Email, Length, regexp
5+
6+
from pkg.password import password_pattern
7+
8+
9+
class PasswordLoginReq(FlaskForm):
10+
"""账号密码登录请求结构"""
11+
12+
email = StringField(
13+
"email",
14+
validators=[
15+
DataRequired("登录邮箱不能为空"),
16+
Email("登录邮箱格式错误"),
17+
Length(min=5, max=254, message="登录邮箱长度在5-254个字符"),
18+
],
19+
)
20+
password = StringField(
21+
"password",
22+
validators=[
23+
DataRequired("账号密码不能为空"),
24+
regexp(
25+
regex=password_pattern,
26+
message="密码最少包含一个字母,一个数字,并且长度为8-16",
27+
),
28+
],
29+
)
30+
31+
32+
class PasswordLoginResp(Schema):
33+
"""账号密码授权认证响应结构"""
34+
35+
access_token = fields.String()
36+
expire_at = fields.Integer()

internal/schema/oauth_schema.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,17 @@
1+
from flask_wtf import FlaskForm
2+
from marshmallow import Schema, fields
3+
from wtforms import StringField
4+
from wtforms.validators import DataRequired
5+
6+
7+
class AuthorizeReq(FlaskForm):
8+
"""第三方授权认证请求体"""
9+
10+
code = StringField("code", validators=[DataRequired("code代码不能为空")])
11+
12+
13+
class AuthorizeResp(Schema):
14+
"""第三方授权认证响应结构"""
15+
16+
access_token = fields.String()
17+
expire_at = fields.Integer()

internal/server/http.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -56,7 +56,7 @@ def __init__(
5656
},
5757
)
5858

59-
self.before_request(middleware.request_loader)
59+
login_manager.request_loader(middleware.request_loader)
6060

6161
router.register_router(self)
6262

internal/service/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
from .conversation_service import ConversationService
1818
from .jwt_service import JwtService
1919
from .account_service import AccountService
20+
from .oauth_service import OAuthService
2021

2122
__all__ = [
2223
"BuiltinToolService",
@@ -38,4 +39,5 @@
3839
"ConversationService",
3940
"JwtService",
4041
"AccountService",
42+
"OAuthService",
4143
]

internal/service/dataset_service.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,7 @@
11
from dataclasses import dataclass
22
from uuid import UUID
3+
4+
from flask_login import login_required
35
from internal.lib.helper import datetime_to_timestamp
46
from injector import inject
57
from sqlalchemy import desc
@@ -179,6 +181,7 @@ def update_dataset(self, dataset_id: UUID, req: UpdateDatasetReq):
179181

180182
return dataset
181183

184+
@login_required
182185
def get_datasets_with_page(self, req: GetDatasetsWithPageReq):
183186
"""根据传递的信息获取知识库列表分页数据"""
184187

0 commit comments

Comments
 (0)