diff --git a/CHANGELOG.md b/CHANGELOG.md index d0130dd..78a7901 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -3,6 +3,20 @@ All notable changes to `citrate-labs-sdk` are documented here. This project adheres to [Semantic Versioning](https://semver.org/). +## [0.6.3] - 2026-09-25 — Hardening + +### Changed + +- `deploy_model` serialises the transaction payload once, guards those exact + bytes (duplicate object keys are refused), and sends the same bytes. + `_send_transaction` accepts an already-serialised JSON payload. +- The key-share guard recognises byte values and their JSON renderings. + +### Tests + +- A differential property test runs the share guard against the strict share + parser and other known hex decoders over generated near-hex input. + ## [0.6.2] - 2026-09-25 — Pre-bounty audit remediation (SECURITY) > **Security advisory: upgrade from 0.6.0 (and any unreleased 0.6.1 build).** diff --git a/citrate_sdk/__init__.py b/citrate_sdk/__init__.py index 829b162..9da8e3f 100644 --- a/citrate_sdk/__init__.py +++ b/citrate_sdk/__init__.py @@ -48,7 +48,7 @@ StakingInfo, ) -__version__ = "0.6.2" +__version__ = "0.6.3" __author__ = "Citrate Team" __all__ = [ diff --git a/citrate_sdk/client.py b/citrate_sdk/client.py index 4f2175a..c22cec6 100644 --- a/citrate_sdk/client.py +++ b/citrate_sdk/client.py @@ -13,7 +13,11 @@ from ._generated import contract as _contract from ._url_security import enforce_transport_security -from .crypto import EncryptionConfig, KeyManager, assert_no_key_share_material +from .crypto import ( + EncryptionConfig, + KeyManager, + assert_payload_has_no_key_share_material, +) from .errors import CitrateError, ModelNotFoundError from .ipfs import upload_to_ipfs from .models import InferenceRequest, InferenceResult, ModelConfig, ModelDeployment @@ -230,16 +234,18 @@ def deploy_model( if encryption_metadata: tx_data["encryption_metadata"] = encryption_metadata - # PBA-L6b-003: this calldata is public. Refuse to send if anything in it, - # including caller-supplied metadata, carries a key-share field. - assert_no_key_share_material(tx_data) + # PBA-L6b-003: this calldata is public. Serialise once, guard the parsed + # payload (duplicate keys refused), and send that exact payload. json.dumps + # only emits JSON types, so the parsed payload is the complete view. + payload = json.dumps(tx_data) + assert_payload_has_no_key_share_material(payload) # Call model deployment precompile. SPY-B-007: the address is read from # the vendored canonical table (ModelDeploy = 0x..0100), NOT a hardcoded # `0x0100..0100` literal. The pre-fix literals were wrong in the high byte # (`0x01`-prefixed), so every deploy dispatched to an address with no # precompile entry and the state change never happened. - tx_hash = self._send_transaction(self._precompile("ModelDeploy"), tx_data) + tx_hash = self._send_transaction(self._precompile("ModelDeploy"), payload) # Wait for confirmation receipt = self._wait_for_receipt(tx_hash) @@ -465,11 +471,16 @@ def _eip155_chain_id(self) -> int: def _send_transaction( self, to_address: str, - data: dict[str, Any], + data: dict[str, Any] | str, value: int = 0, gas_limit: int = 500000 ) -> str: - """Send transaction to blockchain""" + """Send transaction to blockchain. + + ``data`` is either a dict (serialised here) or an already-serialised + JSON string, which is sent byte-for-byte (see ``deploy_model``). + """ + encoded = data if isinstance(data, str) else json.dumps(data) if not self.key_manager: raise CitrateError("Private key required for transactions") @@ -487,7 +498,7 @@ def _send_transaction( "gas": hex(gas_limit), "gasPrice": hex(20_000_000_000), # 20 gwei "nonce": hex(nonce), - "data": "0x" + json.dumps(data).encode().hex(), + "data": "0x" + encoded.encode().hex(), "chainId": self._eip155_chain_id(), } diff --git a/citrate_sdk/crypto.py b/citrate_sdk/crypto.py index f59d747..b029538 100644 --- a/citrate_sdk/crypto.py +++ b/citrate_sdk/crypto.py @@ -64,13 +64,70 @@ def _share_x(x: Any) -> bool: return isinstance(x, str) and x.isascii() and x.isdigit() and 1 <= int(x) <= 255 +def _is_byte_int(v: Any) -> bool: + return isinstance(v, int) and not isinstance(v, bool) and 0 <= v <= 255 + + +def _bytes_like_len(y: Any) -> int: + """Length of ``y`` if it is bytes or a JSON rendering of bytes (an integer + list, the ``{"type": "Buffer", "data": [...]}`` shape, or an object keyed + "0".."n-1" with byte values); otherwise 0.""" + if isinstance(y, (bytes, bytearray)): + return len(y) + if isinstance(y, (list, tuple)): + return len(y) if all(_is_byte_int(v) for v in y) else 0 + if isinstance(y, dict): + if set(y) == {"type", "data"} and y.get("type") == "Buffer": + return _bytes_like_len(list(y["data"])) if isinstance(y.get("data"), list) else 0 + n = len(y) + if n and all(isinstance(k, str) for k in y) and set(y) == {str(i) for i in range(n)}: + return n if all(_is_byte_int(y[str(i)]) for i in range(n)) else 0 + return 0 + + +class _DuplicateKeyError(ValueError): + pass + + +def _no_duplicate_keys(pairs: list[tuple[str, Any]]) -> dict[str, Any]: + out: dict[str, Any] = {} + for k, v in pairs: + if k in out: + raise _DuplicateKeyError(k) + out[k] = v + return out + + +def _loads_strict(text: str) -> Any: + """json.loads that rejects duplicate object keys (raises _DuplicateKeyError).""" + return json.loads(text, object_pairs_hook=_no_duplicate_keys) + + +_DUP_MSG = ( + "deploy_model: refusing to publish JSON with duplicate object keys; decoders disagree on " + "which value wins, so the content cannot be checked for key-share material (PBA-L6b-003)." +) + + +def assert_payload_has_no_key_share_material(payload: str) -> None: + """Guard a serialised JSON payload exactly as it will be sent: parse it + (refusing duplicate keys) and run :func:`assert_no_key_share_material`.""" + try: + decoded = _loads_strict(payload) + except _DuplicateKeyError: + raise CitrateError(_DUP_MSG) + except (ValueError, RecursionError): + raise CitrateError("deploy_model: payload is not valid JSON; refusing to publish it.") + assert_no_key_share_material(decoded) + + def _looks_like_share(d: dict[Any, Any]) -> bool: """A raw Shamir share ({x in 1..255, y of share length as hex or bytes}) or a holder-wrapped share record ({holder_public_key/holderPublicKey, envelope}). Short or coordinate-like values are not treated as shares.""" y = d.get("y") if "x" in d and _share_x(d["x"]): - if isinstance(y, (bytes, bytearray)) and len(y) >= _MIN_SHARE_BYTES: + if _bytes_like_len(y) >= _MIN_SHARE_BYTES: return True if isinstance(y, str) and _share_y_like(y): return True @@ -89,7 +146,9 @@ def assert_no_key_share_material(value: Any, _depth: int = 0) -> None: # Any string that parses as JSON is checked too (no size cap: a padded # blob must not slip through). try: - decoded = json.loads(value) + decoded = _loads_strict(value) + except _DuplicateKeyError: + raise CitrateError(_DUP_MSG) except (ValueError, RecursionError): return assert_no_key_share_material(decoded, _depth + 1) diff --git a/pyproject.toml b/pyproject.toml index 616f249..4f9c403 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -13,7 +13,7 @@ build-backend = "setuptools.build_meta" [project] name = "citrate-labs-sdk" -version = "0.6.2" +version = "0.6.3" description = "Python SDK for the Citrate distributed AI network (chain 40204). The canonical TypeScript SDK is @citratelabs/sdk on npm; this Python client is opt-in and may lag it." readme = "README.md" requires-python = ">=3.10" diff --git a/tests/fixtures/share_guard_vectors.json b/tests/fixtures/share_guard_vectors.json index cb52977..2a083c7 100644 --- a/tests/fixtures/share_guard_vectors.json +++ b/tests/fixtures/share_guard_vectors.json @@ -1,6 +1,6 @@ { - "_comment": "Shared share-guard test vectors. The SAME file is committed to citrate-sdk-js and citrate-sdk-python (tests/fixtures/share_guard_vectors.json); each SDK's guard must refuse every refuse:true entry and accept every refuse:false entry. Keep the two copies byte-identical.", - "version": 1, + "_comment": "Shared share-guard test vectors. The SAME file is committed to citrate-sdk-js and citrate-sdk-python (tests/fixtures/share_guard_vectors.json); each SDK's guard must refuse every refuse:true entry and accept every refuse:false entry. Keep the two copies byte-identical. raw_payloads are serialised payload texts for each SDK's payload guard (strict parse, then guard).", + "version": 2, "vectors": [ { "name": "int x, 64-hex y", @@ -394,6 +394,391 @@ "meta": { "blob": "{not json" } + }, + { + "name": "y as integer list", + "refuse": true, + "meta": { + "a": { + "x": 1, + "y": [ + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171 + ] + } + } + }, + { + "name": "y as 16-integer list", + "refuse": true, + "meta": { + "a": { + "x": 1, + "y": [ + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1 + ] + } + } + }, + { + "name": "y as Buffer JSON shape", + "refuse": true, + "meta": { + "a": { + "x": 1, + "y": { + "type": "Buffer", + "data": [ + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171, + 171 + ] + } + } + } + }, + { + "name": "y as index-keyed byte object", + "refuse": true, + "meta": { + "a": { + "x": 1, + "y": { + "0": 171, + "1": 171, + "2": 171, + "3": 171, + "4": 171, + "5": 171, + "6": 171, + "7": 171, + "8": 171, + "9": 171, + "10": 171, + "11": 171, + "12": 171, + "13": 171, + "14": 171, + "15": 171, + "16": 171, + "17": 171, + "18": 171, + "19": 171, + "20": 171, + "21": 171, + "22": 171, + "23": 171, + "24": 171, + "25": 171, + "26": 171, + "27": 171, + "28": 171, + "29": 171, + "30": 171, + "31": 171 + } + } + } + }, + { + "name": "JSON string with duplicate y", + "refuse": true, + "meta": { + "blob": "{\"x\": 1, \"y\": \"abababababababababababababababababababababababababababababababab\", \"y\": \"10\"}" + } + }, + { + "name": "JSON string with duplicate x", + "refuse": true, + "meta": { + "blob": "{\"x\": 1, \"x\": \"junk\", \"y\": \"abababababababababababababababababababababababababababababababab\"}" + } + }, + { + "name": "15-integer list y", + "refuse": false, + "meta": { + "x": 1, + "y": [ + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1 + ] + } + }, + { + "name": "integer list with values over 255", + "refuse": false, + "meta": { + "x": 1, + "y": [ + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300, + 300 + ] + } + }, + { + "name": "float list y", + "refuse": false, + "meta": { + "x": 1, + "y": [ + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5, + 1.5 + ] + } + }, + { + "name": "Buffer shape with short data", + "refuse": false, + "meta": { + "x": 1, + "y": { + "type": "Buffer", + "data": [ + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1, + 1 + ] + } + } + }, + { + "name": "index-keyed object with a gap", + "refuse": false, + "meta": { + "x": 1, + "y": { + "1": 1, + "2": 1, + "3": 1, + "4": 1, + "5": 1, + "6": 1, + "7": 1, + "8": 1, + "9": 1, + "10": 1, + "11": 1, + "12": 1, + "13": 1, + "14": 1, + "15": 1, + "16": 1, + "17": 1, + "18": 1, + "19": 1, + "20": 1, + "21": 1, + "22": 1, + "23": 1, + "24": 1, + "25": 1, + "26": 1, + "27": 1, + "28": 1, + "29": 1, + "30": 1, + "31": 1, + "32": 1 + } + } + }, + { + "name": "coordinate list y", + "refuse": false, + "meta": { + "x": 1, + "y": [ + 1, + 2 + ] + } + } + ], + "raw_payloads": [ + { + "name": "top-level duplicate y", + "refuse": true, + "text": "{\"a\": {\"x\": 1, \"y\": \"abababababababababababababababababababababababababababababababab\", \"y\": \"10\"}}" + }, + { + "name": "top-level duplicate x", + "refuse": true, + "text": "{\"a\": {\"x\": 1, \"x\": \"junk\", \"y\": \"abababababababababababababababababababababababababababababababab\"}}" + }, + { + "name": "duplicate key at root", + "refuse": true, + "text": "{\"note\": \"a\", \"note\": \"b\"}" + }, + { + "name": "benign payload", + "refuse": false, + "text": "{\"a\": {\"x\": 1, \"y\": \"10\"}, \"note\": \"ok\"}" + }, + { + "name": "same key in sibling objects", + "refuse": false, + "text": "{\"a\": {\"y\": \"10\"}, \"b\": {\"y\": \"10\"}}" } ] } diff --git a/tests/test_guard_parser_differential.py b/tests/test_guard_parser_differential.py new file mode 100644 index 0000000..6bdc1fb --- /dev/null +++ b/tests/test_guard_parser_differential.py @@ -0,0 +1,162 @@ +"""Differential property test: share guard versus share parsers. + +Generates near-hex strings (hex pairs, stray hex digits, ASCII and Unicode +whitespace, junk characters, 0x/0X prefixes) from a fixed seed and checks two +properties on every input: + +1. ``parse_share_y`` agrees with a reference strict grammar (optional 0x/0X, + even-length hex, nothing else). A lenient parser fails this. +2. Any input that ANY known decoder turns into >= 16 bytes (the strict + reference, Python's ``bytes.fromhex``, and the legacy truncating JS + decoder) is refused by the guard. A guard that only matches canonical hex + fails this. +""" +from __future__ import annotations + +import random +import re + +import pytest + +from citrate_sdk import crypto +from citrate_sdk.crypto import assert_no_key_share_material +from citrate_sdk.errors import CitrateError + +N = 40_000 +MIN = 16 +_STRICT = re.compile(r"(?:0[xX])?((?:[0-9a-fA-F]{2})*)") +_HEX = "0123456789abcdefABCDEF" +_WS = [" ", "\t", "\n", "\r", " ", "
", " "] +_JUNK = ["z", "g", "x", "X", "-", ":", "1", "."] + + +def strict_ref(s: str) -> bytes | None: + m = _STRICT.fullmatch(s) + return bytes.fromhex(m.group(1)) if m else None + + +def py_fromhex(s: str) -> bytes | None: + t = s[2:] if s[:2] in ("0x", "0X") else s + try: + return bytes.fromhex(t) + except ValueError: + return None + + +def js_legacy(s: str) -> bytes | None: + """The pre-hardening JS decoder: strip a leading 0x, then read + floor(len/2) pairs with parseInt (a trailing odd character is dropped). + Counted only when every pair it reads is real hex.""" + t = s[2:] if s.startswith("0x") else s + body = t[: len(t) // 2 * 2] + if not re.fullmatch(r"[0-9a-fA-F]*", body): + return None + return bytes.fromhex(body) + + +def _gen(rng: random.Random) -> str: + """Three regimes: canonical hex (optionally prefixed), canonical hex with + one to three perturbations, and free-form near-hex.""" + pairs = [rng.choice(_HEX) + rng.choice(_HEX) for _ in range(rng.randint(12, 40))] + mode = rng.random() + if mode < 0.34: + out = pairs + elif mode < 0.67: + out = list(pairs) + for _ in range(rng.randint(1, 3)): + tok = rng.choice(_WS + _JUNK + list(_HEX)) + out.insert(rng.randint(0, len(out)), tok) + else: + out = [] + for _ in range(rng.randint(10, 40)): + r = rng.random() + if r < 0.80: + out.append(rng.choice(_HEX) + rng.choice(_HEX)) + elif r < 0.88: + out.append(rng.choice(_HEX)) + elif r < 0.95: + out.append(rng.choice(_WS)) + else: + out.append(rng.choice(_JUNK)) + prefix = rng.choice(["0x", "0X"]) if rng.random() < 0.3 else "" + return prefix + "".join(out) + + +def _inputs() -> list[str]: + rng = random.Random(0xC17A7E) + return [_gen(rng) for _ in range(N)] + + +def _guard_refuses(y: str) -> bool: + try: + assert_no_key_share_material({"x": 1, "y": y}) + except CitrateError: + return True + return False + + +def test_parser_matches_the_strict_reference() -> None: + bad = [] + for s in _inputs(): + ref = strict_ref(s) + try: + got: bytes | None = crypto.parse_share_y(s) + except ValueError: + got = None + if got != ref: + bad.append(s) + assert bad == [], f"{len(bad)} disagreements, e.g. {bad[:3]!r}" + + +@pytest.mark.parametrize("decoder", [strict_ref, py_fromhex, js_legacy], ids=["strict", "fromhex", "js-legacy"]) +def test_guard_refuses_everything_any_decoder_accepts(decoder: object) -> None: + accepted = 0 + missed = [] + for s in _inputs(): + b = decoder(s) # type: ignore[operator] + if b is not None and len(b) >= MIN: + accepted += 1 + if not _guard_refuses(s): + missed.append(s) + assert accepted > 1000, "generator produced too few decodable inputs to be meaningful" + assert missed == [], f"{len(missed)} decodable inputs passed the guard, e.g. {missed[:3]!r}" + + +# ---- y given as bytes-like values and their JSON shapes ------------------- + +def _byte_forms(b: bytes) -> list[object]: + return [b, bytearray(b), list(b), {"type": "Buffer", "data": list(b)}, {str(i): v for i, v in enumerate(b)}] + + +def test_guard_refuses_share_length_bytes_in_every_form() -> None: + rng = random.Random(0xB17E5) + for n in (16, 17, 32, 64): + for _ in range(50): + b = bytes(rng.randrange(256) for _ in range(n)) + for y in _byte_forms(b): + with pytest.raises(CitrateError): + assert_no_key_share_material({"x": 1 + n % 200, "y": y}) + + +def test_guard_accepts_short_bytes_in_every_form() -> None: + rng = random.Random(0x5407) + for n in (1, 8, 15): + b = bytes(rng.randrange(256) for _ in range(n)) + for y in _byte_forms(b): + assert_no_key_share_material({"x": 1, "y": y}) + + +def test_duplicate_keys_in_payloads_are_refused() -> None: + rng = random.Random(0xD0B1E) + for _ in range(200): + hexy = bytes(rng.randrange(256) for _ in range(32)).hex() + key = rng.choice(["x", "y", "note"]) + texts = [ + '{"a": {"x": 1, "y": "%s", "%s": "10"}}' % (hexy, key) if key == "y" else + '{"a": {"x": 1, "%s": "junk", "%s": 2, "y": "%s"}}' % (key, key, hexy), + ] + for t in texts: + with pytest.raises(CitrateError): + crypto.assert_payload_has_no_key_share_material(t) + with pytest.raises(CitrateError): + assert_no_key_share_material({"blob": t}) diff --git a/tests/test_hardening_round4.py b/tests/test_hardening_round4.py index 9b34237..412a253 100644 --- a/tests/test_hardening_round4.py +++ b/tests/test_hardening_round4.py @@ -26,7 +26,7 @@ SECP256K1_N = 0xFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFEBAAEDCE6AF48A03BBFD25E8CD0364141 VECTORS = Path(__file__).parent / "fixtures" / "share_guard_vectors.json" -VECTORS_SHA256 = "78d317f898cea7c2a808fc34e6566c144cc88709057175e29ac75867b4a1f8d5" +VECTORS_SHA256 = "8f336f469be58046e7984bf61756d513df9ba497aca8feb9fc8641e2c6a4b805" def _mgr() -> tuple[ClassroomManager, list[Any]]: @@ -123,3 +123,15 @@ def boom(k: str) -> None: monkeypatch.setattr(learning.Account, "from_key", staticmethod(boom)) with pytest.raises(ValueError, match="not a valid secp256k1 key"): learning._invite_account("0x" + "11" * 32) + + +_RAW = json.loads(VECTORS.read_text())["raw_payloads"] + + +@pytest.mark.parametrize("vec", _RAW, ids=[v["name"] for v in _RAW]) +def test_shared_raw_payloads(vec: dict[str, Any]) -> None: + if vec["refuse"]: + with pytest.raises(CitrateError): + crypto.assert_payload_has_no_key_share_material(vec["text"]) + else: + crypto.assert_payload_has_no_key_share_material(vec["text"]) diff --git a/tests/test_payload_guard.py b/tests/test_payload_guard.py new file mode 100644 index 0000000..5c9841a --- /dev/null +++ b/tests/test_payload_guard.py @@ -0,0 +1,112 @@ +"""Deploy payload is guarded as sent. + +deploy_model serialises the payload once, guards the parsed payload, and sends +that exact payload. +""" +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any + +import pytest +import rlp # type: ignore[import-untyped] + +from citrate_sdk import CitrateClient +from citrate_sdk.errors import CitrateError +from citrate_sdk.models import ModelConfig + +OWNER = "0x" + "11" * 32 +Y = "ab" * 32 + + +class _Mapping(dict): + def __getitem__(self, k: Any) -> Any: + return {"x": 0, "y": "10"}.get(k, super().get(k)) + + def get(self, k: Any, default: Any = None) -> Any: + return {"x": 0, "y": "10"}.get(k, default) + + +class _Key(str): + """A str subclass; two distinct instances with the same text.""" + + def __hash__(self) -> int: + return id(self) + + def __eq__(self, other: object) -> bool: + return self is other + + +def _client(tmp_path: Path) -> tuple[CitrateClient, dict[str, Any], Path]: + mp = tmp_path / "m.onnx" + mp.write_bytes(b"w") + client = CitrateClient("http://localhost:8545", private_key=OWNER) + box: dict[str, Any] = {} + + def rpc(method: str, params: Any = None) -> Any: + if method == "eth_sendRawTransaction": + box["raw"] = params[0] + return "0x" + "ab" * 32 + return {"eth_getTransactionCount": "0x0", "eth_chainId": hex(40204), "eth_gasPrice": "0x1"}[method] + + stubs: dict[str, Any] = {"_rpc_call": rpc, "_upload_to_ipfs": lambda d: "bafy", + "_wait_for_receipt": lambda h: {"logs": []}, + "_extract_model_id_from_receipt": lambda r: "0x01"} + for name, fn in stubs.items(): + setattr(client, name, fn) + return client, box, mp + + +def _dup_y() -> dict[Any, Any]: + d: dict[Any, Any] = {"x": 1} + d[_Key("y")] = Y + d[_Key("y")] = "10" + return d + + +def _dup_x() -> dict[Any, Any]: + d: dict[Any, Any] = {} + d[_Key("x")] = 1 + d[_Key("x")] = "junk" + d["y"] = Y + return d + + +@pytest.mark.parametrize("meta", [ + {"m": _Mapping({"x": 1, "y": Y})}, + {"m": [_Mapping({"x": "7", "y": Y})]}, + {"m": _Mapping({"x": 1, "y": Y, "note": "n"})}, + {"m": _dup_y()}, + {"m": _dup_x()}, + {"blob": '{"x": 1, "y": "' + Y + '", "y": "10"}'}, +], ids=["case-1", "case-2", "case-3", "case-4", "case-5", "case-6"]) +def test_payload_guarded_as_sent(tmp_path: Path, meta: dict[str, Any]) -> None: + client, box, mp = _client(tmp_path) + with pytest.raises(CitrateError): + client.deploy_model(mp, ModelConfig(name="m", metadata=meta)) + assert "raw" not in box + + +def test_the_guard_sees_exactly_the_bytes_sent(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + from citrate_sdk import client as client_mod + + seen: list[str] = [] + real = client_mod.assert_payload_has_no_key_share_material + + def spy(payload: str) -> None: + seen.append(payload) + real(payload) + + monkeypatch.setattr(client_mod, "assert_payload_has_no_key_share_material", spy) + client, box, mp = _client(tmp_path) + client.deploy_model(mp, ModelConfig(name="m", metadata={"note": "ok"})) + sent = bytes(rlp.decode(bytes.fromhex(box["raw"][2:]))[5]).decode() + assert seen == [sent] + assert "model_hash" in json.loads(sent) + + +def test_benign_deploy_still_sends(tmp_path: Path) -> None: + client, box, mp = _client(tmp_path) + client.deploy_model(mp, ModelConfig(name="m", metadata={"x": 1, "y": "10"})) + assert "raw" in box diff --git a/uv.lock b/uv.lock index 3519900..4016f98 100644 --- a/uv.lock +++ b/uv.lock @@ -548,7 +548,7 @@ wheels = [ [[package]] name = "citrate-labs-sdk" -version = "0.6.2" +version = "0.6.3" source = { editable = "." } dependencies = [ { name = "cryptography" },