项目文件夹

文件
2026-07-13 13:32:23 +08:00

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