From 1c5ff80f7b6890bc471f4f0475d17d56b1b8ac21 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 24 Sep 2026 03:27:14 +0000 Subject: [PATCH 1/6] test(security): token, Origin/Host and ungated static routes Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01NSpbAKgGb282ixhsEqGQ52 --- tests/test_security_auth.py | 457 ++++++++++++++++++++++++++++++++++++ 1 file changed, 457 insertions(+) create mode 100644 tests/test_security_auth.py diff --git a/tests/test_security_auth.py b/tests/test_security_auth.py new file mode 100644 index 00000000..ebbbd4b1 --- /dev/null +++ b/tests/test_security_auth.py @@ -0,0 +1,457 @@ +"""The HTTP server's own boundary: token, Origin/Host, the ungated routes. + +tests/test_server.py covers the happy shapes - a missing or wrong token is a +401, a foreign Origin a 403, a foreign Host a 400, the query-token allowance +is per route - and tests/test_serve_main.py the `--mcp` refusal on 0.0.0.0. +This file is the adversarial side of the same boundary: path spellings that +might slip past a prefix check, credentials in the wrong place, Origin and +Host values built to fool a hostname comparison, and the routes that are +ungated *on purpose* (`/outputs`, `/inputs`, `/exports`), pinned to serving +their own roots and nothing else. + +Every probe target is under `tmp_path`; "outside" means a sibling directory +the server was never told about. +""" + +import json + +import pytest +from fastapi.testclient import TestClient + +from dw.server.app import create_app +from dw.server.jobs import JobManager + +from .test_server import ScriptedWorkerManager, success_script, valid_workflow + +TOKEN = "s3cr3t-token" +SECRET = "outside-the-roots-probe" + + +@pytest.fixture +def layout(tmp_path): + """A workspace-shaped tree plus a sibling directory holding a secret.""" + paths = { + "workflows": tmp_path / "ws" / "workflows", + "outputs": tmp_path / "ws" / "outputs", + "assets": tmp_path / "ws" / "assets", + "prompts": tmp_path / "ws" / "prompts", + "outside": tmp_path / "outside", + } + for path in paths.values(): + path.mkdir(parents=True) + (paths["workflows"] / "Basic.json").write_text(json.dumps(valid_workflow("b"))) + (paths["outputs"] / "run.png").write_bytes(b"generated") + (paths["assets"] / "iris.png").write_bytes(b"asset") + (paths["outside"] / "secret.txt").write_text(SECRET) + (paths["outside"] / "secret.png").write_text(SECRET) + paths["root"] = tmp_path / "ws" + return paths + + +@pytest.fixture +def make_client(layout, tmp_path): + def make(token=None, host="127.0.0.1", base_url="http://localhost"): + manager = JobManager( + str(layout["outputs"]), + worker_manager=ScriptedWorkerManager(success_script), + history_path=str(tmp_path / "jobs.sqlite"), + ) + app = create_app( + workflow_dir=str(layout["workflows"]), + output_dir=str(layout["outputs"]), + job_manager=manager, + prompt_dir=str(layout["prompts"]), + asset_dir=str(layout["assets"]), + workspace=str(layout["root"]), + token=token, + host=host, + ) + return TestClient(app, base_url=base_url) + + return make + + +def _jobs(client, headers=None): + listing = client.get("/api/jobs", headers=headers or {}).json() + return listing.get("jobs", listing) if isinstance(listing, dict) else listing + + +def _leaks(response): + return response.status_code == 200 and SECRET in response.text + + +# ---------------------------------------------------------------- the token + + +class TestTheTokenGate: + @pytest.mark.parametrize( + "path", + [ + "/api/health", + "/api/health/", + "/api/workflows", + "/api/jobs", + "/api/server", + "/api/gallery", + "/api/prompts", + "/api/assets", + "/%61pi/health", + "/api/%68ealth", + "/api/./health", + "/mcp", + "/mcp/", + ], + ) + def test_every_api_spelling_needs_the_token(self, make_client, path): + with make_client(token=TOKEN) as client: + response = client.get(path) + assert response.status_code == 401, (path, response.status_code) + + @pytest.mark.parametrize("path", ["//api/health", "/API/health", "/api"]) + def test_a_spelling_that_misses_the_prefix_reaches_no_api_route( + self, make_client, path + ): + """If the gate's prefix check misses a spelling, the router must miss + it too - otherwise an ungated spelling of a gated route exists.""" + with make_client(token=TOKEN) as client: + response = client.get(path) + assert response.status_code == 404, (path, response.status_code) + + @pytest.mark.parametrize( + "authorization", + [ + "", + "Bearer", + "Bearer ", + f"Basic {TOKEN}", + f"Token {TOKEN}", + f"{TOKEN}", + f"Bearer {TOKEN}x", + f"Bearer x{TOKEN}", + f"Bearer {TOKEN[:-1]}", + f"Bearer {TOKEN.upper()}", + f"Bearer {TOKEN}\x00", + ], + ) + def test_anything_but_the_exact_bearer_token_is_a_401( + self, make_client, authorization + ): + with make_client(token=TOKEN) as client: + try: + response = client.get( + "/api/health", headers={"Authorization": authorization} + ) + except Exception: + # httpx refuses to send some header values at all - nothing + # reached the server, which is the outcome the test wants + return + assert response.status_code == 401 + + def test_the_scheme_is_case_insensitive_but_the_token_is_not(self, make_client): + with make_client(token=TOKEN) as client: + assert ( + client.get( + "/api/health", headers={"Authorization": f"bearer {TOKEN}"} + ).status_code + == 200 + ) + + @pytest.mark.parametrize( + "method, path", + [ + ("GET", "/api/health"), + ("GET", "/api/workflows"), + ("GET", "/api/jobs"), + ("POST", "/api/validate"), + ("DELETE", "/api/prompts/x"), + ], + ) + def test_the_query_token_is_refused_where_it_is_not_allowed( + self, make_client, method, path + ): + """?token= exists for /EventSource GETs on marked routes only; a + route that is not marked must not accept it.""" + with make_client(token=TOKEN) as client: + response = client.request(method, f"{path}?token={TOKEN}", json={}) + assert response.status_code == 401 + + def test_the_query_token_is_get_only_even_on_a_marked_route( + self, make_client, layout + ): + with make_client(token=TOKEN) as client: + assert ( + client.get(f"/api/gallery/run.png/download?token={TOKEN}").status_code + == 200 + ) + assert ( + client.delete(f"/api/gallery/run.png?token={TOKEN}").status_code == 401 + ) + assert (layout["outputs"] / "run.png").exists() + + def test_a_valid_token_does_not_excuse_a_foreign_origin(self, make_client): + """The Origin check is not an auth check the token can satisfy: a + page that somehow holds the token still may not drive the API.""" + with make_client(token=TOKEN) as client: + response = client.post( + "/api/jobs", + json={"workflow": valid_workflow()}, + headers={ + "Authorization": f"Bearer {TOKEN}", + "Origin": "https://evil.example", + }, + ) + assert response.status_code == 403 + + def test_a_refused_request_queues_nothing(self, make_client): + with make_client(token=TOKEN) as client: + client.post("/api/jobs", json={"workflow": valid_workflow()}) + assert _jobs(client, {"Authorization": f"Bearer {TOKEN}"}) == [] + + +# ---------------------------------------------------------- Origin and Host + + +class TestOriginAndHostSpoofing: + @pytest.mark.parametrize( + "origin", + [ + "http://localhost@evil.example", + "http://localhost:8765@evil.example", + "http://localhost.evil.example", + "http://127.0.0.1.evil.example", + "http://evil.example#localhost", + "http://evil.example/localhost", + "http://evil.example?localhost", + "http://evil-localhost", + "null", + "file://", + ], + ) + def test_a_lookalike_origin_is_refused(self, make_client, origin): + with make_client() as client: + response = client.post( + "/api/jobs", + json={"workflow": valid_workflow()}, + headers={"Origin": origin}, + ) + assert response.status_code == 403, origin + + @pytest.mark.parametrize( + "origin", ["http://[::1].evil.example", "http://[localhost]", "http://["] + ) + def test_an_unparseable_origin_is_not_processed(self, make_client, origin): + """urlparse raises on these; whatever the middleware answers, the + request must not reach a route.""" + app = make_client().app + with TestClient( + app, base_url="http://localhost", raise_server_exceptions=False + ) as client: + response = client.post( + "/api/jobs", + json={"workflow": valid_workflow()}, + headers={"Origin": origin}, + ) + assert response.status_code >= 400 + assert _jobs(client) == [] + + @pytest.mark.xfail( + strict=True, + reason="reject_foreign_origins calls urlparse outside a try, so a " + "bracketed non-IPv6 Origin is a 500 from an unhandled ValueError " + "rather than the 403 every other refused Origin gets", + ) + def test_an_unparseable_origin_is_a_403_not_a_500(self, make_client): + app = make_client().app + with TestClient( + app, base_url="http://localhost", raise_server_exceptions=False + ) as client: + response = client.post( + "/api/jobs", + json={"workflow": valid_workflow()}, + headers={"Origin": "http://[::1].evil.example"}, + ) + assert response.status_code == 403 + + @pytest.mark.parametrize( + "host", + [ + "localhost.evil.example", + "127.0.0.1.evil.example", + "evil.example", + "evil.example:8765", + "127.0.0.2", + "0.0.0.0", + ], + ) + def test_a_lookalike_host_is_refused_on_a_loopback_bind(self, make_client, host): + with make_client() as client: + response = client.get("/api/health", headers={"Host": host}) + assert response.status_code == 400, host + + @pytest.mark.parametrize("path", ["/outputs/run.png", "/inputs/iris.png"]) + def test_the_host_check_covers_the_ungated_routes_too(self, make_client, path): + """DNS rebinding reads through whatever route answers: the static + routes need the Host check as much as /api does.""" + with make_client() as client: + assert client.get(path).status_code == 200 + assert ( + client.get(path, headers={"Host": "rebind.evil.example"}).status_code + == 400 + ) + + +# ------------------------------------------------ the ungated static routes + + +class TestUngatedRoutesServeOnlyTheirRoots: + """/outputs, /inputs and /exports take no token by design - an or + a download link cannot attach one. What they serve is therefore exactly + what anyone who can reach the port can read, so pin it: the workspace's + outputs, its asset search path, and finished exports, never a byte from + anywhere else.""" + + def test_what_they_serve_without_a_token(self, make_client): + with make_client(token=TOKEN) as client: + assert client.get("/outputs/run.png").content == b"generated" + assert client.get("/inputs/iris.png").content == b"asset" + assert client.get("/api/gallery").status_code == 401 + + @pytest.mark.parametrize( + "path", + [ + "/outputs/../outside/secret.txt", + "/outputs/..%2foutside/secret.txt", + "/outputs/%2e%2e/outside/secret.txt", + "/outputs/%2e%2e%2foutside%2fsecret.txt", + "/outputs/..%5coutside%5csecret.txt", + "/outputs/....//outside/secret.txt", + "/outputs//etc/hostname", + "/outputs/%2fetc%2fhostname", + "/outputs/~/secret.txt", + "/inputs/../outside/secret.png", + "/inputs/..%2foutside/secret.png", + "/inputs/%2e%2e%2f%2e%2e%2foutside%2fsecret.png", + "/inputs/..%5coutside%5csecret.png", + "/exports/..%2f..%2foutside.zip", + "/exports/%2e%2e.zip", + ], + ) + def test_traversal_spellings_serve_nothing_outside(self, make_client, path): + with make_client(token=TOKEN) as client: + response = client.get(path) + assert not _leaks(response), path + assert response.status_code in (400, 404, 405), (path, response.status_code) + + def test_an_absolute_path_in_the_name_is_not_joined_as_absolute( + self, make_client, layout + ): + """os.path.join(root, '/abs') is '/abs' - a name that starts with a + separator must still resolve under the root.""" + secret = layout["outside"] / "secret.txt" + with make_client(token=TOKEN) as client: + for route in ("/outputs", "/inputs"): + response = client.get(f"{route}/{secret}") + assert not _leaks(response), route + response = client.get(f"{route}/%2F{str(secret).lstrip('/')}") + assert not _leaks(response), route + + @pytest.mark.parametrize("workspace", ["..", "../outside", "%2e%2e", "/tmp", "."]) + def test_a_hostile_workspace_parameter_is_refused(self, make_client, workspace): + with make_client(token=TOKEN) as client: + response = client.get(f"/outputs/secret.txt?workspace={workspace}") + assert not _leaks(response) + assert response.status_code in (400, 404, 422) + + def test_the_gallery_download_route_is_contained_as_well(self, make_client): + with make_client(token=TOKEN) as client: + for name in ("..%2Foutside%2Fsecret.txt", "%2e%2e/outside/secret.txt"): + response = client.get(f"/api/gallery/{name}/download?token={TOKEN}") + assert not _leaks(response), name + + +# ------------------------------------------------------------- dw.serve CLI + + +@pytest.fixture +def serve(monkeypatch, tmp_path): + """dw.serve.main with the app factory, uvicorn and startup replaced.""" + import uvicorn + + import dw + import dw.serve as serve_module + from dw.server import app as app_module + + calls = {} + + def fake_create_app(**kwargs): + calls["create_app"] = kwargs + return object() + + monkeypatch.setattr(app_module, "create_app", fake_create_app) + monkeypatch.setattr(uvicorn, "run", lambda app, **kwargs: calls.update(ran=True)) + monkeypatch.setattr(dw, "startup", lambda *a, **k: calls.update(started=True)) + monkeypatch.delenv("DW_API_TOKEN", raising=False) + monkeypatch.setenv("DW_PROMPT_DIR", str(tmp_path / "prompts")) + monkeypatch.setenv("DW_ASSET_DIR", str(tmp_path / "assets")) + monkeypatch.setenv("DW_WORKSPACE", str(tmp_path / "workspace")) + monkeypatch.setenv("DW_WORKSPACE_SOURCE", "flag") + (tmp_path / "workflows").mkdir() + + def run(*argv): + calls.clear() + monkeypatch.setattr( + "sys.argv", + ["dw-serve", "--workflow-dir", str(tmp_path / "workflows"), *argv], + ) + serve_module.main() + return calls + + return run + + +class TestServeRefusesAnOpenMcpEndpoint: + @pytest.mark.parametrize( + "host", + [ + "0.0.0.0", + "::", + "10.0.0.2", + "192.168.1.5", + "gpu-box.local", + # loopback in fact, but not a name the check knows - the + # conservative direction, pinned so it stays conservative + "127.0.0.2", + "LOCALHOST", + ], + ) + def test_no_token_means_no_start(self, serve, host): + with pytest.raises(SystemExit) as exit_info: + serve("--mcp", "--host", host) + assert exit_info.value.code == 2 + + @pytest.mark.parametrize("token_args", [["--token", ""], []]) + def test_an_empty_token_is_no_token(self, serve, monkeypatch, token_args): + monkeypatch.setenv("DW_API_TOKEN", "") + with pytest.raises(SystemExit): + serve("--mcp", "--host", "0.0.0.0", *token_args) + + def test_nothing_starts_before_the_refusal(self, serve): + calls = {} + try: + calls = serve("--mcp", "--host", "0.0.0.0") + except SystemExit: + pass + assert "started" not in calls and "ran" not in calls + + def test_the_environment_token_satisfies_it_and_reaches_the_app( + self, serve, monkeypatch + ): + monkeypatch.setenv("DW_API_TOKEN", "from-env") + calls = serve("--mcp", "--host", "0.0.0.0") + assert calls["create_app"]["token"] == "from-env" + assert calls["create_app"]["mcp"] is True + + def test_the_flag_token_wins_over_the_environment(self, serve, monkeypatch): + monkeypatch.setenv("DW_API_TOKEN", "from-env") + calls = serve("--mcp", "--host", "0.0.0.0", "--token", "from-flag") + assert calls["create_app"]["token"] == "from-flag" From 851c051f716336e8c786229c705df58a232a64a8 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 24 Sep 2026 03:27:14 +0000 Subject: [PATCH 2/6] test(security): trust gate refuses before any import, and trust lets it through Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01NSpbAKgGb282ixhsEqGQ52 --- tests/test_security_trust_gate.py | 560 ++++++++++++++++++++++++++++++ 1 file changed, 560 insertions(+) create mode 100644 tests/test_security_trust_gate.py diff --git a/tests/test_security_trust_gate.py b/tests/test_security_trust_gate.py new file mode 100644 index 00000000..3b15f5f3 --- /dev/null +++ b/tests/test_security_trust_gate.py @@ -0,0 +1,560 @@ +"""The code-execution gate, from both sides, proven by what did *not* happen. + +tests/test_workflow_trust.py shows each gate raises. What it cannot show is +*when*: a refusal that arrives as UntrustedWorkflowError after the module was +imported, its source file opened, or the Hub asked for code has already let +the code run - "refused too late" is the same as not refused. So every +untrusted case here names a probe module that exists on sys.path and would +write a marker file the moment its top-level code ran, installs an import +hook that records any attempt to find it, and fails the test on either. + +The trusted half is the other boundary: `--trust-workflows` and +`DW_TRUST_WORKFLOWS` must let each surface through to a real import, or the +documented escape hatch is broken and operators reach for worse ones. + +tests/conftest.py trusts every test by default; each test here states the +posture it runs under. +""" + +import copy +import importlib.abc +import socket +import sys +import textwrap +from unittest.mock import MagicMock + +import pytest + +from dw.security import ( + TRUST_WORKFLOWS_ENV_VAR, + UntrustedWorkflowError, + set_trust_workflows, + workflows_are_trusted, +) + +PROBE = "dw_untrusted_probe_module" + + +class _ImportRecorder(importlib.abc.MetaPathFinder): + """First on sys.meta_path: sees every import that reaches the finders, + records the ones naming the probe, and lets the real finders answer.""" + + def __init__(self, watched): + self.watched = watched + self.attempts = [] + + def find_spec(self, fullname, path=None, target=None): + if fullname.split(".")[0] == self.watched: + self.attempts.append(fullname) + return None + + +@pytest.fixture +def probe(tmp_path, monkeypatch): + """A module that would run if anything imported it, and the evidence. + + Yields an object with `.attempts` (import lookups for the probe) and + `.executed()` (whether its top-level code ran). The module is on + sys.path for real, so a gate that failed open would import it rather + than fail on ModuleNotFoundError and look like a refusal. + """ + package = tmp_path / "probe_path" + package.mkdir() + marker = tmp_path / "EXECUTED" + (package / f"{PROBE}.py").write_text( + textwrap.dedent( + f""" + import pathlib + pathlib.Path({str(marker)!r}).write_text("ran") + + class Thing: + def __init__(self, *args, **kwargs): + self.args = args + self.kwargs = kwargs + + VALUE = 7 + """ + ) + ) + monkeypatch.syspath_prepend(str(package)) + for name in list(sys.modules): + if name == PROBE or name.startswith(PROBE + "."): + monkeypatch.delitem(sys.modules, name) + + recorder = _ImportRecorder(PROBE) + monkeypatch.setattr(sys, "meta_path", [recorder, *sys.meta_path]) + recorder.executed = marker.exists + yield recorder + sys.modules.pop(PROBE, None) + + +@pytest.fixture +def untrusted(monkeypatch): + monkeypatch.setenv(TRUST_WORKFLOWS_ENV_VAR, "0") + + +@pytest.fixture +def no_network(monkeypatch): + """Any attempt to resolve or dial is a failure, not a slow test.""" + attempts = [] + + def refuse(*args, **kwargs): + attempts.append(args) + raise AssertionError("network access attempted") + + monkeypatch.setattr(socket.socket, "connect", refuse) + monkeypatch.setattr(socket, "create_connection", refuse) + monkeypatch.setattr(socket, "getaddrinfo", refuse) + return attempts + + +def assert_nothing_happened(probe, no_network=()): + assert probe.attempts == [], f"import attempted before refusal: {probe.attempts}" + assert not probe.executed(), "the probe module's code ran" + assert list(no_network) == [], "the network was touched before refusal" + + +# ------------------------------------------------------------------ surfaces + + +def _realize(arguments): + """realize_args on a copy - it converts in place, and a parametrized + dict realized once would reach the next case already a class.""" + from dw.arguments import realize_args + + arguments = copy.deepcopy(arguments) + realize_args(arguments) + return arguments + + +TYPE_REFERENCES = [ + pytest.param({"scheduler_type": f"{PROBE}.Thing"}, id="_type"), + pytest.param({"component_type": f"{PROBE}.Thing"}, id="component_type"), + pytest.param({"torch_dtype": f"{PROBE}.Thing"}, id="_dtype"), + pytest.param({"dtype": f"{PROBE}.Thing"}, id="dtype"), + pytest.param( + { + "quantization_config": { + "configuration": {"config_type": f"{PROBE}.Thing"}, + "arguments": {}, + } + }, + id="config_type", + ), + pytest.param( + {"outer": {"inner": [{"weights_dtype": f"{PROBE}.Thing"}]}}, + id="nested", + ), + pytest.param( + {"scheduler_type": f"{PROBE}.sub.Thing"}, + id="submodule", + ), +] + + +def _pipeline(configuration=None, from_pretrained=None, **blocks): + from dw.pipeline_processors.pipeline import Pipeline + + definition = { + "configuration": configuration or {}, + "from_pretrained_arguments": from_pretrained or {"model_name": "a/b"}, + "arguments": {}, + **blocks, + } + return Pipeline(definition, 0, "cpu") + + +class TestUntrustedRefusesBeforeImport: + @pytest.mark.parametrize("arguments", TYPE_REFERENCES) + def test_a_dotted_type_reference(self, untrusted, probe, no_network, arguments): + with pytest.raises(UntrustedWorkflowError, match="trust-workflows"): + _realize(arguments) + assert_nothing_happened(probe, no_network) + + @pytest.mark.parametrize( + "modules", + [ + [PROBE], + [f"{PROBE}.sub"], + # a trusted entry first must not import before the untrusted + # one is looked at: every entry is checked, then any imported + ["json", PROBE], + ], + ) + def test_pre_load_modules(self, untrusted, probe, no_network, modules): + pipeline = _pipeline(configuration={"pre_load_modules": modules}) + imported_before = set(sys.modules) + with pytest.raises(UntrustedWorkflowError, match="pre_load_modules"): + pipeline.load(shared_components={}) + assert_nothing_happened(probe, no_network) + assert PROBE not in set(sys.modules) - imported_before + + def test_pre_load_modules_are_refused_by_the_preflight_too( + self, untrusted, probe, no_network + ): + pipeline = _pipeline(configuration={"pre_load_modules": [PROBE]}) + with pytest.raises(UntrustedWorkflowError): + pipeline.check_trusted() + assert_nothing_happened(probe, no_network) + + @pytest.mark.parametrize( + "reference", + [f"constant:{PROBE}.VALUE", f"constant:{PROBE}.sub.VALUE"], + ) + def test_a_constant_reference(self, untrusted, probe, no_network, reference): + from dw.arguments import fetch_constant + + with pytest.raises(UntrustedWorkflowError): + fetch_constant(reference) + assert_nothing_happened(probe, no_network) + + def test_a_constant_inside_arguments(self, untrusted, probe, no_network): + with pytest.raises(UntrustedWorkflowError): + _realize({"sigmas": f"constant:{PROBE}.VALUE"}) + assert_nothing_happened(probe, no_network) + + @pytest.mark.parametrize( + "extra", + [ + {"trust_remote_code": True}, + {"trust_remote_code": 1}, + {"trust_remote_code": "yes"}, + {"custom_pipeline": "someone/remote-pipeline"}, + {"custom_pipeline": "PROBE_PATH"}, + ], + ) + def test_remote_code_arguments(self, untrusted, probe, no_network, tmp_path, extra): + """Neither goes through our importlib, so the proof is that + from_pretrained - the thing that would fetch and import - is never + called, and no socket is opened.""" + from dw.pipeline_processors.pipeline import load_component + + if extra.get("custom_pipeline") == "PROBE_PATH": + # a local custom pipeline is a .py diffusers would import + extra = {"custom_pipeline": str(tmp_path / "probe_path")} + component_type = MagicMock() + component_type.__name__ = "ProbePipeline" + with pytest.raises(UntrustedWorkflowError): + load_component( + "pipeline", + {"component_type": component_type}, + {"model_name": "a/b", **extra}, + "cpu", + ) + component_type.from_pretrained.assert_not_called() + assert_nothing_happened(probe, no_network) + + @pytest.mark.parametrize("key", ["trust_remote_code", "custom_pipeline"]) + def test_remote_code_in_any_nested_block_is_seen_by_the_preflight( + self, untrusted, no_network, key + ): + """A component inside a list (controlnets, loras, text encoders) is + still a from_pretrained call the gate must see.""" + pipeline = _pipeline( + controlnets=[ + { + "configuration": {}, + "from_pretrained_arguments": {"model_name": "c/d", key: "x/y"}, + } + ], + ) + with pytest.raises(UntrustedWorkflowError, match=key): + pipeline.check_trusted() + assert list(no_network) == [] + + +class TestValidationRefusesBeforeImport: + """POST /api/validate and the pre-queue check call validation_errors - + a refusal there must not itself have imported the module to decide.""" + + def _workflow(self, pipeline, variables=None): + from dw.workflow import Workflow + + definition = { + "id": "trust_probe", + "steps": [ + { + "name": "gen", + "pipeline": pipeline, + "result": {"content_type": "image/png"}, + } + ], + } + if variables: + definition["variables"] = variables + return definition, Workflow + + # The pipeline itself is an escaped '{Fake}' rather than a real diffusers + # class: resolving a real one imports bitsandbytes, whose CPU backend + # asks the Hub for a kernel at import time - a network call these tests + # must not make, and nothing to do with the gate under test + @pytest.mark.parametrize( + "pipeline, path", + [ + pytest.param( + { + "configuration": {"component_type": f"{PROBE}.Thing"}, + "from_pretrained_arguments": {"model_name": "a/b"}, + "arguments": {}, + }, + "steps[0].pipeline.configuration.component_type", + id="component_type", + ), + pytest.param( + { + "configuration": {"component_type": "{Fake}"}, + "from_pretrained_arguments": {"model_name": "a/b"}, + "scheduler": { + "configuration": {"scheduler_type": f"{PROBE}.Thing"} + }, + "arguments": {}, + }, + "steps[0].pipeline.scheduler.configuration.scheduler_type", + id="scheduler_type", + ), + ], + ) + def test_a_dotted_type_is_a_validation_error( + self, untrusted, probe, no_network, tmp_path, pipeline, path + ): + definition, Workflow = self._workflow(pipeline) + errors = Workflow(definition, str(tmp_path), "").validation_errors() + refusals = [e for e in errors if "trust-workflows" in e["message"]] + assert [e["path"] for e in refusals] == [path] + assert_nothing_happened(probe, no_network) + + def test_a_constant_default_is_a_validation_error( + self, untrusted, probe, no_network, tmp_path + ): + definition, Workflow = self._workflow( + { + "configuration": {"component_type": "{Fake}"}, + "from_pretrained_arguments": {"model_name": "a/b"}, + "arguments": {"prompt": "variable:p"}, + }, + variables={"p": f"constant:{PROBE}.VALUE"}, + ) + errors = Workflow(definition, str(tmp_path), "").validation_errors() + assert [e["path"] for e in errors] == ["variables.p"] + assert_nothing_happened(probe, no_network) + + +class TestTheAllowlistIsNotAnEscapeHatch: + """In-ecosystem names are allowed untrusted by top-level package. That is + only safe if nothing reachable under those packages hands a workflow the + code execution the gate exists to deny.""" + + @pytest.mark.xfail( + strict=True, + reason="config_type is called with workflow kwargs and 'torch' is " + "allowlisted, so torch.hub.load(repo_or_dir=..., trust_repo=True) " + "runs a GitHub repo's hubconf.py untrusted", + ) + def test_config_type_cannot_name_a_code_loader_in_an_allowed_package( + self, untrusted, monkeypatch, no_network + ): + import torch.hub + + from dw.pipeline_processors.config_objects import create_quantization_config + + loader = MagicMock(name="torch.hub.load") + monkeypatch.setattr(torch.hub, "load", loader) + definition = { + "configuration": {"config_type": "torch.hub.load"}, + "arguments": { + "repo_or_dir": "attacker/repo", + "model": "anything", + "trust_repo": True, + }, + } + try: + realized = _realize({"quantization_config": definition}) + create_quantization_config(realized["quantization_config"]) + except UntrustedWorkflowError: + pass + loader.assert_not_called() + + @pytest.mark.xfail( + strict=True, + reason="load_constant_from_name walks attributes past the allowlisted " + "top-level module, so constant:torch.os.environ reads the server's " + "environment (DW_API_TOKEN, HF_TOKEN) untrusted", + ) + def test_a_constant_cannot_walk_out_of_an_allowed_package( + self, untrusted, monkeypatch + ): + from dw.arguments import fetch_constant + + monkeypatch.setenv("DW_API_TOKEN", "server-secret-probe") + leaked = None + try: + leaked = fetch_constant("constant:torch.os.environ") + except (UntrustedWorkflowError, ValueError): + pass + assert leaked is None or "server-secret-probe" not in str(leaked) + + @pytest.mark.parametrize( + "name", + [ + "os.system", + "builtins.eval", + "subprocess.Popen", + "importlib.import_module", + ".os", + " torch.os", + "torchx.Thing", + "diffusersx.Thing", + "dwx.Thing", + "Torch.nn.Linear", + ], + ) + def test_near_misses_of_an_allowed_name_are_refused(self, untrusted, name): + from dw.type_helpers import load_type_from_full_name + + with pytest.raises(UntrustedWorkflowError): + load_type_from_full_name(name) + + +# ------------------------------------------------------------------ trusted + + +@pytest.fixture(params=["env", "flag"]) +def trusted(request, monkeypatch): + """Both ways a process ends up trusting workflows: the environment + variable a spawned worker inherits, and the call the CLI flag makes.""" + if request.param == "env": + monkeypatch.setenv(TRUST_WORKFLOWS_ENV_VAR, "1") + else: + monkeypatch.setenv(TRUST_WORKFLOWS_ENV_VAR, "0") + set_trust_workflows(True) + assert workflows_are_trusted() + return request.param + + +class TestTrustedLetsEachSurfaceThrough: + @pytest.mark.parametrize( + "arguments, find", + [ + ({"scheduler_type": f"{PROBE}.Thing"}, lambda a: a["scheduler_type"]), + ({"torch_dtype": f"{PROBE}.Thing"}, lambda a: a["torch_dtype"]), + ({"dtype": f"{PROBE}.Thing"}, lambda a: a["dtype"]), + ], + ) + def test_a_dotted_type_imports(self, trusted, probe, arguments, find): + realized = _realize(arguments) + assert find(realized).__name__ == "Thing" + assert probe.executed() + + def test_a_config_type_imports_and_builds(self, trusted, probe): + from dw.pipeline_processors.config_objects import create_quantization_config + + definition = { + "configuration": {"config_type": f"{PROBE}.Thing"}, + "arguments": {"bits": 4}, + } + realized = _realize({"quantization_config": definition}) + built = create_quantization_config(realized["quantization_config"]) + assert built.kwargs == {"bits": 4} + assert probe.executed() + + def test_pre_load_modules_import(self, trusted, probe, monkeypatch): + from dw.pipeline_processors.pipeline import Pipeline + + pipeline = _pipeline(configuration={"pre_load_modules": [PROBE]}) + monkeypatch.setattr( + Pipeline, + "populate_from_pretrained_arguments", + MagicMock(side_effect=RuntimeError("stop - past the gate")), + ) + with pytest.raises(RuntimeError, match="stop"): + pipeline.load(shared_components={}) + assert PROBE in probe.attempts + assert probe.executed() + + def test_a_constant_reads(self, trusted, probe): + from dw.arguments import fetch_constant + + assert fetch_constant(f"constant:{PROBE}.VALUE") == 7 + assert probe.executed() + + @pytest.mark.parametrize( + "extra", + [{"trust_remote_code": True}, {"custom_pipeline": "someone/pipeline"}], + ) + def test_remote_code_arguments_reach_from_pretrained(self, trusted, extra): + from dw.pipeline_processors.pipeline import load_component + + component_type = MagicMock() + component_type.__name__ = "ProbePipeline" + load_component( + "pipeline", + {"component_type": component_type}, + {"model_name": "a/b", **extra}, + "cpu", + ) + component_type.from_pretrained.assert_called_once() + _, kwargs = component_type.from_pretrained.call_args + for key, value in extra.items(): + assert kwargs[key] == value + + +class TestHowAProcessBecomesTrusted: + """Only the flag trusts: an unset or unrecognized variable is untrusted, + and a server started without --trust-workflows does not inherit trust + from whatever launched it.""" + + @pytest.mark.parametrize("value", ["", "0", "true", "yes", "TRUE", " 1", "1 "]) + def test_only_exactly_1_trusts(self, monkeypatch, value): + monkeypatch.setenv(TRUST_WORKFLOWS_ENV_VAR, value) + assert workflows_are_trusted() is False + + @pytest.fixture + def serve(self, monkeypatch, tmp_path): + import uvicorn + + import dw + import dw.serve as serve_module + from dw.server import app as app_module + + monkeypatch.setattr(app_module, "create_app", lambda **kwargs: object()) + monkeypatch.setattr(uvicorn, "run", lambda app, **kwargs: None) + monkeypatch.setattr(dw, "startup", lambda *args, **kwargs: None) + monkeypatch.delenv("DW_API_TOKEN", raising=False) + monkeypatch.setenv("DW_PROMPT_DIR", str(tmp_path / "prompts")) + monkeypatch.setenv("DW_ASSET_DIR", str(tmp_path / "assets")) + monkeypatch.setenv("DW_WORKSPACE", str(tmp_path / "workspace")) + monkeypatch.setenv("DW_WORKSPACE_SOURCE", "flag") + (tmp_path / "workflows").mkdir() + + def run(*argv): + monkeypatch.setattr( + "sys.argv", + ["dw-serve", "--workflow-dir", str(tmp_path / "workflows"), *argv], + ) + serve_module.main() + return workflows_are_trusted() + + return run + + def test_serve_without_the_flag_overrides_an_inherited_1(self, serve, monkeypatch): + monkeypatch.setenv(TRUST_WORKFLOWS_ENV_VAR, "1") + assert serve() is False + + def test_serve_with_the_flag_trusts(self, serve, monkeypatch): + monkeypatch.setenv(TRUST_WORKFLOWS_ENV_VAR, "0") + assert serve("--trust-workflows") is True + + def test_validate_cli_without_the_flag_is_untrusted(self, monkeypatch, tmp_path): + """dw.validate is how an operator vets a file before running it - + it must vet it under the posture the run will have.""" + import dw.validate as validate_module + + workflow = tmp_path / "w.json" + workflow.write_text('{"id": "w", "steps": []}') + monkeypatch.setenv(TRUST_WORKFLOWS_ENV_VAR, "1") + monkeypatch.setattr("sys.argv", ["dw-validate", str(workflow)]) + try: + validate_module.main() + except SystemExit: + pass + assert workflows_are_trusted() is False From 5b651869665588fa7bb93cca82a1ce91f7664d46 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 24 Sep 2026 03:27:14 +0000 Subject: [PATCH 3/6] test(security): planted symlinks and download_output destinations Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01NSpbAKgGb282ixhsEqGQ52 --- tests/test_security_symlinks.py | 638 ++++++++++++++++++++++++++++++++ 1 file changed, 638 insertions(+) create mode 100644 tests/test_security_symlinks.py diff --git a/tests/test_security_symlinks.py b/tests/test_security_symlinks.py new file mode 100644 index 00000000..8a51c974 --- /dev/null +++ b/tests/test_security_symlinks.py @@ -0,0 +1,638 @@ +"""A symlink planted inside a workspace must not carry anything out of it. + +Every root the server works in - outputs, assets, the shared asset library, +workflows, prompts, exports - is a directory a local user (or a mounted +volume, or an unpacked archive) can put a symlink into. `validate_path` +resolves symlinks before its containment check, so a path that is *checked* +is safe; the question this file asks is which code paths reach the disk +without being checked. For each root: read, write, delete and enumeration +through a link pointing at a sibling directory the server was never told +about, and the same for `download_output`'s destination over a mounted MCP +endpoint (gap 7 of the live suite's "Not covered here" list). + +Everything is under `tmp_path`: "outside" is a sibling of the workspace, and +the "secret" is a marker string, so a failing boundary leaks a marker and +overwrites a scratch file. +""" + +import io +import json +import os +import zipfile + +import httpx +import pytest +from fastapi.testclient import TestClient +from PIL import Image + +from dw.security import TRUST_WORKFLOWS_ENV_VAR +from dw.server.app import create_app +from dw.server.jobs import JobManager + +from .test_server import ScriptedWorkerManager, success_script, valid_workflow + +SECRET = "outside-the-roots-probe" + +pytestmark = pytest.mark.skipif( + not hasattr(os, "symlink") or os.name == "nt", + reason="symlinks need POSIX semantics", +) + + +def _png(path, color="red"): + Image.new("RGB", (4, 4), color).save(path) + return path + + +@pytest.fixture +def tree(tmp_path): + """/ws is the workspace root (default workspace), /outside the + directory nothing may reach. Each outside file carries SECRET.""" + root = tmp_path / "ws" + paths = { + "root": root, + "workflows": root / "workflows", + "outputs": root / "outputs", + "assets": root / "assets", + "prompts": root / "prompts", + "common": root / "common" / "assets", + "outside": tmp_path / "outside", + } + for path in paths.values(): + path.mkdir(parents=True, exist_ok=True) + outside = paths["outside"] + _png(outside / "secret.png") + (outside / "secret.txt").write_text(SECRET) + (outside / "secret.json").write_text( + json.dumps( + { + "id": "secret", + "description": SECRET, + "variables": {SECRET.replace("-", "_"): 1}, + "steps": [], + } + ) + ) + (outside / "prompt.json").write_text(json.dumps({"text": SECRET})) + (outside / "victim.txt").write_text("untouched") + (outside / "dir").mkdir() + _png(outside / "dir" / "inner.png") + (outside / "dir" / "inner.txt").write_text(SECRET) + (paths["workflows"] / "Basic.json").write_text(json.dumps(valid_workflow("b"))) + return paths + + +def link(at, target): + at.parent.mkdir(parents=True, exist_ok=True) + os.symlink(target, at) + return at + + +@pytest.fixture +def client(tree, tmp_path, monkeypatch): + monkeypatch.setenv(TRUST_WORKFLOWS_ENV_VAR, "0") + manager = JobManager( + str(tree["outputs"]), + worker_manager=ScriptedWorkerManager(success_script), + history_path=str(tmp_path / "jobs.sqlite"), + ) + app = create_app( + workflow_dir=str(tree["workflows"]), + output_dir=str(tree["outputs"]), + job_manager=manager, + prompt_dir=str(tree["prompts"]), + asset_dir=str(tree["assets"]), + workspace=str(tree["root"]), + ) + with TestClient(app, base_url="http://localhost") as test_client: + yield test_client + + +def _leaks(response): + return response.status_code < 300 and SECRET.encode() in response.content + + +def _outside_untouched(tree): + outside = tree["outside"] + assert (outside / "victim.txt").read_text() == "untouched" + assert (outside / "secret.txt").read_text() == SECRET + assert (outside / "dir" / "inner.txt").read_text() == SECRET + assert (outside / "secret.png").exists() + assert sorted(p.name for p in outside.iterdir()) == [ + "dir", + "prompt.json", + "secret.json", + "secret.png", + "secret.txt", + "victim.txt", + ], "something was written into the outside directory" + + +# ------------------------------------------------------------------ outputs + + +class TestOutputs: + @pytest.fixture(autouse=True) + def plant(self, tree): + link(tree["outputs"] / "leak.png", tree["outside"] / "secret.png") + link(tree["outputs"] / "leak.txt", tree["outside"] / "secret.txt") + link(tree["outputs"] / "linked_run", tree["outside"] / "dir") + + @pytest.mark.parametrize( + "path", + [ + "/outputs/leak.txt", + "/outputs/leak.png", + "/outputs/linked_run/inner.txt", + "/api/gallery/leak.txt/download", + "/api/gallery/leak.png/download", + "/api/gallery/linked_run/inner.txt/download", + "/api/gallery/leak.png/metadata", + "/api/gallery/leak.png/thumbnail", + "/api/gallery/linked_run/inner.png/thumbnail", + ], + ) + def test_reads_do_not_follow_the_link(self, client, path): + response = client.get(path) + assert not _leaks(response), path + assert response.status_code >= 400, (path, response.status_code) + + def test_the_archive_route_does_not_follow_the_link(self, client): + response = client.post("/api/gallery/archive", json={"names": ["leak.txt"]}) + assert response.status_code >= 400 + assert not _leaks(response) + + @pytest.mark.xfail( + strict=True, + reason="_iter_gallery_files walks outputs with os.walk and os.stat, " + "so GET /api/gallery lists a symlink pointing outside, with the " + "target's size and mtime", + ) + def test_the_gallery_listing_does_not_enumerate_the_link(self, client): + names = [entry["name"] for entry in client.get("/api/gallery").json()["files"]] + assert "leak.png" not in names + + def test_delete_removes_nothing_outside(self, client, tree): + for name in ("leak.txt", "linked_run/inner.txt", "linked_run"): + client.delete(f"/api/gallery/{name}") + _outside_untouched(tree) + + def test_keep_output_does_not_copy_the_target_into_assets(self, client, tree): + response = client.post( + "/api/assets/keep", json={"name": "leak.png", "asset_name": "kept.png"} + ) + assert response.status_code >= 400 + assert not (tree["assets"] / "kept.png").exists() + + def test_an_output_reference_does_not_follow_a_linked_run_directory(self, tree): + from dw.runs import resolve_output_reference + from dw.security import SecurityError + + with pytest.raises((SecurityError, ValueError)): + resolve_output_reference( + "output:linked_run/inner.png", root=str(tree["outputs"]) + ) + + +# ------------------------------------------------------------------- assets + + +class TestAssets: + @pytest.fixture(autouse=True) + def plant(self, tree): + link(tree["assets"] / "leak.png", tree["outside"] / "secret.png") + link(tree["assets"] / "cast", tree["outside"] / "dir") + link(tree["common"] / "shared-leak.png", tree["outside"] / "secret.png") + link(tree["common"] / "shared-dir", tree["outside"] / "dir") + + @pytest.mark.parametrize( + "path", + [ + "/inputs/leak.png", + "/inputs/cast/inner.txt", + "/inputs/cast/inner.png", + "/inputs/shared-leak.png", + "/inputs/shared-dir/inner.txt", + ], + ) + def test_reads_do_not_follow_the_link(self, client, path): + response = client.get(path) + assert response.status_code >= 400, (path, response.status_code) + + def test_the_asset_archive_does_not_follow_the_link(self, client): + for names in (["leak.png"], ["cast/inner.txt"], ["shared-dir/inner.txt"]): + response = client.post("/api/assets/archive", json={"names": names}) + assert response.status_code >= 400, names + assert not _leaks(response) + + @pytest.mark.parametrize( + "reference", + [ + "asset:leak.png", + "asset:cast/inner.png", + "asset:shared-leak.png", + "asset:shared-dir/inner.png", + ], + ) + def test_an_asset_reference_does_not_follow_the_link( + self, tree, monkeypatch, reference + ): + from dw.assets import resolve_asset_reference + from dw.security import SecurityError + + monkeypatch.setenv(TRUST_WORKFLOWS_ENV_VAR, "0") + monkeypatch.setenv("DW_ASSET_DIR", str(tree["assets"])) + monkeypatch.setenv("DW_ASSET_PATH", str(tree["common"])) + with pytest.raises((SecurityError, ValueError)): + resolve_asset_reference(reference) + + @pytest.mark.xfail( + strict=True, + reason="the asset listing walks the library with os.walk/os.stat and " + "lists a symlink that resolves outside it", + ) + def test_the_asset_listing_does_not_enumerate_the_link(self, client): + listing = client.get("/api/assets").json() + assert "leak.png" not in json.dumps(listing) + + def test_an_upload_does_not_write_through_a_linked_name(self, client, tree): + """Uploads land in /uploads/, so that is where a link that + could redirect one would be planted.""" + link(tree["assets"] / "uploads" / "victim.png", tree["outside"] / "victim.txt") + response = client.post( + "/api/uploads?filename=victim.png&asset_name=victim.png", + content=b"\x89PNG overwritten", + ) + assert response.status_code >= 400 + _outside_untouched(tree) + + def test_an_upload_does_not_write_into_a_linked_folder(self, client, tree): + link(tree["assets"] / "uploads" / "cast", tree["outside"] / "dir") + response = client.post( + "/api/uploads?filename=new.png&asset_name=cast/new.png", + content=b"\x89PNG new", + ) + assert response.status_code >= 400 + assert not (tree["outside"] / "dir" / "new.png").exists() + _outside_untouched(tree) + + def test_a_shared_upload_does_not_write_through_a_link_in_common( + self, client, tree + ): + link(tree["common"] / "uploads" / "victim.png", tree["outside"] / "victim.txt") + response = client.post( + "/api/uploads?filename=victim.png&asset_name=victim.png&shared=true", + content=b"\x89PNG overwritten", + ) + assert response.status_code >= 400 + _outside_untouched(tree) + + def test_keep_output_does_not_write_through_a_linked_destination( + self, client, tree + ): + _png(tree["outputs"] / "real.png", "blue") + link(tree["assets"] / "dest.png", tree["outside"] / "victim.txt") + for overwrite in (False, True): + client.post( + "/api/assets/keep", + json={ + "name": "real.png", + "asset_name": "dest.png", + "overwrite": overwrite, + }, + ) + client.post( + "/api/assets/keep", + json={ + "name": "real.png", + "asset_name": "cast/new.png", + "overwrite": overwrite, + }, + ) + _outside_untouched(tree) + + def test_delete_removes_nothing_outside(self, client, tree): + for name in ("leak.png", "cast/inner.png", "cast", "shared-leak.png"): + client.delete(f"/api/assets/{name}") + _outside_untouched(tree) + + +# ------------------------------------------------------ gather_images globs + + +class TestGlobs: + """tests/test_locations.py drops a *file* match that escapes; a linked + *directory* in the pattern's fixed part is refused before expansion.""" + + @pytest.fixture(autouse=True) + def untrusted(self, monkeypatch, tree): + monkeypatch.setenv(TRUST_WORKFLOWS_ENV_VAR, "0") + monkeypatch.setenv("DW_ASSET_DIR", str(tree["assets"])) + link(tree["assets"] / "cast", tree["outside"] / "dir") + link(tree["assets"] / "leak.png", tree["outside"] / "secret.png") + _png(tree["assets"] / "own.png") + + def test_a_linked_directory_in_the_pattern_is_refused(self, tree): + from dw.security import PathTraversalError + from dw.tasks.gather import gather_images + + with pytest.raises(PathTraversalError): + gather_images(glob=str(tree["assets"] / "cast" / "*.png")) + + def test_a_wildcard_does_not_descend_through_a_linked_directory(self, tree): + """The only match is under the linked directory; dropped, it leaves + nothing, which gather_images reports as no images.""" + from dw.tasks.gather import gather_images + + with pytest.raises(ValueError, match="No images found"): + gather_images(glob=str(tree["assets"] / "*" / "*.png")) + + def test_a_linked_file_match_is_dropped(self, tree): + from dw.tasks.gather import gather_images + + images = gather_images(glob=str(tree["assets"] / "*.png")) + assert len(images) == 1 + + def test_gather_videos_drops_a_linked_match(self, tree): + from dw.tasks.gather import gather_videos + + (tree["outside"] / "clip.mp4").write_bytes(b"not really a video") + link(tree["assets"] / "clip.mp4", tree["outside"] / "clip.mp4") + # a decode error here would mean the linked file was opened + with pytest.raises(ValueError, match="No videos found"): + gather_videos(glob=str(tree["assets"] / "*.mp4")) + + +# ------------------------------------------------------- workflows, prompts + + +class TestWorkflows: + @pytest.fixture(autouse=True) + def plant(self, tree): + link(tree["workflows"] / "leak.json", tree["outside"] / "secret.json") + link(tree["workflows"] / "linked", tree["outside"]) + + @pytest.mark.parametrize( + "path", + [ + "/api/workflows/leak", + "/api/workflows/leak/download", + "/api/workflows/leak/variables", + "/api/workflows/linked/secret", + ], + ) + def test_reads_do_not_follow_the_link(self, client, path): + response = client.get(path) + assert not _leaks(response), path + + @pytest.mark.xfail( + strict=True, + reason="workflow_details opens every *.json os.walk finds without a " + "containment check, so GET /api/workflows reads a linked file's " + "description and variable names", + ) + def test_the_listing_does_not_read_through_the_link(self, client): + response = client.get("/api/workflows") + assert SECRET not in response.text + assert SECRET.replace("-", "_") not in response.text + + def test_a_save_does_not_write_through_the_link(self, client, tree): + for name in ("leak", "linked/victim"): + client.put(f"/api/workflows/{name}", json=valid_workflow("overwrite")) + assert json.loads((tree["outside"] / "secret.json").read_text())["id"] == ( + "secret" + ) + _outside_untouched(tree) + + def test_a_delete_removes_nothing_outside(self, client, tree): + for name in ("leak", "linked/secret"): + client.delete(f"/api/workflows/{name}") + _outside_untouched(tree) + + +class TestPrompts: + @pytest.fixture(autouse=True) + def plant(self, tree): + link(tree["prompts"] / "leak.json", tree["outside"] / "prompt.json") + link(tree["prompts"] / "linked", tree["outside"]) + + @pytest.mark.parametrize( + "path", ["/api/prompts/leak", "/api/prompts/leak/download"] + ) + def test_reads_do_not_follow_the_link(self, client, path): + assert not _leaks(client.get(path)), path + + @pytest.mark.xfail( + strict=True, + reason="list_prompts hands every *.json workflow_names finds to " + "prompt_details, which opens it without a containment check - the " + "listing carries a linked file's text", + ) + def test_the_listing_does_not_read_through_the_link(self, client): + assert SECRET not in client.get("/api/prompts").text + + def test_a_prompt_reference_does_not_follow_the_link(self, tree, monkeypatch): + from dw.prompts import fetch_prompt + from dw.security import SecurityError + + for reference in ("prompt:leak", "prompt:linked/prompt"): + with pytest.raises((SecurityError, ValueError)): + fetch_prompt(reference, prompt_dir=str(tree["prompts"])) + + def test_a_save_does_not_write_through_the_link(self, client, tree): + for name in ("leak", "linked/victim"): + client.put(f"/api/prompts/{name}", json={"prompt": {"text": "overwrite"}}) + assert json.loads((tree["outside"] / "prompt.json").read_text()) == { + "text": SECRET + } + _outside_untouched(tree) + + def test_a_delete_removes_nothing_outside(self, client, tree): + for name in ("leak", "linked/prompt"): + client.delete(f"/api/prompts/{name}") + assert (tree["outside"] / "prompt.json").exists() + _outside_untouched(tree) + + +# -------------------------------------------------------- exports, workspaces + + +class TestExportsAndWorkspaces: + @pytest.mark.xfail( + strict=True, + reason="GET /exports/.zip (no token) zips the export directory " + "with os.walk + ZipFile.write, which follows a planted file symlink " + "and archives the target's bytes", + ) + def test_the_export_zip_does_not_follow_a_planted_link(self, client, tree): + export = tree["root"] / "exports" / "job-1" + export.mkdir(parents=True) + (export / "README.md").write_text("an export") + link(export / "leak.txt", tree["outside"] / "secret.txt") + + response = client.get("/exports/job-1.zip") + assert response.status_code == 200 + archive = zipfile.ZipFile(io.BytesIO(response.content)) + for name in archive.namelist(): + assert SECRET.encode() not in archive.read(name), name + + def test_the_export_zip_does_not_follow_a_linked_export_directory( + self, client, tree + ): + link(tree["root"] / "exports" / "job-2", tree["outside"] / "dir") + response = client.get("/exports/job-2.zip") + assert response.status_code == 404 + assert SECRET.encode() not in response.content + + def test_deleting_a_workspace_does_not_follow_a_link_inside_it(self, client, tree): + assert ( + client.post("/api/workspaces", json={"name": "doomed"}).status_code == 201 + ) + doomed = tree["root"] / "doomed" + link(doomed / "outputs" / "out", tree["outside"]) + link(doomed / "assets" / "cast", tree["outside"] / "dir") + client.delete("/api/workspaces/doomed?acknowledged=true") + _outside_untouched(tree) + + def test_the_workspace_size_count_does_not_follow_a_link(self, client, tree): + assert client.post("/api/workspaces", json={"name": "sized"}).status_code == 201 + link(tree["root"] / "sized" / "outputs" / "big", tree["outside"]) + response = client.delete("/api/workspaces/sized") + contents = response.json()["detail"]["contents"] + assert all(entry.get("files", 0) == 0 for entry in contents.values()), contents + + +# ------------------------------------------------------ download_output (7) + + +def _mounted(workspace_root, body=b"downloaded-bytes"): + """A DwClient shaped like the one dw.serve builds for its own /mcp.""" + from dw_mcp.client import DwClient + + def handler(request): + if request.url.path == "/api/server": + return httpx.Response( + 200, json={"directories": {"workspace": str(workspace_root)}} + ) + return httpx.Response(200, content=body, headers={"content-type": "image/png"}) + + client = DwClient(transport=httpx.MockTransport(handler)) + client.mounted = True + return client + + +class TestDownloadOutputDestination: + """Over a mounted endpoint the file lands on the server, so the + destination is confined to the workspace (#113). A legal relative + destination writes inside it and nowhere else.""" + + @pytest.fixture + def workspace(self, tree): + return tree["root"] + + def _download(self, workspace, destination, overwrite=False): + from dw_mcp.media import download_output + + return download_output( + _mounted(workspace), + "run/probe.png", + destination=destination, + overwrite=overwrite, + ) + + @pytest.mark.parametrize( + "destination", + ["kept/probe.png", "probe.png", "kept/", "deep/er/still/probe.png"], + ) + def test_a_relative_destination_lands_inside(self, workspace, tree, destination): + before = {p for p in tree["outside"].rglob("*")} + result = self._download(workspace, destination) + saved = os.path.realpath(result["saved_to"]) + assert saved.startswith(os.path.realpath(workspace) + os.sep) + assert open(saved, "rb").read() == b"downloaded-bytes" + assert {p for p in tree["outside"].rglob("*")} == before + + @pytest.mark.parametrize( + "destination", + [ + "../outside/escaped.png", + "kept/../../outside/escaped.png", + "OUTSIDE_ABSOLUTE", + "~/escaped.png", + ], + ) + def test_dot_dot_absolute_and_home_are_refused( + self, workspace, tree, monkeypatch, destination + ): + from dw_mcp.client import DwApiError + + monkeypatch.setenv("HOME", str(tree["outside"])) + if destination == "OUTSIDE_ABSOLUTE": + destination = str(tree["outside"] / "escaped.png") + with pytest.raises(DwApiError): + self._download(workspace, destination) + assert not (tree["outside"] / "escaped.png").exists() + _outside_untouched(tree) + + @pytest.mark.parametrize("overwrite", [False, True]) + def test_a_linked_parent_is_refused(self, workspace, tree, overwrite): + from dw_mcp.client import DwApiError + + link(workspace / "kept", tree["outside"]) + with pytest.raises(DwApiError): + self._download(workspace, "kept/escaped.png", overwrite=overwrite) + with pytest.raises(DwApiError): + self._download(workspace, "kept/victim.txt", overwrite=overwrite) + _outside_untouched(tree) + + def test_a_linked_grandparent_is_refused_even_when_the_parent_is_missing( + self, workspace, tree + ): + """The parent does not exist yet, so realpath of the destination + alone would stop at the missing segment - the check has to resolve + the nearest existing ancestor, which is the link.""" + from dw_mcp.client import DwApiError + + link(workspace / "kept", tree["outside"]) + with pytest.raises(DwApiError): + self._download(workspace, "kept/new/deeper/escaped.png") + assert not (tree["outside"] / "new").exists() + _outside_untouched(tree) + + def test_overwrite_cannot_clobber_a_linked_file(self, workspace, tree): + from dw_mcp.client import DwApiError + + link(workspace / "victim.png", tree["outside"] / "victim.txt") + with pytest.raises(DwApiError): + self._download(workspace, "victim.png", overwrite=True) + _outside_untouched(tree) + + def test_a_dangling_link_is_replaced_not_followed(self, workspace, tree): + """A link to a file that does not exist yet: writing through it + would create the file outside. The write must replace the link (or + refuse), never create its target.""" + target = tree["outside"] / "created-by-download.png" + link(workspace / "dangling.png", target) + try: + self._download(workspace, "dangling.png", overwrite=True) + except Exception: + pass + assert not target.exists() + _outside_untouched(tree) + + def test_a_workspace_root_that_is_itself_a_link_still_confines( + self, tree, tmp_path + ): + """The operator's workspace may be a symlink (a data volume); the + root is resolved, and a destination is confined to what it resolves + to - not to the link's own parent.""" + from dw_mcp.client import DwApiError + + linked_root = link(tmp_path / "ws-link", tree["root"]) + result = self._download(linked_root, "kept/probe.png") + assert os.path.realpath(result["saved_to"]).startswith( + os.path.realpath(tree["root"]) + os.sep + ) + with pytest.raises(DwApiError): + self._download(linked_root, str(tmp_path / "outside" / "escaped.png")) + _outside_untouched(tree) From 066d20659e088d8d6c7917b56166f2a2736c0847 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 24 Sep 2026 03:27:14 +0000 Subject: [PATCH 4/6] test(security): decoder bombs through get_output_image and gallery routes Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01NSpbAKgGb282ixhsEqGQ52 --- tests/test_security_decoder_bombs.py | 195 +++++++++++++++++++++++++++ 1 file changed, 195 insertions(+) create mode 100644 tests/test_security_decoder_bombs.py diff --git a/tests/test_security_decoder_bombs.py b/tests/test_security_decoder_bombs.py new file mode 100644 index 00000000..d851ae9b --- /dev/null +++ b/tests/test_security_decoder_bombs.py @@ -0,0 +1,195 @@ +"""A crafted PNG that is small on disk and enormous once decoded. + +`get_output_image` (MCP) and the gallery thumbnail route both decode a whole +image a caller names. A PNG header can claim any size; the pixels are a +zlib stream that compresses a flat colour about a thousand to one. Pillow's +own guard (`Image.MAX_IMAGE_PIXELS`, ~89.5M) raises only above *twice* that +and only warns in between - so without a clamp of dw's own, an image just +under 179M pixels is decoded in full: half a gigabyte of RGB per request. + +The probe never lets a decode happen, whatever the boundary does: the +headers here are real, but `ImageFile.load` is replaced for the test with +one that records the attempt and raises before allocating. A test fails when +that record is non-empty - "refused" means refused *before* the decode. +""" + +import struct +import zlib + +import httpx +import pytest +from fastapi.testclient import TestClient +from PIL import Image, ImageFile + +from dw.server.app import create_app +from dw.server.jobs import JobManager +from dw_mcp.client import DwApiError, DwClient +from dw_mcp.media import get_output_image + +from .test_server import ScriptedWorkerManager, success_script + +# Well above anything the tool legitimately returns (a 4K frame is 8.3M) and +# well below Pillow's default, where only dw's own clamp can stop it +UNDER_PILLOWS_ERROR = (12_000, 12_000) # 144M pixels: Pillow only warns +OVER_PILLOWS_ERROR = (20_000, 20_000) # 400M pixels: Pillow refuses on open +DECODE_LIMIT = 50_000_000 + + +def _chunk(kind, data): + body = kind + data + return struct.pack(">I", len(data)) + body + struct.pack(">I", zlib.crc32(body)) + + +def bomb_png(width, height): + """A syntactically valid PNG header claiming width x height grayscale + pixels, with a token IDAT. Tiny on disk; never decoded by this file.""" + header = struct.pack(">IIBBBBB", width, height, 8, 0, 0, 0, 0) + return ( + b"\x89PNG\r\n\x1a\n" + + _chunk(b"IHDR", header) + + _chunk(b"IDAT", zlib.compress(b"\x00" * 64)) + + _chunk(b"IEND", b"") + ) + + +@pytest.fixture +def decodes(monkeypatch): + """Every attempt to decode an image over DECODE_LIMIT pixels, recorded + and stopped before the allocation.""" + attempts = [] + original = ImageFile.ImageFile.load + + def guarded(self): + width, height = self.size + if width * height > DECODE_LIMIT: + attempts.append(self.size) + raise MemoryError("probe: a decompression bomb was decoded") + return original(self) + + monkeypatch.setattr(ImageFile.ImageFile, "load", guarded) + return attempts + + +def test_the_probe_png_parses_to_the_size_it_claims(): + """The header is honest enough for Pillow to believe it - otherwise the + tests below would pass on a parse error rather than a clamp.""" + import io + import warnings + + with warnings.catch_warnings(): + warnings.simplefilter("ignore", Image.DecompressionBombWarning) + with Image.open(io.BytesIO(bomb_png(*UNDER_PILLOWS_ERROR))) as image: + assert image.size == UNDER_PILLOWS_ERROR + + +def test_pillows_guard_has_not_been_disabled(): + """Everything above 2 x MAX_IMAGE_PIXELS rests on Pillow's default; a + module that set it to None would open every size below.""" + assert Image.MAX_IMAGE_PIXELS is not None + assert Image.MAX_IMAGE_PIXELS <= 100_000_000 + + +def _serving(body): + def handler(request): + return httpx.Response(200, content=body, headers={"content-type": "image/png"}) + + return DwClient(transport=httpx.MockTransport(handler)) + + +class TestGetOutputImage: + def test_a_bomb_over_pillows_limit_is_refused_before_decode(self, decodes): + with pytest.raises((DwApiError, Image.DecompressionBombError)): + get_output_image(_serving(bomb_png(*OVER_PILLOWS_ERROR)), "bomb.png") + assert decodes == [] + + @pytest.mark.xfail( + strict=True, + reason="no pixel clamp - gap: get_output_image calls image.load() on " + "whatever Pillow opens, and Pillow only warns below 2x MAX_IMAGE_PIXELS", + ) + def test_a_bomb_under_pillows_limit_is_refused_before_decode(self, decodes): + try: + get_output_image(_serving(bomb_png(*UNDER_PILLOWS_ERROR)), "bomb.png") + except DwApiError: + pass + assert decodes == [] + + @pytest.mark.xfail( + strict=True, + reason="no pixel clamp - gap: a crop is cut from the fully decoded " + "image, so a small crop still decodes the whole bomb", + ) + def test_a_crop_does_not_decode_the_whole_bomb(self, decodes): + try: + get_output_image( + _serving(bomb_png(*UNDER_PILLOWS_ERROR)), "bomb.png", crop=[0, 0, 8, 8] + ) + except DwApiError: + pass + assert decodes == [] + + +@pytest.fixture +def gallery(tmp_path): + outputs = tmp_path / "outputs" + outputs.mkdir() + (tmp_path / "workflows").mkdir() + manager = JobManager( + str(outputs), + worker_manager=ScriptedWorkerManager(success_script), + history_path=str(tmp_path / "jobs.sqlite"), + ) + app = create_app( + workflow_dir=str(tmp_path / "workflows"), + output_dir=str(outputs), + job_manager=manager, + prompt_dir=str(tmp_path / "prompts"), + ) + client = TestClient(app, base_url="http://localhost", raise_server_exceptions=False) + return client, outputs + + +class TestGalleryThumbnail: + def test_a_bomb_over_pillows_limit_is_refused_before_decode(self, gallery, decodes): + client, outputs = gallery + (outputs / "bomb.png").write_bytes(bomb_png(*OVER_PILLOWS_ERROR)) + with client: + response = client.get("/api/gallery/bomb.png/thumbnail") + assert response.status_code >= 400 + assert decodes == [] + + @pytest.mark.xfail( + strict=True, + reason="no pixel clamp - gap: gallery_thumbnail's draft() is a no-op " + "for PNG, so thumbnail() decodes the full image first", + ) + def test_a_bomb_under_pillows_limit_is_refused_before_decode( + self, gallery, decodes + ): + client, outputs = gallery + (outputs / "bomb.png").write_bytes(bomb_png(*UNDER_PILLOWS_ERROR)) + with client: + client.get("/api/gallery/bomb.png/thumbnail") + assert decodes == [] + + @pytest.mark.xfail( + strict=True, + reason="no pixel clamp - gap: read_embedded_metadata (dw/result.py) " + "reads PngImageFile.text, which makes Pillow load() the whole image " + "to reach text chunks after IDAT", + ) + def test_the_gallery_metadata_route_does_not_decode_it(self, gallery, decodes): + """Width and height are in the header; nothing about a request for + metadata needs every pixel of a bomb.""" + client, outputs = gallery + (outputs / "bomb.png").write_bytes(bomb_png(*UNDER_PILLOWS_ERROR)) + with client: + client.get("/api/gallery/bomb.png/metadata") + assert decodes == [] + + def test_the_gallery_listing_does_not_decode_it(self, gallery, decodes): + client, outputs = gallery + (outputs / "bomb.png").write_bytes(bomb_png(*UNDER_PILLOWS_ERROR)) + with client: + assert client.get("/api/gallery").status_code == 200 + assert decodes == [] From 7ee72d253418f8bb344d2b54692d61cbd1009ce7 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 24 Sep 2026 03:27:14 +0000 Subject: [PATCH 5/6] test(security): SSRF address spellings, redirects and HF token scope Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01NSpbAKgGb282ixhsEqGQ52 --- tests/test_security_ssrf.py | 599 ++++++++++++++++++++++++++++++++++++ 1 file changed, 599 insertions(+) create mode 100644 tests/test_security_ssrf.py diff --git a/tests/test_security_ssrf.py b/tests/test_security_ssrf.py new file mode 100644 index 00000000..18765f01 --- /dev/null +++ b/tests/test_security_ssrf.py @@ -0,0 +1,599 @@ +"""SSRF beyond loopback, and where the HuggingFace token may travel. + +tests/test_locations.py pins the host policy for the spellings a workflow +author would write by accident - 127.0.0.1, localhost, 169.254.169.254, one +RFC1918 address each. This file covers the spellings an attacker writes on +purpose: the rest of the private and link-local space, IPv6 and IPv4-mapped +IPv6, the decimal/octal/hex spellings glibc's resolver accepts for an IPv4 +address, DNS names that answer with any of those, and redirects. The second +half pins where `remote_text_encoder` sends this machine's HuggingFace token +(#112) against lookalike hosts, URL-parser tricks and a cross-host redirect. + +Nothing here touches the network. Name resolution goes through +`fake_resolver`, which answers a numeric spelling the way glibc's +`getaddrinfo` does (via `inet_aton`, which is pure computation) and a name +only from its own table; HTTP goes through `_Transport`, which stands in for +requests' `HTTPAdapter.send` - the one place every requests call in the +engine, and diffusers' own `load_image`, reaches the wire. +""" + +import io +import socket +from unittest.mock import patch +from urllib.parse import urlparse + +import pytest +import requests +from PIL import Image + +from dw.locations import token_host_allowed, validate_media_url +from dw.security import InvalidInputError, TRUST_WORKFLOWS_ENV_VAR + + +@pytest.fixture +def untrusted(monkeypatch): + """The posture a deployed server runs on (conftest defaults to trusted).""" + monkeypatch.setenv(TRUST_WORKFLOWS_ENV_VAR, "0") + + +def fake_resolver(names=None): + """A getaddrinfo that never leaves the process. + + A numeric host resolves the way glibc resolves it - `inet_aton` accepts + '2130706433', '0x7f000001', '0177.0.0.1' and '127.1', exactly the + spellings a string-matching SSRF filter misses. Any other name answers + only from `names`, and an unknown one fails the way DNS does. + """ + names = names or {} + + def getaddrinfo(host, port=None, *args, **kwargs): + addresses = names.get(host) + if addresses is None: + try: + addresses = [socket.inet_ntoa(socket.inet_aton(host))] + except OSError: + raise socket.gaierror(socket.EAI_NONAME, "Name or service not known") + answers = [] + for address in addresses: + family = socket.AF_INET6 if ":" in address else socket.AF_INET + sockaddr = (address, port or 0, 0, 0) if ":" in address else (address, 0) + answers.append((family, socket.SOCK_STREAM, 6, "", sockaddr)) + return answers + + return getaddrinfo + + +@pytest.fixture +def no_real_sockets(monkeypatch): + """Belt and braces: a test that forgot a mock fails instead of dialing.""" + + def refuse(*args, **kwargs): + raise AssertionError("a test in this file tried to open a real connection") + + monkeypatch.setattr(socket.socket, "connect", refuse) + monkeypatch.setattr(socket, "create_connection", refuse) + + +def _refused(url, names=None): + with patch("dw.locations.socket.getaddrinfo", fake_resolver(names)): + with pytest.raises(InvalidInputError) as refusal: + validate_media_url(url, "an image argument") + return str(refusal.value) + + +class TestInternalAddressSpellings: + """Every spelling of an address inside the deployment is refused.""" + + @pytest.mark.parametrize( + "url", + [ + # link-local, including the cloud metadata address on other ports + "http://169.254.169.254/latest/meta-data/iam/", + "http://169.254.169.254:8080/", + "http://169.254.0.1/", + # the whole of RFC1918, both ends of each block + "http://10.0.0.1/", + "http://10.255.255.254/", + "http://172.16.0.1/", + "http://172.31.255.254/", + "http://192.168.0.1/", + "http://192.168.255.254/", + # the rest of loopback, and 'this host' + "http://127.0.0.2/", + "http://127.255.255.254/", + "http://0.0.0.0/", + ], + ) + def test_ipv4_literals(self, untrusted, no_real_sockets, url): + assert "inside this deployment" in _refused(url) + + @pytest.mark.parametrize( + "url", + [ + "http://[::1]/", + "http://[0:0:0:0:0:0:0:1]/", + "http://[::]/", + # unique local, including AWS's IPv6 metadata endpoint + "http://[fd00::1]/", + "http://[fd00:ec2::254]/latest/meta-data/", + "http://[fc00::1]/", + # link-local, with and without a zone id + "http://[fe80::1]/", + "http://[fe80::1%25eth0]/", + ], + ) + def test_ipv6_literals(self, untrusted, no_real_sockets, url): + assert "inside this deployment" in _refused(url) + + @pytest.mark.parametrize( + "url", + [ + "http://[::ffff:127.0.0.1]/", + "http://[::ffff:7f00:1]/", + "http://[::ffff:169.254.169.254]/", + "http://[::ffff:a9fe:a9fe]/", + "http://[::ffff:10.0.0.1]/", + "http://[::ffff:192.168.1.1]/", + ], + ) + def test_ipv4_mapped_ipv6(self, untrusted, no_real_sockets, url): + assert "inside this deployment" in _refused(url) + + @pytest.mark.parametrize( + "url", + [ + # 127.0.0.1 as one decimal, as hex, as octal, and shortened + "http://2130706433/", + "http://0x7f000001/", + "http://0x7f.0x0.0x0.0x1/", + "http://0177.0.0.1/", + "http://017700000001/", + "http://127.1/", + "http://0/", + # 169.254.169.254 the same ways + "http://2852039166/", + "http://0xa9fea9fe/", + "http://0251.0376.0251.0376/", + # 10.0.0.1 and 192.168.0.1 + "http://167772161/", + "http://0xc0a80001/", + ], + ) + def test_numeric_spellings_the_resolver_accepts( + self, untrusted, no_real_sockets, url + ): + """Refused because the check runs on what the name resolves to, not + on the string - these never look like an IP address to a regex.""" + assert "inside this deployment" in _refused(url) + + @pytest.mark.parametrize( + "address", + [ + "127.0.0.1", + "169.254.169.254", + "10.1.2.3", + "172.20.0.5", + "192.168.1.10", + "::1", + "fd12:3456::1", + "fe80::1", + "::ffff:127.0.0.1", + ], + ) + def test_a_name_that_resolves_inside(self, untrusted, no_real_sockets, address): + assert "inside this deployment" in _refused( + "http://innocent.example.com/x.png", + {"innocent.example.com": [address]}, + ) + + def test_one_internal_answer_among_public_ones_is_enough( + self, untrusted, no_real_sockets + ): + """The fetch may connect to any of the answers, so every one counts.""" + assert "inside this deployment" in _refused( + "http://round-robin.example.com/x.png", + {"round-robin.example.com": ["93.184.216.34", "10.0.0.7"]}, + ) + + def test_userinfo_does_not_hide_the_real_host(self, untrusted, no_real_sockets): + """'user@host' - the part after the '@' is where the request goes.""" + assert "inside this deployment" in _refused( + "http://example.com@169.254.169.254/latest/meta-data/" + ) + + def test_the_metadata_hostname_is_refused_by_what_it_answers( + self, untrusted, no_real_sockets + ): + assert "inside this deployment" in _refused( + "http://metadata.google.internal/computeMetadata/v1/", + {"metadata.google.internal": ["169.254.169.254"]}, + ) + + @pytest.mark.xfail( + strict=True, + reason="100.64.0.0/10 (CGNAT, Alibaba's 100.100.100.200 metadata) is " + "not is_private in ipaddress, and _is_internal never asks is_global", + ) + def test_shared_address_space_metadata_is_refused(self, untrusted, no_real_sockets): + assert "inside this deployment" in _refused( + "http://100.100.100.200/latest/meta-data/" + ) + + @pytest.mark.parametrize( + "url, names", + [ + ("http://93.184.216.34/x.png", None), + ("https://example.com/x.png", {"example.com": ["93.184.216.34"]}), + ("https://example.com/x.png", {"example.com": ["2606:2800:21f::1"]}), + ], + ) + def test_a_public_address_is_still_allowed( + self, untrusted, no_real_sockets, url, names + ): + """The policy is not 'refuse everything': a public host passes.""" + with patch("dw.locations.socket.getaddrinfo", fake_resolver(names)): + assert validate_media_url(url, "an image argument") == url + + +# ---------------------------------------------------------------- transport + + +def _png_bytes(): + buffer = io.BytesIO() + Image.new("RGB", (2, 2), "red").save(buffer, format="PNG") + return buffer.getvalue() + + +class _Transport: + """Stands in for requests' HTTPAdapter.send: answers from a table of + {url: (status, headers, body)} and records every request that would + have gone on the wire, headers included.""" + + def __init__(self, routes): + self.routes = routes + self.sent = [] + + def send(self, adapter, request, **kwargs): + self.sent.append(request) + status, headers, body = self.routes.get(request.url, (404, {}, b"not found")) + response = requests.Response() + response.status_code = status + response.headers.update(headers) + response.url = request.url + response.request = request + response.raw = io.BytesIO(body) + response.reason = "scripted" + return response + + def install(self, monkeypatch): + transport = self + + def send(adapter, request, **kwargs): + return transport.send(adapter, request, **kwargs) + + monkeypatch.setattr(requests.adapters.HTTPAdapter, "send", send) + return self + + def hosts(self): + return [urlparse(request.url).hostname for request in self.sent] + + +PUBLIC = {"cdn.example.com": ["93.184.216.34"], "evil.example": ["45.33.32.156"]} + + +class TestRedirects: + """A URL is checked once, before the fetch - but the fetch follows + redirects, and the redirect target was never in the document to check.""" + + @pytest.mark.xfail( + strict=True, + reason="fetch_image validates the first URL only; requests follows " + "a 302 from a public host to 169.254.169.254 unchecked", + ) + def test_fetch_image_does_not_follow_a_redirect_inside( + self, untrusted, no_real_sockets, monkeypatch, tmp_path + ): + from dw.arguments import fetch_image + + transport = _Transport( + { + "https://cdn.example.com/x.png": ( + 302, + {"Location": "http://169.254.169.254/latest/meta-data/"}, + b"", + ), + "http://169.254.169.254/latest/meta-data/": ( + 200, + {"Content-Type": "image/png"}, + _png_bytes(), + ), + } + ).install(monkeypatch) + + with patch("dw.locations.socket.getaddrinfo", fake_resolver(PUBLIC)): + try: + fetch_image("https://cdn.example.com/x.png", str(tmp_path)) + except Exception: + pass + + assert "169.254.169.254" not in transport.hosts() + + @pytest.mark.xfail( + strict=True, + reason="audio fetch validates the first URL only; requests follows " + "a redirect to loopback unchecked", + ) + def test_audio_fetch_does_not_follow_a_redirect_inside( + self, untrusted, no_real_sockets, monkeypatch, tmp_path + ): + from dw.tasks.audio_utils import load_audio + + transport = _Transport( + { + "https://cdn.example.com/a.wav": ( + 307, + {"Location": "http://127.0.0.1:8765/api/server"}, + b"", + ), + } + ).install(monkeypatch) + + with patch("dw.locations.socket.getaddrinfo", fake_resolver(PUBLIC)): + try: + load_audio("https://cdn.example.com/a.wav", str(tmp_path)) + except Exception: + pass + + assert "127.0.0.1" not in transport.hosts() + + @pytest.mark.xfail( + strict=True, + reason="validate_media_url checks urlparse's host (example.com, after " + "the backslash) while requests dials urllib3's (169.254.169.254)", + ) + def test_a_backslash_does_not_split_the_checked_host_from_the_dialed_one( + self, untrusted, no_real_sockets, monkeypatch, tmp_path + ): + """Parser differential: the host policy and the HTTP client must + agree on which host a URL names, or the policy checks one and the + request goes to the other.""" + from dw.arguments import fetch_image + + transport = _Transport({}).install(monkeypatch) + with patch( + "dw.locations.socket.getaddrinfo", + fake_resolver({"example.com": ["93.184.216.34"]}), + ): + try: + fetch_image( + "http://169.254.169.254\\@example.com/latest/meta-data/", + str(tmp_path), + ) + except Exception: + pass + assert "169.254.169.254" not in transport.hosts() + + def test_a_refused_first_url_sends_nothing( + self, untrusted, no_real_sockets, monkeypatch, tmp_path + ): + """The refusal comes before the request, not after the response.""" + from dw.arguments import fetch_image + + transport = _Transport({}).install(monkeypatch) + with patch("dw.locations.socket.getaddrinfo", fake_resolver()): + with pytest.raises(InvalidInputError): + fetch_image("http://2852039166/latest/meta-data/", str(tmp_path)) + assert transport.sent == [] + + +# ------------------------------------------------------------- the HF token + +HF_TOKEN = "hf_test_token_never_real" + +HF_HOSTS = { + "api-inference.huggingface.co": ["18.0.0.1"], + "huggingface.co": ["18.0.0.2"], + "abc.us-east-1.aws.endpoints.huggingface.cloud": ["18.0.0.3"], + "someone-encoder.hf.space": ["18.0.0.4"], + "huggingface.co.evil.example": ["45.33.32.156"], + "evilhuggingface.co": ["45.33.32.156"], + "huggingface.co-evil.example": ["45.33.32.156"], + "evil.example": ["45.33.32.156"], + "hf.space.evil.example": ["45.33.32.156"], + "notreallyhf.space": ["45.33.32.156"], + "huggingface.cloud.evil.example": ["45.33.32.156"], +} + + +class _Embeds: + def to(self, device): + return self + + +def _encode(url, monkeypatch, routes=None): + """Run remote_text_encoder against a scripted transport and return what + went on the wire. The token is a fixed fake; torch.load never sees a + real body.""" + from dw.pipeline_processors import remote + + transport = _Transport( + routes + or { + url: (200, {"Content-Type": "application/octet-stream"}, b"embeds"), + } + ).install(monkeypatch) + monkeypatch.setattr(remote, "get_token", lambda: HF_TOKEN) + monkeypatch.setattr(remote.torch, "load", lambda *a, **k: _Embeds()) + with patch("dw.locations.socket.getaddrinfo", fake_resolver(HF_HOSTS)): + try: + remote.remote_text_encoder(["a prompt"], url, "cpu") + except RuntimeError: + # a scripted non-200 at the end of a redirect chain + pass + return transport + + +def _carried_token(request): + return HF_TOKEN in (request.headers.get("Authorization") or "") + + +class TestHuggingFaceTokenScope: + @pytest.mark.parametrize( + "host", + [ + "huggingface.co", + "api-inference.huggingface.co", + "abc.us-east-1.aws.endpoints.huggingface.cloud", + "huggingface.cloud", + "hf.space", + "someone-encoder.hf.space", + "HuggingFace.CO", + ], + ) + def test_the_token_hosts(self, host): + assert token_host_allowed(host) + + @pytest.mark.parametrize( + "host", + [ + "huggingface.co.evil.example", + "evilhuggingface.co", + "huggingface.co-evil.example", + "huggingface-co.evil.example", + "hf.space.evil.example", + "notreallyhf.space", + "huggingface.cloud.evil.example", + "xhuggingface.cloud", + "huggingface.co.", + "huggingface", + "co", + "", + None, + ], + ) + def test_lookalikes_do_not_get_it(self, host): + assert not token_host_allowed(host) + + @pytest.mark.parametrize( + "url", + [ + "https://api-inference.huggingface.co/models/x", + "https://abc.us-east-1.aws.endpoints.huggingface.cloud/", + "https://someone-encoder.hf.space/encode", + ], + ) + def test_a_huggingface_endpoint_gets_the_token( + self, untrusted, no_real_sockets, monkeypatch, url + ): + transport = _encode(url, monkeypatch) + assert len(transport.sent) == 1 + assert _carried_token(transport.sent[0]) + + @pytest.mark.parametrize( + "url", + [ + "https://huggingface.co.evil.example/encode", + "https://evilhuggingface.co/encode", + "https://huggingface.co-evil.example/encode", + "https://evil.example/huggingface.co/encode", + "https://evil.example/?host=huggingface.co", + # userinfo: everything before '@' is credentials, not the host + "https://huggingface.co@evil.example/encode", + "https://huggingface.co:443@evil.example/encode", + "https://api-inference.huggingface.co%2f@evil.example/encode", + # a fragment or a backslash that a browser would read differently + "https://evil.example#@huggingface.co/encode", + "https://evil.example/\\@huggingface.co/encode", + ], + ) + def test_a_lookalike_url_never_carries_it( + self, untrusted, no_real_sockets, monkeypatch, url + ): + transport = _encode(url, monkeypatch) + for request in transport.sent: + assert not _carried_token(request), request.url + + @pytest.mark.parametrize( + "url", + [ + "https://huggingface.co@evil.example/encode", + pytest.param( + "https://evil.example\\@huggingface.co/encode", + marks=pytest.mark.xfail( + strict=True, + reason="remote.py decides on urlparse's host (huggingface.co, " + "after the backslash) but requests dials urllib3's " + "(evil.example) - the token leaves with the request", + ), + ), + "https://evil.example%5c@huggingface.co/encode", + "https://evil.example%40huggingface.co/encode", + ], + ) + def test_the_token_goes_only_where_the_connection_goes( + self, untrusted, no_real_sockets, monkeypatch, url + ): + """Parser differential: the decision reads urllib.parse's idea of the + host, the connection uses requests'/urllib3's. Whatever either one + makes of a URL, a request that carries the token must be addressed + to a HuggingFace host by the parser that actually dials.""" + from urllib3.util import parse_url + + transport = _encode(url, monkeypatch) + for request in transport.sent: + if _carried_token(request): + assert token_host_allowed(parse_url(request.url).host), request.url + + def test_a_redirect_to_another_host_drops_the_token( + self, untrusted, no_real_sockets, monkeypatch + ): + """A HuggingFace endpoint answering 307 elsewhere must not hand the + credential on with the replayed POST.""" + start = "https://api-inference.huggingface.co/models/x" + transport = _encode( + start, + monkeypatch, + routes={ + start: (307, {"Location": "https://evil.example/collect"}, b""), + "https://evil.example/collect": ( + 200, + {"Content-Type": "application/octet-stream"}, + b"embeds", + ), + }, + ) + assert transport.hosts() == ["api-inference.huggingface.co", "evil.example"] + assert _carried_token(transport.sent[0]) + assert not _carried_token(transport.sent[1]) + + def test_a_redirect_to_http_on_the_same_host_drops_the_token( + self, untrusted, no_real_sockets, monkeypatch + ): + """A downgrade to cleartext is a different origin too.""" + start = "https://api-inference.huggingface.co/models/x" + downgrade = "http://api-inference.huggingface.co/models/x" + transport = _encode( + start, + monkeypatch, + routes={ + start: (307, {"Location": downgrade}, b""), + downgrade: ( + 200, + {"Content-Type": "application/octet-stream"}, + b"embeds", + ), + }, + ) + cleartext = [r for r in transport.sent if r.url.startswith("http://")] + assert cleartext, "the scripted redirect was not followed" + assert not any(_carried_token(r) for r in cleartext) + + def test_trust_lifts_the_scope(self, no_real_sockets, monkeypatch): + """--trust-workflows lifts the scope: the documented escape hatch, + pinned so it cannot widen silently into the untrusted default.""" + monkeypatch.setenv(TRUST_WORKFLOWS_ENV_VAR, "1") + transport = _encode("https://evil.example/encode", monkeypatch) + assert _carried_token(transport.sent[0]) + monkeypatch.setenv(TRUST_WORKFLOWS_ENV_VAR, "0") + transport = _encode("https://evil.example/encode", monkeypatch) + assert not _carried_token(transport.sent[0]) From b5aa656f3145f64ea63e04b7cab28dc4c1fd7d36 Mon Sep 17 00:00:00 2001 From: Claude Date: Thu, 24 Sep 2026 03:27:14 +0000 Subject: [PATCH 6/6] test(security): documented input caps refused before any work Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01NSpbAKgGb282ixhsEqGQ52 --- tests/test_security_input_caps.py | 253 ++++++++++++++++++++++++++++++ 1 file changed, 253 insertions(+) create mode 100644 tests/test_security_input_caps.py diff --git a/tests/test_security_input_caps.py b/tests/test_security_input_caps.py new file mode 100644 index 00000000..f3f08d9f --- /dev/null +++ b/tests/test_security_input_caps.py @@ -0,0 +1,253 @@ +"""The documented size limits, refused at validation - before any work. + +`MAX_VARIABLE_VALUE_LENGTH` (20,000 characters), the variable-name cap, the +32-entry `for_each` ceiling, the 50MB workflow-file ceiling and the 200MB +upload ceiling are each a promise that an oversized input costs the server +nothing but the refusal. The live suite can see the refusal; it cannot see +whether a worker was handed the job first, a file was read whole, or a body +was buffered. These tests watch for exactly that: a ScriptedWorkerManager +that records every command it is sent, a sparse file that would be 50MB to +read, and an ASGI receive() that records whether the body was pulled. +""" + +import asyncio +import json + +import pytest +from fastapi.testclient import TestClient + +from dw.security import ( + MAX_JSON_SIZE, + MAX_VARIABLE_NAME_LENGTH, + MAX_VARIABLE_VALUE_LENGTH, + InvalidInputError, +) +from dw.server.app import create_app +from dw.server.jobs import JobManager +from dw.variables import argument_errors + +from .test_server import ScriptedWorkerManager, success_script + +OVERSIZED = "x" * 25_000 + + +def _workflow(variables=None, for_each=None): + step = { + "name": "t", + "task": {"command": "compose_text", "arguments": {"parts": ["variable:p"]}}, + "result": {"content_type": "text/plain"}, + } + if for_each is not None: + step["for_each"] = for_each + step["task"]["arguments"]["parts"] = ["item:prompt"] + return { + "id": "caps", + "variables": {"p": "short", "shots": [], **(variables or {})}, + "steps": [step], + } + + +@pytest.fixture +def server(tmp_path): + (tmp_path / "workflows").mkdir() + worker = ScriptedWorkerManager(success_script) + manager = JobManager( + str(tmp_path / "outputs"), + worker_manager=worker, + history_path=str(tmp_path / "jobs.sqlite"), + ) + app = create_app( + workflow_dir=str(tmp_path / "workflows"), + output_dir=str(tmp_path / "outputs"), + job_manager=manager, + prompt_dir=str(tmp_path / "prompts"), + asset_dir=str(tmp_path / "assets"), + ) + with TestClient(app, base_url="http://localhost") as client: + yield client, worker + + +def _executed(worker): + return [c for c in worker.commands if c.get("type") == "execute"] + + +def test_the_documented_limit_is_the_one_under_test(): + assert MAX_VARIABLE_VALUE_LENGTH == 20_000 + assert len(OVERSIZED) > MAX_VARIABLE_VALUE_LENGTH + + +class TestVariableValueLength: + def test_the_limit_itself_is_accepted(self): + assert ( + argument_errors(_workflow(), {"p": "x" * MAX_VARIABLE_VALUE_LENGTH}) == [] + ) + + def test_one_over_is_refused(self): + errors = argument_errors( + _workflow(), {"p": "x" * (MAX_VARIABLE_VALUE_LENGTH + 1)} + ) + assert [e["path"] for e in errors] == ["arguments.p"] + + @pytest.mark.parametrize( + "arguments, path", + [ + ({"p": OVERSIZED}, "arguments.p"), + ({"shots": [{"name": "a", "prompt": OVERSIZED}]}, "arguments.shots"), + ({"shots": [[OVERSIZED]]}, "arguments.shots"), + ( + {"shots": [{"name": "a", "nested": {"deep": OVERSIZED}}]}, + "arguments.shots", + ), + ], + ) + def test_25000_characters_is_refused_wherever_it_sits(self, arguments, path): + errors = argument_errors(_workflow(), arguments) + assert [e["path"] for e in errors] == [path] + assert "too long" in errors[0]["message"] + + def test_the_validate_route_refuses_it(self, server): + client, worker = server + response = client.post( + "/api/validate", + json={"workflow": _workflow(), "arguments": {"p": OVERSIZED}}, + ) + body = response.json() + assert body["valid"] is False + assert any(e["path"] == "arguments.p" for e in body["errors"]) + assert _executed(worker) == [] + + def test_the_jobs_route_refuses_it_before_the_worker_sees_it(self, server): + client, worker = server + response = client.post( + "/api/jobs", json={"workflow": _workflow(), "arguments": {"p": OVERSIZED}} + ) + assert response.status_code == 400 + assert _executed(worker) == [] + assert client.get("/api/jobs").json() in ([], {"jobs": []}) or not ( + client.get("/api/jobs").json().get("jobs") + ) + + def test_the_cli_refuses_it_before_loading_anything(self, monkeypatch, tmp_path): + import dw.run as run_module + + loaded = [] + monkeypatch.setattr( + run_module, + "workflow_from_file", + lambda *a, **k: loaded.append(a), + raising=False, + ) + workflow = tmp_path / "w.json" + workflow.write_text(json.dumps(_workflow())) + monkeypatch.setattr("sys.argv", ["dw-run", str(workflow), f"p={OVERSIZED}"]) + with pytest.raises(SystemExit) as exit_info: + run_module.main() + assert exit_info.value.code != 0 + assert loaded == [] + + @pytest.mark.xfail( + strict=True, + reason="the cap is applied to caller arguments only - a 25,000-character " + "default written into an inline workflow's own variables validates clean", + ) + def test_a_default_in_the_definition_is_held_to_the_same_cap(self, server): + client, worker = server + response = client.post( + "/api/validate", json={"workflow": _workflow({"p": OVERSIZED})} + ) + assert response.json()["valid"] is False + + +class TestOtherDocumentedLimits: + def test_an_over_long_variable_name_is_refused(self): + name = "v" * (MAX_VARIABLE_NAME_LENGTH + 1) + errors = argument_errors(_workflow({name: "d"}), {name: "value"}) + assert [e["path"] for e in errors] == [f"arguments.{name}"] + + def test_33_for_each_entries_are_a_validation_error(self, server): + client, worker = server + entries = [{"name": f"e{i}", "prompt": "p"} for i in range(33)] + response = client.post( + "/api/validate", + json={"workflow": _workflow(for_each=entries)}, + ) + body = response.json() + assert body["valid"] is False + assert any(e["path"] == "steps[0].for_each" for e in body["errors"]) + assert _executed(worker) == [] + + def test_33_entries_through_arguments_are_refused_before_the_queue(self, server): + client, worker = server + entries = [{"name": f"e{i}", "prompt": "p"} for i in range(33)] + response = client.post( + "/api/jobs", + json={ + "workflow": _workflow(for_each="variable:shots"), + "arguments": {"shots": entries}, + }, + ) + assert response.status_code == 400 + assert _executed(worker) == [] + + def test_an_oversized_workflow_file_is_refused_without_reading_it( + self, tmp_path, monkeypatch + ): + """A sparse file costs no disk; the check is on its size, so a + refusal that opened it first would show up as an open() call.""" + from dw.workflow import workflow_from_file + + big = tmp_path / "big.json" + with open(big, "wb") as file: + file.truncate(MAX_JSON_SIZE + 1) + + opened = [] + real_open = open + + def recording_open(path, *args, **kwargs): + if str(path) == str(big) or str(path).endswith("big.json"): + opened.append(path) + return real_open(path, *args, **kwargs) + + monkeypatch.setattr("builtins.open", recording_open) + with pytest.raises(InvalidInputError, match="too large"): + workflow_from_file(str(big), str(tmp_path / "out")) + assert opened == [] + + def test_an_upload_over_the_cap_is_refused_from_its_declared_length(self, server): + """The 200MB ceiling is checked on Content-Length before the body is + read - driven over raw ASGI so the declared length can exceed what + is actually sent, and so the test sees whether receive() was called.""" + client, worker = server + app = client.app + pulled = [] + sent = [] + + async def receive(): + pulled.append(True) + return {"type": "http.request", "body": b"x", "more_body": False} + + async def send(message): + sent.append(message) + + scope = { + "type": "http", + "asgi": {"version": "3.0"}, + "http_version": "1.1", + "method": "POST", + "scheme": "http", + "path": "/api/uploads", + "raw_path": b"/api/uploads", + "query_string": b"filename=huge.png", + "root_path": "", + "headers": [ + (b"host", b"localhost"), + (b"content-length", str(300 * 1024 * 1024).encode()), + (b"content-type", b"application/octet-stream"), + ], + "client": ("127.0.0.1", 1), + "server": ("localhost", 80), + } + asyncio.run(app(scope, receive, send)) + start = next(m for m in sent if m["type"] == "http.response.start") + assert start["status"] == 413 + assert pulled == []