dev-bucket-sync / tests /test_caller_write.py
cmpatino's picture
cmpatino HF Staff
PR #15 backend: caller-token handshake writes in a child process (pre-review build)
be19956 verified
Raw History Blame Contribute Delete
9.73 kB
"""Caller-token bucket writes run in a child process (app/caller_write.py), so
a refused Xet upload cannot poison the backend's own Xet session."""
import json
import os
import subprocess
import sys
import threading
import time
from pathlib import Path
import httpx
import pytest
from huggingface_hub.errors import HfHubHTTPError
import app.caller_write as caller_write
import app.hub as hub_module
TOKEN = "hf_CALLERsecretTokenDoNotLeak42"
ADMIN = "hf_ADMINsecretTokenDoNotLeak42"
BUCKET = "test-org/test-agent-9"
BACKEND_DIR = Path(caller_write.__file__).resolve().parent.parent
# Stand-in child: behaves according to the bucket name, and reports what it
# could see of the token and the environment.
FAKE_CHILD = r"""
import json, os, sys, time
req = json.loads(sys.stdin.read())
seen = {"token_on_stdin": req["token"] == %r,
"token_in_argv": any(%r in a for a in sys.argv),
"admin_in_env": any(%r in v for v in os.environ.values()),
"hf_token_path": os.environ.get("HF_TOKEN_PATH")}
kind = req["bucket"].split("/")[1]
if kind == "slow":
time.sleep(30)
if kind == "crash":
sys.stderr.write("Traceback ... token=" + req["token"] + "\n")
sys.exit(1)
if kind == "garbage":
print("not json")
sys.exit(0)
result = {"ok": {"result": "ok"},
"forbid": {"result": "forbidden", "status": 403},
"fail": {"result": "failed", "type": "RuntimeError", "status": None}}.get(kind, {"result": "ok"})
result["seen"] = seen
print(json.dumps(result))
""" % (TOKEN, TOKEN, ADMIN)
@pytest.fixture
def client(env, monkeypatch):
"""A real HubClient whose child is the stand-in above, with the pre-check
neutral and in-process uploads forbidden outright."""
monkeypatch.setenv("HF_TOKEN", ADMIN)
monkeypatch.setattr(hub_module, "_CALLER_WRITE_CMD", [sys.executable, "-c", FAKE_CHILD])
def no_in_process_upload(**k):
raise AssertionError("caller write ran in the backend process")
monkeypatch.setattr(hub_module, "batch_bucket_files", no_in_process_upload)
c = hub_module.HubClient(env.settings)
monkeypatch.setattr(c, "caller_may_write", lambda bucket, token: None)
return c
def _hub_error(status):
resp = httpx.Response(status, request=httpx.Request("GET", "https://hf.co/api/x"))
return HfHubHTTPError(f"{status}", response=resp)
# ── classification (shared by the child and create_bucket_as) ────────────────
@pytest.mark.parametrize("exc, expected", [
(_hub_error(403), {"result": "forbidden", "status": 403}),
(_hub_error(401), {"result": "forbidden", "status": 401}),
(ConnectionError("Network error: Request error: HTTP status client error (403 Forbidden), "
"domain: https://huggingface.co/api/buckets/o/b/xet-write-token"),
{"result": "forbidden", "status": 403}),
# A bare 401/403 elsewhere in the text (a bucket name, a URL) is not a refusal.
(ConnectionError("Network error: connection reset, domain: https://huggingface.co/api/buckets/o/dev-403/x"),
{"result": "failed", "type": "ConnectionError", "status": None}),
(ConnectionError("HTTP status client error (404 Not Found)"),
{"result": "failed", "type": "ConnectionError", "status": None}),
(RuntimeError("Previous task error: HTTP status client error (403 Forbidden)"),
{"result": "failed", "type": "RuntimeError", "status": None}),
(_hub_error(500), {"result": "failed", "type": "HfHubHTTPError", "status": 500}),
])
def test_classify(exc, expected):
assert caller_write.classify(exc) == expected
def test_child_run_reports_structured_results_without_the_token(monkeypatch):
import huggingface_hub
calls = []
outcomes = iter([None, _hub_error(403), RuntimeError("boom " + TOKEN)])
def fake_upload(**k):
calls.append(k)
e = next(outcomes)
if e:
raise e
monkeypatch.setattr(huggingface_hub, "batch_bucket_files", fake_upload)
req = {"token": TOKEN, "bucket": BUCKET, "path": ".bucket-sync-handshake", "text": "u"}
results = [caller_write.run(req) for _ in range(3)]
assert [r["result"] for r in results] == ["ok", "forbidden", "failed"]
assert calls[0]["token"] == TOKEN and calls[0]["add"] == [(b"u", ".bucket-sync-handshake")]
assert TOKEN not in json.dumps(results)
def test_real_child_runs_isolated_and_returns_a_result(monkeypatch):
# The real module end to end, with no network: an invalid bucket id is
# rejected by huggingface_hub before any request (a closed port instead
# would sit in its retry backoff for minutes). The token goes in on stdin
# and must come back out nowhere.
env = {k: v for k, v in os.environ.items() if k not in ("HF_TOKEN",)}
env.update(HF_ENDPOINT="http://127.0.0.1:9", HF_TOKEN_PATH=os.devnull, HF_HUB_DISABLE_PROGRESS_BARS="1")
proc = subprocess.run(
[sys.executable, "-m", "app.caller_write"],
input=json.dumps({"token": TOKEN, "bucket": "not a/valid bucket id", "path": "p", "text": "t"}),
capture_output=True, text=True, timeout=120, cwd=BACKEND_DIR, env=env,
)
result = json.loads(proc.stdout.strip().splitlines()[-1])
assert result["result"] == "failed"
assert TOKEN not in proc.stdout and TOKEN not in proc.stderr
bad = subprocess.run([sys.executable, "-m", "app.caller_write"], input="{", capture_output=True,
text=True, timeout=60, cwd=BACKEND_DIR, env=env)
assert json.loads(bad.stdout)["type"] == "BadRequest"
# ── parent: HubClient.write_text_as ──────────────────────────────────────────
def _seen_by_child(monkeypatch):
seen = {}
real = hub_module._run_caller_write
def spy(*a):
r = real(*a)
seen.update(r.get("seen", {}))
return r
monkeypatch.setattr(hub_module, "_run_caller_write", spy)
return seen
def test_ok_write_uses_stdin_token_and_no_admin_credential(client, monkeypatch):
seen = _seen_by_child(monkeypatch)
client.write_text_as("test-org/ok", ".bucket-sync-handshake", "u", TOKEN)
assert seen == {"token_on_stdin": True, "token_in_argv": False,
"admin_in_env": False, "hf_token_path": os.devnull}
def test_refused_write_is_permission_error(client):
with pytest.raises(PermissionError) as exc:
client.write_text_as("test-org/forbid", "p", "t", TOKEN)
assert TOKEN not in str(exc.value)
@pytest.mark.parametrize("kind", ["fail", "crash", "garbage"])
def test_failed_crashed_or_unreadable_child_is_upstream_error(client, caplog, kind):
caplog.set_level("DEBUG")
with pytest.raises(hub_module.HubUnreachable) as exc:
client.write_text_as(f"test-org/{kind}", "p", "t", TOKEN)
assert TOKEN not in str(exc.value)
assert TOKEN not in caplog.text # the crashing child printed it to stderr
def test_child_timeout_is_upstream_error(client, monkeypatch):
monkeypatch.setattr(hub_module, "CALLER_WRITE_TIMEOUT_S", 1.0)
t0 = time.monotonic()
with pytest.raises(hub_module.HubUnreachable):
client.write_text_as("test-org/slow", "p", "t", TOKEN)
assert time.monotonic() - t0 < 15
def test_concurrent_writes_are_bounded(client, monkeypatch):
monkeypatch.setattr(hub_module, "CALLER_WRITE_SLOTS", threading.BoundedSemaphore(1))
monkeypatch.setattr(hub_module, "CALLER_WRITE_TIMEOUT_S", 5.0)
monkeypatch.setattr(hub_module, "CALLER_WRITE_SLOT_WAIT_S", 0.5)
slow = threading.Thread(target=lambda: pytest.raises(
hub_module.HubUnreachable, client.write_text_as, "test-org/slow", "p", "t", TOKEN))
slow.start()
time.sleep(1.0) # the slow child holds the only slot
with pytest.raises(hub_module.HubUnreachable):
client.write_text_as("test-org/ok", "p", "t", TOKEN)
slow.join()
client.write_text_as("test-org/ok", "p", "t", TOKEN) # slot released
def test_precheck_refusal_never_starts_a_child(client, monkeypatch):
monkeypatch.setattr(client, "caller_may_write", lambda bucket, token: False)
def no_child(*a):
raise AssertionError("child started after a refused pre-check")
monkeypatch.setattr(hub_module, "_run_caller_write", no_child)
with pytest.raises(PermissionError):
client.write_text_as("test-org/ok", "p", "t", TOKEN)
def test_precheck_success_then_upload_failure_is_upstream_error(client, monkeypatch):
monkeypatch.setattr(client, "caller_may_write", lambda bucket, token: True)
with pytest.raises(hub_module.HubUnreachable):
client.write_text_as("test-org/fail", "p", "t", TOKEN)
with pytest.raises(PermissionError): # permissions changed after the check
client.write_text_as("test-org/forbid", "p", "t", TOKEN)
class _FakeSession:
def __init__(self, outcome):
self.outcome, self.calls = outcome, []
def get(self, url, headers, timeout):
self.calls.append((url, headers))
if isinstance(self.outcome, Exception):
raise self.outcome
return httpx.Response(self.outcome, request=httpx.Request("GET", url))
@pytest.mark.parametrize("outcome, expected", [
(403, False), (401, False), (200, True), (500, None), (404, None), (httpx.ConnectError("down"), None),
])
def test_caller_may_write(env, monkeypatch, outcome, expected):
session = _FakeSession(outcome)
monkeypatch.setattr(hub_module, "get_session", lambda: session)
c = hub_module.HubClient(env.settings)
assert c.caller_may_write(BUCKET, TOKEN) is expected
url, headers = session.calls[0]
assert url.endswith(f"/api/buckets/{BUCKET}/xet-write-token")
assert headers["authorization"] == f"Bearer {TOKEN}"