summaryrefslogtreecommitdiff
path: root/server/szurubooru/search/executor.py
diff options
context:
space:
mode:
Diffstat (limited to 'server/szurubooru/search/executor.py')
-rw-r--r--server/szurubooru/search/executor.py124
1 files changed, 106 insertions, 18 deletions
diff --git a/server/szurubooru/search/executor.py b/server/szurubooru/search/executor.py
index a5ef962..5302b14 100644
--- a/server/szurubooru/search/executor.py
+++ b/server/szurubooru/search/executor.py
@@ -31,6 +31,8 @@ class Executor:
Class for search parsing and execution. Handles plaintext parsing and
delegates sqlalchemy filter decoration to SearchConfig instances.
"""
+ AROUND_NEXT = "up"
+ AROUND_PREV = "down"
def __init__(self, search_config: BaseSearchConfig) -> None:
self.config = search_config
@@ -38,30 +40,29 @@ class Executor:
def get_around(
self, query_text: str, entity_id: int
- ) -> Tuple[model.Base, model.Base]:
+ ) -> Tuple[model.Base, model.Base, model.Base]:
search_query = self.parser.parse(query_text)
self.config.on_search_query_parsed(search_query)
- filter_query = self.config.create_around_query().options(
- sa.orm.lazyload("*")
- )
- filter_query = self._prepare_db_query(
- filter_query, search_query, False
- )
+ filter_query = (
+ self.config
+ .create_around_query()
+ .options(sa.orm.lazyload("*")))
prev_filter_query = (
- filter_query.filter(self.config.id_column > entity_id)
- .order_by(None)
- .order_by(sa.func.abs(self.config.id_column - entity_id).asc())
- .limit(1)
- )
+ self._prepare_sorted_around_query(
+ filter_query, search_query, entity_id, self.AROUND_PREV
+ ).limit(1))
next_filter_query = (
- filter_query.filter(self.config.id_column < entity_id)
- .order_by(None)
- .order_by(sa.func.abs(self.config.id_column - entity_id).asc())
- .limit(1)
- )
+ self._prepare_sorted_around_query(
+ filter_query, search_query, entity_id, self.AROUND_NEXT
+ ).limit(1))
+ # random post
+ if "sort:random" not in query_text:
+ query_text = "sort:random " + query_text
+ count, random_entities = self.execute(query_text, 0, 1)
return (
prev_filter_query.one_or_none(),
next_filter_query.one_or_none(),
+ random_entities[0] if random_entities else None
)
def get_around_and_serialize(
@@ -76,6 +77,7 @@ class Executor:
return {
"prev": serializer(entities[0]),
"next": serializer(entities[1]),
+ "random": serializer(entities[2]),
}
def execute(
@@ -94,7 +96,7 @@ class Executor:
disable_eager_loads = True
key = (id(self.config), hash(search_query), offset, limit)
- if cache.has(key):
+ if not disable_eager_loads and cache.has(key):
return cache.get(key)
filter_query = self.config.create_filter_query(disable_eager_loads)
@@ -131,6 +133,20 @@ class Executor:
"results": list([serializer(entity) for entity in entities]),
}
+ def count(self, query_text:str) -> int:
+ search_query = self.parser.parse(query_text)
+ self.config.on_search_query_parsed(search_query)
+ count_query = self.config.create_count_query(True)
+ count_query = count_query.options(sa.orm.lazyload("*"))
+ count_query = self._prepare_db_query(count_query, search_query, False)
+ count_statement = (
+ count_query
+ .statement
+ .with_only_columns([sa.func.count()])
+ .order_by(None))
+ count = db.session.execute(count_statement).scalar()
+ return count
+
def _prepare_db_query(
self, db_query: SaQuery, search_query: SearchQuery, use_sort: bool
) -> SaQuery:
@@ -192,3 +208,75 @@ class Executor:
db_query = self.config.finalize_query(db_query)
return db_query
+
+ def _prepare_sorted_around_query(
+ self,
+ db_query: SaQuery,
+ search_query: SearchQuery,
+ entity_id: int,
+ direction: str):
+ db_query = self._prepare_db_query(db_query, search_query, False)
+ db_query = db_query.order_by(None)
+ found_sort_column = False
+
+ for sort_token in search_query.sort_tokens:
+ if sort_token.name == "random":
+ continue
+ if sort_token.name not in self.config.sort_columns:
+ raise errors.SearchError(
+ "Unknown sort token: %r. "
+ "Available sort tokens: %r." % (
+ sort_token.name,
+ _format_dict_keys(self.config.sort_columns)))
+ column, default_order = (
+ self.config.sort_columns[sort_token.name])
+ order = _get_order(sort_token.order, default_order)
+
+ # the order column may be joined, so we need to query its value:
+ column_query = (
+ db.session.query(self.config.id_column, column)
+ .options(sa.orm.lazyload("*")))
+ column_query = (
+ # empty search query because we already know entity id
+ self._prepare_db_query(column_query, SearchQuery(), False)
+ .filter(self.config.id_column == entity_id))
+ id, column_value = column_query.one_or_none()
+ # it's possible that this entity doesn't have the column
+ if not column_value:
+ continue
+ found_sort_column = True
+
+ if order == sort_token.SORT_ASC:
+ if direction == self.AROUND_NEXT:
+ db_query = (
+ db_query
+ .order_by(column.asc())
+ .filter(column > column_value))
+ elif direction == self.AROUND_PREV:
+ db_query = (
+ db_query
+ .order_by(column.desc())
+ .filter(column < column_value))
+ elif order == sort_token.SORT_DESC:
+ if direction == self.AROUND_NEXT:
+ db_query = (
+ db_query
+ .order_by(column.desc())
+ .filter(column < column_value))
+ elif direction == self.AROUND_PREV:
+ db_query = (
+ db_query
+ .order_by(column.asc())
+ .filter(column > column_value))
+
+ if not found_sort_column:
+ # no sorting, use default sorting by id
+ if direction == self.AROUND_NEXT:
+ db_query = db_query.filter(self.config.id_column < entity_id)
+ elif direction == self.AROUND_PREV:
+ db_query = db_query.filter(self.config.id_column > entity_id)
+ db_query = db_query.order_by(
+ sa.func.abs(self.config.id_column - entity_id).asc())
+ return db_query
+
+ return db_query