#!/usr/bin/env python3
import argparse
import base64
import json
import os
import sys
import time
import urllib.error
import urllib.request
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--deadline", type=float, default=120.0)
args = parser.parse_args()
client = Client(args.deadline)
try:
run_variant(client, expect_failure=False, disconnect=True)
run_variant(client, expect_failure=True, disconnect=False)
print("Observed the successful and nonzero test outcomes.")
finally:
client.cancel_owned()
def run_variant(client, expect_failure, disconnect):
label = "expected failure" if expect_failure else "success"
operation_id = client.create_test_run(expect_failure)
print(f"\n--- {label}: operation {operation_id} ---", flush=True)
cursor = None
if disconnect:
cursor, _ = client.consume_events(operation_id, disconnect_after=5)
print(f"[client] resume after event {cursor}", flush=True)
_, completed = client.consume_events(operation_id, cursor=cursor)
result = client.final_result(operation_id)
process = verify_result(result, 1 if expect_failure else 0)
stdout = decode_stream(process["stdout"])
stderr = decode_stream(process["stderr"])
if expect_failure:
if "FAIL deliberate-failure" not in stderr:
raise RuntimeError('Failing run did not report the expected test error')
if "SUMMARY passed=2 failed=1" not in stdout:
raise RuntimeError('Failing run returned an unexpected summary')
else:
if "SUMMARY passed=2 failed=0" not in stdout:
raise RuntimeError('Successful run returned an unexpected summary')
print(f"[final] operation={result['status']} process_exit={process['state']['exit_code']}")
if not completed:
print("[final] stream ended without a completion frame; final status supplied the outcome")
# Replace this fixture suite with your build or test command.
TEST_SUITE = r"""
set -u
passed=0
failed=0
check() {
label=$1
shift
if "$@"; then
printf 'PASS %s\n' "$label"
passed=$((passed + 1))
else
printf 'FAIL %s\n' "$label" >&2
failed=$((failed + 1))
fi
sleep 1
}
check arithmetic test $((2 + 2)) -eq 4
printf 'alpha\n' > /tmp/fixture.txt
check file-content grep -qx alpha /tmp/fixture.txt
# OPTIONAL_FAILURE
printf 'SUMMARY passed=%s failed=%s\n' "$passed" "$failed"
test "$failed" -eq 0
"""
def read_sse_frames(response, check_deadline):
"""Yield complete SSE frames; a partial frame never advances the cursor."""
frame_lines = []
while True:
check_deadline()
raw = response.readline()
if not raw:
return
line = raw.decode("utf-8").rstrip("\r\n")
if line:
frame_lines.append(line)
continue
if not frame_lines:
continue
fields = {}
for item in frame_lines:
if item.startswith(":"):
continue
key, separator, value = item.partition(":")
if separator:
# SSE removes at most one space after the colon.
if value.startswith(" "):
value = value[1:]
fields.setdefault(key, []).append(value)
frame_lines = []
yield fields
TERMINAL = {"SUCCESS", "FAILED", "CANCELLED"}
def required(name):
value = os.environ.get(name)
if not value:
raise SystemExit(f"Set {name}; see https://docs.tokenfactory.nebius.com/sandboxes/start/set-up-access")
return value
def retry_delay(headers):
try:
return max(0.1, float(headers.get("Retry-After", "1")))
except ValueError:
return 1.0
def decode_stream(stream):
value = stream.get("value", "")
encoding = stream.get("encoding", "ascii")
if encoding == "base64":
return base64.b64decode(value).decode("utf-8", errors="replace")
if encoding == "ascii":
return value
raise RuntimeError(f"Unsupported output encoding: {encoding}")
class Client:
def __init__(self, deadline_seconds):
self.base = os.environ.get(
"CONTREE_BASE_URL", "https://api.tokenfactory.nebius.com/sandboxes"
).rstrip("/")
self.token = required("NEBIUS_API_KEY")
self.project = required("NEBIUS_PROJECT_ID")
self.image = required("IMAGE_UUID")
self.deadline = time.monotonic() + deadline_seconds
self.owned = set()
def remaining(self):
value = self.deadline - time.monotonic()
if value <= 0:
raise TimeoutError("Recipe deadline exceeded")
return min(15.0, value)
def request(self, method, path, body=None, extra_headers=None):
headers = {
"Authorization": f"Bearer {self.token}",
"Project": self.project,
"Accept": "application/json",
}
data = None
if body is not None:
data = json.dumps(body).encode()
headers["Content-Type"] = "application/json"
if extra_headers:
headers.update(extra_headers)
request = urllib.request.Request(
self.base + "/v1" + path, data=data, headers=headers, method=method
)
with urllib.request.urlopen(request, timeout=self.remaining()) as response:
payload = response.read()
return response.status, response.headers, json.loads(payload) if payload else None
def create_test_run(self, expect_failure):
command = TEST_SUITE.replace(
"# OPTIONAL_FAILURE", "check deliberate-failure false" if expect_failure else ""
)
status, _, result = self.request(
"POST",
"/instances",
{
"image": self.image,
"command": command,
"shell": True,
"disposable": True,
"timeout": 30,
"truncate_output_at": 1048576,
},
)
if status != 201:
raise RuntimeError(f"Expected create status 201, got {status}")
operation_id = result["uuid"]
self.owned.add(operation_id)
return operation_id
def open_events(self, operation_id, last_event_id=None):
while True:
headers = {
"Authorization": f"Bearer {self.token}",
"Project": self.project,
"Accept": "text/event-stream",
}
if last_event_id is not None:
headers["Last-Event-Id"] = str(last_event_id)
request = urllib.request.Request(
f"{self.base}/v1/operations/{operation_id}/events?follow=1",
headers=headers,
)
try:
return urllib.request.urlopen(request, timeout=self.remaining())
except urllib.error.HTTPError as error:
if error.code not in {410, 425, 502, 504}:
raise
time.sleep(min(retry_delay(error.headers), self.remaining()))
def consume_events(self, operation_id, cursor=None, disconnect_after=None):
response = self.open_events(operation_id, cursor)
processed = 0
saw_completion = False
try:
for fields in read_sse_frames(response, self.remaining):
event = fields.get("event", ["message"])[0]
if event == "sse_error":
raise ConnectionError("SSE server reported an in-band error")
if "id" not in fields or "data" not in fields:
continue
event_id = int(fields["id"][0])
payload = json.loads("\n".join(fields["data"]))
if cursor is not None and event_id <= cursor:
continue
if payload.get("type") in {"stdout", "stderr"}:
channel = payload["type"]
text = decode_stream(payload["data"])
print(f"[{channel}] {text}", end="", flush=True)
elif payload.get("type") in {"spawn", "exit", "completion"}:
print(f"[event] {payload['type']} id={event_id}", flush=True)
# The frame is complete, valid, and displayed. Commit its cursor now.
cursor = event_id
processed += 1
saw_completion = saw_completion or payload.get("type") == "completion"
if disconnect_after is not None and processed >= disconnect_after:
print(f"[client] disconnect after processed event {cursor}", flush=True)
break
finally:
response.close()
return cursor, saw_completion
def final_result(self, operation_id):
while True:
_, headers, result = self.request("GET", f"/operations/{operation_id}")
if result["status"] in TERMINAL:
self.owned.discard(operation_id)
return result
time.sleep(min(retry_delay(headers), self.remaining()))
def cancel_owned(self):
self.deadline = max(self.deadline, time.monotonic() + 15.0)
for operation_id in list(self.owned):
try:
_, _, current = self.request("GET", f"/operations/{operation_id}")
if current["status"] not in TERMINAL:
self.request("DELETE", f"/operations/{operation_id}")
self.final_result(operation_id)
except Exception as error:
print(f"Cleanup failed for {operation_id}: {error}", file=sys.stderr)
def verify_result(result, expected_exit):
process = result.get("metadata", {}).get("result") or {}
state = process.get("state") or {}
actual = state.get("exit_code")
if result.get("status") != "SUCCESS":
raise RuntimeError(f"Operation failed: {result.get('status')} {result.get('error')}")
if actual != expected_exit:
raise RuntimeError(f"Expected process exit {expected_exit}, got {actual}")
if state.get("timed_out") or state.get("signal") not in (-1, None):
raise RuntimeError(f"Process was interrupted: {state}")
for channel in ("stdout", "stderr"):
if (process.get(channel) or {}).get("truncated"):
raise RuntimeError(f"{channel} was truncated")
return process
if __name__ == "__main__":
main()