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"
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 == []
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 == []
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])
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)
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