excel.py 6.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178
  1. import random
  2. import string
  3. from fastapi import APIRouter, File, UploadFile, Form, BackgroundTasks, Depends, Request, WebSocket
  4. from fastapi.responses import JSONResponse, FileResponse
  5. from sqlalchemy.orm import Session
  6. # from starlette.websockets import WebSocket
  7. from app.api import get_current_user, get_current_user_websocket, Response
  8. from app.models import UserModel, AgentType
  9. from app.models.base_model import get_db
  10. from app.service.session import SessionService
  11. from app.utils.excelmerge.conformity import run_conformity
  12. import shutil
  13. import os
  14. router = APIRouter()
  15. ALLOWED_EXTENSIONS = {'xlsx'}
  16. EXCEL_FILES_PATH = 'data/output'
  17. SOURCE_FILES_PATH = 'data/source'
  18. def allowed_file(filename: str) -> bool:
  19. return '.' in filename and filename.rsplit('.', 1)[1].lower() in ALLOWED_EXTENSIONS
  20. def create_dir_if_not_exists(path: str):
  21. if not os.path.exists(path):
  22. os.makedirs(path)
  23. def clear_directory(path: str) -> dict:
  24. for filename in os.listdir(path):
  25. file_path = os.path.join(path, filename)
  26. try:
  27. if os.path.isfile(file_path) or os.path.islink(file_path):
  28. os.unlink(file_path)
  29. elif os.path.isdir(file_path):
  30. shutil.rmtree(file_path)
  31. except Exception as e:
  32. return {"error": "清空出错"}
  33. return {"message": "目录已清空"}
  34. def user_file_path(userid: str, path: str) -> str:
  35. return os.path.join(path, userid)
  36. def generate_db_id(prefix: str = "me") -> str:
  37. random_part = ''.join(random.choices(string.ascii_letters + string.digits, k=13))
  38. return prefix + random_part
  39. def db_create_session(db: Session, user_id: str, message:str, upload_filenames: list):
  40. db_id = generate_db_id()
  41. session = SessionService(db).create_session(
  42. db_id,
  43. message,
  44. "basic_excel_merge",
  45. AgentType.BASIC,
  46. int(user_id),
  47. {"role": "user", "content": message, "upload_filenames": upload_filenames}
  48. )
  49. return session
  50. @router.post('/excel/upload', response_model=Response)
  51. async def upload_file(files: list[UploadFile] = File(...), current_user: UserModel = Depends(get_current_user)):
  52. user_id = str(current_user.id)
  53. if not any(file.filename for file in files):
  54. return Response(code=400, msg="没有文件部分", data={})
  55. if not user_id:
  56. return Response(code=400, msg="缺少参数user_id", data={})
  57. user_source = user_file_path(user_id, SOURCE_FILES_PATH)
  58. user_excel = EXCEL_FILES_PATH
  59. create_dir_if_not_exists(user_source)
  60. create_dir_if_not_exists(user_excel)
  61. clear_directory(user_source)
  62. save_path_list = []
  63. for file in files:
  64. if file and allowed_file(file.filename):
  65. save_path = os.path.join(user_source, file.filename)
  66. with open(save_path, 'wb') as buffer:
  67. shutil.copyfileobj(file.file, buffer)
  68. save_path_list.append(save_path)
  69. else:
  70. return Response(code=400, msg="不允许的文件类型", data={})
  71. return Response(code=200, msg="上传成功", data={})
  72. # ws://localhost:9201/api/document/ws/excel
  73. @router.websocket("/ws/excel")
  74. async def ws_excel(websocket: WebSocket,
  75. current_user: UserModel = Depends(get_current_user_websocket),
  76. db: Session = Depends(get_db)):
  77. await websocket.accept()
  78. user_id = str(current_user.id)
  79. user_source = user_file_path(user_id, SOURCE_FILES_PATH)
  80. user_excel = EXCEL_FILES_PATH
  81. create_dir_if_not_exists(user_source)
  82. create_dir_if_not_exists(user_excel)
  83. while True:
  84. # data = await websocket.receive_text()git
  85. receive_message = await websocket.receive_json()
  86. try:
  87. if receive_message.get("message") == "合并Excel":
  88. upload_filenames = receive_message.get('upload_filenames', [])
  89. merge_file = run_conformity(user_source, user_excel)
  90. if merge_file is not None:
  91. await websocket.send_json({
  92. "type": "stream",
  93. "files": [
  94. {
  95. "file_name": "Excel",
  96. "file_url": f"./api/document/download/{merge_file}.xlsx?file_type=excel",
  97. }
  98. ]
  99. })
  100. await websocket.send_json({
  101. "message": "合并成功",
  102. "type": "close",
  103. })
  104. # 创建会话记录
  105. session = db_create_session(db, user_id, receive_message.get("message"), upload_filenames)
  106. # 更新会话记录
  107. if session:
  108. session_id = session.id
  109. new_message = {
  110. "role": "assistant",
  111. "content": {
  112. "message": "\u5408\u5e76\u6210\u529f",
  113. "type": "message",
  114. "file_name": "Excel",
  115. "file_url": f"/api/document/download/{merge_file}.xlsx?file_type=excel"
  116. }
  117. }
  118. session_service = SessionService(db)
  119. session_service.update_session(session_id, message=new_message)
  120. else:
  121. await websocket.send_json({"error": "合并失败", "type": "stream", "files": []})
  122. await websocket.close()
  123. else:
  124. print(f"Received data: {receive_message.get('message')}")
  125. await websocket.send_json({"error": "未知指令", "data": str(receive_message.get('message'))})
  126. await websocket.close()
  127. except Exception as e:
  128. await websocket.send_json({"error": str(e)})
  129. await websocket.close()
  130. @router.get("/download/{file_full_name}")
  131. async def download_file(file_full_name: str):
  132. file_name = os.path.basename(file_full_name)
  133. user_excel = EXCEL_FILES_PATH
  134. file_path = os.path.join(user_excel, file_full_name)
  135. if not os.path.exists(file_path):
  136. return JSONResponse(content={"error": "文件不存在"}, status_code=404)
  137. return FileResponse(
  138. path=file_path,
  139. filename="Excel.xlsx",
  140. media_type='application/octet-stream',
  141. )
  142. # def delete_file():
  143. # try:
  144. # os.unlink(file_path)
  145. # except OSError as e:
  146. # print(f"Deleting file error")
  147. # 待下载完成后删除生成的文件
  148. # background_tasks.add_task(delete_file)
  149. # return FileResponse(path=file_path, filename=file_name,
  150. # media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet")