Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions Dockerfile
Original file line number Diff line number Diff line change
Expand Up @@ -110,5 +110,9 @@ RUN --mount=type=secret,id=attestation_secret,required=true \
# Build-time only (BuildKit secret mount → file, not ENV):
# attestation_secret → /run/prism/attestation_hmac_key (mode 0400)

# Sidecar listen mode (optional). Publish this port in the Lium template's
# internal_ports so BASE can dial POST /v1/sidecar/attest on the instance.
EXPOSE 8787

ENTRYPOINT ["prism-recipe"]
CMD ["preflight"]
4 changes: 4 additions & 0 deletions Dockerfile.cuda
Original file line number Diff line number Diff line change
Expand Up @@ -88,5 +88,9 @@ RUN --mount=type=secret,id=attestation_secret,required=true \
&& test -s /run/prism/attestation_hmac_key \
&& touch -h -d "@${SOURCE_DATE_EPOCH}" /run/prism /run/prism/attestation_hmac_key

# Sidecar listen mode (optional). Publish this port in the Lium template's
# internal_ports so BASE can dial POST /v1/sidecar/attest on the instance.
EXPOSE 8787

ENTRYPOINT ["prism-recipe"]
CMD ["preflight"]
105 changes: 64 additions & 41 deletions src/prism_recipe/sidecar/__main__.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@

from prism_recipe.sidecar.config import SidecarConfig
from prism_recipe.sidecar.errors import SidecarError, SidecarReachabilityError
from prism_recipe.sidecar.listen import serve_forever
from prism_recipe.sidecar.service import AttestationSidecar
from prism_recipe.sidecar.transport import FakeChallengeTransport, HttpxChallengeTransport
from prism_recipe.sidecar.types import Challenge, ChallengePhase
Expand Down Expand Up @@ -50,21 +51,42 @@ def main(argv: Sequence[str] | None = None) -> int:

p_once = sub.add_parser("answer-once", help="Answer a single challenge JSON")
p_once.add_argument("--nonce", required=True)
p_once.add_argument(
"--phase", choices=("start", "interval", "end"), default="start"
)
p_once.add_argument("--phase", choices=("start", "interval", "end"), default="start")
p_once.add_argument("--pod-id", required=True)
p_once.add_argument("--digest", required=True)
p_once.add_argument("--variant", choices=("cpu", "cuda"), default="cpu")
p_once.add_argument("--secret-path", type=Path, required=True)
p_once.add_argument("--root", type=Path, default=None)

p_serve = sub.add_parser(
"serve",
help="Listen for BASE attest requests (stdlib HTTP; no pull loop)",
)
p_serve.add_argument(
"--host",
default="0.0.0.0",
help="Bind address (explicit; default 0.0.0.0 for container publish)",
)
p_serve.add_argument(
"--port",
type=int,
default=8787,
help="Bind port (default 8787; publish via Lium internal_ports)",
)
p_serve.add_argument("--pod-id", default=None)
p_serve.add_argument("--digest", default=None)
p_serve.add_argument("--variant", choices=("cpu", "cuda"), default=None)
p_serve.add_argument("--secret-path", type=Path, default=None)
p_serve.add_argument("--root", type=Path, default=None)

args = parser.parse_args(list(sys.argv[1:] if argv is None else argv))

if args.command == "answer-once":
return _cmd_answer_once(args)
if args.command == "run":
return _cmd_run(args)
if args.command == "serve":
return _cmd_serve(args)
return 2


Expand Down Expand Up @@ -100,39 +122,48 @@ def _cmd_run(args: argparse.Namespace) -> int:
try:
if args.fake:
cfg = _config_from_args_or_env(args, require_identity=True)
transport: FakeChallengeTransport | HttpxChallengeTransport = (
FakeChallengeTransport(
challenges_by_phase={
ChallengePhase.START: Challenge(
nonce="fake-nonce-start", phase=ChallengePhase.START
),
ChallengePhase.INTERVAL: Challenge(
nonce="fake-nonce-interval", phase=ChallengePhase.INTERVAL
),
ChallengePhase.END: Challenge(
nonce="fake-nonce-end", phase=ChallengePhase.END
),
}
)
transport: FakeChallengeTransport | HttpxChallengeTransport = FakeChallengeTransport(
challenges_by_phase={
ChallengePhase.START: Challenge(
nonce="fake-nonce-start", phase=ChallengePhase.START
),
ChallengePhase.INTERVAL: Challenge(
nonce="fake-nonce-interval", phase=ChallengePhase.INTERVAL
),
ChallengePhase.END: Challenge(nonce="fake-nonce-end", phase=ChallengePhase.END),
}
)
else:
cfg = _config_from_args_or_env(args, require_identity=True)
base_url = args.base_url or cfg.base_url
if not base_url:
raise SidecarError("base_url required (env PRISM_SIDECAR_BASE_URL)")
transport = HttpxChallengeTransport(base_url=base_url)
code = AttestationSidecar(cfg).run(
transport=transport, rng=random.Random(args.seed)
)
code = AttestationSidecar(cfg).run(transport=transport, rng=random.Random(args.seed))
return code
except SidecarError as exc:
sys.stderr.write(f"sidecar error: {exc}\n")
return 1


def _config_from_args_or_env(
args: argparse.Namespace, *, require_identity: bool
) -> SidecarConfig:
def _cmd_serve(args: argparse.Namespace) -> int:
try:
cfg = _config_from_args_or_env(args, require_identity=True)
# Attach optional CLI-only knobs that _config_from_args_or_env may omit
# when falling back to env (base_url unused in listen mode).
serve_forever(cfg, host=str(args.host), port=int(args.port))
return 0
except (SidecarError, SidecarReachabilityError, ValueError, OSError) as exc:
sys.stderr.write(f"sidecar error: {exc}\n")
return 1


def _config_from_args_or_env(args: argparse.Namespace, *, require_identity: bool) -> SidecarConfig:
retry_budget = getattr(args, "retry_budget", None)
interval_count = getattr(args, "interval_count", None)
interval_min_s = getattr(args, "interval_min_s", None)
interval_max_s = getattr(args, "interval_max_s", None)
base_url = getattr(args, "base_url", None)
if args.pod_id and args.digest:
variant = args.variant or "cpu"
secret = args.secret_path or SidecarConfig.default_secret_path()
Expand All @@ -142,13 +173,11 @@ def _config_from_args_or_env(
variant=variant,
secret_path=secret,
root=args.root,
retry_budget=args.retry_budget or 3,
interval_count=(
args.interval_count if args.interval_count is not None else 0
),
interval_min_s=args.interval_min_s if args.interval_min_s is not None else 0.0,
interval_max_s=args.interval_max_s if args.interval_max_s is not None else 0.0,
base_url=args.base_url,
retry_budget=retry_budget or 3,

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🎯 Functional Correctness | 🟠 Major | ⚡ Quick win

retry_budget=0 is silently overridden to the default.

Both branches use retry_budget or <default>, but interval_count/interval_min_s/interval_max_s on the same lines correctly use an explicit is not None check to preserve 0. Since 0 is a falsy-but-valid value for retry_budget (fail-fast, no retries), 0 or 3 (and 0 or env_cfg.retry_budget) silently discards an explicit --retry-budget 0.

🐛 Proposed fix
-            retry_budget=retry_budget or 3,
+            retry_budget=retry_budget if retry_budget is not None else 3,
-            retry_budget=retry_budget or env_cfg.retry_budget,
+            retry_budget=retry_budget if retry_budget is not None else env_cfg.retry_budget,

Also applies to: 191-191

🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.

In `@src/prism_recipe/sidecar/__main__.py` at line 176, Preserve an explicit zero
retry budget in the retry configuration branches by replacing the retry_budget
fallback expressions with an explicit None check, matching the neighboring
interval_count, interval_min_s, and interval_max_s handling. Update both
retry_budget assignments around the affected configuration logic so None still
selects the default while 0 remains fail-fast.

interval_count=(interval_count if interval_count is not None else 0),
interval_min_s=interval_min_s if interval_min_s is not None else 0.0,
interval_max_s=interval_max_s if interval_max_s is not None else 0.0,
base_url=base_url,
)
if require_identity and not (args.pod_id and args.digest):
# Fall back to env for container runs.
Expand All @@ -159,23 +188,17 @@ def _config_from_args_or_env(
variant=args.variant or env_cfg.variant,
secret_path=args.secret_path or env_cfg.secret_path,
root=args.root,
retry_budget=args.retry_budget or env_cfg.retry_budget,
retry_budget=retry_budget or env_cfg.retry_budget,
interval_count=(
args.interval_count
if args.interval_count is not None
else env_cfg.interval_count
interval_count if interval_count is not None else env_cfg.interval_count
),
interval_min_s=(
args.interval_min_s
if args.interval_min_s is not None
else env_cfg.interval_min_s
interval_min_s if interval_min_s is not None else env_cfg.interval_min_s
),
interval_max_s=(
args.interval_max_s
if args.interval_max_s is not None
else env_cfg.interval_max_s
interval_max_s if interval_max_s is not None else env_cfg.interval_max_s
),
base_url=args.base_url or env_cfg.base_url,
base_url=base_url or env_cfg.base_url,
)
raise SidecarError("pod_id and digest required")

Expand Down
208 changes: 208 additions & 0 deletions src/prism_recipe/sidecar/listen.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,208 @@
"""Stdlib HTTP listen mode for the in-image attestation sidecar.

BASE dials the running instance and POSTs a fresh nonce. No third-party
server stack — ``http.server.ThreadingHTTPServer`` only (hermetic image).
"""

from __future__ import annotations

import json
import logging
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from typing import Any, Final
from urllib.parse import urlparse

from prism_recipe.sidecar.config import SidecarConfig
from prism_recipe.sidecar.errors import SidecarError
from prism_recipe.sidecar.service import AttestationSidecar, ChallengeAnswer
from prism_recipe.sidecar.types import Challenge, ChallengePhase
from prism_recipe.sidecar.wire import signed_attestation_to_wire

logger = logging.getLogger(__name__)

ATTEST_PATH: Final[str] = "/v1/sidecar/attest"
HEALTHZ_PATH: Final[str] = "/healthz"
MAX_BODY_BYTES: Final[int] = 64 * 1024
_VALID_PHASES: Final[frozenset[str]] = frozenset({"start", "interval", "end"})


def build_server(
config: SidecarConfig,
*,
host: str,
port: int,
) -> ThreadingHTTPServer:
"""Bind an explicit host/port and return a ready ``ThreadingHTTPServer``.

``host`` must be caller-provided (no implicit ``0.0.0.0`` default here).
Pass ``port=0`` in tests for an ephemeral port.
"""
if not host.strip():
msg = "host must be non-empty"
raise ValueError(msg)
sidecar = AttestationSidecar(config)
handler = _make_handler(sidecar)
return ThreadingHTTPServer((host, port), handler)


def serve_forever(config: SidecarConfig, *, host: str, port: int) -> None:
"""Build the server and block serving until interrupted."""
server = build_server(config, host=host, port=port)
bound_host, bound_port = server.server_address[:2]
logger.info("sidecar listen mode on http://%s:%s", bound_host, bound_port)
try:
server.serve_forever()
finally:
server.server_close()


def answer_to_wire(answer: ChallengeAnswer) -> dict[str, Any]:
"""Serialize a ``ChallengeAnswer`` to the answer-once wire shape."""
wire = signed_attestation_to_wire(answer.signed)
wire["phase"] = answer.phase.value
wire["baked_manifest_match"] = answer.baked_manifest_match
wire["mismatched_paths"] = list(answer.mismatched_paths)
return wire


def _make_handler(sidecar: AttestationSidecar) -> type[BaseHTTPRequestHandler]:
class SidecarHTTPRequestHandler(BaseHTTPRequestHandler):
server_version = "PrismSidecar/1.0"
protocol_version = "HTTP/1.1"

def log_message(self, fmt: str, *args: object) -> None:
# Access log only — never include body/secret material.
logger.info("%s - %s", self.address_string(), fmt % args)

def do_GET(self) -> None: # noqa: N802 — stdlib handler API
path = urlparse(self.path).path
if path == HEALTHZ_PATH:
self._json_response(200, {"status": "ok"})
return
if path == ATTEST_PATH:
self._json_response(405, {"error": "method not allowed"})
return
self._json_response(404, {"error": "not found"})

def do_POST(self) -> None: # noqa: N802 — stdlib handler API
path = urlparse(self.path).path
if path != ATTEST_PATH:
if path == HEALTHZ_PATH:
self._json_response(405, {"error": "method not allowed"})
return
self._json_response(404, {"error": "not found"})
return
self._handle_attest()

def do_PUT(self) -> None: # noqa: N802
self._method_not_allowed_or_404()

def do_DELETE(self) -> None: # noqa: N802
self._method_not_allowed_or_404()

def do_PATCH(self) -> None: # noqa: N802
self._method_not_allowed_or_404()

def _method_not_allowed_or_404(self) -> None:
path = urlparse(self.path).path
if path in {ATTEST_PATH, HEALTHZ_PATH}:
self._json_response(405, {"error": "method not allowed"})
return
self._json_response(404, {"error": "not found"})

def _handle_attest(self) -> None:
try:
raw = self._read_body()
except _ClientBodyError as exc:
self._json_response(exc.status, {"error": exc.message})
return
try:
challenge = _parse_attest_body(raw)
except ValueError as exc:
self._json_response(400, {"error": str(exc)})
return
try:
answer = sidecar.answer_challenge(challenge)
wire = answer_to_wire(answer)
except SidecarError:
logger.exception("sidecar attest failed")
self._json_response(500, {"error": "internal error"})
return
except Exception:
logger.exception("unexpected attest failure")
self._json_response(500, {"error": "internal error"})
return
self._json_response(200, wire)

def _read_body(self) -> bytes:
length_hdr = self.headers.get("Content-Length")
if length_hdr is None:
# No body / chunked not supported — treat as empty.
return b""
try:
length = int(length_hdr)
except ValueError as exc:
raise _ClientBodyError(400, "invalid Content-Length") from exc
if length < 0:
raise _ClientBodyError(400, "invalid Content-Length")
if length > MAX_BODY_BYTES:
raise _ClientBodyError(413, "request body too large")
return self.rfile.read(length)

def _json_response(self, status: int, body: dict[str, Any]) -> None:
payload = json.dumps(body, separators=(",", ":"), sort_keys=True).encode("utf-8")
self.send_response(status)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(payload)))
self.send_header("Connection", "close")
self.end_headers()
self.wfile.write(payload)

return SidecarHTTPRequestHandler


def _parse_attest_body(raw: bytes) -> Challenge:
"""Parse POST /v1/sidecar/attest JSON into a ``Challenge`` (boundary)."""
if not raw.strip():
msg = "request body required"
raise ValueError(msg)
try:
data = json.loads(raw.decode("utf-8"))
except (UnicodeDecodeError, json.JSONDecodeError) as exc:
msg = "malformed JSON"
raise ValueError(msg) from exc
if not isinstance(data, dict):
msg = "JSON body must be an object"
raise ValueError(msg)
nonce = data.get("nonce")
if not isinstance(nonce, str) or not nonce.strip():
msg = "nonce must be a non-empty string"
raise ValueError(msg)
phase_raw = data.get("phase", "start")
if not isinstance(phase_raw, str):
msg = "phase must be a string"
raise ValueError(msg)
phase_key = phase_raw.strip().lower()
if phase_key not in _VALID_PHASES:
msg = f"invalid phase: {phase_raw!r}"
raise ValueError(msg)
match phase_key:
case "start":
phase = ChallengePhase.START
case "interval":
phase = ChallengePhase.INTERVAL
case "end":
phase = ChallengePhase.END
case unreachable:
msg = f"invalid phase: {unreachable!r}"
raise ValueError(msg)
return Challenge(nonce=nonce.strip(), phase=phase)


class _ClientBodyError(Exception):
"""Client-caused body read failure (maps to 4xx)."""

def __init__(self, status: int, message: str) -> None:
super().__init__(message)
self.status = status
self.message = message
Loading
Loading