RamEx / api /views /data_views.py
zdy10046's picture
deploy RamEx to Hugging Face without binary files
e657e99
Raw
History Blame Contribute Delete
17.3 kB
from rest_framework import status, permissions
from rest_framework.views import APIView
from rest_framework.response import Response
from api.models import Project, DataFile
from api.serializers import DataFileSerializer
from api.utils.r_executor import RExecutor
import os
import zipfile
import shutil
from django.conf import settings
import logging
import json
class DataUploadView(APIView):
"""数据上传视图"""
permission_classes = [permissions.IsAuthenticated]
def post(self, request, project_id):
try:
# 获取项目实例
project = Project.objects.get(id=project_id, user=request.user)
except Project.DoesNotExist:
return Response(
{"error": {"code": "project_not_found", "message": "项目不存在"}},
status=status.HTTP_404_NOT_FOUND
)
# 获取上传的文件
if 'files' not in request.FILES:
return Response(
{"error": {"code": "no_file", "message": "未提供数据文件"}},
status=status.HTTP_400_BAD_REQUEST
)
uploaded_file = request.FILES['files']
# 验证是否为zip文件
if not uploaded_file.name.lower().endswith('.zip'):
return Response(
{"error": {"code": "invalid_file", "message": "请上传zip格式的压缩文件"}},
status=status.HTTP_400_BAD_REQUEST
)
# 获取group_index参数
group_index = request.data.get('group_index', '2')
try:
group_index = int(group_index)
if group_index <= 0:
return Response(
{"error": {"code": "invalid_group_index", "message": "分组索引必须是大于0的整数"}},
status=status.HTTP_400_BAD_REQUEST
)
except ValueError:
return Response(
{"error": {"code": "invalid_group_index", "message": "分组索引必须是整数"}},
status=status.HTTP_400_BAD_REQUEST
)
# 获取波数范围参数
wavenumber_range = None
if 'Wavenumber range' in request.data and request.data['Wavenumber range'].strip():
try:
wavenumber_values = request.data['Wavenumber range'].split(',')
if len(wavenumber_values) == 2:
wavenumber_range = [float(wavenumber_values[0]), float(wavenumber_values[1])]
else:
return Response(
{"error": {"code": "invalid_wavenumber", "message": "波数范围格式无效,应为两个数值,如'500,3150'"}},
status=status.HTTP_400_BAD_REQUEST
)
except (ValueError, IndexError):
return Response(
{"error": {"code": "invalid_wavenumber", "message": "波数范围格式无效,应为两个数值,如'500,3150'"}},
status=status.HTTP_400_BAD_REQUEST
)
# 保存上传的文件
data_file = DataFile(
project=project,
file=uploaded_file,
filename=uploaded_file.name,
group_index=group_index # 保存group_index
)
if wavenumber_range:
data_file.wavenumber_min = wavenumber_range[0]
data_file.wavenumber_max = wavenumber_range[1]
data_file.save()
# 如果这是项目的第一个数据文件,将其设置为活动文件
if project.data_files.count() == 1:
project.active_data_file = data_file
project.save()
# 使用R处理上传的文件
try:
# 解压缩并处理文件
result = RExecutor.process_upload(
project_id=project_id,
file_path=data_file.file.path,
wavenumber_range=wavenumber_range,
group_index=group_index, # 添加group_index参数
is_zip=True # 标记为zip文件
)
# 如果处理失败,返回错误
if result.get('status') == 'error':
# 删除保存的文件
data_file.file.delete()
data_file.delete()
return Response(
{"error": {"code": "processing_error", "message": result.get('message', '处理文件时出错')}},
status=status.HTTP_500_INTERNAL_SERVER_ERROR
)
# 获取分组信息
group_names = result.get('group_names', [])
if project.active_data_file_id and project.status == 'initialized':
project.status = 'data_loaded'
project.save(update_fields=['status'])
# 添加next_steps字段
response_data = {
"status": "success",
"group_names": group_names, # 在响应中包含分组信息
"next_steps": [
{
"step": "preprocessing",
"url": f"/api/projects/{project_id}/analysis/",
"description": "数据预处理",
"frontend_url": f"/projects/{project_id}/analysis/"
}
],
"message": "数据上传成功,请确认分组信息"
}
return Response(response_data, status=status.HTTP_200_OK)
except Exception as e:
# 处理失败,删除保存的文件
data_file.file.delete()
data_file.delete()
return Response(
{"error": {"code": "processing_error", "message": f"处理文件时出错: {str(e)}"}},
status=status.HTTP_500_INTERNAL_SERVER_ERROR
)
class GroupInfoView(APIView):
"""获取分组信息视图"""
permission_classes = [permissions.IsAuthenticated]
def get(self, request, project_id):
try:
# 获取项目实例
project = Project.objects.get(id=project_id, user=request.user)
except Project.DoesNotExist:
return Response(
{"error": {"code": "project_not_found", "message": "项目不存在"}},
status=status.HTTP_404_NOT_FOUND
)
# 检查项目是否有活动数据文件
if not project.active_data_file:
return Response(
{"error": {"code": "no_active_data", "message": "项目没有活动数据文件,请先上传并选择数据"}},
status=status.HTTP_400_BAD_REQUEST
)
# 获取分组信息
try:
result = RExecutor.get_group_info(project_id=project_id)
# 如果处理失败,返回错误
if result.get('status') == 'error':
return Response(
{"error": {"code": "processing_error", "message": result.get('message', '获取分组信息时出错')}},
status=status.HTTP_500_INTERNAL_SERVER_ERROR
)
# 返回分组信息
return Response(
{
"status": "success",
"group_names": result.get('group_names', [])
},
status=status.HTTP_200_OK
)
except Exception as e:
import traceback
error_details = traceback.format_exc()
logger = logging.getLogger(__name__)
logger.error(f"分组信息错误: {str(e)}\n{error_details}")
return Response(
{"error": {"code": "processing_error", "message": f"获取分组信息时出错: {str(e)}"}},
status=status.HTTP_500_INTERNAL_SERVER_ERROR
)
class SelectDataFileView(APIView):
"""选择活动数据文件视图"""
permission_classes = [permissions.IsAuthenticated]
def post(self, request, project_id, file_id):
try:
# 获取项目实例
project = Project.objects.get(id=project_id, user=request.user)
except Project.DoesNotExist:
return Response(
{"error": {"code": "project_not_found", "message": "项目不存在"}},
status=status.HTTP_404_NOT_FOUND
)
try:
# 获取数据文件实例
data_file = DataFile.objects.get(id=file_id, project=project)
except DataFile.DoesNotExist:
return Response(
{"error": {"code": "file_not_found", "message": "数据文件不存在"}},
status=status.HTTP_404_NOT_FOUND
)
# 设置为当前活动文件
project.active_data_file = data_file
# 选定活动数据文件后,项目应回到可执行预处理/分析的状态
project.status = 'data_loaded'
project.save()
# 返回成功响应
return Response(
{
"status": "success",
"message": f"已将文件 '{data_file.filename}' 设置为当前活动文件",
"active_file": {
"id": data_file.id,
"filename": data_file.filename,
"uploaded_at": data_file.uploaded_at
}
},
status=status.HTTP_200_OK
)
class UpdateDataFileView(APIView):
"""更新数据文件视图"""
permission_classes = [permissions.IsAuthenticated]
def post(self, request, project_id, file_id):
try:
# 获取项目实例
project = Project.objects.get(id=project_id, user=request.user)
except Project.DoesNotExist:
return Response(
{"error": {"code": "project_not_found", "message": "项目不存在"}},
status=status.HTTP_404_NOT_FOUND
)
try:
# 获取数据文件实例
data_file = DataFile.objects.get(id=file_id, project=project)
except DataFile.DoesNotExist:
return Response(
{"error": {"code": "file_not_found", "message": "数据文件不存在"}},
status=status.HTTP_404_NOT_FOUND
)
# 获取波数范围参数
if 'wavenumber_range' in request.data:
try:
wavenumber_values = request.data['wavenumber_range'].split(',')
data_file.wavenumber_min = float(wavenumber_values[0])
data_file.wavenumber_max = float(wavenumber_values[1])
data_file.save()
except (ValueError, IndexError):
return Response(
{"error": {"code": "invalid_wavenumber", "message": "波数范围格式无效,应为两个数值,如'500,3150'"}},
status=status.HTTP_400_BAD_REQUEST
)
# 如果这是当前活动的数据文件,需要重新生成ramex_data.rds文件
if project.active_data_file == data_file:
try:
# 调用update_wavenumber_range方法重新生成ramex_data.rds文件
result = RExecutor.update_wavenumber_range(
project_id=project_id,
file_id=file_id,
wavenumber_range=[data_file.wavenumber_min, data_file.wavenumber_max]
)
if result["status"] != "success":
return Response(
{"error": {"code": "processing_error", "message": result["message"]}},
status=status.HTTP_500_INTERNAL_SERVER_ERROR
)
except Exception as e:
return Response(
{"error": {"code": "processing_error", "message": f"更新波数范围时出错: {str(e)}"}},
status=status.HTTP_500_INTERNAL_SERVER_ERROR
)
# 返回成功响应
return Response(
{
"status": "success",
"message": f"已成功更新文件 '{data_file.filename}' 的波数范围",
"data_file": {
"id": data_file.id,
"filename": data_file.filename,
"wavenumber_min": data_file.wavenumber_min,
"wavenumber_max": data_file.wavenumber_max
}
},
status=status.HTTP_200_OK
)
class DeleteDataFileView(APIView):
"""删除数据文件视图"""
permission_classes = [permissions.IsAuthenticated]
def delete(self, request, project_id, file_id):
try:
# 获取项目实例
project = Project.objects.get(id=project_id, user=request.user)
except Project.DoesNotExist:
return Response(
{"error": {"code": "project_not_found", "message": "项目不存在"}},
status=status.HTTP_404_NOT_FOUND
)
try:
# 获取数据文件实例
data_file = DataFile.objects.get(id=file_id, project=project)
except DataFile.DoesNotExist:
return Response(
{"error": {"code": "file_not_found", "message": "数据文件不存在"}},
status=status.HTTP_404_NOT_FOUND
)
# 检查是否是活动文件
if project.active_data_file == data_file:
# 将活动文件设置为None
project.active_data_file = None
project.save()
# 如果项目只有这一个文件并且已删除,将项目状态重置为初始化
if project.data_files.count() == 1:
project.status = 'initialized'
project.save()
# 删除实际的文件
if data_file.file:
try:
if os.path.exists(data_file.file.path):
os.remove(data_file.file.path)
except Exception as e:
# 即使文件删除失败,也继续删除数据库记录
pass
# 删除数据库记录
data_file.delete()
# 返回成功响应
return Response(
{"status": "success", "message": f"数据文件 '{data_file.filename}' 已被删除"},
status=status.HTTP_200_OK
)
class SaveRdsDataView(APIView):
"""保存确认的RDS数据视图"""
permission_classes = [permissions.IsAuthenticated]
def post(self, request, project_id):
try:
# 获取项目实例
project = Project.objects.get(id=project_id, user=request.user)
except Project.DoesNotExist:
return Response(
{"error": {"code": "project_not_found", "message": "项目不存在"}},
status=status.HTTP_404_NOT_FOUND
)
# 检查项目是否有活动数据文件
if not project.active_data_file:
return Response(
{"error": {"code": "no_active_data", "message": "项目没有活动数据文件,请先上传并选择数据"}},
status=status.HTTP_400_BAD_REQUEST
)
# 获取请求中的分组顺序 - 直接从request.data中获取,避免直接访问request.body
group_order = None
if hasattr(request, 'data') and isinstance(request.data, dict):
group_order = request.data.get('group_order')
# 调用R执行器的save_rds方法,传入分组顺序
result = RExecutor.save_rds(project_id, group_order)
if result.get('status') != 'success':
return Response({
"status": "error",
"error": {"message": result.get('message', '保存RDS数据时出错')}
}, status=status.HTTP_500_INTERNAL_SERVER_ERROR)
# 获取用户ID和项目目录
user_id = project.user.id
project_dir = os.path.join(settings.MEDIA_ROOT, f'users/{user_id}/projects/{project_id}/data')
# 检查临时RamEx数据文件是否存在
ramex_data_path = os.path.join(project_dir, 'ramex_data.rds')
if not os.path.exists(ramex_data_path):
return Response(
{"error": {"code": "no_data", "message": "未找到RamEx数据文件,请重新上传数据"}},
status=status.HTTP_400_BAD_REQUEST
)
# 确认分组后,项目应处于可执行预处理/分析的状态
project.status = 'data_loaded'
project.save()
# 返回成功响应
return Response({
"status": "success",
"message": "分组信息已确认,数据已保存",
"next_url": f"/projects/{project_id}/analysis/"
}, status=status.HTTP_200_OK)