From 1107a610ae03d62c4c646be3cc7fefb5f3287bd4 Mon Sep 17 00:00:00 2001 From: SaulBuilds Date: Fri, 25 Sep 2026 08:22:27 -0700 Subject: [PATCH 1/4] fix(transport): stricter endpoint URL validation (PBA-L6b-026 hardening) The transport gate now refuses characters outside RFC 3986 and any userinfo, parses with urllib3 (the parser requests uses) and requires it to agree with urllib.parse on the host, and reports unparseable URLs as InsecureTransportError. urllib3 is now an explicit dependency (already locked via requests). All gated clients and verify_wallet_address_on_chain (PBA-L6b-027) are covered by tests, including a live listener check. test_network_error_handling now uses a valid port, because an unparseable URL is refused at construction. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01FYQkdsk54yob6FD24jAT8P --- citrate_sdk/_url_security.py | 43 ++++- pyproject.toml | 3 + tests/test_integration.py | 13 +- ...test_pba_r2_l6b_026_parser_differential.py | 149 ++++++++++++++++++ uv.lock | 2 + 5 files changed, 204 insertions(+), 6 deletions(-) create mode 100644 tests/test_pba_r2_l6b_026_parser_differential.py diff --git a/citrate_sdk/_url_security.py b/citrate_sdk/_url_security.py index 8cb382f..aaab12a 100644 --- a/citrate_sdk/_url_security.py +++ b/citrate_sdk/_url_security.py @@ -14,9 +14,17 @@ """ import ipaddress +import re import unicodedata from urllib.parse import urlparse +from urllib3.exceptions import LocationParseError +from urllib3.util import parse_url + +# PBA-L6b-026 follow-up: RFC 3986 characters only (unreserved, reserved and +# "%"). URL parsers disagree on anything else, so it is refused. +_RFC3986_CHARS = re.compile(r"^[A-Za-z0-9\-._~:/?#\[\]@!$&'()*+,;=%]*$") + # PBA-L6b-026: only these schemes are ever accepted. _ALLOWED_SCHEMES = frozenset({"https", "http"}) @@ -87,24 +95,51 @@ def enforce_transport_security(url: str, *, allow_insecure_http: bool = False) - "lets an http:// endpoint slip past this check (PBA-L6b-026)." ) - parsed = urlparse(url) + if not _RFC3986_CHARS.match(url): + bad = next(ch for ch in url if not _RFC3986_CHARS.match(ch)) + raise InsecureTransportError( + f"Refusing endpoint URL {url!r}: character {bad!r} is not allowed in a URL " + "(RFC 3986) (PBA-L6b-026)." + ) + + try: + parsed = urlparse(url) + parsed_host = parsed.hostname + requests_view = parse_url(url) + except (ValueError, LocationParseError): + raise InsecureTransportError(f"Refusing unparseable endpoint URL {url!r} (PBA-L6b-026).") scheme = parsed.scheme.lower() if scheme not in _ALLOWED_SCHEMES: raise InsecureTransportError( f"Refusing endpoint URL {url!r}: scheme {parsed.scheme!r} is not https or " "http (PBA-L6b-026)." ) - if not parsed.hostname: + if not parsed_host: raise InsecureTransportError(f"Refusing endpoint URL {url!r}: no host (PBA-L6b-026).") + # Credentials in the URL are refused outright; pass them as headers. + if "@" in parsed.netloc or parsed.username is not None or requests_view.auth is not None: + raise InsecureTransportError( + f"Refusing endpoint URL {url!r}: userinfo (user@ / user:pass@) is not " + "allowed in an endpoint URL; pass credentials as headers (PBA-L6b-026)." + ) + # Gate on the host exactly as urllib3 (what requests connects to) parses + # it, and require urllib.parse to agree. + transport_host = (requests_view.host or "").strip("[]").lower() + if transport_host != parsed_host.lower(): + raise InsecureTransportError( + f"Refusing endpoint URL {url!r}: urllib.parse sees host {parsed_host!r} but " + f"requests would connect to {requests_view.host!r} (PBA-L6b-026)." + ) + if scheme == "https": return url - if _is_local_host(parsed.hostname): + if _is_local_host(parsed_host): return url if not allow_insecure_http: raise InsecureTransportError( - f"Refusing to connect to remote host {parsed.hostname!r} over " + f"Refusing to connect to remote host {parsed_host!r} over " f"plaintext http:// ({url!r}). Traffic (including signed " f"transactions, private inputs, and bearer credentials) would be " f"sent in cleartext and could be intercepted or tampered with. Use " diff --git a/pyproject.toml b/pyproject.toml index c48b24c..616f249 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -55,6 +55,9 @@ classifiers = [ # HKDF envelope (crypto.py) and RSA ID-token verification (identity/jwt.py). dependencies = [ "requests~=2.33", + # PBA-L6b-026: the transport gate parses URLs with urllib3 (what requests + # connects with) to rule out parser differentials; pinned explicitly. + "urllib3>=2.0,<3", "cryptography>=48.0.1,<51", "eth-account~=0.9", "web3~=7.15", diff --git a/tests/test_integration.py b/tests/test_integration.py index 2f2314b..d06f351 100644 --- a/tests/test_integration.py +++ b/tests/test_integration.py @@ -17,6 +17,7 @@ import pytest # Import SDK modules +from citrate_sdk._url_security import InsecureTransportError from citrate_sdk.client import CitrateClient from citrate_sdk.crypto import KeyManager from citrate_sdk.errors import CitrateError @@ -486,8 +487,16 @@ def test_invalid_address_error(self, client): client.get_balance('not_a_valid_address') def test_network_error_handling(self): - """Network errors are handled gracefully""" - bad_client = CitrateClient(rpc_url='http://localhost:99999') + """Network errors are handled gracefully. + + PBA-L6b-026 follow-up: this used port 99999, which is not a valid port. + The transport gate now parses URLs the way urllib3/requests do and refuses + an unparseable one at construction, so the network-error path is + exercised with a valid port that has no listener instead. + """ + with pytest.raises(InsecureTransportError): + CitrateClient(rpc_url='http://localhost:99999') + bad_client = CitrateClient(rpc_url='http://localhost:1') with pytest.raises(CitrateError): bad_client.get_chain_id() diff --git a/tests/test_pba_r2_l6b_026_parser_differential.py b/tests/test_pba_r2_l6b_026_parser_differential.py new file mode 100644 index 0000000..33e984f --- /dev/null +++ b/tests/test_pba_r2_l6b_026_parser_differential.py @@ -0,0 +1,149 @@ +"""Transport-gate hardening (PBA-L6b-026 follow-up). + +The gate must judge the same host the HTTP stack will connect to. It refuses +characters outside RFC 3986 and any userinfo, and requires urllib3 (what +requests uses) and urllib.parse to agree on the host. The live test runs a +real listener on this host's non-loopback address and checks that a refused +URL never reaches it. +""" +from __future__ import annotations + +import http.server +import json +import socket +import threading +from typing import Any +from unittest import mock + +import pytest +import requests + +from citrate_sdk import CitrateClient +from citrate_sdk._url_security import InsecureTransportError, enforce_transport_security +from citrate_sdk.gateway import GatewayClient +from citrate_sdk.identity import wallet +from citrate_sdk.ipfs import IPFSClient +from citrate_sdk.memory import MemoryClient + +BYPASSES = [ + "http://evil.com\\@localhost", + "http://evil.com\\@127.0.0.1:8545", + "http://evil.com:8545\\x@localhost", + "http://evil.com\\@[::1]:8545", + "https://evil.com\\@localhost", + "http://localhost\\@evil.com", + "http://u:p@localhost:8545", + "http://evil.com@localhost", + "http://localhost@evil.com", + "http://[::1]@evil.com", + "http://evil.com%5C@localhost", + "http://evil.com^@localhost", + "http://evil.com`@localhost", + "http://ev{il}.com", + "http://evil.com|x@localhost", + 'http://evil.com"@localhost', + "http://evil.com@localhost", +] + + +@pytest.mark.parametrize("url", BYPASSES) +def test_parser_differential_inputs_are_refused(url: str) -> None: + with pytest.raises(InsecureTransportError): + enforce_transport_security(url) + with pytest.raises(InsecureTransportError): + enforce_transport_security(url, allow_insecure_http=True) + + +@pytest.mark.parametrize("url", [ + "http://localhost:8545", "http://127.0.0.1:5001/api", "http://[::1]:8545", "https://rpc.citrate.ai", + "https://rpc.citrate.ai/path?q=1&x=%20#frag", "HTTPS://RPC.CITRATE.AI", "https://rpc.citrate.ai:443/a-b_c.d~e", +]) +def test_ordinary_urls_still_pass(url: str) -> None: + assert enforce_transport_security(url) == url + + +def test_remote_http_opt_in_still_works() -> None: + assert enforce_transport_security("http://10.0.0.5:8545", allow_insecure_http=True) == "http://10.0.0.5:8545" + + +def test_client_never_sends_to_the_urllib3_host() -> None: + sent: list[str] = [] + + def fake_send(self: Any, request: Any, **kw: Any) -> requests.Response: + sent.append(request.url) + r = requests.Response() + r.status_code = 200 + r._content = b'{"jsonrpc":"2.0","id":1,"result":"0x9d0c"}' + return r + + with mock.patch.object(requests.adapters.HTTPAdapter, "send", fake_send): + with pytest.raises(InsecureTransportError): + CitrateClient("http://evil.com\\@localhost").get_chain_id() + assert sent == [] + + +def _non_loopback_ipv4() -> str: + s = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) + try: + s.connect(("192.0.2.1", 9)) # TEST-NET-1; UDP connect sends nothing + return str(s.getsockname()[0]) + finally: + s.close() + + +def test_live_listener_on_a_non_loopback_address_receives_nothing() -> None: + """Live check against a real listener on a non-loopback address: refused + URLs never reach it; the explicit opt-in path does (so the listener works).""" + ip = _non_loopback_ipv4() + assert not ip.startswith("127."), ip + hits: list[str] = [] + + class H(http.server.BaseHTTPRequestHandler): + def do_POST(self) -> None: # noqa: N802 + hits.append(self.path) + body = json.dumps({"jsonrpc": "2.0", "id": 1, "result": "0x9d0c"}).encode() + self.send_response(200) + self.send_header("content-type", "application/json") + self.send_header("content-length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def log_message(self, *a: Any) -> None: + pass + + srv = http.server.HTTPServer((ip, 0), H) + port = srv.server_address[1] + t = threading.Thread(target=srv.serve_forever, daemon=True) + t.start() + try: + # Sanity: the plain remote URL is refused by the gate. + with pytest.raises(InsecureTransportError): + CitrateClient(f"http://{ip}:{port}") + # A mixed-parser URL must be refused, and nothing may reach the listener. + with pytest.raises(InsecureTransportError): + CitrateClient(f"http://{ip}:{port}\\@localhost").get_chain_id() + # And the opt-in path really does reach it (the listener works). + assert CitrateClient(f"http://{ip}:{port}", allow_insecure_http=True).get_chain_id() == 40204 + assert hits == ["/"] + finally: + srv.shutdown() + srv.server_close() + + +@pytest.mark.parametrize("make", [ + lambda u: IPFSClient(u), + lambda u: GatewayClient("cgk_test", base_url=u), + lambda u: MemoryClient(u), +]) +def test_every_gated_client_refuses_the_backslash_form(make: Any) -> None: + with pytest.raises(InsecureTransportError): + make("http://evil.com\\@localhost") + + +def test_wallet_verification_inherits_the_fix() -> None: + """PBA-L6b-027 re-test: verify_wallet_address_on_chain uses the same gate.""" + calls: list[Any] = [] + with mock.patch.object(wallet.requests, "post", lambda *a, **k: calls.append(a)): + with pytest.raises(InsecureTransportError): + wallet.verify_wallet_address_on_chain("0x" + "42" * 32, rpc_url="http://evil.com\\@localhost:8545") + assert calls == [] diff --git a/uv.lock b/uv.lock index ef090b4..24cf0ad 100644 --- a/uv.lock +++ b/uv.lock @@ -541,6 +541,7 @@ dependencies = [ { name = "numpy", version = "2.5.3", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.12'" }, { name = "requests" }, { name = "typing-extensions" }, + { name = "urllib3" }, { name = "web3" }, ] @@ -591,6 +592,7 @@ requires-dist = [ { name = "sphinx-rtd-theme", marker = "extra == 'docs'", specifier = ">=1.0" }, { name = "tomli", marker = "extra == 'dev'", specifier = ">=2.0" }, { name = "typing-extensions", specifier = "~=4.0" }, + { name = "urllib3", specifier = ">=2.0,<3" }, { name = "web3", specifier = "~=7.15" }, ] provides-extras = ["dev", "docs"] From 516cb7257f82f0d340052a7e8bac708dfd74438f Mon Sep 17 00:00:00 2001 From: SaulBuilds Date: Fri, 25 Sep 2026 08:22:27 -0700 Subject: [PATCH 2/4] fix(ipfs): make content verification mandatory (PBA-L6b-030 hardening) Downloads whose address cannot verify its own content now require expected_sha256, or an explicit verify=False opt-out that logs a warning. The opt-out never skips a supplied hash or the size cap. Earlier size-cap tests pass verify=False. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01FYQkdsk54yob6FD24jAT8P --- citrate_sdk/ipfs.py | 36 ++++++-- tests/test_pba_r2_l6b_030_ipfs.py | 8 +- tests/test_pba_r2_l6b_030_mandatory_verify.py | 83 +++++++++++++++++++ 3 files changed, 118 insertions(+), 9 deletions(-) create mode 100644 tests/test_pba_r2_l6b_030_mandatory_verify.py diff --git a/citrate_sdk/ipfs.py b/citrate_sdk/ipfs.py index 679cfea..34ab8e8 100644 --- a/citrate_sdk/ipfs.py +++ b/citrate_sdk/ipfs.py @@ -6,6 +6,7 @@ import hashlib import hmac import json +import logging from typing import Any, cast import requests @@ -13,6 +14,8 @@ from ._url_security import enforce_transport_security from .errors import IPFSError +_log = logging.getLogger(__name__) + #: PBA-L6b-030: default ceiling on a single download (bytes). Pass #: ``max_bytes`` to raise or lower it per call. DEFAULT_MAX_DOWNLOAD_BYTES = 1 << 30 # 1 GiB @@ -124,6 +127,7 @@ def download_bytes( *, expected_sha256: str | None = None, max_bytes: int | None = None, + verify: bool = True, ) -> bytes: """ Download bytes from IPFS using hash @@ -139,10 +143,18 @@ def download_bytes( - self-describing addresses are verified with no hint: ``sha256:`` pointers and CIDv1 raw/sha2-256 (``bafkrei...``). + Verification is MANDATORY (PBA-L6b-030 follow-up): if the address + cannot verify itself and no ``expected_sha256`` is given, the download + is refused before any request. ``verify=False`` is the explicit opt-out + for callers who accept unverified bytes; it logs a warning. It never + skips a supplied ``expected_sha256`` or the size cap. + Args: ipfs_hash: IPFS hash (CID) - expected_sha256: optional hex sha256 the content must match + expected_sha256: hex sha256 the content must match (required for + dag-pb CIDs unless ``verify=False``) max_bytes: optional size ceiling for this download + verify: set False to accept unverifiable content (logged) Returns: Downloaded bytes @@ -151,6 +163,19 @@ def download_bytes( IPFSError: If download fails, exceeds the cap, or fails verification """ limit = DEFAULT_MAX_DOWNLOAD_BYTES if max_bytes is None else max_bytes + committed = _expected_digest_from_address(ipfs_hash) + if committed is None and expected_sha256 is None: + if verify: + raise IPFSError( + f"refusing to download {ipfs_hash}: this address does not verify its own " + "content (dag-pb CIDs hash a chunked DAG, not the bytes), so pass " + "expected_sha256 (e.g. the on-chain model_hash), or verify=False to accept " + "unverified bytes explicitly (PBA-L6b-030)." + ) + _log.warning( + "IPFS download of %s is UNVERIFIED: verify=False and no expected_sha256; the " + "gateway can return arbitrary bytes (PBA-L6b-030).", ipfs_hash, + ) try: response = self.session.post( f"{self.api_url}/api/v0/cat", @@ -177,7 +202,6 @@ def download_bytes( want = expected_sha256[2:] if expected_sha256.startswith("0x") else expected_sha256 if not hmac.compare_digest(digest.hex(), want.lower()): raise IPFSError(f"IPFS content for {ipfs_hash} does not match expected_sha256") - committed = _expected_digest_from_address(ipfs_hash) if committed is not None and not hmac.compare_digest(digest, committed): raise IPFSError(f"IPFS content does not match its address {ipfs_hash}") return data @@ -362,6 +386,7 @@ def download( *, expected_sha256: str | None = None, max_bytes: int | None = None, + verify: bool = True, ) -> bytes: """ Download data with automatic fallback @@ -382,7 +407,7 @@ def download( if self.active_client: try: return self.active_client.download_bytes( - ipfs_hash, expected_sha256=expected_sha256, max_bytes=max_bytes) + ipfs_hash, expected_sha256=expected_sha256, max_bytes=max_bytes, verify=verify) except IPFSError: pass @@ -391,7 +416,7 @@ def download( try: if client.is_available(): return client.download_bytes( - ipfs_hash, expected_sha256=expected_sha256, max_bytes=max_bytes) + ipfs_hash, expected_sha256=expected_sha256, max_bytes=max_bytes, verify=verify) except IPFSError as e: last_error = e @@ -471,6 +496,7 @@ def download_from_ipfs( *, expected_sha256: str | None = None, max_bytes: int | None = None, + verify: bool = True, ) -> bytes: """ Convenience function to download data from IPFS @@ -483,4 +509,4 @@ def download_from_ipfs( Downloaded bytes """ manager = get_ipfs_manager(ipfs_urls) - return manager.download(ipfs_hash, expected_sha256=expected_sha256, max_bytes=max_bytes) + return manager.download(ipfs_hash, expected_sha256=expected_sha256, max_bytes=max_bytes, verify=verify) diff --git a/tests/test_pba_r2_l6b_030_ipfs.py b/tests/test_pba_r2_l6b_030_ipfs.py index be80635..e2e773a 100644 --- a/tests/test_pba_r2_l6b_030_ipfs.py +++ b/tests/test_pba_r2_l6b_030_ipfs.py @@ -70,16 +70,16 @@ def test_sha256_pointer_is_verified() -> None: def test_size_cap_stops_the_stream() -> None: with pytest.raises(IPFSError, match="exceeds max_bytes"): - _client(b"a" * 1025).download_bytes("QmWhatever", max_bytes=1024) - assert _client(b"a" * 1024).download_bytes("QmWhatever", max_bytes=1024) == b"a" * 1024 + _client(b"a" * 1025).download_bytes("QmWhatever", max_bytes=1024, verify=False) + assert _client(b"a" * 1024).download_bytes("QmWhatever", max_bytes=1024, verify=False) == b"a" * 1024 def test_default_cap_exists() -> None: c = _client(b"") with mock.patch("citrate_sdk.ipfs.DEFAULT_MAX_DOWNLOAD_BYTES", 10): with pytest.raises(IPFSError, match="exceeds max_bytes"): - _client(b"a" * 11).download_bytes("QmWhatever") - assert c.download_bytes("QmWhatever") == b"" + _client(b"a" * 11).download_bytes("QmWhatever", verify=False) + assert c.download_bytes("QmWhatever", verify=False) == b"" def test_manager_threads_the_checks_through() -> None: diff --git a/tests/test_pba_r2_l6b_030_mandatory_verify.py b/tests/test_pba_r2_l6b_030_mandatory_verify.py new file mode 100644 index 0000000..bf149bc --- /dev/null +++ b/tests/test_pba_r2_l6b_030_mandatory_verify.py @@ -0,0 +1,83 @@ +"""IPFS hardening (PBA-L6b-030 follow-up): verification is mandatory. + +A download whose address cannot verify its own content requires +``expected_sha256``, or an explicit ``verify=False`` opt-out that logs a +warning. +""" +from __future__ import annotations + +import hashlib +import io +import logging +from typing import Any +from unittest import mock + +import pytest +import requests + +from citrate_sdk.errors import IPFSError +from citrate_sdk.ipfs import IPFSClient, IPFSManager, download_from_ipfs + +EVIL = b"ATTACKER-BYTES" +DAG_PB = ["QmYwAPJzv5CZsnA625s3Xf2nemtYgPpHdWEz79ojWnPbdG", + "bafybeigdyrzt5sfp7udm7hu76uh7y26nf3efuylqabf3oclgtqy55fbzdi"] + + +def _client(body: bytes = EVIL) -> IPFSClient: + c = IPFSClient("http://127.0.0.1:5001") + + def post(url: str, params: Any = None, timeout: Any = None, stream: bool = False) -> requests.Response: + r = requests.Response() + r.status_code = 200 + r.raw = io.BytesIO(body) + return r + c.session.post = post # type: ignore[method-assign,assignment] + return c + + +@pytest.mark.parametrize("cid", DAG_PB) +def test_unverifiable_cid_without_hash_is_refused(cid: str) -> None: + with pytest.raises(IPFSError, match="expected_sha256"): + _client().download_bytes(cid) + + +@pytest.mark.parametrize("cid", DAG_PB) +def test_expected_sha256_still_works(cid: str) -> None: + assert _client().download_bytes(cid, expected_sha256=hashlib.sha256(EVIL).hexdigest()) == EVIL + with pytest.raises(IPFSError, match="does not match"): + _client().download_bytes(cid, expected_sha256="00" * 32) + + +def test_explicit_opt_out_returns_bytes_and_warns(caplog: pytest.LogCaptureFixture) -> None: + with caplog.at_level(logging.WARNING, logger="citrate_sdk.ipfs"): + assert _client().download_bytes(DAG_PB[0], verify=False) == EVIL + assert any("UNVERIFIED" in r.getMessage() and DAG_PB[0] in r.getMessage() for r in caplog.records) + + +def test_opt_out_does_not_skip_a_supplied_hash_or_the_cap() -> None: + with pytest.raises(IPFSError, match="does not match"): + _client().download_bytes(DAG_PB[0], verify=False, expected_sha256="00" * 32) + with pytest.raises(IPFSError, match="exceeds max_bytes"): + _client().download_bytes(DAG_PB[0], verify=False, max_bytes=3) + + +def test_manager_and_helper_enforce_it_too() -> None: + m = IPFSManager("http://127.0.0.1:5001") + m.primary = _client() + m.primary.is_available = lambda: True # type: ignore[method-assign] + with pytest.raises(IPFSError, match="expected_sha256"): + m.download(DAG_PB[0]) + assert m.download(DAG_PB[0], verify=False) == EVIL + with mock.patch("citrate_sdk.ipfs.get_ipfs_manager", return_value=m): + with pytest.raises(IPFSError, match="expected_sha256"): + download_from_ipfs(DAG_PB[1]) + assert download_from_ipfs(DAG_PB[1], verify=False) == EVIL + + +def test_refusal_happens_before_any_request() -> None: + c = IPFSClient("http://127.0.0.1:5001") + calls: list[Any] = [] + c.session.post = lambda *a, **k: calls.append(a) # type: ignore[method-assign,assignment] + with pytest.raises(IPFSError, match="expected_sha256"): + c.download_bytes(DAG_PB[0]) + assert calls == [] From f48de0e4c9dd4e97cfa75b94cddcbdd27c2a9d5a Mon Sep 17 00:00:00 2001 From: SaulBuilds Date: Fri, 25 Sep 2026 08:22:27 -0700 Subject: [PATCH 3/4] fix(crypto): key-sharing hardening (PBA-L6b-003 follow-up) Holder dedupe is pinned by a canonical-encoding test. threshold_shares=1 needs allow_single_holder_recovery=True. The deploy guard also refuses share-shaped values and JSON-encoded blobs, not only the known field names. Plumbing tests that used unverifiable CIDs pass verify=False; the 1-of-1 boundary test passes the new opt-in. Mutation-checked with mutmut: every remaining survivor changes only message text or is equivalent. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01FYQkdsk54yob6FD24jAT8P --- citrate_sdk/crypto.py | 46 +++++++- tests/test_pba_r2_followup_shares.py | 136 ++++++++++++++++++++++++ tests/test_pba_r2_mutation_hardening.py | 20 ++-- 3 files changed, 192 insertions(+), 10 deletions(-) create mode 100644 tests/test_pba_r2_followup_shares.py diff --git a/citrate_sdk/crypto.py b/citrate_sdk/crypto.py index c98d376..3c9e484 100644 --- a/citrate_sdk/crypto.py +++ b/citrate_sdk/crypto.py @@ -5,6 +5,7 @@ import hashlib import hmac import json +import re import secrets from typing import Any, cast @@ -27,8 +28,35 @@ _MAX_GUARD_DEPTH = 32 +_HEX_RE = re.compile(r"^(0x)?[0-9a-fA-F]+$") + + +def _looks_like_share(d: dict[Any, Any]) -> bool: + """A raw Shamir share ({x, y} with y as hex/bytes) or a holder-wrapped + share record ({holder_public_key/holderPublicKey, envelope}).""" + y = d.get("y") + if "x" in d and (isinstance(y, (bytes, bytearray)) or (isinstance(y, str) and _HEX_RE.match(y))): + return True + return "envelope" in d and ("holder_public_key" in d or "holderPublicKey" in d) + + def assert_no_key_share_material(value: Any, _depth: int = 0) -> None: - """Raise CitrateError if ``value`` carries a key-share field (PBA-L6b-003).""" + """Raise CitrateError if ``value`` carries key-share material (PBA-L6b-003). + + Refuses the known share field names AND, by structure, any object that + looks like a share ({x, y-hex}) or a wrapped share record, at any depth, + including inside JSON-encoded string values (a renamed field such as + ``myShares`` is caught too). + """ + if isinstance(value, str): + # Any string that parses as JSON is checked too (no size cap: a padded + # blob must not slip through). + try: + decoded = json.loads(value) + except (ValueError, RecursionError): + return + assert_no_key_share_material(decoded, _depth + 1) + return if isinstance(value, dict): items = list(value.items()) elif isinstance(value, (list, tuple)): @@ -40,6 +68,12 @@ def assert_no_key_share_material(value: Any, _depth: int = 0) -> None: f"deploy_model: metadata is nested more than {_MAX_GUARD_DEPTH} levels deep; refusing to " "publish calldata that cannot be fully checked for key-share material (PBA-L6b-003)." ) + if isinstance(value, dict) and _looks_like_share(value): + raise CitrateError( + "deploy_model: refusing to publish a value shaped like a key share ({x, y} or a " + "wrapped share record) in public deploy calldata. Deliver key shares to their " + "holders off-chain (PBA-L6b-003)." + ) for k, v in items: if k in SHARE_FIELD_DENYLIST: raise CitrateError( @@ -259,6 +293,12 @@ def _plan_key_sharing(config: 'EncryptionConfig') -> tuple[int, int, list[str]] f"total_shares={total!r}); need integers with 1 <= threshold_shares <= " "total_shares <= 255." ) + if threshold == 1 and not config.allow_single_holder_recovery: + raise CitrateError( + "encrypt_model: threshold_shares=1 lets ANY single holder recover the model key, " + "which is not threshold sharing. Pass allow_single_holder_recovery=True if that " + "is really intended (PBA-L6b-003 follow-up)." + ) holders = config.share_holder_public_keys if not holders: raise CitrateError( @@ -724,6 +764,7 @@ def __init__( threshold_shares: int = 0, total_shares: int = 0, share_holder_public_keys: list[str] | None = None, + allow_single_holder_recovery: bool = False, ): """ Args: @@ -734,6 +775,8 @@ def __init__( (PBA-L6b-003): one distinct secp256k1 public key per share. Each share is ECDH-wrapped to its holder and returned for off-chain delivery; shares are never written to deploy calldata. + allow_single_holder_recovery: threshold_shares=1 (any one holder + alone recovers the key) is refused unless this is True. """ self.algorithm = algorithm self.key_derivation = key_derivation @@ -741,6 +784,7 @@ def __init__( self.threshold_shares = threshold_shares self.total_shares = total_shares self.share_holder_public_keys = share_holder_public_keys + self.allow_single_holder_recovery = allow_single_holder_recovery def generate_model_key() -> str: diff --git a/tests/test_pba_r2_followup_shares.py b/tests/test_pba_r2_followup_shares.py new file mode 100644 index 0000000..feb7225 --- /dev/null +++ b/tests/test_pba_r2_followup_shares.py @@ -0,0 +1,136 @@ +"""Key-sharing hardening (PBA-L6b-003 follow-up). + +1. Holder de-duplication runs on canonical keys (same key in compressed, + uncompressed and 0x-prefixed forms counts once). +2. threshold_shares=1 needs allow_single_holder_recovery=True. +3. The deploy guard recognises share-shaped values by structure as well as by + field name. +""" +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any + +import pytest + +from citrate_sdk import CitrateClient, KeyManager +from citrate_sdk.crypto import EncryptionConfig, assert_no_key_share_material +from citrate_sdk.errors import CitrateError +from citrate_sdk.finite_field import split_secret_bytes +from citrate_sdk.models import ModelConfig + +OWNER = "0x" + "11" * 32 +H = [KeyManager("0x" + b * 32) for b in ("22", "33", "44")] + + +def _cfg(pubs: list[str], t: int = 2, **kw: Any) -> EncryptionConfig: + return EncryptionConfig(threshold_shares=t, total_shares=len(pubs), share_holder_public_keys=pubs, **kw) + + +class TestCanonicalDedupe: + def test_same_holder_compressed_and_uncompressed_is_refused(self) -> None: + comp = H[0].get_public_key() + unc = H[0].ecdh_manager.get_public_key_uncompressed().hex() + assert comp != unc + with pytest.raises(CitrateError, match="distinct"): + KeyManager(OWNER).encrypt_model_with_key_shares(b"m", _cfg([comp, unc, H[1].get_public_key()])) + + def test_same_holder_with_and_without_0x_is_refused(self) -> None: + comp = H[0].get_public_key() + with pytest.raises(CitrateError, match="distinct"): + KeyManager(OWNER).encrypt_model_with_key_shares(b"m", _cfg([comp, "0x" + comp, H[1].get_public_key()])) + + def test_distinct_holders_in_mixed_encodings_are_accepted(self) -> None: + pubs = [H[0].get_public_key(), H[1].ecdh_manager.get_public_key_uncompressed().hex(), + "0x" + H[2].get_public_key()] + _, _, envs = KeyManager(OWNER).encrypt_model_with_key_shares(b"m", _cfg(pubs)) + assert len({e["holder_public_key"] for e in envs}) == 3 + + +class TestThresholdOne: + def test_threshold_one_is_refused_by_default(self) -> None: + with pytest.raises(CitrateError, match=r"threshold_shares=1 lets ANY single holder recover.*allow_single_holder_recovery=True"): + KeyManager(OWNER).encrypt_model_with_key_shares(b"m", _cfg([h.get_public_key() for h in H], t=1)) + + def test_threshold_one_with_explicit_opt_in(self) -> None: + _, meta, envs = KeyManager(OWNER).encrypt_model_with_key_shares( + b"m", _cfg([h.get_public_key() for h in H], t=1, allow_single_holder_recovery=True)) + assert meta["key_sharing"]["threshold"] == 1 and len(envs) == 3 + + def test_threshold_two_needs_no_opt_in(self) -> None: + KeyManager(OWNER).encrypt_model_with_key_shares(b"m", _cfg([h.get_public_key() for h in H], t=2)) + + +class TestStructuralGuard: + shares = [{"x": str(x), "y": y.hex()} for x, y in split_secret_bytes(b"k" * 32, 2, 3)] + + @pytest.mark.parametrize("meta", [ + {"myShares": shares}, + {"a": [{"x": 1, "y": "ab12"}]}, + {"a": {"x": "1", "y": "0xab"}}, + {"blob": json.dumps({"parts": shares})}, + {"blob": json.dumps([{"x": 2, "y": "cd"}])}, + {"w": {"holder_public_key": "02" + "11" * 32, "envelope": "{}"}}, + {"w": [{"holderPublicKey": "02" + "11" * 32, "envelope": "{}"}]}, + {"y": {"x": 1, "y": b"\\x01"}}, + {"padded": ' {"x": 1, "y": "ab"} '}, + {"big": json.dumps({"pad": "a" * 1_100_000, "s": {"x": 1, "y": "ab"}})}, + ], ids=lambda m: str(list(m)[0])) + def test_share_shaped_values_are_refused(self, meta: dict[str, Any]) -> None: + with pytest.raises(CitrateError, match="key share|key-share"): + assert_no_key_share_material(meta) + + @pytest.mark.parametrize("meta", [ + {"x": 1, "y": 2}, + {"x": 1, "y": "not hex"}, + {"point": {"x": 1}}, + {"blob": "{not json"}, + {"envelope": "e"}, + {"text": "[1, 2, 3]"}, + {"x": 1, "y": "zz12"}, + {"x": 1, "y": "12zz"}, + {"n": "123"}, + ]) + def test_ordinary_metadata_passes(self, meta: dict[str, Any]) -> None: + assert_no_key_share_material(meta) + + def test_json_in_string_nesting_counts_toward_depth(self) -> None: + v: Any = {"leaf": 1} + for _ in range(20): + v = json.dumps({"n": v}) + with pytest.raises(CitrateError, match="nested more than 32 levels"): + assert_no_key_share_material({"v": v}) + + def test_message(self) -> None: + with pytest.raises(CitrateError, match=r"shaped like a key share \(\{x, y\} or a wrapped share record\) in public deploy calldata\. Deliver key shares"): + assert_no_key_share_material({"a": {"x": 1, "y": "ab"}}) + + def test_deploy_refuses_renamed_share_field(self, tmp_path: Path) -> None: + mp = tmp_path / "m.onnx" + mp.write_bytes(b"w") + client = CitrateClient("http://localhost:8545", private_key=OWNER) + sent: list[Any] = [] + + def rpc(method: str, params: Any = None) -> Any: + sent.append(method) + return "0x0" + stubs: dict[str, Any] = {"_rpc_call": rpc, "_upload_to_ipfs": lambda d: "bafy"} + for name, fn in stubs.items(): + setattr(client, name, fn) + with pytest.raises(CitrateError, match="key share"): + client.deploy_model(mp, ModelConfig(name="m", metadata={"myShares": self.shares})) + assert "eth_sendRawTransaction" not in sent + + def test_real_encrypted_deploy_metadata_is_not_a_false_positive(self) -> None: + pubs = [h.get_public_key() for h in H] + _, meta, _ = KeyManager(OWNER).encrypt_model_with_key_shares(b"m", _cfg(pubs)) + assert_no_key_share_material({"encryption_metadata": meta}) + + +@pytest.mark.parametrize(("t", "n"), [(True, 3), (2, True), (True, True)]) +def test_bool_share_parameters_are_parameter_errors(t: Any, n: Any) -> None: + pubs = [h.get_public_key() for h in H] + cfg = EncryptionConfig(threshold_shares=t, total_shares=n, share_holder_public_keys=pubs) + with pytest.raises(CitrateError, match="invalid share parameters"): + KeyManager(OWNER).encrypt_model_with_key_shares(b"m", cfg) diff --git a/tests/test_pba_r2_mutation_hardening.py b/tests/test_pba_r2_mutation_hardening.py index 97b77ce..2e233b0 100644 --- a/tests/test_pba_r2_mutation_hardening.py +++ b/tests/test_pba_r2_mutation_hardening.py @@ -95,7 +95,9 @@ def test_invalid_share_parameters(self, t: Any, n: Any) -> None: @pytest.mark.parametrize(("t", "n"), [(1, 1), (3, 3), (2, 255)]) def test_boundary_share_parameters_accepted(self, t: int, n: int) -> None: - cfg = EncryptionConfig(threshold_shares=t, total_shares=n, share_holder_public_keys=self._pubs(n)) + # threshold 1 needs the explicit opt-in since the R2 follow-up. + cfg = EncryptionConfig(threshold_shares=t, total_shares=n, share_holder_public_keys=self._pubs(n), + allow_single_holder_recovery=(t == 1)) _, meta, envs = KeyManager(OWNER).encrypt_model_with_key_shares(b"m", cfg) assert meta["key_sharing"] == {"threshold": t, "total_shares": n} assert [e["x"] for e in envs] == list(range(1, n + 1)) @@ -441,7 +443,7 @@ def test_request_shape(self) -> None: calls: list[Any] = [] c = IPFSClient("http://127.0.0.1:5001", timeout=7.0) c.session.post = _session(DATA, calls) # type: ignore[method-assign] - assert c.download_bytes("QmX") == DATA + assert c.download_bytes("QmX", verify=False) == DATA assert calls == [("http://127.0.0.1:5001/api/v0/cat", {"arg": "QmX"}, 7.0, True)] def test_0x_prefixed_expected_hash(self) -> None: @@ -453,20 +455,20 @@ def test_http_error_and_connection_error(self) -> None: c = IPFSClient("http://127.0.0.1:5001") c.session.post = _session(b"", [], status=500) # type: ignore[method-assign] with pytest.raises(IPFSError, match="HTTP 500"): - c.download_bytes("QmX") + c.download_bytes("QmX", verify=False) def boom(*a: Any, **k: Any) -> Any: raise requests.exceptions.ConnectionError("refused") c.session.post = boom # type: ignore[method-assign] with pytest.raises(IPFSError, match="connection error"): - c.download_bytes("QmX") + c.download_bytes("QmX", verify=False) def test_non_raw_cidv1_is_not_misverified(self) -> None: mh = bytes([0x12, 0x20]) + hashlib.sha256(b"other").digest() cid = "b" + base64.b32encode(bytes([0x01, 0x70]) + mh).decode().lower().rstrip("=") c = IPFSClient("http://127.0.0.1:5001") c.session.post = _session(DATA, []) # type: ignore[method-assign] - assert c.download_bytes(cid) == DATA + assert c.download_bytes(cid, verify=False) == DATA def test_manager_falls_back_and_threads_parameters(self) -> None: m = IPFSManager("http://127.0.0.1:5001", fallback_urls=["http://127.0.0.1:5002"]) @@ -481,7 +483,7 @@ def test_manager_active_client_gets_the_parameters(self) -> None: m = _single(DATA) m.active_client = m.primary with pytest.raises(IPFSError, match="exceeds max_bytes"): - m.download("QmX", max_bytes=3) + m.download("QmX", max_bytes=3, verify=False) with pytest.raises(IPFSError, match="expected_sha256"): m.download("QmX", expected_sha256="00" * 32) @@ -493,7 +495,7 @@ def test_download_from_ipfs_threads_parameters(self) -> None: with pytest.raises(IPFSError, match="expected_sha256"): download_from_ipfs("QmX", expected_sha256="00" * 32) with pytest.raises(IPFSError, match="exceeds max_bytes"): - download_from_ipfs("QmX", max_bytes=3) + download_from_ipfs("QmX", max_bytes=3, verify=False) def _single(body: bytes) -> IPFSManager: @@ -522,13 +524,13 @@ def test_active_client_receives_the_hash(self) -> None: m = IPFSManager("http://127.0.0.1:5001") m.primary.session.post = _session(DATA, calls) # type: ignore[method-assign] m.active_client = m.primary - assert m.download("QmY") == DATA + assert m.download("QmY", verify=False) == DATA assert calls[0][1] == {"arg": "QmY"} def test_download_from_ipfs_passes_the_urls(self) -> None: with mock.patch("citrate_sdk.ipfs.get_ipfs_manager", return_value=_single(DATA)) as g: from citrate_sdk.ipfs import download_from_ipfs - download_from_ipfs("QmX", ["http://127.0.0.1:5009"]) + download_from_ipfs("QmX", ["http://127.0.0.1:5009"], verify=False) g.assert_called_once_with(["http://127.0.0.1:5009"]) def test_pinned_send_verifies_when_the_flag_is_absent(self) -> None: From 66b1a8fde733e3503b5bf40efb5fef8d691092e4 Mon Sep 17 00:00:00 2001 From: SaulBuilds Date: Fri, 25 Sep 2026 08:22:27 -0700 Subject: [PATCH 4/4] docs(changelog): record the follow-up hardening under 0.6.2 0.6.2 is still unpublished, so no version bump. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_01FYQkdsk54yob6FD24jAT8P --- CHANGELOG.md | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/CHANGELOG.md b/CHANGELOG.md index 4b28627..5596e36 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -48,6 +48,17 @@ All notable changes to `citrate-labs-sdk` are documented here. This project adhe verify `expected_sha256` and self-describing CIDs. - **PBA-L6b-042:** manager writes assert `eth_chainId` against the pinned chain. +- **Verifier follow-ups (still 0.6.2, unreleased):** + - Transport-gate hardening: characters outside RFC 3986 and userinfo are + refused, and the gate checks the same host the HTTP stack connects to + (PBA-L6b-026 / PBA-L6b-027). + - IPFS downloads of CIDs that cannot verify their own content need + `expected_sha256`, or an explicit `verify=False`, which logs a warning + (PBA-L6b-030). + - `threshold_shares=1` needs `allow_single_holder_recovery=True`. + - The deploy guard also refuses values shaped like shares, not only the + known field names. + ### Changed (breaking) - `IdentityClient.refresh(refresh_token, expected_sub)`: `expected_sub` is @@ -58,6 +69,10 @@ All notable changes to `citrate-labs-sdk` are documented here. This project adhe `{"kind": "redirect" | "token", ...}` (the authority never returned the access/refresh tokens the old client expected). - `reconstruct_key_from_shares(shares, threshold)`: threshold is required. +- Endpoint URLs with credentials (`user:pass@host`) or non-RFC-3986 + characters are refused; pass credentials as headers. +- `download_bytes("Qm...")` without `expected_sha256` raises unless + `verify=False`. - Unknown `access` / `tier` / `mode` strings raise `ValueError` (PBA-L6b-031) instead of silently becoming option 0.