diff --git a/aws_lambda_powertools/event_handler/middlewares/openapi_validation.py b/aws_lambda_powertools/event_handler/middlewares/openapi_validation.py index a0311184a51..b0217147ae2 100644 --- a/aws_lambda_powertools/event_handler/middlewares/openapi_validation.py +++ b/aws_lambda_powertools/event_handler/middlewares/openapi_validation.py @@ -5,8 +5,8 @@ import json import logging import warnings -from typing import TYPE_CHECKING, Any, Callable, Mapping, MutableMapping, Sequence, Union, cast -from urllib.parse import parse_qs +from typing import TYPE_CHECKING, Any, Callable, Collection, Mapping, MutableMapping, Sequence, Union, cast +from urllib.parse import parse_qs, unquote from pydantic import BaseModel from typing_extensions import get_args, get_origin @@ -29,6 +29,9 @@ ) from aws_lambda_powertools.event_handler.openapi.params import Param, UploadFile from aws_lambda_powertools.event_handler.openapi.types import UnionType +from aws_lambda_powertools.utilities.data_classes.alb_event import ALBEvent +from aws_lambda_powertools.utilities.data_classes.common import CaseInsensitiveDict +from aws_lambda_powertools.utilities.data_classes.vpc_lattice import VPCLatticeEventV2 if TYPE_CHECKING: from pydantic.fields import FieldInfo @@ -39,6 +42,7 @@ from aws_lambda_powertools.event_handler.openapi.compat import ModelField from aws_lambda_powertools.event_handler.openapi.types import IncEx from aws_lambda_powertools.event_handler.types import EventHandlerInstance + from aws_lambda_powertools.utilities.data_classes.common import BaseProxyEvent logger = logging.getLogger(__name__) @@ -79,6 +83,7 @@ def handler(self, app: EventHandlerInstance, next_middleware: NextMiddleware) -> query_string = _normalize_multi_params( app.current_event.resolved_query_string_parameters, route.dependant.query_params, + multi_value_params=_get_multi_value_query_params(app.current_event), ) # Process query values @@ -91,6 +96,7 @@ def handler(self, app: EventHandlerInstance, next_middleware: NextMiddleware) -> headers = _normalize_multi_params( app.current_event.resolved_headers_field, route.dependant.header_params, + multi_value_params=CaseInsensitiveDict(app.current_event.get("multiValueHeaders")), ) # Process header values @@ -583,9 +589,31 @@ def _get_embed_body( return received_body, field_alias_omitted +def _get_multi_value_query_params(event: BaseProxyEvent) -> Collection[str]: + """Identify native repeated parameters so they are not mistaken for split scalar strings.""" + if raw_query := event.get("rawQueryString"): + return {name for name, values in parse_qs(raw_query, keep_blank_values=True).items() if len(values) > 1} + + if isinstance(event, VPCLatticeEventV2): + return { + name + for name, values in (event.get("queryStringParameters") or {}).items() + if isinstance(values, list) and len(values) > 1 + } + + params = event.multi_value_query_string_parameters + if isinstance(event, ALBEvent) and event.decode_query_parameters: + # Follow the same merge and decoding order as ALBEvent, including decoded key collisions. + decoded_sources = {unquote(name): name in params for name in {**event.query_string_parameters, **params}} + return {name for name, is_multi_value in decoded_sources.items() if is_multi_value} + return params.keys() + + def _normalize_multi_params( input_dict: MutableMapping[str, Any], params: Sequence[ModelField], + *, + multi_value_params: Collection[str] = (), ) -> MutableMapping[str, Any]: """ Extract and normalize query string or header parameters with Pydantic model support. @@ -596,6 +624,8 @@ def _normalize_multi_params( A dictionary containing the initial query string or header parameters. params: Sequence[ModelField] A sequence of ModelField objects representing parameters. + multi_value_params: Collection[str] + Names supplied as native multiple values rather than comma-separated strings. Returns ------- @@ -604,23 +634,44 @@ def _normalize_multi_params( """ for param in params: if is_scalar_field(param): - _process_scalar_param(input_dict, param) + _process_scalar_param(input_dict, param, multi_value_params) elif lenient_issubclass(param.field_info.annotation, BaseModel): - _process_model_param(input_dict, param) + _process_model_param(input_dict, param, multi_value_params) return input_dict -def _process_scalar_param(input_dict: MutableMapping[str, Any], param: ModelField) -> None: - """Process a scalar parameter by normalizing single-item lists.""" +def _restore_scalar_parameter(value: Any, name: str, multi_value_params: Collection[str]) -> Any: + """Reconstruct a scalar string split by the event, leaving native multiple values intact.""" + if ( + isinstance(value, list) + and len(value) > 1 + and name not in multi_value_params + and all(isinstance(item, str) for item in value) + ): + return ",".join(value) + return value + + +def _process_scalar_param( + input_dict: MutableMapping[str, Any], + param: ModelField, + multi_value_params: Collection[str], +) -> None: + """Restore scalar strings and unwrap single-item lists.""" try: - value = input_dict[param.alias] + value = _restore_scalar_parameter(input_dict[param.alias], param.alias, multi_value_params) if isinstance(value, list) and len(value) == 1: - input_dict[param.alias] = value[0] + value = value[0] + input_dict[param.alias] = value except KeyError: pass -def _process_model_param(input_dict: MutableMapping[str, Any], param: ModelField) -> None: +def _process_model_param( + input_dict: MutableMapping[str, Any], + param: ModelField, + multi_value_params: Collection[str], +) -> None: """Process a Pydantic model parameter by extracting model fields.""" model_class = cast(type[BaseModel], param.field_info.annotation) @@ -630,6 +681,9 @@ def _process_model_param(input_dict: MutableMapping[str, Any], param: ModelField value = _get_param_value(input_dict, field_alias, field_name, model_class) if value is not None: + if not _is_or_contains_sequence(field_info.annotation): + source_name = field_alias if input_dict.get(field_alias) is not None else field_name + value = _restore_scalar_parameter(value, source_name, multi_value_params) model_data[field_alias] = _normalize_field_value(value=value, field_info=field_info) input_dict[param.alias] = model_data diff --git a/tests/functional/event_handler/_pydantic/test_openapi_comma_parameters.py b/tests/functional/event_handler/_pydantic/test_openapi_comma_parameters.py new file mode 100644 index 00000000000..cccca599794 --- /dev/null +++ b/tests/functional/event_handler/_pydantic/test_openapi_comma_parameters.py @@ -0,0 +1,367 @@ +import json +from copy import deepcopy +from typing import Annotated + +import pytest +from pydantic import BaseModel, ConfigDict, Field + +from aws_lambda_powertools.event_handler import ( + ALBResolver, + APIGatewayHttpResolver, + APIGatewayRestResolver, + LambdaFunctionUrlResolver, + VPCLatticeResolver, + VPCLatticeV2Resolver, +) +from aws_lambda_powertools.event_handler.openapi.params import Body, Form, Header, Query +from tests.functional.utils import load_event + +RESOLVERS = [ + (APIGatewayHttpResolver, "apiGatewayProxyV2Event.json"), + (LambdaFunctionUrlResolver, "lambdaFunctionUrlEventWithHeaders.json"), + (APIGatewayRestResolver, "apiGatewayProxyEvent.json"), + (ALBResolver, "albEvent.json"), + (VPCLatticeResolver, "vpcLatticeEvent.json"), + (VPCLatticeV2Resolver, "vpcLatticeV2EventWithHeaders.json"), +] + + +@pytest.fixture(params=RESOLVERS, ids=lambda entry: entry[0].__name__) +def resolver_event(request): + resolver, fixture = request.param + app = resolver(enable_validation=True) + event = load_event(fixture) + event.pop("multiValueQueryStringParameters", None) + event.pop("multiValueHeaders", None) + event.update(path="/search", rawPath="/search", raw_path="/search", httpMethod="GET", method="GET") + if "http" in event.get("requestContext", {}): + event["requestContext"]["http"].update(method="GET", path="/search") + event["requestContext"]["stage"] = "$default" + event["headers"] = {} + event["queryStringParameters"] = {} + event["query_string_parameters"] = {} + event["rawQueryString"] = "" + event["body"] = None + return app, event + + +def set_query(app, event, query): + if isinstance(app, VPCLatticeV2Resolver): + event["queryStringParameters"] = {key: [value] for key, value in query.items()} + elif isinstance(app, VPCLatticeResolver): + event["query_string_parameters"] = query + else: + event["queryStringParameters"] = query + + +@pytest.mark.parametrize("value", ["hello,world", ",hello", "hello,", "hello,, world", "", "plain"]) +@pytest.mark.parametrize("location", [Query, Header], ids=["query", "header"]) +def test_scalar_parameters_preserve_commas(resolver_event, value, location): + app, event = resolver_event + + @app.get("/search") + def handler(search: Annotated[str, location(alias="X-Search")]): + return {"search": search} + + if location is Header: + event["headers"] = {"x-search": value} + else: + set_query(app, event, {"X-Search": value}) + original_event = deepcopy(event) + + result = app.resolve(event, {}) + + assert result["statusCode"] == 200 + assert json.loads(result["body"]) == {"search": value} + assert event == original_event + + +@pytest.mark.parametrize("value", ["hello,world", ",hello", "hello,", "hello,, world", ""]) +@pytest.mark.parametrize("location", [Query, Header], ids=["query", "header"]) +def test_model_parameters_preserve_commas(resolver_event, value, location): + app, event = resolver_event + + class Params(BaseModel): + search: str = Field(alias="x-search") + + @app.get("/search") + def handler(params: Annotated[Params, location()]): + return params.model_dump() + + if location is Header: + event["headers"] = {"X-Search": value} + else: + set_query(app, event, {"x-search": value}) + + result = app.resolve(event, {}) + + assert result["statusCode"] == 200 + assert json.loads(result["body"]) == {"search": value} + + +def test_query_model_preserves_scalar_and_sequence_fields(resolver_event): + app, event = resolver_event + + class Params(BaseModel): + model_config = ConfigDict(populate_by_name=True) + search: str | None = Field(default=None, alias="term") + tags: list[str] | None = None + + @app.get("/search") + def handler(params: Annotated[Params, Query()]): + return params.model_dump() + + set_query(app, event, {"search": "hello, world", "tags": "a,,b"}) + + result = app.resolve(event, {}) + + assert result["statusCode"] == 200 + assert json.loads(result["body"]) == {"search": "hello, world", "tags": ["a", "", "b"]} + + +def test_sequence_query_parameters_still_split_commas(resolver_event): + app, event = resolver_event + + @app.get("/search") + def handler(tags: Annotated[list[str], Query()]): + return {"tags": tags} + + set_query(app, event, {"tags": "a, b,,c"}) + + result = app.resolve(event, {}) + + assert result["statusCode"] == 200 + assert json.loads(result["body"]) == {"tags": ["a", " b", "", "c"]} + + +@pytest.mark.parametrize("annotation,value", [(int, "1,2"), (float, "1.5,2.5"), (bool, "true,false")]) +def test_query_model_validates_the_entire_scalar(resolver_event, annotation, value): + app, event = resolver_event + + class Params(BaseModel): + search: annotation + + @app.get("/search") + def handler(params: Annotated[Params, Query()]): + pytest.fail("Invalid scalar must not be truncated into a valid value") + + set_query(app, event, {"search": value}) + + result = app.resolve(event, {}) + + assert result["statusCode"] == 422 + + +@pytest.mark.parametrize("resolver", [APIGatewayHttpResolver, LambdaFunctionUrlResolver]) +@pytest.mark.parametrize("model", [False, True], ids=["scalar", "model"]) +def test_header_sequences_keep_splitting_commas(resolver, model): + app = resolver(enable_validation=True) + event = load_event("apiGatewayProxyV2Event.json") + event.update(rawPath="/search", body=None) + event["requestContext"]["http"]["method"] = "GET" + event["headers"] = {"X-Tags": "a, b,,c"} + + if model: + + class Params(BaseModel): + tags: list[str] = Field(alias="x-tags") + + @app.get("/search") + def model_handler(params: Annotated[Params, Header()]): + return params.model_dump() + else: + + @app.get("/search") + def scalar_handler(tags: Annotated[list[str], Header(alias="X-Tags")]): + return {"tags": tags} + + result = app.resolve(event, {}) + + assert result["statusCode"] == 200 + assert json.loads(result["body"]) == {"tags": ["a", " b", "", "c"]} + + +@pytest.mark.parametrize("resolver", [APIGatewayHttpResolver, LambdaFunctionUrlResolver]) +@pytest.mark.parametrize( + "raw_query,query,expected_status", + [ + ("search=hello,world", "hello,world", 200), + ("search=hello%2Cworld", "hello,world", 200), + ("search=hello&search=world", "hello,world", 422), + ("%73earch=hello&search=world", "hello,world", 422), + ("search=&search=world", ",world", 422), + ], +) +def test_http_v2_distinguishes_repeated_query_values(resolver, raw_query, query, expected_status): + app = resolver(enable_validation=True) + event = load_event("apiGatewayProxyV2Event.json") + event.update(rawPath="/search", rawQueryString=raw_query, body=None) + event["requestContext"]["http"]["method"] = "GET" + event["queryStringParameters"] = {"search": query} + + @app.get("/search") + def handler(search: str): + return {"search": search} + + result = app.resolve(event, {}) + + assert result["statusCode"] == expected_status + if expected_status == 200: + assert json.loads(result["body"]) == {"search": query} + + +@pytest.mark.parametrize("model", [False, True], ids=["scalar", "model"]) +@pytest.mark.parametrize("location", [Query, Header], ids=["query", "header"]) +@pytest.mark.parametrize("resolver", [APIGatewayRestResolver, ALBResolver]) +def test_native_repeated_parameters_keep_existing_behavior(resolver, location, model): + app = resolver(enable_validation=True) + event = load_event("albEvent.json" if resolver is ALBResolver else "apiGatewayProxyEvent.json") + event.update(path="/search", httpMethod="GET") + event["headers"] = {} + event["multiValueHeaders"] = {} + event["queryStringParameters"] = {"search": "ignored"} + event["multiValueQueryStringParameters"] = {} + event["body"] = None + key = "multiValueHeaders" if location is Header else "multiValueQueryStringParameters" + event[key] = {"SEARCH" if location is Header else "search": ["first", "second"]} + + if model: + + class Params(BaseModel): + model_config = ConfigDict(populate_by_name=True) + search: str = Field(alias="term") + + @app.get("/search") + def model_handler(params: Annotated[Params, location()]): + return params.model_dump() + else: + + @app.get("/search") + def scalar_handler(search: Annotated[str, location()]): + return {"search": search} + + result = app.resolve(event, {}) + + if model: + assert result["statusCode"] == 200 + assert json.loads(result["body"]) == {"search": "first"} + else: + assert result["statusCode"] == 422 + + +@pytest.mark.parametrize("resolver", [APIGatewayRestResolver, ALBResolver]) +def test_native_multivalue_query_preserves_literal_commas_and_fallback(resolver): + app = resolver(enable_validation=True) + event = load_event("albEvent.json" if resolver is ALBResolver else "apiGatewayProxyEvent.json") + event.update(path="/search", httpMethod="GET") + event["queryStringParameters"] = {"search": "ignored", "fallback": "hello,world"} + event["multiValueQueryStringParameters"] = {"search": ["first,second"], "tags": ["a,b", "c"]} + event["body"] = None + + @app.get("/search") + def handler(search: str, fallback: str, tags: Annotated[list[str], Query()]): + return {"search": search, "fallback": fallback, "tags": tags} + + result = app.resolve(event, {}) + + assert result["statusCode"] == 200 + assert json.loads(result["body"]) == { + "search": "first,second", + "fallback": "hello,world", + "tags": ["a,b", "c"], + } + + +@pytest.mark.parametrize("value", ["hello%2Cworld", "hello,world", "hello%2Cworld, again"]) +def test_alb_decoding_preserves_commas(value): + app = ALBResolver(enable_validation=True, decode_query_parameters=True) + event = load_event("albEvent.json") + event["path"] = "/search" + event["body"] = None + event["queryStringParameters"] = {"%73earch": value} + + @app.get("/search") + def handler(search: str): + return {"search": search} + + result = app.resolve(event, {}) + + assert result["statusCode"] == 200 + assert json.loads(result["body"]) == {"search": value.replace("%2C", ",")} + + +def test_alb_decoding_keeps_native_repeated_parameters(): + app = ALBResolver(enable_validation=True, decode_query_parameters=True) + event = load_event("albEvent.json") + event["path"] = "/search" + event["body"] = None + event["queryStringParameters"] = {"search": "ignored"} + event["multiValueQueryStringParameters"] = {"%73earch": ["first", "second"]} + + @app.get("/search") + def handler(search: str): + return {"search": search} + + result = app.resolve(event, {}) + + assert result["statusCode"] == 422 + + +def test_alb_decoding_preserves_precedence_when_parameter_names_collide(): + app = ALBResolver(enable_validation=True, decode_query_parameters=True) + event = load_event("albEvent.json") + event["path"] = "/search" + event["body"] = None + event["queryStringParameters"] = {"search": "ignored", "%73earch": "hello,world"} + event["multiValueQueryStringParameters"] = {"search": ["first", "second"]} + + @app.get("/search") + def handler(search: str): + return {"search": search} + + result = app.resolve(event, {}) + + # The later encoded key wins when ALB decodes the merged parameter mapping. + assert result["statusCode"] == 200 + assert json.loads(result["body"]) == {"search": "hello,world"} + + +def test_lattice_v2_keeps_native_repeated_parameters(): + app = VPCLatticeV2Resolver(enable_validation=True) + event = load_event("vpcLatticeV2EventWithHeaders.json") + event.update(path="/search", method="GET", body=None) + event["queryStringParameters"] = {"search": ["first", "second"]} + + @app.get("/search") + def handler(search: str): + return {"search": search} + + result = app.resolve(event, {}) + + assert result["statusCode"] == 422 + + +@pytest.mark.parametrize("content_type", ["application/json", "application/x-www-form-urlencoded"]) +def test_body_scalar_lists_keep_existing_normalization(content_type): + app = APIGatewayRestResolver(enable_validation=True) + event = load_event("apiGatewayProxyEvent.json") + event.update(path="/search", httpMethod="POST", isBase64Encoded=False) + event["headers"] = {"Content-Type": content_type} + event["multiValueHeaders"] = {} + event["queryStringParameters"] = {} + event["multiValueQueryStringParameters"] = {} + if content_type == "application/json": + param = Body(embed=True) + event["body"] = json.dumps({"search": ["first", "second"]}) + else: + param = Form() + event["body"] = "search=first&search=second" + + @app.post("/search") + def handler(search: Annotated[str, param]): + return {"search": search} + + result = app.resolve(event, {}) + + assert result["statusCode"] == 200 + assert json.loads(result["body"]) == {"search": "first"}