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
17 changes: 6 additions & 11 deletions src/postgrest/src/postgrest/base_request_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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]]
Expand All @@ -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))

Expand All @@ -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:
Expand All @@ -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:
Expand Down
31 changes: 30 additions & 1 deletion src/postgrest/src/postgrest/utils.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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("%", "*"))

Expand Down
42 changes: 40 additions & 2 deletions src/postgrest/tests/_async/test_filter_request_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
42 changes: 40 additions & 2 deletions src/postgrest/tests/_sync/test_filter_request_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down