From 87021c8cbaa444d2df764b3c7cbdc6bfcef868ee Mon Sep 17 00:00:00 2001 From: daniel Date: Thu, 3 Sep 2026 21:10:09 -0400 Subject: [PATCH] fix(postgrest): escape special chars in array filters so values with commas don't split I ran into this while filtering on a tags column - contains("tags", ["a,b", "c"]) was sending tags=cs.{a,b,c}, which Postgres reads as three elements instead of two. Silent wrong results, no error, took me a while to spot. This adds a small sanitize_array_element helper following the Postgres array literal rules (quote on empty, NULL, braces, comma, quote, backslash, whitespace, with backslash-escaping inside quotes) and uses it in cs, cd, contains, contained_by, and ov/overlaps. Simple values like ["a","b"] and [1,2,3] are untouched. Fixes #1592 --- .../src/postgrest/base_request_builder.py | 12 ++++++------ src/postgrest/src/postgrest/utils.py | 19 +++++++++++++++++++ 2 files changed, 25 insertions(+), 6 deletions(-) diff --git a/src/postgrest/src/postgrest/base_request_builder.py b/src/postgrest/src/postgrest/base_request_builder.py index 2a562859..b0df0c4a 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_element, sanitize_param class QueryArgs(NamedTuple): @@ -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( @@ -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)) @@ -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)) @@ -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)) diff --git a/src/postgrest/src/postgrest/utils.py b/src/postgrest/src/postgrest/utils.py index 0fa2cbec..925158a0 100644 --- a/src/postgrest/src/postgrest/utils.py +++ b/src/postgrest/src/postgrest/utils.py @@ -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("%", "*"))