Coverage for documents/management/commands/base.py: 0%
175 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 09:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 09:07 +0000
1"""
2Base command class for Paperless-ngx management commands.
4Provides automatic progress bar and multiprocessing support with minimal boilerplate.
5"""
7from __future__ import annotations
9import logging
10import os
11from collections.abc import Callable
12from collections.abc import Iterable
13from collections.abc import Sized
14from concurrent.futures import ProcessPoolExecutor
15from concurrent.futures import as_completed
16from contextlib import contextmanager
17from dataclasses import dataclass
18from typing import TYPE_CHECKING
19from typing import Any
20from typing import ClassVar
21from typing import Generic
22from typing import TypeVar
24import django
25from django import db
26from django.core.management import CommandError
27from django.db.models import QuerySet
28from django_rich.management import RichCommand
29from rich import box
30from rich.console import Console
31from rich.console import Group
32from rich.console import RenderableType
33from rich.live import Live
34from rich.progress import BarColumn
35from rich.progress import MofNCompleteColumn
36from rich.progress import Progress
37from rich.progress import SpinnerColumn
38from rich.progress import TextColumn
39from rich.progress import TimeElapsedColumn
40from rich.progress import TimeRemainingColumn
41from rich.table import Table
42from rich.text import Text
44if TYPE_CHECKING:
45 from collections.abc import Generator
46 from collections.abc import Sequence
48 from django.core.management import CommandParser
50T = TypeVar("T")
51R = TypeVar("R")
54@dataclass(slots=True, frozen=True)
55class _BufferedRecord:
56 level: int
57 name: str
58 message: str
61class BufferingLogHandler(logging.Handler):
62 """Captures log records during a command run for deferred rendering.
64 Attach to a logger before a long operation and call ``render()``
65 afterwards to emit the buffered records via Rich, optionally filtered
66 by minimum level.
67 """
69 def __init__(self) -> None:
70 super().__init__()
71 self._records: list[_BufferedRecord] = []
73 def emit(self, record: logging.LogRecord) -> None:
74 self._records.append(
75 _BufferedRecord(
76 level=record.levelno,
77 name=record.name,
78 message=self.format(record),
79 ),
80 )
82 def render(
83 self,
84 console: Console,
85 *,
86 min_level: int = logging.DEBUG,
87 title: str = "Log Output",
88 ) -> None:
89 records = [r for r in self._records if r.level >= min_level]
90 if not records:
91 return
93 table = Table(
94 title=title,
95 show_header=True,
96 header_style="bold",
97 show_lines=False,
98 box=box.SIMPLE,
99 )
100 table.add_column("Level", style="bold", width=8)
101 table.add_column("Logger", style="dim")
102 table.add_column("Message", no_wrap=False)
104 _level_styles: dict[int, str] = {
105 logging.DEBUG: "dim",
106 logging.INFO: "cyan",
107 logging.WARNING: "yellow",
108 logging.ERROR: "red",
109 logging.CRITICAL: "bold red",
110 }
112 for record in records:
113 style = _level_styles.get(record.level, "")
114 table.add_row(
115 Text(logging.getLevelName(record.level), style=style),
116 record.name,
117 record.message,
118 )
120 console.print(table)
122 def clear(self) -> None:
123 self._records.clear()
126@dataclass(frozen=True, slots=True)
127class ProcessResult(Generic[T, R]):
128 """
129 Result of processing a single item in parallel.
131 Attributes:
132 item: The input item that was processed.
133 result: The return value from the processing function, or None if an error occurred.
134 error: The exception if processing failed, or None on success.
135 """
137 item: T
138 result: R | None
139 error: BaseException | None
141 @property
142 def success(self) -> bool:
143 """Return True if the item was processed successfully."""
144 return self.error is None
147class PaperlessCommand(RichCommand):
148 """
149 Base command class with automatic progress bar and multiprocessing support.
151 Features are opt-in via class attributes:
152 supports_progress_bar: Adds --no-progress-bar argument (default: True)
153 supports_multiprocessing: Adds --processes argument (default: False)
155 Example usage:
157 class Command(PaperlessCommand):
158 help = "Process all documents"
160 def handle(self, *args, **options):
161 documents = Document.objects.all()
162 for doc in self.track(documents, description="Processing..."):
163 process_document(doc)
165 class Command(PaperlessCommand):
166 help = "Regenerate thumbnails"
167 supports_multiprocessing = True
169 def handle(self, *args, **options):
170 ids = list(Document.objects.values_list("id", flat=True))
171 for result in self.process_parallel(process_doc, ids):
172 if result.error:
173 self.console.print(f"[red]Failed: {result.error}[/red]")
175 class Command(PaperlessCommand):
176 help = "Import documents with live stats"
178 def handle(self, *args, **options):
179 stats = ImportStats()
181 def render_stats() -> Table:
182 ... # build Rich Table from stats
184 for item in self.track_with_stats(
185 items,
186 description="Importing...",
187 stats_renderer=render_stats,
188 ):
189 result = import_item(item)
190 stats.imported += 1
191 """
193 supports_progress_bar: ClassVar[bool] = True
194 supports_multiprocessing: ClassVar[bool] = False
196 # Instance attributes set by execute() before handle() runs
197 no_progress_bar: bool
198 process_count: int
200 def add_arguments(self, parser: CommandParser) -> None:
201 """Add arguments based on supported features."""
202 super().add_arguments(parser)
204 if self.supports_progress_bar:
205 parser.add_argument(
206 "--no-progress-bar",
207 default=False,
208 action="store_true",
209 help="Disable the progress bar",
210 )
212 if self.supports_multiprocessing:
213 default_processes = max(1, (os.cpu_count() or 1) // 4)
214 parser.add_argument(
215 "--processes",
216 default=default_processes,
217 type=int,
218 help=f"Number of processes to use (default: {default_processes})",
219 )
221 def execute(self, *args: Any, **options: Any) -> str | None:
222 """
223 Set up instance state before handle() is called.
225 This is called by Django's command infrastructure after argument parsing
226 but before handle(). We use it to set instance attributes from options.
227 """
228 if self.supports_progress_bar:
229 self.no_progress_bar = options.get("no_progress_bar", False)
230 else:
231 self.no_progress_bar = True
233 if self.supports_multiprocessing:
234 self.process_count = options.get("processes", 1)
235 if self.process_count < 1:
236 raise CommandError("--processes must be at least 1")
237 else:
238 self.process_count = 1
240 return super().execute(*args, **options)
242 @contextmanager
243 def buffered_logging(
244 self,
245 *logger_names: str,
246 level: int = logging.DEBUG,
247 ) -> Generator[BufferingLogHandler, None, None]:
248 """Context manager that captures log output from named loggers.
250 Installs a ``BufferingLogHandler`` on each named logger for the
251 duration of the block, suppressing propagation to avoid interleaving
252 with the Rich live display. The handler is removed on exit regardless
253 of whether an exception occurred.
255 Usage::
257 with self.buffered_logging("paperless", "documents") as log_buf:
258 # ... run progress loop ...
259 if options["verbose"]:
260 log_buf.render(self.console)
261 """
262 handler = BufferingLogHandler()
263 handler.setFormatter(logging.Formatter("%(message)s"))
265 loggers: list[logging.Logger] = []
266 original_propagate: dict[str, bool] = {}
268 for name in logger_names:
269 log = logging.getLogger(name)
270 log.addHandler(handler)
271 original_propagate[name] = log.propagate
272 log.propagate = False
273 loggers.append(log)
275 try:
276 yield handler
277 finally:
278 for log in loggers:
279 log.removeHandler(handler)
280 log.propagate = original_propagate[log.name]
282 @staticmethod
283 def _progress_columns() -> tuple[Any, ...]:
284 """
285 Return the standard set of progress bar columns.
287 Extracted so both _create_progress (standalone) and track_with_stats
288 (inside Live) use identical column configuration without duplication.
289 """
290 return (
291 SpinnerColumn(),
292 TextColumn("[progress.description]{task.description}"),
293 BarColumn(),
294 MofNCompleteColumn(),
295 TimeElapsedColumn(),
296 TimeRemainingColumn(),
297 )
299 def _create_progress(self, description: str) -> Progress:
300 """
301 Create a standalone Progress instance with its own stderr Console.
303 Use this for track(). For track_with_stats(), Progress is created
304 directly inside a Live context instead.
306 Progress output is directed to stderr to match the convention that
307 progress bars are transient UI feedback, not command output. This
308 mirrors the convention that progress bars are transient UI feedback and prevents progress bar rendering
309 from interfering with stdout-based assertions in tests or piped
310 command output.
312 Args:
313 description: Text to display alongside the progress bar.
315 Returns:
316 A Progress instance configured with appropriate columns.
317 """
318 return Progress(
319 *self._progress_columns(),
320 console=Console(stderr=True),
321 transient=False,
322 )
324 def _get_iterable_length(self, iterable: Iterable[object]) -> int | None:
325 """
326 Attempt to determine the length of an iterable without consuming it.
328 Tries .count() first (for Django querysets - executes SELECT COUNT(*)),
329 then falls back to len() for sequences.
331 Args:
332 iterable: The iterable to measure.
334 Returns:
335 The length if determinable, None otherwise.
336 """
337 if isinstance(iterable, QuerySet):
338 return iterable.count()
340 if isinstance(iterable, Sized):
341 return len(iterable)
343 return None
345 def track(
346 self,
347 iterable: Iterable[T],
348 *,
349 description: str = "Processing...",
350 total: int | None = None,
351 ) -> Generator[T, None, None]:
352 """
353 Iterate over items with an optional progress bar.
355 Respects --no-progress-bar flag. When disabled, simply yields items
356 without any progress display.
358 Args:
359 iterable: The items to iterate over.
360 description: Text to display alongside the progress bar.
361 total: Total number of items. If None, attempts to determine
362 automatically via .count() (for querysets) or len().
364 Yields:
365 Items from the iterable.
367 Example:
368 for doc in self.track(documents, description="Renaming..."):
369 process(doc)
370 """
371 if self.no_progress_bar:
372 yield from iterable
373 return
375 if total is None:
376 total = self._get_iterable_length(iterable)
378 with self._create_progress(description) as progress:
379 task_id = progress.add_task(description, total=total)
380 for item in iterable:
381 yield item
382 progress.advance(task_id)
384 def track_with_stats(
385 self,
386 iterable: Iterable[T],
387 *,
388 description: str = "Processing...",
389 stats_renderer: Callable[[], RenderableType],
390 total: int | None = None,
391 ) -> Generator[T, None, None]:
392 """
393 Iterate over items with a progress bar and a live-updating stats display.
395 The progress bar and stats renderable are combined in a single Live
396 context, so the stats panel re-renders in place below the progress bar
397 after each item is processed.
399 Respects --no-progress-bar flag. When disabled, yields items without
400 any display (stats are still updated by the caller's loop body, so
401 they will be accurate for any post-loop summary the caller prints).
403 Args:
404 iterable: The items to iterate over.
405 description: Text to display alongside the progress bar.
406 stats_renderer: Zero-argument callable that returns a Rich
407 renderable. Called after each item to refresh the display.
408 The caller typically closes over a mutable dataclass and
409 rebuilds a Table from it on each call.
410 total: Total number of items. If None, attempts to determine
411 automatically via .count() (for querysets) or len().
413 Yields:
414 Items from the iterable.
416 Example:
417 @dataclass
418 class Stats:
419 processed: int = 0
420 failed: int = 0
422 stats = Stats()
424 def render_stats() -> Table:
425 table = Table(box=None)
426 table.add_column("Processed")
427 table.add_column("Failed")
428 table.add_row(str(stats.processed), str(stats.failed))
429 return table
431 for item in self.track_with_stats(
432 items,
433 description="Importing...",
434 stats_renderer=render_stats,
435 ):
436 try:
437 import_item(item)
438 stats.processed += 1
439 except Exception:
440 stats.failed += 1
441 """
442 if self.no_progress_bar:
443 yield from iterable
444 return
446 if total is None:
447 total = self._get_iterable_length(iterable)
449 stderr_console = Console(stderr=True)
451 # Progress is created without its own console so Live controls rendering.
452 progress = Progress(*self._progress_columns())
453 task_id = progress.add_task(description, total=total)
455 with Live(
456 Group(progress, stats_renderer()),
457 console=stderr_console,
458 refresh_per_second=4,
459 ) as live:
460 for item in iterable:
461 yield item
462 progress.advance(task_id)
463 live.update(Group(progress, stats_renderer()))
465 def process_parallel(
466 self,
467 fn: Callable[[T], R],
468 items: Sequence[T],
469 *,
470 description: str = "Processing...",
471 ) -> Generator[ProcessResult[T, R], None, None]:
472 """
473 Process items in parallel with progress tracking.
475 When --processes=1, runs sequentially in the main process without
476 spawning subprocesses. This is critical for testing, as multiprocessing
477 breaks fixtures, mocks, and database transactions.
479 When --processes > 1, uses ProcessPoolExecutor and automatically closes
480 database connections before spawning workers (required for PostgreSQL).
482 Args:
483 fn: Function to apply to each item. Must be picklable for parallel
484 execution (i.e., defined at module level, not a lambda or closure).
485 items: Sequence of items to process.
486 description: Text to display alongside the progress bar.
488 Yields:
489 ProcessResult for each item, containing the item, result, and any error.
491 Example:
492 def regenerate_thumbnail(doc_id: int) -> Path:
493 ...
495 for result in self.process_parallel(regenerate_thumbnail, doc_ids):
496 if result.error:
497 self.console.print(f"[red]Failed {result.item}[/red]")
498 """
499 total = len(items)
501 if self.process_count == 1:
502 # Sequential execution in main process - critical for testing, so we don't fork in fork, etc
503 yield from self._process_sequential(fn, items, description, total)
504 else:
505 # Parallel execution with ProcessPoolExecutor
506 yield from self._process_parallel(fn, items, description, total)
508 def _process_sequential(
509 self,
510 fn: Callable[[T], R],
511 items: Sequence[T],
512 description: str,
513 total: int,
514 ) -> Generator[ProcessResult[T, R], None, None]:
515 """Process items sequentially in the main process."""
516 for item in self.track(items, description=description, total=total):
517 try:
518 result = fn(item)
519 yield ProcessResult(item=item, result=result, error=None)
520 except Exception as e:
521 yield ProcessResult(item=item, result=None, error=e)
523 def _process_parallel(
524 self,
525 fn: Callable[[T], R],
526 items: Sequence[T],
527 description: str,
528 total: int,
529 ) -> Generator[ProcessResult[T, R], None, None]:
530 """Process items in parallel using ProcessPoolExecutor."""
532 # Close database connections before forking - required for PostgreSQL
533 db.connections.close_all()
535 with self._create_progress(description) as progress:
536 task_id = progress.add_task(description, total=total)
538 with ProcessPoolExecutor(
539 max_workers=self.process_count,
540 initializer=django.setup,
541 ) as executor:
542 # Submit all tasks and map futures back to items
543 future_to_item = {executor.submit(fn, item): item for item in items}
545 # Yield results as they complete
546 for future in as_completed(future_to_item):
547 item = future_to_item[future]
548 try:
549 result = future.result()
550 yield ProcessResult(item=item, result=result, error=None)
551 except Exception as e:
552 yield ProcessResult(item=item, result=None, error=e)
553 finally:
554 progress.advance(task_id)