96 lines
3.5 KiB
Python
96 lines
3.5 KiB
Python
from fastapi import APIRouter, File, UploadFile, Query
|
||
from fastapi.responses import FileResponse
|
||
from app.schemas import Success, SuccessExtra, Fail
|
||
from app.controllers.finance_parse import task_controller, parse_finance_data
|
||
import os, time
|
||
from app.schemas.task import DecodeTaskParams, DecodeTaskResult, TaskResponse
|
||
from app.models.automation import TaskStatus, TaskType
|
||
from datetime import datetime
|
||
router = APIRouter()
|
||
|
||
import shutil
|
||
from pathlib import Path
|
||
|
||
from app.utils.excel_utils import get_sheets_and_headers
|
||
|
||
|
||
@router.post("/upload", summary="上传文件")
|
||
async def upload_file(file: UploadFile = File(...)):
|
||
# file.filename: 原始文件名
|
||
# file.content_type: MIME 类型
|
||
# file.file: 类文件对象(SpooledTemporaryFile)
|
||
|
||
# 保存文件到本地(示例:保存到 ./uploads/)
|
||
upload_dir = Path("uploads")
|
||
upload_dir.mkdir(exist_ok=True)
|
||
filename = f"{datetime.now().strftime('%Y%m%d%H%M%S')}_{file.filename}"
|
||
|
||
file_path = upload_dir / filename
|
||
with open(file_path, "wb") as buffer:
|
||
shutil.copyfileobj(file.file, buffer)
|
||
|
||
# 获取所有 sheet 名称及其表头
|
||
sheets_headers = get_sheets_and_headers(file_path)
|
||
|
||
# # 打印结果
|
||
# for sheet_name, headers in sheets_headers.items():
|
||
# print(f"Sheet: {sheet_name}")
|
||
# print(f"Headers: {headers}\n")
|
||
|
||
data = {
|
||
"filename": filename,
|
||
"data": sheets_headers,
|
||
}
|
||
return Success(data=data)
|
||
|
||
|
||
@router.post("/parse", summary="解析文件")
|
||
async def parse_file(data: dict):
|
||
upload_dir = Path("uploads")
|
||
filename = upload_dir / data["filename"]
|
||
sheet = data["sheet"]
|
||
header = data["header"]
|
||
parse_type = data["parse_type"]
|
||
taskname = data["filename"] + "_" + sheet
|
||
|
||
print(filename, sheet, header, parse_type)
|
||
task, task_obj = await task_controller.create_task(name=taskname, obj_in=DecodeTaskParams(filename=str(filename), sheet_name=sheet, header=header, decode_type=parse_type))
|
||
total_amount = 0
|
||
start = time.time()
|
||
|
||
try:
|
||
result_file_path, total_amount = parse_finance_data(
|
||
filename,
|
||
target_index=header,
|
||
is_horizontal=(parse_type == "horizontal"),
|
||
sheet_name=sheet
|
||
)
|
||
except Exception as e:
|
||
return Fail(msg=f"解析失败: {str(e)}", code=400)
|
||
|
||
if not os.path.exists(result_file_path):
|
||
return Fail(msg=f"解析结果文件未生成", code=404)
|
||
|
||
# 提取原始文件名(不含路径),用于下载时的默认文件名
|
||
download_filename = os.path.basename(result_file_path)
|
||
await task_controller.update_task(task_obj, task.id, TaskStatus.SUCCESS, DecodeTaskResult(
|
||
filename=str(result_file_path),
|
||
spend=time.time() - start,
|
||
rows=total_amount,
|
||
))
|
||
|
||
return FileResponse(
|
||
path=result_file_path,
|
||
filename=download_filename, # 浏览器下载时显示的文件名
|
||
media_type='application/vnd.openxmlformats-officedocument.spreadsheetml.sheet' # .xlsx
|
||
)
|
||
|
||
@router.get("/list", summary="获取任务列表")
|
||
async def get_tasks(page: int = Query(1, description="页码"),
|
||
page_size: int = Query(10, description="每页数量"),
|
||
name: str = Query("", description="任务名称,用于查询"),
|
||
type: TaskType = Query(None, description="任务类型,用于查询")):
|
||
total, tasks = await task_controller.list(name, type, page, page_size)
|
||
data = [TaskResponse.from_orm(task).model_dump() for task in tasks]
|
||
return SuccessExtra(data=data, total=total, page=page, page_size=page_size)
|