diff options
Diffstat (limited to 'server/szurubooru/search/executor.py')
| -rw-r--r-- | server/szurubooru/search/executor.py | 124 |
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 |