Skip to content

Commit c0f1fb8

Browse files
committed
adding supports for tokenizer_params
1 parent 2814c8a commit c0f1fb8

1 file changed

Lines changed: 15 additions & 7 deletions

File tree

paradedb/sqlalchemy/search.py

Lines changed: 15 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,7 @@
33
import json
44
import re
55
from collections.abc import Sequence
6-
from typing import Any
6+
from typing import Any, Tuple
77

88
from sqlalchemy import Text, cast, func, literal, literal_column, or_
99
from sqlalchemy.dialects.postgresql import ARRAY, array
@@ -35,6 +35,7 @@
3535
_PROXIMITY_ORDERED: Any = operators.custom_op("##>", precedence=5)
3636
_PDB_IDENTIFIER_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
3737
_TextClause = str | ClauseElement
38+
_TOKENIZER_PARAMS = Sequence[Any]
3839

3940

4041
def _text_literal(value: str) -> ClauseElement:
@@ -86,11 +87,14 @@ def _validate_pdb_identifier(name: str, *, field_name: str) -> str:
8687
return name
8788

8889

89-
def _apply_tokenizer(expr: ClauseElement, tokenizer: str | None) -> ClauseElement:
90+
def _apply_tokenizer(
91+
expr: ClauseElement, tokenizer: str | None,
92+
tokenizer_params: _TOKENIZER_PARAMS = (),
93+
) -> ClauseElement:
9094
if tokenizer is None:
9195
return expr
9296
tokenizer_name = _validate_pdb_identifier(tokenizer, field_name="tokenizer")
93-
return PDBCast(expr, tokenizer_name)
97+
return PDBCast(expr, tokenizer_name, tokenizer_params)
9498

9599

96100
def _to_phrase_payload(value: _TextClause | Sequence[_TextClause]) -> ClauseElement:
@@ -147,10 +151,11 @@ def match_all(
147151
prefix: bool = False,
148152
transpose_cost_one: bool = False,
149153
tokenizer: str | None = None,
154+
tokenizer_params: _TOKENIZER_PARAMS = (),
150155
) -> ColumnElement[bool]:
151156
payload = _to_term_payload(*terms)
152157
payload = _apply_fuzzy(payload, distance=distance, prefix=prefix, transpose_cost_one=transpose_cost_one)
153-
payload = _apply_tokenizer(payload, tokenizer)
158+
payload = _apply_tokenizer(payload, tokenizer, tokenizer_params)
154159
payload = _apply_score_tuning(payload, boost=boost, const=const)
155160
return field.operate(_MATCH_ALL, payload)
156161

@@ -164,10 +169,11 @@ def match_any(
164169
prefix: bool = False,
165170
transpose_cost_one: bool = False,
166171
tokenizer: str | None = None,
172+
tokenizer_params: _TOKENIZER_PARAMS = (),
167173
) -> ColumnElement[bool]:
168174
payload = _to_term_payload(*terms)
169175
payload = _apply_fuzzy(payload, distance=distance, prefix=prefix, transpose_cost_one=transpose_cost_one)
170-
payload = _apply_tokenizer(payload, tokenizer)
176+
payload = _apply_tokenizer(payload, tokenizer, tokenizer_params)
171177
payload = _apply_score_tuning(payload, boost=boost, const=const)
172178
return field.operate(_MATCH_ANY, payload)
173179

@@ -182,10 +188,11 @@ def term(
182188
prefix: bool = False,
183189
transpose_cost_one: bool = False,
184190
tokenizer: str | None = None,
191+
tokenizer_params: _TOKENIZER_PARAMS = (),
185192
) -> ColumnElement[bool]:
186193
payload: ClauseElement = _to_text_clause(value)
187194
payload = _apply_fuzzy(payload, distance=distance, prefix=prefix, transpose_cost_one=transpose_cost_one)
188-
payload = _apply_tokenizer(payload, tokenizer)
195+
payload = _apply_tokenizer(payload, tokenizer, tokenizer_params)
189196
payload = _apply_score_tuning(payload, boost=boost, const=const)
190197
return field.operate(_TERM, payload)
191198

@@ -198,12 +205,13 @@ def phrase(
198205
boost: float | None = None,
199206
const: float | None = None,
200207
tokenizer: str | None = None,
208+
tokenizer_params: _TOKENIZER_PARAMS = (),
201209
) -> ColumnElement[bool]:
202210
if slop is not None:
203211
require_non_negative(slop, field_name="slop")
204212
is_token_array = isinstance(value, Sequence) and not isinstance(value, (str, ClauseElement))
205213
payload: ClauseElement = _to_phrase_payload(value)
206-
payload = _apply_tokenizer(payload, tokenizer)
214+
payload = _apply_tokenizer(payload, tokenizer, tokenizer_params)
207215
if slop is not None:
208216
# psycopg binds array elements as VARCHAR by default; slop cast requires TEXT[].
209217
if is_token_array:

0 commit comments

Comments
 (0)