Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 23 additions & 1 deletion st2api/tests/unit/controllers/v1/test_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,11 +48,13 @@ def test_origin(self):
self.assertEqual(
response.headers["Access-Control-Allow-Origin"], "http://127.0.0.1:3000"
)
self.assertEqual(response.headers["Access-Control-Allow-Credentials"], "true")

def test_additional_origin(self):
response = self.app.get("/", headers={"origin": "http://dev"})
self.assertEqual(response.status_int, 200)
self.assertEqual(response.headers["Access-Control-Allow-Origin"], "http://dev")
self.assertEqual(response.headers["Access-Control-Allow-Credentials"], "true")

def test_wrong_origin(self):
# Invalid origin (not specified in the config), we return first allowed origin specified
Expand All @@ -62,6 +64,7 @@ def test_wrong_origin(self):
self.assertEqual(
response.headers.get("Access-Control-Allow-Origin"), "http://127.0.0.1:3000"
)
self.assertNotIn("Access-Control-Allow-Credentials", response.headers)

invalid_origins = [
"http://",
Expand All @@ -78,6 +81,7 @@ def test_wrong_origin(self):
response.headers.get("Access-Control-Allow-Origin"),
"http://127.0.0.1:3000",
)
self.assertNotIn("Access-Control-Allow-Credentials", response.headers)

def test_wildcard_origin(self):
try:
Expand All @@ -86,7 +90,25 @@ def test_wildcard_origin(self):
finally:
cfg.CONF.clear_override("allow_origin", "api")
self.assertEqual(response.status_int, 200)
self.assertEqual(response.headers["Access-Control-Allow-Origin"], "http://xss")
# Must return wildcard origin "*", never reflecting the untrusted origin
self.assertEqual(response.headers["Access-Control-Allow-Origin"], "*")
# Must NOT include Access-Control-Allow-Credentials
self.assertNotIn("Access-Control-Allow-Credentials", response.headers)

def test_hardcoded_localhost_origins_not_automatically_allowed(self):
try:
cfg.CONF.set_override(
"allow_origin", ["http://custom-origin.example.com"], "api"
)
response = self.app.get("/", headers={"origin": "http://localhost:8080"})
self.assertEqual(response.status_int, 200)
self.assertEqual(
response.headers["Access-Control-Allow-Origin"],
"http://custom-origin.example.com",
)
self.assertNotIn("Access-Control-Allow-Credentials", response.headers)
finally:
cfg.CONF.clear_override("allow_origin", "api")

def test_valid_status_code_is_returned_on_invalid_path(self):
# TypeError: get_all() takes exactly 1 argument (2 given)
Expand Down
3 changes: 2 additions & 1 deletion st2common/st2common/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -491,7 +491,8 @@ def register_opts(ignore_errors=False):
cfg.ListOpt(
"allow_origin",
default=["http://127.0.0.1:3000"],
help="List of origins allowed for api, auth and stream",
help="List of origins allowed for api, auth and stream. Note: If '*' is specified, "
"Access-Control-Allow-Credentials will not be set.",
),
cfg.IntOpt(
"max_page_size",
Expand Down
66 changes: 49 additions & 17 deletions st2common/st2common/middleware/cors.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,32 +43,49 @@ def custom_start_response(status, headers, exc_info=None):
headers = ResponseHeaders(headers)

origin = request.headers.get("Origin")
origins = OrderedSet(cfg.CONF.api.allow_origin)
raw_origins = cfg.CONF.api.allow_origin or []
origins = OrderedSet(
[o.strip() for o in raw_origins if isinstance(o, str) and o.strip()]
)

# Build a list of the default allowed origins
public_api_url = cfg.CONF.auth.api_url

# Default gulp development server WebUI URL
origins.add("http://127.0.0.1:3000")

# By default WebUI simple http server listens on 8080
origins.add("http://localhost:8080")
origins.add("http://127.0.0.1:8080")

if public_api_url:
if (
public_api_url
and isinstance(public_api_url, str)
and public_api_url.strip()
):
# Public API URL
origins.add(public_api_url)
origins.add(public_api_url.strip())

origins = list(origins)

origin_allowed = None
allow_credentials = False
vary_origin = False

if origin:
if "*" in origins:
if origin in origins and origin != "*":
origin_allowed = origin
allow_credentials = True
vary_origin = True
elif "*" in origins:
origin_allowed = "*"
allow_credentials = False
elif origins:
# Origin is not allowed; return first configured origin (per commit 66605b7b)
# so browser CORS check rejects it, while not enabling credentials.
origin_allowed = origins[0]
allow_credentials = False
vary_origin = True
elif origins:
# No Origin header was provided (e.g. non-browser client / direct request).
if "*" in origins:
origin_allowed = "*"
else:
# See http://www.w3.org/TR/cors/#access-control-allow-origin-response-header
origin_allowed = origin if origin in origins else list(origins)[0]
else:
origin_allowed = list(origins)[0]
origin_allowed = origins[0]
allow_credentials = False

methods_allowed = ["GET", "POST", "PUT", "DELETE", "OPTIONS"]
request_headers_allowed = [
Expand All @@ -85,10 +102,25 @@ def custom_start_response(status, headers, exc_info=None):
REQUEST_ID_HEADER,
]

headers["Access-Control-Allow-Origin"] = origin_allowed
if origin_allowed:
headers["Access-Control-Allow-Origin"] = origin_allowed
if vary_origin:
existing_vary = headers.get("Vary")
if existing_vary:
vary_tokens = [
v.strip().lower() for v in existing_vary.split(",")
]
if "origin" not in vary_tokens and "*" not in vary_tokens:
headers["Vary"] = "%s, Origin" % existing_vary
else:
headers["Vary"] = "Origin"

headers["Access-Control-Allow-Methods"] = ",".join(methods_allowed)
headers["Access-Control-Allow-Headers"] = ",".join(request_headers_allowed)
headers["Access-Control-Allow-Credentials"] = "true"

if allow_credentials:
headers["Access-Control-Allow-Credentials"] = "true"

headers["Access-Control-Expose-Headers"] = ",".join(
response_headers_allowed
)
Expand Down
3 changes: 3 additions & 0 deletions st2common/tests/unit/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,9 @@ python_tests(
"test_param_utils.py": dict(
uses=["system_user"],
),
"test_cors_middleware.py": dict(
uses=[],
),
},
)

Expand Down
Loading
Loading