diff --git a/aws_lambda_powertools/utilities/data_classes/common.py b/aws_lambda_powertools/utilities/data_classes/common.py index c8908eda9ce..ded7f7c2fc6 100644 --- a/aws_lambda_powertools/utilities/data_classes/common.py +++ b/aws_lambda_powertools/utilities/data_classes/common.py @@ -63,6 +63,27 @@ def update(self, data=None, **kwargs): super().update((k.lower(), v) for k, v in data) super().update((k.lower(), v) for k, v in kwargs.items()) + def copy(self): + return CaseInsensitiveDict(self) + + def __or__(self, other): + if not isinstance(other, Mapping): + return NotImplemented + new = self.copy() + new.update(other) + return new + + def __ror__(self, other): + if not isinstance(other, Mapping): + return NotImplemented + new = CaseInsensitiveDict(other) + new.update(self) + return new + + def __ior__(self, other): + self.update(other) + return self + def __contains__(self, k): return super().__contains__(k.lower()) diff --git a/tests/unit/data_classes/required_dependencies/test_common.py b/tests/unit/data_classes/required_dependencies/test_common.py index 2a2ad1f0f93..eb157f188c4 100644 --- a/tests/unit/data_classes/required_dependencies/test_common.py +++ b/tests/unit/data_classes/required_dependencies/test_common.py @@ -17,3 +17,31 @@ def test_case_insensitive_dict_update_with_kwargs(): assert headers["AB"] == "1" assert headers["user_agent"] == "test" assert len(headers) == 3 + + +def test_case_insensitive_dict_copy_keeps_case_insensitive_lookup(): + headers = CaseInsensitiveDict({"Content-Type": "application/json"}) + copied = headers.copy() + + assert isinstance(copied, CaseInsensitiveDict) + assert copied.get("Content-Type") == "application/json" + + copied["X-Trace-Id"] = "abc" + assert "x-trace-id" not in headers + + +def test_case_insensitive_dict_merge_operators(): + headers = CaseInsensitiveDict({"Content-Type": "application/json"}) + + merged = headers | {"X-Trace-Id": "abc"} + assert isinstance(merged, CaseInsensitiveDict) + assert merged["x-trace-id"] == "abc" + assert merged["CONTENT-TYPE"] == "application/json" + + merged = {"X-Trace-Id": "abc"} | headers + assert isinstance(merged, CaseInsensitiveDict) + assert merged["X-TRACE-ID"] == "abc" + + headers |= {"X-Trace-Id": "abc"} + assert headers["x-trace-id"] == "abc" + assert list(headers) == ["content-type", "x-trace-id"]