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. 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/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/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/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_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_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/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 == [] 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: 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"]