From a48dd521396db3a98992d71fd8b653d4dd9a4dff Mon Sep 17 00:00:00 2001 From: SaulBuilds Date: Fri, 25 Sep 2026 16:59:59 -0700 Subject: [PATCH 1/3] compat(compute): track the current compute contract ABI - post_job: jobs verified under the ZK tier (tier "ZK", or "Commitment" above 10 SALT) post the circuit's 32-byte BN254 input commitment as inputHash via the new input_commitment argument; the client checks it is a non-zero scalar below the field modulus (and the model hash is too) before sending. Other tiers are unchanged. - ComputePool exit is two steps: request_leave(), then leave_pool() once LEAVE_COOLDOWN (150 blocks) has elapsed; leave_status() reports it. join_pool() takes the required stake. - Refund claims: claim_refund / refund_owed (InferenceRouter), claim_native_refund / native_refund_owed (ComputeMarketplace), claim_requester_refund / requester_refund_pending (ComputePoolTraining). - postJob targets computeMarketplace when configured (falls back to computePool). - Env-gated end-to-end test against the contracts on a local anvil. Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_012cD3fDq5vhh2YWZPU2SV6H --- citrate_sdk/compute.py | 301 +++++++++++++++++++++++++++- tests/test_compute.py | 12 +- tests/test_pba_r2_l6b_031_enums.py | 3 +- tests/test_r2_compute_anvil.py | 152 ++++++++++++++ tests/test_r2_compute_compat.py | 310 +++++++++++++++++++++++++++++ 5 files changed, 764 insertions(+), 14 deletions(-) create mode 100644 tests/test_r2_compute_anvil.py create mode 100644 tests/test_r2_compute_compat.py diff --git a/citrate_sdk/compute.py b/citrate_sdk/compute.py index 0e13866..7f3d14c 100644 --- a/citrate_sdk/compute.py +++ b/citrate_sdk/compute.py @@ -25,7 +25,7 @@ from ._chain_guard import expected_chain_id, pinned_send from .abi import AbiInterface, enum_index, from_wei, to_wei -from .errors import ConfigurationError +from .errors import CitrateError, ConfigurationError, ValidationError from .types import ComputeJob, ComputePool, Dispute, ProviderInfo # ============================================================================ @@ -47,6 +47,26 @@ "function createPool(string name, uint8 mode, uint256 minProviders, uint256 throughput, uint256 price) payable returns (uint256 poolId)", "function joinPool(uint256 poolId, uint256 gpuCount) payable", "function leavePool(uint256 poolId)", + "function requestLeave(uint256 poolId)", + "function leaveRequestedAt(uint256 poolId, address provider) view returns (uint256)", +] + +# ComputeMarketplace: native escrow refunds that could not be pushed. +MARKETPLACE_REFUND_ABI = [ + "function nativeRefundOwed(address requester) view returns (uint256)", + "function claimNativeRefund()", +] + +# InferenceRouter: refunds credited on cancel / expiry. +INFERENCE_ROUTER_REFUND_ABI = [ + "function refundOwed(address requester) view returns (uint256)", + "function claimRefund()", +] + +# ComputePoolTraining: requester escrow refunds deferred at finalize. +TRAINING_REFUND_ABI = [ + "function requesterRefundPending(uint256 jobId) view returns (uint128)", + "function claimRequesterRefund(uint256 jobId)", ] DISPUTE_ABI = [ @@ -65,6 +85,70 @@ DISPUTE_STATES = ("Open", "ChallengerWins", "DefenderWins", "Settled") DISPUTE_OUTCOMES = ("Pending", "ChallengerWins", "DefenderWins", "Draw") +#: BN254 scalar field modulus (ComputeVerifier.BN254_SCALAR_MODULUS). +BN254_SCALAR_MODULUS = ( + 21888242871839275222246405745257275088548364400416034343698204186575808495617 +) +#: ComputeVerifier.VALUE_THRESHOLD: a Commitment-tier request whose maxPrice is +#: strictly above this is verified under the ZK tier. +ZK_AUTO_UPGRADE_THRESHOLD_WEI = 10 * 10**18 +#: ComputePool.LEAVE_COOLDOWN: blocks between requestLeave and leavePool. +LEAVE_COOLDOWN_BLOCKS = 150 +#: ComputeMarketplace.DISPUTE_WINDOW: blocks after a Valid verification before +#: completeJob is accepted. +DISPUTE_WINDOW_BLOCKS = 100 + + +class ZKCommitmentError(ValidationError): + """A ZK-tier job needs a canonical 32-byte BN254 input commitment.""" + + +class LeaveNotReadyError(CitrateError): + """leavePool would revert: the LEAVE_COOLDOWN has not elapsed yet.""" + + def __init__(self, pool_id: int, executable_at: int, current_block: int) -> None: + self.pool_id = pool_id + self.executable_at = executable_at + self.current_block = current_block + super().__init__( + f"pool {pool_id}: leave cooldown ends at block {executable_at} " + f"(current block {current_block})" + ) + + +class NothingToClaimError(CitrateError): + """The claim would revert: nothing is owed to the caller.""" + + +def effective_tier(tier: str, max_price_wei: int) -> str: + """The tier ComputeVerifier settles a job under. + + A Commitment request above ZK_AUTO_UPGRADE_THRESHOLD_WEI is verified as ZK; + every other request keeps its tier. + """ + name = JOB_TIERS[enum_index(JOB_TIERS, tier, "tier")] + if name == "Commitment" and max_price_wei > ZK_AUTO_UPGRADE_THRESHOLD_WEI: + return "ZK" + return name + + +def _canonical_field_element(value: str | bytes, what: str) -> bytes: + """Return ``value`` as 32 bytes if it is a non-zero BN254 scalar < r.""" + if isinstance(value, bytes): + raw = value + else: + text = value[2:] if value.startswith(("0x", "0X")) else value + try: + raw = bytes.fromhex(text) + except ValueError as exc: + raise ZKCommitmentError(f"{what} is not hex") from exc + if len(raw) != 32: + raise ZKCommitmentError(f"{what} must be exactly 32 bytes, got {len(raw)}") + n = int.from_bytes(raw, "big") + if n == 0 or n >= BN254_SCALAR_MODULUS: + raise ZKCommitmentError(f"{what} must be a non-zero BN254 scalar below the field modulus") + return raw + # ============================================================================ # ComputeManager @@ -102,8 +186,16 @@ def __init__( self._compute_addr = addrs.get("computePool") self._dispute_addr = addrs.get("disputeResolution") + # ComputeMarketplace hosts postJob; older configs point computePool at it. + self._marketplace_addr = addrs.get("computeMarketplace") or self._compute_addr + self._router_addr = addrs.get("inferenceRouter") + self._training_addr = addrs.get("computePoolTraining") + self._compute_iface = AbiInterface(COMPUTE_POOL_ABI) self._dispute_iface = AbiInterface(DISPUTE_ABI) + self._market_refund_iface = AbiInterface(MARKETPLACE_REFUND_ABI) + self._router_refund_iface = AbiInterface(INFERENCE_ROUTER_REFUND_ABI) + self._training_refund_iface = AbiInterface(TRAINING_REFUND_ABI) # --- Internal helpers --- @@ -121,6 +213,42 @@ def _require_dispute_address(self) -> str: ) return self._dispute_addr + def _require_marketplace_address(self) -> str: + if not self._marketplace_addr: + raise ConfigurationError( + "ComputeMarketplace contract address not configured. Set computeMarketplace " + "(or computePool) via contract_addresses." + ) + return self._marketplace_addr + + def _require_router_address(self) -> str: + if not self._router_addr: + raise ConfigurationError( + "InferenceRouter contract address not configured. Set it via contract_addresses." + ) + return self._router_addr + + def _require_training_address(self) -> str: + if not self._training_addr: + raise ConfigurationError( + "ComputePoolTraining contract address not configured. Set it via contract_addresses." + ) + return self._training_addr + + def _account(self, address: str | None) -> str: + who = address or self._default_account + if not who: + raise ConfigurationError("No address provided and no defaultAccount configured.") + return who + + def _block_number(self) -> int: + return int(cast(str, self._rpc_call("eth_blockNumber", [])), 16) + + def _read_uint(self, iface: AbiInterface, to: str, fn: str, args: list[Any]) -> int: + result = self._eth_call(to, iface.encode_function_data(fn, args)) + (value,) = iface.decode_function_result(fn, result) + return int(value) + def _eth_call(self, to: str, data: str) -> str: result = self._rpc_call("eth_call", [{"to": to, "data": data}, "latest"]) return cast(str, result) @@ -148,26 +276,52 @@ def post_job( input_data: str, max_price: str, tier: str, + *, + input_commitment: str | bytes | None = None, ) -> str: """Post a new compute job to the marketplace. - Data source: ComputePool.postJob() via eth_sendTransaction with msg.value. + Data source: ComputeMarketplace.postJob() via eth_sendTransaction with msg.value. Args: model_hash: 32-byte model hash (hex string, optional 0x prefix). - input_data: UTF-8 input payload for the compute job. + input_data: UTF-8 input payload for the compute job. Posted as + ``inputHash`` for Commitment and TEE jobs; for ZK jobs the input + is delivered off-chain and only ``input_commitment`` is posted. max_price: Maximum price in ether-denominated decimal string. tier: One of 'Commitment', 'ZK', 'TEE'. + input_commitment: Required when the job is verified under the ZK + tier (tier 'ZK', or 'Commitment' above 10 SALT): the inference + circuit's 32-byte input commitment, a non-zero BN254 scalar + below the field modulus, produced by the prover tooling. A hash + of the raw input is not accepted by the verifier. Returns: Transaction hash (job ID emitted in JobPosted event). + + Raises: + ZKCommitmentError: ZK job without a canonical commitment, a + non-canonical model hash on a ZK job, or a commitment passed + for a non-ZK job. Raised before any transaction is sent. """ - addr = self._require_compute_address() + addr = self._require_marketplace_address() h = model_hash if model_hash.startswith("0x") else f"0x{model_hash}" hash_bytes = bytes.fromhex(h[2:]) tier_num = enum_index(JOB_TIERS, tier, "tier") max_price_wei = to_wei(max_price) - input_bytes = input_data.encode("utf-8") + if effective_tier(tier, max_price_wei) == "ZK": + if input_commitment is None: + raise ZKCommitmentError( + "this job is verified under the ZK tier (tier 'ZK', or 'Commitment' above " + "10 SALT): pass input_commitment, the circuit's 32-byte BN254 input commitment" + ) + input_bytes = _canonical_field_element(input_commitment, "input_commitment") + if int.from_bytes(hash_bytes, "big") >= BN254_SCALAR_MODULUS: + raise ZKCommitmentError("model_hash must be below the BN254 field modulus for a ZK job") + elif input_commitment is not None: + raise ZKCommitmentError("input_commitment is only used for ZK-tier jobs") + else: + input_bytes = input_data.encode("utf-8") data = self._compute_iface.encode_function_data("postJob", [ hash_bytes, input_bytes, max_price_wei, tier_num, 50, 100, ]) @@ -376,31 +530,83 @@ def create_pool( ]) return self._send_transaction(addr, data) - def join_pool(self, pool_id: int, gpu_count: int) -> str: + def join_pool(self, pool_id: int, gpu_count: int, *, stake: str | None = None) -> str: """Join a compute pool with GPU allocation. - Data source: ComputePool.joinPool() via eth_sendTransaction. + Data source: ComputePool.joinPool() via eth_sendTransaction with msg.value. Args: pool_id: The pool identifier. gpu_count: Number of GPUs to allocate. + stake: SALT staked with the membership, as a decimal string. The + pool requires at least ``gpu_count * MIN_STAKE_PER_GPU`` + (10 SALT per GPU); without it the join reverts. Returns: Transaction hash. """ addr = self._require_compute_address() data = self._compute_iface.encode_function_data("joinPool", [pool_id, gpu_count]) + value = hex(to_wei(stake)) if stake is not None else "0x0" + return self._send_transaction(addr, data, value) + + def request_leave(self, pool_id: int) -> str: + """Start leaving a compute pool (step 1 of 2). + + Data source: ComputePool.requestLeave() via eth_sendTransaction. The + member stays active and slashable until ``leavePool`` completes the + exit, which is accepted LEAVE_COOLDOWN_BLOCKS blocks later. + + Returns: + Transaction hash. + """ + addr = self._require_compute_address() + data = self._compute_iface.encode_function_data("requestLeave", [pool_id]) return self._send_transaction(addr, data) + def leave_status(self, pool_id: int, address: str | None = None) -> dict[str, Any]: + """Pending-exit state for ``address`` (default account) in ``pool_id``. + + Data source: ComputePool.leaveRequestedAt(uint256,address) + eth_blockNumber. + + Returns: + ``{"requested_at": int, "executable_at": int | None, "ready": bool}``; + ``requested_at`` is 0 when no exit is pending. + """ + addr = self._require_compute_address() + who = self._account(address) + requested_at = self._read_uint( + self._compute_iface, addr, "leaveRequestedAt", [pool_id, who] + ) + if requested_at == 0: + return {"requested_at": 0, "executable_at": None, "ready": False} + executable_at = requested_at + LEAVE_COOLDOWN_BLOCKS + return { + "requested_at": requested_at, + "executable_at": executable_at, + "ready": self._block_number() >= executable_at, + } + def leave_pool(self, pool_id: int) -> str: - """Leave a compute pool and reclaim stake. + """Leave a compute pool and reclaim stake (two-step exit). - Data source: ComputePool.leavePool() via eth_sendTransaction. + Data source: ComputePool.leaveRequestedAt / requestLeave / leavePool. + + * No exit pending: sends ``requestLeave`` (step 1) and returns its hash. + Call again once the cooldown has elapsed. + * Exit pending and LEAVE_COOLDOWN_BLOCKS elapsed: sends ``leavePool``. + * Exit pending inside the cooldown: raises LeaveNotReadyError (no + transaction is sent). Returns: Transaction hash. """ addr = self._require_compute_address() + status = self.leave_status(pool_id) + if status["requested_at"] == 0: + return self.request_leave(pool_id) + if not status["ready"]: + raise LeaveNotReadyError(pool_id, status["executable_at"], self._block_number()) data = self._compute_iface.encode_function_data("leavePool", [pool_id]) return self._send_transaction(addr, data) @@ -455,6 +661,83 @@ def get_pools(self) -> list[ComputePool]: return pools + # ------------------------------------------------------------------- + # Refund claims + # ------------------------------------------------------------------- + + def refund_owed(self, address: str | None = None) -> int: + """Wei credited to ``address`` (default account) by InferenceRouter. + + Data source: InferenceRouter.refundOwed(address) via eth_call. + """ + return self._read_uint( + self._router_refund_iface, self._require_router_address(), "refundOwed", + [self._account(address)], + ) + + def claim_refund(self) -> str: + """Withdraw InferenceRouter refunds credited to the default account. + + Data source: InferenceRouter.claimRefund() via eth_sendTransaction. + + Raises: + NothingToClaimError: nothing is owed (the call would revert). + """ + addr = self._require_router_address() + if self.refund_owed() == 0: + raise NothingToClaimError("InferenceRouter: no refund owed") + return self._send_transaction(addr, self._router_refund_iface.encode_function_data("claimRefund")) + + def native_refund_owed(self, address: str | None = None) -> int: + """Wei of job escrow refunds ComputeMarketplace could not push to ``address``. + + Data source: ComputeMarketplace.nativeRefundOwed(address) via eth_call. + """ + return self._read_uint( + self._market_refund_iface, self._require_marketplace_address(), "nativeRefundOwed", + [self._account(address)], + ) + + def claim_native_refund(self) -> str: + """Withdraw ComputeMarketplace escrow refunds owed to the default account. + + Data source: ComputeMarketplace.claimNativeRefund() via eth_sendTransaction. + + Raises: + NothingToClaimError: nothing is owed (the call would revert). + """ + addr = self._require_marketplace_address() + if self.native_refund_owed() == 0: + raise NothingToClaimError("ComputeMarketplace: no native refund owed") + return self._send_transaction( + addr, self._market_refund_iface.encode_function_data("claimNativeRefund") + ) + + def requester_refund_pending(self, job_id: int) -> int: + """Wei of a training job's requester escrow refund deferred at finalize. + + Data source: ComputePoolTraining.requesterRefundPending(uint256) via eth_call. + """ + return self._read_uint( + self._training_refund_iface, self._require_training_address(), + "requesterRefundPending", [job_id], + ) + + def claim_requester_refund(self, job_id: int) -> str: + """Claim a training job's deferred requester refund (requester only). + + Data source: ComputePoolTraining.claimRequesterRefund(uint256) via eth_sendTransaction. + + Raises: + NothingToClaimError: nothing is pending for the job (the call would revert). + """ + addr = self._require_training_address() + if self.requester_refund_pending(job_id) == 0: + raise NothingToClaimError(f"ComputePoolTraining: no refund pending for job {job_id}") + return self._send_transaction( + addr, self._training_refund_iface.encode_function_data("claimRequesterRefund", [job_id]) + ) + # ------------------------------------------------------------------- # Disputes # ------------------------------------------------------------------- diff --git a/tests/test_compute.py b/tests/test_compute.py index 621be27..93db977 100644 --- a/tests/test_compute.py +++ b/tests/test_compute.py @@ -70,14 +70,18 @@ def test_post_job_calldata(self): """post_job sends correct calldata, value, and contract address.""" rpc = chain_rpc("0xtx") mgr = self._make_manager(rpc) - model_hash = "0x" + "ab" * 32 - mgr.post_job(model_hash, "hello world", "1.5", "ZK") + model_hash = "0x" + "0b" * 32 # ZK jobs need a model hash below the BN254 modulus + # ZK tier: inputHash is the circuit's 32-byte BN254 input commitment. + commitment = "0x" + (42).to_bytes(32, "big").hex() + mgr.post_job(model_hash, "hello world", "1.5", "ZK", input_commitment=commitment) tx = rpc.call_args_list[-1][0][1][0] assert tx["to"] == FAKE_COMPUTE_ADDR assert int(tx["value"], 16) == to_wei("1.5") # Verify function selector expected_selector = keccak256(b"postJob(bytes32,bytes,uint256,uint8,uint256,uint256)")[:4].hex() assert tx["data"][2:10] == expected_selector + assert commitment[2:] in tx["data"] + assert b"hello world".hex() not in tx["data"] def test_bid_on_job_calldata(self): """bid_on_job sends correct calldata.""" @@ -150,9 +154,9 @@ def test_leave_pool_calldata(self): """leave_pool sends correct calldata.""" rpc = chain_rpc("0xtx") mgr = self._make_manager(rpc) - mgr.leave_pool(pool_id=9) + mgr.request_leave(pool_id=9) tx = rpc.call_args_list[-1][0][1][0] - expected = _compute_iface.encode_function_data("leavePool", [9]) + expected = "0x" + keccak256(b"requestLeave(uint256)")[:4].hex() + (9).to_bytes(32, "big").hex() assert tx["data"] == expected def test_dispute_result_calldata(self): diff --git a/tests/test_pba_r2_l6b_031_enums.py b/tests/test_pba_r2_l6b_031_enums.py index 12e9797..a07b0f2 100644 --- a/tests/test_pba_r2_l6b_031_enums.py +++ b/tests/test_pba_r2_l6b_031_enums.py @@ -59,7 +59,8 @@ def test_unknown_mode_raises() -> None: def test_known_tiers_encode_their_index(tier: str, index: int) -> None: sent: list[Any] = [] mgr = ComputeManager(_rpc(sent), default_account=ACCT, contract_addresses=ADDRS) - mgr.post_job("0x" + "aa" * 32, "in", "1", tier) + commitment = "0x" + (1).to_bytes(32, "big").hex() if index == 1 else None + mgr.post_job("0x" + "0a" * 32, "in", "1", tier, input_commitment=commitment) data = bytes.fromhex(sent[-1][1][0]["data"][2:]) # postJob(bytes32,bytes,uint256,uint8,...): tier is the 4th head word. assert int.from_bytes(data[4 + 96: 4 + 128], "big") == index diff --git a/tests/test_r2_compute_anvil.py b/tests/test_r2_compute_anvil.py new file mode 100644 index 0000000..e7dc7cd --- /dev/null +++ b/tests/test_r2_compute_anvil.py @@ -0,0 +1,152 @@ +"""End-to-end compute compat against the chain-main contracts on a local anvil. + +Skipped unless ``CITRATE_R2_ANVIL_ADDRS`` points at the addresses JSON written +by the chain-main deploy (keys: ComputeMarketplace, ComputePool, +ComputePoolTraining, InferenceRouter) and ``CITRATE_R2_ANVIL_RPC`` (default +http://127.0.0.1:8599) serves a chain whose id is 40204 with anvil's unlocked +default accounts. Uses default accounts 3 and 4 only. +""" +from __future__ import annotations + +import json +import os +import time +import urllib.request +from typing import Any, cast + +import pytest + +from citrate_sdk.abi import keccak256 +from citrate_sdk.compute import ( + LEAVE_COOLDOWN_BLOCKS, + ComputeManager, + LeaveNotReadyError, + NothingToClaimError, + ZKCommitmentError, +) + +ADDRS_PATH = os.environ.get("CITRATE_R2_ANVIL_ADDRS") +RPC_URL = os.environ.get("CITRATE_R2_ANVIL_RPC", "http://127.0.0.1:8599") + +pytestmark = [ + pytest.mark.integration, + pytest.mark.skipif(not ADDRS_PATH, reason="CITRATE_R2_ANVIL_ADDRS not set"), +] + +ACCT3 = "0x90F79bf6EB2c4f870365E785982E1f101E93b906" +ACCT4 = "0x15d34AAf54267DB7D7c367839AAf71A00a2C6A65" +MODEL = "0x" + (0xC0FFEE).to_bytes(32, "big").hex() +COMMIT = "0x" + (0x1234567).to_bytes(32, "big").hex() + + +def rpc(method: str, params: list[Any]) -> Any: + body = json.dumps({"jsonrpc": "2.0", "id": 1, "method": method, "params": params}).encode() + req = urllib.request.Request(RPC_URL, body, {"Content-Type": "application/json"}) + with urllib.request.urlopen(req, timeout=30) as r: # noqa: S310 - local anvil only + out = json.loads(r.read()) + if "error" in out: + raise RuntimeError(out["error"]) + return out["result"] + + +def receipt(tx: str) -> dict[str, Any]: + # The anvil is shared with other lanes; poll briefly for the receipt. + for _ in range(100): + r = rpc("eth_getTransactionReceipt", [tx]) + if r is not None: + return cast(dict[str, Any], r) + time.sleep(0.2) + raise AssertionError(f"no receipt for {tx}") + + +def addrs() -> dict[str, str]: + assert ADDRS_PATH + with open(ADDRS_PATH) as f: + return cast(dict[str, str], json.load(f)) + + +def mgr(account: str) -> ComputeManager: + a = addrs() + return ComputeManager( + rpc, + default_account=account, + gas_limit=3_000_000, + contract_addresses={ + "computePool": a["ComputePool"], + "computeMarketplace": a["ComputeMarketplace"], + "inferenceRouter": a["InferenceRouter"], + "computePoolTraining": a["ComputePoolTraining"], + }, + chain_id=40204, + ) + + +def job_input_hash(job_id: int) -> str: + a = addrs() + sel = keccak256(b"jobBinding(uint256)")[:4].hex() + data = "0x" + sel + job_id.to_bytes(32, "big").hex() + out = str(rpc("eth_call", [{"to": a["ComputeVerifier"], "data": data}, "latest"])) + return "0x" + out[2:66] + + +def next_job_id() -> int: + a = addrs() + sel = keccak256(b"nextJobId()")[:4].hex() + return int(rpc("eth_call", [{"to": a["ComputeMarketplace"], "data": "0x" + sel}, "latest"]), 16) + + +@pytest.mark.parametrize(("tier", "price"), [("ZK", "1"), ("Commitment", "10.5")]) +def test_zk_and_auto_zk_jobs_post_with_bn254_commitment(tier: str, price: str) -> None: + jid = next_job_id() + tx = mgr(ACCT3).post_job(MODEL, "hello world", price, tier, input_commitment=COMMIT) + assert int(receipt(tx)["status"], 16) == 1 + # The verifier bound exactly the commitment as the job's input commitment. + assert job_input_hash(jid) == COMMIT + + +def test_zk_job_without_commitment_never_reaches_chain() -> None: + with pytest.raises(ZKCommitmentError): + mgr(ACCT3).post_job(MODEL, "hello world", "1", "ZK") + + +def test_commitment_tier_job_still_posts_raw_input() -> None: + tx = mgr(ACCT3).post_job(MODEL, "hello world", "1", "Commitment") + assert int(receipt(tx)["status"], 16) == 1 + + +def test_two_step_pool_exit() -> None: + m = mgr(ACCT4) + tx = m.create_pool("r2-compat", "InferencePool", 1, 1, "0.01") + assert int(receipt(tx)["status"], 16) == 1 + pool_id = _last_pool_id() + assert int(receipt(m.join_pool(pool_id, 1, stake="10"))["status"], 16) == 1 + + # Step 1: leave_pool with nothing pending sends requestLeave. + assert int(receipt(m.leave_pool(pool_id))["status"], 16) == 1 + st = m.leave_status(pool_id) + assert st["requested_at"] > 0 and st["ready"] is False + with pytest.raises(LeaveNotReadyError): + m.leave_pool(pool_id) + + rpc("anvil_mine", [hex(LEAVE_COOLDOWN_BLOCKS)]) + assert m.leave_status(pool_id)["ready"] is True + before = int(rpc("eth_getBalance", [ACCT4, "latest"]), 16) + assert int(receipt(m.leave_pool(pool_id))["status"], 16) == 1 + assert m.leave_status(pool_id)["requested_at"] == 0 + assert int(rpc("eth_getBalance", [ACCT4, "latest"]), 16) > before # stake returned + + +def _last_pool_id() -> int: + a = addrs() + sel = keccak256(b"nextPoolId()")[:4].hex() + return int(rpc("eth_call", [{"to": a["ComputePool"], "data": "0x" + sel}, "latest"]), 16) - 1 + + +def test_refund_views_and_claims_match_deployed_abi() -> None: + m = mgr(ACCT3) + assert m.refund_owed() == 0 + assert m.native_refund_owed() == 0 + assert m.requester_refund_pending(0) == 0 + for claim, args in (("claim_refund", ()), ("claim_native_refund", ()), ("claim_requester_refund", (0,))): + with pytest.raises(NothingToClaimError): + getattr(m, claim)(*args) diff --git a/tests/test_r2_compute_compat.py b/tests/test_r2_compute_compat.py new file mode 100644 index 0000000..af9b6c4 --- /dev/null +++ b/tests/test_r2_compute_compat.py @@ -0,0 +1,310 @@ +"""Compute marketplace client compatibility with the chain-main contracts. + +Covers the client side of the current ComputeMarketplace / ComputeVerifier / +ComputePool / ComputePoolTraining / InferenceRouter ABI: + +* ZK-tier jobs (requested ZK, or Commitment above 10 SALT, which the verifier + settles as ZK) carry a 32-byte canonical BN254 input commitment as + ``inputHash``; the client refuses to post anything else. +* Leaving a ComputePool is two steps: ``requestLeave`` then ``leavePool`` + after ``LEAVE_COOLDOWN`` blocks. +* Refund claims: ``InferenceRouter.claimRefund``, + ``ComputeMarketplace.claimNativeRefund``, + ``ComputePoolTraining.claimRequesterRefund``. +""" +from __future__ import annotations + +from typing import Any + +import pytest + +from citrate_sdk._generated import contract as _contract +from citrate_sdk.abi import AbiInterface, keccak256, to_wei +from citrate_sdk.compute import ( + BN254_SCALAR_MODULUS, + COMPUTE_POOL_ABI, + DISPUTE_WINDOW_BLOCKS, + LEAVE_COOLDOWN_BLOCKS, + ZK_AUTO_UPGRADE_THRESHOLD_WEI, + ComputeManager, + LeaveNotReadyError, + NothingToClaimError, + ZKCommitmentError, + effective_tier, +) + +POOL = "0x" + "aa" * 20 +MARKET = "0x" + "ab" * 20 +ROUTER = "0x" + "cc" * 20 +TRAINING = "0x" + "dd" * 20 +ACCT = "0x" + "11" * 20 +MODEL = "0x" + "0a" * 32 +COMMIT = "0x" + (12345).to_bytes(32, "big").hex() + +_iface = AbiInterface(COMPUTE_POOL_ABI) + + +def _word(v: int) -> str: + return "0x" + v.to_bytes(32, "big").hex() + + +class Rpc: + """eth_chainId + eth_blockNumber + per-selector eth_call answers; records sends.""" + + def __init__(self, calls: dict[str, int] | None = None, block: int = 1_000) -> None: + self.calls = calls or {} + self.block = block + self.sent: list[dict[str, Any]] = [] + + def __call__(self, method: str, params: list[Any]) -> Any: + if method == "eth_chainId": + return hex(_contract.chain_id()) + if method == "eth_blockNumber": + return hex(self.block) + if method == "eth_call": + sel = params[0]["data"][2:10] + return _word(self.calls.get(sel, 0)) + if method == "eth_sendTransaction": + self.sent.append(params[0]) + return "0x" + "ee" * 32 + raise AssertionError(f"unexpected rpc {method}") + + +def _sel(sig: str) -> str: + return keccak256(sig.encode())[:4].hex() + + +def _mgr(rpc: Rpc) -> ComputeManager: + return ComputeManager( + rpc, + default_account=ACCT, + contract_addresses={ + "computePool": POOL, + "computeMarketplace": MARKET, + "inferenceRouter": ROUTER, + "computePoolTraining": TRAINING, + }, + ) + + +def _posted_input_hash(tx: dict[str, Any]) -> bytes: + data = bytes.fromhex(tx["data"][2:])[4:] + off = int.from_bytes(data[32:64], "big") + ln = int.from_bytes(data[off: off + 32], "big") + return data[off + 32: off + 32 + ln] + + +# -------------------------------------------------------------------------- +# constants mirror the contracts +# -------------------------------------------------------------------------- + +def test_constants_match_chain_main() -> None: + assert BN254_SCALAR_MODULUS == int( + "21888242871839275222246405745257275088548364400416034343698204186575808495617" + ) + assert ZK_AUTO_UPGRADE_THRESHOLD_WEI == 10 * 10**18 + assert LEAVE_COOLDOWN_BLOCKS == 150 + assert DISPUTE_WINDOW_BLOCKS == 100 + + +@pytest.mark.parametrize( + ("tier", "price", "expected"), + [ + ("Commitment", "10", "Commitment"), + ("Commitment", "10.000000000000000001", "ZK"), + ("Commitment", "1", "Commitment"), + ("ZK", "1", "ZK"), + ("TEE", "50", "TEE"), + ], +) +def test_effective_tier(tier: str, price: str, expected: str) -> None: + assert effective_tier(tier, to_wei(price)) == expected + + +# -------------------------------------------------------------------------- +# post_job: BN254 input commitment for ZK / auto-ZK +# -------------------------------------------------------------------------- + +def test_zk_job_posts_the_commitment_not_the_input() -> None: + rpc = Rpc() + _mgr(rpc).post_job(MODEL, "hello world", "1", "ZK", input_commitment=COMMIT) + assert _posted_input_hash(rpc.sent[-1]) == bytes.fromhex(COMMIT[2:]) + assert rpc.sent[-1]["to"] == MARKET + + +def test_auto_upgraded_job_requires_commitment() -> None: + rpc = Rpc() + with pytest.raises(ZKCommitmentError, match="input_commitment"): + _mgr(rpc).post_job(MODEL, "hello world", "10.5", "Commitment") + assert rpc.sent == [] + _mgr(rpc).post_job(MODEL, "hello world", "10.5", "Commitment", input_commitment=COMMIT) + assert _posted_input_hash(rpc.sent[-1]) == bytes.fromhex(COMMIT[2:]) + + +def test_zk_job_without_commitment_is_refused_before_sending() -> None: + rpc = Rpc() + with pytest.raises(ZKCommitmentError): + _mgr(rpc).post_job(MODEL, "hello world", "1", "ZK") + assert rpc.sent == [] + + +@pytest.mark.parametrize( + "bad", + [ + "0x" + "00" * 32, # zero + "0x" + BN254_SCALAR_MODULUS.to_bytes(32, "big").hex(), # == r + "0x" + "ff" * 32, # > r + "0x" + "01" * 31, # 31 bytes + "0x" + "01" * 33, # 33 bytes + "0xzz" + "01" * 31, # not hex + ], +) +def test_non_canonical_commitment_is_refused(bad: str) -> None: + rpc = Rpc() + with pytest.raises(ZKCommitmentError): + _mgr(rpc).post_job(MODEL, "x", "1", "ZK", input_commitment=bad) + assert rpc.sent == [] + + +def test_largest_canonical_commitment_is_accepted() -> None: + rpc = Rpc() + top = "0x" + (BN254_SCALAR_MODULUS - 1).to_bytes(32, "big").hex() + _mgr(rpc).post_job(MODEL, "x", "1", "ZK", input_commitment=top) + assert _posted_input_hash(rpc.sent[-1]) == bytes.fromhex(top[2:]) + + +def test_commitment_as_bytes_is_accepted() -> None: + rpc = Rpc() + _mgr(rpc).post_job(MODEL, "x", "1", "ZK", input_commitment=(7).to_bytes(32, "big")) + assert _posted_input_hash(rpc.sent[-1]) == (7).to_bytes(32, "big") + + +def test_zk_job_model_hash_must_be_canonical() -> None: + rpc = Rpc() + with pytest.raises(ZKCommitmentError, match="model_hash"): + _mgr(rpc).post_job("0x" + "ff" * 32, "x", "1", "ZK", input_commitment=COMMIT) + assert rpc.sent == [] + + +def test_commitment_on_non_zk_job_is_refused() -> None: + rpc = Rpc() + with pytest.raises(ZKCommitmentError, match="ZK"): + _mgr(rpc).post_job(MODEL, "x", "1", "Commitment", input_commitment=COMMIT) + assert rpc.sent == [] + + +def test_commitment_tier_job_keeps_input_bytes() -> None: + rpc = Rpc() + _mgr(rpc).post_job(MODEL, "hello world", "10", "Commitment") + assert _posted_input_hash(rpc.sent[-1]) == b"hello world" + + +def test_marketplace_address_falls_back_to_compute_pool() -> None: + rpc = Rpc() + ComputeManager(rpc, default_account=ACCT, contract_addresses={"computePool": POOL}).post_job( + MODEL, "x", "1", "Commitment" + ) + assert rpc.sent[-1]["to"] == POOL + + +# -------------------------------------------------------------------------- +# leave: requestLeave → cooldown → leavePool +# -------------------------------------------------------------------------- + +LEAVE_REQ = _sel("leaveRequestedAt(uint256,address)") + + +def test_request_leave_calldata() -> None: + rpc = Rpc() + _mgr(rpc).request_leave(9) + assert rpc.sent[-1]["data"][2:10] == _sel("requestLeave(uint256)") + assert rpc.sent[-1]["to"] == POOL + + +def test_leave_pool_without_request_sends_request_leave() -> None: + rpc = Rpc({LEAVE_REQ: 0}) + _mgr(rpc).leave_pool(9) + assert rpc.sent[-1]["data"][2:10] == _sel("requestLeave(uint256)") + + +def test_leave_pool_inside_cooldown_raises_with_ready_block() -> None: + rpc = Rpc({LEAVE_REQ: 900}, block=1_049) + with pytest.raises(LeaveNotReadyError) as ei: + _mgr(rpc).leave_pool(9) + assert ei.value.executable_at == 1_050 + assert rpc.sent == [] + + +def test_leave_pool_after_cooldown_completes() -> None: + for head in (1_050, 1_051): + rpc = Rpc({LEAVE_REQ: 900}, block=head) + _mgr(rpc).leave_pool(9) + assert rpc.sent[-1]["data"] == _iface.encode_function_data("leavePool", [9]) + + +def test_leave_status() -> None: + rpc = Rpc({LEAVE_REQ: 900}, block=1_000) + st = _mgr(rpc).leave_status(9) + assert st == {"requested_at": 900, "executable_at": 1_050, "ready": False} + rpc.block = 1_050 + assert _mgr(rpc).leave_status(9)["ready"] is True + assert _mgr(Rpc({LEAVE_REQ: 0})).leave_status(9) == { + "requested_at": 0, "executable_at": None, "ready": False, + } + + +# -------------------------------------------------------------------------- +# refund claims +# -------------------------------------------------------------------------- + +@pytest.mark.parametrize( + ("claim", "view", "view_sig", "claim_sig", "to", "args"), + [ + ("claim_refund", "refund_owed", "refundOwed(address)", "claimRefund()", ROUTER, ()), + ("claim_native_refund", "native_refund_owed", "nativeRefundOwed(address)", + "claimNativeRefund()", MARKET, ()), + ("claim_requester_refund", "requester_refund_pending", "requesterRefundPending(uint256)", + "claimRequesterRefund(uint256)", TRAINING, (4,)), + ], +) +def test_refund_claims(claim: str, view: str, view_sig: str, claim_sig: str, to: str, + args: tuple[int, ...]) -> None: + rpc = Rpc({_sel(view_sig): 5}) + mgr = _mgr(rpc) + assert getattr(mgr, view)(*args) == 5 + getattr(mgr, claim)(*args) + tx = rpc.sent[-1] + assert tx["to"] == to + assert tx["data"][2:10] == _sel(claim_sig) + if args: + assert int(tx["data"][10:74], 16) == args[0] + # Nothing owed → refuse before sending a transaction that would revert. + rpc2 = Rpc({_sel(view_sig): 0}) + with pytest.raises(NothingToClaimError): + getattr(_mgr(rpc2), claim)(*args) + assert rpc2.sent == [] + + +def test_refund_views_default_to_the_configured_account() -> None: + seen: list[str] = [] + + class Spy(Rpc): + def __call__(self, method: str, params: list[Any]) -> Any: + if method == "eth_call": + seen.append(params[0]["data"]) + return super().__call__(method, params) + + _mgr(Spy()).refund_owed() + assert seen[-1][-40:] == ACCT[2:] + _mgr(Spy()).native_refund_owed("0x" + "22" * 20) + assert seen[-1][-40:] == "22" * 20 + + +def test_missing_refund_contract_addresses_raise() -> None: + from citrate_sdk.errors import ConfigurationError + + mgr = ComputeManager(Rpc(), default_account=ACCT, contract_addresses={}) + with pytest.raises(ConfigurationError, match="InferenceRouter"): + mgr.claim_refund() + with pytest.raises(ConfigurationError, match="ComputePoolTraining"): + mgr.claim_requester_refund(1) From 474e3690092a6e645355b059770158f513907b97 Mon Sep 17 00:00:00 2001 From: SaulBuilds Date: Fri, 25 Sep 2026 17:34:56 -0700 Subject: [PATCH 2/3] test(compat): argument and address plumbing cases for compute compat Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_012cD3fDq5vhh2YWZPU2SV6H --- tests/test_r2_compute_compat.py | 135 +++++++++++++++++++++++++++++++- 1 file changed, 132 insertions(+), 3 deletions(-) diff --git a/tests/test_r2_compute_compat.py b/tests/test_r2_compute_compat.py index af9b6c4..904a6cf 100644 --- a/tests/test_r2_compute_compat.py +++ b/tests/test_r2_compute_compat.py @@ -48,21 +48,39 @@ def _word(v: int) -> str: return "0x" + v.to_bytes(32, "big").hex() +# Which contract answers each read selector. A read sent anywhere else fails. +_READ_HOME = { + "leaveRequestedAt(uint256,address)": POOL, + "refundOwed(address)": ROUTER, + "nativeRefundOwed(address)": MARKET, + "requesterRefundPending(uint256)": TRAINING, +} + + class Rpc: """eth_chainId + eth_blockNumber + per-selector eth_call answers; records sends.""" - def __init__(self, calls: dict[str, int] | None = None, block: int = 1_000) -> None: + def __init__(self, calls: dict[str, int] | None = None, block: int = 1_000, + chain_id: int | None = None) -> None: self.calls = calls or {} self.block = block + self.chain_id = chain_id if chain_id is not None else _contract.chain_id() self.sent: list[dict[str, Any]] = [] + self.reads: list[dict[str, Any]] = [] + self.homes = {_sel(k): v for k, v in _READ_HOME.items()} def __call__(self, method: str, params: list[Any]) -> Any: + assert isinstance(params, list), f"{method}: params must be a JSON array" if method == "eth_chainId": - return hex(_contract.chain_id()) + return hex(self.chain_id) if method == "eth_blockNumber": + assert params == [] return hex(self.block) if method == "eth_call": - sel = params[0]["data"][2:10] + call = params[0] + sel = call["data"][2:10] + assert call["to"] == self.homes[sel], f"read {sel} sent to {call['to']}" + self.reads.append(call) return _word(self.calls.get(sel, 0)) if method == "eth_sendTransaction": self.sent.append(params[0]) @@ -308,3 +326,114 @@ def test_missing_refund_contract_addresses_raise() -> None: mgr.claim_refund() with pytest.raises(ConfigurationError, match="ComputePoolTraining"): mgr.claim_requester_refund(1) + + +# -------------------------------------------------------------------------- +# argument / address plumbing +# -------------------------------------------------------------------------- + +def _head_words(tx: dict[str, Any]) -> list[int]: + data = bytes.fromhex(tx["data"][2:])[4:] + return [int.from_bytes(data[i: i + 32], "big") for i in range(0, len(data), 32)] + + +def test_post_job_windows_and_unprefixed_model_hash() -> None: + rpc = Rpc() + _mgr(rpc).post_job(MODEL[2:], "x", "1", "Commitment") + w = _head_words(rpc.sent[-1]) + assert w[0] == int(MODEL, 16) # model hash, no 0x prefix given + assert w[3] == 0 # tier + assert (w[4], w[5]) == (50, 100) # bid / exec windows + + +def test_post_job_default_gas_and_price() -> None: + rpc = Rpc() + _mgr(rpc).post_job(MODEL, "x", "1", "Commitment") + assert rpc.sent[-1]["gas"] == hex(500_000) + assert rpc.sent[-1]["gasPrice"] == "0x3b9aca00" + + +def test_writes_refuse_a_foreign_chain() -> None: + from citrate_sdk.errors import CitrateError + + rpc = Rpc(chain_id=1) + with pytest.raises(CitrateError): + _mgr(rpc).post_job(MODEL, "x", "1", "Commitment") + assert rpc.sent == [] + + +def test_explicit_chain_id_is_enforced() -> None: + from citrate_sdk.errors import CitrateError + + rpc = Rpc(chain_id=_contract.chain_id()) + mgr = ComputeManager(rpc, default_account=ACCT, contract_addresses={"computePool": POOL}, + chain_id=31337) + with pytest.raises(CitrateError): + mgr.request_leave(1) + + +def test_zk_model_hash_equal_to_modulus_is_refused() -> None: + rpc = Rpc() + at_r = "0x" + BN254_SCALAR_MODULUS.to_bytes(32, "big").hex() + with pytest.raises(ZKCommitmentError, match="model_hash"): + _mgr(rpc).post_job(at_r, "x", "1", "ZK", input_commitment=COMMIT) + below = "0x" + (BN254_SCALAR_MODULUS - 1).to_bytes(32, "big").hex() + _mgr(rpc).post_job(below, "x", "1", "ZK", input_commitment=COMMIT) + assert len(rpc.sent) == 1 + + +@pytest.mark.parametrize("form", ["unprefixed", "upper"]) +def test_commitment_hex_forms(form: str) -> None: + rpc = Rpc() + body = (99).to_bytes(32, "big").hex() + value = body if form == "unprefixed" else "0X" + body + _mgr(rpc).post_job(MODEL, "x", "1", "ZK", input_commitment=value) + assert _posted_input_hash(rpc.sent[-1]) == (99).to_bytes(32, "big") + + +def test_join_pool_sends_stake_to_the_pool() -> None: + rpc = Rpc() + _mgr(rpc).join_pool(3, 2, stake="20") + tx = rpc.sent[-1] + assert tx["to"] == POOL + assert int(tx["value"], 16) == to_wei("20") + assert tx["data"] == _iface.encode_function_data("joinPool", [3, 2]) + _mgr(rpc).join_pool(3, 2) + assert rpc.sent[-1]["value"] == "0x0" + + +def test_leave_pool_targets_the_pool_and_reports_blocks() -> None: + rpc = Rpc({LEAVE_REQ: 900}, block=1_060) + _mgr(rpc).leave_pool(9) + assert rpc.sent[-1]["to"] == POOL + rpc = Rpc({LEAVE_REQ: 900}, block=1_010) + with pytest.raises(LeaveNotReadyError) as ei: + _mgr(rpc).leave_pool(9) + assert (ei.value.pool_id, ei.value.executable_at, ei.value.current_block) == (9, 1_050, 1_010) + assert "1050" in str(ei.value) and "1010" in str(ei.value) + + +def test_views_use_the_given_address() -> None: + other = "0x" + "33" * 20 + rpc = Rpc({LEAVE_REQ: 0}) + _mgr(rpc).leave_status(9, other) + assert rpc.reads[-1]["data"].endswith(other[2:]) + _mgr(rpc).refund_owed(other) + assert rpc.reads[-1]["data"].endswith(other[2:]) + _mgr(rpc).requester_refund_pending(12) + assert int(rpc.reads[-1]["data"][10:], 16) == 12 + + +@pytest.mark.parametrize( + ("claim", "args", "needle"), + [("claim_refund", (), "InferenceRouter"), ("claim_native_refund", (), "ComputeMarketplace"), + ("claim_requester_refund", (7,), "job 7")], +) +def test_nothing_to_claim_names_the_source(claim: str, args: tuple[int, ...], needle: str) -> None: + with pytest.raises(NothingToClaimError, match=needle): + getattr(_mgr(Rpc()), claim)(*args) + + +def test_effective_tier_rejects_unknown_tier() -> None: + with pytest.raises(ValueError, match="tier"): + effective_tier("zkp", 0) From c75d0e9e6eeeaddb6d553666aa9c07fc12a4e222 Mon Sep 17 00:00:00 2001 From: SaulBuilds Date: Fri, 25 Sep 2026 18:28:40 -0700 Subject: [PATCH 3/3] test(compat): pin the ZK-commitment error messages Co-Authored-By: Claude Opus 5.5 (1M context) Claude-Session: https://claude.ai/code/session_012cD3fDq5vhh2YWZPU2SV6H --- tests/test_r2_compute_compat.py | 37 +++++++++++++++++++++++++++++++++ 1 file changed, 37 insertions(+) diff --git a/tests/test_r2_compute_compat.py b/tests/test_r2_compute_compat.py index 904a6cf..d3d87f2 100644 --- a/tests/test_r2_compute_compat.py +++ b/tests/test_r2_compute_compat.py @@ -437,3 +437,40 @@ def test_nothing_to_claim_names_the_source(claim: str, args: tuple[int, ...], ne def test_effective_tier_rejects_unknown_tier() -> None: with pytest.raises(ValueError, match="tier"): effective_tier("zkp", 0) + + +# -------------------------------------------------------------------------- +# user-facing error messages (the caller acts on these) +# -------------------------------------------------------------------------- + +@pytest.mark.parametrize( + ("value", "needle"), + [ + ("0xzz", "input_commitment is not hex"), + ("0x" + "01" * 31, "input_commitment must be exactly 32 bytes, got 31"), + ("0x" + "00" * 32, "input_commitment must be a non-zero BN254 scalar"), + ], +) +def test_bad_commitment_messages_name_the_argument(value: str, needle: str) -> None: + with pytest.raises(ZKCommitmentError) as ei: + _mgr(Rpc()).post_job(MODEL, "x", "1", "ZK", input_commitment=value) + assert needle in str(ei.value) + + +def test_missing_commitment_message_explains_the_zk_rule() -> None: + with pytest.raises(ZKCommitmentError) as ei: + _mgr(Rpc()).post_job(MODEL, "x", "11", "Commitment") + msg = str(ei.value) + assert "ZK tier" in msg and "'Commitment' above 10 SALT" in msg + assert "input_commitment" in msg and "BN254" in msg + + +def test_commitment_on_non_zk_job_message() -> None: + with pytest.raises(ZKCommitmentError, match="only used for ZK-tier jobs"): + _mgr(Rpc()).post_job(MODEL, "x", "1", "Commitment", input_commitment=COMMIT) + + +def test_zk_model_hash_message() -> None: + at_r = "0x" + BN254_SCALAR_MODULUS.to_bytes(32, "big").hex() + with pytest.raises(ZKCommitmentError, match="model_hash must be below the BN254 field modulus"): + _mgr(Rpc()).post_job(at_r, "x", "1", "ZK", input_commitment=COMMIT)