summaryrefslogtreecommitdiff
path: root/server/szurubooru/search/executor.py
diff options
context:
space:
mode:
authorHunternif <hunternif@gmail.com>2021-08-04 01:56:00 +0100
committerHunternif <hunternif@gmail.com>2021-08-04 01:56:00 +0100
commitca861cdc44ca476cec2949f237bbc7503862ad58 (patch)
tree520f918a9a77432bf283b46b4296812eb4332ffe /server/szurubooru/search/executor.py
parentdd03540398fcab11bfc2c3cebed2fb7a78592c16 (diff)
parent59452711668b9f7a7aeea3c57301b22882987c2b (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.py123
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())