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

1""" 

2Base command class for Paperless-ngx management commands. 

3 

4Provides automatic progress bar and multiprocessing support with minimal boilerplate. 

5""" 

6 

7from __future__ import annotations 

8 

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 

23 

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 

43 

44if TYPE_CHECKING: 

45 from collections.abc import Generator 

46 from collections.abc import Sequence 

47 

48 from django.core.management import CommandParser 

49 

50T = TypeVar("T") 

51R = TypeVar("R") 

52 

53 

54@dataclass(slots=True, frozen=True) 

55class _BufferedRecord: 

56 level: int 

57 name: str 

58 message: str 

59 

60 

61class BufferingLogHandler(logging.Handler): 

62 """Captures log records during a command run for deferred rendering. 

63 

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 """ 

68 

69 def __init__(self) -> None: 

70 super().__init__() 

71 self._records: list[_BufferedRecord] = [] 

72 

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 ) 

81 

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 

92 

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) 

103 

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 } 

111 

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 ) 

119 

120 console.print(table) 

121 

122 def clear(self) -> None: 

123 self._records.clear() 

124 

125 

126@dataclass(frozen=True, slots=True) 

127class ProcessResult(Generic[T, R]): 

128 """ 

129 Result of processing a single item in parallel. 

130 

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 """ 

136 

137 item: T 

138 result: R | None 

139 error: BaseException | None 

140 

141 @property 

142 def success(self) -> bool: 

143 """Return True if the item was processed successfully.""" 

144 return self.error is None 

145 

146 

147class PaperlessCommand(RichCommand): 

148 """ 

149 Base command class with automatic progress bar and multiprocessing support. 

150 

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) 

154 

155 Example usage: 

156 

157 class Command(PaperlessCommand): 

158 help = "Process all documents" 

159 

160 def handle(self, *args, **options): 

161 documents = Document.objects.all() 

162 for doc in self.track(documents, description="Processing..."): 

163 process_document(doc) 

164 

165 class Command(PaperlessCommand): 

166 help = "Regenerate thumbnails" 

167 supports_multiprocessing = True 

168 

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]") 

174 

175 class Command(PaperlessCommand): 

176 help = "Import documents with live stats" 

177 

178 def handle(self, *args, **options): 

179 stats = ImportStats() 

180 

181 def render_stats() -> Table: 

182 ... # build Rich Table from stats 

183 

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 """ 

192 

193 supports_progress_bar: ClassVar[bool] = True 

194 supports_multiprocessing: ClassVar[bool] = False 

195 

196 # Instance attributes set by execute() before handle() runs 

197 no_progress_bar: bool 

198 process_count: int 

199 

200 def add_arguments(self, parser: CommandParser) -> None: 

201 """Add arguments based on supported features.""" 

202 super().add_arguments(parser) 

203 

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 ) 

211 

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 ) 

220 

221 def execute(self, *args: Any, **options: Any) -> str | None: 

222 """ 

223 Set up instance state before handle() is called. 

224 

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 

232 

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 

239 

240 return super().execute(*args, **options) 

241 

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. 

249 

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. 

254 

255 Usage:: 

256 

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")) 

264 

265 loggers: list[logging.Logger] = [] 

266 original_propagate: dict[str, bool] = {} 

267 

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) 

274 

275 try: 

276 yield handler 

277 finally: 

278 for log in loggers: 

279 log.removeHandler(handler) 

280 log.propagate = original_propagate[log.name] 

281 

282 @staticmethod 

283 def _progress_columns() -> tuple[Any, ...]: 

284 """ 

285 Return the standard set of progress bar columns. 

286 

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 ) 

298 

299 def _create_progress(self, description: str) -> Progress: 

300 """ 

301 Create a standalone Progress instance with its own stderr Console. 

302 

303 Use this for track(). For track_with_stats(), Progress is created 

304 directly inside a Live context instead. 

305 

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. 

311 

312 Args: 

313 description: Text to display alongside the progress bar. 

314 

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 ) 

323 

324 def _get_iterable_length(self, iterable: Iterable[object]) -> int | None: 

325 """ 

326 Attempt to determine the length of an iterable without consuming it. 

327 

328 Tries .count() first (for Django querysets - executes SELECT COUNT(*)), 

329 then falls back to len() for sequences. 

330 

331 Args: 

332 iterable: The iterable to measure. 

333 

334 Returns: 

335 The length if determinable, None otherwise. 

336 """ 

337 if isinstance(iterable, QuerySet): 

338 return iterable.count() 

339 

340 if isinstance(iterable, Sized): 

341 return len(iterable) 

342 

343 return None 

344 

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. 

354 

355 Respects --no-progress-bar flag. When disabled, simply yields items 

356 without any progress display. 

357 

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(). 

363 

364 Yields: 

365 Items from the iterable. 

366 

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 

374 

375 if total is None: 

376 total = self._get_iterable_length(iterable) 

377 

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) 

383 

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. 

394 

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. 

398 

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). 

402 

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(). 

412 

413 Yields: 

414 Items from the iterable. 

415 

416 Example: 

417 @dataclass 

418 class Stats: 

419 processed: int = 0 

420 failed: int = 0 

421 

422 stats = Stats() 

423 

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 

430 

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 

445 

446 if total is None: 

447 total = self._get_iterable_length(iterable) 

448 

449 stderr_console = Console(stderr=True) 

450 

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) 

454 

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())) 

464 

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. 

474 

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. 

478 

479 When --processes > 1, uses ProcessPoolExecutor and automatically closes 

480 database connections before spawning workers (required for PostgreSQL). 

481 

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. 

487 

488 Yields: 

489 ProcessResult for each item, containing the item, result, and any error. 

490 

491 Example: 

492 def regenerate_thumbnail(doc_id: int) -> Path: 

493 ... 

494 

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) 

500 

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) 

507 

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) 

522 

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.""" 

531 

532 # Close database connections before forking - required for PostgreSQL 

533 db.connections.close_all() 

534 

535 with self._create_progress(description) as progress: 

536 task_id = progress.add_task(description, total=total) 

537 

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} 

544 

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)