personal memory agent
0

Configure Feed

Select the types of activity you want to include in your feed.

solstone / scripts / spp_ratls_loopback_e2e.py
21 kB 666 lines
1#!/usr/bin/env python3 2# SPDX-License-Identifier: AGPL-3.0-only 3# Copyright (c) 2026 sol pbc 4 5from __future__ import annotations 6 7import argparse 8import importlib.util 9import json 10import secrets 11import socket 12import subprocess 13import sys 14import threading 15import time 16from dataclasses import dataclass 17from datetime import datetime, timezone 18from pathlib import Path 19from types import ModuleType 20from typing import Callable, cast 21 22from cryptography.hazmat.primitives import serialization 23from cryptography.hazmat.primitives.asymmetric import ec 24from OpenSSL import SSL 25 26from solstone.think.services.spp_attest.composite import ( 27 CompositeVerdict, 28 verify_composite, 29) 30from solstone.think.services.spp_attest.nvgpu.binary import locate_nvattest 31from solstone.think.services.spp_attest.nvgpu.errors import GpuAppraisalError 32from solstone.think.services.spp_attest.ratls import verify as ratls_verify 33from solstone.think.services.spp_attest.ratls.channel import ( 34 AttestedChannel, 35 RatlsEndpoint, 36 establish_attested_channel, 37) 38from solstone.think.services.spp_attest.ratls.contract import ( 39 EXPORTER_BYTES, 40 EXPORTER_LABEL, 41 OWNER_NONCE_BYTES, 42 PREFACE_MAGIC, 43 exporter_context, 44) 45 46DEFAULT_REQUEST_BODY = b"{}" 47DEFAULT_BANNER = ( 48 "mode: default synthetic loopback; verifier stubs installed; " 49 "protocol regression only, not proof of real appraisal" 50) 51REAL_BANNER = "mode: real loopback; production verifiers against live hardware" 52 53 54@dataclass(frozen=True, slots=True) 55class RealModeConfig: 56 nvattest_dir: Path 57 upstream_port: int 58 request_body: bytes 59 60 61@dataclass(frozen=True, slots=True) 62class RunContext: 63 establish: Callable[[int], AttestedChannel] 64 upstream_port: int | None 65 request_body: bytes 66 banner: str 67 68 69def build_parser() -> argparse.ArgumentParser: 70 parser = argparse.ArgumentParser( 71 description=( 72 "Run the SPP RA-TLS loopback harness. Default mode runs the five " 73 "protocol-regression cases with synthetic upstream traffic and " 74 "verifier stubs; it is not proof of real hardware appraisal." 75 ), 76 epilog=( 77 "Real mode (--real) runs the same five cases with production verifiers. " 78 "It requires live confidential hardware, a real collector and gateway, " 79 "an nvattest install, and a separately-running loopback upstream " 80 "listening on 127.0.0.1:<upstream-port>. --real requires " 81 "--nvattest-dir, --upstream-port, and --request-body; those flags " 82 "are rejected without --real." 83 ), 84 ) 85 parser.add_argument("--gateway", type=Path, required=True) 86 parser.add_argument("--collector", type=Path, required=True) 87 parser.add_argument( 88 "--real", 89 action="store_true", 90 help="run the same five cases with production verifiers against live hardware", 91 ) 92 parser.add_argument( 93 "--nvattest-dir", 94 type=Path, 95 help=( 96 "nvattest install root containing bin/nvattest, lib/, and " 97 "share/ca/ca-bundle.pem (requires --real)" 98 ), 99 ) 100 parser.add_argument( 101 "--upstream-port", 102 type=int, 103 help=( 104 "port of the separately-running 127.0.0.1 loopback upstream " 105 "(requires --real)" 106 ), 107 ) 108 parser.add_argument( 109 "--request-body", 110 type=Path, 111 help=( 112 "path to the raw JSON request body sent to /v1/chat/completions " 113 "(requires --real)" 114 ), 115 ) 116 return parser 117 118 119def validate_runtime_args( 120 args: argparse.Namespace, parser: argparse.ArgumentParser 121) -> RealModeConfig | None: 122 real_only = { 123 "--nvattest-dir": args.nvattest_dir, 124 "--upstream-port": args.upstream_port, 125 "--request-body": args.request_body, 126 } 127 provided_real_only = [ 128 name for name, value in real_only.items() if value is not None 129 ] 130 if not args.real: 131 if provided_real_only: 132 parser.error(f"{', '.join(provided_real_only)} require --real") 133 return None 134 135 missing = [name for name, value in real_only.items() if value is None] 136 if missing: 137 parser.error(f"--real requires {', '.join(missing)}") 138 139 try: 140 request_body = args.request_body.read_bytes() 141 except OSError: 142 parser.error(f"unable to read --request-body: {args.request_body}") 143 try: 144 json.loads(request_body) 145 except ValueError: 146 parser.error("--request-body must contain valid JSON") 147 148 try: 149 locate_nvattest(args.nvattest_dir) 150 except GpuAppraisalError: 151 parser.error( 152 "--nvattest-dir must contain bin/nvattest, lib/, and share/ca/ca-bundle.pem" 153 ) 154 155 return RealModeConfig( 156 nvattest_dir=args.nvattest_dir.resolve(), 157 upstream_port=args.upstream_port, 158 request_body=request_body, 159 ) 160 161 162def build_chat_completions_request(body: bytes, *, host: str = "spp-engine") -> bytes: 163 return ( 164 b"POST /v1/chat/completions HTTP/1.1\r\n" 165 + f"Host: {host}\r\n".encode("ascii") 166 + f"Content-Length: {len(body)}\r\n\r\n".encode("ascii") 167 + body 168 ) 169 170 171def validate_chat_completion_envelope(response: bytes) -> None: 172 try: 173 _head, body = response.split(b"\r\n\r\n", 1) 174 data = json.loads(body) 175 except ValueError: 176 raise RuntimeError("response_envelope_invalid") from None 177 if not isinstance(data, dict): 178 raise RuntimeError("response_envelope_invalid") 179 if not isinstance(data.get("id"), str) or not data["id"]: 180 raise RuntimeError("response_envelope_invalid") 181 if data.get("object") != "chat.completion": 182 raise RuntimeError("response_envelope_invalid") 183 choices = data.get("choices") 184 if not isinstance(choices, list) or not choices: 185 raise RuntimeError("response_envelope_invalid") 186 choice = choices[0] 187 if not isinstance(choice, dict): 188 raise RuntimeError("response_envelope_invalid") 189 message = choice.get("message") 190 if not isinstance(message, dict): 191 raise RuntimeError("response_envelope_invalid") 192 if message.get("role") != "assistant": 193 raise RuntimeError("response_envelope_invalid") 194 if "content" not in message: 195 raise RuntimeError("response_envelope_invalid") 196 finish_reason = choice.get("finish_reason") 197 if not isinstance(finish_reason, str) or not finish_reason: 198 raise RuntimeError("response_envelope_invalid") 199 200 201class Upstream: 202 def __init__(self) -> None: 203 self.listener = socket.socket() 204 self.listener.bind(("127.0.0.1", 0)) 205 self.listener.listen(1) 206 self.port = self.listener.getsockname()[1] 207 self.request = b"" 208 self.thread = threading.Thread(target=self._run, daemon=True) 209 210 def start(self) -> None: 211 self.thread.start() 212 213 def close(self) -> None: 214 try: 215 self.listener.close() 216 except OSError: 217 pass 218 219 def _run(self) -> None: 220 try: 221 conn, _addr = self.listener.accept() 222 except OSError: 223 return 224 with conn: 225 self.request = _recv_http_request(conn) 226 body = b'{"id":"ok","choices":[]}' 227 conn.sendall( 228 b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\n" 229 + f"Content-Length: {len(body)}\r\nConnection: close\r\n\r\n".encode( 230 "ascii" 231 ) 232 + body 233 ) 234 self.close() 235 236 237class GatewayProcess: 238 def __init__( 239 self, gateway_path: Path, collector_path: Path, upstream_port: int 240 ) -> None: 241 self.process = subprocess.Popen( 242 [ 243 sys.executable, 244 str(gateway_path), 245 "--listen-host", 246 "127.0.0.1", 247 "--listen-port", 248 "0", 249 "--upstream-port", 250 str(upstream_port), 251 "--collector-command", 252 f"{sys.executable} {collector_path}", 253 ], 254 stdout=subprocess.PIPE, 255 stderr=subprocess.DEVNULL, 256 text=True, 257 ) 258 assert self.process.stdout is not None 259 ready = json.loads(self.process.stdout.readline()) 260 self.port = int(ready["port"]) 261 262 def close(self) -> None: 263 self.process.terminate() 264 try: 265 self.process.communicate(timeout=5) 266 except subprocess.TimeoutExpired: 267 self.process.kill() 268 self.process.communicate(timeout=5) 269 270 271class AdversarialGateway: 272 def __init__( 273 self, 274 gateway_module: ModuleType, 275 collector_path: Path, 276 mode: str, 277 timeout_s: int, 278 ) -> None: 279 self.gateway = gateway_module 280 self.collector_path = collector_path 281 self.mode = mode 282 self.timeout_s = timeout_s 283 self.listener = socket.socket() 284 self.listener.bind(("127.0.0.1", 0)) 285 self.listener.listen(1) 286 self.port = self.listener.getsockname()[1] 287 self.thread = threading.Thread(target=self._run, daemon=True) 288 289 def start(self) -> None: 290 self.thread.start() 291 292 def close(self) -> None: 293 try: 294 self.listener.close() 295 except OSError: 296 pass 297 298 def _collector(self): 299 return self.gateway.CommandCollector( 300 [sys.executable, str(self.collector_path)], self.timeout_s 301 ) 302 303 def _run(self) -> None: 304 try: 305 raw, _addr = self.listener.accept() 306 except OSError: 307 return 308 collector = self._collector() 309 connection = None 310 try: 311 preface = self.gateway._recv_exact( 312 raw, len(PREFACE_MAGIC) + OWNER_NONCE_BYTES 313 ) 314 owner_nonce = preface[len(PREFACE_MAGIC) :] 315 key_a = ec.generate_private_key(ec.SECP256R1()) 316 key_b = ec.generate_private_key(ec.SECP256R1()) 317 spki_a = _spki_der(key_a) 318 319 if self.mode == "relay": 320 evidence = collector.collect_composite(owner_nonce, spki_a) 321 tls_key = key_b 322 elif self.mode == "splice": 323 foreign_nonce = b"s" * OWNER_NONCE_BYTES 324 evidence = collector.collect_composite(foreign_nonce, spki_a) 325 tls_key = key_a 326 else: 327 evidence = collector.collect_composite(owner_nonce, spki_a) 328 tls_key = key_a 329 330 cert = self.gateway._make_certificate(tls_key, evidence.to_der()) 331 connection = SSL.Connection(self.gateway._tls_context(tls_key, cert), raw) 332 connection.setblocking(1) 333 connection.set_accept_state() 334 connection.do_handshake() 335 336 if self.mode != "stale": 337 return 338 339 tls_exporter = connection.export_keying_material( 340 EXPORTER_LABEL, 341 EXPORTER_BYTES, 342 exporter_context(owner_nonce, spki_a), 343 ) 344 self.gateway._recv_proof_request(connection) 345 proof = collector.collect_exporter_proof( 346 owner_nonce, 347 spki_a, 348 b"x" * len(tls_exporter), 349 evidence.gpu_envelope, 350 ) 351 self.gateway._send_proof(connection, proof.to_der()) 352 except Exception: 353 return 354 finally: 355 if connection is not None: 356 try: 357 connection.close() 358 except Exception: 359 pass 360 else: 361 try: 362 raw.close() 363 except OSError: 364 pass 365 self.close() 366 367 368def _spki_der(key: ec.EllipticCurvePrivateKey) -> bytes: 369 return key.public_key().public_bytes( 370 serialization.Encoding.DER, 371 serialization.PublicFormat.SubjectPublicKeyInfo, 372 ) 373 374 375def _recv_http_request(conn: socket.socket) -> bytes: 376 data = bytearray() 377 while b"\r\n\r\n" not in data: 378 chunk = conn.recv(4096) 379 if not chunk: 380 return bytes(data) 381 data.extend(chunk) 382 head, body = bytes(data).split(b"\r\n\r\n", 1) 383 length = 0 384 for line in head.split(b"\r\n")[1:]: 385 name, _, value = line.partition(b":") 386 if name.lower() == b"content-length": 387 length = int(value.strip()) 388 while len(body) < length: 389 chunk = conn.recv(length - len(body)) 390 if not chunk: 391 break 392 body += chunk 393 return head + b"\r\n\r\n" + body[:length] 394 395 396def _recv_http_response(connection: SSL.Connection) -> bytes: 397 data = bytearray() 398 while b"\r\n\r\n" not in data: 399 data.extend(connection.recv(4096)) 400 head, body = bytes(data).split(b"\r\n\r\n", 1) 401 length = 0 402 for line in head.split(b"\r\n")[1:]: 403 name, _, value = line.partition(b":") 404 if name.lower() == b"content-length": 405 length = int(value.strip()) 406 while len(body) < length: 407 body += connection.recv(length - len(body)) 408 return head + b"\r\n\r\n" + body[:length] 409 410 411def _load_gateway_module(gateway_path: Path) -> ModuleType: 412 sys.path.insert(0, str(gateway_path.parent)) 413 spec = importlib.util.spec_from_file_location( 414 "spp_ratls_gateway_harness", gateway_path 415 ) 416 if spec is None or spec.loader is None: 417 raise RuntimeError("gateway_load_failed") 418 module = importlib.util.module_from_spec(spec) 419 spec.loader.exec_module(module) 420 return module 421 422 423def _stub_verifiers() -> None: 424 ratls_verify.verify_quote = lambda **_kwargs: None 425 426 427def _stub_composite_verifier(_bundle, **_kwargs): 428 return object() 429 430 431def _establish( 432 port: int, 433 *, 434 owner_nonce: bytes, 435 nvattest_dir: Path, 436 composite_verifier: Callable[..., CompositeVerdict], 437) -> AttestedChannel: 438 return establish_attested_channel( 439 RatlsEndpoint("127.0.0.1", port), 440 owner_nonce=owner_nonce, 441 nvattest_dir=nvattest_dir, 442 now=datetime.now(timezone.utc), 443 composite_verifier=composite_verifier, 444 monotonic_now=time.monotonic, 445 epoch=0, 446 ) 447 448 449def _establish_default(port: int) -> AttestedChannel: 450 return _establish( 451 port, 452 owner_nonce=b"n" * OWNER_NONCE_BYTES, 453 nvattest_dir=Path("."), 454 composite_verifier=cast( 455 Callable[..., CompositeVerdict], 456 _stub_composite_verifier, 457 ), 458 ) 459 460 461def _establish_real(port: int, *, nvattest_dir: Path) -> AttestedChannel: 462 return _establish( 463 port, 464 owner_nonce=secrets.token_bytes(OWNER_NONCE_BYTES), 465 nvattest_dir=nvattest_dir, 466 composite_verifier=verify_composite, 467 ) 468 469 470def make_run_context(config: RealModeConfig | None) -> RunContext: 471 if config is None: 472 _stub_verifiers() 473 return RunContext( 474 establish=_establish_default, 475 upstream_port=None, 476 request_body=DEFAULT_REQUEST_BODY, 477 banner=DEFAULT_BANNER, 478 ) 479 return RunContext( 480 establish=lambda port: _establish_real(port, nvattest_dir=config.nvattest_dir), 481 upstream_port=config.upstream_port, 482 request_body=config.request_body, 483 banner=REAL_BANNER, 484 ) 485 486 487def run_positive(gateway_path: Path, collector_path: Path, context: RunContext) -> str: 488 request = build_chat_completions_request(context.request_body) 489 upstream = None 490 upstream_port = context.upstream_port 491 if upstream_port is None: 492 upstream = Upstream() 493 upstream.start() 494 upstream_port = upstream.port 495 496 gateway = GatewayProcess(gateway_path, collector_path, upstream_port) 497 channel = None 498 try: 499 channel = context.establish(gateway.port) 500 channel.tls.sendall(request) 501 response = _recv_http_response(channel.tls) 502 if b"HTTP/1.1 200 OK" not in response: 503 raise RuntimeError("positive_http_failed") 504 if upstream is not None: 505 upstream.thread.join(timeout=5) 506 # The forwarded request line, not the verbatim bytes: a gateway may 507 # rewrite or add headers while proxying. Byte-exact construction of 508 # the caller's request is a property of build_chat_completions_request 509 # and is proven there. 510 if not upstream.request.startswith(b"POST /v1/chat/completions"): 511 raise RuntimeError("positive_proxy_failed") 512 else: 513 validate_chat_completion_envelope(response) 514 substrate = channel.verdict.substrate 515 if not substrate: 516 raise RuntimeError("positive_substrate_missing") 517 return f"verified substrate={substrate}" 518 return "verified" 519 finally: 520 if channel is not None: 521 channel.close() 522 gateway.close() 523 if upstream is not None: 524 upstream.close() 525 526 527def run_adversarial( 528 gateway_module: ModuleType, 529 collector_path: Path, 530 context: RunContext, 531 mode: str, 532 expected_reason: str, 533) -> str: 534 gateway = AdversarialGateway(gateway_module, collector_path, mode, timeout_s=5) 535 gateway.start() 536 try: 537 channel = context.establish(gateway.port) 538 except Exception as exc: 539 reason = getattr(exc, "reason_code", type(exc).__name__) 540 if reason != expected_reason: 541 raise RuntimeError( 542 f"expected {expected_reason}, observed {reason}" 543 ) from exc 544 return reason 545 else: 546 channel.close() 547 raise RuntimeError("unexpected_success") 548 finally: 549 gateway.close() 550 551 552def run_premature_inference( 553 gateway_path: Path, 554 collector_path: Path, 555 context: RunContext, 556) -> str: 557 request = build_chat_completions_request(context.request_body) 558 upstream = None 559 upstream_port = context.upstream_port 560 if upstream_port is None: 561 upstream = Upstream() 562 upstream.start() 563 upstream_port = upstream.port 564 565 gateway = GatewayProcess(gateway_path, collector_path, upstream_port) 566 raw = socket.create_connection(("127.0.0.1", gateway.port), timeout=5) 567 connection = None 568 try: 569 raw.sendall(PREFACE_MAGIC + b"p" * OWNER_NONCE_BYTES) 570 tls_context = SSL.Context(SSL.TLS_CLIENT_METHOD) 571 tls_context.set_min_proto_version(SSL.TLS1_3_VERSION) 572 tls_context.set_max_proto_version(SSL.TLS1_3_VERSION) 573 tls_context.set_verify(SSL.VERIFY_NONE, lambda *_args: True) 574 connection = SSL.Connection(tls_context, raw) 575 connection.setblocking(1) 576 connection.set_connect_state() 577 connection.set_tlsext_host_name(b"spp-engine") 578 connection.do_handshake() 579 connection.sendall(request) 580 try: 581 if connection.recv(1) != b"": 582 raise RuntimeError("premature_inference_not_rejected") 583 except (SSL.ZeroReturnError, SSL.SysCallError, ConnectionError, OSError): 584 pass 585 if upstream is not None: 586 upstream.thread.join(timeout=0.5) 587 if upstream.request: 588 raise RuntimeError("premature_inference_reached_upstream") 589 return "protocol_or_tls_rejected" 590 finally: 591 if connection is not None: 592 connection.close() 593 raw.close() 594 gateway.close() 595 if upstream is not None: 596 upstream.close() 597 598 599def build_cases( 600 gateway_path: Path, 601 collector_path: Path, 602 gateway_module: ModuleType, 603 context: RunContext, 604) -> list[tuple[str, Callable[[], str]]]: 605 return [ 606 ("positive", lambda: run_positive(gateway_path, collector_path, context)), 607 ( 608 "relay", 609 lambda: run_adversarial( 610 gateway_module, 611 collector_path, 612 context, 613 "relay", 614 "spki_mismatch", 615 ), 616 ), 617 ( 618 "splice", 619 lambda: run_adversarial( 620 gateway_module, 621 collector_path, 622 context, 623 "splice", 624 "nonce_mismatch", 625 ), 626 ), 627 ( 628 "stale", 629 lambda: run_adversarial( 630 gateway_module, 631 collector_path, 632 context, 633 "stale", 634 "exporter_mismatch", 635 ), 636 ), 637 ( 638 "premature-inference", 639 lambda: run_premature_inference(gateway_path, collector_path, context), 640 ), 641 ] 642 643 644def main(argv: list[str] | None = None) -> int: 645 parser = build_parser() 646 args = parser.parse_args(argv) 647 config = validate_runtime_args(args, parser) 648 context = make_run_context(config) 649 gateway_module = _load_gateway_module(args.gateway) 650 cases = build_cases(args.gateway, args.collector, gateway_module, context) 651 652 print(context.banner) 653 failed = False 654 for name, runner in cases: 655 try: 656 reason = runner() 657 except Exception as exc: 658 failed = True 659 print(f"{name}: FAIL {type(exc).__name__}") 660 else: 661 print(f"{name}: PASS {reason}") 662 return 1 if failed else 0 663 664 665if __name__ == "__main__": 666 raise SystemExit(main())