cvat-ai--cvat
1180 行
43 KiB
Python
1180 行
43 KiB
Python
# Copyright (C) 2020-2022 Intel Corporation
|
|
# Copyright (C) CVAT.ai Corporation
|
|
#
|
|
# SPDX-License-Identifier: MIT
|
|
|
|
from __future__ import annotations
|
|
|
|
import io
|
|
import os
|
|
import os.path
|
|
import pickle # nosec
|
|
import tempfile
|
|
import time
|
|
import zipfile
|
|
import zlib
|
|
from collections.abc import Callable, Collection, Generator, Iterator, Sequence
|
|
from contextlib import ExitStack, closing
|
|
from datetime import datetime, timezone
|
|
from itertools import groupby, pairwise
|
|
from pathlib import Path, PurePath
|
|
from typing import Any, TypeAlias, overload
|
|
|
|
import attrs
|
|
import av
|
|
import django_rq
|
|
import PIL.Image
|
|
import PIL.ImageOps
|
|
import rq
|
|
from django.conf import settings
|
|
from django.core.cache import caches
|
|
from django.db import models as django_models
|
|
from django.utils import timezone as django_tz
|
|
from redis.exceptions import LockError
|
|
from rest_framework.exceptions import NotFound, ValidationError
|
|
from rq.job import JobStatus as RQJobStatus
|
|
|
|
from cvat.apps.engine import models
|
|
from cvat.apps.engine.cache_signals import cache_item_created_signal, cache_item_read_signal
|
|
from cvat.apps.engine.cloud_provider import db_storage_to_storage_instance
|
|
from cvat.apps.engine.log import ServerLogManager
|
|
from cvat.apps.engine.media_extractors import (
|
|
ImageReaderWithManifest,
|
|
ValidateDimension,
|
|
VideoReader,
|
|
VideoReaderWithManifest,
|
|
ZipCompressedChunkWriter,
|
|
load_image,
|
|
)
|
|
from cvat.apps.engine.rq import RQMetaWithFailureInfo
|
|
from cvat.apps.engine.utils import (
|
|
CvatChunkTimestampMismatchError,
|
|
format_list,
|
|
get_rq_lock_for_job,
|
|
md5_hash,
|
|
)
|
|
from cvat.utils import django_database as db_utils
|
|
from cvat.utils.paths import join_untrusted_path
|
|
from utils.dataset_manifest import ImageManifestManager
|
|
from utils.dataset_manifest.utils import Openable
|
|
|
|
slogger = ServerLogManager(__name__)
|
|
|
|
|
|
DataWithMime: TypeAlias = tuple[io.BytesIO, str]
|
|
_CacheItem: TypeAlias = tuple[io.BytesIO, str, int, datetime | None]
|
|
_RQ_JOB_ORIGIN_ATTRIBUTE = "origin"
|
|
|
|
ASSETS_DIR = Path(__file__).parent / "assets"
|
|
|
|
|
|
class CacheTooLargeDataError(Exception):
|
|
pass
|
|
|
|
|
|
class ChunkCreationError(Exception):
|
|
pass
|
|
|
|
|
|
def _build_chunk_job_failure_exception(
|
|
rq_job: rq.job.Job, job_meta: RQMetaWithFailureInfo
|
|
) -> Exception:
|
|
exc_type = job_meta.exc_type or ChunkCreationError
|
|
exc_args = job_meta.exc_args or ("Cannot create chunk",)
|
|
|
|
try:
|
|
return exc_type(*exc_args)
|
|
except TypeError:
|
|
exc_name = exc_type.__name__
|
|
details = job_meta.formatted_exception or repr(exc_args)
|
|
return ChunkCreationError(
|
|
f"Chunk job {rq_job.id} failed with {exc_name}: {details.strip()}"
|
|
)
|
|
|
|
|
|
def enqueue_create_chunk_job(
|
|
queue: rq.Queue,
|
|
rq_job_id: str,
|
|
create_callback: Callback,
|
|
*,
|
|
rq_job_result_ttl: int = 60,
|
|
rq_job_failure_ttl: int = 3600 * 24 * 14, # 2 weeks
|
|
) -> rq.job.Job:
|
|
try:
|
|
with get_rq_lock_for_job(queue, rq_job_id):
|
|
rq_job = queue.fetch_job(rq_job_id)
|
|
|
|
if not rq_job or (
|
|
# Enqueue the job if the chunk was deleted but the RQ job still exists.
|
|
# This can happen in cases involving jobs with honeypots and
|
|
# if the job wasn't collected by the requesting process for any reason.
|
|
rq_job.get_status(refresh=False)
|
|
in {RQJobStatus.FINISHED, RQJobStatus.FAILED, RQJobStatus.CANCELED}
|
|
):
|
|
rq_job = queue.enqueue(
|
|
create_callback,
|
|
job_id=rq_job_id,
|
|
result_ttl=rq_job_result_ttl,
|
|
failure_ttl=rq_job_failure_ttl,
|
|
)
|
|
except LockError:
|
|
raise TimeoutError(f"Cannot acquire lock for {rq_job_id}")
|
|
|
|
return rq_job
|
|
|
|
|
|
def wait_for_rq_job(rq_job: rq.job.Job):
|
|
retries = settings.CVAT_CHUNK_CREATE_TIMEOUT // settings.CVAT_CHUNK_CREATE_CHECK_INTERVAL or 1
|
|
while retries > 0:
|
|
job_status = rq_job.get_status()
|
|
if job_status in ("finished",):
|
|
return
|
|
elif job_status in ("failed",):
|
|
rq_job.get_meta() # refresh from Redis
|
|
job_meta = RQMetaWithFailureInfo.for_job(rq_job)
|
|
raise _build_chunk_job_failure_exception(rq_job, job_meta)
|
|
|
|
time.sleep(settings.CVAT_CHUNK_CREATE_CHECK_INTERVAL)
|
|
retries -= 1
|
|
|
|
raise TimeoutError(f"Chunk processing takes too long {rq_job.id}")
|
|
|
|
|
|
def _is_run_inside_rq() -> bool:
|
|
return rq.get_current_job() is not None
|
|
|
|
|
|
def _get_current_rq_queue_name() -> str | None:
|
|
return getattr(rq.get_current_job(), _RQ_JOB_ORIGIN_ATTRIBUTE, None)
|
|
|
|
|
|
def _convert_args_for_callback(func_args: list[Any]) -> list[Any]:
|
|
result = []
|
|
for func_arg in func_args:
|
|
if _is_run_inside_rq():
|
|
result.append(func_arg)
|
|
else:
|
|
if isinstance(
|
|
func_arg,
|
|
django_models.Model,
|
|
):
|
|
result.append(func_arg.id)
|
|
elif isinstance(func_arg, list):
|
|
result.append(_convert_args_for_callback(func_arg))
|
|
else:
|
|
result.append(func_arg)
|
|
|
|
return result
|
|
|
|
|
|
@attrs.frozen
|
|
class Callback:
|
|
_callable: Callable[..., DataWithMime] = attrs.field(
|
|
validator=attrs.validators.is_callable(),
|
|
)
|
|
_args: list[Any] = attrs.field(
|
|
factory=list,
|
|
validator=attrs.validators.instance_of(list),
|
|
converter=_convert_args_for_callback,
|
|
)
|
|
_kwargs: dict[str, bool | int | float | str | None] = attrs.field(
|
|
factory=dict,
|
|
validator=attrs.validators.deep_mapping(
|
|
key_validator=attrs.validators.instance_of(str),
|
|
value_validator=attrs.validators.instance_of((bool, int, float, str, type(None))),
|
|
mapping_validator=attrs.validators.instance_of(dict),
|
|
),
|
|
)
|
|
|
|
def __call__(self) -> DataWithMime:
|
|
return self._callable(*self._args, **self._kwargs)
|
|
|
|
|
|
class MediaCache:
|
|
_QUEUE_NAME = settings.CVAT_QUEUES.CHUNKS.value
|
|
_QUEUE_JOB_PREFIX_TASK = "chunks:prepare-item-"
|
|
_CACHE_NAME = "media"
|
|
_PREVIEW_TTL = settings.CVAT_PREVIEW_CACHE_TTL
|
|
|
|
@staticmethod
|
|
def _cache():
|
|
return caches[MediaCache._CACHE_NAME]
|
|
|
|
@staticmethod
|
|
def _get_checksum(value: bytes) -> int:
|
|
return zlib.crc32(value)
|
|
|
|
@staticmethod
|
|
def _get_cache_item_size(item: _CacheItem) -> int:
|
|
return item[0].getbuffer().nbytes
|
|
|
|
def _get_or_set_cache_item(
|
|
self,
|
|
key: str,
|
|
create_callback: Callback,
|
|
*,
|
|
cache_item_ttl: int | None = None,
|
|
) -> _CacheItem:
|
|
item = self._get_cache_item(key)
|
|
if item:
|
|
return item
|
|
|
|
return self._create_cache_item(
|
|
key,
|
|
create_callback,
|
|
cache_item_ttl=cache_item_ttl,
|
|
)
|
|
|
|
@classmethod
|
|
def _get_queue(cls) -> rq.Queue:
|
|
return django_rq.get_queue(cls._QUEUE_NAME)
|
|
|
|
@classmethod
|
|
def _make_queue_job_id(cls, key: str) -> str:
|
|
return f"{cls._QUEUE_JOB_PREFIX_TASK}{key}"
|
|
|
|
@staticmethod
|
|
def _drop_return_value(func: Callable[..., DataWithMime], *args: Any, **kwargs: Any):
|
|
func(*args, **kwargs)
|
|
|
|
@classmethod
|
|
def _create_and_set_cache_item(
|
|
cls,
|
|
key: str,
|
|
create_callback: Callback,
|
|
cache_item_ttl: int | None = None,
|
|
) -> DataWithMime:
|
|
timestamp = django_tz.now()
|
|
item_data = create_callback()
|
|
item_data_bytes = item_data[0].getvalue()
|
|
item = (item_data[0], item_data[1], cls._get_checksum(item_data_bytes), timestamp)
|
|
|
|
# allow empty data to be set in cache to prevent
|
|
# future rq jobs from being enqueued to prepare the item
|
|
cache = cls._cache()
|
|
with get_rq_lock_for_job(
|
|
cls._get_queue(),
|
|
key,
|
|
):
|
|
cached_item = cache.get(key)
|
|
if cached_item is not None:
|
|
cache_item_read_signal.send(
|
|
sender=cls,
|
|
item_key=key,
|
|
item_data_size=cls._get_cache_item_size(cached_item),
|
|
rq_queue=_get_current_rq_queue_name(),
|
|
)
|
|
|
|
if timestamp <= cached_item[3]:
|
|
item = cached_item
|
|
if cached_item is None or timestamp > cached_item[3]:
|
|
item_size = cls._get_cache_item_size(item)
|
|
if item_size > settings.CVAT_CACHE_ITEM_MAX_SIZE:
|
|
raise CacheTooLargeDataError(
|
|
f"Chunk data size {item_size} exceeds the maximum allowed size "
|
|
f"{settings.CVAT_CACHE_ITEM_MAX_SIZE}."
|
|
)
|
|
cache.set(key, item, timeout=cache_item_ttl or cache.default_timeout)
|
|
|
|
cache_item_created_signal.send(
|
|
sender=cls,
|
|
item_key=key,
|
|
item_data_size=item_size,
|
|
rq_queue=_get_current_rq_queue_name(),
|
|
)
|
|
|
|
return item
|
|
|
|
def _create_cache_item(
|
|
self,
|
|
key: str,
|
|
create_callback: Callback,
|
|
*,
|
|
cache_item_ttl: int | None = None,
|
|
) -> _CacheItem:
|
|
slogger.glob.info(f"Starting to prepare chunk: key {key}")
|
|
if _is_run_inside_rq():
|
|
item = self._create_and_set_cache_item(
|
|
key,
|
|
create_callback,
|
|
cache_item_ttl=cache_item_ttl,
|
|
)
|
|
else:
|
|
rq_job = enqueue_create_chunk_job(
|
|
queue=self._get_queue(),
|
|
rq_job_id=self._make_queue_job_id(key),
|
|
create_callback=Callback(
|
|
callable=self._drop_return_value,
|
|
args=[
|
|
self._create_and_set_cache_item,
|
|
key,
|
|
create_callback,
|
|
],
|
|
kwargs={
|
|
"cache_item_ttl": cache_item_ttl,
|
|
},
|
|
),
|
|
)
|
|
wait_for_rq_job(rq_job)
|
|
item = self._get_cache_item(key)
|
|
|
|
slogger.glob.info(f"Ending to prepare chunk: key {key}")
|
|
|
|
return item
|
|
|
|
def _delete_cache_item(self, key: str):
|
|
self._cache().delete(key)
|
|
slogger.glob.info(f"Removed the cache key {key}")
|
|
|
|
def _bulk_delete_cache_items(self, keys: Sequence[str]):
|
|
self._cache().delete_many(keys)
|
|
slogger.glob.info(f"Removed the cache keys {format_list(keys)}")
|
|
|
|
def _get_cache_item(self, key: str) -> _CacheItem | None:
|
|
rq_queue = _get_current_rq_queue_name()
|
|
try:
|
|
item = self._cache().get(key)
|
|
except pickle.UnpicklingError:
|
|
slogger.glob.error(f"Unable to get item from cache: key {key}", exc_info=True)
|
|
return None
|
|
|
|
if not item:
|
|
return None
|
|
|
|
item_data = item[0].getbuffer() if isinstance(item[0], io.BytesIO) else item[0]
|
|
item_checksum = item[2] if len(item) == 4 else None
|
|
cache_item_read_signal.send(
|
|
sender=self.__class__,
|
|
item_key=key,
|
|
item_data_size=self._get_cache_item_size(item),
|
|
rq_queue=rq_queue,
|
|
)
|
|
|
|
if item_checksum != self._get_checksum(item_data):
|
|
slogger.glob.info(f"Cache item {key} checksum mismatch")
|
|
return None
|
|
|
|
return item
|
|
|
|
def _validate_cache_item_timestamp(
|
|
self, item: _CacheItem, expected_timestamp: datetime
|
|
) -> _CacheItem:
|
|
if item[3] < expected_timestamp:
|
|
raise CvatChunkTimestampMismatchError(
|
|
f"Cache timestamp mismatch. Item_ts: {item[3]}, expected_ts: {expected_timestamp}"
|
|
)
|
|
|
|
return item
|
|
|
|
@classmethod
|
|
def _has_key(cls, key: str) -> bool:
|
|
return cls._cache().has_key(key)
|
|
|
|
@staticmethod
|
|
def _make_cache_key_prefix(
|
|
obj: models.Task | models.Segment | models.Job | models.CloudStorage,
|
|
) -> str:
|
|
if isinstance(obj, models.Task):
|
|
return f"task_{obj.id}"
|
|
elif isinstance(obj, models.Segment):
|
|
return f"segment_{obj.id}"
|
|
elif isinstance(obj, models.Job):
|
|
return f"job_{obj.id}"
|
|
elif isinstance(obj, models.CloudStorage):
|
|
return f"cloudstorage_{obj.id}"
|
|
else:
|
|
assert False, f"Unexpected object type {type(obj)}"
|
|
|
|
@classmethod
|
|
def _make_chunk_key(
|
|
cls,
|
|
db_obj: models.Task | models.Segment | models.Job,
|
|
chunk_number: int,
|
|
*,
|
|
quality: models.FrameQuality,
|
|
) -> str:
|
|
return f"{cls._make_cache_key_prefix(db_obj)}_chunk_{chunk_number}_{quality}"
|
|
|
|
def _make_preview_key(self, db_obj: models.Segment | models.CloudStorage) -> str:
|
|
return f"{self._make_cache_key_prefix(db_obj)}_preview"
|
|
|
|
def _make_segment_task_chunk_key(
|
|
self,
|
|
db_obj: models.Segment,
|
|
chunk_number: int,
|
|
*,
|
|
quality: models.FrameQuality,
|
|
) -> str:
|
|
return f"{self._make_cache_key_prefix(db_obj)}_task_chunk_{chunk_number}_{quality}"
|
|
|
|
def _make_frame_context_images_chunk_key(self, db_data: models.Data, frame_number: int) -> str:
|
|
return f"context_images_{db_data.id}_{frame_number}"
|
|
|
|
@overload
|
|
def _to_data_with_mime(self, cache_item: _CacheItem) -> DataWithMime: ...
|
|
|
|
@overload
|
|
def _to_data_with_mime(
|
|
self, cache_item: _CacheItem | None, *, allow_none: bool = False
|
|
) -> DataWithMime | None: ...
|
|
|
|
def _to_data_with_mime(
|
|
self, cache_item: _CacheItem | None, *, allow_none: bool = False
|
|
) -> DataWithMime | None:
|
|
if not cache_item:
|
|
if allow_none:
|
|
return None
|
|
|
|
raise ValueError("A cache item is not allowed to be None")
|
|
|
|
return cache_item[:2]
|
|
|
|
def get_or_set_segment_chunk(
|
|
self, db_segment: models.Segment, chunk_number: int, *, quality: models.FrameQuality
|
|
) -> DataWithMime:
|
|
|
|
item = self._get_or_set_cache_item(
|
|
self._make_chunk_key(db_segment, chunk_number, quality=quality),
|
|
Callback(
|
|
callable=self.prepare_segment_chunk,
|
|
args=[db_segment, chunk_number],
|
|
kwargs={"quality": quality},
|
|
),
|
|
)
|
|
db_segment.refresh_from_db(fields=["chunks_updated_date"])
|
|
|
|
return self._to_data_with_mime(
|
|
self._validate_cache_item_timestamp(item, db_segment.chunks_updated_date)
|
|
)
|
|
|
|
def get_task_chunk(
|
|
self, db_task: models.Task, chunk_number: int, *, quality: models.FrameQuality
|
|
) -> DataWithMime | None:
|
|
return self._to_data_with_mime(
|
|
self._get_cache_item(
|
|
key=self._make_chunk_key(db_task, chunk_number, quality=quality),
|
|
),
|
|
allow_none=True,
|
|
)
|
|
|
|
def get_or_set_task_chunk(
|
|
self,
|
|
db_task: models.Task,
|
|
chunk_number: int,
|
|
set_callback: Callback,
|
|
*,
|
|
quality: models.FrameQuality,
|
|
) -> DataWithMime:
|
|
|
|
item = self._get_or_set_cache_item(
|
|
self._make_chunk_key(db_task, chunk_number, quality=quality),
|
|
set_callback,
|
|
)
|
|
|
|
if db_utils.is_field_cached(db_task, "segment_set"):
|
|
# Refresh segments to report actual dates if they were fetched previously
|
|
# Doing so without a check leads to an error if the related object is not prefetched
|
|
db_task.refresh_from_db(fields=["segment_set"])
|
|
|
|
return self._to_data_with_mime(
|
|
self._validate_cache_item_timestamp(item, db_task.get_chunks_updated_date())
|
|
)
|
|
|
|
def get_segment_task_chunk(
|
|
self, db_segment: models.Segment, chunk_number: int, *, quality: models.FrameQuality
|
|
) -> DataWithMime | None:
|
|
return self._to_data_with_mime(
|
|
self._get_cache_item(
|
|
key=self._make_segment_task_chunk_key(db_segment, chunk_number, quality=quality),
|
|
),
|
|
allow_none=True,
|
|
)
|
|
|
|
def get_or_set_segment_task_chunk(
|
|
self,
|
|
db_segment: models.Segment,
|
|
chunk_number: int,
|
|
*,
|
|
quality: models.FrameQuality,
|
|
set_callback: Callback,
|
|
) -> DataWithMime:
|
|
|
|
item = self._get_or_set_cache_item(
|
|
self._make_segment_task_chunk_key(db_segment, chunk_number, quality=quality),
|
|
set_callback,
|
|
)
|
|
db_segment.refresh_from_db(fields=["chunks_updated_date"])
|
|
|
|
return self._to_data_with_mime(
|
|
self._validate_cache_item_timestamp(item, db_segment.chunks_updated_date),
|
|
)
|
|
|
|
def get_or_set_selective_job_chunk(
|
|
self, db_job: models.Job, chunk_number: int, *, quality: models.FrameQuality
|
|
) -> DataWithMime:
|
|
return self._to_data_with_mime(
|
|
self._get_or_set_cache_item(
|
|
self._make_chunk_key(db_job, chunk_number, quality=quality),
|
|
Callback(
|
|
callable=self.prepare_masked_range_segment_chunk,
|
|
args=[db_job.segment, chunk_number],
|
|
kwargs={
|
|
"quality": quality,
|
|
},
|
|
),
|
|
)
|
|
)
|
|
|
|
def get_or_set_segment_preview(self, db_segment: models.Segment) -> DataWithMime:
|
|
return self._to_data_with_mime(
|
|
self._get_or_set_cache_item(
|
|
self._make_preview_key(db_segment),
|
|
Callback(
|
|
callable=self._prepare_segment_preview,
|
|
args=[db_segment],
|
|
),
|
|
cache_item_ttl=self._PREVIEW_TTL,
|
|
)
|
|
)
|
|
|
|
def remove_task_chunk(
|
|
self, db_task: models.Task, chunk_number: int, *, quality: models.FrameQuality
|
|
) -> None:
|
|
self._delete_cache_item(
|
|
self._make_chunk_key(db_task, chunk_number, quality=quality),
|
|
)
|
|
|
|
def remove_segment_preview(self, db_segment: models.Segment) -> None:
|
|
self._delete_cache_item(self._make_preview_key(db_segment))
|
|
|
|
def remove_segment_chunk(
|
|
self, db_segment: models.Segment, chunk_number: str, *, quality: str
|
|
) -> None:
|
|
self._delete_cache_item(
|
|
self._make_chunk_key(db_segment, chunk_number=chunk_number, quality=quality)
|
|
)
|
|
|
|
def remove_context_images_chunk(self, db_data: models.Data, frame_number: str) -> None:
|
|
self._delete_cache_item(
|
|
self._make_frame_context_images_chunk_key(db_data, frame_number=frame_number)
|
|
)
|
|
|
|
def remove_segments_chunks(self, params: Sequence[dict[str, Any]]) -> None:
|
|
"""
|
|
Removes several segment chunks from the cache.
|
|
|
|
The function expects a sequence of remove_segment_chunk() parameters as dicts.
|
|
"""
|
|
# TODO: add a version of this function
|
|
# that removes related cache elements as well (context images, previews, ...)
|
|
# to provide encapsulation
|
|
|
|
# TODO: add a generic bulk cleanup function for different objects, including related ones
|
|
# (likely a bulk key aggregator should be used inside to reduce requests count)
|
|
|
|
keys_to_remove = []
|
|
for item_params in params:
|
|
db_obj = item_params.pop("db_segment")
|
|
keys_to_remove.append(self._make_chunk_key(db_obj, **item_params))
|
|
|
|
self._bulk_delete_cache_items(keys_to_remove)
|
|
|
|
def remove_context_images_chunks(self, params: Sequence[dict[str, Any]]) -> None:
|
|
"""
|
|
Removes several context image chunks from the cache.
|
|
|
|
The function expects a sequence of remove_context_images_chunk() parameters as dicts.
|
|
"""
|
|
|
|
keys_to_remove = []
|
|
for item_params in params:
|
|
db_obj = item_params.pop("db_data")
|
|
keys_to_remove.append(self._make_frame_context_images_chunk_key(db_obj, **item_params))
|
|
|
|
self._bulk_delete_cache_items(keys_to_remove)
|
|
|
|
def get_cloud_preview(self, db_storage: models.CloudStorage) -> DataWithMime | None:
|
|
return self._to_data_with_mime(
|
|
self._get_cache_item(self._make_preview_key(db_storage)), allow_none=True
|
|
)
|
|
|
|
def get_or_set_cloud_preview(self, db_storage: models.CloudStorage) -> DataWithMime:
|
|
return self._to_data_with_mime(
|
|
self._get_or_set_cache_item(
|
|
self._make_preview_key(db_storage),
|
|
Callback(
|
|
callable=self._prepare_cloud_preview,
|
|
args=[db_storage],
|
|
),
|
|
cache_item_ttl=self._PREVIEW_TTL,
|
|
)
|
|
)
|
|
|
|
def get_or_set_frame_context_images_chunk(
|
|
self, db_data: models.Data, frame_number: int
|
|
) -> DataWithMime:
|
|
return self._to_data_with_mime(
|
|
self._get_or_set_cache_item(
|
|
self._make_frame_context_images_chunk_key(db_data, frame_number),
|
|
Callback(
|
|
callable=self.prepare_context_images_chunk,
|
|
args=[db_data, frame_number],
|
|
),
|
|
)
|
|
)
|
|
|
|
@staticmethod
|
|
def read_raw_audio(db_task: models.Task) -> tuple[Openable, str]:
|
|
db_data = db_task.require_data()
|
|
assert db_data.storage in (
|
|
models.StorageChoice.LOCAL,
|
|
models.StorageChoice.SHARE,
|
|
), db_data.storage
|
|
|
|
filename = str(db_data.audio.path)
|
|
data_dir = db_data.get_raw_data_dirname()
|
|
return (data_dir / filename, filename)
|
|
|
|
@staticmethod
|
|
def read_raw_images(
|
|
db_task: models.Task, frame_ids: Sequence[int], *, decode: bool = True
|
|
) -> Generator[tuple[PIL.Image.Image | str, str], None, None]:
|
|
db_data = db_task.require_data()
|
|
manifest_path = db_data.get_manifest_path()
|
|
|
|
def requested_db_images() -> Iterator[str]:
|
|
# TODO: find a way to use prefetched results, if provided
|
|
db_images = (
|
|
db_data.images.order_by("frame")
|
|
.filter(frame__gte=frame_ids[0], frame__lte=frame_ids[-1])
|
|
.values_list("frame", "path")
|
|
)
|
|
|
|
requested_frame_iter = iter(frame_ids)
|
|
next_requested_frame_id = next(requested_frame_iter, None)
|
|
if next_requested_frame_id is None:
|
|
return
|
|
|
|
for frame_id, frame_path in db_images:
|
|
if frame_id == next_requested_frame_id:
|
|
yield frame_path
|
|
next_requested_frame_id = next(requested_frame_iter, None)
|
|
|
|
if next_requested_frame_id is None:
|
|
return
|
|
|
|
assert False, f"frame #{next_requested_frame_id} is missing from DB"
|
|
|
|
if storage_client := db_data.get_cloud_storage_instance():
|
|
with ExitStack() as es:
|
|
tmp_dir = Path(es.enter_context(tempfile.TemporaryDirectory(prefix="cvat")))
|
|
# (storage filename, output filename)
|
|
files_to_download: list[tuple[str, PurePath]] = []
|
|
checksums = []
|
|
media = []
|
|
if db_data.local_storage_backing_cs_id:
|
|
for frame_path in requested_db_images():
|
|
abs_frame_path = join_untrusted_path(tmp_dir, frame_path)
|
|
|
|
files_to_download.append((frame_path, abs_frame_path))
|
|
checksums.append(None)
|
|
media.append((abs_frame_path, os.fspath(abs_frame_path)))
|
|
|
|
else:
|
|
assert manifest_path.is_file()
|
|
reader = ImageReaderWithManifest(manifest_path)
|
|
for item in reader.iterate_frames(frame_ids):
|
|
frame_path = item.get("meta", {}).get(
|
|
"original_name", f"{item['name']}{item['extension']}"
|
|
)
|
|
abs_frame_path = join_untrusted_path(tmp_dir, frame_path)
|
|
|
|
files_to_download.append((frame_path, abs_frame_path))
|
|
checksums.append(item.get("checksum", None))
|
|
media.append((abs_frame_path, os.fspath(abs_frame_path)))
|
|
|
|
storage_client.bulk_download_to_dir(files=files_to_download, upload_dir=tmp_dir)
|
|
|
|
for checksum, media_item in zip(checksums, media):
|
|
frame_path = media_item[1]
|
|
if checksum and not md5_hash(frame_path) == checksum:
|
|
slogger.task[db_task.id].warning(
|
|
"Hash sums of files {} do not match".format(frame_path)
|
|
)
|
|
|
|
if db_task.dimension == models.DimensionType.DIM_3D and (
|
|
frame_path.endswith(".bin")
|
|
):
|
|
frame_path = ValidateDimension().convert_bin_to_pcd(
|
|
frame_path,
|
|
# one file can be used several times for honeypots
|
|
delete_source=False,
|
|
)
|
|
media_item = (frame_path, frame_path)
|
|
|
|
if db_task.dimension == models.DimensionType.DIM_2D and decode:
|
|
media_item = load_image(media_item)
|
|
|
|
yield media_item
|
|
|
|
else:
|
|
raw_data_dir = db_data.get_raw_data_dirname()
|
|
media = []
|
|
for frame_path in requested_db_images():
|
|
source_path = join_untrusted_path(raw_data_dir, frame_path)
|
|
media.append((source_path, source_path))
|
|
|
|
if db_task.dimension == models.DimensionType.DIM_2D and decode:
|
|
media = map(load_image, media)
|
|
|
|
yield from media
|
|
|
|
@classmethod
|
|
def read_raw_context_images(
|
|
cls,
|
|
db_data: models.Data,
|
|
frame_ids: Sequence[int],
|
|
*,
|
|
truncate_common_filename_prefix: bool = True, # should be done on the UI, probably
|
|
decode: bool = True,
|
|
) -> Generator[tuple[int, tuple[PIL.Image.Image | str, str]], None, None]:
|
|
raw_data_dir = db_data.get_raw_data_dirname()
|
|
|
|
with ExitStack() as es:
|
|
ThroughModel = models.RelatedFile.images.through
|
|
|
|
db_related_files = (
|
|
ThroughModel.objects.filter(relatedfile__data=db_data, image__frame__in=frame_ids)
|
|
.order_by("image__frame", "relatedfile__path")
|
|
.values_list("image__frame", "relatedfile__path")
|
|
)
|
|
|
|
media = [
|
|
(frame_id, [ri_path for _, ri_path in frame_ris])
|
|
for frame_id, frame_ris in groupby(db_related_files, key=lambda v: v[0])
|
|
]
|
|
|
|
if storage_client := db_data.get_cloud_storage_instance():
|
|
tmp_dir = Path(es.enter_context(tempfile.TemporaryDirectory(prefix="cvat")))
|
|
files_to_download: list[tuple[str, PurePath]] = []
|
|
for _, frame_media in media:
|
|
for ri_path in frame_media:
|
|
abs_ri_path = join_untrusted_path(tmp_dir, ri_path)
|
|
files_to_download.append((ri_path, abs_ri_path))
|
|
|
|
storage_client.bulk_download_to_dir(files_to_download, upload_dir=tmp_dir)
|
|
media_base_dir = tmp_dir
|
|
else:
|
|
media_base_dir = raw_data_dir
|
|
|
|
for frame_id, frame_media in media:
|
|
if truncate_common_filename_prefix:
|
|
# Truncate RI prefixes on the per-frame basis
|
|
common_prefix = os.path.commonpath(os.path.dirname(m) for m in frame_media)
|
|
|
|
frame_media_tuples = [
|
|
(
|
|
os.fspath(join_untrusted_path(media_base_dir, m)),
|
|
os.path.relpath(m, common_prefix),
|
|
)
|
|
for m in frame_media
|
|
]
|
|
else:
|
|
frame_media_tuples = [
|
|
(os.fspath(join_untrusted_path(media_base_dir, m)), os.fspath(m))
|
|
for m in frame_media
|
|
]
|
|
|
|
for m in frame_media_tuples:
|
|
if decode:
|
|
m = load_image(m)
|
|
|
|
yield frame_id, m
|
|
|
|
@staticmethod
|
|
def _read_raw_frames(
|
|
db_task: models.Task | int, frame_ids: Sequence[int]
|
|
) -> Generator[tuple[av.VideoFrame | PIL.Image.Image | str, str | None], None, None]:
|
|
if isinstance(db_task, int):
|
|
db_task = models.Task.objects.get(pk=db_task)
|
|
|
|
for prev_frame, cur_frame in pairwise(frame_ids):
|
|
assert (
|
|
prev_frame <= cur_frame
|
|
), f"Requested frame ids must be sorted, got a ({prev_frame}, {cur_frame}) pair"
|
|
|
|
db_data = db_task.require_data()
|
|
|
|
if hasattr(db_data, "video"):
|
|
source_path = db_data.get_raw_data_dirname() / db_data.video.path
|
|
|
|
manifest_path = db_data.get_manifest_path()
|
|
reader = VideoReaderWithManifest(
|
|
manifest_path=manifest_path,
|
|
source_path=source_path,
|
|
allow_threading=False,
|
|
)
|
|
if not os.path.isfile(manifest_path):
|
|
try:
|
|
reader.manifest.link(source_path, force=True)
|
|
reader.manifest.create()
|
|
except Exception as e:
|
|
slogger.task[db_task.id].warning(
|
|
f"Failed to create video manifest: {e}", exc_info=True
|
|
)
|
|
reader = None
|
|
|
|
if reader:
|
|
for frame in reader.iterate_frames(frame_filter=frame_ids):
|
|
yield (frame, None)
|
|
else:
|
|
reader = VideoReader([source_path], allow_threading=False)
|
|
|
|
yield from reader.iterate_frames(frame_filter=frame_ids)
|
|
else:
|
|
yield from MediaCache.read_raw_images(db_task, frame_ids)
|
|
|
|
def prepare_segment_chunk(
|
|
self,
|
|
db_segment: models.Segment | int,
|
|
chunk_number: int,
|
|
*,
|
|
quality: models.FrameQuality,
|
|
) -> DataWithMime:
|
|
if isinstance(db_segment, int):
|
|
db_segment = models.Segment.objects.get(pk=db_segment)
|
|
|
|
if db_segment.type == models.SegmentType.RANGE:
|
|
return self.prepare_range_segment_chunk(db_segment, chunk_number, quality=quality)
|
|
elif db_segment.type == models.SegmentType.SPECIFIC_FRAMES:
|
|
return self.prepare_masked_range_segment_chunk(
|
|
db_segment, chunk_number, quality=quality
|
|
)
|
|
else:
|
|
assert False, f"Unknown segment type {db_segment.type}"
|
|
|
|
def prepare_range_segment_chunk(
|
|
self, db_segment: models.Segment, chunk_number: int, *, quality: models.FrameQuality
|
|
) -> DataWithMime:
|
|
db_task = db_segment.task
|
|
db_data = db_task.require_data()
|
|
|
|
chunk_size = db_data.chunk_size
|
|
chunk_start_frame_index = chunk_size * chunk_number
|
|
chunk_end_frame_index = chunk_size * (chunk_number + 1)
|
|
chunk_frame_range = db_segment.frame_set[chunk_start_frame_index:chunk_end_frame_index]
|
|
return self.prepare_custom_range_segment_chunk(
|
|
db_task, chunk_frame_range, quality=quality, cache=self
|
|
)
|
|
|
|
@classmethod
|
|
def prepare_custom_range_segment_chunk(
|
|
cls,
|
|
db_task: models.Task,
|
|
frame_ids: Sequence[int],
|
|
*,
|
|
quality: models.FrameQuality,
|
|
cache: MediaCache | None = None,
|
|
) -> DataWithMime:
|
|
# TODO: refactor all chunk building into another class
|
|
|
|
match db_task.media_type:
|
|
case models.MediaType.AUDIO:
|
|
from cvat.apps.engine.media_io.audio_provider import TaskAudioProvider
|
|
|
|
assert cache
|
|
|
|
return TaskAudioProvider._build_audio_chunk(
|
|
db_task=db_task,
|
|
chunk_frames=(frame_ids[0], frame_ids[-1]),
|
|
quality=quality,
|
|
cache=cache,
|
|
)
|
|
|
|
case models.MediaType.IMAGE | models.MediaType.POINT_CLOUD:
|
|
from cvat.apps.engine.media_io.frame_provider import prepare_image_chunk
|
|
|
|
with closing(cls._read_raw_frames(db_task, frame_ids=frame_ids)) as frame_iter:
|
|
return prepare_image_chunk(frame_iter, quality=quality, db_task=db_task)
|
|
case _ as media_type:
|
|
assert False, f"Unknown media type '{media_type}'"
|
|
|
|
def prepare_masked_range_segment_chunk(
|
|
self, db_segment: models.Segment, chunk_number: int, *, quality: models.FrameQuality
|
|
) -> DataWithMime:
|
|
db_task = db_segment.task
|
|
db_data = db_task.require_data()
|
|
|
|
chunk_size = db_data.chunk_size
|
|
chunk_frame_ids = sorted(db_segment.frame_set)[
|
|
chunk_size * chunk_number : chunk_size * (chunk_number + 1)
|
|
]
|
|
|
|
assert db_task.media_type != models.MediaType.AUDIO
|
|
return self.prepare_custom_masked_range_segment_chunk(
|
|
db_task, chunk_frame_ids, chunk_number, quality=quality
|
|
)
|
|
|
|
@classmethod
|
|
def prepare_custom_masked_range_segment_chunk(
|
|
cls,
|
|
db_task: models.Task | int,
|
|
frame_ids: Collection[int],
|
|
chunk_number: int,
|
|
*,
|
|
quality: models.FrameQuality,
|
|
insert_placeholders: bool = False,
|
|
) -> DataWithMime:
|
|
if isinstance(db_task, int):
|
|
db_task = models.Task.objects.get(pk=db_task)
|
|
|
|
db_data = db_task.require_data()
|
|
|
|
frame_step = db_data.get_frame_step()
|
|
|
|
image_quality = 100 if quality == models.FrameQuality.ORIGINAL else db_data.image_quality
|
|
writer = ZipCompressedChunkWriter(quality=image_quality, dimension=db_task.dimension)
|
|
|
|
dummy_frame = io.BytesIO()
|
|
PIL.Image.new("RGB", (1, 1)).save(dummy_frame, writer.IMAGE_EXT)
|
|
|
|
# Optimize frame access if all the required frames are already cached
|
|
# Otherwise we might need to download files.
|
|
# This is not needed for video tasks, as it will reduce performance,
|
|
# because of reading multiple files (chunks)
|
|
from cvat.apps.engine.media_io.frame_provider import FrameOutputType, make_frame_provider
|
|
|
|
task_frame_provider = make_frame_provider(db_task)
|
|
|
|
use_cached_data = False
|
|
if db_task.mode != models.TaskMode.INTERPOLATION:
|
|
required_frame_set = set(frame_ids)
|
|
available_chunks = []
|
|
for db_segment in db_task.segment_set.filter(type=models.SegmentType.RANGE).all():
|
|
segment_frame_provider = make_frame_provider(db_segment)
|
|
|
|
for i, chunk_frames in groupby(
|
|
sorted(required_frame_set.intersection(db_segment.frame_set)),
|
|
key=lambda abs_frame: (
|
|
segment_frame_provider.validate_frame_number(
|
|
task_frame_provider.get_rel_frame_number(abs_frame)
|
|
)[1]
|
|
),
|
|
):
|
|
if not list(chunk_frames):
|
|
continue
|
|
|
|
chunk_available = cls._has_key(
|
|
cls._make_chunk_key(db_segment, i, quality=quality)
|
|
)
|
|
available_chunks.append(chunk_available)
|
|
|
|
use_cached_data = bool(available_chunks) and all(available_chunks)
|
|
|
|
if hasattr(db_data, "video"):
|
|
frame_size = (db_data.video.width, db_data.video.height)
|
|
else:
|
|
frame_size = None
|
|
|
|
def get_frames():
|
|
with ExitStack() as es:
|
|
es.callback(task_frame_provider.unload)
|
|
|
|
if insert_placeholders:
|
|
frame_range = (
|
|
(
|
|
db_data.start_frame
|
|
+ (chunk_number * db_data.chunk_size + chunk_frame_idx) * frame_step
|
|
)
|
|
for chunk_frame_idx in range(db_data.chunk_size)
|
|
)
|
|
else:
|
|
frame_range = frame_ids
|
|
|
|
if not use_cached_data:
|
|
frames_gen = cls._read_raw_frames(db_task, frame_ids)
|
|
frames_iter = iter(es.enter_context(closing(frames_gen)))
|
|
|
|
for abs_frame_idx in frame_range:
|
|
if db_data.stop_frame < abs_frame_idx:
|
|
break
|
|
|
|
if abs_frame_idx in frame_ids:
|
|
if use_cached_data:
|
|
frame_data = task_frame_provider.get_frame(
|
|
task_frame_provider.get_rel_frame_number(abs_frame_idx),
|
|
quality=quality,
|
|
out_type=FrameOutputType.BUFFER,
|
|
)
|
|
frame = frame_data.data
|
|
else:
|
|
frame, _ = next(frames_iter)
|
|
|
|
if hasattr(db_data, "video"):
|
|
# Decoded video frames can have different size, restore the original one
|
|
|
|
if isinstance(frame, av.VideoFrame):
|
|
frame = frame.to_image()
|
|
else:
|
|
frame = PIL.Image.open(frame)
|
|
|
|
if frame.size != frame_size:
|
|
frame = frame.resize(frame_size)
|
|
else:
|
|
# Populate skipped frames with placeholder data,
|
|
# this is required for video chunk decoding implementation in UI
|
|
frame = io.BytesIO(dummy_frame.getvalue())
|
|
|
|
yield (frame, None)
|
|
|
|
buff = io.BytesIO()
|
|
with closing(get_frames()) as frame_iter:
|
|
writer.save_as_chunk(
|
|
frame_iter,
|
|
buff,
|
|
zip_compress_level=1,
|
|
# there are likely to be many skips with repeated placeholder frames
|
|
# in SPECIFIC_FRAMES segments, it makes sense to compress the archive
|
|
)
|
|
|
|
buff.seek(0)
|
|
return buff, writer.CHUNK_MIME_TYPE
|
|
|
|
def _prepare_segment_preview(self, db_segment: models.Segment | int) -> DataWithMime:
|
|
if isinstance(db_segment, int):
|
|
db_segment = models.Segment.objects.get(pk=db_segment)
|
|
|
|
match db_segment.task.media_type:
|
|
case models.MediaType.POINT_CLOUD:
|
|
preview = PIL.Image.open(ASSETS_DIR / "point_cloud_default_preview.png")
|
|
case models.MediaType.AUDIO:
|
|
from cvat.apps.engine.media_extractors import AudioReader
|
|
|
|
preview = None
|
|
if db_segment.task.data.audio.has_cover_image:
|
|
source_audio = self.read_raw_audio(db_segment.task)[0]
|
|
reader = AudioReader([source_audio])
|
|
preview = reader.get_preview_image()
|
|
else:
|
|
preview = PIL.Image.open(ASSETS_DIR / "audio_default_preview.png")
|
|
case models.MediaType.IMAGE:
|
|
from cvat.apps.engine.media_io.frame_provider import ( # avoid circular import
|
|
FrameOutputType,
|
|
make_frame_provider,
|
|
)
|
|
|
|
task_frame_provider = make_frame_provider(db_segment.task)
|
|
segment_frame_provider = make_frame_provider(db_segment)
|
|
|
|
preview = segment_frame_provider.get_frame(
|
|
task_frame_provider.get_rel_frame_number(min(db_segment.frame_set)),
|
|
quality=models.FrameQuality.COMPRESSED,
|
|
out_type=FrameOutputType.PIL,
|
|
).data
|
|
case _ as media_type:
|
|
assert False, f"Unknown media type {media_type}"
|
|
|
|
return prepare_preview_image(preview)
|
|
|
|
def _prepare_cloud_preview(self, db_storage: models.CloudStorage | int) -> DataWithMime:
|
|
if isinstance(db_storage, int):
|
|
db_storage = models.CloudStorage.objects.get(pk=db_storage)
|
|
|
|
storage = db_storage_to_storage_instance(db_storage)
|
|
if not db_storage.manifests.count():
|
|
raise ValidationError("Cannot get the cloud storage preview. There is no manifest file")
|
|
|
|
preview_path = None
|
|
for db_manifest in db_storage.manifests.all():
|
|
manifest_prefix = os.path.dirname(db_manifest.filename)
|
|
|
|
full_manifest_path = join_untrusted_path(
|
|
db_storage.get_storage_dirname(), db_manifest.filename
|
|
)
|
|
|
|
if not full_manifest_path.exists() or datetime.fromtimestamp(
|
|
full_manifest_path.stat().st_mtime, tz=timezone.utc
|
|
) < storage.get_file_last_modified(db_manifest.filename):
|
|
storage.download_file(db_manifest.filename, full_manifest_path)
|
|
|
|
manifest = ImageManifestManager(full_manifest_path, db_storage.get_storage_dirname())
|
|
# need to update index
|
|
manifest.set_index()
|
|
if not len(manifest):
|
|
continue
|
|
|
|
preview_info = manifest[0]
|
|
preview_filename = "".join([preview_info["name"], preview_info["extension"]])
|
|
preview_path = os.path.join(manifest_prefix, preview_filename)
|
|
break
|
|
|
|
if not preview_path:
|
|
msg = "Cloud storage {} does not contain any images".format(db_storage.pk)
|
|
slogger.cloud_storage[db_storage.pk].info(msg)
|
|
raise NotFound(msg)
|
|
|
|
preview_bytes = storage.download_fileobj(preview_path)
|
|
image = PIL.Image.open(io.BytesIO(preview_bytes))
|
|
return prepare_preview_image(image)
|
|
|
|
def prepare_context_images_chunk(
|
|
self, db_data: models.Data | int, frame_number: int
|
|
) -> DataWithMime:
|
|
if isinstance(db_data, int):
|
|
db_data = models.Data.objects.get(pk=db_data)
|
|
|
|
zip_buffer = io.BytesIO()
|
|
mime_type = ""
|
|
|
|
with (
|
|
closing(self.read_raw_context_images(db_data, frame_ids=[frame_number])) as ri_iter,
|
|
zipfile.ZipFile(zip_buffer, "a", zipfile.ZIP_DEFLATED, False) as zip_file,
|
|
):
|
|
for _, (image, path) in ri_iter:
|
|
name = os.path.splitext(path)[0]
|
|
|
|
try:
|
|
if image.mode != "RGB" and image.mode != "L":
|
|
image = image.convert("RGB")
|
|
|
|
image_file = io.BytesIO()
|
|
image.save(image_file, format="JPEG", quality=100, optimize=True)
|
|
image_file.seek(0)
|
|
except OSError as e:
|
|
raise Exception('Failed to encode image to ".jpeg" format') from e
|
|
|
|
zip_file.writestr(f"{name}.jpg", image_file.getbuffer())
|
|
|
|
if not mime_type:
|
|
mime_type = "application/zip"
|
|
|
|
zip_buffer.seek(0)
|
|
return zip_buffer, mime_type
|
|
|
|
|
|
def prepare_preview_image(image: PIL.Image.Image) -> DataWithMime:
|
|
PREVIEW_SIZE = (256, 256)
|
|
|
|
ALLOWED_FORMATS = {"PNG", "JPEG"}
|
|
|
|
def get_mime(format_name: str) -> str:
|
|
return PIL.Image.MIME[format_name]
|
|
|
|
# format is erased by exif_transpose(), keep it if possible
|
|
image_format = image.format
|
|
image = PIL.ImageOps.exif_transpose(image)
|
|
|
|
if image.size != PREVIEW_SIZE:
|
|
image.thumbnail(PREVIEW_SIZE)
|
|
|
|
output_buf = io.BytesIO()
|
|
if image_format in ALLOWED_FORMATS:
|
|
image.save(output_buf, format=image_format)
|
|
mime = get_mime(image_format)
|
|
else:
|
|
image.convert("RGB").save(output_buf, format="JPEG")
|
|
mime = get_mime("JPEG")
|
|
|
|
output_buf.seek(0)
|
|
return output_buf, mime
|