Skip to content
File

Blob: scripts/flash_cache.py

python123 lines
1"""Private successful-flash snapshots for esptool's verified sector comparison."""
2from contextlib import contextmanager
3import fcntl
4import hashlib
5import json
6import os
7from pathlib import Path
8import shutil
9import tempfile
10 
11from project import ROOT
12 
13 
14def digest(path):
15 with Path(path).open("rb") as source:
16 return hashlib.file_digest(source, "sha256").hexdigest()
17 
18 
19def layout_digest(table):
20 # IDF reserves 0xc00 bytes for the table; a USB read may include sector padding.
21 return hashlib.sha256(table[:0xC00].ljust(0xC00, b"\xff")).hexdigest()
22 
23 
24class FlashCache:
25 def __init__(self, identity, layout, root=None):
26 self.identity = identity
27 self.layout = layout
28 root = Path(root) if root is not None else ROOT / "artifacts/flash-cache"
29 self.directory = root / hashlib.sha256(identity.encode()).hexdigest()[:24]
30 
31 @contextmanager
32 def locked(self):
33 self.directory.mkdir(parents=True, exist_ok=True, mode=0o700)
34 lock = os.open(self.directory / ".lock", os.O_RDWR | os.O_CREAT, 0o600)
35 try:
36 try:
37 fcntl.flock(lock, fcntl.LOCK_EX | fcntl.LOCK_NB)
38 except BlockingIOError:
39 raise ValueError("Another flash operation is already using this device cache") from None
40 yield self
41 finally:
42 os.close(lock)
43 
44 def record(self):
45 try:
46 value = json.loads((self.directory / "verified.json").read_text())
47 if (value.get("version") == 1 and value.get("identity") == self.identity
48 and value.get("layout") == self.layout and isinstance(value.get("images"), dict)):
49 return value["images"]
50 except (OSError, ValueError, AttributeError):
51 pass
52 return {}
53 
54 def bases(self, addresses):
55 record = self.record()
56 result = []
57 for address in addresses:
58 path = self.directory / f"{address:08x}.bin"
59 try:
60 valid = record.get(str(address)) == digest(path)
61 except OSError:
62 valid = False
63 result.append(path if valid else None)
64 return result
65 
66 def remember(self, images):
67 """Record the immutable input copies only after successful device verification."""
68 records = self.record()
69 for address, source in images:
70 destination = self.directory / f"{address:08x}.bin"
71 self._replace(destination, lambda output: _copy(source, output))
72 records[str(address)] = digest(destination)
73 value = {"version": 1, "identity": self.identity, "layout": self.layout, "images": records}
74 self._replace(self.directory / "verified.json", lambda output: output.write((json.dumps(value, indent=2) + "\n").encode()))
75 descriptor = os.open(self.directory, os.O_RDONLY)
76 try:
77 os.fsync(descriptor)
78 finally:
79 os.close(descriptor)
80 
81 def _replace(self, destination, write):
82 with tempfile.NamedTemporaryFile(dir=self.directory, delete=False) as temporary:
83 path = Path(temporary.name)
84 try:
85 write(temporary)
86 temporary.flush()
87 os.fsync(temporary.fileno())
88 path.replace(destination)
89 finally:
90 path.unlink(missing_ok=True)
91 
92 
93def _copy(source, output):
94 with Path(source).open("rb") as source_file:
95 shutil.copyfileobj(source_file, output)
96 
97 
98def flash_images(images, base_command, run, cache=None, full=False):
99 """Stage inputs so concurrent catalog preparation cannot change a running flash."""
100 with tempfile.TemporaryDirectory(prefix="radio-flash-") as directory:
101 staged = []
102 for address, source in sorted(images):
103 target = Path(directory) / f"{address:08x}.bin"
104 with target.open("xb") as output:
105 os.chmod(target, 0o600)
106 _copy(source, output)
107 staged.append((address, target))
108 bases = cache.bases([address for address, _ in staged]) if cache else []
109 options = []
110 if not full:
111 if any(bases):
112 options = ["--diff-with", *(str(path) if path else "skip" for path in bases)]
113 else:
114 options = ["--skip-flashed"]
115 # A following option terminates esptool's variadic --diff-with list.
116 command = base_command + ["write-flash", *options, "--flash-mode", "dout", "--flash-size", "32MB", "--flash-freq", "80m"]
117 for address, path in staged:
118 command.extend([hex(address), str(path)])
119 result = run(command)
120 if result.returncode == 0 and cache:
121 cache.remember(staged)
122 return result.returncode