thunk 0.1.0__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.
thunk/__init__.py ADDED
@@ -0,0 +1,37 @@
1
+ """Generate automatic save and load methods for functions that manipulate arrays"""
2
+
3
+ from importlib.metadata import version
4
+
5
+ from ._errors import (
6
+ OutputCodecError,
7
+ SchemaMismatchError,
8
+ SerializerContractError,
9
+ SpecError,
10
+ StorageFormatError,
11
+ ThunkError,
12
+ ValueTypeError,
13
+ )
14
+ from ._markers import Data, DataSerializer, DataValidator, Skip, Static
15
+ from ._persistence import Extras, FunctionPersistence, Missing, fn
16
+
17
+ __version__ = version("thunk")
18
+
19
+ __all__ = [
20
+ "Data",
21
+ "DataSerializer",
22
+ "DataValidator",
23
+ "Extras",
24
+ "FunctionPersistence",
25
+ "Missing",
26
+ "OutputCodecError",
27
+ "SchemaMismatchError",
28
+ "SerializerContractError",
29
+ "Skip",
30
+ "SpecError",
31
+ "Static",
32
+ "StorageFormatError",
33
+ "ThunkError",
34
+ "ValueTypeError",
35
+ "__version__",
36
+ "fn",
37
+ ]
thunk/_atomic.py ADDED
@@ -0,0 +1,59 @@
1
+ """Atomic file publication via a temporary file in the destination directory."""
2
+
3
+ import contextlib
4
+ import os
5
+ import tempfile
6
+ from collections.abc import Callable, Sequence
7
+ from pathlib import Path
8
+
9
+ type Writer = Callable[[Path], None]
10
+
11
+
12
+ def _umask() -> int:
13
+ mask = os.umask(0)
14
+ os.umask(mask)
15
+ return mask
16
+
17
+
18
+ def _new_temp(dest: Path) -> Path:
19
+ fd, name = tempfile.mkstemp(dir=dest.parent, prefix=f".{dest.name}.", suffix=".tmp")
20
+ os.close(fd)
21
+ return Path(name)
22
+
23
+
24
+ def _publish(temp: Path, dest: Path) -> None:
25
+ try:
26
+ mode = dest.stat().st_mode & 0o777
27
+ except FileNotFoundError:
28
+ mode = 0o666 & ~_umask()
29
+ os.chmod(temp, mode)
30
+ os.replace(temp, dest)
31
+
32
+
33
+ def atomic_write(path: os.PathLike[str] | str, writer: Writer) -> None:
34
+ publish_all([(path, writer)])
35
+
36
+
37
+ def publish_all(items: Sequence[tuple[os.PathLike[str] | str, Writer]]) -> None:
38
+ """Encode every item to a temp file, then replace the destinations.
39
+
40
+ No destination is touched unless every writer succeeds. The replacements
41
+ themselves are sequential, so they are not one transaction.
42
+ """
43
+ dests = [Path(p) for p, _ in items]
44
+ resolved = [d.resolve() for d in dests]
45
+ if len(set(resolved)) != len(resolved):
46
+ raise ValueError("destination paths must be distinct")
47
+
48
+ temps: list[Path] = []
49
+ try:
50
+ for dest, (_, writer) in zip(dests, items, strict=True):
51
+ temp = _new_temp(dest)
52
+ temps.append(temp)
53
+ writer(temp)
54
+ for temp, dest in zip(temps, dests, strict=True):
55
+ _publish(temp, dest)
56
+ finally:
57
+ for temp in temps:
58
+ with contextlib.suppress(FileNotFoundError):
59
+ temp.unlink()
thunk/_errors.py ADDED
@@ -0,0 +1,29 @@
1
+ """Exception hierarchy for thunk."""
2
+
3
+
4
+ class ThunkError(Exception):
5
+ """Base class for every error raised deliberately by thunk."""
6
+
7
+
8
+ class SpecError(ThunkError, TypeError):
9
+ """A callable or annotation cannot be compiled into a persistence spec."""
10
+
11
+
12
+ class ValueTypeError(ThunkError, TypeError):
13
+ """A runtime value does not match its annotation-derived spec."""
14
+
15
+
16
+ class SchemaMismatchError(ThunkError, ValueError):
17
+ """Stored names, roles, or fingerprints disagree with the current callable."""
18
+
19
+
20
+ class StorageFormatError(ThunkError, ValueError):
21
+ """A file is not a valid thunk file or has an unsupported storage version."""
22
+
23
+
24
+ class SerializerContractError(ThunkError, ValueError):
25
+ """A custom serializer or validator violated its contract."""
26
+
27
+
28
+ class OutputCodecError(ThunkError, TypeError):
29
+ """No supported output codec can be derived from the return annotation."""
thunk/_fingerprint.py ADDED
@@ -0,0 +1,162 @@
1
+ """Schema fingerprints and content digests."""
2
+
3
+ import hashlib
4
+ import json
5
+ import struct
6
+ from collections.abc import Mapping
7
+ from typing import Any
8
+
9
+ import numpy as np
10
+
11
+ from . import _jax
12
+ from ._signature import Param
13
+ from ._spec import (
14
+ Array,
15
+ CustomData,
16
+ CustomStatic,
17
+ DataclassNode,
18
+ DictNode,
19
+ FixedTupleNode,
20
+ JaxArray,
21
+ ListNode,
22
+ LiteralNode,
23
+ Node,
24
+ NoneNode,
25
+ OptionalNode,
26
+ Scalar,
27
+ VarTupleNode,
28
+ )
29
+
30
+
31
+ def canonical_json(obj: Any) -> str:
32
+ return json.dumps(obj, sort_keys=True, separators=(",", ":"), allow_nan=False)
33
+
34
+
35
+ def _sha256(text: str) -> str:
36
+ return hashlib.sha256(text.encode()).hexdigest()
37
+
38
+
39
+ def fingerprint(name: str, role: str, node: Node) -> str:
40
+ return _sha256(
41
+ canonical_json({"name": name, "role": role, "spec": node.describe()})
42
+ )
43
+
44
+
45
+ def param_fingerprint(param: Param) -> str:
46
+ assert param.node is not None
47
+ return fingerprint(param.name, param.role, param.node)
48
+
49
+
50
+ def output_fingerprint(node: Node) -> str:
51
+ return fingerprint("output", "output", node)
52
+
53
+
54
+ class _Hasher:
55
+ """Length-framed, type-tagged feed into SHA-256."""
56
+
57
+ def __init__(self) -> None:
58
+ self.h = hashlib.sha256()
59
+
60
+ def tag(self, tag: bytes) -> None:
61
+ self.h.update(tag)
62
+
63
+ def blob(self, tag: bytes, data: bytes) -> None:
64
+ self.h.update(tag + struct.pack(">Q", len(data)) + data)
65
+
66
+ def scalar(self, value: Any) -> None:
67
+ match value:
68
+ case bool():
69
+ self.blob(b"b", b"1" if value else b"0")
70
+ case int():
71
+ self.blob(b"i", str(value).encode())
72
+ case float():
73
+ self.blob(b"f", struct.pack(">d", value))
74
+ case str():
75
+ self.blob(b"s", value.encode())
76
+ case None:
77
+ self.tag(b"N")
78
+ case _:
79
+ raise TypeError(f"cannot digest scalar {value!r}")
80
+
81
+ def array(self, value: np.ndarray) -> None:
82
+ arr = np.asarray(value, order="C")
83
+ dtype = arr.dtype.newbyteorder("<")
84
+ arr = arr.astype(dtype, copy=False)
85
+ self.blob(b"a", f"{dtype.str}|{arr.shape}".encode())
86
+ self.blob(b"d", arr.tobytes())
87
+
88
+ def nested(self, value: Any) -> None:
89
+ """Digest an arbitrary (already validated) serializer-output tree."""
90
+ if isinstance(value, Mapping):
91
+ self.blob(b"{", struct.pack(">Q", len(value)))
92
+ for key, item in value.items():
93
+ self.scalar(key)
94
+ self.nested(item)
95
+ elif isinstance(value, np.ndarray):
96
+ self.array(value)
97
+ else:
98
+ self.scalar(value)
99
+
100
+ def feed(self, node: Node, value: Any) -> None:
101
+ match node:
102
+ case Scalar() | LiteralNode():
103
+ self.scalar(value)
104
+ case NoneNode():
105
+ self.tag(b"N")
106
+ case Array():
107
+ self.array(value)
108
+ case JaxArray():
109
+ data, implementation = _jax.to_host(value, "content digest")
110
+ if implementation is None:
111
+ self.tag(b"J")
112
+ else:
113
+ self.blob(b"K", implementation.encode())
114
+ self.array(data)
115
+ case OptionalNode(inner=inner):
116
+ if value is None:
117
+ self.tag(b"N")
118
+ else:
119
+ self.tag(b"S")
120
+ self.feed(inner, value)
121
+ case ListNode(inner=inner) | VarTupleNode(inner=inner):
122
+ self.blob(b"[", struct.pack(">Q", len(value)))
123
+ for item in value:
124
+ self.feed(inner, item)
125
+ case FixedTupleNode(items=items):
126
+ self.tag(b"(")
127
+ for sub, item in zip(items, value, strict=True):
128
+ self.feed(sub, item)
129
+ case DictNode(value=inner):
130
+ self.blob(b"{", struct.pack(">Q", len(value)))
131
+ for key, item in value.items():
132
+ self.scalar(key)
133
+ self.feed(inner, item)
134
+ case DataclassNode(fields=fields):
135
+ self.tag(b"D")
136
+ for name, sub in fields:
137
+ self.feed(sub, getattr(value, name))
138
+ case CustomData(serializer=serializer):
139
+ self.tag(b"C")
140
+ self.nested(serializer.func(value))
141
+ case CustomStatic():
142
+ raise TypeError("custom static values are digested through JSON")
143
+ case _:
144
+ raise TypeError(f"cannot digest node {node!r}")
145
+
146
+ def hexdigest(self) -> str:
147
+ return self.h.hexdigest()
148
+
149
+
150
+ def group_digest(params: tuple[Param, ...], values: Mapping[str, Any]) -> str:
151
+ """Content digest of a group of ``Data`` parameter values."""
152
+ hasher = _Hasher()
153
+ for p in params:
154
+ assert p.node is not None
155
+ hasher.blob(b"p", p.name.encode())
156
+ hasher.feed(p.node, values[p.name])
157
+ return hasher.hexdigest()
158
+
159
+
160
+ def json_digest(encoded: Mapping[str, Any]) -> str:
161
+ """Digest of an opts group's JSON encoding (insertion order is significant)."""
162
+ return _sha256(json.dumps(encoded, separators=(",", ":"), allow_nan=False))