diff options
| author | Hunternif <hunternif@gmail.com> | 2021-08-04 01:56:00 +0100 |
|---|---|---|
| committer | Hunternif <hunternif@gmail.com> | 2021-08-04 01:56:00 +0100 |
| commit | ca861cdc44ca476cec2949f237bbc7503862ad58 (patch) | |
| tree | 520f918a9a77432bf283b46b4296812eb4332ffe /server/szurubooru/search/executor.py | |
| parent | dd03540398fcab11bfc2c3cebed2fb7a78592c16 (diff) | |
| parent | 59452711668b9f7a7aeea3c57301b22882987c2b (diff) | |
Merge remote-tracking branch 'origin/master' into hunternif
# Conflicts:
# client/css/post-content-control.styl
# client/css/post-list-view.styl
# client/html/post_edit_sidebar.tpl
# client/js/controllers/post_list_controller.js
# client/js/controllers/post_main_controller.js
# client/js/controllers/post_upload_controller.js
# client/js/controllers/tag_controller.js
# client/js/controllers/user_list_controller.js
# client/js/controls/expander_control.js
# client/js/controls/post_content_control.js
# client/js/controls/post_edit_sidebar_control.js
# client/js/controls/post_readonly_sidebar_control.js
# client/js/controls/tag_input_control.js
# client/js/main.js
# client/js/models/abstract_list.js
# client/js/models/post.js
# client/js/models/post_list.js
# client/js/models/settings.js
# client/js/models/tag.js
# client/js/models/tag_list.js
# client/js/tags.js
# client/js/util/search.js
# client/js/util/touch.js
# client/js/util/uri.js
# client/js/util/views.js
# client/js/views/post_main_view.js
# client/js/views/post_upload_view.js
# client/js/views/posts_header_view.js
# client/js/views/posts_page_view.js
# client/js/views/settings_view.js
# client/js/views/tag_view.js
# client/package-lock.json
# client/package.json
# server/config.yaml.dist
# server/szurubooru/api/__init__.py
# server/szurubooru/api/post_api.py
# server/szurubooru/api/tag_api.py
# server/szurubooru/func/posts.py
# server/szurubooru/func/tags.py
# server/szurubooru/model/__init__.py
# server/szurubooru/model/post.py
# server/szurubooru/model/tag.py
# server/szurubooru/search/configs/__init__.py
# server/szurubooru/search/configs/post_search_config.py
# server/szurubooru/search/executor.py
# server/szurubooru/tests/api/test_post_retrieving.py
# server/szurubooru/tests/api/test_post_updating.py
# server/szurubooru/tests/api/test_tag_updating.py
# server/szurubooru/tests/conftest.py
# server/szurubooru/tests/func/test_posts.py
# server/szurubooru/tests/func/test_tags.py
# server/szurubooru/tests/search/configs/test_post_search_config.py
Diffstat (limited to 'server/szurubooru/search/executor.py')
| -rw-r--r-- | server/szurubooru/search/executor.py | 123 |
1 files changed, 64 insertions, 59 deletions
diff --git a/server/szurubooru/search/executor.py b/server/szurubooru/search/executor.py index 3a8ec68..92a7e2f 100644 --- a/server/szurubooru/search/executor.py +++ b/server/szurubooru/search/executor.py @@ -1,11 +1,13 @@ -from typing import Union, Tuple, List, Dict, Callable +from typing import Callable, Dict, List, Tuple, Union + import sqlalchemy as sa -from szurubooru import db, model, errors, rest + +from szurubooru import db, errors, model, rest from szurubooru.func import cache -from szurubooru.search import tokens, parser -from szurubooru.search.typing import SaQuery -from szurubooru.search.query import SearchQuery +from szurubooru.search import parser, tokens from szurubooru.search.configs.base_search_config import BaseSearchConfig +from szurubooru.search.query import SearchQuery +from szurubooru.search.typing import SaQuery def _format_dict_keys(source: Dict) -> List[str]: @@ -25,10 +27,10 @@ def _get_order(order: str, default_order: str) -> Union[bool, str]: 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' @@ -37,9 +39,8 @@ class Executor: self.parser = parser.Parser() def get_around( - self, - query_text: str, - entity_id: int) -> Tuple[model.Base, model.Base, model.Base]: + self, query_text: str, entity_id: int + ) -> Tuple[model.Base, model.Base, model.Base]: search_query = self.parser.parse(query_text) self.config.on_search_query_parsed(search_query) filter_query = ( @@ -68,21 +69,19 @@ class Executor: self, ctx: rest.Context, entity_id: int, - serializer: Callable[[model.Base], rest.Response] + serializer: Callable[[model.Base], rest.Response], ) -> rest.Response: entities = self.get_around( - ctx.get_param_as_string('query', default=''), entity_id) + ctx.get_param_as_string("query", default=""), entity_id + ) return { - 'prev': serializer(entities[0]), - 'next': serializer(entities[1]), - 'random': serializer(entities[2]), + "prev": serializer(entities[0]), + "next": serializer(entities[1]), + "random": serializer(entities[2]), } def execute( - self, - query_text: str, - offset: int, - limit: int + self, query_text: str, offset: int, limit: int ) -> Tuple[int, List[model.Base]]: search_query = self.parser.parse(query_text) self.config.on_search_query_parsed(search_query) @@ -93,7 +92,7 @@ class Executor: disable_eager_loads = False for token in search_query.sort_tokens: - if token.name == 'random': + if token.name == "random": disable_eager_loads = True key = (id(self.config), hash(search_query), offset, limit) @@ -101,22 +100,16 @@ class Executor: return cache.get(key) filter_query = self.config.create_filter_query(disable_eager_loads) - filter_query = filter_query.options(sa.orm.lazyload('*')) + filter_query = filter_query.options(sa.orm.lazyload("*")) filter_query = self._prepare_db_query(filter_query, search_query, True) - entities = ( - filter_query - .offset(offset) - .limit(limit) - .all()) + entities = filter_query.offset(offset).limit(limit).all() count_query = self.config.create_count_query(disable_eager_loads) - count_query = count_query.options(sa.orm.lazyload('*')) + 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_statement = count_query.statement.with_only_columns( + [sa.func.count()] + ).order_by(None) count = db.session.execute(count_statement).scalar() ret = (count, entities) @@ -126,18 +119,18 @@ class Executor: def execute_and_serialize( self, ctx: rest.Context, - serializer: Callable[[model.Base], rest.Response] + serializer: Callable[[model.Base], rest.Response], ) -> rest.Response: - query = ctx.get_param_as_string('query', default='') - offset = ctx.get_param_as_int('offset', default=0, min=0) - limit = ctx.get_param_as_int('limit', default=100, min=1, max=100) + query = ctx.get_param_as_string("query", default="") + offset = ctx.get_param_as_int("offset", default=0, min=0) + limit = ctx.get_param_as_int("limit", default=100, min=1, max=100) count, entities = self.execute(query, offset, limit) return { - 'query': query, - 'offset': offset, - 'limit': limit, - 'total': count, - 'results': list([serializer(entity) for entity in entities]), + "query": query, + "offset": offset, + "limit": limit, + "total": count, + "results": list([serializer(entity) for entity in entities]), } def count(self, query_text:str) -> int: @@ -155,46 +148,58 @@ class Executor: return count def _prepare_db_query( - self, - db_query: SaQuery, - search_query: SearchQuery, - use_sort: bool) -> SaQuery: + self, db_query: SaQuery, search_query: SearchQuery, use_sort: bool + ) -> SaQuery: for anon_token in search_query.anonymous_tokens: if not self.config.anonymous_filter: raise errors.SearchError( - 'Anonymous tokens are not valid in this context.') + "Anonymous tokens are not valid in this context." + ) db_query = self.config.anonymous_filter( - db_query, anon_token.criterion, anon_token.negated) + db_query, anon_token.criterion, anon_token.negated + ) for named_token in search_query.named_tokens: if named_token.name not in self.config.named_filters: raise errors.SearchError( - 'Unknown named token: %r. Available named tokens: %r.' % ( + "Unknown named token: %r. Available named tokens: %r." + % ( named_token.name, - _format_dict_keys(self.config.named_filters))) + _format_dict_keys(self.config.named_filters), + ) + ) db_query = self.config.named_filters[named_token.name]( - db_query, named_token.criterion, named_token.negated) + db_query, named_token.criterion, named_token.negated + ) for sp_token in search_query.special_tokens: if sp_token.value not in self.config.special_filters: raise errors.SearchError( - 'Unknown special token: %r. ' - 'Available special tokens: %r.' % ( + "Unknown special token: %r. " + "Available special tokens: %r." + % ( sp_token.value, - _format_dict_keys(self.config.special_filters))) + _format_dict_keys(self.config.special_filters), + ) + ) db_query = self.config.special_filters[sp_token.value]( - db_query, None, sp_token.negated) + db_query, None, sp_token.negated + ) if use_sort: for sort_token in search_query.sort_tokens: if sort_token.name not in self.config.sort_columns: raise errors.SearchError( - 'Unknown sort token: %r. ' - 'Available sort tokens: %r.' % ( + "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]) + _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) if order == sort_token.SORT_ASC: db_query = db_query.order_by(column.asc()) |