Skip to content
Closed
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
12 changes: 6 additions & 6 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_element, sanitize_param


class QueryArgs(NamedTuple):
Expand Down Expand Up @@ -455,11 +455,11 @@ 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)
values = ",".join(sanitize_array_element(v) for v in values)
return self.filter(column, Filters.CS, f"{{{values}}}")

def cd(self: Self, column: str, values: Iterable[Any]) -> Self:
values = ",".join(str(v) for v in values)
values = ",".join(sanitize_array_element(v) for v in values)
return self.filter(column, Filters.CD, f"{{{values}}}")

def contains(
Expand All @@ -471,7 +471,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)
stringified_values = ",".join(sanitize_array_element(v) for v in value)
return self.filter(column, Filters.CS, f"{{{stringified_values}}}")

return self.filter(column, Filters.CS, json.dumps(value))
Expand All @@ -483,7 +483,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)
stringified_values = ",".join(sanitize_array_element(v) for v in value)
return self.filter(column, Filters.CD, f"{{{stringified_values}}}")
return self.filter(column, Filters.CD, json.dumps(value))

Expand All @@ -494,7 +494,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)
stringified_values = ",".join(sanitize_array_element(v) for v in value)
return self.filter(column, Filters.OV, f"{{{stringified_values}}}")
return self.filter(column, Filters.OV, json.dumps(value))

Expand Down
19 changes: 19 additions & 0 deletions src/postgrest/src/postgrest/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,25 @@ def sanitize_param(param: Any) -> str:
return param_str


def sanitize_array_element(param: Any) -> str:
"""Quote a single value for use inside a PostgreSQL array literal `{...}`.

Follows https://www.postgresql.org/docs/current/arrays.html :
quote when empty, NULL (case-insensitive), or containing
braces, comma, double quote, backslash, or whitespace.
Embedded backslashes and double quotes are backslash-escaped.
"""
param_str = str(param)
if (
param_str == ""
or param_str.upper() == "NULL"
or any(char in param_str for char in '{},"\\ \t\n\r')
):
param_str = param_str.replace("\\", "\\\\").replace('"', '\\"')
return f'"{param_str}"'
return param_str


def sanitize_pattern_param(pattern: str) -> str:
return sanitize_param(pattern.replace("%", "*"))

Expand Down