auth.py 7.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188
  1. import json
  2. from fastapi import APIRouter, Depends
  3. from sqlalchemy.orm import Session
  4. from sqlalchemy.ext.asyncio import AsyncSession
  5. from Log import logger
  6. from app.api import Response, pwd_context, get_current_user
  7. from app.config.config import settings
  8. from app.config.const import RAGFLOW, BISHENG, DIFY
  9. from app.models.app_token_model import AppToken
  10. from app.models.base_model import get_db
  11. from app.models.postgresql_base_model import get_pdb
  12. from app.models.token_model import upsert_token, get_token, update_token
  13. from app.models.user import UserCreate, LoginData
  14. from app.models.user_model import UserModel
  15. from app.service.auth import authenticate_user, create_access_token
  16. from app.service.bisheng import BishengService
  17. from app.service.common.app_register import AppRegisterDao
  18. from app.service.ragflow import RagflowService
  19. from sqlalchemy.future import select
  20. router = APIRouter()
  21. @router.post("/register", response_model=Response)
  22. async def register(user: UserCreate, db=Depends(get_db)):
  23. db_user = db.query(UserModel).filter(UserModel.username == user.username).first()
  24. if db_user:
  25. return Response(code=200, msg="Username already registered")
  26. bisheng_service = BishengService(settings.sgb_base_url)
  27. ragflow_service = RagflowService(settings.fwr_base_url)
  28. # 注册到毕昇
  29. try:
  30. bisheng_info = await bisheng_service.register(user.username, user.password)
  31. except Exception as e:
  32. return Response(code=500, msg=f"Failed to register with Bisheng: {str(e)}")
  33. # 注册到ragflow
  34. try:
  35. ragflow_info = await ragflow_service.register(user.username, user.password)
  36. except Exception as e:
  37. return Response(code=500, msg=f"Failed to register with Ragflow: {str(e)}")
  38. # 存储用户信息
  39. hashed_password = pwd_context.hash(user.password)
  40. db_user = UserModel(username=user.username, hashed_password=hashed_password, email=ragflow_info.get("email", f"{user.username}@example.com"),ragflow_id=ragflow_info.get("id"),bisheng_id=bisheng_info.get("user_id"))
  41. db_user.password = db_user.encrypted_password(user.password)
  42. db.add(db_user)
  43. db.commit()
  44. db.refresh(db_user)
  45. return Response(code=200, msg="User registered successfully",data={"username": db_user.username})
  46. @router.post("/login", response_model=Response)
  47. async def login(login_data: LoginData, db: Session = Depends(get_db)):
  48. user = authenticate_user(db, login_data.username, login_data.password)
  49. if not user:
  50. return Response(code=400, msg="Incorrect username or password")
  51. bisheng_service = BishengService(settings.sgb_base_url)
  52. ragflow_service = RagflowService(settings.fwr_base_url)
  53. # 登录到毕昇
  54. try:
  55. bisheng_token = await bisheng_service.login(login_data.username, login_data.password)
  56. except Exception as e:
  57. return Response(code=500, msg=f"Failed to login with Bisheng: {str(e)}")
  58. # 登录到ragflow
  59. try:
  60. ragflow_token = await ragflow_service.login(login_data.username, login_data.password)
  61. except Exception as e:
  62. return Response(code=500, msg=f"Failed to login with Ragflow: {str(e)}")
  63. # 创建本地token
  64. access_token = create_access_token(data={"sub": user.username, "user_id": user.id})
  65. upsert_token(db, user.id, access_token, bisheng_token, ragflow_token)
  66. return Response(code=200, msg="Login successful", data={
  67. "access_token": access_token,
  68. "token_type": "bearer",
  69. "username": user.username,
  70. "nickname": "",
  71. "user": user.to_login_json()
  72. })
  73. @router.get("/token", response_model=Response)
  74. async def token_api(db: Session = Depends(get_db), current_user: UserModel = Depends(get_current_user)):
  75. # 查询现有记录
  76. token = get_token(db, current_user.id)
  77. if token is None:
  78. return Response(code=400, msg="token not found")
  79. return Response(code=200, msg="success", data={
  80. "ragflow_token": token.ragflow_token,
  81. "bisheng_token": token.bisheng_token,
  82. })
  83. @router.post("/v2/login", response_model=Response)
  84. async def login_test(login_data: LoginData, db: Session = Depends(get_db), pdb: AsyncSession = Depends(get_pdb)):
  85. user = authenticate_user(db, login_data.username, login_data.password)
  86. if not user:
  87. return Response(code=400, msg="Incorrect username or password")
  88. app_register = AppRegisterDao(db).get_apps()
  89. token_dict = {}
  90. for app in app_register:
  91. if app["id"] == RAGFLOW:
  92. service = RagflowService(settings.fwr_base_url)
  93. elif app["id"] == BISHENG:
  94. service = BishengService(settings.sgb_base_url)
  95. elif app["id"] == DIFY:
  96. continue
  97. else:
  98. logger.error("未知注册应用---")
  99. continue
  100. try:
  101. token = await service.login(login_data.username, login_data.password)
  102. token_dict[app["id"]] = token
  103. except Exception as e:
  104. return Response(code=500, msg=f"Failed to login with {app['id']}: {str(e)}")
  105. # 创建本地token
  106. access_token = create_access_token(data={"sub": user.username, "user_id": user.id})
  107. await update_token(db, user.id, access_token, token_dict)
  108. result = await pdb.execute(select(AppToken).where(AppToken.id == user.id))
  109. db_app_token = result.scalars().first()
  110. if not db_app_token:
  111. app_token_str = json.dumps(token_dict)
  112. # print(app_token_str)
  113. app_token = AppToken(id=user.id, token=access_token.decode(), app_token=app_token_str)
  114. pdb.add(app_token)
  115. await pdb.commit()
  116. await pdb.refresh(app_token)
  117. else:
  118. db_app_token.token = access_token.decode()
  119. db_app_token.app_token = json.dumps(token_dict)
  120. await pdb.commit()
  121. await pdb.refresh(db_app_token)
  122. return Response(code=200, msg="Login successful", data={
  123. "access_token": access_token,
  124. "token_type": "bearer",
  125. "username": user.username,
  126. "nickname": "",
  127. # "user": user.to_login_json()
  128. })
  129. @router.post("/v2/register", response_model=Response)
  130. async def register_test(user: UserCreate, db=Depends(get_db)):
  131. db_user = db.query(UserModel).filter(UserModel.username == user.username).first()
  132. if db_user:
  133. return Response(code=200, msg="Username already registered")
  134. app_register = AppRegisterDao(db).get_apps()
  135. register_dict = {}
  136. for app in app_register:
  137. if app["id"] == RAGFLOW:
  138. service = RagflowService(settings.fwr_base_url)
  139. elif app["id"] == BISHENG:
  140. service = BishengService(settings.sgb_base_url)
  141. elif app["id"] == DIFY:
  142. continue
  143. else:
  144. logger.error("未知注册应用---")
  145. continue
  146. try:
  147. register_info = await service.register(user.username, user.password)
  148. register_dict[app['id']] = register_info.get("id") if app['id'] == RAGFLOW else register_info.get("user_id") if app['id'] == BISHENG else ""
  149. except Exception as e:
  150. return Response(code=500, msg=f"Failed to register with {app['id']}: {str(e)}")
  151. # 存储用户信息
  152. hashed_password = pwd_context.hash(user.password)
  153. db_user = UserModel(username=user.username, hashed_password=hashed_password, email=user.email)
  154. db_user.password = db_user.encrypted_password(user.password)
  155. for k, v in register_dict.items():
  156. setattr(db_user, k.replace("app", "id"), v)
  157. db.add(db_user)
  158. db.commit()
  159. db.refresh(db_user)
  160. return Response(code=200, msg="User registered successfully",data={"username": db_user.username})