model_train_dm/backend/apps/datasets/views/dataset_sample.py

218 lines
7.8 KiB
Python
Raw Permalink Normal View History

2026-07-27 17:51:49 +08:00
import hashlib
import base64
import mimetypes
from django.db import transaction
from django.http import StreamingHttpResponse
from django.shortcuts import get_object_or_404, get_list_or_404
from rest_framework import status
from rest_framework.decorators import api_view, permission_classes
from rest_framework.pagination import PageNumberPagination
from rest_framework.permissions import IsAuthenticated
from rest_framework.response import Response
from apps.common.serializers import AiDatasetSampleSerializer
from apps.common.minio_client import minio_storage
from apps.core.models import AiDataset, AiDatasetSample
from apps.common.utils.annotation_store import get_dataset_sample_annotation_content
from apps.datasets.dataset_stats import refresh_dataset_sample_counters
from apps.datasets.storage_service import DATASET_BUCKET, delete_sample_storage
def _sample_object_name(sample: AiDatasetSample) -> str:
return f"{str(sample.saved_path or '').strip('/').replace('\\', '/')}/{sample.saved_filename}".strip("/")
def _stream_object(response, chunk_size: int = 1024 * 1024):
try:
while True:
chunk = response.read(chunk_size)
if not chunk:
break
yield chunk
finally:
response.close()
response.release_conn()
def _build_sample_stream_response(sample: AiDatasetSample, as_attachment: bool):
object_name = _sample_object_name(sample)
response = minio_storage.get_object(DATASET_BUCKET, object_name)
content_type = (
mimetypes.guess_type(sample.original_filename or sample.saved_filename)[0]
or "application/octet-stream"
)
filename = sample.original_filename or sample.saved_filename
stream = StreamingHttpResponse(_stream_object(response), content_type=content_type)
disposition = "attachment" if as_attachment else "inline"
stream["Content-Disposition"] = f'{disposition}; filename="{filename}"'
return stream
def _get_sample_md5(sample: AiDatasetSample) -> str | None:
object_name = _sample_object_name(sample)
try:
info = minio_storage.get_object_info(DATASET_BUCKET, object_name)
etag = str(info.get("etag") or "").strip('"')
if etag:
return etag
file_bytes = minio_storage.download_bytes(DATASET_BUCKET, object_name)
return hashlib.md5(file_bytes).hexdigest()
except Exception as e:
print(f"Error reading sample {sample.id} from MinIO: {e}")
return None
@api_view(["GET"])
@permission_classes([IsAuthenticated])
def search_dataset_samples(request):
dataset_id = request.GET.get("dataset_id", None)
original_filename = request.GET.get("original_filename", None)
start_time = request.GET.get("start_time", None)
end_time = request.GET.get("end_time", None)
status = request.GET.get("status", None)
page_size = request.GET.get("page_size", None)
filter_kwargs = {}
if dataset_id:
filter_kwargs["dataset_id"] = dataset_id
if original_filename:
filter_kwargs["original_filename__contains"] = original_filename
if start_time and end_time:
filter_kwargs["create_time__range"] = [start_time, end_time]
if start_time and end_time is None:
filter_kwargs["create_time__gte"] = start_time
if start_time is None and end_time:
filter_kwargs["create_time__lte"] = end_time
if status:
filter_kwargs["status"] = status
samples = AiDatasetSample.objects.filter(**filter_kwargs).order_by("saved_filename")
paginator = PageNumberPagination()
paginator.page_size = page_size or 10
result_page = paginator.paginate_queryset(samples, request)
serializer = AiDatasetSampleSerializer(result_page, many=True)
return paginator.get_paginated_response(serializer.data)
@api_view(["DELETE"])
@permission_classes([IsAuthenticated])
@transaction.atomic
def delete_invalid_samples(request):
dataset_id = request.GET.get("dataset_id", None)
sample_ids = request.GET.get("ids", None)
if sample_ids is None:
return Response({"message": "未提供样本 IDs"}, status=400)
sample_ids_list = [id for id in sample_ids.split(",")]
samples = get_list_or_404(AiDatasetSample, id__in=sample_ids_list)
for sample in samples:
delete_sample_storage(sample.saved_path, sample.saved_filename)
sample.delete()
refresh_dataset_sample_counters(dataset_id)
return Response({"message": "无效样本删除成功!"})
@api_view(["DELETE"])
@permission_classes([IsAuthenticated])
@transaction.atomic
def delete_duplicate_samples(request):
dataset_id = request.GET.get("dataset_id", None)
if dataset_id is None:
return Response({"message": "未选择数据集"}, status=400)
samples = AiDatasetSample.objects.filter(dataset_id=dataset_id).order_by("create_time", "id")
if not samples.exists():
return Response({"message": "当前数据集没有样本,无需去重!"})
unique_images = {}
deleted_count = 0
failed_samples = []
for sample in samples:
md5_hash = _get_sample_md5(sample)
if not md5_hash:
failed_samples.append(sample.id)
continue
if md5_hash in unique_images:
try:
failed_objects = delete_sample_storage(sample.saved_path, sample.saved_filename)
if failed_objects:
failed_samples.append(sample.id)
continue
sample.delete()
deleted_count += 1
except Exception:
failed_samples.append(sample.id)
else:
unique_images[md5_hash] = sample.id
refresh_dataset_sample_counters(dataset_id)
if failed_samples:
return Response(
{
"message": f'样本去重完成,共删除了"{deleted_count}"条重复数据,"{len(failed_samples)}"条样本处理失败!',
"failed_ids": failed_samples,
}
)
return Response({"message": f'样本去重完成,共删除了"{deleted_count}"条重复数据!'})
@api_view(["PUT"])
@permission_classes([IsAuthenticated])
def update_dataset_sample(request):
data = request.data
sample = get_object_or_404(AiDatasetSample, id=data.get("id"))
serializer = AiDatasetSampleSerializer(instance=sample, data=data, partial=True)
if serializer.is_valid():
serializer.save()
refresh_dataset_sample_counters(sample.dataset_id)
return Response({"message": "样本更新成功"})
return Response(
{"error": "无效的数据", "details": serializer.errors},
status=status.HTTP_400_BAD_REQUEST,
)
@api_view(["GET"])
@permission_classes([IsAuthenticated])
def read_dataset_sample(request):
sample_id = request.GET.get("id", None)
sample = get_object_or_404(AiDatasetSample, id=sample_id)
object_name = _sample_object_name(sample)
file_bytes = minio_storage.download_bytes(DATASET_BUCKET, object_name)
content_type = (
mimetypes.guess_type(sample.original_filename or sample.saved_filename)[0]
or "application/octet-stream"
)
image_base64 = f"data:{content_type};base64,{base64.b64encode(file_bytes).decode('utf-8')}"
response_data = {
"id": sample.id,
"image_base64": image_base64,
"annotation_content": get_dataset_sample_annotation_content(sample),
}
return Response(response_data)
@api_view(["GET"])
@permission_classes([IsAuthenticated])
def preview_dataset_sample(request):
sample_id = request.GET.get("id", None)
sample = get_object_or_404(AiDatasetSample, id=sample_id)
return _build_sample_stream_response(sample, as_attachment=False)
@api_view(["GET"])
@permission_classes([IsAuthenticated])
def download_dataset_sample(request):
sample_id = request.GET.get("id", None)
sample = get_object_or_404(AiDatasetSample, id=sample_id)
return _build_sample_stream_response(sample, as_attachment=True)