diff --git a/src/postgrest/src/postgrest/base_request_builder.py b/src/postgrest/src/postgrest/base_request_builder.py index 2a562859..94910dea 100644 --- a/src/postgrest/src/postgrest/base_request_builder.py +++ b/src/postgrest/src/postgrest/base_request_builder.py @@ -40,7 +40,7 @@ from .base_client import BasePostgrestClient from .types import JSON, CountMethod, Filters, JSONAdapter, RequestMethod, ReturnMethod -from .utils import sanitize_param +from .utils import sanitize_array_param, sanitize_param class QueryArgs(NamedTuple): @@ -455,12 +455,10 @@ def in_(self: Self, column: str, values: Iterable[Any]) -> Self: return self.filter(column, Filters.IN, f"({values})") def cs(self: Self, column: str, values: Iterable[Any]) -> Self: - values = ",".join(str(v) for v in values) - return self.filter(column, Filters.CS, f"{{{values}}}") + return self.filter(column, Filters.CS, sanitize_array_param(values)) def cd(self: Self, column: str, values: Iterable[Any]) -> Self: - values = ",".join(str(v) for v in values) - return self.filter(column, Filters.CD, f"{{{values}}}") + return self.filter(column, Filters.CD, sanitize_array_param(values)) def contains( self: Self, column: str, value: Union[Iterable[Any], str, Dict[Any, Any]] @@ -471,8 +469,7 @@ def contains( return self.filter(column, Filters.CS, value) if not isinstance(value, dict) and isinstance(value, Iterable): # Expected to be some type of iterable - stringified_values = ",".join(str(v) for v in value) - return self.filter(column, Filters.CS, f"{{{stringified_values}}}") + return self.filter(column, Filters.CS, sanitize_array_param(value)) return self.filter(column, Filters.CS, json.dumps(value)) @@ -483,8 +480,7 @@ def contained_by( # range return self.filter(column, Filters.CD, value) if not isinstance(value, dict) and isinstance(value, Iterable): - stringified_values = ",".join(str(v) for v in value) - return self.filter(column, Filters.CD, f"{{{stringified_values}}}") + return self.filter(column, Filters.CD, sanitize_array_param(value)) return self.filter(column, Filters.CD, json.dumps(value)) def ov(self: Self, column: str, value: Iterable[Any]) -> Self: @@ -494,8 +490,7 @@ def ov(self: Self, column: str, value: Iterable[Any]) -> Self: return self.filter(column, Filters.OV, value) if not isinstance(value, dict) and isinstance(value, Iterable): # Expected to be some type of iterable - stringified_values = ",".join(str(v) for v in value) - return self.filter(column, Filters.OV, f"{{{stringified_values}}}") + return self.filter(column, Filters.OV, sanitize_array_param(value)) return self.filter(column, Filters.OV, json.dumps(value)) def sl(self: Self, column: str, range: Tuple[int, int]) -> Self: diff --git a/src/postgrest/src/postgrest/utils.py b/src/postgrest/src/postgrest/utils.py index 0fa2cbec..bc4ec6d3 100644 --- a/src/postgrest/src/postgrest/utils.py +++ b/src/postgrest/src/postgrest/utils.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Any, Type, TypeVar, cast, get_origin +from typing import Any, Iterable, Type, TypeVar, cast, get_origin from urllib.parse import urlparse from deprecation import deprecated @@ -37,6 +37,35 @@ def sanitize_param(param: Any) -> str: return param_str +def sanitize_array_element(element: Any) -> str: + """Quote a single value for use inside a PostgreSQL array literal ``{...}``. + + PostgreSQL quotes an array element when it is empty, matches ``NULL`` + case-insensitively, or contains a brace, the comma delimiter, a double + quote, a backslash, or whitespace; embedded double quotes and backslashes + are backslash-escaped. Without this, a value such as ``"a,b"`` is emitted + bare and parsed by PostgREST as two separate elements. + """ + element_str = str(element) + needs_quoting = ( + element_str == "" + or element_str.lower() == "null" + or any(char in element_str for char in '{},"\\') + or any(char.isspace() for char in element_str) + ) + if not needs_quoting: + return element_str + escaped = element_str.replace("\\", "\\\\").replace('"', '\\"') + return f'"{escaped}"' + + +def sanitize_array_param(values: Iterable[Any]) -> str: + """Render an iterable as a PostgreSQL array literal ``{elem,elem,...}``, + quoting each element as needed so values containing the comma delimiter (or + other special characters) survive as single elements.""" + return f"{{{','.join(sanitize_array_element(value) for value in values)}}}" + + def sanitize_pattern_param(pattern: str) -> str: return sanitize_param(pattern.replace("%", "*")) diff --git a/src/postgrest/tests/_async/test_filter_request_builder.py b/src/postgrest/tests/_async/test_filter_request_builder.py index f6e6fecd..c1a4346a 100644 --- a/src/postgrest/tests/_async/test_filter_request_builder.py +++ b/src/postgrest/tests/_async/test_filter_request_builder.py @@ -155,10 +155,48 @@ def test_contains_in_list(filter_request_builder): def test_contained_by_mixed_items(filter_request_builder): + # The second element is the literal string '["b", "c"]'. It contains the + # array delimiter, so it has to be quoted (with its inner quotes escaped); + # emitting it bare as {a,["b", "c"]} produced a corrupted literal that + # PostgREST would have split into several elements. builder = filter_request_builder.contained_by("x", ["a", '["b", "c"]']) - # {a,["b",+"c"]} - assert str(builder.request.params) == "x=cd.%7Ba%2C%5B%22b%22%2C+%22c%22%5D%7D" + assert builder.request.params["x"] == 'cd.{a,"[\\"b\\", \\"c\\"]"}' + + +def test_contains_quotes_element_containing_the_delimiter(filter_request_builder): + # A value containing the comma delimiter must be quoted, otherwise + # PostgREST reads {a,b,c} as three elements rather than the two passed. + builder = filter_request_builder.contains("x", ["a,b", "c"]) + + assert builder.request.params["x"] == 'cs.{"a,b",c}' + + +def test_cs_quotes_element_containing_the_delimiter(filter_request_builder): + builder = filter_request_builder.cs("x", ["a,b"]) + + assert builder.request.params["x"] == 'cs.{"a,b"}' + + +def test_cd_quotes_element_containing_the_delimiter(filter_request_builder): + builder = filter_request_builder.cd("x", ["a,b"]) + + assert builder.request.params["x"] == 'cd.{"a,b"}' + + +def test_overlaps_quotes_element_containing_the_delimiter(filter_request_builder): + builder = filter_request_builder.overlaps("x", ["a,b"]) + + assert builder.request.params["x"] == 'ov.{"a,b"}' + + +def test_contains_quotes_braces_quotes_whitespace_empty_and_null(filter_request_builder): + # Braces, embedded double quotes, whitespace, the empty string and the word + # NULL all force quoting; embedded quotes are backslash-escaped so each + # element round-trips as a single value. + builder = filter_request_builder.contains("x", ["a}b", 'x"y', "a b", "", "NULL"]) + + assert builder.request.params["x"] == 'cs.{"a}b","x\\"y","a b","","NULL"}' def test_range_greater_than(filter_request_builder): diff --git a/src/postgrest/tests/_sync/test_filter_request_builder.py b/src/postgrest/tests/_sync/test_filter_request_builder.py index 11cc9959..4a9fa036 100644 --- a/src/postgrest/tests/_sync/test_filter_request_builder.py +++ b/src/postgrest/tests/_sync/test_filter_request_builder.py @@ -155,10 +155,48 @@ def test_contains_in_list(filter_request_builder): def test_contained_by_mixed_items(filter_request_builder): + # The second element is the literal string '["b", "c"]'. It contains the + # array delimiter, so it has to be quoted (with its inner quotes escaped); + # emitting it bare as {a,["b", "c"]} produced a corrupted literal that + # PostgREST would have split into several elements. builder = filter_request_builder.contained_by("x", ["a", '["b", "c"]']) - # {a,["b",+"c"]} - assert str(builder.request.params) == "x=cd.%7Ba%2C%5B%22b%22%2C+%22c%22%5D%7D" + assert builder.request.params["x"] == 'cd.{a,"[\\"b\\", \\"c\\"]"}' + + +def test_contains_quotes_element_containing_the_delimiter(filter_request_builder): + # A value containing the comma delimiter must be quoted, otherwise + # PostgREST reads {a,b,c} as three elements rather than the two passed. + builder = filter_request_builder.contains("x", ["a,b", "c"]) + + assert builder.request.params["x"] == 'cs.{"a,b",c}' + + +def test_cs_quotes_element_containing_the_delimiter(filter_request_builder): + builder = filter_request_builder.cs("x", ["a,b"]) + + assert builder.request.params["x"] == 'cs.{"a,b"}' + + +def test_cd_quotes_element_containing_the_delimiter(filter_request_builder): + builder = filter_request_builder.cd("x", ["a,b"]) + + assert builder.request.params["x"] == 'cd.{"a,b"}' + + +def test_overlaps_quotes_element_containing_the_delimiter(filter_request_builder): + builder = filter_request_builder.overlaps("x", ["a,b"]) + + assert builder.request.params["x"] == 'ov.{"a,b"}' + + +def test_contains_quotes_braces_quotes_whitespace_empty_and_null(filter_request_builder): + # Braces, embedded double quotes, whitespace, the empty string and the word + # NULL all force quoting; embedded quotes are backslash-escaped so each + # element round-trips as a single value. + builder = filter_request_builder.contains("x", ["a}b", 'x"y', "a b", "", "NULL"]) + + assert builder.request.params["x"] == 'cs.{"a}b","x\\"y","a b","","NULL"}' def test_range_greater_than(filter_request_builder):