caption-flow 0.4.2__tar.gz → 0.5.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (58) hide show
  1. {caption_flow-0.4.2/src/caption_flow.egg-info → caption_flow-0.5.0}/PKG-INFO +7 -6
  2. {caption_flow-0.4.2 → caption_flow-0.5.0}/README.md +2 -3
  3. {caption_flow-0.4.2 → caption_flow-0.5.0}/pyproject.toml +7 -5
  4. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/cli.py +32 -5
  5. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/models.py +14 -0
  6. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/workers/caption.py +106 -3
  7. {caption_flow-0.4.2 → caption_flow-0.5.0/src/caption_flow.egg-info}/PKG-INFO +7 -6
  8. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow.egg-info/requires.txt +2 -1
  9. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_cli.py +121 -0
  10. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_worker_caption.py +215 -4
  11. {caption_flow-0.4.2 → caption_flow-0.5.0}/LICENSE +0 -0
  12. {caption_flow-0.4.2 → caption_flow-0.5.0}/setup.cfg +0 -0
  13. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/__init__.py +0 -0
  14. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/monitor.py +0 -0
  15. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/orchestrator.py +0 -0
  16. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/processors/__init__.py +0 -0
  17. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/processors/base.py +0 -0
  18. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/processors/huggingface.py +0 -0
  19. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/processors/local_filesystem.py +0 -0
  20. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/processors/webdataset.py +0 -0
  21. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/storage/__init__.py +0 -0
  22. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/storage/exporter.py +0 -0
  23. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/storage/manager.py +0 -0
  24. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/utils/__init__.py +0 -0
  25. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/utils/auth.py +0 -0
  26. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/utils/caption_utils.py +0 -0
  27. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/utils/certificates.py +0 -0
  28. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/utils/checkpoint_tracker.py +0 -0
  29. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/utils/chunk_tracker.py +0 -0
  30. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/utils/image_processor.py +0 -0
  31. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/utils/json_utils.py +0 -0
  32. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/utils/prompt_template.py +0 -0
  33. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/utils/vllm_config.py +0 -0
  34. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/viewer.py +0 -0
  35. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/workers/base.py +0 -0
  36. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow/workers/data.py +0 -0
  37. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow.egg-info/SOURCES.txt +0 -0
  38. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow.egg-info/dependency_links.txt +0 -0
  39. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow.egg-info/entry_points.txt +0 -0
  40. {caption_flow-0.4.2 → caption_flow-0.5.0}/src/caption_flow.egg-info/top_level.txt +0 -0
  41. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_caption_utils.py +0 -0
  42. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_certificates.py +0 -0
  43. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_config_reload.py +0 -0
  44. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_duplicate_job_assignments.py +0 -0
  45. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_exporter.py +0 -0
  46. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_fix_verification.py +0 -0
  47. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_huggingface_ranges.py +0 -0
  48. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_json_utils.py +0 -0
  49. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_main.py +0 -0
  50. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_monitor.py +0 -0
  51. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_processors.py +0 -0
  52. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_range_level_distribution.py +0 -0
  53. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_storage_components.py +0 -0
  54. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_viewer.py +0 -0
  55. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_vllm_config.py +0 -0
  56. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_webdataset_ranges.py +0 -0
  57. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_worker_reconnection_complete.py +0 -0
  58. {caption_flow-0.4.2 → caption_flow-0.5.0}/tests/test_worker_reconnection_sequence.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: caption-flow
3
- Version: 0.4.2
3
+ Version: 0.5.0
4
4
  Summary: Self-contained distributed community captioning system
5
5
  Author-email: bghira <bghira@users.github.com>
6
6
  License: MIT
@@ -11,7 +11,8 @@ Classifier: License :: OSI Approved :: MIT License
11
11
  Classifier: Programming Language :: Python :: 3
12
12
  Classifier: Programming Language :: Python :: 3.11
13
13
  Classifier: Programming Language :: Python :: 3.12
14
- Requires-Python: <3.13,>=3.11
14
+ Classifier: Programming Language :: Python :: 3.13
15
+ Requires-Python: <3.14,>=3.11
15
16
  Description-Content-Type: text/markdown
16
17
  License-File: LICENSE
17
18
  Requires-Dist: websockets>=12.0
@@ -25,7 +26,8 @@ Requires-Dist: pyyaml>=6.0
25
26
  Requires-Dist: certbot>=2.0.0
26
27
  Requires-Dist: numpy>=1.24.0
27
28
  Requires-Dist: pillow>=10.0.0
28
- Requires-Dist: vllm<0.11.0,>=0.10.0
29
+ Requires-Dist: vllm<0.20.0,>=0.19.0
30
+ Requires-Dist: transformers<6.0.0,>=5.0.0
29
31
  Requires-Dist: webdataset<2.0.0,>=1.0.2
30
32
  Requires-Dist: pandas<3.0.0,>=2.3.1
31
33
  Requires-Dist: arrow<2.0.0,>=1.3.0
@@ -64,13 +66,12 @@ a fast websocket-based orchestrator paired with lightweight gpu workers achieves
64
66
 
65
67
  ---
66
68
 
67
- ## install
69
+ ## install from pypi
68
70
 
69
71
  ```bash
70
72
  python -m venv .venv
71
73
  source .venv/bin/activate # windows: .venv\Scripts\activate
72
- pip install --upgrade pip
73
- pip install -e . # installs the `caption-flow` command
74
+ pip install caption-flow
74
75
  ```
75
76
 
76
77
  ## quickstart (single box)
@@ -16,13 +16,12 @@ a fast websocket-based orchestrator paired with lightweight gpu workers achieves
16
16
 
17
17
  ---
18
18
 
19
- ## install
19
+ ## install from pypi
20
20
 
21
21
  ```bash
22
22
  python -m venv .venv
23
23
  source .venv/bin/activate # windows: .venv\Scripts\activate
24
- pip install --upgrade pip
25
- pip install -e . # installs the `caption-flow` command
24
+ pip install caption-flow
26
25
  ```
27
26
 
28
27
  ## quickstart (single box)
@@ -1,9 +1,9 @@
1
1
  [project]
2
2
  name = "caption-flow"
3
- version = "0.4.2"
3
+ version = "0.5.0"
4
4
  description = "Self-contained distributed community captioning system"
5
5
  readme = "README.md"
6
- requires-python = ">=3.11,<3.13"
6
+ requires-python = ">=3.11,<3.14"
7
7
  license = { text = "MIT" }
8
8
  authors = [{ name = "bghira", email = "bghira@users.github.com" }]
9
9
  keywords = ["captioning", "distributed", "vllm", "dataset", "community"]
@@ -14,6 +14,7 @@ classifiers = [
14
14
  "Programming Language :: Python :: 3",
15
15
  "Programming Language :: Python :: 3.11",
16
16
  "Programming Language :: Python :: 3.12",
17
+ "Programming Language :: Python :: 3.13",
17
18
  ]
18
19
 
19
20
  dependencies = [
@@ -28,7 +29,8 @@ dependencies = [
28
29
  "certbot>=2.0.0",
29
30
  "numpy>=1.24.0",
30
31
  "pillow>=10.0.0",
31
- "vllm (>=0.10.0,<0.11.0)",
32
+ "vllm (>=0.19.0,<0.20.0)",
33
+ "transformers (>=5.0.0,<6.0.0)",
32
34
  "webdataset (>=1.0.2,<2.0.0)",
33
35
  "pandas (>=2.3.1,<3.0.0)",
34
36
  "arrow (>=1.3.0,<2.0.0)",
@@ -64,7 +66,7 @@ where = ["src"]
64
66
 
65
67
  [tool.black]
66
68
  line-length = 100
67
- target-version = ['py310']
69
+ target-version = ['py311', 'py312', 'py313']
68
70
 
69
71
  [tool.ruff]
70
72
  line-length = 100
@@ -85,7 +87,7 @@ ignore = [
85
87
  ]
86
88
 
87
89
  [tool.mypy]
88
- python_version = "3.11"
90
+ python_version = "3.13"
89
91
  warn_return_any = true
90
92
  warn_unused_configs = true
91
93
  disallow_untyped_defs = true
@@ -214,6 +214,7 @@ def validate_orchestrator_auth_config(
214
214
 
215
215
  Raises:
216
216
  ValueError: If no auth configuration is found
217
+
217
218
  """
218
219
  # Check if auth is already in the orchestrator section
219
220
  if "auth" in config_data:
@@ -277,7 +278,7 @@ def validate_orchestrator_auth_config(
277
278
  "[yellow]Warning: No admin tokens configured - admin operations will be unavailable[/yellow]"
278
279
  )
279
280
 
280
- console.print(f"[green]✓ Auth validation passed[/green]")
281
+ console.print("[green]✓ Auth validation passed[/green]")
281
282
  return config_data
282
283
 
283
284
 
@@ -363,6 +364,13 @@ def orchestrator(ctx, config: Optional[str], **kwargs):
363
364
  @click.option("--gpu-id", type=int, help="GPU device ID (for vLLM)")
364
365
  @click.option("--precision", help="Model precision (for vLLM)")
365
366
  @click.option("--model", help="Model name (for vLLM)")
367
+ @click.option(
368
+ "--when_finished",
369
+ type=click.Choice(["stay_connected", "shutdown", "post_exec_hook"]),
370
+ default="stay_connected",
371
+ help="Action when all captions are complete (default: stay_connected)",
372
+ )
373
+ @click.option("--post_exec_hook", help="Path to executable for post_exec_hook action")
366
374
  @click.pass_context
367
375
  def worker(ctx, config: Optional[str], **kwargs):
368
376
  """Start a worker node."""
@@ -376,7 +384,18 @@ def worker(ctx, config: Optional[str], **kwargs):
376
384
  config_data = base_config
377
385
 
378
386
  # Apply CLI overrides (only non-None values)
379
- for key in ["server", "token", "name", "batch_size", "gpu_id", "precision", "model"]:
387
+ cli_overrides = [
388
+ "server",
389
+ "token",
390
+ "name",
391
+ "batch_size",
392
+ "gpu_id",
393
+ "precision",
394
+ "model",
395
+ "when_finished",
396
+ "post_exec_hook",
397
+ ]
398
+ for key in cli_overrides:
380
399
  if kwargs.get(key) is not None:
381
400
  config_data[key] = kwargs[key]
382
401
 
@@ -391,6 +410,14 @@ def worker(ctx, config: Optional[str], **kwargs):
391
410
  console.print("[red]Error: --token required (or set in config)[/red]")
392
411
  sys.exit(1)
393
412
 
413
+ # Validate when_finished logic
414
+ when_finished = config_data.get("when_finished", "stay_connected")
415
+ if when_finished == "post_exec_hook" and not config_data.get("post_exec_hook"):
416
+ console.print(
417
+ "[red]Error: --post_exec_hook required when --when_finished=post_exec_hook[/red]"
418
+ )
419
+ sys.exit(1)
420
+
394
421
  # Choose worker type
395
422
  if kwargs.get("vllm") or config_data.get("vllm"):
396
423
  from .workers.caption import CaptionWorker
@@ -853,7 +880,7 @@ def add(ctx, role: str, name: str, token_value: str, no_reload: bool):
853
880
  console.print("[yellow]No admin token specified, skipping orchestrator reload[/yellow]")
854
881
  console.print("[dim]Use --token to reload orchestrator config[/dim]")
855
882
  else:
856
- console.print(f"[cyan]Reloading orchestrator config...[/cyan]")
883
+ console.print("[cyan]Reloading orchestrator config...[/cyan]")
857
884
  success = asyncio.run(
858
885
  _reload_orchestrator_config(server, admin_token, config_data, no_verify_ssl)
859
886
  )
@@ -923,7 +950,7 @@ def remove(ctx, role: str, identifier: str, no_reload: bool):
923
950
  elif not admin_token:
924
951
  console.print("[yellow]No admin token specified, skipping orchestrator reload[/yellow]")
925
952
  else:
926
- console.print(f"[cyan]Reloading orchestrator config...[/cyan]")
953
+ console.print("[cyan]Reloading orchestrator config...[/cyan]")
927
954
  success = asyncio.run(
928
955
  _reload_orchestrator_config(server, admin_token, config_data, no_verify_ssl)
929
956
  )
@@ -1008,7 +1035,7 @@ def reload_config(
1008
1035
  sys.exit(1)
1009
1036
 
1010
1037
  # Validate and normalize auth configuration for reload
1011
- console.print(f"[cyan]Validating configuration...[/cyan]")
1038
+ console.print("[cyan]Validating configuration...[/cyan]")
1012
1039
  if "orchestrator" in new_cfg:
1013
1040
  orchestrator_config = new_cfg["orchestrator"]
1014
1041
  try:
@@ -253,6 +253,20 @@ class StorageContents:
253
253
  logger.warning(f"Row missing columns: {missing_cols}")
254
254
 
255
255
 
256
+ class WhenFinished(Enum):
257
+ """Actions to take when caption worker finishes all work."""
258
+
259
+ STAY_CONNECTED = "stay_connected"
260
+ SHUTDOWN = "shutdown"
261
+ POST_EXEC_HOOK = "post_exec_hook"
262
+
263
+ def __str__(self):
264
+ return self.value
265
+
266
+ def to_json(self):
267
+ return self.value
268
+
269
+
256
270
  class ExportError(Exception):
257
271
  """Base exception for export-related errors."""
258
272
 
@@ -18,7 +18,7 @@ import websockets
18
18
  from huggingface_hub import get_token
19
19
  from PIL import Image
20
20
 
21
- from ..models import ProcessingStage, StageResult
21
+ from ..models import ProcessingStage, StageResult, WhenFinished
22
22
  from ..processors import (
23
23
  HuggingFaceDatasetWorkerProcessor,
24
24
  LocalFilesystemWorkerProcessor,
@@ -93,9 +93,11 @@ class MultiStageVLLMManager:
93
93
 
94
94
  # Load tokenizer and processor
95
95
  self.tokenizers[model_name] = AutoTokenizer.from_pretrained(
96
- model_name, trust_remote_code=True, use_fast=True
96
+ model_name, trust_remote_code=True
97
+ )
98
+ self.processors[model_name] = AutoProcessor.from_pretrained(
99
+ model_name, trust_remote_code=True
97
100
  )
98
- self.processors[model_name] = AutoProcessor.from_pretrained(model_name)
99
101
 
100
102
  # Initialize LLM
101
103
  vllm_params = {
@@ -217,6 +219,15 @@ class CaptionWorker(BaseWorker):
217
219
  # Processing control
218
220
  self.should_stop_processing = Event()
219
221
 
222
+ # When finished behavior
223
+ when_finished_str = config.get("when_finished", "stay_connected")
224
+ self.when_finished = WhenFinished(when_finished_str)
225
+ self.post_exec_hook = config.get("post_exec_hook")
226
+
227
+ # Track consecutive no work responses to detect completion
228
+ self.consecutive_no_work = 0
229
+ self.no_work_threshold = 3 # Consider work complete after 3 consecutive no_work responses
230
+
220
231
  def _init_metrics(self):
221
232
  """Initialize worker metrics."""
222
233
  self.items_processed = 0
@@ -347,13 +358,105 @@ class CaptionWorker(BaseWorker):
347
358
  self.assigned_units.append(unit)
348
359
  logger.info(f"Received {len(assignment.units)} work units")
349
360
 
361
+ # Reset no work counter since we got work
362
+ self.consecutive_no_work = 0
363
+
350
364
  elif msg_type == "no_work":
351
365
  logger.info("No work available")
366
+ self.consecutive_no_work += 1
367
+
368
+ # Check if all work appears to be complete
369
+ if self.consecutive_no_work >= self.no_work_threshold:
370
+ logger.info(
371
+ f"Received {self.consecutive_no_work} consecutive 'no work' responses, "
372
+ "assuming all captions are complete"
373
+ )
374
+ await self._handle_work_completion()
375
+ return
376
+
352
377
  await asyncio.sleep(10)
353
378
 
354
379
  if self.websocket and self.connected.is_set():
355
380
  await self.websocket.send(json.dumps({"type": "get_work_units", "count": 2}))
356
381
 
382
+ async def _handle_work_completion(self):
383
+ """Handle when all captions are complete based on when_finished setting."""
384
+ logger.info(f"Handling work completion with action: {self.when_finished}")
385
+
386
+ if self.when_finished == WhenFinished.STAY_CONNECTED:
387
+ logger.info("Staying connected and continuing to poll for work")
388
+ return
389
+
390
+ elif self.when_finished == WhenFinished.SHUTDOWN:
391
+ logger.info("All captions complete, shutting down worker")
392
+ await self._shutdown_worker()
393
+
394
+ elif self.when_finished == WhenFinished.POST_EXEC_HOOK:
395
+ if not self.post_exec_hook:
396
+ logger.error("post_exec_hook action specified but no hook path provided")
397
+ return
398
+
399
+ logger.info(f"All captions complete, executing hook: {self.post_exec_hook}")
400
+ await self._execute_post_hook()
401
+
402
+ async def _shutdown_worker(self):
403
+ """Clean up and shut down the worker."""
404
+ logger.info("Shutting down worker...")
405
+
406
+ # Stop processing
407
+ self.should_stop_processing.set()
408
+
409
+ # Set running to False to exit the main loop
410
+ self.running = False
411
+
412
+ # Disconnect gracefully
413
+ if self.websocket:
414
+ try:
415
+ await self.websocket.close()
416
+ except Exception as e:
417
+ logger.warning(f"Error closing websocket: {e}")
418
+
419
+ async def _execute_post_hook(self):
420
+ """Execute the post-execution hook."""
421
+ import os
422
+ from pathlib import Path
423
+
424
+ hook_path = Path(self.post_exec_hook)
425
+
426
+ if not hook_path.exists():
427
+ logger.error(f"Post-exec hook not found: {hook_path}")
428
+ return
429
+
430
+ if not os.access(hook_path, os.X_OK):
431
+ logger.error(f"Post-exec hook is not executable: {hook_path}")
432
+ return
433
+
434
+ try:
435
+ logger.info(f"Executing post-exec hook: {hook_path}")
436
+
437
+ # Execute the hook
438
+ process = await asyncio.create_subprocess_exec(
439
+ str(hook_path), stdout=asyncio.subprocess.PIPE, stderr=asyncio.subprocess.PIPE
440
+ )
441
+
442
+ stdout, stderr = await process.communicate()
443
+
444
+ if process.returncode == 0:
445
+ logger.info("Post-exec hook completed successfully")
446
+ if stdout:
447
+ logger.info(f"Hook output: {stdout.decode()}")
448
+ else:
449
+ logger.error(f"Post-exec hook failed with return code {process.returncode}")
450
+ if stderr:
451
+ logger.error(f"Hook error: {stderr.decode()}")
452
+
453
+ except Exception as e:
454
+ logger.error(f"Error executing post-exec hook: {e}")
455
+
456
+ finally:
457
+ # Shut down after executing hook
458
+ await self._shutdown_worker()
459
+
357
460
  def _parse_stages_config(self, vllm_config: Dict[str, Any]) -> List[ProcessingStage]:
358
461
  """Parse stages configuration from vLLM config."""
359
462
  stages_config = vllm_config.get("stages", [])
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: caption-flow
3
- Version: 0.4.2
3
+ Version: 0.5.0
4
4
  Summary: Self-contained distributed community captioning system
5
5
  Author-email: bghira <bghira@users.github.com>
6
6
  License: MIT
@@ -11,7 +11,8 @@ Classifier: License :: OSI Approved :: MIT License
11
11
  Classifier: Programming Language :: Python :: 3
12
12
  Classifier: Programming Language :: Python :: 3.11
13
13
  Classifier: Programming Language :: Python :: 3.12
14
- Requires-Python: <3.13,>=3.11
14
+ Classifier: Programming Language :: Python :: 3.13
15
+ Requires-Python: <3.14,>=3.11
15
16
  Description-Content-Type: text/markdown
16
17
  License-File: LICENSE
17
18
  Requires-Dist: websockets>=12.0
@@ -25,7 +26,8 @@ Requires-Dist: pyyaml>=6.0
25
26
  Requires-Dist: certbot>=2.0.0
26
27
  Requires-Dist: numpy>=1.24.0
27
28
  Requires-Dist: pillow>=10.0.0
28
- Requires-Dist: vllm<0.11.0,>=0.10.0
29
+ Requires-Dist: vllm<0.20.0,>=0.19.0
30
+ Requires-Dist: transformers<6.0.0,>=5.0.0
29
31
  Requires-Dist: webdataset<2.0.0,>=1.0.2
30
32
  Requires-Dist: pandas<3.0.0,>=2.3.1
31
33
  Requires-Dist: arrow<2.0.0,>=1.3.0
@@ -64,13 +66,12 @@ a fast websocket-based orchestrator paired with lightweight gpu workers achieves
64
66
 
65
67
  ---
66
68
 
67
- ## install
69
+ ## install from pypi
68
70
 
69
71
  ```bash
70
72
  python -m venv .venv
71
73
  source .venv/bin/activate # windows: .venv\Scripts\activate
72
- pip install --upgrade pip
73
- pip install -e . # installs the `caption-flow` command
74
+ pip install caption-flow
74
75
  ```
75
76
 
76
77
  ## quickstart (single box)
@@ -9,7 +9,8 @@ pyyaml>=6.0
9
9
  certbot>=2.0.0
10
10
  numpy>=1.24.0
11
11
  pillow>=10.0.0
12
- vllm<0.11.0,>=0.10.0
12
+ vllm<0.20.0,>=0.19.0
13
+ transformers<6.0.0,>=5.0.0
13
14
  webdataset<2.0.0,>=1.0.2
14
15
  pandas<3.0.0,>=2.3.1
15
16
  arrow<2.0.0,>=1.3.0
@@ -857,5 +857,126 @@ class TestReloadConfigCommand:
857
857
  assert "✓ Auth validation passed" in result.output
858
858
 
859
859
 
860
+ class TestWhenFinishedCLI:
861
+ """Test CLI validation for when_finished parameters."""
862
+
863
+ def test_worker_command_when_finished_validation_success(self, runner):
864
+ """Test successful worker command with when_finished parameters."""
865
+ with patch("caption_flow.workers.caption.CaptionWorker") as mock_worker:
866
+ mock_worker_instance = Mock()
867
+ mock_worker.return_value = mock_worker_instance
868
+ mock_worker_instance.start = AsyncMock()
869
+
870
+ # Test valid combinations
871
+ result = runner.invoke(
872
+ main,
873
+ [
874
+ "worker",
875
+ "--server",
876
+ "ws://localhost:8765",
877
+ "--token",
878
+ "test-token",
879
+ "--vllm",
880
+ "--when_finished",
881
+ "stay_connected",
882
+ ],
883
+ )
884
+ # This will fail due to async mocking, but should pass validation
885
+ assert "--post_exec_hook required" not in result.output
886
+
887
+ def test_worker_command_when_finished_shutdown(self, runner):
888
+ """Test worker command with shutdown option."""
889
+ with patch("caption_flow.workers.caption.CaptionWorker") as mock_worker:
890
+ mock_worker_instance = Mock()
891
+ mock_worker.return_value = mock_worker_instance
892
+ mock_worker_instance.start = AsyncMock()
893
+
894
+ result = runner.invoke(
895
+ main,
896
+ [
897
+ "worker",
898
+ "--server",
899
+ "ws://localhost:8765",
900
+ "--token",
901
+ "test-token",
902
+ "--vllm",
903
+ "--when_finished",
904
+ "shutdown",
905
+ ],
906
+ )
907
+ assert "--post_exec_hook required" not in result.output
908
+
909
+ def test_worker_command_when_finished_post_hook_missing(self, runner):
910
+ """Test worker command fails when post_exec_hook is missing."""
911
+ result = runner.invoke(
912
+ main,
913
+ [
914
+ "worker",
915
+ "--server",
916
+ "ws://localhost:8765",
917
+ "--token",
918
+ "test-token",
919
+ "--vllm",
920
+ "--when_finished",
921
+ "post_exec_hook",
922
+ # Missing --post_exec_hook
923
+ ],
924
+ )
925
+
926
+ assert result.exit_code == 1
927
+ assert "--post_exec_hook required when --when_finished=post_exec_hook" in result.output
928
+
929
+ def test_worker_command_when_finished_post_hook_valid(self, runner, tmp_path):
930
+ """Test worker command with valid post_exec_hook."""
931
+ hook_script = tmp_path / "test_hook.sh"
932
+ hook_script.write_text("#!/bin/bash\necho 'test'")
933
+ hook_script.chmod(0o755)
934
+
935
+ with patch("caption_flow.workers.caption.CaptionWorker") as mock_worker:
936
+ mock_worker_instance = Mock()
937
+ mock_worker.return_value = mock_worker_instance
938
+ mock_worker_instance.start = AsyncMock()
939
+
940
+ result = runner.invoke(
941
+ main,
942
+ [
943
+ "worker",
944
+ "--server",
945
+ "ws://localhost:8765",
946
+ "--token",
947
+ "test-token",
948
+ "--vllm",
949
+ "--when_finished",
950
+ "post_exec_hook",
951
+ "--post_exec_hook",
952
+ str(hook_script),
953
+ ],
954
+ )
955
+ # Should pass validation
956
+ assert "--post_exec_hook required" not in result.output
957
+
958
+ def test_worker_command_when_finished_default(self, runner):
959
+ """Test worker command uses default when_finished value."""
960
+ with patch("caption_flow.workers.caption.CaptionWorker") as mock_worker:
961
+ mock_worker_instance = Mock()
962
+ mock_worker.return_value = mock_worker_instance
963
+ mock_worker_instance.start = AsyncMock()
964
+
965
+ result = runner.invoke(
966
+ main,
967
+ [
968
+ "worker",
969
+ "--server",
970
+ "ws://localhost:8765",
971
+ "--token",
972
+ "test-token",
973
+ "--vllm",
974
+ # No when_finished specified, should default to stay_connected
975
+ ],
976
+ )
977
+ # Should work without errors
978
+ assert "--post_exec_hook required" not in result.output
979
+
980
+
860
981
  if __name__ == "__main__":
861
982
  pytest.main([__file__])
@@ -12,7 +12,7 @@ from PIL import Image
12
12
  # Import pytest-asyncio
13
13
  pytest_plugins = ("pytest_asyncio",)
14
14
  import pytest_asyncio
15
- from caption_flow.models import Caption, JobId, ProcessingStage
15
+ from caption_flow.models import Caption, JobId, ProcessingStage, WhenFinished
16
16
  from caption_flow.processors import WorkAssignment, WorkUnit
17
17
  from caption_flow.storage import StorageManager
18
18
 
@@ -859,9 +859,9 @@ class TestCaptionWorker:
859
859
  work_failed_call = sent_data
860
860
  break
861
861
 
862
- assert (
863
- work_failed_call is not None
864
- ), "work_failed message should have been sent for incomplete unit"
862
+ assert work_failed_call is not None, (
863
+ "work_failed message should have been sent for incomplete unit"
864
+ )
865
865
  assert work_failed_call["unit_id"] == "unit1"
866
866
  assert "Processing incomplete" in work_failed_call["error"]
867
867
  assert "1/3 items processed" in work_failed_call["error"]
@@ -1119,5 +1119,216 @@ class TestCaptionWorkerConfigReload:
1119
1119
  assert worker.items_failed == 1
1120
1120
 
1121
1121
 
1122
+ class TestWhenFinishedFunctionality:
1123
+ """Test suite for when_finished functionality."""
1124
+
1125
+ @pytest.fixture
1126
+ def worker_config_stay_connected(self):
1127
+ """Create test worker config with stay_connected behavior."""
1128
+ return {
1129
+ "name": "test_worker",
1130
+ "token": "test_token",
1131
+ "server": "ws://localhost:8765",
1132
+ "server_url": "ws://localhost:8765",
1133
+ "gpu_id": 0,
1134
+ "batch_image_processing": True,
1135
+ "when_finished": "stay_connected",
1136
+ }
1137
+
1138
+ @pytest.fixture
1139
+ def worker_config_shutdown(self):
1140
+ """Create test worker config with shutdown behavior."""
1141
+ return {
1142
+ "name": "test_worker",
1143
+ "token": "test_token",
1144
+ "server": "ws://localhost:8765",
1145
+ "server_url": "ws://localhost:8765",
1146
+ "gpu_id": 0,
1147
+ "batch_image_processing": True,
1148
+ "when_finished": "shutdown",
1149
+ }
1150
+
1151
+ @pytest.fixture
1152
+ def worker_config_post_hook(self, tmp_path):
1153
+ """Create test worker config with post_exec_hook behavior."""
1154
+ hook_script = tmp_path / "test_hook.sh"
1155
+ hook_script.write_text("#!/bin/bash\necho 'Hook executed'\nexit 0")
1156
+ hook_script.chmod(0o755)
1157
+
1158
+ return {
1159
+ "name": "test_worker",
1160
+ "token": "test_token",
1161
+ "server": "ws://localhost:8765",
1162
+ "server_url": "ws://localhost:8765",
1163
+ "gpu_id": 0,
1164
+ "batch_image_processing": True,
1165
+ "when_finished": "post_exec_hook",
1166
+ "post_exec_hook": str(hook_script),
1167
+ }
1168
+
1169
+ def test_when_finished_initialization_stay_connected(self, worker_config_stay_connected):
1170
+ """Test worker initialization with stay_connected setting."""
1171
+ worker = create_fast_caption_worker(worker_config_stay_connected)
1172
+
1173
+ assert worker.when_finished == WhenFinished.STAY_CONNECTED
1174
+ assert worker.post_exec_hook is None
1175
+ assert worker.consecutive_no_work == 0
1176
+ assert worker.no_work_threshold == 3
1177
+
1178
+ def test_when_finished_initialization_shutdown(self, worker_config_shutdown):
1179
+ """Test worker initialization with shutdown setting."""
1180
+ worker = create_fast_caption_worker(worker_config_shutdown)
1181
+
1182
+ assert worker.when_finished == WhenFinished.SHUTDOWN
1183
+ assert worker.post_exec_hook is None
1184
+
1185
+ def test_when_finished_initialization_post_hook(self, worker_config_post_hook):
1186
+ """Test worker initialization with post_exec_hook setting."""
1187
+ worker = create_fast_caption_worker(worker_config_post_hook)
1188
+
1189
+ assert worker.when_finished == WhenFinished.POST_EXEC_HOOK
1190
+ assert worker.post_exec_hook is not None
1191
+ assert "test_hook.sh" in worker.post_exec_hook
1192
+
1193
+ @pytest.mark.asyncio
1194
+ async def test_consecutive_no_work_counter(self, worker_config_stay_connected):
1195
+ """Test consecutive no_work counter logic."""
1196
+ worker = create_fast_caption_worker(worker_config_stay_connected)
1197
+
1198
+ # Mock websocket
1199
+ mock_websocket = AsyncMock()
1200
+ worker.websocket = mock_websocket
1201
+ worker.connected.set()
1202
+
1203
+ # Simulate receiving work assignment - should reset counter
1204
+ assignment_data = {
1205
+ "type": "work_assignment",
1206
+ "assignment": {
1207
+ "assignment_id": "test_assignment",
1208
+ "worker_id": "test_worker",
1209
+ "units": [],
1210
+ "assigned_at": datetime.now().isoformat(),
1211
+ },
1212
+ }
1213
+
1214
+ worker.consecutive_no_work = 2
1215
+ await worker._handle_message(assignment_data)
1216
+ assert worker.consecutive_no_work == 0
1217
+
1218
+ # Simulate receiving no_work messages
1219
+ no_work_data = {"type": "no_work"}
1220
+
1221
+ # First no_work - should increment counter but not trigger completion
1222
+ await worker._handle_message(no_work_data)
1223
+ assert worker.consecutive_no_work == 1
1224
+
1225
+ # Second no_work - should increment counter but not trigger completion
1226
+ await worker._handle_message(no_work_data)
1227
+ assert worker.consecutive_no_work == 2
1228
+
1229
+ # Mock the completion handler to avoid actual shutdown
1230
+ with patch.object(worker, "_handle_work_completion") as mock_completion:
1231
+ # Third no_work - should trigger completion
1232
+ await worker._handle_message(no_work_data)
1233
+ assert worker.consecutive_no_work == 3
1234
+ mock_completion.assert_called_once()
1235
+
1236
+ @pytest.mark.asyncio
1237
+ async def test_handle_work_completion_stay_connected(self, worker_config_stay_connected):
1238
+ """Test work completion handling with stay_connected action."""
1239
+ worker = create_fast_caption_worker(worker_config_stay_connected)
1240
+
1241
+ # Should just log and return
1242
+ await worker._handle_work_completion()
1243
+ # No assertions needed - just verify it doesn't crash
1244
+
1245
+ @pytest.mark.asyncio
1246
+ async def test_handle_work_completion_shutdown(self, worker_config_shutdown):
1247
+ """Test work completion handling with shutdown action."""
1248
+ worker = create_fast_caption_worker(worker_config_shutdown)
1249
+
1250
+ # Mock websocket
1251
+ mock_websocket = AsyncMock()
1252
+ worker.websocket = mock_websocket
1253
+
1254
+ await worker._handle_work_completion()
1255
+
1256
+ # Should set running to False and stop processing
1257
+ assert worker.running is False
1258
+ assert worker.should_stop_processing.is_set()
1259
+
1260
+ @pytest.mark.asyncio
1261
+ async def test_handle_work_completion_post_hook_success(self, worker_config_post_hook):
1262
+ """Test work completion handling with successful post_exec_hook."""
1263
+ worker = create_fast_caption_worker(worker_config_post_hook)
1264
+
1265
+ # Mock websocket
1266
+ mock_websocket = AsyncMock()
1267
+ worker.websocket = mock_websocket
1268
+
1269
+ await worker._handle_work_completion()
1270
+
1271
+ # Should execute hook and then shutdown
1272
+ assert worker.running is False
1273
+ assert worker.should_stop_processing.is_set()
1274
+
1275
+ @pytest.mark.asyncio
1276
+ async def test_handle_work_completion_post_hook_missing_path(self, worker_config_shutdown):
1277
+ """Test work completion handling when post_exec_hook path is missing."""
1278
+ worker = create_fast_caption_worker(worker_config_shutdown)
1279
+ worker.when_finished = WhenFinished.POST_EXEC_HOOK
1280
+ worker.post_exec_hook = None # Missing hook path
1281
+
1282
+ # Store initial state
1283
+ initial_running_state = worker.running
1284
+
1285
+ # Should log error and return without crashing
1286
+ await worker._handle_work_completion()
1287
+
1288
+ # Should not change running state when hook path is missing
1289
+ assert worker.running == initial_running_state
1290
+ assert worker.should_stop_processing.is_set() is False
1291
+
1292
+ @pytest.mark.asyncio
1293
+ async def test_execute_post_hook_nonexistent_file(self, tmp_path):
1294
+ """Test post-exec hook with nonexistent file."""
1295
+ config = {
1296
+ "name": "test_worker",
1297
+ "token": "test_token",
1298
+ "server": "ws://localhost:8765",
1299
+ "server_url": "ws://localhost:8765",
1300
+ "gpu_id": 0,
1301
+ "when_finished": "post_exec_hook",
1302
+ "post_exec_hook": str(tmp_path / "nonexistent.sh"),
1303
+ }
1304
+
1305
+ worker = create_fast_caption_worker(config)
1306
+
1307
+ # Should handle missing file gracefully
1308
+ await worker._execute_post_hook()
1309
+
1310
+ @pytest.mark.asyncio
1311
+ async def test_execute_post_hook_not_executable(self, tmp_path):
1312
+ """Test post-exec hook with non-executable file."""
1313
+ hook_script = tmp_path / "non_executable.sh"
1314
+ hook_script.write_text("#!/bin/bash\necho 'test'")
1315
+ # Don't make it executable
1316
+
1317
+ config = {
1318
+ "name": "test_worker",
1319
+ "token": "test_token",
1320
+ "server": "ws://localhost:8765",
1321
+ "server_url": "ws://localhost:8765",
1322
+ "gpu_id": 0,
1323
+ "when_finished": "post_exec_hook",
1324
+ "post_exec_hook": str(hook_script),
1325
+ }
1326
+
1327
+ worker = create_fast_caption_worker(config)
1328
+
1329
+ # Should handle non-executable file gracefully
1330
+ await worker._execute_post_hook()
1331
+
1332
+
1122
1333
  if __name__ == "__main__":
1123
1334
  pytest.main([__file__, "-v", "-s"])
File without changes
File without changes