Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
70 changes: 46 additions & 24 deletions src/aris/video/processing.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,8 @@
import logging
import queue
import threading

from typing import Any, Callable, Iterable
from collections.abc import Callable, Iterable
from typing import Any

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -99,7 +99,9 @@ def cancel(self):
threads to exit when they next check the stop flag. Queued items may be
left unprocessed.
"""
logger.info(f"Pipeline cancel requested{f' ({self.task_name})' if self.task_name else ''}")
logger.info(
f"Pipeline cancel requested{f' ({self.task_name})' if self.task_name else ''}"
)
self._cancelled = True
self._stop_event.set()

Expand Down Expand Up @@ -132,7 +134,9 @@ def _notify_task_done(self):
if self._on_task_done:
with self._lock:
self._completed += 1
cancel_requested = self._on_task_done(self._completed, self._total_items)
cancel_requested = self._on_task_done(
self._completed, self._total_items
)
if cancel_requested:
self.cancel()

Expand All @@ -142,11 +146,15 @@ def _producer_thread(self):
if producer is None or not callable(producer):
raise RuntimeError("Producer thread requires a callable producer")
task_prefix = f"{self.task_name}->" if self.task_name else ""
logger.info(f"--- Starting producer thread (target: {task_prefix}{producer.__name__}) ---")
logger.info(
f"--- Starting producer thread (target: {task_prefix}{producer.__name__}) ---"
)
try:
for item in producer():
if self._stop_event.is_set():
logger.info(f"--- Producer thread ({task_prefix}{producer.__name__}) cancelled. ---")
logger.info(
f"--- Producer thread ({task_prefix}{producer.__name__}) cancelled. ---"
)
break
self._put(self.task_queue, item)
else:
Expand All @@ -159,12 +167,16 @@ def _producer_thread(self):
exc_info=True,
)
self._fail(e)
logger.info(f"--- Producer thread ({task_prefix}{producer.__name__}) finished. ---")
logger.info(
f"--- Producer thread ({task_prefix}{producer.__name__}) finished. ---"
)

def _consumer_thread(self, thread_num: int):
"""Internal method for consumer workers."""
task_prefix = f"{self.task_name}->" if self.task_name else ""
logger.info(f"--- Consumer thread {thread_num}/{self.num_workers} started (target: {task_prefix}{self.consumer.__name__}) ---")
logger.info(
f"--- Consumer thread {thread_num}/{self.num_workers} started (target: {task_prefix}{self.consumer.__name__}) ---"
)
while True:
if self._stop_event.is_set():
break
Expand All @@ -183,7 +195,10 @@ def _consumer_thread(self, thread_num: int):
if self.writer and result is not None and writer_queue is not None:
self._put(writer_queue, result)
except Exception as e:
logger.error(f"Exception in consumer thread {thread_num} processing task ({type(e).__name__}): {e}", exc_info=True)
logger.error(
f"Exception in consumer thread {thread_num} processing task ({type(e).__name__}): {e}",
exc_info=True,
)
self._fail(e)
finally:
# Always mark task as done to prevent queue.join() from hanging
Expand All @@ -198,7 +213,9 @@ def _writer_thread(self):
if writer is None or writer_queue is None:
raise RuntimeError("Writer thread requires a writer and writer queue")
task_prefix = f"{self.task_name}->" if self.task_name else ""
logger.info(f"--- Starting writer thread (target: {task_prefix}{writer.__name__}) ---")
logger.info(
f"--- Starting writer thread (target: {task_prefix}{writer.__name__}) ---"
)
while not self._stop_event.is_set():
try:
item_to_write = writer_queue.get(timeout=0.1)
Expand All @@ -209,7 +226,10 @@ def _writer_thread(self):
try:
writer(item_to_write)
except Exception as e:
logger.error(f"Exception in writer thread ({type(e).__name__}): {e}", exc_info=True)
logger.error(
f"Exception in writer thread ({type(e).__name__}): {e}",
exc_info=True,
)
writer_queue.task_done()
self._fail(e)
break
Expand All @@ -235,17 +255,16 @@ def _monitor_queues(self, interval=0.1):

def _run_sync(self):
"""Run the pipeline synchronously on the calling thread.

Used when num_workers == 0. Executes producer→consumer→writer
sequentially without any threading overhead.
"""
import time

start_time = time.time()

task_desc = f" ({self.task_name})" if self.task_name else ""
logger.info(
f"Running pipeline{task_desc} synchronously..."
)
logger.info(f"Running pipeline{task_desc} synchronously...")

for item in self.producer() if self.producer else ():
if self._stop_event.is_set():
Expand All @@ -256,13 +275,13 @@ def _run_sync(self):
break
if self.writer and result is not None:
self.writer(result)

self._processing_time = time.time() - start_time
logger.info(f"Pipeline{task_desc} processing complete!")

def run(self):
"""Starts and runs the entire processing pipeline.

If num_workers is 0, runs synchronously on the calling thread.
"""
self._stop_event.clear()
Expand All @@ -273,7 +292,7 @@ def run(self):
# Call on_task_done immediately with 0 progress for instant feedback
if self._on_task_done and self._total_items > 0:
self._on_task_done(0, self._total_items)

if self.num_workers == 0:
return self._run_sync()

Expand All @@ -297,10 +316,13 @@ def run(self):
self._writer_queue_empty_time = 0

import time

start_time = time.time()

task_desc = f" ({self.task_name})" if self.task_name else ""
logger.info(f"Starting pipeline{task_desc} with {self.num_workers} worker threads...")
logger.info(
f"Starting pipeline{task_desc} with {self.num_workers} worker threads..."
)
worker_threads = []
writer_thread = None
producer_thread = None
Expand All @@ -322,7 +344,7 @@ def run(self):
thread = threading.Thread(
target=self._consumer_thread,
args=(i + 1,),
name=f"ConsumerThread-{i+1}",
name=f"ConsumerThread-{i + 1}",
)
thread.start()
worker_threads.append(thread)
Expand All @@ -345,7 +367,7 @@ def run(self):
# Stop monitor (does not mark the run as cancelled)
self._stop_event.set()
monitor_thread.join()

self._processing_time = time.time() - start_time

task_desc = f" ({self.task_name})" if self.task_name else ""
Expand Down Expand Up @@ -377,7 +399,7 @@ def get_queue_stats(self) -> dict[str, Any]:
"avg_depth": avg_depth,
"empty_time": self._writer_queue_empty_time,
}

stats["processing_time"] = self._processing_time

return stats
4 changes: 3 additions & 1 deletion tests/integration/test_video_export_fps.py
Original file line number Diff line number Diff line change
Expand Up @@ -90,7 +90,9 @@ def pytest_generate_tests(metafunc: pytest.Metafunc) -> None:


@pytest.mark.slow
def test_video_export_source_fps_writes_two_frames(aris_path: Path | None, tmp_path: Path):
def test_video_export_source_fps_writes_two_frames(
aris_path: Path | None, tmp_path: Path
):
"""Export two frames at each file's native fps; must complete and write."""
if aris_path is None:
pytest.skip(_aris_skip_reason())
Expand Down
Loading