3737_PDB_IDENTIFIER_RE = re .compile (r"^[A-Za-z_][A-Za-z0-9_]*$" )
3838
3939
40+ def _text_literal (value : str ) -> ClauseElement :
41+ return literal (value , type_ = Text ())
42+
43+
44+ def _text_array (values : Sequence [str ]) -> ClauseElement :
45+ return array ([str (value ) for value in values ], type_ = Text ())
46+
47+
4048def _inline_string_literal (value : str ) -> ClauseElement :
4149 return literal_column ("'" + value .replace ("'" , "''" ) + "'" , Text ())
4250
@@ -46,8 +54,8 @@ def _to_term_payload(*terms: str) -> ClauseElement:
4654 raise InvalidArgumentError ("at least one search term is required" )
4755 require_non_empty_strings (terms , field_name = "terms" )
4856 if len (terms ) == 1 :
49- return literal (terms [0 ])
50- return array ( list ( terms ), type_ = Text () )
57+ return _text_literal (terms [0 ])
58+ return _text_array ( terms )
5159
5260
5361def _apply_boost (expr : ClauseElement , boost : float | None ) -> ClauseElement :
@@ -83,13 +91,13 @@ def _apply_tokenizer(expr: ClauseElement, tokenizer: str | None) -> ClauseElemen
8391def _to_phrase_payload (value : str | Sequence [str ]) -> ClauseElement :
8492 if isinstance (value , str ):
8593 require_non_empty_string (value , field_name = "value" )
86- return literal (value )
94+ return _text_literal (value )
8795 if not isinstance (value , Sequence ):
8896 raise InvalidArgumentError ("value must be a non-empty string or a sequence of non-empty strings" )
8997 if not value :
9098 raise InvalidArgumentError ("value must contain at least one token" )
9199 require_non_empty_strings (value , field_name = "value" )
92- return array ( list ( value ), type_ = Text () )
100+ return _text_array ( value )
93101
94102
95103def _apply_score_tuning (
@@ -171,7 +179,7 @@ def term(
171179 tokenizer : str | None = None ,
172180) -> ColumnElement [bool ]:
173181 require_non_empty_string (value , field_name = "value" )
174- payload : ClauseElement = literal (value )
182+ payload : ClauseElement = _text_literal (value )
175183 payload = _apply_fuzzy (payload , distance = distance , prefix = prefix , transpose_cost_one = transpose_cost_one )
176184 payload = _apply_tokenizer (payload , tokenizer )
177185 payload = _apply_score_tuning (payload , boost = boost , const = const )
@@ -209,7 +217,7 @@ def regex(
209217 const : float | None = None ,
210218) -> ColumnElement [bool ]:
211219 require_non_empty_string (pattern , field_name = "pattern" )
212- payload : ClauseElement = func .pdb .regex (pattern )
220+ payload : ClauseElement = func .pdb .regex (_text_literal ( pattern ) )
213221 payload = _apply_score_tuning (payload , boost = boost , const = const )
214222 return field .operate (_QUERY , payload )
215223
@@ -239,8 +247,8 @@ def prox_regex(pattern: str, max_expansions: int | None = None) -> ProximityExpr
239247 require_non_empty_string (pattern , field_name = "pattern" )
240248 if max_expansions is not None :
241249 require_non_negative (max_expansions , field_name = "max_expansions" )
242- return ProximityExpr (func .pdb .prox_regex (pattern , max_expansions ))
243- return ProximityExpr (func .pdb .prox_regex (pattern ))
250+ return ProximityExpr (func .pdb .prox_regex (_text_literal ( pattern ) , max_expansions ))
251+ return ProximityExpr (func .pdb .prox_regex (_text_literal ( pattern ) ))
244252
245253
246254def prox_array (* clauses : str | ProximityExpr ) -> ProximityExpr :
@@ -294,7 +302,7 @@ def parse(
294302 field : ColumnElement , query : str , * , lenient : bool = False , conjunction_mode : bool = False
295303) -> ColumnElement [bool ]:
296304 require_non_empty_string (query , field_name = "query" )
297- return field .operate (_QUERY , func .pdb .parse (query , lenient , conjunction_mode ))
305+ return field .operate (_QUERY , func .pdb .parse (_text_literal ( query ) , lenient , conjunction_mode ))
298306
299307
300308def phrase_prefix (field : ColumnElement , terms : list [str ], * , max_expansions : int | None = None ) -> ColumnElement [bool ]:
@@ -303,9 +311,9 @@ def phrase_prefix(field: ColumnElement, terms: list[str], *, max_expansions: int
303311 require_non_empty_strings (terms , field_name = "terms" )
304312 if max_expansions is not None :
305313 require_positive (max_expansions , field_name = "max_expansions" )
306- return field .operate (_QUERY , func .pdb .phrase_prefix (array (terms , type_ = Text () ), max_expansions ))
314+ return field .operate (_QUERY , func .pdb .phrase_prefix (_text_array (terms ), max_expansions ))
307315 else :
308- return field .operate (_QUERY , func .pdb .phrase_prefix (array (terms , type_ = Text () )))
316+ return field .operate (_QUERY , func .pdb .phrase_prefix (_text_array (terms )))
309317
310318
311319def regex_phrase (
@@ -321,9 +329,9 @@ def regex_phrase(
321329 require_non_negative (slop , field_name = "slop" )
322330 if max_expansions is not None :
323331 require_positive (max_expansions , field_name = "max_expansions" )
324- return field .operate (_QUERY , func .pdb .regex_phrase (array (terms , type_ = Text () ), slop , max_expansions ))
332+ return field .operate (_QUERY , func .pdb .regex_phrase (_text_array (terms ), slop , max_expansions ))
325333 else :
326- return field .operate (_QUERY , func .pdb .regex_phrase (array (terms , type_ = Text () ), slop ))
334+ return field .operate (_QUERY , func .pdb .regex_phrase (_text_array (terms ), slop ))
327335
328336
329337def range_term (
@@ -373,7 +381,7 @@ def range_term(
373381 escaped = bounds .replace ("'" , "''" )
374382 range_bounds_arg : ClauseElement = literal_column (f"'{ escaped } '::{ range_type } " )
375383 else :
376- range_bounds_arg = literal (bounds )
384+ range_bounds_arg = _text_literal (bounds )
377385 return field .operate (_QUERY , func .pdb .range_term (range_bounds_arg , relation_arg ))
378386
379387
@@ -461,12 +469,12 @@ def more_like_this(
461469 if boost_factor is not None :
462470 named_options .append (("boost_factor" , boost_factor ))
463471 if stopwords is not None :
464- named_options .append (("stopwords" , array (stopwords , type_ = Text () )))
472+ named_options .append (("stopwords" , _text_array (stopwords )))
465473
466474 def _build_mlt_call (source_arg : ClauseElement , * , include_fields : bool ) -> ClauseElement :
467475 positional_args : list [ClauseElement ] = [source_arg ]
468476 if include_fields and fields is not None :
469- positional_args .append (array (fields , type_ = Text () ))
477+ positional_args .append (_text_array (fields ))
470478 return PDBFunctionWithNamedArgs ("more_like_this" , positional_args , named_options )
471479
472480 if document_ids is not None :
@@ -479,4 +487,5 @@ def _build_mlt_call(source_arg: ClauseElement, *, include_fields: bool) -> Claus
479487 return field .operate (_QUERY , _build_mlt_call (literal (document_id ), include_fields = True ))
480488
481489 payload = document if isinstance (document , str ) else json .dumps (document , separators = ("," , ":" ), sort_keys = True )
482- return field .operate (_QUERY , _build_mlt_call (literal (payload ), include_fields = False ))
490+ payload_arg = _text_literal (payload ) if isinstance (payload , str ) else literal (payload )
491+ return field .operate (_QUERY , _build_mlt_call (payload_arg , include_fields = False ))
0 commit comments