diff --git a/flask_cors/core.py b/flask_cors/core.py index 1eeb20a..acab973 100644 --- a/flask_cors/core.py +++ b/flask_cors/core.py @@ -5,6 +5,7 @@ from collections.abc import Iterable, Mapping from dataclasses import dataclass from datetime import timedelta +from ipaddress import IPv6Address from typing import Any, TypedDict, Union, cast from flask import Blueprint, Flask, Response, current_app, request @@ -508,6 +509,21 @@ def _resolve_patterns(patterns: Iterable[ResourcePattern], *, ignore_case: bool) return [_resolve_pattern(p, ignore_case=ignore_case) for p in patterns] +def _resolve_origin(origin: ResourcePattern) -> ResourcePattern: + # IPv6 URL brackets delimit a host, not a regex character class. Restrict + # this exception to complete literal origins; regex origins still compile. + if isinstance(origin, str): + match = re.fullmatch(r"[a-zA-Z][a-zA-Z0-9+.-]*://\[([0-9a-fA-F:.]+)\](?::[0-9]+)?", origin) + if match: + try: + IPv6Address(match.group(1)) + except ValueError: + pass + else: + return origin + return _resolve_pattern(origin, ignore_case=True) + + def serialize_options(opts: Mapping[str, Any]) -> _ComputedCorsOptions: """ Normalize a raw options mapping into a strongly-typed :class:`_ComputedCorsOptions`. @@ -537,7 +553,7 @@ def serialize_options(opts: Mapping[str, Any]) -> _ComputedCorsOptions: "http://www.w3.org/TR/cors/#resource-requests" ) - origins = _resolve_patterns(sanitized_origins, ignore_case=True) + origins = [_resolve_origin(origin) for origin in sanitized_origins] allow_headers = _resolve_patterns(sanitize_regex_param(opts.get("allow_headers")), ignore_case=True) methods = flexible_str(opts.get("methods")) diff --git a/tests/extension/test_ipv6.py b/tests/extension/test_ipv6.py new file mode 100644 index 0000000..029e404 --- /dev/null +++ b/tests/extension/test_ipv6.py @@ -0,0 +1,49 @@ +import re + +import pytest +from flask import Flask + +from flask_cors import CORS, cross_origin + + +@pytest.mark.parametrize("origin", ["http://[::1]:8000", "https://[2001:db8::1]", "http://[::ffff:192.0.2.1]:8080"]) +@pytest.mark.parametrize("use_decorator", [False, True]) +def test_literal_ipv6_origin(origin, use_decorator): + app = Flask(__name__) + + def index(): + return "ok" + + if use_decorator: + index = cross_origin(origins=[origin])(index) + else: + CORS(app, origins=[origin]) + app.add_url_rule("/", view_func=index) + client = app.test_client() + + for requested in (origin, origin.upper()): + response = client.get("/", headers={"Origin": requested}) + assert response.headers["Access-Control-Allow-Origin"] == requested + for requested in (origin + "0", "http://[::2]:8000", "http://:8000"): + response = client.get("/", headers={"Origin": requested}) + assert "Access-Control-Allow-Origin" not in response.headers + + +@pytest.mark.parametrize("origin", [r"http://\[::1\]:\d+$", re.compile(r"http://\[::1\]:\d+$")]) +def test_ipv6_regex_origin(origin): + app = Flask(__name__) + CORS(app, origins=[origin]) + app.add_url_rule("/", view_func=lambda: "ok") + client = app.test_client() + assert ( + client.get("/", headers={"Origin": "http://[::1]:8000"}).headers["Access-Control-Allow-Origin"] + == "http://[::1]:8000" + ) + assert "Access-Control-Allow-Origin" not in client.get("/", headers={"Origin": "http://[::2]:8000"}).headers + + +def test_literal_ipv6_origin_without_request_origin(): + app = Flask(__name__) + CORS(app, origins=["http://[::1]:8000"]) + app.add_url_rule("/", view_func=lambda: "ok") + assert app.test_client().get("/").headers["Access-Control-Allow-Origin"] == "http://[::1]:8000"