218 lines
7.8 KiB
Python
218 lines
7.8 KiB
Python
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)
|