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("%", "*"))