diffusers-workflow 0.4.0a3__py3-none-any.whl
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.
- diffusers_workflow-0.4.0a3.dist-info/METADATA +310 -0
- diffusers_workflow-0.4.0a3.dist-info/RECORD +171 -0
- diffusers_workflow-0.4.0a3.dist-info/WHEEL +5 -0
- diffusers_workflow-0.4.0a3.dist-info/entry_points.txt +6 -0
- diffusers_workflow-0.4.0a3.dist-info/licenses/LICENSE +201 -0
- diffusers_workflow-0.4.0a3.dist-info/top_level.txt +1 -0
- dw/__init__.py +353 -0
- dw/arguments.py +906 -0
- dw/cache_blocks.json +16 -0
- dw/cache_blocks.py +145 -0
- dw/community_pipelines/pipeline_flux_rf_inversion.py +1184 -0
- dw/events.py +78 -0
- dw/hub_cache.py +289 -0
- dw/introspection.py +458 -0
- dw/log_setup.py +45 -0
- dw/pipeline_processors/chain.py +750 -0
- dw/pipeline_processors/config_objects.py +235 -0
- dw/pipeline_processors/pipeline.py +1687 -0
- dw/pipeline_processors/remote.py +18 -0
- dw/previous_results.py +259 -0
- dw/prompt_weighting.py +378 -0
- dw/repl.py +298 -0
- dw/repl_commands.py +808 -0
- dw/repl_worker.py +129 -0
- dw/result.py +850 -0
- dw/run.py +92 -0
- dw/schema.py +24 -0
- dw/security.py +379 -0
- dw/serve.py +70 -0
- dw/server/__init__.py +2 -0
- dw/server/app.py +588 -0
- dw/server/jobs.py +547 -0
- dw/server/ui/assets/abap-08VXUWAP.js +1 -0
- dw/server/ui/assets/apex-BWPQTe0t.js +1 -0
- dw/server/ui/assets/azcli-Bc_sGQ0U.js +1 -0
- dw/server/ui/assets/bat-i0X4ZdIN.js +1 -0
- dw/server/ui/assets/bicep-B5-_aFwp.js +2 -0
- dw/server/ui/assets/cameligo-DMUM7wLl.js +1 -0
- dw/server/ui/assets/clojure-Cm7r79vr.js +1 -0
- dw/server/ui/assets/codicon-Brq4_Ui5.ttf +0 -0
- dw/server/ui/assets/coffee-Ba7i2nA0.js +1 -0
- dw/server/ui/assets/cpp-C7h46wYY.js +1 -0
- dw/server/ui/assets/csharp-BKxtCVv1.js +1 -0
- dw/server/ui/assets/csp-bTuwJoIa.js +1 -0
- dw/server/ui/assets/css-DIMkf-bt.js +3 -0
- dw/server/ui/assets/css.worker-B3ciXF_0.js +93 -0
- dw/server/ui/assets/cssMode-CEh6hWi2.js +1 -0
- dw/server/ui/assets/cypher-CVaqCwHa.js +1 -0
- dw/server/ui/assets/dart-onAF5SnQ.js +1 -0
- dw/server/ui/assets/dockerfile-DZFCIeNp.js +1 -0
- dw/server/ui/assets/ecl-D05T4iGw.js +1 -0
- dw/server/ui/assets/editor-jjEx9u7D.css +1 -0
- dw/server/ui/assets/editor.api-CExg3_mM.js +847 -0
- dw/server/ui/assets/editor.worker-q-txB4vs.js +30 -0
- dw/server/ui/assets/elixir-6RTg0lbw.js +1 -0
- dw/server/ui/assets/flow9-C5_-GSwl.js +1 -0
- dw/server/ui/assets/freemarker2-DH6orYh2.js +3 -0
- dw/server/ui/assets/fsharp-C8Ef5oNN.js +1 -0
- dw/server/ui/assets/go-C-y9NEjX.js +1 -0
- dw/server/ui/assets/graphql-fmXr3nnJ.js +1 -0
- dw/server/ui/assets/handlebars-CbrMVW4Q.js +1 -0
- dw/server/ui/assets/hcl-CpzslTdj.js +1 -0
- dw/server/ui/assets/html-YDNPZw2M.js +1 -0
- dw/server/ui/assets/html.worker-C93Ht9o9.js +506 -0
- dw/server/ui/assets/htmlMode-B_zSGWO2.js +1 -0
- dw/server/ui/assets/index-B7-VcYS-.css +1 -0
- dw/server/ui/assets/index-D_EiPU3b.js +13 -0
- dw/server/ui/assets/ini-sBoK_t0W.js +1 -0
- dw/server/ui/assets/java-BEtHBSE6.js +1 -0
- dw/server/ui/assets/javascript-dYuBvioq.js +1 -0
- dw/server/ui/assets/json.worker-B2V3pomh.js +62 -0
- dw/server/ui/assets/jsonMode-CUqLM39V.js +7 -0
- dw/server/ui/assets/julia-Bri6UV-V.js +1 -0
- dw/server/ui/assets/kotlin-BOotOW0E.js +1 -0
- dw/server/ui/assets/less-B9JPFI3C.js +2 -0
- dw/server/ui/assets/lexon-CfSJPG6W.js +1 -0
- dw/server/ui/assets/liquid-D6vxBzMv.js +1 -0
- dw/server/ui/assets/lspLanguageFeatures-1WJ2palX.js +4 -0
- dw/server/ui/assets/lua-CsQS60Ue.js +1 -0
- dw/server/ui/assets/m3-D-oSqn_W.js +1 -0
- dw/server/ui/assets/markdown-Cimd5fb3.js +1 -0
- dw/server/ui/assets/mdx-SHQb6vmD.js +1 -0
- dw/server/ui/assets/mips-CIPQ_RoX.js +1 -0
- dw/server/ui/assets/monaco--ixms01u.css +1 -0
- dw/server/ui/assets/monaco-CP-s5rcP.js +56 -0
- dw/server/ui/assets/msdax-DauUninz.js +1 -0
- dw/server/ui/assets/mysql-SOo6toE5.js +1 -0
- dw/server/ui/assets/objective-c-FvmIjYaQ.js +1 -0
- dw/server/ui/assets/pascal-DrH0SRf2.js +1 -0
- dw/server/ui/assets/pascaligo-D-ptJ9y-.js +1 -0
- dw/server/ui/assets/perl-oz_6vUea.js +1 -0
- dw/server/ui/assets/pgsql-DTj74zXo.js +1 -0
- dw/server/ui/assets/php-nr791fC2.js +1 -0
- dw/server/ui/assets/pla-CopQ2nXW.js +1 -0
- dw/server/ui/assets/postiats-43DmfD33.js +1 -0
- dw/server/ui/assets/powerquery-D3hlyOfw.js +1 -0
- dw/server/ui/assets/powershell-DmHpPYUd.js +1 -0
- dw/server/ui/assets/protobuf-C531GsRP.js +2 -0
- dw/server/ui/assets/pug-Z5eAx3Zn.js +1 -0
- dw/server/ui/assets/python-x0_EGHq9.js +1 -0
- dw/server/ui/assets/qsharp-DkqhCAOL.js +1 -0
- dw/server/ui/assets/r-BwWrilGY.js +1 -0
- dw/server/ui/assets/razor-BZC4LQDP.js +1 -0
- dw/server/ui/assets/redis-ClamHrr6.js +1 -0
- dw/server/ui/assets/redshift-DT7zqm-g.js +1 -0
- dw/server/ui/assets/restructuredtext-BYgofb2h.js +1 -0
- dw/server/ui/assets/ruby-DezsRK8O.js +1 -0
- dw/server/ui/assets/rust-DdL9SqIa.js +1 -0
- dw/server/ui/assets/sb-CcwsVR0C.js +1 -0
- dw/server/ui/assets/scala-DHpiXF5c.js +1 -0
- dw/server/ui/assets/scheme-BeGwcela.js +1 -0
- dw/server/ui/assets/scss-gp-XZpBa.js +3 -0
- dw/server/ui/assets/shell-CC2rA5mh.js +1 -0
- dw/server/ui/assets/solidity-BEEn4gHE.js +1 -0
- dw/server/ui/assets/sophia-CRfGWb83.js +1 -0
- dw/server/ui/assets/sparql-D_Lu-MrJ.js +1 -0
- dw/server/ui/assets/sql-NEE52Syq.js +1 -0
- dw/server/ui/assets/st-DbInun42.js +1 -0
- dw/server/ui/assets/swift-Bxkupp3x.js +1 -0
- dw/server/ui/assets/systemverilog-Bz4Y3fRF.js +1 -0
- dw/server/ui/assets/tcl-DISqw1ZD.js +1 -0
- dw/server/ui/assets/ts.worker-D7T1-Ig5.js +67738 -0
- dw/server/ui/assets/tsMode-BTfA6SbD.js +11 -0
- dw/server/ui/assets/twig-De2hgUGE.js +1 -0
- dw/server/ui/assets/typescript-CWA4MsNk.js +1 -0
- dw/server/ui/assets/typespec-B8J7ngcE.js +1 -0
- dw/server/ui/assets/vb-DV3o63ZY.js +1 -0
- dw/server/ui/assets/wgsl-DpFanUEy.js +298 -0
- dw/server/ui/assets/workers-CWU0uvj5.js +1 -0
- dw/server/ui/assets/xml-KmfTm3rg.js +1 -0
- dw/server/ui/assets/yaml-nFO_dDS6.js +1 -0
- dw/server/ui/index.html +17 -0
- dw/settings.py +77 -0
- dw/step.py +132 -0
- dw/tasks/audio_utils.py +266 -0
- dw/tasks/background_remover.py +43 -0
- dw/tasks/borders.py +113 -0
- dw/tasks/concat_videos.py +80 -0
- dw/tasks/depth_estimator.py +54 -0
- dw/tasks/diffusion_upscale.py +109 -0
- dw/tasks/format_messages.py +24 -0
- dw/tasks/gather.py +139 -0
- dw/tasks/image_to_text.py +43 -0
- dw/tasks/image_utils.py +661 -0
- dw/tasks/interpolate_frames.py +227 -0
- dw/tasks/model_cache.py +39 -0
- dw/tasks/pair_audio.py +58 -0
- dw/tasks/qr_code.py +19 -0
- dw/tasks/restore_faces.py +175 -0
- dw/tasks/rife_model.py +192 -0
- dw/tasks/segment.py +121 -0
- dw/tasks/task.py +474 -0
- dw/tasks/tensor_image.py +57 -0
- dw/tasks/text_generation.py +168 -0
- dw/tasks/text_sections.py +80 -0
- dw/tasks/upscale.py +203 -0
- dw/tasks/video_utils.py +154 -0
- dw/tasks/zoe_depth.py +71 -0
- dw/teacache.py +376 -0
- dw/teacache_models.json +99 -0
- dw/test.py +29 -0
- dw/type_helpers.py +68 -0
- dw/validate.py +43 -0
- dw/variables.py +153 -0
- dw/worker.py +517 -0
- dw/workflow.py +553 -0
- dw/workflow_schema.json +1157 -0
- dw/workflows/augment_prompt.json +65 -0
- dw/workflows/describe_image.json +58 -0
- dw/workflows/h3_context_ir.json +57 -0
- dw/workflows/test.json +31 -0
dw/events.py
ADDED
|
@@ -0,0 +1,78 @@
|
|
|
1
|
+
"""Progress reporting and cooperative cancellation for workflow runs.
|
|
2
|
+
|
|
3
|
+
A RunContext travels ambiently (contextvars) through a run: Workflow.run
|
|
4
|
+
activates it, and Step/Pipeline reach it with get_context() rather than
|
|
5
|
+
threading a parameter through every action signature. Callers that never
|
|
6
|
+
pass a context get a no-op one, so the CLI path pays nothing.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
import logging
|
|
10
|
+
import threading
|
|
11
|
+
import contextvars
|
|
12
|
+
|
|
13
|
+
logger = logging.getLogger("dw")
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class WorkflowCancelled(Exception):
|
|
17
|
+
"""Raised inside a run when its RunContext has been cancelled."""
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class RunContext:
|
|
21
|
+
"""Carries the event sink and the cancellation flag for one workflow run.
|
|
22
|
+
|
|
23
|
+
cancel() may be called from any thread (the worker's command watcher);
|
|
24
|
+
everything else runs on the thread executing the workflow.
|
|
25
|
+
"""
|
|
26
|
+
|
|
27
|
+
def __init__(self, on_event=None):
|
|
28
|
+
self._on_event = on_event
|
|
29
|
+
self._cancel = threading.Event()
|
|
30
|
+
# Pipeline cache keys this run resolved - the worker evicts entries a
|
|
31
|
+
# run no longer touches, so an edited workflow drops stale models
|
|
32
|
+
self.touched_pipelines = set()
|
|
33
|
+
|
|
34
|
+
def emit(self, event_type, **data):
|
|
35
|
+
if self._on_event is None:
|
|
36
|
+
return
|
|
37
|
+
try:
|
|
38
|
+
self._on_event({"event": event_type, **data})
|
|
39
|
+
except Exception as e:
|
|
40
|
+
# A broken sink must not kill a run that is otherwise fine
|
|
41
|
+
logger.warning(f"Progress event sink failed on '{event_type}': {e}")
|
|
42
|
+
|
|
43
|
+
def cancel(self):
|
|
44
|
+
self._cancel.set()
|
|
45
|
+
|
|
46
|
+
@property
|
|
47
|
+
def cancelled(self):
|
|
48
|
+
return self._cancel.is_set()
|
|
49
|
+
|
|
50
|
+
def check_cancelled(self):
|
|
51
|
+
if self._cancel.is_set():
|
|
52
|
+
raise WorkflowCancelled("Workflow run was cancelled")
|
|
53
|
+
|
|
54
|
+
def touch_pipeline(self, cache_key):
|
|
55
|
+
self.touched_pipelines.add(cache_key)
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
_active_context = contextvars.ContextVar("dw_run_context", default=None)
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def get_context():
|
|
62
|
+
"""The active run's context, or a no-op one outside any run."""
|
|
63
|
+
context = _active_context.get()
|
|
64
|
+
return context if context is not None else RunContext()
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def current_context():
|
|
68
|
+
"""The active run's context, or None outside any run."""
|
|
69
|
+
return _active_context.get()
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def activate_context(context):
|
|
73
|
+
"""Make a context the active one; returns a token for deactivate_context."""
|
|
74
|
+
return _active_context.set(context)
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def deactivate_context(token):
|
|
78
|
+
_active_context.reset(token)
|
dw/hub_cache.py
ADDED
|
@@ -0,0 +1,289 @@
|
|
|
1
|
+
"""Inventory and deletion for the Hugging Face hub cache.
|
|
2
|
+
|
|
3
|
+
Every from_pretrained download lands in the hub cache and stays there until
|
|
4
|
+
something deletes it - a few video models is a few hundred GB. This wraps
|
|
5
|
+
huggingface_hub's own cache scanner so the server can show what is on disk
|
|
6
|
+
and free it, without inventing any path handling of its own: deletion goes
|
|
7
|
+
through scan_cache_dir's delete_revisions strategy, which only ever removes
|
|
8
|
+
revisions it found inside the cache directory.
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
import logging
|
|
12
|
+
import shutil
|
|
13
|
+
import threading
|
|
14
|
+
import time
|
|
15
|
+
import uuid
|
|
16
|
+
|
|
17
|
+
from huggingface_hub import constants, scan_cache_dir
|
|
18
|
+
from huggingface_hub.utils import CacheNotFound
|
|
19
|
+
|
|
20
|
+
logger = logging.getLogger("dw")
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def _resolved_cache_dir(cache_dir):
|
|
24
|
+
return str(cache_dir) if cache_dir else constants.HF_HUB_CACHE
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def scan_models(cache_dir=None):
|
|
28
|
+
"""The cache's contents as plain data: repos sorted largest-first,
|
|
29
|
+
with per-revision detail, plus disk totals for the volume it lives on."""
|
|
30
|
+
resolved = _resolved_cache_dir(cache_dir)
|
|
31
|
+
try:
|
|
32
|
+
scan = scan_cache_dir(resolved)
|
|
33
|
+
except CacheNotFound:
|
|
34
|
+
# No cache directory yet - nothing downloaded is a state, not an error
|
|
35
|
+
return {
|
|
36
|
+
"cache_dir": resolved,
|
|
37
|
+
"size_on_disk": 0,
|
|
38
|
+
"repos": [],
|
|
39
|
+
"warnings": [],
|
|
40
|
+
"disk_free": None,
|
|
41
|
+
"disk_total": None,
|
|
42
|
+
}
|
|
43
|
+
|
|
44
|
+
repos = []
|
|
45
|
+
for repo in scan.repos:
|
|
46
|
+
revisions = sorted(
|
|
47
|
+
repo.revisions, key=lambda rev: rev.last_modified or 0, reverse=True
|
|
48
|
+
)
|
|
49
|
+
repos.append(
|
|
50
|
+
{
|
|
51
|
+
"repo_id": repo.repo_id,
|
|
52
|
+
"repo_type": repo.repo_type,
|
|
53
|
+
"size_on_disk": repo.size_on_disk,
|
|
54
|
+
"nb_files": repo.nb_files,
|
|
55
|
+
"last_accessed": repo.last_accessed,
|
|
56
|
+
"last_modified": repo.last_modified,
|
|
57
|
+
"revisions": [
|
|
58
|
+
{
|
|
59
|
+
"commit_hash": revision.commit_hash,
|
|
60
|
+
"size_on_disk": revision.size_on_disk,
|
|
61
|
+
"refs": sorted(revision.refs),
|
|
62
|
+
"last_modified": revision.last_modified,
|
|
63
|
+
}
|
|
64
|
+
for revision in revisions
|
|
65
|
+
],
|
|
66
|
+
}
|
|
67
|
+
)
|
|
68
|
+
repos.sort(key=lambda entry: entry["size_on_disk"], reverse=True)
|
|
69
|
+
|
|
70
|
+
usage = shutil.disk_usage(resolved)
|
|
71
|
+
return {
|
|
72
|
+
"cache_dir": resolved,
|
|
73
|
+
"size_on_disk": scan.size_on_disk,
|
|
74
|
+
"repos": repos,
|
|
75
|
+
"warnings": [str(warning) for warning in scan.warnings],
|
|
76
|
+
"disk_free": usage.free,
|
|
77
|
+
"disk_total": usage.total,
|
|
78
|
+
}
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def delete_model(repo_id, cache_dir=None):
|
|
82
|
+
"""Delete every cached revision of repo_id. Returns the bytes freed.
|
|
83
|
+
|
|
84
|
+
Raises ValueError when the repo is not in the cache - the caller typed
|
|
85
|
+
or raced something, and nothing was deleted.
|
|
86
|
+
"""
|
|
87
|
+
scan = scan_cache_dir(_resolved_cache_dir(cache_dir))
|
|
88
|
+
repo = next((r for r in scan.repos if r.repo_id == repo_id), None)
|
|
89
|
+
if repo is None:
|
|
90
|
+
raise ValueError(f"'{repo_id}' is not in the hub cache")
|
|
91
|
+
|
|
92
|
+
strategy = scan.delete_revisions(
|
|
93
|
+
*[revision.commit_hash for revision in repo.revisions]
|
|
94
|
+
)
|
|
95
|
+
freed = strategy.expected_freed_size
|
|
96
|
+
strategy.execute()
|
|
97
|
+
return freed
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
# ------------------------------------------------------------------ downloads
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
class DownloadCancelled(Exception):
|
|
104
|
+
"""Raised inside a download's progress callback to abort it."""
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
class DownloadManager:
|
|
108
|
+
"""Background snapshot downloads into the hub cache, with progress.
|
|
109
|
+
|
|
110
|
+
One thread per download; progress is fed by a tqdm-compatible tracker
|
|
111
|
+
that snapshot_download instantiates per file, so the counters aggregate
|
|
112
|
+
across the file pool. Cancellation raises out of the next progress tick;
|
|
113
|
+
huggingface_hub's partial files remain resumable, so a cancelled or
|
|
114
|
+
failed download picks up where it stopped when retried.
|
|
115
|
+
"""
|
|
116
|
+
|
|
117
|
+
KEEP_FINISHED = 20
|
|
118
|
+
|
|
119
|
+
def __init__(self, download_fn=None, info_fn=None):
|
|
120
|
+
# Injectable for tests - the defaults reach the network
|
|
121
|
+
if download_fn is None or info_fn is None:
|
|
122
|
+
from huggingface_hub import HfApi, snapshot_download
|
|
123
|
+
|
|
124
|
+
download_fn = download_fn or snapshot_download
|
|
125
|
+
info_fn = info_fn or (
|
|
126
|
+
lambda repo_id: HfApi().repo_info(repo_id, files_metadata=True)
|
|
127
|
+
)
|
|
128
|
+
self._download_fn = download_fn
|
|
129
|
+
self._info_fn = info_fn
|
|
130
|
+
self._lock = threading.Lock()
|
|
131
|
+
self._downloads = {}
|
|
132
|
+
|
|
133
|
+
def start(self, repo_id):
|
|
134
|
+
"""Begin downloading repo_id; returns the download's status dict.
|
|
135
|
+
Raises ValueError for an invalid repo id or one already in flight."""
|
|
136
|
+
from huggingface_hub.utils import HFValidationError, validate_repo_id
|
|
137
|
+
|
|
138
|
+
try:
|
|
139
|
+
validate_repo_id(repo_id)
|
|
140
|
+
except HFValidationError as e:
|
|
141
|
+
raise ValueError(str(e))
|
|
142
|
+
|
|
143
|
+
with self._lock:
|
|
144
|
+
for entry in self._downloads.values():
|
|
145
|
+
if entry["repo_id"] == repo_id and entry["status"] == "downloading":
|
|
146
|
+
raise ValueError(f"'{repo_id}' is already downloading")
|
|
147
|
+
entry = {
|
|
148
|
+
"id": uuid.uuid4().hex[:12],
|
|
149
|
+
"repo_id": repo_id,
|
|
150
|
+
"status": "downloading",
|
|
151
|
+
"downloaded": 0,
|
|
152
|
+
"total": None,
|
|
153
|
+
"error": None,
|
|
154
|
+
"started_at": time.time(),
|
|
155
|
+
"finished_at": None,
|
|
156
|
+
}
|
|
157
|
+
self._downloads[entry["id"]] = entry
|
|
158
|
+
self._prune()
|
|
159
|
+
|
|
160
|
+
cancel_event = threading.Event()
|
|
161
|
+
entry["_cancel"] = cancel_event
|
|
162
|
+
thread = threading.Thread(
|
|
163
|
+
target=self._run, args=(entry, cancel_event), daemon=True
|
|
164
|
+
)
|
|
165
|
+
thread.start()
|
|
166
|
+
return self.status(entry["id"])
|
|
167
|
+
|
|
168
|
+
def _run(self, entry, cancel_event):
|
|
169
|
+
try:
|
|
170
|
+
try:
|
|
171
|
+
info = self._info_fn(entry["repo_id"])
|
|
172
|
+
total = sum(sibling.size or 0 for sibling in (info.siblings or []))
|
|
173
|
+
with self._lock:
|
|
174
|
+
entry["total"] = total or None
|
|
175
|
+
except Exception as e:
|
|
176
|
+
# Size is cosmetic; the download itself decides success
|
|
177
|
+
logger.debug(f"No size metadata for {entry['repo_id']}: {e}")
|
|
178
|
+
|
|
179
|
+
self._download_fn(
|
|
180
|
+
entry["repo_id"], tqdm_class=_tracker_class(self, entry, cancel_event)
|
|
181
|
+
)
|
|
182
|
+
self._finish(entry, "completed")
|
|
183
|
+
except DownloadCancelled:
|
|
184
|
+
self._finish(entry, "cancelled")
|
|
185
|
+
except Exception as e:
|
|
186
|
+
self._finish(entry, "failed", str(e))
|
|
187
|
+
|
|
188
|
+
def _finish(self, entry, status, error=None):
|
|
189
|
+
with self._lock:
|
|
190
|
+
entry["status"] = status
|
|
191
|
+
entry["error"] = error
|
|
192
|
+
entry["finished_at"] = time.time()
|
|
193
|
+
|
|
194
|
+
def _add_progress(self, entry, n):
|
|
195
|
+
with self._lock:
|
|
196
|
+
entry["downloaded"] += n
|
|
197
|
+
|
|
198
|
+
def cancel(self, download_id):
|
|
199
|
+
"""Request cancellation; returns the status dict or None if unknown.
|
|
200
|
+
Takes effect at the download's next progress tick."""
|
|
201
|
+
with self._lock:
|
|
202
|
+
entry = self._downloads.get(download_id)
|
|
203
|
+
if entry is None:
|
|
204
|
+
return None
|
|
205
|
+
if entry["status"] == "downloading":
|
|
206
|
+
entry["_cancel"].set()
|
|
207
|
+
return self.status(download_id)
|
|
208
|
+
|
|
209
|
+
def status(self, download_id):
|
|
210
|
+
with self._lock:
|
|
211
|
+
entry = self._downloads.get(download_id)
|
|
212
|
+
return _public(entry) if entry else None
|
|
213
|
+
|
|
214
|
+
def status_list(self):
|
|
215
|
+
"""Every tracked download, newest first."""
|
|
216
|
+
with self._lock:
|
|
217
|
+
entries = sorted(
|
|
218
|
+
self._downloads.values(),
|
|
219
|
+
key=lambda e: e["started_at"],
|
|
220
|
+
reverse=True,
|
|
221
|
+
)
|
|
222
|
+
return [_public(entry) for entry in entries]
|
|
223
|
+
|
|
224
|
+
def is_active(self):
|
|
225
|
+
with self._lock:
|
|
226
|
+
return any(
|
|
227
|
+
entry["status"] == "downloading" for entry in self._downloads.values()
|
|
228
|
+
)
|
|
229
|
+
|
|
230
|
+
def _prune(self):
|
|
231
|
+
# Called with the lock held: drop the oldest finished entries
|
|
232
|
+
finished = sorted(
|
|
233
|
+
(e for e in self._downloads.values() if e["status"] != "downloading"),
|
|
234
|
+
key=lambda e: e["started_at"],
|
|
235
|
+
)
|
|
236
|
+
excess = len(finished) - self.KEEP_FINISHED
|
|
237
|
+
for entry in finished[: max(0, excess)]:
|
|
238
|
+
del self._downloads[entry["id"]]
|
|
239
|
+
|
|
240
|
+
|
|
241
|
+
def _public(entry):
|
|
242
|
+
return {key: value for key, value in entry.items() if not key.startswith("_")}
|
|
243
|
+
|
|
244
|
+
|
|
245
|
+
def _tracker_class(manager, entry, cancel_event):
|
|
246
|
+
"""A tqdm stand-in snapshot_download instantiates per file; every update
|
|
247
|
+
feeds the shared counters and honours cancellation."""
|
|
248
|
+
|
|
249
|
+
class Tracker:
|
|
250
|
+
def __init__(self, *args, **kwargs):
|
|
251
|
+
self.n = 0
|
|
252
|
+
self.total = kwargs.get("total")
|
|
253
|
+
|
|
254
|
+
def update(self, n=1):
|
|
255
|
+
if cancel_event.is_set():
|
|
256
|
+
raise DownloadCancelled()
|
|
257
|
+
if n:
|
|
258
|
+
self.n += n
|
|
259
|
+
manager._add_progress(entry, n)
|
|
260
|
+
return True
|
|
261
|
+
|
|
262
|
+
def close(self):
|
|
263
|
+
pass
|
|
264
|
+
|
|
265
|
+
def refresh(self):
|
|
266
|
+
pass
|
|
267
|
+
|
|
268
|
+
def set_description(self, *args, **kwargs):
|
|
269
|
+
pass
|
|
270
|
+
|
|
271
|
+
def set_postfix(self, *args, **kwargs):
|
|
272
|
+
pass
|
|
273
|
+
|
|
274
|
+
def __enter__(self):
|
|
275
|
+
return self
|
|
276
|
+
|
|
277
|
+
def __exit__(self, *exc):
|
|
278
|
+
self.close()
|
|
279
|
+
return False
|
|
280
|
+
|
|
281
|
+
@staticmethod
|
|
282
|
+
def get_lock():
|
|
283
|
+
return threading.RLock()
|
|
284
|
+
|
|
285
|
+
@staticmethod
|
|
286
|
+
def set_lock(lock):
|
|
287
|
+
pass
|
|
288
|
+
|
|
289
|
+
return Tracker
|