Compare commits
3 commits
main
...
studio-sec
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
57e888dc7b | ||
|
|
28f5b81853 |
||
|
|
7698667162 |
2 changed files with 169 additions and 26 deletions
|
|
@ -492,11 +492,8 @@ app.add_middleware(LoggingMiddleware)
|
||||||
|
|
||||||
# img/media-src allow any https origin so HF model-card assets render (mirrors
|
# img/media-src allow any https origin so HF model-card assets render (mirrors
|
||||||
# tauri.conf.json); scripts/frames/connect-src stay same-origin + HF.
|
# tauri.conf.json); scripts/frames/connect-src stay same-origin + HF.
|
||||||
from starlette.middleware.base import BaseHTTPMiddleware # noqa: E402
|
|
||||||
from starlette.requests import Request as _StarletteRequest # noqa: E402
|
|
||||||
|
|
||||||
|
|
||||||
_CSP_SCRIPT_NONCE_HEADER = "x-internal-script-nonce"
|
_CSP_SCRIPT_NONCE_HEADER = "x-internal-script-nonce"
|
||||||
|
_CSP_SCRIPT_NONCE_HEADER_BYTES = _CSP_SCRIPT_NONCE_HEADER.encode("latin-1")
|
||||||
_ARTIFACT_PREVIEW_FRAME_PATH = "/api/inference/artifact-preview-frame"
|
_ARTIFACT_PREVIEW_FRAME_PATH = "/api/inference/artifact-preview-frame"
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -549,28 +546,64 @@ def _build_csp(script_nonce: "str | None" = None) -> str:
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class SecurityHeadersMiddleware(BaseHTTPMiddleware):
|
class SecurityHeadersMiddleware:
|
||||||
"""Set baseline security headers; splice per-response inline-script nonces into CSP."""
|
"""Set baseline security headers and splice per-response script nonces into CSP.
|
||||||
|
|
||||||
async def dispatch(self, request: _StarletteRequest, call_next):
|
Pure ASGI, not BaseHTTPMiddleware: the latter wraps streaming responses in its
|
||||||
response = await call_next(request)
|
own anyio task group, breaking request.is_disconnected() and raising cancel
|
||||||
# Strip the internal nonce hand-off header so it never reaches the client
|
scope errors on the /v1/chat/completions SSE stream. Mirrors LoggingMiddleware.
|
||||||
nonce = response.headers.get(_CSP_SCRIPT_NONCE_HEADER)
|
"""
|
||||||
if nonce is not None:
|
|
||||||
del response.headers[_CSP_SCRIPT_NONCE_HEADER]
|
def __init__(self, app):
|
||||||
response.headers.setdefault("Content-Security-Policy", _build_csp(nonce))
|
self.app = app
|
||||||
# Omit X-Frame-Options in Colab — CSP frame-ancestors handles it, and
|
|
||||||
# DENY would block serve_kernel_port_as_iframe regardless of CSP.
|
async def __call__(self, scope, receive, send):
|
||||||
if not _IS_COLAB and request.url.path != _ARTIFACT_PREVIEW_FRAME_PATH:
|
if scope["type"] != "http":
|
||||||
response.headers.setdefault("X-Frame-Options", "DENY")
|
await self.app(scope, receive, send)
|
||||||
response.headers.setdefault("X-Content-Type-Options", "nosniff")
|
return
|
||||||
response.headers.setdefault("Referrer-Policy", "no-referrer")
|
|
||||||
response.headers.setdefault(
|
path = scope.get("path", "")
|
||||||
"Permissions-Policy",
|
|
||||||
"camera=(), microphone=(self), geolocation=()",
|
async def send_wrapper(message):
|
||||||
)
|
if message["type"] != "http.response.start":
|
||||||
response.headers["server"] = "unsloth-studio"
|
await send(message)
|
||||||
return response
|
return
|
||||||
|
|
||||||
|
# Drop the internal nonce handoff header; splice its value into the CSP.
|
||||||
|
nonce = None
|
||||||
|
headers = []
|
||||||
|
for name, value in message.get("headers") or []:
|
||||||
|
if name.lower() == _CSP_SCRIPT_NONCE_HEADER_BYTES:
|
||||||
|
nonce = value.decode("latin-1")
|
||||||
|
continue
|
||||||
|
headers.append((name, value))
|
||||||
|
|
||||||
|
present = {name.lower() for name, _ in headers}
|
||||||
|
|
||||||
|
def _setdefault(name: bytes, value: str) -> None:
|
||||||
|
if name not in present:
|
||||||
|
headers.append((name, value.encode("latin-1")))
|
||||||
|
present.add(name)
|
||||||
|
|
||||||
|
_setdefault(b"content-security-policy", _build_csp(nonce))
|
||||||
|
# Skip X-Frame-Options in Colab: CSP frame-ancestors covers it; DENY
|
||||||
|
# would block serve_kernel_port_as_iframe.
|
||||||
|
if not _IS_COLAB and path != _ARTIFACT_PREVIEW_FRAME_PATH:
|
||||||
|
_setdefault(b"x-frame-options", "DENY")
|
||||||
|
_setdefault(b"x-content-type-options", "nosniff")
|
||||||
|
_setdefault(b"referrer-policy", "no-referrer")
|
||||||
|
_setdefault(
|
||||||
|
b"permissions-policy",
|
||||||
|
"camera=(), microphone=(self), geolocation=()",
|
||||||
|
)
|
||||||
|
# Hard override (was response.headers["server"] = ...).
|
||||||
|
headers = [(n, v) for n, v in headers if n.lower() != b"server"]
|
||||||
|
headers.append((b"server", b"unsloth-studio"))
|
||||||
|
|
||||||
|
message["headers"] = headers
|
||||||
|
await send(message)
|
||||||
|
|
||||||
|
await self.app(scope, receive, send_wrapper)
|
||||||
|
|
||||||
|
|
||||||
app.add_middleware(SecurityHeadersMiddleware)
|
app.add_middleware(SecurityHeadersMiddleware)
|
||||||
|
|
|
||||||
|
|
@ -293,6 +293,116 @@ class TestSecurityHeadersMiddleware:
|
||||||
# not read directive-string `in` membership as URL sanitisation.
|
# not read directive-string `in` membership as URL sanitisation.
|
||||||
assert any(src == "https:" for src in directives[name])
|
assert any(src == "https:" for src in directives[name])
|
||||||
|
|
||||||
|
def test_is_pure_asgi_not_basehttp_middleware(self, main_module):
|
||||||
|
# Regression: as a BaseHTTPMiddleware this wrapped the SSE stream in its
|
||||||
|
# own anyio task group, breaking disconnect detection (GPU stuck at 100%)
|
||||||
|
# and raising cancel scope errors. Must stay pure ASGI.
|
||||||
|
from starlette.middleware.base import BaseHTTPMiddleware
|
||||||
|
|
||||||
|
cls = main_module.SecurityHeadersMiddleware
|
||||||
|
assert not issubclass(cls, BaseHTTPMiddleware)
|
||||||
|
assert not hasattr(cls, "dispatch")
|
||||||
|
|
||||||
|
def test_forwards_receive_channel_unchanged(self, main_module):
|
||||||
|
# Must forward the ASGI receive channel untouched so client disconnects
|
||||||
|
# reach the streaming handler (BaseHTTPMiddleware swapped in its own).
|
||||||
|
seen = {}
|
||||||
|
|
||||||
|
async def inner_app(scope, receive, send):
|
||||||
|
seen["receive"] = receive
|
||||||
|
await send({"type": "http.response.start", "status": 200, "headers": []})
|
||||||
|
await send({"type": "http.response.body", "body": b"ok", "more_body": False})
|
||||||
|
|
||||||
|
mw = main_module.SecurityHeadersMiddleware(inner_app)
|
||||||
|
sentinel_receive = object() # forwarded verbatim, never wrapped/awaited
|
||||||
|
sent = []
|
||||||
|
|
||||||
|
async def send(message):
|
||||||
|
sent.append(message)
|
||||||
|
|
||||||
|
async def run():
|
||||||
|
await mw(
|
||||||
|
{"type": "http", "path": "/plain", "headers": []},
|
||||||
|
sentinel_receive,
|
||||||
|
send,
|
||||||
|
)
|
||||||
|
|
||||||
|
asyncio.run(run())
|
||||||
|
assert seen["receive"] is sentinel_receive
|
||||||
|
start = next(m for m in sent if m["type"] == "http.response.start")
|
||||||
|
names = {n.lower() for n, _ in start["headers"]}
|
||||||
|
assert b"content-security-policy" in names
|
||||||
|
assert b"server" in names
|
||||||
|
|
||||||
|
def test_streaming_response_survives_client_disconnect(self, main_module):
|
||||||
|
# A StreamingResponse that polls is_disconnected() (like gguf_tool_stream)
|
||||||
|
# must unwind cleanly on client disconnect: no cancel scope error, the
|
||||||
|
# generator's finally runs, and security headers are still applied.
|
||||||
|
from fastapi import FastAPI, Request
|
||||||
|
from fastapi.responses import StreamingResponse
|
||||||
|
|
||||||
|
state = {"cleaned_up": False}
|
||||||
|
app = FastAPI()
|
||||||
|
app.add_middleware(main_module.SecurityHeadersMiddleware)
|
||||||
|
|
||||||
|
@app.get("/v1/chat/completions")
|
||||||
|
async def stream(request: Request):
|
||||||
|
async def gen():
|
||||||
|
try:
|
||||||
|
for i in range(1000):
|
||||||
|
if await request.is_disconnected():
|
||||||
|
break
|
||||||
|
yield f"data: {i}\n\n".encode()
|
||||||
|
await asyncio.sleep(0.01)
|
||||||
|
finally:
|
||||||
|
state["cleaned_up"] = True
|
||||||
|
|
||||||
|
return StreamingResponse(gen(), media_type = "text/event-stream")
|
||||||
|
|
||||||
|
scope = {
|
||||||
|
"type": "http",
|
||||||
|
"asgi": {"version": "3.0", "spec_version": "2.3"},
|
||||||
|
"http_version": "1.1",
|
||||||
|
"method": "GET",
|
||||||
|
"path": "/v1/chat/completions",
|
||||||
|
"raw_path": b"/v1/chat/completions",
|
||||||
|
"query_string": b"",
|
||||||
|
"root_path": "",
|
||||||
|
"scheme": "http",
|
||||||
|
"headers": [(b"host", b"testserver")],
|
||||||
|
"client": ("127.0.0.1", 50000),
|
||||||
|
"server": ("127.0.0.1", 80),
|
||||||
|
}
|
||||||
|
|
||||||
|
async def run():
|
||||||
|
body_started = asyncio.Event()
|
||||||
|
calls = {"n": 0}
|
||||||
|
|
||||||
|
async def receive():
|
||||||
|
calls["n"] += 1
|
||||||
|
if calls["n"] == 1:
|
||||||
|
return {"type": "http.request", "body": b"", "more_body": False}
|
||||||
|
await body_started.wait() # client clicks Stop after tokens stream
|
||||||
|
return {"type": "http.disconnect"}
|
||||||
|
|
||||||
|
sent = []
|
||||||
|
|
||||||
|
async def send(message):
|
||||||
|
sent.append(message)
|
||||||
|
if message["type"] == "http.response.body" and message.get("body"):
|
||||||
|
body_started.set()
|
||||||
|
|
||||||
|
# Must return without raising the anyio cancel-scope RuntimeError.
|
||||||
|
await asyncio.wait_for(app(scope, receive, send), timeout = 5.0)
|
||||||
|
return sent
|
||||||
|
|
||||||
|
sent = asyncio.run(run())
|
||||||
|
assert state["cleaned_up"] is True
|
||||||
|
start = next(m for m in sent if m["type"] == "http.response.start")
|
||||||
|
names = {n.lower() for n, _ in start["headers"]}
|
||||||
|
assert b"content-security-policy" in names
|
||||||
|
assert b"server" in names
|
||||||
|
|
||||||
|
|
||||||
# /api/health auth gate
|
# /api/health auth gate
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue