/
/
1"""Task execution context helpers for long running background tasks."""
2
3from __future__ import annotations
4
5from contextvars import ContextVar
6from dataclasses import dataclass
7from typing import TYPE_CHECKING
8
9if TYPE_CHECKING:
10 from collections.abc import Callable
11
12 from music_assistant_models.background_task import BackgroundTask
13
14
15@dataclass(slots=True)
16class TaskExecutionContext:
17 """Runtime task context exposed to long-running task code."""
18
19 task_id: str
20 get_task: Callable[[str], BackgroundTask]
21 update_progress: Callable[[str, int | None, str | None], None]
22 update_progress_text: Callable[[str, str | None], None]
23 add_failure: Callable[[str, str], None]
24
25 @property
26 def task(self) -> BackgroundTask:
27 """Return the attached background task object."""
28 return self.get_task(self.task_id)
29
30 def set_progress(self, progress: int | None, text: str | None = None) -> None:
31 """Set an absolute progress percentage and optional phase text."""
32 self.update_progress(self.task_id, progress, text)
33
34 def set_progress_text(self, text: str | None) -> None:
35 """Update the human-readable progress text only."""
36 self.update_progress_text(self.task_id, text)
37
38 def set_progress_from_index(self, current: int, total: int, text: str | None = None) -> int:
39 """Set progress from the current item index and total item count."""
40 progress = calculate_progress(current, total)
41 self.update_progress(self.task_id, progress, text)
42 return progress
43
44 def record_failure(self, message: str) -> None:
45 """Record a non-fatal failure for the current task."""
46 self.add_failure(self.task_id, message)
47
48
49ACTIVE_TASK_CONTEXT: ContextVar[TaskExecutionContext | None] = ContextVar(
50 "active_background_task_context",
51 default=None,
52)
53
54
55def calculate_progress(current: int, total: int) -> int:
56 """Convert the current item index and total item count into a percentage."""
57 if total <= 0:
58 raise ValueError("Task progress total must be > 0")
59 current = max(current, 0)
60 return min(int((current * 100) / total), 100)
61
62
63def get_current_task_context() -> TaskExecutionContext | None:
64 """Return the task context active in the current async/thread context."""
65 return ACTIVE_TASK_CONTEXT.get()
66
67
68def get_current_task() -> BackgroundTask | None:
69 """Return the active background task for the current async/thread context."""
70 if task_context := get_current_task_context():
71 return task_context.task
72 return None
73
74
75def get_current_task_id() -> str | None:
76 """Return the active task id for the current async/thread context."""
77 if task_context := get_current_task_context():
78 return task_context.task_id
79 return None
80
81
82def update_current_task_progress(progress: int | None, text: str | None = None) -> None:
83 """Update progress for the task active in the current async/thread context."""
84 if task_context := get_current_task_context():
85 task_context.set_progress(progress, text)
86
87
88def update_current_task_progress_text(text: str | None) -> None:
89 """Update progress text for the task active in the current async/thread context."""
90 if task_context := get_current_task_context():
91 task_context.set_progress_text(text)
92
93
94def update_current_task_progress_from_index(
95 current: int, total: int, text: str | None = None
96) -> int | None:
97 """Update progress from item counts for the current async/thread context."""
98 if task_context := get_current_task_context():
99 return task_context.set_progress_from_index(current, total, text)
100 return None
101
102
103def report_current_task_failure(message: str) -> None:
104 """Record a non-fatal failure for the current async/thread context."""
105 if task_context := get_current_task_context():
106 task_context.record_failure(message)
107