Skip to content

Commit 8760604

Browse files
authored
fix: Accept DB functions in queries (#138)
1 parent a585e7e commit 8760604

5 files changed

Lines changed: 133 additions & 61 deletions

File tree

CHANGELOG.md

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,12 @@ All notable changes to this project will be documented in this file. The format
66

77
## Unreleased
88

9+
## [0.9.0] - 2026-06-18
10+
11+
### Added
12+
13+
- Added support for passing DB functions into search queries.
14+
915
## [0.8.0] - 2026-05-21
1016

1117
### Changed
@@ -203,6 +209,7 @@ All notable changes to this project will be documented in this file. The format
203209
- JSON field key indexing support
204210
- Full Django ORM integration with `Q` objects and standard filters
205211

212+
[0.9.0]: https://github.com/paradedb/django-paradedb/compare/v0.8.0...v0.9.0
206213
[0.8.0]: https://github.com/paradedb/django-paradedb/compare/v0.7.0...v0.8.0
207214
[0.7.0]: https://github.com/paradedb/django-paradedb/compare/v0.6.0...v0.7.0
208215
[0.6.0]: https://github.com/paradedb/django-paradedb/compare/v0.5.0...v0.6.0

paradedb/search.py

Lines changed: 62 additions & 57 deletions
Original file line numberDiff line numberDiff line change
@@ -65,7 +65,7 @@
6565
PDB_TYPE_TOKENIZER_WHITESPACE,
6666
)
6767

68-
SearchValue: TypeAlias = "str | list[str] | tuple[str, ...] | Modifier"
68+
SearchValue: TypeAlias = "str | list[str] | tuple[str, ...] | Modifier | Expression"
6969
Modifiable: TypeAlias = "SearchValue | QueryExpression"
7070

7171

@@ -720,19 +720,6 @@ def _validate_numeric_params(self) -> None:
720720
raise ValueError(f"MoreLikeThis {param_name} must be >= 1.")
721721

722722

723-
def render_more_like_this(term: MoreLikeThis, lhs_sql: str) -> tuple[str, list[Any]]:
724-
params: list[Any] = []
725-
726-
if term.id is not None:
727-
mlt_sql, mlt_params = _render_more_like_this_call(term, term.id)
728-
params.extend(mlt_params)
729-
return f"{lhs_sql} {OP_SEARCH} {mlt_sql}", params
730-
731-
mlt_sql, mlt_params = _render_more_like_this_call(term, term.document)
732-
params.extend(mlt_params)
733-
return f"{lhs_sql} {OP_SEARCH} {mlt_sql}", params
734-
735-
736723
def _render_more_like_this_call(
737724
term: MoreLikeThis, value: object
738725
) -> tuple[str, list[Any]]:
@@ -854,14 +841,12 @@ def resolve_expression(
854841

855842
def as_sql(
856843
self,
857-
_compiler: SQLCompiler,
844+
compiler: SQLCompiler,
858845
_connection: BaseDatabaseWrapper,
859846
lhs_sql: str,
860847
) -> tuple[str, list[object]]:
861-
if isinstance(self._term, MoreLikeThis):
862-
return render_more_like_this(self._term, lhs_sql)
863-
rendered = self._render_term(self._term)
864-
return f"{lhs_sql} {rendered}", []
848+
rendered, params = self._render_term(self._term, compiler)
849+
return f"{lhs_sql} {rendered}", params
865850

866851
@staticmethod
867852
def _unwrap_term(term: TermType) -> TermType:
@@ -881,34 +866,40 @@ def _quote_range_literal(
881866
return f"{_quote_term(literal)}::{safe_range_type}"
882867

883868
@staticmethod
884-
def _render_search_value(value: object) -> str:
869+
def _render_search_value(
870+
value: object, compiler: SQLCompiler
871+
) -> tuple[str, list[Any]]:
885872
if isinstance(value, Boost):
886-
rendered = ParadeDB._render_search_value(value.value)
887-
return f"{rendered}::{PDB_TYPE_BOOST}({value.factor})"
873+
rendered, params = ParadeDB._render_search_value(value.value, compiler)
874+
return f"{rendered}::{PDB_TYPE_BOOST}({value.factor})", params
888875
if isinstance(value, Const):
889-
rendered = ParadeDB._render_search_value(value.value)
890-
return f"{rendered}::{PDB_TYPE_CONST}({value.score})"
876+
rendered, params = ParadeDB._render_search_value(value.value, compiler)
877+
return f"{rendered}::{PDB_TYPE_CONST}({value.score})", params
891878
if isinstance(value, Fuzzy):
892-
rendered = ParadeDB._render_search_value(value.value)
879+
rendered, params = ParadeDB._render_search_value(value.value, compiler)
893880
fuzzy_args = [str(value.distance)]
894881
if value.prefix is not None:
895882
fuzzy_args.append("t" if value.prefix else "f")
896883
if value.transposition_cost_one is not None:
897884
fuzzy_args.append("t" if value.transposition_cost_one else "f")
898-
return f"{rendered}::{PDB_TYPE_FUZZY}({', '.join(fuzzy_args)})"
885+
return f"{rendered}::{PDB_TYPE_FUZZY}({', '.join(fuzzy_args)})", params
899886
if isinstance(value, Slop):
900-
rendered = ParadeDB._render_search_value(value.value)
901-
return f"{rendered}::{PDB_TYPE_SLOP}({value.distance})"
887+
rendered, params = ParadeDB._render_search_value(value.value, compiler)
888+
return f"{rendered}::{PDB_TYPE_SLOP}({value.distance})", params
902889
if isinstance(value, Tokenized):
903-
rendered = ParadeDB._render_search_value(value.value)
904-
return f"{rendered}::{value.tokenizer.render()}"
890+
rendered, params = ParadeDB._render_search_value(value.value, compiler)
891+
return f"{rendered}::{value.tokenizer.render()}", params
905892
if isinstance(value, QueryExpression):
906-
return ParadeDB(value)._render_term(value)
893+
return ParadeDB(value)._render_term(value, compiler)
894+
if isinstance(value, Expression):
895+
expression = value.resolve_expression(compiler.query)
896+
sql, expression_params = compiler.compile(expression)
897+
return sql, list(expression_params)
907898
if isinstance(value, str):
908-
return _quote_term(value)
899+
return _quote_term(value), []
909900
if isinstance(value, list | tuple):
910901
quoted = [_quote_term(item) for item in value]
911-
return f"ARRAY[{', '.join(quoted)}]"
902+
return f"ARRAY[{', '.join(quoted)}]", []
912903
raise TypeError(f"Unsupported search value type. {value}")
913904

914905
def _render_proximity_node(self, node: ProximityNode) -> str:
@@ -941,71 +932,85 @@ def _render_proximity_term(self, term: ProximityTerm) -> str:
941932
return f"{FN_PROX_ARRAY}({', '.join(parts)})"
942933
raise AssertionError(f"Unhandled proximity term: {term!r}")
943934

944-
def _render_term(self, term: TermType) -> str:
935+
def _render_term(
936+
self, term: TermType, compiler: SQLCompiler
937+
) -> tuple[str, list[Any]]:
945938
if isinstance(term, Boost | Const | Fuzzy | Slop | Tokenized):
946-
return self._render_search_value(term)
939+
return self._render_search_value(term, compiler)
947940
if isinstance(term, Phrase):
948-
rendered = self._render_search_value(
949-
term.terms[0] if len(term.terms) == 1 else term.terms
941+
rendered, params = self._render_search_value(
942+
term.terms[0] if len(term.terms) == 1 else term.terms, compiler
950943
)
951-
return f"{OP_PHRASE} {rendered}"
944+
return f"{OP_PHRASE} {rendered}", params
952945
if isinstance(term, ProximityNode):
953-
return self._render_proximity_node(term)
946+
return self._render_proximity_node(term), []
954947
if isinstance(term, Parse):
955948
rendered = (
956949
f"{OP_SEARCH} {FN_PARSE}({_quote_term(term.query)}"
957950
f"{self._render_options({'lenient': term.lenient, 'conjunction_mode': term.conjunction_mode})})"
958951
)
959-
return rendered
952+
return rendered, []
960953
if isinstance(term, PhrasePrefix):
961954
phrases_sql = ", ".join(_quote_term(phrase) for phrase in term.phrases)
962955
return (
963956
f"{OP_SEARCH} {FN_PHRASE_PREFIX}(ARRAY[{phrases_sql}]"
964957
f"{self._render_options({'max_expansion': term.max_expansion})})"
965-
)
958+
), []
966959
if isinstance(term, RegexPhrase):
967960
regex_sql = ", ".join(_quote_term(regex) for regex in term.regexes)
968961
return (
969962
f"{OP_SEARCH} {FN_REGEX_PHRASE}(ARRAY[{regex_sql}]"
970963
f"{self._render_options({'slop': term.slop, 'max_expansions': term.max_expansions})})"
971-
)
964+
), []
972965
if isinstance(term, RangeTerm):
973966
if term.relation is None:
974-
return f"{OP_SEARCH} {FN_RANGE_TERM}({self._render_value(term.value)})"
967+
return (
968+
f"{OP_SEARCH} {FN_RANGE_TERM}({self._render_value(term.value)})",
969+
[],
970+
)
975971
else:
976972
assert term.range_type is not None
977973
return (
978974
f"{OP_SEARCH} {FN_RANGE_TERM}("
979975
f"{self._quote_range_literal(term.value, term.range_type)}, "
980976
f"{_quote_term(term.relation)}"
981977
")"
982-
)
978+
), []
983979
if isinstance(term, Term):
984-
return f"{OP_SEARCH} {FN_TERM}({self._render_search_value(term.value)})"
980+
rendered, params = self._render_search_value(term.value, compiler)
981+
return f"{OP_SEARCH} {FN_TERM}({rendered})", params
985982
if isinstance(term, Regex):
986-
return f"{OP_SEARCH} {FN_REGEX}({_quote_term(term.pattern)})"
983+
return f"{OP_SEARCH} {FN_REGEX}({_quote_term(term.pattern)})", []
987984
if isinstance(term, Exists):
988-
return f"{OP_SEARCH} {FN_EXISTS}()"
985+
return f"{OP_SEARCH} {FN_EXISTS}()", []
989986
if isinstance(term, FuzzyTerm):
990987
if term.value is not None:
991-
return f"{OP_SEARCH} {FN_FUZZY_TERM}({_quote_term(term.value)})"
988+
return f"{OP_SEARCH} {FN_FUZZY_TERM}({_quote_term(term.value)})", []
992989
else:
993-
return f"{OP_SEARCH} {FN_FUZZY_TERM}()"
990+
return f"{OP_SEARCH} {FN_FUZZY_TERM}()", []
994991
if isinstance(term, TermSet):
995992
array_sql = self._render_term_set_array(term.terms)
996-
return f"{OP_SEARCH} {FN_TERM_SET}({array_sql})"
993+
return f"{OP_SEARCH} {FN_TERM_SET}({array_sql})", []
997994
if isinstance(term, All):
998-
return f"{OP_SEARCH} {FN_ALL}()"
995+
return f"{OP_SEARCH} {FN_ALL}()", []
999996
if isinstance(term, MatchAll):
1000-
rendered = self._render_search_value(
1001-
term.terms[0] if len(term.terms) == 1 else term.terms
997+
rendered, params = self._render_search_value(
998+
term.terms[0] if len(term.terms) == 1 else term.terms, compiler
1002999
)
1003-
return f"{OP_AND} {rendered}"
1000+
return f"{OP_AND} {rendered}", params
10041001
if isinstance(term, MatchAny):
1005-
rendered = self._render_search_value(
1006-
term.terms[0] if len(term.terms) == 1 else term.terms
1002+
rendered, params = self._render_search_value(
1003+
term.terms[0] if len(term.terms) == 1 else term.terms, compiler
10071004
)
1008-
return f"{OP_OR} {rendered}"
1005+
return f"{OP_OR} {rendered}", params
1006+
if isinstance(term, MoreLikeThis):
1007+
if term.id is not None:
1008+
mlt_sql, mlt_params = _render_more_like_this_call(term, term.id)
1009+
return f"{OP_SEARCH} {mlt_sql}", mlt_params
1010+
else:
1011+
mlt_sql, mlt_params = _render_more_like_this_call(term, term.document)
1012+
return f"{OP_SEARCH} {mlt_sql}", mlt_params
1013+
10091014
raise TypeError(f"Unsupported ParadeDB term type. {term}")
10101015

10111016
@staticmethod

pyproject.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
44

55
[project]
66
name = "django-paradedb"
7-
version = "0.8.0"
7+
version = "0.9.0"
88
description = "Official ParadeDB integration for Django"
99
readme = "README.md"
1010
license = "MIT"

tests/test_queries.py

Lines changed: 62 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -7,9 +7,10 @@
77
"""
88

99
import pytest
10+
from django.contrib.postgres.fields import ArrayField
1011
from django.db import connection
11-
from django.db.models import F, Q, Window
12-
from django.db.models.functions import Coalesce
12+
from django.db.models import F, Func, Q, TextField, Value, Window
13+
from django.db.models.functions import Cast, Coalesce, Trim
1314

1415
from paradedb.functions import Agg, Score, Snippet, SnippetPositions, Snippets
1516
from paradedb.search import (
@@ -262,6 +263,65 @@ def test_and_search_three_terms(self) -> None:
262263
)
263264
_run_query(queryset)
264265

266+
def test_match_all_accepts_database_function_value(self) -> None:
267+
queryset = MockItem.objects.filter(
268+
description=ParadeDB(MatchAll(Trim(Value("shoes "))))
269+
)
270+
sql, params = queryset.query.sql_with_params()
271+
assert (
272+
sql
273+
== 'SELECT "mock_items"."id", "mock_items"."description", "mock_items"."category", "mock_items"."rating", "mock_items"."in_stock", "mock_items"."created_at", "mock_items"."metadata" FROM "mock_items" WHERE "mock_items"."description" &&& TRIM(%s)'
274+
)
275+
assert params == ("shoes ",)
276+
_run_query(queryset)
277+
278+
def test_match_all_accepts_cast_database_function_value(self) -> None:
279+
queryset = MockItem.objects.filter(
280+
description=ParadeDB(
281+
MatchAll(
282+
Cast(
283+
Func(
284+
Value("running,shoes"),
285+
Value(","),
286+
function="string_to_array",
287+
),
288+
ArrayField(TextField()),
289+
)
290+
)
291+
)
292+
)
293+
sql, params = queryset.query.sql_with_params()
294+
assert (
295+
sql
296+
== 'SELECT "mock_items"."id", "mock_items"."description", "mock_items"."category", "mock_items"."rating", "mock_items"."in_stock", "mock_items"."created_at", "mock_items"."metadata" FROM "mock_items" WHERE "mock_items"."description" &&& (string_to_array(%s, %s))::text[]'
297+
)
298+
assert params == ("running,shoes", ",")
299+
_run_query(queryset)
300+
301+
def test_match_any_accepts_database_function_value(self) -> None:
302+
queryset = MockItem.objects.filter(
303+
description=ParadeDB(MatchAny(Trim(Value("shoes "))))
304+
)
305+
sql, params = queryset.query.sql_with_params()
306+
assert (
307+
sql
308+
== 'SELECT "mock_items"."id", "mock_items"."description", "mock_items"."category", "mock_items"."rating", "mock_items"."in_stock", "mock_items"."created_at", "mock_items"."metadata" FROM "mock_items" WHERE "mock_items"."description" ||| TRIM(%s)'
309+
)
310+
assert params == ("shoes ",)
311+
_run_query(queryset)
312+
313+
def test_term_accepts_database_function_value(self) -> None:
314+
queryset = MockItem.objects.filter(
315+
description=ParadeDB(Term(Trim(Value("shoes "))))
316+
)
317+
sql, params = queryset.query.sql_with_params()
318+
assert (
319+
sql
320+
== 'SELECT "mock_items"."id", "mock_items"."description", "mock_items"."category", "mock_items"."rating", "mock_items"."in_stock", "mock_items"."created_at", "mock_items"."metadata" FROM "mock_items" WHERE "mock_items"."description" @@@ pdb.term(TRIM(%s))'
321+
)
322+
assert params == ("shoes ",)
323+
_run_query(queryset)
324+
265325

266326
class TestMoreLikeThis:
267327
"""Test MoreLikeThis SQL generation."""

uv.lock

Lines changed: 1 addition & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

0 commit comments

Comments
 (0)