mirror of
https://github.com/baketnk/frame-yap.git
synced 2026-10-06 01:00:04 +02:00
fix(installer): unwind owned model downloads on cancellation
This commit is contained in:
1 parent
4a2c524b17
commit
ebaae02b97
3 files changed
+140
-52
No files matched your search
+42
-26
@@ -32,6 +32,7 @@ from pathlib import Path
|
|||||||
import platform
|
import platform
|
||||||
import re
|
import re
|
||||||
import shlex
|
import shlex
|
||||||
|
import signal
|
||||||
import shutil
|
import shutil
|
||||||
import subprocess
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
@@ -535,6 +536,20 @@ def model_target(args, root):
|
|||||||
return root / "models" / args.backend
|
return root / "models" / args.backend
|
||||||
|
|
||||||
|
|
||||||
|
@contextlib.contextmanager
|
||||||
|
def model_download_cancellation():
|
||||||
|
# Install-model runs on the main thread. Raising unwinds the owned temp's
|
||||||
|
# finally block; unlike SIGKILL this does not leave a partial download.
|
||||||
|
def terminate(_signal, _frame):
|
||||||
|
raise ValueError("model installation cancelled")
|
||||||
|
|
||||||
|
previous = signal.signal(signal.SIGTERM, terminate)
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
signal.signal(signal.SIGTERM, previous)
|
||||||
|
|
||||||
|
|
||||||
def install_model(args, root):
|
def install_model(args, root):
|
||||||
module, backend, manifest_sha = installed_backend(root, args.backend)
|
module, backend, manifest_sha = installed_backend(root, args.backend)
|
||||||
check_expected_manifest(args, manifest_sha) # under .model.lock, before model mkdir/network
|
check_expected_manifest(args, manifest_sha) # under .model.lock, before model mkdir/network
|
||||||
@@ -559,32 +574,33 @@ def install_model(args, root):
|
|||||||
url = f"{backend.source}/resolve/{quote(backend.revision, safe='')}/{quote(item.path, safe='/')}"
|
url = f"{backend.source}/resolve/{quote(backend.revision, safe='')}/{quote(item.path, safe='/')}"
|
||||||
if args.json:
|
if args.json:
|
||||||
print(json.dumps({"ok": True, "event": "model_file", "file": item.path, "state": "downloading", "bytes": item.size}), flush=True)
|
print(json.dumps({"ok": True, "event": "model_file", "file": item.path, "state": "downloading", "bytes": item.size}), flush=True)
|
||||||
fd, temp = tempfile.mkstemp(prefix=".download-", dir=target.parent)
|
with model_download_cancellation():
|
||||||
try:
|
fd, temp = tempfile.mkstemp(prefix=".download-", dir=target.parent)
|
||||||
with os.fdopen(fd, "wb") as output:
|
try:
|
||||||
# The shared verifier is used for both pre-existing and downloaded files.
|
with os.fdopen(fd, "wb") as output:
|
||||||
request = urllib.request.Request(url, headers={"User-Agent": "frameyap-installer"})
|
# The shared verifier is used for both pre-existing and downloaded files.
|
||||||
with urllib.request.urlopen(request, timeout=60) as response:
|
request = urllib.request.Request(url, headers={"User-Agent": "frameyap-installer"})
|
||||||
if response.geturl().split(":", 1)[0] != "https":
|
with urllib.request.urlopen(request, timeout=60) as response:
|
||||||
fail("model URL redirected away from HTTPS")
|
if response.geturl().split(":", 1)[0] != "https":
|
||||||
total = 0
|
fail("model URL redirected away from HTTPS")
|
||||||
while chunk := response.read(1024 * 1024):
|
total = 0
|
||||||
total += len(chunk)
|
while chunk := response.read(1024 * 1024):
|
||||||
if total > item.size:
|
total += len(chunk)
|
||||||
fail(f"model download exceeded pinned size: {item.path}")
|
if total > item.size:
|
||||||
output.write(chunk)
|
fail(f"model download exceeded pinned size: {item.path}")
|
||||||
output.flush()
|
output.write(chunk)
|
||||||
os.fsync(output.fileno())
|
output.flush()
|
||||||
temporary = type(item)(Path(temp).name, item.size, item.sha256)
|
os.fsync(output.fileno())
|
||||||
if module.check_file(target.parent, temporary)[0] is not None:
|
temporary = type(item)(Path(temp).name, item.size, item.sha256)
|
||||||
fail(f"pinned SHA-256/size mismatch: {item.path}")
|
if module.check_file(target.parent, temporary)[0] is not None:
|
||||||
if target.exists() or target.is_symlink():
|
fail(f"pinned SHA-256/size mismatch: {item.path}")
|
||||||
fail(f"model destination changed during download: {target}")
|
if target.exists() or target.is_symlink():
|
||||||
os.replace(temp, target)
|
fail(f"model destination changed during download: {target}")
|
||||||
if args.json:
|
os.replace(temp, target)
|
||||||
print(json.dumps({"ok": True, "event": "model_file", "file": item.path, "state": "verified"}), flush=True)
|
if args.json:
|
||||||
finally:
|
print(json.dumps({"ok": True, "event": "model_file", "file": item.path, "state": "verified"}), flush=True)
|
||||||
Path(temp).unlink(missing_ok=True)
|
finally:
|
||||||
|
Path(temp).unlink(missing_ok=True)
|
||||||
state = module.check_model(backend, dest)
|
state = module.check_model(backend, dest)
|
||||||
if state["state"] != "installed_verified":
|
if state["state"] != "installed_verified":
|
||||||
fail(f"model did not verify: {state['reason']}")
|
fail(f"model did not verify: {state['reason']}")
|
||||||
|
|||||||
+42
-26
@@ -12,6 +12,7 @@ from pathlib import Path
|
|||||||
import platform
|
import platform
|
||||||
import re
|
import re
|
||||||
import shlex
|
import shlex
|
||||||
|
import signal
|
||||||
import shutil
|
import shutil
|
||||||
import subprocess
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
@@ -515,6 +516,20 @@ def model_target(args, root):
|
|||||||
return root / "models" / args.backend
|
return root / "models" / args.backend
|
||||||
|
|
||||||
|
|
||||||
|
@contextlib.contextmanager
|
||||||
|
def model_download_cancellation():
|
||||||
|
# Install-model runs on the main thread. Raising unwinds the owned temp's
|
||||||
|
# finally block; unlike SIGKILL this does not leave a partial download.
|
||||||
|
def terminate(_signal, _frame):
|
||||||
|
raise ValueError("model installation cancelled")
|
||||||
|
|
||||||
|
previous = signal.signal(signal.SIGTERM, terminate)
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
signal.signal(signal.SIGTERM, previous)
|
||||||
|
|
||||||
|
|
||||||
def install_model(args, root):
|
def install_model(args, root):
|
||||||
module, backend, manifest_sha = installed_backend(root, args.backend)
|
module, backend, manifest_sha = installed_backend(root, args.backend)
|
||||||
check_expected_manifest(args, manifest_sha) # under .model.lock, before model mkdir/network
|
check_expected_manifest(args, manifest_sha) # under .model.lock, before model mkdir/network
|
||||||
@@ -539,32 +554,33 @@ def install_model(args, root):
|
|||||||
url = f"{backend.source}/resolve/{quote(backend.revision, safe='')}/{quote(item.path, safe='/')}"
|
url = f"{backend.source}/resolve/{quote(backend.revision, safe='')}/{quote(item.path, safe='/')}"
|
||||||
if args.json:
|
if args.json:
|
||||||
print(json.dumps({"ok": True, "event": "model_file", "file": item.path, "state": "downloading", "bytes": item.size}), flush=True)
|
print(json.dumps({"ok": True, "event": "model_file", "file": item.path, "state": "downloading", "bytes": item.size}), flush=True)
|
||||||
fd, temp = tempfile.mkstemp(prefix=".download-", dir=target.parent)
|
with model_download_cancellation():
|
||||||
try:
|
fd, temp = tempfile.mkstemp(prefix=".download-", dir=target.parent)
|
||||||
with os.fdopen(fd, "wb") as output:
|
try:
|
||||||
# The shared verifier is used for both pre-existing and downloaded files.
|
with os.fdopen(fd, "wb") as output:
|
||||||
request = urllib.request.Request(url, headers={"User-Agent": "frameyap-installer"})
|
# The shared verifier is used for both pre-existing and downloaded files.
|
||||||
with urllib.request.urlopen(request, timeout=60) as response:
|
request = urllib.request.Request(url, headers={"User-Agent": "frameyap-installer"})
|
||||||
if response.geturl().split(":", 1)[0] != "https":
|
with urllib.request.urlopen(request, timeout=60) as response:
|
||||||
fail("model URL redirected away from HTTPS")
|
if response.geturl().split(":", 1)[0] != "https":
|
||||||
total = 0
|
fail("model URL redirected away from HTTPS")
|
||||||
while chunk := response.read(1024 * 1024):
|
total = 0
|
||||||
total += len(chunk)
|
while chunk := response.read(1024 * 1024):
|
||||||
if total > item.size:
|
total += len(chunk)
|
||||||
fail(f"model download exceeded pinned size: {item.path}")
|
if total > item.size:
|
||||||
output.write(chunk)
|
fail(f"model download exceeded pinned size: {item.path}")
|
||||||
output.flush()
|
output.write(chunk)
|
||||||
os.fsync(output.fileno())
|
output.flush()
|
||||||
temporary = type(item)(Path(temp).name, item.size, item.sha256)
|
os.fsync(output.fileno())
|
||||||
if module.check_file(target.parent, temporary)[0] is not None:
|
temporary = type(item)(Path(temp).name, item.size, item.sha256)
|
||||||
fail(f"pinned SHA-256/size mismatch: {item.path}")
|
if module.check_file(target.parent, temporary)[0] is not None:
|
||||||
if target.exists() or target.is_symlink():
|
fail(f"pinned SHA-256/size mismatch: {item.path}")
|
||||||
fail(f"model destination changed during download: {target}")
|
if target.exists() or target.is_symlink():
|
||||||
os.replace(temp, target)
|
fail(f"model destination changed during download: {target}")
|
||||||
if args.json:
|
os.replace(temp, target)
|
||||||
print(json.dumps({"ok": True, "event": "model_file", "file": item.path, "state": "verified"}), flush=True)
|
if args.json:
|
||||||
finally:
|
print(json.dumps({"ok": True, "event": "model_file", "file": item.path, "state": "verified"}), flush=True)
|
||||||
Path(temp).unlink(missing_ok=True)
|
finally:
|
||||||
|
Path(temp).unlink(missing_ok=True)
|
||||||
state = module.check_model(backend, dest)
|
state = module.check_model(backend, dest)
|
||||||
if state["state"] != "installed_verified":
|
if state["state"] != "installed_verified":
|
||||||
fail(f"model did not verify: {state['reason']}")
|
fail(f"model did not verify: {state['reason']}")
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ from unittest.mock import patch
|
|||||||
import shutil
|
import shutil
|
||||||
import pty
|
import pty
|
||||||
import select
|
import select
|
||||||
|
import signal
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
REPO = Path(__file__).resolve().parents[1]
|
REPO = Path(__file__).resolve().parents[1]
|
||||||
@@ -850,6 +851,61 @@ class InstallTests(unittest.TestCase):
|
|||||||
fetch.assert_not_called()
|
fetch.assert_not_called()
|
||||||
self.assertFalse((root / "models/other").exists())
|
self.assertFalse((root / "models/other").exists())
|
||||||
|
|
||||||
|
def test_model_sigterm_removes_owned_partial_download(self):
|
||||||
|
# The child uses a mocked blocking response. No network/download occurs.
|
||||||
|
shutil.copyfile(REPO / "python/frameyap/model_files.py",
|
||||||
|
self.stage / "python/frameyap/model_files.py")
|
||||||
|
manifest = json.loads((REPO / "assets/backends/redux.json").read_text())
|
||||||
|
manifest["model"]["files"] = [{"path": "weights.bin", "size": 9,
|
||||||
|
"sha256": hashlib.sha256(b"ninebytes").hexdigest()}]
|
||||||
|
manifests = self.stage / "assets/backends"
|
||||||
|
manifests.mkdir()
|
||||||
|
(manifests / "redux.json").write_text(json.dumps(manifest))
|
||||||
|
archive, digest = self.package("0.1.202609241530")
|
||||||
|
self.install("0.1.202609241530", archive, digest, "--without-model")
|
||||||
|
dest = self.base / "cancelled-model"
|
||||||
|
script = '''import importlib.util, signal, sys
|
||||||
|
from unittest.mock import patch
|
||||||
|
spec = importlib.util.spec_from_file_location("installer", sys.argv[1])
|
||||||
|
module = importlib.util.module_from_spec(spec)
|
||||||
|
spec.loader.exec_module(module)
|
||||||
|
class Response:
|
||||||
|
def __enter__(self): return self
|
||||||
|
def __exit__(self, *args): pass
|
||||||
|
def geturl(self): return "https://example.org/model"
|
||||||
|
def read(self, size):
|
||||||
|
print("fixture-read-started", flush=True)
|
||||||
|
signal.pause() # parent delivers SIGTERM during the owned download
|
||||||
|
with patch.object(module, "check_host"), patch.object(module.urllib.request, "urlopen", return_value=Response()):
|
||||||
|
sys.exit(module.cli(["--install-model", "--backend", "redux", "--model-dir", sys.argv[2], "--yes", "--json"]))
|
||||||
|
'''
|
||||||
|
child = subprocess.Popen([sys.executable, "-c", script,
|
||||||
|
str(REPO / "scripts/install_payload.py"), str(dest)],
|
||||||
|
stdout=subprocess.PIPE, stderr=subprocess.PIPE,
|
||||||
|
env={**os.environ, "PYTHONDONTWRITEBYTECODE": "1"})
|
||||||
|
try:
|
||||||
|
output = b""
|
||||||
|
for _ in range(8):
|
||||||
|
readable, _, _ = select.select([child.stdout], [], [], 8)
|
||||||
|
self.assertTrue(readable, "mocked download did not reach response.read")
|
||||||
|
output += os.read(child.stdout.fileno(), 4096)
|
||||||
|
if b"fixture-read-started\n" in output:
|
||||||
|
break
|
||||||
|
self.assertIn(b"fixture-read-started\n", output)
|
||||||
|
self.assertEqual(len(list(dest.glob(".download-*"))), 1)
|
||||||
|
os.kill(child.pid, signal.SIGTERM) # only our fixture process
|
||||||
|
rest, error = child.communicate(timeout=5)
|
||||||
|
self.assertEqual(child.returncode, 1, error)
|
||||||
|
event = json.loads(rest.strip())
|
||||||
|
self.assertFalse(event["ok"])
|
||||||
|
self.assertIn("cancelled", event["message"])
|
||||||
|
self.assertEqual(list(dest.glob(".download-*")), [])
|
||||||
|
self.assertFalse((dest / "weights.bin").exists())
|
||||||
|
finally:
|
||||||
|
if child.poll() is None:
|
||||||
|
child.kill()
|
||||||
|
child.wait(timeout=5)
|
||||||
|
|
||||||
def test_model_install_manifest_verification_and_explicit_consent(self):
|
def test_model_install_manifest_verification_and_explicit_consent(self):
|
||||||
# Tiny pinned files via the shared verifier; urllib is mocked, no network.
|
# Tiny pinned files via the shared verifier; urllib is mocked, no network.
|
||||||
source_module = REPO / "python/frameyap/model_files.py"
|
source_module = REPO / "python/frameyap/model_files.py"
|
||||||
|
|||||||
Reference in new issue
Block a user