File size: 6,026 Bytes
66ee87e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
"""API token, body limit and security headers (app/security.py): a pure ASGI middleware, tested without a web framework."""
import asyncio
import json
import unittest

from app.security import SECURITY_HEADERS, Gate, SecurityMiddleware



async def echo_app(scope, receive, send):
    body = b""
    while True:
        msg = await receive()
        body += msg.get("body", b"")
        if not msg.get("more_body"):
            break
    await send({"type": "http.response.start", "status": 200, "headers": [(b"content-type", b"application/json")]})
    await send({"type": "http.response.body", "body": json.dumps({"path": scope["path"], "len": len(body)}).encode()})


def call(mw, path, method="GET", headers=None, body=b"", chunks=None, query=b""):
    scope = {"type": "http", "method": method, "path": path, "query_string": query,
             "headers": [(k.lower().encode(), v.encode()) for k, v in (headers or {}).items()]}
    parts = chunks if chunks is not None else [body]
    queue = [{"type": "http.request", "body": c, "more_body": i < len(parts) - 1} for i, c in enumerate(parts)]
    sent = []

    async def receive():
        return queue.pop(0) if queue else {"type": "http.disconnect"}

    async def send(msg):
        sent.append(msg)

    asyncio.run(mw(scope, receive, send))
    start = next(m for m in sent if m["type"] == "http.response.start")
    body_out = b"".join(m.get("body", b"") for m in sent if m["type"] == "http.response.body")
    return start["status"], {k.decode(): v.decode() for k, v in start["headers"]}, body_out


def mw(max_body=1000):
    return SecurityMiddleware(echo_app, max_body=max_body)


class OpenApiTest(unittest.TestCase):
    """No token auth anywhere (operator ruling 2026-09-28, Hugging Face Space): the page and the API are open."""

    def test_api_needs_no_token(self):
        for path in ("/api/status", "/api/demos", "/api/health"):
            with self.subTest(path=path):
                self.assertEqual(call(mw(), path)[0], 200)
        self.assertEqual(call(mw(), "/api/decide", "POST", body=b"{}")[0], 200)

    def test_no_token_or_cookie_machinery_remains(self):
        _, headers, _ = call(mw(), "/")
        self.assertNotIn("set-cookie", headers)
        self.assertNotIn("www-authenticate", call(mw(), "/api/status")[1])

    def test_constructor_takes_no_token(self):
        with self.assertRaises(TypeError):
            SecurityMiddleware(echo_app, token="x" * 40, max_body=10)


class BodyLimitTest(unittest.TestCase):
    def test_body_under_limit_reaches_the_app_intact(self):
        status, _, body = call(mw(), "/api/decide", "POST", chunks=[b"a" * 400, b"b" * 400])
        self.assertEqual((status, json.loads(body)["len"]), (200, 800))

    def test_declared_length_over_limit_is_413_before_reading(self):
        status, _, body = call(mw(), "/api/decide", "POST", headers={"Content-Length": "5000"}, body=b"x")
        self.assertEqual(status, 413)
        self.assertEqual(json.loads(body)["detail"], "Request body is over the 1000-byte limit.")

    def test_streamed_body_over_limit_is_413(self):
        self.assertEqual(call(mw(), "/api/decide", "POST", chunks=[b"a" * 600, b"b" * 600])[0], 413)

    def test_gradio_requests_are_left_to_gradio(self):
        status, _, body = call(mw(), "/gradio/gradio_api/call/decide", "POST", chunks=[b"a" * 600, b"b" * 600])
        self.assertEqual((status, json.loads(body)["len"]), (200, 1200))


class HeadersTest(unittest.TestCase):
    def test_every_lab_response_carries_the_security_headers(self):
        for path in ("/", "/static/app.js", "/api/status"):
            _, headers, _ = call(mw(), path)
            for k, v in SECURITY_HEADERS.items():
                with self.subTest(path=path, header=k):
                    self.assertEqual(headers[k], v)

    def test_csp_is_strict_and_only_huggingface_may_frame_the_page(self):
        csp = SECURITY_HEADERS["content-security-policy"]
        for part in ("default-src 'self'", "script-src 'self'", "style-src 'self'", "object-src 'none'",
                     "base-uri 'none'", "form-action 'none'",
                     "frame-ancestors 'self' https://huggingface.co https://*.hf.space"):
            self.assertIn(part, csp)
        self.assertNotIn("unsafe", csp)

    def test_no_x_frame_options_so_the_space_can_be_shown_on_huggingface(self):
        self.assertNotIn("x-frame-options", call(mw(), "/")[1])

    def test_the_other_protective_headers_are_present(self):
        _, headers, _ = call(mw(), "/")
        self.assertEqual((headers["x-content-type-options"], headers["referrer-policy"]), ("nosniff", "no-referrer"))

    def test_gradio_pages_keep_their_own_headers(self):
        _, headers, _ = call(mw(), "/gradio/")
        self.assertNotIn("content-security-policy", headers)

    def test_api_responses_are_not_cached(self):
        self.assertEqual(call(mw(), "/api/status")[1]["cache-control"], "no-store")

    def test_page_and_static_files_are_revalidated_on_every_load(self):
        for path in ("/", "/static/app.js", "/static/app.css"):
            with self.subTest(path=path):
                self.assertEqual(call(mw(), path)[1]["cache-control"], "no-cache")

    def test_non_http_scopes_pass_through(self):
        seen = []

        async def app(scope, receive, send):
            seen.append(scope["type"])

        asyncio.run(SecurityMiddleware(app, max_body=10)({"type": "lifespan"}, None, None))
        self.assertEqual(seen, ["lifespan"])


class GateTest(unittest.TestCase):
    def test_gate_admits_up_to_its_size_then_refuses(self):
        g = Gate(2)
        self.assertEqual([g.try_enter(), g.try_enter(), g.try_enter()], [True, True, False])
        g.leave()
        self.assertTrue(g.try_enter())

    def test_gate_of_one_is_a_non_blocking_lock(self):
        g = Gate(1)
        self.assertTrue(g.try_enter())
        self.assertFalse(g.try_enter())
        g.leave()
        self.assertTrue(g.try_enter())


if __name__ == "__main__":
    unittest.main()