Voice / vlib /net.py
Wiself's picture
Update voice.py, vlib, tests, bundles
fcff745
Raw History Blame Contribute Delete
17.7 kB
"""HTTP, auth, retries, and ranged downloads."""
from pathlib import Path
import json
import os
import re
import sys
import time
import urllib.error
import urllib.request
from urllib.parse import urlparse
from vlib import ctx
from vlib.ui import _fail, _note, _progress_pref, _say, _transform
_HF_TOKEN_CACHE = {"tok": None, "ts": 0.0}
def _get_hf_token():
# cache 60s to avoid file thrash on hot path
now = time.time()
if now - _HF_TOKEN_CACHE["ts"] < 60 and _HF_TOKEN_CACHE["tok"] is not None:
return _HF_TOKEN_CACHE["tok"]
tok = None
for k in ("HF_TOKEN", "HF_HUB_TOKEN", "HUGGING_FACE_HUB_TOKEN"):
v = os.environ.get(k)
if v and v.strip().startswith("hf_"):
tok = v.strip()
break
if tok is None:
for p in (ctx.HF_TOKEN_PATH, Path.home() / ".cache" / "huggingface" / "token", Path.home() / ".huggingface" / "token"):
try:
if p.exists():
t = p.read_text().strip()
if t.startswith("hf_"):
tok = t
break
except Exception:
continue
_HF_TOKEN_CACHE["tok"] = tok
_HF_TOKEN_CACHE["ts"] = now
return tok
_HF_TRUSTED_HOSTS = frozenset({"huggingface.co", "cdn-lfs.huggingface.co", "cdn-lfs-us-1.huggingface.co"})
from functools import lru_cache
@lru_cache(maxsize=256)
def _is_hf_url(url):
try:
host = (urlparse(url).hostname or "").lower()
except Exception:
return False
if host in _HF_TRUSTED_HOSTS:
return True
return host == "huggingface.co" or host.endswith(".huggingface.co")
def _hf_auth_headers(url, extra=None):
"""Auth + UA. Never sends Authorization to non-HF hosts; never overwrites caller header for HF either."""
h = {"User-Agent": "voice/2.0"}
if extra:
# Strip caller Authorization if target is not HF - prevents token leak to arbitrary URL
if not _is_hf_url(url) and "Authorization" in extra:
extra = {k: v for k, v in extra.items() if k != "Authorization"}
h.update(extra)
if "Authorization" not in h:
tok = _get_hf_token()
if tok and _is_hf_url(url):
h["Authorization"] = f"Bearer {tok}"
# If caller supplied Authorization for non-HF, drop it
if "Authorization" in h and not _is_hf_url(url):
del h["Authorization"]
return h
class _PreserveAuthRedirectHandler(urllib.request.HTTPRedirectHandler):
def redirect_request(self, req, fp, code, msg, hdrs, newurl):
result = super().redirect_request(req, fp, code, msg, hdrs, newurl)
if result is None:
return None
if _is_hf_url(newurl):
if "Authorization" not in result.headers:
tok = _get_hf_token()
if tok:
result.add_header("Authorization", f"Bearer {tok}")
else:
# stdlib forwards the original headers across hosts: strip the token
# (bridge/CDN URLs are pre-signed, so it is never needed there either)
try:
result.remove_header("Authorization")
except Exception:
pass
return result
urllib.request.install_opener(urllib.request.build_opener(_PreserveAuthRedirectHandler()))
def _parse_retry_after(headers):
try:
v = headers.get("Retry-After") or headers.get("retry-after")
return int(str(v).split(",")[0].strip()) if v is not None else None
except Exception:
return None
def _bounded_retry_after(headers, fallback, cap=30):
"""Retry-After with a cap: a 429 must never look like a hang (e.g. Retry-After: 3600)."""
try:
ra = _parse_retry_after(headers)
wait = ra if ra is not None else fallback
except Exception:
wait = fallback
try:
wait = float(wait)
except Exception:
wait = float(fallback)
if wait < 0:
wait = float(fallback)
return min(wait, cap)
def _timeout(base):
try:
env = int(os.environ.get("HF_HUB_DOWNLOAD_TIMEOUT", "0"))
return max(base, env) if env > 0 else base
except Exception:
return base
def http_range(url, start, end, timeout=60):
for attempt in range(3):
hdrs = _hf_auth_headers(url, {"Range": f"bytes={start}-{end}"})
try:
with urllib.request.urlopen(urllib.request.Request(url, headers=hdrs), timeout=_timeout(timeout)) as resp:
if resp.status == 429 and attempt < 2:
wait = _bounded_retry_after(resp.headers, 1.5 * (2 ** attempt))
if wait >= 5:
_note(f" Rate limited (429) — waiting {wait:.0f}s before retry…")
time.sleep(wait)
continue
if resp.status not in (200, 206):
raise urllib.error.HTTPError(url, resp.status, f"status {resp.status}", resp.headers, None)
return resp.read(), resp.status, resp.headers.get("ETag")
except urllib.error.HTTPError as e:
if e.code == 429 and attempt < 2:
wait = _bounded_retry_after(getattr(e, "headers", None) or {}, 1.5 * (2 ** attempt))
if wait >= 5:
_note(f" Rate limited (429) — waiting {wait:.0f}s before retry…")
time.sleep(wait)
continue
raise
def http_get_json(url, timeout=30):
req = urllib.request.Request(url, headers=_hf_auth_headers(url))
try:
with urllib.request.urlopen(req, timeout=timeout) as resp:
return json.loads(resp.read())
except urllib.error.HTTPError as e:
try:
e.close() # else the response fp dangles until GC (ResourceWarning)
except Exception:
pass
raise
def _safe_id(model_id, shard_url=""):
mid = re.sub(r"[^a-z0-9._-]", "_", model_id.strip().lower().replace("..", "_").replace("/", "__").replace("\\", "_").replace(":", "_"))
mid = re.sub(r"_+", "_", mid).strip("_")[:120]
if shard_url:
shard = re.sub(r"[^a-z0-9._-]", "_", shard_url.split("/")[-1].replace(".safetensors", "").replace(".gguf", "")).lower()[:40]
return f"{mid}__{shard}"
return mid
def _acquire_parts_lock(part_dir):
"""Best-effort cross-process lock for .parts. Returns lock path or None."""
lock = part_dir / ".lock"
try:
# O_EXCL to avoid double acquisition
fd = os.open(str(lock), os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600)
try:
os.write(fd, str(os.getpid()).encode())
finally:
os.close(fd)
return lock
except FileExistsError:
try:
# stale lock (>10 min) - reclaim
if time.time() - lock.stat().st_mtime > 600:
lock.unlink(missing_ok=True)
fd = os.open(str(lock), os.O_WRONLY | os.O_CREAT | os.O_EXCL, 0o600)
os.close(fd)
return lock
except Exception:
pass
return None
except Exception:
return None
def download_range(url, abs_start, abs_end, part_dir, etag, label="", n_parts=None):
import threading
total = abs_end - abs_start + 1
part_dir.mkdir(parents=True, exist_ok=True)
try:
os.chmod(part_dir.parent, 0o700)
except Exception:
pass
meta_path = part_dir / "meta.json"
old_etag = None
if meta_path.exists():
try:
old_etag = json.loads(meta_path.read_text()).get("etag")
except Exception:
old_etag = None
# acquire lock BEFORE writing meta so racers can't interleave etag/url
_parts_lock = _acquire_parts_lock(part_dir)
try:
meta_path.write_text(json.dumps({"url": url, "abs_start": abs_start, "total": total, "etag": etag}, indent=2))
except Exception:
pass
try:
import shutil
free = shutil.disk_usage(part_dir).free
if free < int(total * 1.25):
_fail(f" ✗ Not enough disk space: need {total*1.25/1e9:.1f} GB, free {free/1e9:.1f} GB.")
except SystemExit:
raise
except Exception:
pass
if n_parts is None:
n_parts = 8 if total > 512 * 1024 * 1024 else (4 if total > 128 * 1024 * 1024 else 1)
chunk_total = total // n_parts
ranges = []
for i in range(n_parts):
s = i * chunk_total
e = (s + chunk_total - 1) if i < n_parts - 1 else total - 1
ranges.append((abs_start + s, abs_start + e, part_dir / f"chunk_{i:03d}.part", e - s + 1))
downloaded = 0
for _rs, _re, p, csz in ranges:
if p.exists():
try:
sz = p.stat().st_size
# If ETag changed (or one side unknown), invalidate stale chunk
if old_etag != etag and (old_etag is not None or etag is not None):
p.unlink(missing_ok=True)
continue
if sz > csz:
p.unlink(missing_ok=True)
continue
downloaded += sz
except Exception:
pass
t0 = time.time()
lock = threading.RLock() # reentrant: progress() also takes it, and runs inside locked regions
stop = threading.Event() # set on fatal error / interrupt: workers wind down, partials kept
errors = []
last = [0.0]
def progress(force=False):
now = time.time()
with lock:
cur_downloaded = downloaded
if not force and now - last[0] < 0.25 and cur_downloaded < total:
return
last[0] = now
elapsed = now - t0 or 1e-6
mbps = cur_downloaded / elapsed / 1e6
pct = cur_downloaded / total * 100 if total else 100
remain = total - cur_downloaded
eta = _fmt_eta(remain, mbps)
msg = f" → {label} {pct:3.0f}% · {cur_downloaded/1e6:.0f}/{total/1e6:.0f} MB · {mbps:.1f} MB/s · {eta} "
sys.stdout.write(_transform(msg) + ("\n" if _progress_pref() else "\r"))
sys.stdout.flush()
# show 0% immediately so user sees waiting + speed
progress(force=True)
def chunk(r_start, r_end, p, csz):
nonlocal downloaded
existing = p.stat().st_size if p.exists() else 0
if existing == csz:
return # already complete (counted in pre-scan); re-requesting would 416
eff = r_start + existing
for attempt in range(3):
if stop.is_set():
return
try:
hdrs = _hf_auth_headers(url, {"Range": f"bytes={eff}-{r_end}"})
if etag:
hdrs["If-Match"] = etag
with urllib.request.urlopen(urllib.request.Request(url, headers=hdrs), timeout=_timeout(120)) as resp:
if resp.status not in (200, 206):
raise urllib.error.HTTPError(url, resp.status, "bad status", None, None)
if resp.status == 200 and (r_end - r_start + 1) != total:
# multi-part chunk but server ignored Range: body is wrong data — fail fast
_fail(f" ✗ Server ignored the Range request for {label or url} (got 200, need 206). Resume unsupported — try again later.")
if existing and resp.status == 200:
# server ignored Range - we already counted existing in downloaded, correct it
with lock:
downloaded -= existing
p.unlink(missing_ok=True)
existing = 0
eff = r_start
mode = "ab" if existing else "wb"
with open(p, mode) as f:
while not stop.is_set():
b = resp.read(8 << 20)
if not b:
break
f.write(b)
with lock:
downloaded += len(b)
progress()
f.flush()
try:
os.fsync(f.fileno())
except Exception:
pass
final = p.stat().st_size if p.exists() else 0
if final != csz and attempt < 2:
existing = final
eff = r_start + final
time.sleep(0.5 * (attempt + 1))
continue
if final != csz:
raise OSError(f"{p.name} incomplete {final}/{csz}")
return
except urllib.error.HTTPError as e:
if attempt < 2:
wait = _bounded_retry_after(getattr(e, "headers", None) or {}, 0.7 * (2 ** attempt))
if wait >= 5:
_note(f" Rate limited ({e.code}) — waiting {wait:.0f}s before retry…")
time.sleep(wait)
if p.exists():
existing = p.stat().st_size
eff = r_start + existing
continue
raise
except (urllib.error.URLError, TimeoutError, OSError):
if attempt < 2:
time.sleep(0.7 * (2 ** attempt))
if p.exists():
existing = p.stat().st_size
eff = r_start + existing
continue
raise
raise OSError(f"{p.name} failed after 3 retries")
def worker(rs, re_, p, csz):
try:
chunk(rs, re_, p, csz)
except SystemExit as e:
errors.append(e) # _fail() already printed the reason
stop.set()
except Exception as e: # noqa: BLE001 — transported to main thread below
errors.append(e)
stop.set()
try:
if n_parts == 1:
rs, re_, p, csz = ranges[0]
chunk(rs, re_, p, csz)
else:
# daemon threads: Ctrl-C returns instantly, partial chunks kept for resume
threads = [threading.Thread(target=worker, args=(rs, re_, p, csz), daemon=True)
for rs, re_, p, csz in ranges]
for t in threads:
t.start()
while any(t.is_alive() for t in threads):
for t in threads:
t.join(timeout=0.25)
if stop.is_set():
break
if stop.is_set():
break
if errors:
raise errors[0]
except KeyboardInterrupt:
stop.set()
sys.stdout.write("\n")
_fail(f" ✗ Download cancelled ({downloaded/1e6:.0f}/{total/1e6:.0f} MB kept). Run again to resume.")
except (urllib.error.HTTPError, urllib.error.URLError, TimeoutError, OSError):
sys.stdout.write("\n")
_fail(f" ✗ Download interrupted ({downloaded/1e6:.0f}/{total/1e6:.0f} MB kept). Run again to resume.")
except Exception as e:
sys.stdout.write("\n")
if ctx.VERBOSE:
import traceback
_say(traceback.format_exc())
_fail(f" ✗ Download failed: {e}")
# ensure 100% with speed is shown before combining
progress(force=True)
for _rs, _re, p, csz in ranges:
try:
if p.stat().st_size != csz:
_fail(f" ✗ Incomplete download ({p.name}). Run again to resume.")
except Exception:
_fail(f" ✗ Missing chunk {p.name}. Run again to resume.")
combined = part_dir / "combined.part"
with open(combined, "wb") as out:
for _rs, _re, p, _csz in ranges:
with open(p, "rb") as inp:
while True:
b = inp.read(8 << 20)
if not b:
break
out.write(b)
out.flush()
try:
os.fsync(out.fileno())
except Exception:
pass
# clean up lock
try:
if '_parts_lock' in locals() and _parts_lock and _parts_lock.exists():
_parts_lock.unlink(missing_ok=True)
except Exception:
pass
sys.stdout.write("\n")
sys.stdout.flush()
return combined
_PROGRESS_BYTES = 16 * 1024 * 1024
def _fetch_bytes_progress(url, start, end, label):
"""Bytes for one HTTP range, with progress + resume when large.
Small ranges keep the single http_range read (no flicker). Large ranges
go through the multipart downloader into a stable part dir (resume across
runs), are read back, and cleaned up on success ONLY — parts stay on
failure so the next run resumes. Returns bytes; callers unchanged, and
peak RAM matches the single-read path (file is transient on disk)."""
total = end - start + 1
if total < _PROGRESS_BYTES:
data, _st, _e = http_range(url, start, end, timeout=300)
return data
part_dir = ctx.VOICES_DIR / ".parts" / _safe_id(url, f"{start}-{end}")
combined = download_range(url, start, end, part_dir, None, label=label)
try:
with open(combined, "rb") as f:
return f.read()
finally:
_cleanup_tmp(part_dir)
def _fmt_eta(remaining, mbps):
if mbps < 1e-6:
return "--:--"
sec = int(remaining / (mbps * 1e6))
return f"{sec}s" if sec < 60 else f"{sec//60}m {sec%60}s"
def _cleanup_tmp(tmp_dir):
try:
if tmp_dir.exists():
import shutil
shutil.rmtree(tmp_dir)
except Exception:
pass