From e1fa323932a63fad8ddc1b07fed4cf090940ba4d Mon Sep 17 00:00:00 2001 From: Xuan Son Nguyen Date: Tue, 14 Jul 2026 12:28:32 +0200 Subject: [PATCH] add tests --- tools/server/tests/unit/test_security.py | 61 ++++++++++++++++++++++++ tools/server/tests/utils.py | 3 ++ 2 files changed, 64 insertions(+) diff --git a/tools/server/tests/unit/test_security.py b/tools/server/tests/unit/test_security.py index a0c3e214ae..fa22c5dcda 100644 --- a/tools/server/tests/unit/test_security.py +++ b/tools/server/tests/unit/test_security.py @@ -107,6 +107,67 @@ def test_cors_options(origin: str, cors_header: str, cors_header_value: str): assert res.headers[cors_header] == cors_header_value +@pytest.mark.parametrize("origin", [ + "http://localhost", + "http://localhost:8080", + "http://127.0.0.1", + "http://127.0.0.1:3000", + "http://[::1]", + "http://[::1]:3000", +]) +def test_cors_origins_localhost_reflects(origin: str): + server = ServerPreset.router() + server.cors_origins = "localhost" + server.start() + res = server.make_request("OPTIONS", "/completions", headers={ + "Origin": origin, + "Access-Control-Request-Method": "POST", + "Access-Control-Request-Headers": "Authorization", + }) + assert res.status_code == 200 + assert res.headers["Access-Control-Allow-Origin"] == origin + + +@pytest.mark.parametrize("origin", [ + "http://web.mydomain.fr", + "http://evil.com", + "http://notlocalhost", + "http://localhost.evil.com", +]) +def test_cors_origins_localhost_rejects(origin: str): + server = ServerPreset.router() + server.cors_origins = "localhost" + server.start() + res = server.make_request("OPTIONS", "/completions", headers={ + "Origin": origin, + "Access-Control-Request-Method": "POST", + "Access-Control-Request-Headers": "Authorization", + }) + assert res.status_code == 200 + assert "Access-Control-Allow-Origin" not in res.headers + + +def test_cors_origins_defaults_to_localhost_with_tools_enabled(): + server = ServerPreset.router() + server.server_tools = "all" + server.start() + res = server.make_request("OPTIONS", "/completions", headers={ + "Origin": "http://localhost:8080", + "Access-Control-Request-Method": "POST", + "Access-Control-Request-Headers": "Authorization", + }) + assert res.status_code == 200 + assert res.headers["Access-Control-Allow-Origin"] == "http://localhost:8080" + + res = server.make_request("OPTIONS", "/completions", headers={ + "Origin": "http://evil.com", + "Access-Control-Request-Method": "POST", + "Access-Control-Request-Headers": "Authorization", + }) + assert res.status_code == 200 + assert "Access-Control-Allow-Origin" not in res.headers + + def test_cors_proxy_only_forwards_explicit_proxy_headers(): class CaptureHeadersHandler(BaseHTTPRequestHandler): def do_GET(self): diff --git a/tools/server/tests/utils.py b/tools/server/tests/utils.py index 8c0de384f6..b32a037d82 100644 --- a/tools/server/tests/utils.py +++ b/tools/server/tests/utils.py @@ -114,6 +114,7 @@ class ServerProcess: backend_sampling: bool = False gcp_compat: bool = False server_tools: str | None = None + cors_origins: str | None = None # session variables process: subprocess.Popen | None = None @@ -170,6 +171,8 @@ class ServerProcess: server_args.extend(["--models-max", self.models_max]) if self.models_preset: server_args.extend(["--models-preset", self.models_preset]) + if self.cors_origins: + server_args.extend(["--cors-origins", self.cors_origins]) if self.n_batch: server_args.extend(["--batch-size", self.n_batch]) if self.n_ubatch: