diff options
| author | Hunternif | 2021-08-04 01:56:00 +0100 |
|---|---|---|
| committer | Hunternif | 2021-08-04 01:56:00 +0100 |
| commit | ca861cdc44ca476cec2949f237bbc7503862ad58 (patch) | |
| tree | 520f918a9a77432bf283b46b4296812eb4332ffe /server/szurubooru/func | |
| 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/func')
24 files changed, 2049 insertions, 1084 deletions
diff --git a/server/szurubooru/func/auth.py b/server/szurubooru/func/auth.py index 65be79a..d013775 100644 --- a/server/szurubooru/func/auth.py +++ b/server/szurubooru/func/auth.py @@ -1,60 +1,67 @@ -from typing import Tuple, Optional import hashlib import random import uuid from collections import OrderedDict from datetime import datetime +from typing import Optional, Tuple + from nacl import pwhash from nacl.exceptions import InvalidkeyError -from szurubooru import config, db, model, errors -from szurubooru.func import util +from szurubooru import config, db, errors, model +from szurubooru.func import util -RANK_MAP = OrderedDict([ - (model.User.RANK_ANONYMOUS, 'anonymous'), - (model.User.RANK_RESTRICTED, 'restricted'), - (model.User.RANK_REGULAR, 'regular'), - (model.User.RANK_POWER, 'power'), - (model.User.RANK_MODERATOR, 'moderator'), - (model.User.RANK_ADMINISTRATOR, 'administrator'), - (model.User.RANK_NOBODY, 'nobody'), -]) +RANK_MAP = OrderedDict( + [ + (model.User.RANK_ANONYMOUS, "anonymous"), + (model.User.RANK_RESTRICTED, "restricted"), + (model.User.RANK_REGULAR, "regular"), + (model.User.RANK_POWER, "power"), + (model.User.RANK_MODERATOR, "moderator"), + (model.User.RANK_ADMINISTRATOR, "administrator"), + (model.User.RANK_NOBODY, "nobody"), + ] +) def get_password_hash(salt: str, password: str) -> Tuple[str, int]: - ''' Retrieve argon2id password hash. ''' - return pwhash.argon2id.str( - (config.config['secret'] + salt + password).encode('utf8') - ).decode('utf8'), 3 + """ Retrieve argon2id password hash. """ + return ( + pwhash.argon2id.str( + (config.config["secret"] + salt + password).encode("utf8") + ).decode("utf8"), + 3, + ) def get_sha256_legacy_password_hash( - salt: str, password: str) -> Tuple[str, int]: - ''' Retrieve old-style sha256 password hash. ''' + salt: str, password: str +) -> Tuple[str, int]: + """ Retrieve old-style sha256 password hash. """ digest = hashlib.sha256() - digest.update(config.config['secret'].encode('utf8')) - digest.update(salt.encode('utf8')) - digest.update(password.encode('utf8')) + digest.update(config.config["secret"].encode("utf8")) + digest.update(salt.encode("utf8")) + digest.update(password.encode("utf8")) return digest.hexdigest(), 2 def get_sha1_legacy_password_hash(salt: str, password: str) -> Tuple[str, int]: - ''' Retrieve old-style sha1 password hash. ''' + """ Retrieve old-style sha1 password hash. """ digest = hashlib.sha1() - digest.update(b'1A2/$_4xVa') - digest.update(salt.encode('utf8')) - digest.update(password.encode('utf8')) + digest.update(b"1A2/$_4xVa") + digest.update(salt.encode("utf8")) + digest.update(password.encode("utf8")) return digest.hexdigest(), 1 def create_password() -> str: alphabet = { - 'c': list('bcdfghijklmnpqrstvwxyz'), - 'v': list('aeiou'), - 'n': list('0123456789'), + "c": list("bcdfghijklmnpqrstvwxyz"), + "v": list("aeiou"), + "n": list("0123456789"), } - pattern = 'cvcvnncvcv' - return ''.join(random.choice(alphabet[l]) for l in list(pattern)) + pattern = "cvcvnncvcv" + return "".join(random.choice(alphabet[type]) for type in list(pattern)) def is_valid_password(user: model.User, password: str) -> bool: @@ -63,12 +70,13 @@ def is_valid_password(user: model.User, password: str) -> bool: try: return pwhash.verify( - user.password_hash.encode('utf8'), - (config.config['secret'] + salt + password).encode('utf8')) + user.password_hash.encode("utf8"), + (config.config["secret"] + salt + password).encode("utf8"), + ) except InvalidkeyError: possible_hashes = [ get_sha256_legacy_password_hash(salt, password)[0], - get_sha1_legacy_password_hash(salt, password)[0] + get_sha1_legacy_password_hash(salt, password)[0], ] if valid_hash in possible_hashes: # Convert the user password hash to the new hash @@ -82,16 +90,18 @@ def is_valid_password(user: model.User, password: str) -> bool: def is_valid_token(user_token: Optional[model.UserToken]) -> bool: - ''' + """ Token must be enabled and if it has an expiration, it must be greater than now. - ''' + """ if user_token is None: return False if not user_token.enabled: return False - if (user_token.expiration_time is not None - and user_token.expiration_time < datetime.utcnow()): + if ( + user_token.expiration_time is not None + and user_token.expiration_time < datetime.utcnow() + ): return False return True @@ -99,26 +109,27 @@ def is_valid_token(user_token: Optional[model.UserToken]) -> bool: def has_privilege(user: model.User, privilege_name: str) -> bool: assert user all_ranks = list(RANK_MAP.keys()) - assert privilege_name in config.config['privileges'] + assert privilege_name in config.config["privileges"] assert user.rank in all_ranks minimal_rank = util.flip(RANK_MAP)[ - config.config['privileges'][privilege_name]] - good_ranks = all_ranks[all_ranks.index(minimal_rank):] + config.config["privileges"][privilege_name] + ] + good_ranks = all_ranks[all_ranks.index(minimal_rank) :] return user.rank in good_ranks def verify_privilege(user: model.User, privilege_name: str) -> None: assert user if not has_privilege(user, privilege_name): - raise errors.AuthError('Insufficient privileges to do this.') + raise errors.AuthError("Insufficient privileges to do this.") def generate_authentication_token(user: model.User) -> str: - ''' Generate nonguessable challenge (e.g. links in password reminder). ''' + """ Generate nonguessable challenge (e.g. links in password reminder). """ assert user digest = hashlib.md5() - digest.update(config.config['secret'].encode('utf8')) - digest.update(user.password_salt.encode('utf8')) + digest.update(config.config["secret"].encode("utf8")) + digest.update(user.password_salt.encode("utf8")) return digest.hexdigest() diff --git a/server/szurubooru/func/cache.py b/server/szurubooru/func/cache.py index 01e4659..65e99e1 100644 --- a/server/szurubooru/func/cache.py +++ b/server/szurubooru/func/cache.py @@ -1,5 +1,5 @@ -from typing import Any, List, Dict from datetime import datetime +from typing import Any, Dict, List class LruCacheItem: @@ -18,12 +18,11 @@ class LruCache: def insert_item(self, item: LruCacheItem) -> None: if item.key in self.hash: item_index = next( - i - for i, v in enumerate(self.item_list) - if v.key == item.key) + i for i, v in enumerate(self.item_list) if v.key == item.key + ) self.item_list[:] = ( - self.item_list[:item_index] + - self.item_list[item_index + 1:]) + self.item_list[:item_index] + self.item_list[item_index + 1 :] + ) self.item_list.insert(0, item) else: if len(self.item_list) > self.length: diff --git a/server/szurubooru/func/comments.py b/server/szurubooru/func/comments.py index 9f88283..5eb7c8e 100644 --- a/server/szurubooru/func/comments.py +++ b/server/szurubooru/func/comments.py @@ -1,7 +1,8 @@ from datetime import datetime -from typing import Any, Optional, List, Dict, Callable -from szurubooru import db, model, errors, rest -from szurubooru.func import users, scores, serialization +from typing import Any, Callable, Dict, List, Optional + +from szurubooru import db, errors, model, rest +from szurubooru.func import scores, serialization, users class InvalidCommentIdError(errors.ValidationError): @@ -23,15 +24,15 @@ class CommentSerializer(serialization.BaseSerializer): def _serializers(self) -> Dict[str, Callable[[], Any]]: return { - 'id': self.serialize_id, - 'user': self.serialize_user, - 'postId': self.serialize_post_id, - 'version': self.serialize_version, - 'text': self.serialize_text, - 'creationTime': self.serialize_creation_time, - 'lastEditTime': self.serialize_last_edit_time, - 'score': self.serialize_score, - 'ownScore': self.serialize_own_score, + "id": self.serialize_id, + "user": self.serialize_user, + "postId": self.serialize_post_id, + "version": self.serialize_version, + "text": self.serialize_text, + "creationTime": self.serialize_creation_time, + "lastEditTime": self.serialize_last_edit_time, + "score": self.serialize_score, + "ownScore": self.serialize_own_score, } def serialize_id(self) -> Any: @@ -63,9 +64,8 @@ class CommentSerializer(serialization.BaseSerializer): def serialize_comment( - comment: model.Comment, - auth_user: model.User, - options: List[str] = []) -> rest.Response: + comment: model.Comment, auth_user: model.User, options: List[str] = [] +) -> rest.Response: if comment is None: return None return CommentSerializer(comment, auth_user).serialize(options) @@ -74,21 +74,22 @@ def serialize_comment( def try_get_comment_by_id(comment_id: int) -> Optional[model.Comment]: comment_id = int(comment_id) return ( - db.session - .query(model.Comment) + db.session.query(model.Comment) .filter(model.Comment.comment_id == comment_id) - .one_or_none()) + .one_or_none() + ) def get_comment_by_id(comment_id: int) -> model.Comment: comment = try_get_comment_by_id(comment_id) if comment: return comment - raise CommentNotFoundError('Comment %r not found.' % comment_id) + raise CommentNotFoundError("Comment %r not found." % comment_id) def create_comment( - user: model.User, post: model.Post, text: str) -> model.Comment: + user: model.User, post: model.Post, text: str +) -> model.Comment: comment = model.Comment() comment.user = user comment.post = post @@ -100,5 +101,5 @@ def create_comment( def update_comment_text(comment: model.Comment, text: str) -> None: assert comment if not text: - raise EmptyCommentTextError('Comment text cannot be empty.') + raise EmptyCommentTextError("Comment text cannot be empty.") comment.text = text diff --git a/server/szurubooru/func/diff.py b/server/szurubooru/func/diff.py index 90014f7..3282ebb 100644 --- a/server/szurubooru/func/diff.py +++ b/server/szurubooru/func/diff.py @@ -1,4 +1,4 @@ -from typing import List, Dict, Any +from typing import Any, Dict, List def get_list_diff(old: List[Any], new: List[Any]) -> Any: @@ -16,8 +16,11 @@ def get_list_diff(old: List[Any], new: List[Any]) -> Any: equal = False added.append(item) - return None if equal else { - 'type': 'list change', 'added': added, 'removed': removed} + return ( + None + if equal + else {"type": "list change", "added": added, "removed": removed} + ) def get_dict_diff(old: Dict[str, Any], new: Dict[str, Any]) -> Any: @@ -40,23 +43,20 @@ def get_dict_diff(old: Dict[str, Any], new: Dict[str, Any]) -> Any: else: equal = False value[key] = { - 'type': 'primitive change', - 'old-value': old[key], - 'new-value': new[key], + "type": "primitive change", + "old-value": old[key], + "new-value": new[key], } else: equal = False - value[key] = { - 'type': 'deleted property', - 'value': old[key] - } + value[key] = {"type": "deleted property", "value": old[key]} for key in new.keys(): if key not in old: equal = False value[key] = { - 'type': 'added property', - 'value': new[key], + "type": "added property", + "value": new[key], } - return None if equal else {'type': 'object change', 'value': value} + return None if equal else {"type": "object change", "value": value} diff --git a/server/szurubooru/func/favorites.py b/server/szurubooru/func/favorites.py index f567bfa..a012537 100644 --- a/server/szurubooru/func/favorites.py +++ b/server/szurubooru/func/favorites.py @@ -1,6 +1,7 @@ -from typing import Any, Optional, Callable, Tuple from datetime import datetime -from szurubooru import db, model, errors +from typing import Any, Callable, Optional, Tuple + +from szurubooru import db, errors, model class InvalidFavoriteTargetError(errors.ValidationError): @@ -8,10 +9,11 @@ class InvalidFavoriteTargetError(errors.ValidationError): def _get_table_info( - entity: model.Base) -> Tuple[model.Base, Callable[[model.Base], Any]]: + entity: model.Base, +) -> Tuple[model.Base, Callable[[model.Base], Any]]: assert entity resource_type, _, _ = model.util.get_resource_info(entity) - if resource_type == 'post': + if resource_type == "post": return model.PostFavorite, lambda table: table.post_id raise InvalidFavoriteTargetError() @@ -38,6 +40,7 @@ def unset_favorite(entity: model.Base, user: Optional[model.User]) -> None: def set_favorite(entity: model.Base, user: Optional[model.User]) -> None: from szurubooru.func import scores + assert entity assert user try: diff --git a/server/szurubooru/func/file_uploads.py b/server/szurubooru/func/file_uploads.py index e7f93d8..800397d 100644 --- a/server/szurubooru/func/file_uploads.py +++ b/server/szurubooru/func/file_uploads.py @@ -1,25 +1,25 @@ -from typing import Optional from datetime import datetime, timedelta -from szurubooru.func import files, util +from typing import Optional +from szurubooru.func import files, util MAX_MINUTES = 60 def _get_path(checksum: str) -> str: - return 'temporary-uploads/%s.dat' % checksum + return "temporary-uploads/%s.dat" % checksum def purge_old_uploads() -> None: now = datetime.now() - for file in files.scan('temporary-uploads'): + for file in files.scan("temporary-uploads"): file_time = datetime.fromtimestamp(file.stat().st_ctime) if now - file_time > timedelta(minutes=MAX_MINUTES): - files.delete('temporary-uploads/%s' % file.name) + files.delete("temporary-uploads/%s" % file.name) def get(checksum: str) -> Optional[bytes]: - return files.get('temporary-uploads/%s.dat' % checksum) + return files.get("temporary-uploads/%s.dat" % checksum) def save(content: bytes) -> str: diff --git a/server/szurubooru/func/files.py b/server/szurubooru/func/files.py index fa9f36f..6a89826 100644 --- a/server/szurubooru/func/files.py +++ b/server/szurubooru/func/files.py @@ -1,10 +1,11 @@ -from typing import Any, Optional, List import os +from typing import Any, List, Optional + from szurubooru import config def _get_full_path(path: str) -> str: - return os.path.join(config.config['data_dir'], path) + return os.path.join(config.config["data_dir"], path) def delete(path: str) -> None: @@ -31,12 +32,12 @@ def get(path: str) -> Optional[bytes]: full_path = _get_full_path(path) if not os.path.exists(full_path): return None - with open(full_path, 'rb') as handle: + with open(full_path, "rb") as handle: return handle.read() def save(path: str, content: bytes) -> None: full_path = _get_full_path(path) os.makedirs(os.path.dirname(full_path), exist_ok=True) - with open(full_path, 'wb') as handle: + with open(full_path, "wb") as handle: handle.write(content) diff --git a/server/szurubooru/func/image_hash.py b/server/szurubooru/func/image_hash.py index b89b218..fc7d141 100644 --- a/server/szurubooru/func/image_hash.py +++ b/server/szurubooru/func/image_hash.py @@ -1,14 +1,14 @@ import logging -from io import BytesIO +import math from datetime import datetime -from typing import Any, Optional, Tuple, Set, List, Callable -import elasticsearch -import elasticsearch_dsl +from io import BytesIO +from typing import Any, Callable, List, Optional, Set, Tuple + import numpy as np from PIL import Image + from szurubooru import config, errors -# pylint: disable=invalid-name logger = logging.getLogger(__name__) # Math based on paper from H. Chi Wong, Marshall Bern and David Goldberg @@ -17,87 +17,96 @@ logger = logging.getLogger(__name__) LOWER_PERCENTILE = 5 UPPER_PERCENTILE = 95 -IDENTICAL_TOLERANCE = 2 / 255. +IDENTICAL_TOLERANCE = 2 / 255.0 DISTANCE_CUTOFF = 0.45 N_LEVELS = 2 N = 9 P = None SAMPLE_WORDS = 16 MAX_WORDS = 63 -ES_DOC_TYPE = 'image' -ES_MAX_RESULTS = 100 - -Window = Tuple[Tuple[float, float], Tuple[float, float]] -NpMatrix = Any +SIG_CHUNK_BITS = 32 +SIG_NUMS = 8 * N * N +SIG_BASE = 2 * N_LEVELS + 2 +SIG_CHUNK_WIDTH = int(SIG_CHUNK_BITS / math.log2(SIG_BASE)) +SIG_CHUNK_NUMS = SIG_NUMS / SIG_CHUNK_WIDTH +assert SIG_NUMS % SIG_CHUNK_WIDTH == 0 -def _get_session() -> elasticsearch.Elasticsearch: - return elasticsearch.Elasticsearch([{ - 'host': config.config['elasticsearch']['host'], - 'port': config.config['elasticsearch']['port'], - }]) +Window = Tuple[Tuple[float, float], Tuple[float, float]] +NpMatrix = np.ndarray def _preprocess_image(content: bytes) -> NpMatrix: - img = Image.open(BytesIO(content)) - return np.asarray(img.convert('L'), dtype=np.uint8) + try: + img = Image.open(BytesIO(content)) + return np.asarray(img.convert("L"), dtype=np.uint8) + except IOError: + raise errors.ProcessingError( + "Unable to generate a signature hash " "for this image." + ) def _crop_image( - image: NpMatrix, - lower_percentile: float, - upper_percentile: float) -> Window: + image: NpMatrix, lower_percentile: float, upper_percentile: float +) -> Window: rw = np.cumsum(np.sum(np.abs(np.diff(image, axis=1)), axis=1)) cw = np.cumsum(np.sum(np.abs(np.diff(image, axis=0)), axis=0)) upper_column_limit = np.searchsorted( - cw, np.percentile(cw, upper_percentile), side='left') + cw, np.percentile(cw, upper_percentile), side="left" + ) lower_column_limit = np.searchsorted( - cw, np.percentile(cw, lower_percentile), side='right') + cw, np.percentile(cw, lower_percentile), side="right" + ) upper_row_limit = np.searchsorted( - rw, np.percentile(rw, upper_percentile), side='left') + rw, np.percentile(rw, upper_percentile), side="left" + ) lower_row_limit = np.searchsorted( - rw, np.percentile(rw, lower_percentile), side='right') + rw, np.percentile(rw, lower_percentile), side="right" + ) if lower_row_limit > upper_row_limit: - lower_row_limit = int(lower_percentile / 100. * image.shape[0]) - upper_row_limit = int(upper_percentile / 100. * image.shape[0]) + lower_row_limit = int(lower_percentile / 100.0 * image.shape[0]) + upper_row_limit = int(upper_percentile / 100.0 * image.shape[0]) if lower_column_limit > upper_column_limit: - lower_column_limit = int(lower_percentile / 100. * image.shape[1]) - upper_column_limit = int(upper_percentile / 100. * image.shape[1]) + lower_column_limit = int(lower_percentile / 100.0 * image.shape[1]) + upper_column_limit = int(upper_percentile / 100.0 * image.shape[1]) return ( (lower_row_limit, upper_row_limit), - (lower_column_limit, upper_column_limit)) + (lower_column_limit, upper_column_limit), + ) def _normalize_and_threshold( - diff_array: NpMatrix, - identical_tolerance: float, - n_levels: int) -> None: + diff_array: NpMatrix, identical_tolerance: float, n_levels: int +) -> None: mask = np.abs(diff_array) < identical_tolerance - diff_array[mask] = 0. + diff_array[mask] = 0.0 if np.all(mask): return positive_cutoffs = np.percentile( - diff_array[diff_array > 0.], np.linspace(0, 100, n_levels + 1)) + diff_array[diff_array > 0.0], np.linspace(0, 100, n_levels + 1) + ) negative_cutoffs = np.percentile( - diff_array[diff_array < 0.], np.linspace(100, 0, n_levels + 1)) + diff_array[diff_array < 0.0], np.linspace(100, 0, n_levels + 1) + ) for level, interval in enumerate( - positive_cutoffs[i:i + 2] - for i in range(positive_cutoffs.shape[0] - 1)): + positive_cutoffs[i : i + 2] + for i in range(positive_cutoffs.shape[0] - 1) + ): diff_array[ - (diff_array >= interval[0]) & (diff_array <= interval[1])] = \ - level + 1 + (diff_array >= interval[0]) & (diff_array <= interval[1]) + ] = (level + 1) for level, interval in enumerate( - negative_cutoffs[i:i + 2] - for i in range(negative_cutoffs.shape[0] - 1)): + negative_cutoffs[i : i + 2] + for i in range(negative_cutoffs.shape[0] - 1) + ): diff_array[ - (diff_array <= interval[0]) & (diff_array >= interval[1])] = \ - -(level + 1) + (diff_array <= interval[0]) & (diff_array >= interval[1]) + ] = -(level + 1) def _compute_grid_points( - image: NpMatrix, - n: float, - window: Window = None) -> Tuple[NpMatrix, NpMatrix]: + image: NpMatrix, n: float, window: Window = None +) -> Tuple[NpMatrix, NpMatrix]: if window is None: window = ((0, image.shape[0]), (0, image.shape[1])) x_coords = np.linspace(window[0][0], window[0][1], n + 2, dtype=int)[1:-1] @@ -106,12 +115,10 @@ def _compute_grid_points( def _compute_mean_level( - image: NpMatrix, - x_coords: NpMatrix, - y_coords: NpMatrix, - p: Optional[float]) -> NpMatrix: + image: NpMatrix, x_coords: NpMatrix, y_coords: NpMatrix, p: Optional[float] +) -> NpMatrix: if p is None: - p = max([2.0, int(0.5 + min(image.shape) / 20.)]) + p = max([2.0, int(0.5 + min(image.shape) / 20.0)]) avg_grey = np.zeros((x_coords.shape[0], y_coords.shape[0])) for i, x in enumerate(x_coords): lower_x_lim = int(max([x - p / 2, 0])) @@ -120,7 +127,8 @@ def _compute_mean_level( lower_y_lim = int(max([y - p / 2, 0])) upper_y_lim = int(min([lower_y_lim + p, image.shape[1]])) avg_grey[i, j] = np.mean( - image[lower_x_lim:upper_x_lim, lower_y_lim:upper_y_lim]) + image[lower_x_lim:upper_x_lim, lower_y_lim:upper_y_lim] + ) return avg_grey @@ -130,95 +138,117 @@ def _compute_differentials(grey_level_matrix: NpMatrix) -> NpMatrix: ( np.diff(grey_level_matrix), ( - np.zeros(grey_level_matrix.shape[0]) - .reshape((grey_level_matrix.shape[0], 1)) - ) - ), axis=1) + np.zeros(grey_level_matrix.shape[0]).reshape( + (grey_level_matrix.shape[0], 1) + ) + ), + ), + axis=1, + ) down_neighbors = -np.concatenate( ( np.diff(grey_level_matrix, axis=0), ( - np.zeros(grey_level_matrix.shape[1]) - .reshape((1, grey_level_matrix.shape[1])) - ) - )) + np.zeros(grey_level_matrix.shape[1]).reshape( + (1, grey_level_matrix.shape[1]) + ) + ), + ) + ) left_neighbors = -np.concatenate( - (right_neighbors[:, -1:], right_neighbors[:, :-1]), axis=1) + (right_neighbors[:, -1:], right_neighbors[:, :-1]), axis=1 + ) up_neighbors = -np.concatenate((down_neighbors[-1:], down_neighbors[:-1])) diagonals = np.arange( - -grey_level_matrix.shape[0] + 1, grey_level_matrix.shape[0]) - upper_left_neighbors = sum([ - np.diagflat(np.insert(np.diff(np.diag(grey_level_matrix, i)), 0, 0), i) - for i in diagonals]) - upper_right_neighbors = sum([ - np.diagflat(np.insert(np.diff(np.diag(flipped, i)), 0, 0), i) - for i in diagonals]) + -grey_level_matrix.shape[0] + 1, grey_level_matrix.shape[0] + ) + upper_left_neighbors = sum( + [ + np.diagflat( + np.insert(np.diff(np.diag(grey_level_matrix, i)), 0, 0), i + ) + for i in diagonals + ] + ) + upper_right_neighbors = sum( + [ + np.diagflat(np.insert(np.diff(np.diag(flipped, i)), 0, 0), i) + for i in diagonals + ] + ) lower_right_neighbors = -np.pad( - upper_left_neighbors[1:, 1:], (0, 1), mode='constant') + upper_left_neighbors[1:, 1:], (0, 1), mode="constant" + ) lower_left_neighbors = -np.pad( - upper_right_neighbors[1:, 1:], (0, 1), mode='constant') - return np.dstack(np.array([ - upper_left_neighbors, - up_neighbors, - np.fliplr(upper_right_neighbors), - left_neighbors, - right_neighbors, - np.fliplr(lower_left_neighbors), - down_neighbors, - lower_right_neighbors])) + upper_right_neighbors[1:, 1:], (0, 1), mode="constant" + ) + return np.dstack( + np.array( + [ + upper_left_neighbors, + up_neighbors, + np.fliplr(upper_right_neighbors), + left_neighbors, + right_neighbors, + np.fliplr(lower_left_neighbors), + down_neighbors, + lower_right_neighbors, + ] + ) + ) -def _generate_signature(content: bytes) -> NpMatrix: - im_array = _preprocess_image(content) - image_limits = _crop_image( - im_array, - lower_percentile=LOWER_PERCENTILE, - upper_percentile=UPPER_PERCENTILE) - x_coords, y_coords = _compute_grid_points( - im_array, n=N, window=image_limits) - avg_grey = _compute_mean_level(im_array, x_coords, y_coords, p=P) - diff_matrix = _compute_differentials(avg_grey) - _normalize_and_threshold( - diff_matrix, - identical_tolerance=IDENTICAL_TOLERANCE, - n_levels=N_LEVELS) - return np.ravel(diff_matrix).astype('int8') +def _words_to_int(word_array: NpMatrix) -> List[int]: + width = word_array.shape[1] + coding_vector = 3 ** np.arange(width) + return np.dot(word_array + 1, coding_vector).astype(int).tolist() def _get_words(array: NpMatrix, k: int, n: int) -> NpMatrix: - word_positions = np.linspace( - 0, array.shape[0], n, endpoint=False).astype('int') + word_positions = np.linspace(0, array.shape[0], n, endpoint=False).astype( + "int" + ) assert k <= array.shape[0] assert word_positions.shape[0] <= array.shape[0] - words = np.zeros((n, k)).astype('int8') + words = np.zeros((n, k)).astype("int8") for i, pos in enumerate(word_positions): if pos + k <= array.shape[0]: - words[i] = array[pos:pos + k] + words[i] = array[pos : pos + k] else: temp = array[pos:].copy() - temp.resize(k) + temp.resize(k, refcheck=False) words[i] = temp - _max_contrast(words) - words = _words_to_int(words) + words[words > 0] = 1 + words[words < 0] = -1 return words -def _words_to_int(word_array: NpMatrix) -> NpMatrix: - width = word_array.shape[1] - coding_vector = 3**np.arange(width) - return np.dot(word_array + 1, coding_vector) +def generate_signature(content: bytes) -> NpMatrix: + im_array = _preprocess_image(content) + image_limits = _crop_image( + im_array, + lower_percentile=LOWER_PERCENTILE, + upper_percentile=UPPER_PERCENTILE, + ) + x_coords, y_coords = _compute_grid_points( + im_array, n=N, window=image_limits + ) + avg_grey = _compute_mean_level(im_array, x_coords, y_coords, p=P) + diff_matrix = _compute_differentials(avg_grey) + _normalize_and_threshold( + diff_matrix, identical_tolerance=IDENTICAL_TOLERANCE, n_levels=N_LEVELS + ) + return np.ravel(diff_matrix).astype("int8") -def _max_contrast(array: NpMatrix) -> None: - array[array > 0] = 1 - array[array < 0] = -1 +def generate_words(signature: NpMatrix) -> List[int]: + return _words_to_int(_get_words(signature, k=SAMPLE_WORDS, n=MAX_WORDS)) -def _normalized_distance( - target_array: NpMatrix, - vec: NpMatrix, - nan_value: float = 1.0) -> List[float]: - target_array = target_array.astype(int) +def normalized_distance( + target_array: Any, vec: NpMatrix, nan_value: float = 1.0 +) -> List[float]: + target_array = np.array(target_array).astype(int) vec = vec.astype(int) topvec = np.linalg.norm(vec - target_array, axis=1) norm1 = np.linalg.norm(vec, axis=0) @@ -228,124 +258,50 @@ def _normalized_distance( return finvec -def _safety_blanket(default_param_factory: Callable[[], Any]) -> Callable: - def wrapper_outer(target_function: Callable) -> Callable: - def wrapper_inner(*args: Any, **kwargs: Any) -> Any: - try: - return target_function(*args, **kwargs) - except elasticsearch.exceptions.NotFoundError: - # index not yet created, will be created dynamically by - # add_image() - return default_param_factory() - except elasticsearch.exceptions.ElasticsearchException as ex: - logger.warning('Problem with elastic search: %s', ex) - raise errors.ThirdPartyError( - 'Error connecting to elastic search.') - except IOError: - raise errors.ProcessingError('Not an image.') - except Exception as ex: - raise errors.ThirdPartyError('Unknown error (%s).' % ex) - return wrapper_inner - return wrapper_outer - - -class Lookalike: - def __init__(self, score: int, distance: float, path: Any) -> None: - self.score = score - self.distance = distance - self.path = path - +def pack_signature(signature: NpMatrix) -> bytes: + """ + Serializes the signature vector for efficient storage in a database. -@_safety_blanket(lambda: None) -def add_image(path: str, image_content: bytes) -> None: - assert path - assert image_content - signature = _generate_signature(image_content) - words = _get_words(signature, k=SAMPLE_WORDS, n=MAX_WORDS) + Shifts the range of the signature vector from [-N_LEVELS,+N_LEVELS] + to [0, base] - record = { - 'signature': signature.tolist(), - 'path': path, - 'timestamp': datetime.now(), - } - for i in range(MAX_WORDS): - record['simple_word_' + str(i)] = words[i].tolist() - - _get_session().index( - index=config.config['elasticsearch']['index'], - doc_type=ES_DOC_TYPE, - body=record, - refresh=True) - - -@_safety_blanket(lambda: None) -def delete_image(path: str) -> None: - assert path - _get_session().delete_by_query( - index=config.config['elasticsearch']['index'], - doc_type=ES_DOC_TYPE, - body={'query': {'term': {'path': path}}}) - - -@_safety_blanket(lambda: []) -def search_by_image(image_content: bytes) -> List[Lookalike]: - signature = _generate_signature(image_content) - words = _get_words(signature, k=SAMPLE_WORDS, n=MAX_WORDS) - - res = _get_session().search( - index=config.config['elasticsearch']['index'], - doc_type=ES_DOC_TYPE, - body={ - 'query': - { - 'bool': - { - 'should': - [ - {'term': {'simple_word_%d' % i: word.tolist()}} - for i, word in enumerate(words) - ] - } - }, - '_source': {'excludes': ['simple_word_*']}}, - size=ES_MAX_RESULTS, - timeout='10s')['hits']['hits'] - - if len(res) == 0: - return [] - - sigs = np.array([x['_source']['signature'] for x in res]) - dists = _normalized_distance(sigs, np.array(signature)) - - ids = set() # type: Set[int] - ret = [] - for item, dist in zip(res, dists): - id = item['_id'] - score = item['_score'] - path = item['_source']['path'] - if id in ids: - continue - ids.add(id) - if dist < DISTANCE_CUTOFF: - ret.append(Lookalike(score=score, distance=dist, path=path)) - return ret + The vector can then be broken up into chunks, with each chunk + consisting of SIG_CHUNK_WIDTH digits of radix `base`. + This is then converted into a more packed array consisting of + uint32 elements (for SIG_CHUNK_BITS = 32). + """ + coding_vector = np.flipud(SIG_BASE ** np.arange(SIG_CHUNK_WIDTH)) + return ( + np.array( + [ + np.dot(x, coding_vector) + for x in np.reshape( + signature + N_LEVELS, (-1, SIG_CHUNK_WIDTH) + ) + ] + ) + .astype(f"uint{SIG_CHUNK_BITS}") + .tobytes() + ) -@_safety_blanket(lambda: None) -def purge() -> None: - _get_session().delete_by_query( - index=config.config['elasticsearch']['index'], - doc_type=ES_DOC_TYPE, - body={'query': {'match_all': {}}}, - refresh=True) +def unpack_signature(packed: bytes) -> NpMatrix: + """ + Deserializes the signature vector once recieved from the database. -@_safety_blanket(lambda: set()) -def get_all_paths() -> Set[str]: - search = ( - elasticsearch_dsl.Search( - using=_get_session(), - index=config.config['elasticsearch']['index'], - doc_type=ES_DOC_TYPE) - .source(['path'])) - return set(h.path for h in search.scan()) + Functions as an inverse transformation of pack_signature() + """ + return np.ravel( + np.array( + [ + [ + int(digit) - N_LEVELS + for digit in np.base_repr(e, base=SIG_BASE).zfill( + SIG_CHUNK_WIDTH + ) + ] + for e in np.frombuffer(packed, dtype=f"uint{SIG_CHUNK_BITS}") + ] + ).astype("int8") + ) diff --git a/server/szurubooru/func/images.py b/server/szurubooru/func/images.py index 6e88ab4..6413ac8 100644 --- a/server/szurubooru/func/images.py +++ b/server/szurubooru/func/images.py @@ -1,14 +1,14 @@ -from typing import List -import logging import json -import shlex -import subprocess +import logging import math import re +import shlex +import subprocess +from typing import List + from szurubooru import errors from szurubooru.func import mime, util - logger = logging.getLogger(__name__) @@ -19,97 +19,142 @@ class Image: @property def width(self) -> int: - return self.info['streams'][0]['width'] + return self.info["streams"][0]["width"] @property def height(self) -> int: - return self.info['streams'][0]['height'] + return self.info["streams"][0]["height"] @property def frames(self) -> int: - return self.info['streams'][0]['nb_read_frames'] + return self.info["streams"][0]["nb_read_frames"] def resize_fill(self, width: int, height: int) -> None: width_greater = self.width > self.height width, height = (-1, height) if width_greater else (width, -1) cli = [ - '-i', '{path}', - '-f', 'image2', - '-filter:v', "scale='{width}:{height}'".format( - width=width, height=height), - '-map', '0:v:0', - '-vframes', '1', - '-vcodec', 'png', - '-', + "-i", + "{path}", + "-f", + "image2", + "-filter:v", + "scale='{width}:{height}'".format(width=width, height=height), + "-map", + "0:v:0", + "-vframes", + "1", + "-vcodec", + "png", + "-", ] - if 'duration' in self.info['format'] \ - and self.info['format']['format_name'] != 'swf': - duration = float(self.info['format']['duration']) + if ( + "duration" in self.info["format"] + and self.info["format"]["format_name"] != "swf" + ): + duration = float(self.info["format"]["duration"]) if duration > 3: cli = [ - '-ss', - '%d' % math.floor(duration * 0.3), + "-ss", + "%d" % math.floor(duration * 0.3), ] + cli content = self._execute(cli, ignore_error_if_data=True) if not content: - raise errors.ProcessingError('Error while resizing image.') + raise errors.ProcessingError("Error while resizing image.") self.content = content self._reload_info() def to_png(self) -> bytes: - return self._execute([ - '-i', '{path}', - '-f', 'image2', - '-map', '0:v:0', - '-vframes', '1', - '-vcodec', 'png', - '-', - ]) + return self._execute( + [ + "-i", + "{path}", + "-f", + "image2", + "-map", + "0:v:0", + "-vframes", + "1", + "-vcodec", + "png", + "-", + ] + ) def to_jpeg(self) -> bytes: - return self._execute([ - '-f', 'lavfi', - '-i', 'color=white:s=%dx%d' % (self.width, self.height), - '-i', '{path}', - '-f', 'image2', - '-filter_complex', 'overlay', - '-map', '0:v:0', - '-vframes', '1', - '-vcodec', 'mjpeg', - '-', - ]) + return self._execute( + [ + "-f", + "lavfi", + "-i", + "color=white:s=%dx%d" % (self.width, self.height), + "-i", + "{path}", + "-f", + "image2", + "-filter_complex", + "overlay", + "-map", + "0:v:0", + "-vframes", + "1", + "-vcodec", + "mjpeg", + "-", + ] + ) def to_webm(self) -> bytes: - with util.create_temp_file_path(suffix='.log') as phase_log_path: + with util.create_temp_file_path(suffix=".log") as phase_log_path: # Pass 1 - self._execute([ - '-i', '{path}', - '-pass', '1', - '-passlogfile', phase_log_path, - '-vcodec', 'libvpx-vp9', - '-crf', '4', - '-b:v', '2500K', - '-acodec', 'libvorbis', - '-f', 'webm', - '-y', '/dev/null' - ]) + self._execute( + [ + "-i", + "{path}", + "-pass", + "1", + "-passlogfile", + phase_log_path, + "-vcodec", + "libvpx-vp9", + "-crf", + "4", + "-b:v", + "2500K", + "-acodec", + "libvorbis", + "-f", + "webm", + "-y", + "/dev/null", + ] + ) # Pass 2 - return self._execute([ - '-i', '{path}', - '-pass', '2', - '-passlogfile', phase_log_path, - '-vcodec', 'libvpx-vp9', - '-crf', '4', - '-b:v', '2500K', - '-acodec', 'libvorbis', - '-f', 'webm', - '-' - ]) + return self._execute( + [ + "-i", + "{path}", + "-pass", + "2", + "-passlogfile", + phase_log_path, + "-vcodec", + "libvpx-vp9", + "-crf", + "4", + "-b:v", + "2500K", + "-acodec", + "libvorbis", + "-f", + "webm", + "-", + ] + ) def to_mp4(self) -> bytes: - with util.create_temp_file_path(suffix='.dat') as mp4_temp_path: + with util.create_temp_file_path(suffix=".dat") as mp4_temp_path: width = self.width height = self.height altered_dimensions = False @@ -123,96 +168,138 @@ class Image: altered_dimensions = True args = [ - '-i', '{path}', - '-vcodec', 'libx264', - '-preset', 'slow', - '-crf', '22', - '-b:v', '200K', - '-profile:v', 'main', - '-pix_fmt', 'yuv420p', - '-acodec', 'aac', - '-f', 'mp4' + "-i", + "{path}", + "-vcodec", + "libx264", + "-preset", + "slow", + "-crf", + "22", + "-b:v", + "200K", + "-profile:v", + "main", + "-pix_fmt", + "yuv420p", + "-acodec", + "aac", + "-f", + "mp4", ] if altered_dimensions: - args += ['-filter:v', 'scale=\'%d:%d\'' % (width, height)] + args += ["-filter:v", "scale='%d:%d'" % (width, height)] - self._execute(args + ['-y', mp4_temp_path]) + self._execute(args + ["-y", mp4_temp_path]) - with open(mp4_temp_path, 'rb') as mp4_temp: + with open(mp4_temp_path, "rb") as mp4_temp: return mp4_temp.read() def check_for_sound(self) -> bool: - audioinfo = json.loads(self._execute([ - '-i', '{path}', - '-of', 'json', - '-select_streams', 'a', - '-show_streams', - ], program='ffprobe').decode('utf-8')) - assert 'streams' in audioinfo - if len(audioinfo['streams']) < 1: + audioinfo = json.loads( + self._execute( + [ + "-i", + "{path}", + "-of", + "json", + "-select_streams", + "a", + "-show_streams", + ], + program="ffprobe", + ).decode("utf-8") + ) + assert "streams" in audioinfo + if len(audioinfo["streams"]) < 1: return False - log = self._execute([ - '-hide_banner', - '-progress', '-', - '-i', '{path}', - '-af', 'volumedetect', - '-max_muxing_queue_size', '99999', - '-vn', '-sn', - '-f', 'null', - '-y', '/dev/null', - ], get_logs=True).decode('utf-8', errors='replace') - log_match = re.search(r'.*volumedetect.*mean_volume: (.*) dB', log) - assert log_match - assert log_match.groups() + log = self._execute( + [ + "-hide_banner", + "-progress", + "-", + "-i", + "{path}", + "-af", + "volumedetect", + "-max_muxing_queue_size", + "99999", + "-vn", + "-sn", + "-f", + "null", + "-y", + "/dev/null", + ], + get_logs=True, + ).decode("utf-8", errors="replace") + log_match = re.search(r".*volumedetect.*mean_volume: (.*) dB", log) + if not log_match or not log_match.groups(): + raise errors.ProcessingError( + "A problem occured when trying to check for audio" + ) meanvol = float(log_match.groups()[0]) # -91.0 dB is the minimum for 16-bit audio, assume sound if > -80.0 dB return meanvol > -80.0 def _execute( - self, - cli: List[str], - program: str = 'ffmpeg', - ignore_error_if_data: bool = False, - get_logs: bool = False) -> bytes: + self, + cli: List[str], + program: str = "ffmpeg", + ignore_error_if_data: bool = False, + get_logs: bool = False, + ) -> bytes: extension = mime.get_extension(mime.get_mime_type(self.content)) assert extension - with util.create_temp_file(suffix='.' + extension) as handle: + with util.create_temp_file(suffix="." + extension) as handle: handle.write(self.content) handle.flush() - cli = [program, '-loglevel', '32' if get_logs else '24'] + cli + cli = [program, "-loglevel", "32" if get_logs else "24"] + cli cli = [part.format(path=handle.name) for part in cli] proc = subprocess.Popen( cli, stdout=subprocess.PIPE, stdin=subprocess.PIPE, - stderr=subprocess.PIPE) + stderr=subprocess.PIPE, + ) out, err = proc.communicate(input=self.content) if proc.returncode != 0: logger.warning( - 'Failed to execute ffmpeg command (cli=%r, err=%r)', - ' '.join(shlex.quote(arg) for arg in cli), - err) - if ((len(out) > 0 and not ignore_error_if_data) - or len(out) == 0): + "Failed to execute ffmpeg command (cli=%r, err=%r)", + " ".join(shlex.quote(arg) for arg in cli), + err, + ) + if (len(out) > 0 and not ignore_error_if_data) or len( + out + ) == 0: raise errors.ProcessingError( - 'Error while processing image.\n' - + err.decode('utf-8')) + "Error while processing image.\n" + err.decode("utf-8") + ) return err if get_logs else out def _reload_info(self) -> None: - self.info = json.loads(self._execute([ - '-i', '{path}', - '-of', 'json', - '-select_streams', 'v', - '-show_format', - '-show_streams', - ], program='ffprobe').decode('utf-8')) - assert 'format' in self.info - assert 'streams' in self.info - if len(self.info['streams']) < 1: - logger.warning('The video contains no video streams.') + self.info = json.loads( + self._execute( + [ + "-i", + "{path}", + "-of", + "json", + "-select_streams", + "v", + "-show_format", + "-show_streams", + ], + program="ffprobe", + ).decode("utf-8") + ) + assert "format" in self.info + assert "streams" in self.info + if len(self.info["streams"]) < 1: + logger.warning("The video contains no video streams.") raise errors.ProcessingError( - 'The video contains no video streams.') + "The video contains no video streams." + ) diff --git a/server/szurubooru/func/mailer.py b/server/szurubooru/func/mailer.py index 76682f1..c4cf9db 100644 --- a/server/szurubooru/func/mailer.py +++ b/server/szurubooru/func/mailer.py @@ -1,16 +1,18 @@ -import smtplib import email.mime.text +import smtplib + from szurubooru import config def send_mail(sender: str, recipient: str, subject: str, body: str) -> None: msg = email.mime.text.MIMEText(body) - msg['Subject'] = subject - msg['From'] = sender - msg['To'] = recipient + msg["Subject"] = subject + msg["From"] = sender + msg["To"] = recipient smtp = smtplib.SMTP( - config.config['smtp']['host'], int(config.config['smtp']['port'])) - smtp.login(config.config['smtp']['user'], config.config['smtp']['pass']) + config.config["smtp"]["host"], int(config.config["smtp"]["port"]) + ) + smtp.login(config.config["smtp"]["user"], config.config["smtp"]["pass"]) smtp.send_message(msg) smtp.quit() diff --git a/server/szurubooru/func/mime.py b/server/szurubooru/func/mime.py index c83f744..5f6279b 100644 --- a/server/szurubooru/func/mime.py +++ b/server/szurubooru/func/mime.py @@ -4,55 +4,66 @@ from typing import Optional def get_mime_type(content: bytes) -> str: if not content: - return 'application/octet-stream' + return "application/octet-stream" - if content[0:3] in (b'CWS', b'FWS', b'ZWS'): - return 'application/x-shockwave-flash' + if content[0:3] in (b"CWS", b"FWS", b"ZWS"): + return "application/x-shockwave-flash" - if content[0:3] == b'\xFF\xD8\xFF': - return 'image/jpeg' + if content[0:3] == b"\xFF\xD8\xFF": + return "image/jpeg" - if content[0:6] == b'\x89PNG\x0D\x0A': - return 'image/png' + if content[0:6] == b"\x89PNG\x0D\x0A": + return "image/png" - if content[0:6] in (b'GIF87a', b'GIF89a'): - return 'image/gif' + if content[0:6] in (b"GIF87a", b"GIF89a"): + return "image/gif" - if content[0:4] == b'\x1A\x45\xDF\xA3': - return 'video/webm' + if content[8:12] == b"WEBP": + return "image/webp" - if content[4:12] in (b'ftypisom', b'ftypmp42'): - return 'video/mp4' + if content[0:4] == b"\x1A\x45\xDF\xA3": + return "video/webm" - return 'application/octet-stream' + if content[4:12] in (b"ftypisom", b"ftypiso5", b"ftypmp42"): + return "video/mp4" + + return "application/octet-stream" def get_extension(mime_type: str) -> Optional[str]: extension_map = { - 'application/x-shockwave-flash': 'swf', - 'image/gif': 'gif', - 'image/jpeg': 'jpg', - 'image/png': 'png', - 'video/mp4': 'mp4', - 'video/webm': 'webm', - 'application/octet-stream': 'dat', + "application/x-shockwave-flash": "swf", + "image/gif": "gif", + "image/jpeg": "jpg", + "image/png": "png", + "image/webp": "webp", + "video/mp4": "mp4", + "video/webm": "webm", + "application/octet-stream": "dat", } - return extension_map.get((mime_type or '').strip().lower(), None) + return extension_map.get((mime_type or "").strip().lower(), None) def is_flash(mime_type: str) -> bool: - return mime_type.lower() == 'application/x-shockwave-flash' + return mime_type.lower() == "application/x-shockwave-flash" def is_video(mime_type: str) -> bool: - return mime_type.lower() in ('application/ogg', 'video/mp4', 'video/webm') + return mime_type.lower() in ("application/ogg", "video/mp4", "video/webm") def is_image(mime_type: str) -> bool: - return mime_type.lower() in ('image/jpeg', 'image/png', 'image/gif') + return mime_type.lower() in ( + "image/jpeg", + "image/png", + "image/gif", + "image/webp", + ) def is_animated_gif(content: bytes) -> bool: - pattern = b'\x21\xF9\x04[\x00-\xFF]{4}\x00[\x2C\x21]' - return get_mime_type(content) == 'image/gif' \ + pattern = b"\x21\xF9\x04[\x00-\xFF]{4}\x00[\x2C\x21]" + return ( + get_mime_type(content) == "image/gif" and len(re.findall(pattern, content)) > 1 + ) diff --git a/server/szurubooru/func/net.py b/server/szurubooru/func/net.py index e6326c0..4e4c222 100644 --- a/server/szurubooru/func/net.py +++ b/server/szurubooru/func/net.py @@ -1,16 +1,94 @@ +import json +import logging +import os +import urllib.error import urllib.request -from szurubooru import config -from szurubooru import errors +from tempfile import NamedTemporaryFile +from threading import Thread +from typing import Any, Dict, List +from youtube_dl import YoutubeDL +from youtube_dl.utils import YoutubeDLError -def download(url: str) -> bytes: +from szurubooru import config, errors +from szurubooru.func import mime, util + +logger = logging.getLogger(__name__) + + +def download(url: str, use_video_downloader: bool = False) -> bytes: assert url request = urllib.request.Request(url) - if config.config['user_agent']: - request.add_header('User-Agent', config.config['user_agent']) - request.add_header('Referer', url) + if config.config["user_agent"]: + request.add_header("User-Agent", config.config["user_agent"]) + request.add_header("Referer", url) try: with urllib.request.urlopen(request) as handle: - return handle.read() + content = handle.read() except Exception as ex: - raise errors.ProcessingError('Error downloading %s (%s)' % (url, ex)) + raise errors.ProcessingError("Error downloading %s (%s)" % (url, ex)) + if ( + use_video_downloader + and mime.get_mime_type(content) == "application/octet-stream" + ): + return _youtube_dl_wrapper(url) + return content + + +def _youtube_dl_wrapper(url: str) -> bytes: + outpath = os.path.join( + config.config["data_dir"], + "temporary-uploads", + "youtubedl-" + util.get_sha1(url)[0:8] + ".dat", + ) + options = { + "ignoreerrors": False, + "format": "best[ext=webm]/best[ext=mp4]/best[ext=flv]", + "logger": logger, + "max_filesize": config.config["max_dl_filesize"], + "max_downloads": 1, + "outtmpl": outpath, + } + try: + with YoutubeDL(options) as ydl: + ydl.extract_info(url, download=True) + with open(outpath, "rb") as f: + return f.read() + except YoutubeDLError as ex: + raise errors.ThirdPartyError( + "Error downloading video %s (%s)" % (url, ex) + ) + except FileNotFoundError: + raise errors.ThirdPartyError( + "Error downloading video %s (file could not be saved)" % (url) + ) + + +def post_to_webhooks(payload: Dict[str, Any]) -> List[Thread]: + threads = [ + Thread(target=_post_to_webhook, args=(webhook, payload)) + for webhook in (config.config["webhooks"] or []) + ] + for thread in threads: + thread.daemon = False + thread.start() + return threads + + +def _post_to_webhook(webhook: str, payload: Dict[str, Any]) -> None: + req = urllib.request.Request(webhook) + req.data = json.dumps( + payload, + default=lambda x: x.isoformat("T") + "Z", + ).encode("utf-8") + req.add_header("Content-Type", "application/json") + try: + res = urllib.request.urlopen(req) + if not 200 <= res.status <= 299: + logger.warning( + f"Webhook {webhook} returned {res.status} {res.reason}" + ) + return res.status + except urllib.error.URLError as e: + logger.warning(f"Unable to call webhook {webhook}: {str(e)}") + return 400 diff --git a/server/szurubooru/func/pool_categories.py b/server/szurubooru/func/pool_categories.py new file mode 100644 index 0000000..c13857c --- /dev/null +++ b/server/szurubooru/func/pool_categories.py @@ -0,0 +1,207 @@ +import re +from typing import Any, Callable, Dict, List, Optional + +import sqlalchemy as sa + +from szurubooru import config, db, errors, model, rest +from szurubooru.func import cache, serialization, util + +DEFAULT_CATEGORY_NAME_CACHE_KEY = "default-pool-category" + + +class PoolCategoryNotFoundError(errors.NotFoundError): + pass + + +class PoolCategoryAlreadyExistsError(errors.ValidationError): + pass + + +class PoolCategoryIsInUseError(errors.ValidationError): + pass + + +class InvalidPoolCategoryNameError(errors.ValidationError): + pass + + +class InvalidPoolCategoryColorError(errors.ValidationError): + pass + + +def _verify_name_validity(name: str) -> None: + name_regex = config.config["pool_category_name_regex"] + if not re.match(name_regex, name): + raise InvalidPoolCategoryNameError( + "Name must satisfy regex %r." % name_regex + ) + + +class PoolCategorySerializer(serialization.BaseSerializer): + def __init__(self, category: model.PoolCategory) -> None: + self.category = category + + def _serializers(self) -> Dict[str, Callable[[], Any]]: + return { + "name": self.serialize_name, + "version": self.serialize_version, + "color": self.serialize_color, + "usages": self.serialize_usages, + "default": self.serialize_default, + } + + def serialize_name(self) -> Any: + return self.category.name + + def serialize_version(self) -> Any: + return self.category.version + + def serialize_color(self) -> Any: + return self.category.color + + def serialize_usages(self) -> Any: + return self.category.pool_count + + def serialize_default(self) -> Any: + return self.category.default + + +def serialize_category( + category: Optional[model.PoolCategory], options: List[str] = [] +) -> Optional[rest.Response]: + if not category: + return None + return PoolCategorySerializer(category).serialize(options) + + +def create_category(name: str, color: str) -> model.PoolCategory: + category = model.PoolCategory() + update_category_name(category, name) + update_category_color(category, color) + if not get_all_categories(): + category.default = True + return category + + +def update_category_name(category: model.PoolCategory, name: str) -> None: + assert category + if not name: + raise InvalidPoolCategoryNameError("Name cannot be empty.") + expr = sa.func.lower(model.PoolCategory.name) == name.lower() + if category.pool_category_id: + expr = expr & ( + model.PoolCategory.pool_category_id != category.pool_category_id + ) + already_exists = ( + db.session.query(model.PoolCategory).filter(expr).count() > 0 + ) + if already_exists: + raise PoolCategoryAlreadyExistsError( + "A category with this name already exists." + ) + if util.value_exceeds_column_size(name, model.PoolCategory.name): + raise InvalidPoolCategoryNameError("Name is too long.") + _verify_name_validity(name) + category.name = name + cache.remove(DEFAULT_CATEGORY_NAME_CACHE_KEY) + + +def update_category_color(category: model.PoolCategory, color: str) -> None: + assert category + if not color: + raise InvalidPoolCategoryColorError("Color cannot be empty.") + if not re.match(r"^#?[0-9a-z]+$", color): + raise InvalidPoolCategoryColorError("Invalid color.") + if util.value_exceeds_column_size(color, model.PoolCategory.color): + raise InvalidPoolCategoryColorError("Color is too long.") + category.color = color + + +def try_get_category_by_name( + name: str, lock: bool = False +) -> Optional[model.PoolCategory]: + query = db.session.query(model.PoolCategory).filter( + sa.func.lower(model.PoolCategory.name) == name.lower() + ) + if lock: + query = query.with_for_update() + return query.one_or_none() + + +def get_category_by_name(name: str, lock: bool = False) -> model.PoolCategory: + category = try_get_category_by_name(name, lock) + if not category: + raise PoolCategoryNotFoundError("Pool category %r not found." % name) + return category + + +def get_all_category_names() -> List[str]: + return [cat.name for cat in get_all_categories()] + + +def get_all_categories() -> List[model.PoolCategory]: + return ( + db.session.query(model.PoolCategory) + .order_by(model.PoolCategory.name.asc()) + .all() + ) + + +def try_get_default_category( + lock: bool = False, +) -> Optional[model.PoolCategory]: + query = db.session.query(model.PoolCategory).filter( + model.PoolCategory.default + ) + if lock: + query = query.with_for_update() + category = query.first() + # if for some reason (e.g. as a result of migration) there's no default + # category, get the first record available. + if not category: + query = db.session.query(model.PoolCategory).order_by( + model.PoolCategory.pool_category_id.asc() + ) + if lock: + query = query.with_for_update() + category = query.first() + return category + + +def get_default_category(lock: bool = False) -> model.PoolCategory: + category = try_get_default_category(lock) + if not category: + raise PoolCategoryNotFoundError("No pool category created yet.") + return category + + +def get_default_category_name() -> str: + if cache.has(DEFAULT_CATEGORY_NAME_CACHE_KEY): + return cache.get(DEFAULT_CATEGORY_NAME_CACHE_KEY) + default_category = get_default_category() + default_category_name = default_category.name + cache.put(DEFAULT_CATEGORY_NAME_CACHE_KEY, default_category_name) + return default_category_name + + +def set_default_category(category: model.PoolCategory) -> None: + assert category + old_category = try_get_default_category(lock=True) + if old_category: + db.session.refresh(old_category) + old_category.default = False + db.session.refresh(category) + category.default = True + cache.remove(DEFAULT_CATEGORY_NAME_CACHE_KEY) + + +def delete_category(category: model.PoolCategory) -> None: + assert category + if len(get_all_category_names()) == 1: + raise PoolCategoryIsInUseError("Cannot delete the last category.") + if (category.pool_count or 0) > 0: + raise PoolCategoryIsInUseError( + "Pool category has some usages and cannot be deleted. " + + "Please remove this category from relevant pools first." + ) + db.session.delete(category) diff --git a/server/szurubooru/func/pools.py b/server/szurubooru/func/pools.py new file mode 100644 index 0000000..c3ea9f0 --- /dev/null +++ b/server/szurubooru/func/pools.py @@ -0,0 +1,337 @@ +import re +from datetime import datetime +from typing import Any, Callable, Dict, List, Optional, Tuple + +import sqlalchemy as sa + +from szurubooru import config, db, errors, model, rest +from szurubooru.func import pool_categories, posts, serialization, util + + +class PoolNotFoundError(errors.NotFoundError): + pass + + +class PoolAlreadyExistsError(errors.ValidationError): + pass + + +class PoolIsInUseError(errors.ValidationError): + pass + + +class InvalidPoolNameError(errors.ValidationError): + pass + + +class InvalidPoolDuplicateError(errors.ValidationError): + pass + + +class InvalidPoolCategoryError(errors.ValidationError): + pass + + +class InvalidPoolDescriptionError(errors.ValidationError): + pass + + +class InvalidPoolRelationError(errors.ValidationError): + pass + + +class InvalidPoolNonexistentPostError(errors.ValidationError): + pass + + +def _verify_name_validity(name: str) -> None: + if util.value_exceeds_column_size(name, model.PoolName.name): + raise InvalidPoolNameError("Name is too long.") + name_regex = config.config["pool_name_regex"] + if not re.match(name_regex, name): + raise InvalidPoolNameError("Name must satisfy regex %r." % name_regex) + + +def _get_names(pool: model.Pool) -> List[str]: + assert pool + return [pool_name.name for pool_name in pool.names] + + +def _lower_list(names: List[str]) -> List[str]: + return [name.lower() for name in names] + + +def _check_name_intersection( + names1: List[str], names2: List[str], case_sensitive: bool +) -> bool: + if not case_sensitive: + names1 = _lower_list(names1) + names2 = _lower_list(names2) + return len(set(names1).intersection(names2)) > 0 + + +def _duplicates(a: List[int]) -> List[int]: + seen = set() + dupes = [] + for x in a: + if x not in seen: + seen.add(x) + else: + dupes.append(x) + return dupes + + +def sort_pools(pools: List[model.Pool]) -> List[model.Pool]: + default_category_name = pool_categories.get_default_category_name() + return sorted( + pools, + key=lambda pool: ( + default_category_name == pool.category.name, + pool.category.name, + pool.names[0].name, + ), + ) + + +class PoolSerializer(serialization.BaseSerializer): + def __init__(self, pool: model.Pool) -> None: + self.pool = pool + + def _serializers(self) -> Dict[str, Callable[[], Any]]: + return { + "id": self.serialize_id, + "names": self.serialize_names, + "category": self.serialize_category, + "version": self.serialize_version, + "description": self.serialize_description, + "creationTime": self.serialize_creation_time, + "lastEditTime": self.serialize_last_edit_time, + "postCount": self.serialize_post_count, + "posts": self.serialize_posts, + } + + def serialize_id(self) -> Any: + return self.pool.pool_id + + def serialize_names(self) -> Any: + return [pool_name.name for pool_name in self.pool.names] + + def serialize_category(self) -> Any: + return self.pool.category.name + + def serialize_version(self) -> Any: + return self.pool.version + + def serialize_description(self) -> Any: + return self.pool.description + + def serialize_creation_time(self) -> Any: + return self.pool.creation_time + + def serialize_last_edit_time(self) -> Any: + return self.pool.last_edit_time + + def serialize_post_count(self) -> Any: + return self.pool.post_count + + def serialize_posts(self) -> Any: + return [ + post + for post in [ + posts.serialize_micro_post(rel, None) + for rel in self.pool.posts + ] + ] + + +def serialize_pool( + pool: model.Pool, options: List[str] = [] +) -> Optional[rest.Response]: + if not pool: + return None + return PoolSerializer(pool).serialize(options) + + +def serialize_micro_pool(pool: model.Pool) -> Optional[rest.Response]: + return serialize_pool( + pool, options=["id", "names", "category", "description", "postCount"] + ) + + +def try_get_pool_by_id(pool_id: int) -> Optional[model.Pool]: + return ( + db.session.query(model.Pool) + .filter(model.Pool.pool_id == pool_id) + .one_or_none() + ) + + +def get_pool_by_id(pool_id: int) -> model.Pool: + pool = try_get_pool_by_id(pool_id) + if not pool: + raise PoolNotFoundError("Pool %r not found." % pool_id) + return pool + + +def try_get_pool_by_name(name: str) -> Optional[model.Pool]: + return ( + db.session.query(model.Pool) + .join(model.PoolName) + .filter(sa.func.lower(model.PoolName.name) == name.lower()) + .one_or_none() + ) + + +def get_pool_by_name(name: str) -> model.Pool: + pool = try_get_pool_by_name(name) + if not pool: + raise PoolNotFoundError("Pool %r not found." % name) + return pool + + +def get_pools_by_names(names: List[str]) -> List[model.Pool]: + names = util.icase_unique(names) + if len(names) == 0: + return [] + return ( + db.session.query(model.Pool) + .join(model.PoolName) + .filter( + sa.sql.or_( + sa.func.lower(model.PoolName.name) == name.lower() + for name in names + ) + ) + .all() + ) + + +def get_or_create_pools_by_names( + names: List[str], +) -> Tuple[List[model.Pool], List[model.Pool]]: + names = util.icase_unique(names) + existing_pools = get_pools_by_names(names) + new_pools = [] + pool_category_name = pool_categories.get_default_category_name() + for name in names: + found = False + for existing_pool in existing_pools: + if _check_name_intersection( + _get_names(existing_pool), [name], False + ): + found = True + break + if not found: + new_pool = create_pool( + names=[name], category_name=pool_category_name, post_ids=[] + ) + db.session.add(new_pool) + new_pools.append(new_pool) + return existing_pools, new_pools + + +def delete(source_pool: model.Pool) -> None: + assert source_pool + db.session.delete(source_pool) + + +def merge_pools(source_pool: model.Pool, target_pool: model.Pool) -> None: + assert source_pool + assert target_pool + if source_pool.pool_id == target_pool.pool_id: + raise InvalidPoolRelationError("Cannot merge pool with itself.") + + def merge_pool_posts(source_pool_id: int, target_pool_id: int) -> None: + alias1 = model.PoolPost + alias2 = sa.orm.util.aliased(model.PoolPost) + update_stmt = sa.sql.expression.update(alias1).where( + alias1.pool_id == source_pool_id + ) + update_stmt = update_stmt.where( + ~sa.exists() + .where(alias1.post_id == alias2.post_id) + .where(alias2.pool_id == target_pool_id) + ) + update_stmt = update_stmt.values(pool_id=target_pool_id) + db.session.execute(update_stmt) + + merge_pool_posts(source_pool.pool_id, target_pool.pool_id) + delete(source_pool) + + +def create_pool( + names: List[str], category_name: str, post_ids: List[int] +) -> model.Pool: + pool = model.Pool() + pool.creation_time = datetime.utcnow() + update_pool_names(pool, names) + update_pool_category_name(pool, category_name) + update_pool_posts(pool, post_ids) + return pool + + +def update_pool_category_name(pool: model.Pool, category_name: str) -> None: + assert pool + pool.category = pool_categories.get_category_by_name(category_name) + + +def update_pool_names(pool: model.Pool, names: List[str]) -> None: + # sanitize + assert pool + names = util.icase_unique([name for name in names if name]) + if not len(names): + raise InvalidPoolNameError("At least one name must be specified.") + for name in names: + _verify_name_validity(name) + + # check for existing pools + expr = sa.sql.false() + for name in names: + expr = expr | (sa.func.lower(model.PoolName.name) == name.lower()) + if pool.pool_id: + expr = expr & (model.PoolName.pool_id != pool.pool_id) + existing_pools = db.session.query(model.PoolName).filter(expr).all() + if len(existing_pools): + raise PoolAlreadyExistsError( + "One of names is already used by another pool." + ) + + # remove unwanted items + for pool_name in pool.names[:]: + if not _check_name_intersection([pool_name.name], names, True): + pool.names.remove(pool_name) + # add wanted items + for name in names: + if not _check_name_intersection(_get_names(pool), [name], True): + pool.names.append(model.PoolName(name, -1)) + + # set alias order to match the request + for i, name in enumerate(names): + for pool_name in pool.names: + if pool_name.name.lower() == name.lower(): + pool_name.order = i + + +def update_pool_description(pool: model.Pool, description: str) -> None: + assert pool + if util.value_exceeds_column_size(description, model.Pool.description): + raise InvalidPoolDescriptionError("Description is too long.") + pool.description = description or None + + +def update_pool_posts(pool: model.Pool, post_ids: List[int]) -> None: + assert pool + dupes = _duplicates(post_ids) + if len(dupes) > 0: + dupes = ", ".join(list(str(x) for x in dupes)) + raise InvalidPoolDuplicateError("Duplicate post(s) in pool: " + dupes) + ret = posts.get_posts_by_ids(post_ids) + if len(post_ids) != len(ret): + missing = set(post_ids) - set(post.post_id for post in ret) + missing = ", ".join(list(str(x) for x in missing)) + raise InvalidPoolNonexistentPostError( + "The following posts do not exist: " + missing + ) + pool.posts.clear() + for post in ret: + pool.posts.append(post) diff --git a/server/szurubooru/func/posts.py b/server/szurubooru/func/posts.py index b6816b0..ee7c31a 100644 --- a/server/szurubooru/func/posts.py +++ b/server/szurubooru/func/posts.py @@ -1,17 +1,35 @@ import hmac -from typing import Any, Optional, Tuple, List, Dict, Callable +import logging from datetime import datetime +from typing import Any, Callable, Dict, List, Optional, Tuple + import sqlalchemy as sa -from szurubooru import config, db, model, errors, rest + +from szurubooru import config, db, errors, model, rest from szurubooru.func import ( - users, scores, comments, tags, metrics, util, - mime, images, files, image_hash, serialization, snapshots) + comments, + files, + image_hash, + images, + metrics, + mime, + pools, + scores, + serialization, + snapshots, + tags, + users, + util, +) + +logger = logging.getLogger(__name__) EMPTY_PIXEL = ( - b'\x47\x49\x46\x38\x39\x61\x01\x00\x01\x00\x80\x01\x00\x00\x00\x00' - b'\xff\xff\xff\x21\xf9\x04\x01\x00\x00\x01\x00\x2c\x00\x00\x00\x00' - b'\x01\x00\x01\x00\x00\x02\x02\x4c\x01\x00\x3b') + b"\x47\x49\x46\x38\x39\x61\x01\x00\x01\x00\x80\x01\x00\x00\x00\x00" + b"\xff\xff\xff\x21\xf9\x04\x01\x00\x00\x01\x00\x2c\x00\x00\x00\x00" + b"\x01\x00\x01\x00\x00\x02\x02\x4c\x01\x00\x3b" +) class PostNotFoundError(errors.NotFoundError): @@ -25,11 +43,12 @@ class PostAlreadyFeaturedError(errors.ValidationError): class PostAlreadyUploadedError(errors.ValidationError): def __init__(self, other_post: model.Post) -> None: super().__init__( - 'Post already uploaded (%d)' % other_post.post_id, + "Post already uploaded (%d)" % other_post.post_id, { - 'otherPostUrl': get_post_content_url(other_post), - 'otherPostId': other_post.post_id, - }) + "otherPostUrl": get_post_content_url(other_post), + "otherPostId": other_post.post_id, + }, + ) class InvalidPostIdError(errors.ValidationError): @@ -60,81 +79,83 @@ class InvalidPostFlagError(errors.ValidationError): pass -class PostLookalike(image_hash.Lookalike): - def __init__(self, score: int, distance: float, post: model.Post) -> None: - super().__init__(score, distance, post.post_id) - self.post = post - - SAFETY_MAP = { - model.Post.SAFETY_SAFE: 'safe', - model.Post.SAFETY_SKETCHY: 'sketchy', - model.Post.SAFETY_UNSAFE: 'unsafe', + model.Post.SAFETY_SAFE: "safe", + model.Post.SAFETY_SKETCHY: "sketchy", + model.Post.SAFETY_UNSAFE: "unsafe", } TYPE_MAP = { - model.Post.TYPE_IMAGE: 'image', - model.Post.TYPE_ANIMATION: 'animation', - model.Post.TYPE_VIDEO: 'video', - model.Post.TYPE_FLASH: 'flash', + model.Post.TYPE_IMAGE: "image", + model.Post.TYPE_ANIMATION: "animation", + model.Post.TYPE_VIDEO: "video", + model.Post.TYPE_FLASH: "flash", } FLAG_MAP = { - model.Post.FLAG_LOOP: 'loop', - model.Post.FLAG_SOUND: 'sound', + model.Post.FLAG_LOOP: "loop", + model.Post.FLAG_SOUND: "sound", } def get_post_security_hash(id: int) -> str: return hmac.new( - config.config['secret'].encode('utf8'), - str(id).encode('utf-8')).hexdigest()[0:16] + config.config["secret"].encode("utf8"), + msg=str(id).encode("utf-8"), + digestmod="md5", + ).hexdigest()[0:16] def get_post_content_url(post: model.Post) -> str: assert post - return '%s/posts/%d_%s.%s' % ( - config.config['data_url'].rstrip('/'), + return "%s/posts/%d_%s.%s" % ( + config.config["data_url"].rstrip("/"), post.post_id, get_post_security_hash(post.post_id), - mime.get_extension(post.mime_type) or 'dat') + mime.get_extension(post.mime_type) or "dat", + ) def get_post_thumbnail_url(post: model.Post) -> str: assert post - return '%s/generated-thumbnails/%d_%s.jpg' % ( - config.config['data_url'].rstrip('/'), + return "%s/generated-thumbnails/%d_%s.jpg" % ( + config.config["data_url"].rstrip("/"), post.post_id, - get_post_security_hash(post.post_id)) + get_post_security_hash(post.post_id), + ) def get_post_content_path(post: model.Post) -> str: assert post assert post.post_id - return 'posts/%d_%s.%s' % ( + return "posts/%d_%s.%s" % ( post.post_id, get_post_security_hash(post.post_id), - mime.get_extension(post.mime_type) or 'dat') + mime.get_extension(post.mime_type) or "dat", + ) def get_post_thumbnail_path(post: model.Post) -> str: assert post - return 'generated-thumbnails/%d_%s.jpg' % ( + return "generated-thumbnails/%d_%s.jpg" % ( post.post_id, - get_post_security_hash(post.post_id)) + get_post_security_hash(post.post_id), + ) def get_post_thumbnail_backup_path(post: model.Post) -> str: assert post - return 'posts/custom-thumbnails/%d_%s.dat' % ( - post.post_id, get_post_security_hash(post.post_id)) + return "posts/custom-thumbnails/%d_%s.dat" % ( + post.post_id, + get_post_security_hash(post.post_id), + ) def serialize_note(note: model.PostNote) -> rest.Response: assert note return { - 'polygon': note.polygon, - 'text': note.text, + "polygon": note.polygon, + "text": note.text, } @@ -145,40 +166,41 @@ class PostSerializer(serialization.BaseSerializer): def _serializers(self) -> Dict[str, Callable[[], Any]]: return { - 'id': self.serialize_id, - 'version': self.serialize_version, - 'creationTime': self.serialize_creation_time, - 'lastEditTime': self.serialize_last_edit_time, - 'safety': self.serialize_safety, - 'source': self.serialize_source, - 'type': self.serialize_type, - 'mimeType': self.serialize_mime, - 'checksum': self.serialize_checksum, - 'fileSize': self.serialize_file_size, - 'canvasWidth': self.serialize_canvas_width, - 'canvasHeight': self.serialize_canvas_height, - 'contentUrl': self.serialize_content_url, - 'thumbnailUrl': self.serialize_thumbnail_url, - 'flags': self.serialize_flags, - 'tags': self.serialize_tags, - 'relations': self.serialize_relations, - 'user': self.serialize_user, - 'score': self.serialize_score, - 'ownScore': self.serialize_own_score, - 'ownFavorite': self.serialize_own_favorite, - 'tagCount': self.serialize_tag_count, - 'favoriteCount': self.serialize_favorite_count, - 'commentCount': self.serialize_comment_count, - 'noteCount': self.serialize_note_count, - 'relationCount': self.serialize_relation_count, - 'featureCount': self.serialize_feature_count, - 'lastFeatureTime': self.serialize_last_feature_time, - 'favoritedBy': self.serialize_favorited_by, - 'hasCustomThumbnail': self.serialize_has_custom_thumbnail, - 'notes': self.serialize_notes, - 'comments': self.serialize_comments, - 'metrics': self.serialize_metrics, - 'metricRanges': self.serialize_metric_ranges, + "id": self.serialize_id, + "version": self.serialize_version, + "creationTime": self.serialize_creation_time, + "lastEditTime": self.serialize_last_edit_time, + "safety": self.serialize_safety, + "source": self.serialize_source, + "type": self.serialize_type, + "mimeType": self.serialize_mime, + "checksum": self.serialize_checksum, + "fileSize": self.serialize_file_size, + "canvasWidth": self.serialize_canvas_width, + "canvasHeight": self.serialize_canvas_height, + "contentUrl": self.serialize_content_url, + "thumbnailUrl": self.serialize_thumbnail_url, + "flags": self.serialize_flags, + "tags": self.serialize_tags, + "relations": self.serialize_relations, + "user": self.serialize_user, + "score": self.serialize_score, + "ownScore": self.serialize_own_score, + "ownFavorite": self.serialize_own_favorite, + "tagCount": self.serialize_tag_count, + "favoriteCount": self.serialize_favorite_count, + "commentCount": self.serialize_comment_count, + "noteCount": self.serialize_note_count, + "relationCount": self.serialize_relation_count, + "featureCount": self.serialize_feature_count, + "lastFeatureTime": self.serialize_last_feature_time, + "favoritedBy": self.serialize_favorited_by, + "hasCustomThumbnail": self.serialize_has_custom_thumbnail, + "notes": self.serialize_notes, + "comments": self.serialize_comments, + "metrics": self.serialize_metrics, + "metricRanges": self.serialize_metric_ranges, + "pools": self.serialize_pools, } def serialize_id(self) -> Any: @@ -229,25 +251,28 @@ class PostSerializer(serialization.BaseSerializer): def serialize_tags(self) -> Any: return [ { - 'names': [name.name for name in tag.names], - 'category': tag.category.name, - 'usages': tag.post_count, - 'metric': { - 'min': tag.metric.min, - 'max': tag.metric.max + "names": [name.name for name in tag.names], + "category": tag.category.name, + "usages": tag.post_count, + "metric": { + "min": tag.metric.min, + "max": tag.metric.max } if tag.metric else None, } - for tag in tags.sort_tags(self.post.tags)] + for tag in tags.sort_tags(self.post.tags) + ] def serialize_relations(self) -> Any: return sorted( { - post['id']: post + post["id"]: post for post in [ serialize_micro_post(rel, self.auth_user) - for rel in self.post.relations] + for rel in self.post.relations + ] }.values(), - key=lambda post: post['id']) + key=lambda post: post["id"], + ) def serialize_user(self) -> Any: return users.serialize_micro_user(self.post.user, self.auth_user) @@ -259,10 +284,16 @@ class PostSerializer(serialization.BaseSerializer): return scores.get_score(self.post, self.auth_user) def serialize_own_favorite(self) -> Any: - return len([ - user for user in self.post.favorited_by - if user.user_id == self.auth_user.user_id] - ) > 0 + return ( + len( + [ + user + for user in self.post.favorited_by + if user.user_id == self.auth_user.user_id + ] + ) + > 0 + ) def serialize_tag_count(self) -> Any: return self.post.tag_count @@ -297,43 +328,58 @@ class PostSerializer(serialization.BaseSerializer): def serialize_notes(self) -> Any: return sorted( [serialize_note(note) for note in self.post.notes], - key=lambda x: x['polygon']) + key=lambda x: x["polygon"], + ) def serialize_comments(self) -> Any: return [ comments.serialize_comment(comment, self.auth_user) for comment in sorted( - self.post.comments, - key=lambda comment: comment.creation_time)] + self.post.comments, key=lambda comment: comment.creation_time + ) + ] + + def serialize_pools(self) -> List[Any]: + return [ + pools.serialize_micro_pool(pool) + for pool in sorted( + self.post.pools, key=lambda pool: pool.creation_time + ) + ] def serialize_metrics(self) -> Any: return [ metrics.serialize_post_metric(metric) for metric in sorted( self.post.metrics, - key=lambda metric: metric.metric.tag_name)] + key=lambda metric: metric.metric.tag_name + ) + ] def serialize_metric_ranges(self) -> Any: return [ metrics.serialize_post_metric_range(metric_range) for metric_range in sorted( self.post.metric_ranges, - key=lambda metric_range: metric_range.metric.tag_name)] + key=lambda metric_range: metric_range.metric.tag_name + ) + ] def serialize_post( - post: Optional[model.Post], - auth_user: model.User, - options: List[str] = []) -> Optional[rest.Response]: + post: Optional[model.Post], auth_user: model.User, options: List[str] = [] +) -> Optional[rest.Response]: if not post: return None return PostSerializer(post, auth_user).serialize(options) def serialize_micro_post( - post: model.Post, auth_user: model.User) -> Optional[rest.Response]: + post: model.Post, auth_user: model.User +) -> Optional[rest.Response]: return serialize_post( - post, auth_user=auth_user, options=['id', 'thumbnailUrl']) + post, auth_user=auth_user, options=["id", "thumbnailUrl"] + ) def get_post_count() -> int: @@ -342,25 +388,37 @@ def get_post_count() -> int: def try_get_post_by_id(post_id: int) -> Optional[model.Post]: return ( - db.session - .query(model.Post) + db.session.query(model.Post) .filter(model.Post.post_id == post_id) - .one_or_none()) + .one_or_none() + ) def get_post_by_id(post_id: int) -> model.Post: post = try_get_post_by_id(post_id) if not post: - raise PostNotFoundError('Post %r not found.' % post_id) + raise PostNotFoundError("Post %r not found." % post_id) return post +def get_posts_by_ids(ids: List[int]) -> List[model.Post]: + if len(ids) == 0: + return [] + posts = ( + db.session.query(model.Post) + .filter(sa.sql.or_(model.Post.post_id == post_id for post_id in ids)) + .all() + ) + id_order = {v: k for k, v in enumerate(ids)} + return sorted(posts, key=lambda post: id_order.get(post.post_id)) + + def try_get_current_post_feature() -> Optional[model.PostFeature]: return ( - db.session - .query(model.PostFeature) + db.session.query(model.PostFeature) .order_by(model.PostFeature.time.desc()) - .first()) + .first() + ) def try_get_featured_post() -> Optional[model.Post]: @@ -369,22 +427,22 @@ def try_get_featured_post() -> Optional[model.Post]: def create_post( - content: bytes, - tag_names: List[str], - user: Optional[model.User]) -> Tuple[model.Post, List[model.Tag]]: + content: bytes, tag_names: List[str], user: Optional[model.User] +) -> Tuple[model.Post, List[model.Tag]]: post = model.Post() post.safety = model.Post.SAFETY_SAFE post.user = user post.creation_time = datetime.utcnow() post.flags = [] - post.type = '' - post.checksum = '' - post.mime_type = '' - db.session.add(post) + post.type = "" + post.checksum = "" + post.mime_type = "" update_post_content(post, content) new_tags = update_post_tags(post, tag_names) + + db.session.add(post) return post, new_tags @@ -393,35 +451,38 @@ def update_post_safety(post: model.Post, safety: str) -> None: safety = util.flip(SAFETY_MAP).get(safety, None) if not safety: raise InvalidPostSafetyError( - 'Safety can be either of %r.' % list(SAFETY_MAP.values())) + "Safety can be either of %r." % list(SAFETY_MAP.values()) + ) post.safety = safety def update_post_source(post: model.Post, source: Optional[str]) -> None: assert post if util.value_exceeds_column_size(source, model.Post.source): - raise InvalidPostSourceError('Source is too long.') + raise InvalidPostSourceError("Source is too long.") post.source = source or None -@sa.events.event.listens_for(model.Post, 'after_insert') +@sa.events.event.listens_for(model.Post, "after_insert") def _after_post_insert( - _mapper: Any, _connection: Any, post: model.Post) -> None: + _mapper: Any, _connection: Any, post: model.Post +) -> None: _sync_post_content(post) -@sa.events.event.listens_for(model.Post, 'after_update') +@sa.events.event.listens_for(model.Post, "after_update") def _after_post_update( - _mapper: Any, _connection: Any, post: model.Post) -> None: + _mapper: Any, _connection: Any, post: model.Post +) -> None: _sync_post_content(post) -@sa.events.event.listens_for(model.Post, 'before_delete') +@sa.events.event.listens_for(model.Post, "before_delete") def _before_post_delete( - _mapper: Any, _connection: Any, post: model.Post) -> None: + _mapper: Any, _connection: Any, post: model.Post +) -> None: if post.post_id: - image_hash.delete_image(post.post_id) - if config.config['delete_source_files']: + if config.config["delete_source_files"]: files.delete(get_post_content_path(post)) files.delete(get_post_thumbnail_path(post)) @@ -429,54 +490,50 @@ def _before_post_delete( def _sync_post_content(post: model.Post) -> None: regenerate_thumb = False - if hasattr(post, '__content'): - content = getattr(post, '__content') + if hasattr(post, "__content"): + content = getattr(post, "__content") files.save(get_post_content_path(post), content) - delattr(post, '__content') + delattr(post, "__content") regenerate_thumb = True - if post.post_id and post.type in ( - model.Post.TYPE_IMAGE, model.Post.TYPE_ANIMATION): - image_hash.delete_image(post.post_id) - image_hash.add_image(post.post_id, content) - if hasattr(post, '__thumbnail'): - if getattr(post, '__thumbnail'): + if hasattr(post, "__thumbnail"): + if getattr(post, "__thumbnail"): files.save( get_post_thumbnail_backup_path(post), - getattr(post, '__thumbnail')) + getattr(post, "__thumbnail"), + ) else: files.delete(get_post_thumbnail_backup_path(post)) - delattr(post, '__thumbnail') + delattr(post, "__thumbnail") regenerate_thumb = True if regenerate_thumb: generate_post_thumbnail(post) -def generate_alternate_formats(post: model.Post, content: bytes) \ - -> List[Tuple[model.Post, List[model.Tag]]]: +def generate_alternate_formats( + post: model.Post, content: bytes +) -> List[Tuple[model.Post, List[model.Tag]]]: assert post assert content new_posts = [] if mime.is_animated_gif(content): tag_names = [tag.first_name for tag in post.tags] - if config.config['convert']['gif']['to_mp4']: + if config.config["convert"]["gif"]["to_mp4"]: mp4_post, new_tags = create_post( - images.Image(content).to_mp4(), - tag_names, - post.user) - update_post_flags(mp4_post, ['loop']) + images.Image(content).to_mp4(), tag_names, post.user + ) + update_post_flags(mp4_post, ["loop"]) update_post_safety(mp4_post, post.safety) update_post_source(mp4_post, post.source) new_posts += [(mp4_post, new_tags)] - if config.config['convert']['gif']['to_webm']: + if config.config["convert"]["gif"]["to_webm"]: webm_post, new_tags = create_post( - images.Image(content).to_webm(), - tag_names, - post.user) - update_post_flags(webm_post, ['loop']) + images.Image(content).to_webm(), tag_names, post.user + ) + update_post_flags(webm_post, ["loop"]) update_post_safety(webm_post, post.safety) update_post_source(webm_post, post.source) new_posts += [(webm_post, new_tags)] @@ -492,25 +549,75 @@ def generate_alternate_formats(post: model.Post, content: bytes) \ return new_posts -def test_sound(post: model.Post, content: bytes) -> None: - assert post +def get_default_flags(content: bytes) -> List[str]: assert content + ret = [] if mime.is_video(mime.get_mime_type(content)): + ret.append(model.Post.FLAG_LOOP) if images.Image(content).check_for_sound(): - flags = post.flags - if model.Post.FLAG_SOUND not in flags: - flags.append(model.Post.FLAG_SOUND) - update_post_flags(post, flags) + ret.append(model.Post.FLAG_SOUND) + return ret + + +def purge_post_signature(post: model.Post) -> None: + ( + db.session.query(model.PostSignature) + .filter(model.PostSignature.post_id == post.post_id) + .delete() + ) + + +def generate_post_signature(post: model.Post, content: bytes) -> None: + try: + unpacked_signature = image_hash.generate_signature(content) + packed_signature = image_hash.pack_signature(unpacked_signature) + words = image_hash.generate_words(unpacked_signature) + + db.session.add( + model.PostSignature( + post=post, signature=packed_signature, words=words + ) + ) + except errors.ProcessingError: + if not config.config["allow_broken_uploads"]: + raise InvalidPostContentError( + "Unable to generate image hash data." + ) + + +def update_all_post_signatures() -> None: + posts_to_hash = ( + db.session.query(model.Post) + .filter( + (model.Post.type == model.Post.TYPE_IMAGE) + | (model.Post.type == model.Post.TYPE_ANIMATION) + ) + .filter(model.Post.signature == None) # noqa: E711 + .order_by(model.Post.post_id.asc()) + .all() + ) + for post in posts_to_hash: + try: + generate_post_signature( + post, files.get(get_post_content_path(post)) + ) + db.session.commit() + logger.info("Hashed Post %d", post.post_id) + except Exception as ex: + logger.exception(ex) def update_post_content(post: model.Post, content: Optional[bytes]) -> None: assert post if not content: - raise InvalidPostContentError('Post content missing.') + raise InvalidPostContentError("Post content missing.") + + update_signature = False post.mime_type = mime.get_mime_type(content) if mime.is_flash(post.mime_type): post.type = model.Post.TYPE_FLASH elif mime.is_image(post.mime_type): + update_signature = True if mime.is_animated_gif(content): post.type = model.Post.TYPE_ANIMATION else: @@ -519,39 +626,56 @@ def update_post_content(post: model.Post, content: Optional[bytes]) -> None: post.type = model.Post.TYPE_VIDEO else: raise InvalidPostContentError( - 'Unhandled file type: %r' % post.mime_type) + "Unhandled file type: %r" % post.mime_type + ) post.checksum = util.get_sha1(content) other_post = ( - db.session - .query(model.Post) + db.session.query(model.Post) .filter(model.Post.checksum == post.checksum) .filter(model.Post.post_id != post.post_id) - .one_or_none()) - if other_post \ - and other_post.post_id \ - and other_post.post_id != post.post_id: + .one_or_none() + ) + if ( + other_post + and other_post.post_id + and other_post.post_id != post.post_id + ): raise PostAlreadyUploadedError(other_post) + if update_signature: + purge_post_signature(post) + post.signature = generate_post_signature(post, content) + post.file_size = len(content) try: image = images.Image(content) post.canvas_width = image.width post.canvas_height = image.height except errors.ProcessingError: - post.canvas_width = None - post.canvas_height = None - if (post.canvas_width is not None and post.canvas_width <= 0) \ - or (post.canvas_height is not None and post.canvas_height <= 0): - post.canvas_width = None - post.canvas_height = None - setattr(post, '__content', content) + if not config.config["allow_broken_uploads"]: + raise InvalidPostContentError("Unable to process image metadata") + else: + post.canvas_width = None + post.canvas_height = None + if (post.canvas_width is not None and post.canvas_width <= 0) or ( + post.canvas_height is not None and post.canvas_height <= 0 + ): + if not config.config["allow_broken_uploads"]: + raise InvalidPostContentError( + "Invalid image dimensions returned during processing" + ) + else: + post.canvas_width = None + post.canvas_height = None + setattr(post, "__content", content) def update_post_thumbnail( - post: model.Post, content: Optional[bytes] = None) -> None: + post: model.Post, content: Optional[bytes] = None +) -> None: assert post - setattr(post, '__thumbnail', content) + setattr(post, "__thumbnail", content) def generate_post_thumbnail(post: model.Post) -> None: @@ -564,15 +688,17 @@ def generate_post_thumbnail(post: model.Post) -> None: assert content image = images.Image(content) image.resize_fill( - int(config.config['thumbnails']['post_width']), - int(config.config['thumbnails']['post_height'])) + int(config.config["thumbnails"]["post_width"]), + int(config.config["thumbnails"]["post_height"]), + ) files.save(get_post_thumbnail_path(post), image.to_jpeg()) except errors.ProcessingError: files.save(get_post_thumbnail_path(post), EMPTY_PIXEL) def update_post_tags( - post: model.Post, tag_names: List[str]) -> List[model.Tag]: + post: model.Post, tag_names: List[str] +) -> List[model.Tag]: assert post existing_tags, new_tags = tags.get_or_create_tags_by_names(tag_names) post.tags = existing_tags + new_tags @@ -584,22 +710,21 @@ def update_post_relations(post: model.Post, new_post_ids: List[int]) -> None: try: new_post_ids = [int(id) for id in new_post_ids] except ValueError: - raise InvalidPostRelationError( - 'A relation must be numeric post ID.') + raise InvalidPostRelationError("A relation must be numeric post ID.") old_posts = post.relations old_post_ids = [int(p.post_id) for p in old_posts] if new_post_ids: new_posts = ( - db.session - .query(model.Post) + db.session.query(model.Post) .filter(model.Post.post_id.in_(new_post_ids)) - .all()) + .all() + ) else: new_posts = [] if len(new_posts) != len(new_post_ids): - raise InvalidPostRelationError('One of relations does not exist.') + raise InvalidPostRelationError("One of relations does not exist.") if post.post_id in new_post_ids: - raise InvalidPostRelationError('Post cannot relate to itself.') + raise InvalidPostRelationError("Post cannot relate to itself.") relations_to_del = [p for p in old_posts if p.post_id not in new_post_ids] relations_to_add = [p for p in new_posts if p.post_id not in old_post_ids] @@ -615,37 +740,44 @@ def update_post_notes(post: model.Post, notes: Any) -> None: assert post post.notes = [] for note in notes: - for field in ('polygon', 'text'): + for field in ("polygon", "text"): if field not in note: - raise InvalidPostNoteError('Note is missing %r field.' % field) - if not note['text']: - raise InvalidPostNoteError('A note\'s text cannot be empty.') - if not isinstance(note['polygon'], (list, tuple)): + raise InvalidPostNoteError("Note is missing %r field." % field) + if not note["text"]: + raise InvalidPostNoteError("A note's text cannot be empty.") + if not isinstance(note["polygon"], (list, tuple)): raise InvalidPostNoteError( - 'A note\'s polygon must be a list of points.') - if len(note['polygon']) < 3: + "A note's polygon must be a list of points." + ) + if len(note["polygon"]) < 3: raise InvalidPostNoteError( - 'A note\'s polygon must have at least 3 points.') - for point in note['polygon']: + "A note's polygon must have at least 3 points." + ) + for point in note["polygon"]: if not isinstance(point, (list, tuple)): raise InvalidPostNoteError( - 'A note\'s polygon point must be a list of length 2.') + "A note's polygon point must be a list of length 2." + ) if len(point) != 2: raise InvalidPostNoteError( - 'A point in note\'s polygon must have two coordinates.') + "A point in note's polygon must have two coordinates." + ) try: pos_x = float(point[0]) pos_y = float(point[1]) if not 0 <= pos_x <= 1 or not 0 <= pos_y <= 1: raise InvalidPostNoteError( - 'All points must fit in the image (0..1 range).') + "All points must fit in the image (0..1 range)." + ) except ValueError: raise InvalidPostNoteError( - 'A point in note\'s polygon must be numeric.') - if util.value_exceeds_column_size(note['text'], model.PostNote.text): - raise InvalidPostNoteError('Note text is too long.') + "A point in note's polygon must be numeric." + ) + if util.value_exceeds_column_size(note["text"], model.PostNote.text): + raise InvalidPostNoteError("Note text is too long.") post.notes.append( - model.PostNote(polygon=note['polygon'], text=str(note['text']))) + model.PostNote(polygon=note["polygon"], text=str(note["text"])) + ) def update_post_flags(post: model.Post, flags: List[str]) -> None: @@ -655,7 +787,8 @@ def update_post_flags(post: model.Post, flags: List[str]) -> None: flag = util.flip(FLAG_MAP).get(flag, None) if not flag: raise InvalidPostFlagError( - 'Flag must be one of %r.' % list(FLAG_MAP.values())) + "Flag must be one of %r." % list(FLAG_MAP.values()) + ) target_flags.append(flag) post.flags = target_flags @@ -675,32 +808,31 @@ def delete(post: model.Post) -> None: def merge_posts( - source_post: model.Post, - target_post: model.Post, - replace_content: bool) -> None: + source_post: model.Post, target_post: model.Post, replace_content: bool +) -> None: assert source_post assert target_post if source_post.post_id == target_post.post_id: - raise InvalidPostRelationError('Cannot merge post with itself.') + raise InvalidPostRelationError("Cannot merge post with itself.") def merge_tables( - table: model.Base, - anti_dup_func: Optional[Callable[[model.Base, model.Base], bool]], - source_post_id: int, - target_post_id: int) -> None: + table: model.Base, + anti_dup_func: Optional[Callable[[model.Base, model.Base], bool]], + source_post_id: int, + target_post_id: int, + ) -> None: alias1 = table alias2 = sa.orm.util.aliased(table) - update_stmt = ( - sa.sql.expression.update(alias1) - .where(alias1.post_id == source_post_id)) + update_stmt = sa.sql.expression.update(alias1).where( + alias1.post_id == source_post_id + ) if anti_dup_func is not None: - update_stmt = ( - update_stmt - .where( - ~sa.exists() - .where(anti_dup_func(alias1, alias2)) - .where(alias2.post_id == target_post_id))) + update_stmt = update_stmt.where( + ~sa.exists() + .where(anti_dup_func(alias1, alias2)) + .where(alias2.post_id == target_post_id) + ) update_stmt = update_stmt.values(post_id=target_post_id) db.session.execute(update_stmt) @@ -710,21 +842,24 @@ def merge_posts( model.PostTag, lambda alias1, alias2: alias1.tag_id == alias2.tag_id, source_post_id, - target_post_id) + target_post_id, + ) def merge_scores(source_post_id: int, target_post_id: int) -> None: merge_tables( model.PostScore, lambda alias1, alias2: alias1.user_id == alias2.user_id, source_post_id, - target_post_id) + target_post_id, + ) def merge_favorites(source_post_id: int, target_post_id: int) -> None: merge_tables( model.PostFavorite, lambda alias1, alias2: alias1.user_id == alias2.user_id, source_post_id, - target_post_id) + target_post_id, + ) def merge_comments(source_post_id: int, target_post_id: int) -> None: merge_tables(model.Comment, None, source_post_id, target_post_id) @@ -739,8 +874,10 @@ def merge_posts( .where( ~sa.exists() .where(alias2.child_id == alias1.child_id) - .where(alias2.parent_id == target_post_id)) - .values(parent_id=target_post_id)) + .where(alias2.parent_id == target_post_id) + ) + .values(parent_id=target_post_id) + ) db.session.execute(update_stmt) update_stmt = ( @@ -750,26 +887,31 @@ def merge_posts( .where( ~sa.exists() .where(alias2.parent_id == alias1.parent_id) - .where(alias2.child_id == target_post_id)) - .values(child_id=target_post_id)) + .where(alias2.child_id == target_post_id) + ) + .values(child_id=target_post_id) + ) db.session.execute(update_stmt) - def transfer_flags(source_post_id: int, target_post_id: int) -> None: - target = get_post_by_id(target_post_id) - source = get_post_by_id(source_post_id) - target.flags = source.flags - merge_tags(source_post.post_id, target_post.post_id) merge_comments(source_post.post_id, target_post.post_id) merge_scores(source_post.post_id, target_post.post_id) merge_favorites(source_post.post_id, target_post.post_id) merge_relations(source_post.post_id, target_post.post_id) + def transfer_flags(source_post_id: int, target_post_id: int) -> None: + target = get_post_by_id(target_post_id) + source = get_post_by_id(source_post_id) + target.flags = source.flags + db.session.flush() + content = None if replace_content: content = files.get(get_post_content_path(source_post)) transfer_flags(source_post.post_id, target_post.post_id) + # fixes unknown issue with SA's cascade deletions + purge_post_signature(source_post) delete(source_post) db.session.flush() @@ -780,44 +922,50 @@ def merge_posts( def search_by_image_exact(image_content: bytes) -> Optional[model.Post]: checksum = util.get_sha1(image_content) return ( - db.session - .query(model.Post) + db.session.query(model.Post) .filter(model.Post.checksum == checksum) - .one_or_none()) + .one_or_none() + ) -def search_by_image(image_content: bytes) -> List[PostLookalike]: - ret = [] - for result in image_hash.search_by_image(image_content): - post = try_get_post_by_id(result.path) - if post: - ret.append(PostLookalike( - score=result.score, - distance=result.distance, - post=post)) - return ret +def search_by_image(image_content: bytes) -> List[Tuple[float, model.Post]]: + query_signature = image_hash.generate_signature(image_content) + query_words = image_hash.generate_words(query_signature) + """ + The unnest function is used here to expand one row containing the 'words' + array into multiple rows each containing a singular word. -def populate_reverse_search() -> None: - excluded_post_ids = image_hash.get_all_paths() + Documentation of the unnest function can be found here: + https://www.postgresql.org/docs/9.2/functions-array.html + """ - post_ids_to_hash = ( - db.session - .query(model.Post.post_id) - .filter( - (model.Post.type == model.Post.TYPE_IMAGE) | - (model.Post.type == model.Post.TYPE_ANIMATION)) - .filter(~model.Post.post_id.in_(excluded_post_ids)) - .order_by(model.Post.post_id.asc()) - .all()) + dbquery = """ + SELECT s.post_id, s.signature, count(a.query) AS score + FROM post_signature AS s, unnest(s.words, :q) AS a(word, query) + WHERE a.word = a.query + GROUP BY s.post_id + ORDER BY score DESC LIMIT 100; + """ - for post_ids_chunk in util.chunks(post_ids_to_hash, 100): - posts_chunk = ( - db.session - .query(model.Post) - .filter(model.Post.post_id.in_(post_ids_chunk)) - .all()) - for post in posts_chunk: - content_path = get_post_content_path(post) - if files.has(content_path): - image_hash.add_image(post.post_id, files.get(content_path)) + candidates = db.session.execute(dbquery, {"q": query_words}) + data = tuple( + zip( + *[ + (post_id, image_hash.unpack_signature(packedsig)) + for post_id, packedsig, score in candidates + ] + ) + ) + if data: + candidate_post_ids, sigarray = data + distances = image_hash.normalized_distance(sigarray, query_signature) + return [ + (distance, try_get_post_by_id(candidate_post_id)) + for candidate_post_id, distance in zip( + candidate_post_ids, distances + ) + if distance < image_hash.DISTANCE_CUTOFF + ] + else: + return [] diff --git a/server/szurubooru/func/scores.py b/server/szurubooru/func/scores.py index 615fd98..b095f48 100644 --- a/server/szurubooru/func/scores.py +++ b/server/szurubooru/func/scores.py @@ -1,6 +1,7 @@ import datetime -from typing import Any, Tuple, Callable -from szurubooru import db, model, errors +from typing import Any, Callable, Tuple + +from szurubooru import db, errors, model class InvalidScoreTargetError(errors.ValidationError): @@ -12,12 +13,13 @@ class InvalidScoreValueError(errors.ValidationError): def _get_table_info( - entity: model.Base) -> Tuple[model.Base, Callable[[model.Base], Any]]: + entity: model.Base, +) -> Tuple[model.Base, Callable[[model.Base], Any]]: assert entity resource_type, _, _ = model.util.get_resource_info(entity) - if resource_type == 'post': + if resource_type == "post": return model.PostScore, lambda table: table.post_id - elif resource_type == 'comment': + elif resource_type == "comment": return model.CommentScore, lambda table: table.comment_id raise InvalidScoreTargetError() @@ -40,16 +42,17 @@ def get_score(entity: model.Base, user: model.User) -> int: assert user table, get_column = _get_table_info(entity) row = ( - db.session - .query(table.score) + db.session.query(table.score) .filter(get_column(table) == get_column(entity)) .filter(table.user_id == user.user_id) - .one_or_none()) + .one_or_none() + ) return row[0] if row else 0 def set_score(entity: model.Base, user: model.User, score: int) -> None: from szurubooru.func import favorites + assert entity assert user if not score: @@ -61,7 +64,8 @@ def set_score(entity: model.Base, user: model.User, score: int) -> None: return if score not in (-1, 1): raise InvalidScoreValueError( - 'Score %r is invalid. Valid scores: %r.' % (score, (-1, 1))) + "Score %r is invalid. Valid scores: %r." % (score, (-1, 1)) + ) score_entity = _get_score_entity(entity, user) if score_entity: score_entity.score = score diff --git a/server/szurubooru/func/serialization.py b/server/szurubooru/func/serialization.py index 699fb47..d2fadc0 100644 --- a/server/szurubooru/func/serialization.py +++ b/server/szurubooru/func/serialization.py @@ -1,9 +1,10 @@ -from typing import Any, List, Dict, Callable -from szurubooru import model, rest, errors +from typing import Any, Callable, Dict, List + +from szurubooru import errors, model, rest def get_serialization_options(ctx: rest.Context) -> List[str]: - return ctx.get_param_as_list('fields', default=[]) + return ctx.get_param_as_list("fields", default=[]) class BaseSerializer: @@ -17,8 +18,9 @@ class BaseSerializer: for key in options: if key not in field_factories: raise errors.ValidationError( - 'Invalid key: %r. Valid keys: %r.' % ( - key, list(sorted(field_factories.keys())))) + "Invalid key: %r. Valid keys: %r." + % (key, list(sorted(field_factories.keys()))) + ) factory = field_factories[key] ret[key] = factory() return ret diff --git a/server/szurubooru/func/snapshots.py b/server/szurubooru/func/snapshots.py index 240c3bc..afb26ea 100644 --- a/server/szurubooru/func/snapshots.py +++ b/server/szurubooru/func/snapshots.py @@ -1,73 +1,111 @@ -from typing import Any, Optional, Dict, Callable from datetime import datetime +from typing import Any, Callable, Dict, Optional + +import sqlalchemy as sa + from szurubooru import db, model -from szurubooru.func import diff, users +from szurubooru.func import diff, net, users def get_tag_category_snapshot(category: model.TagCategory) -> Dict[str, Any]: assert category return { - 'name': category.name, - 'color': category.color, - 'default': True if category.default else False, + "name": category.name, + "color": category.color, + "default": True if category.default else False, } def get_tag_snapshot(tag: model.Tag) -> Dict[str, Any]: assert tag return { - 'names': [tag_name.name for tag_name in tag.names], - 'category': tag.category.name, - 'suggestions': sorted(rel.first_name for rel in tag.suggestions), - 'implications': sorted(rel.first_name for rel in tag.implications), + "names": [tag_name.name for tag_name in tag.names], + "category": tag.category.name, + "suggestions": sorted(rel.first_name for rel in tag.suggestions), + "implications": sorted(rel.first_name for rel in tag.implications), + } + + +def get_pool_category_snapshot(category: model.PoolCategory) -> Dict[str, Any]: + assert category + return { + "name": category.name, + "color": category.color, + "default": True if category.default else False, + } + + +def get_pool_snapshot(pool: model.Pool) -> Dict[str, Any]: + assert pool + return { + "names": [pool_name.name for pool_name in pool.names], + "category": pool.category.name, + "posts": [post.post_id for post in pool.posts], } def get_post_snapshot(post: model.Post) -> Dict[str, Any]: assert post return { - 'source': post.source, - 'safety': post.safety, - 'checksum': post.checksum, - 'flags': post.flags, - 'featured': post.is_featured, - 'tags': sorted([tag.first_name for tag in post.tags]), - 'relations': sorted([rel.post_id for rel in post.relations]), - 'notes': sorted([{ - 'polygon': [[point[0], point[1]] for point in note.polygon], - 'text': note.text, - } for note in post.notes], key=lambda x: x['polygon']), + "source": post.source, + "safety": post.safety, + "checksum": post.checksum, + "flags": post.flags, + "featured": post.is_featured, + "tags": sorted([tag.first_name for tag in post.tags]), + "relations": sorted([rel.post_id for rel in post.relations]), + "notes": sorted( + [ + { + "polygon": [ + [point[0], point[1]] for point in note.polygon + ], + "text": note.text, + } + for note in post.notes + ], + key=lambda x: x["polygon"], + ), } _snapshot_factories = { # lambdas allow mocking target functions in the tests - # pylint: disable=unnecessary-lambda - 'tag_category': lambda entity: get_tag_category_snapshot(entity), - 'tag': lambda entity: get_tag_snapshot(entity), - 'post': lambda entity: get_post_snapshot(entity), + "tag_category": lambda entity: get_tag_category_snapshot(entity), + "tag": lambda entity: get_tag_snapshot(entity), + "post": lambda entity: get_post_snapshot(entity), + "pool_category": lambda entity: get_pool_category_snapshot(entity), + "pool": lambda entity: get_pool_snapshot(entity), } # type: Dict[model.Base, Callable[[model.Base], Dict[str ,Any]]] def serialize_snapshot( - snapshot: model.Snapshot, auth_user: model.User) -> Dict[str, Any]: + snapshot: model.Snapshot, auth_user: model.User +) -> Dict[str, Any]: assert snapshot return { - 'operation': snapshot.operation, - 'type': snapshot.resource_type, - 'id': snapshot.resource_name, - 'user': users.serialize_micro_user(snapshot.user, auth_user), - 'data': snapshot.data, - 'time': snapshot.creation_time, + "operation": snapshot.operation, + "type": snapshot.resource_type, + "id": snapshot.resource_name, + "user": users.serialize_micro_user(snapshot.user, auth_user), + "data": snapshot.data, + "time": snapshot.creation_time, } +def _post_to_webhooks(snapshot: model.Snapshot) -> None: + webhook_user = model.User() + webhook_user.name = None + webhook_user.rank = "anonymous" + net.post_to_webhooks(serialize_snapshot(snapshot, webhook_user)) + + def _create( - operation: str, - entity: model.Base, - auth_user: Optional[model.User]) -> model.Snapshot: - resource_type, resource_pkey, resource_name = ( - model.util.get_resource_info(entity)) + operation: str, entity: model.Base, auth_user: Optional[model.User] +) -> model.Snapshot: + resource_type, resource_pkey, resource_name = model.util.get_resource_info( + entity + ) snapshot = model.Snapshot() snapshot.creation_time = datetime.utcnow() @@ -85,9 +123,9 @@ def create(entity: model.Base, auth_user: Optional[model.User]) -> None: snapshot_factory = _snapshot_factories[snapshot.resource_type] snapshot.data = snapshot_factory(entity) db.session.add(snapshot) + _post_to_webhooks(snapshot) -# pylint: disable=protected-access def modify(entity: model.Base, auth_user: Optional[model.User]) -> None: assert entity @@ -95,18 +133,19 @@ def modify(entity: model.Base, auth_user: Optional[model.User]) -> None: ( cls for cls in model.Base._decl_class_registry.values() - if hasattr(cls, '__table__') + if hasattr(cls, "__table__") and cls.__table__.fullname == entity.__table__.fullname ), - None) + None, + ) assert table snapshot = _create(model.Snapshot.OPERATION_MODIFIED, entity, auth_user) snapshot_factory = _snapshot_factories[snapshot.resource_type] - detached_session = db.sessionmaker() + detached_session = sa.orm.sessionmaker(bind=db.session.get_bind())() detached_entity = detached_session.query(table).get(snapshot.resource_pkey) - assert detached_entity, 'Entity not found in DB, have you committed it?' + assert detached_entity, "Entity not found in DB, have you committed it?" detached_snapshot = snapshot_factory(detached_entity) detached_session.close() @@ -116,6 +155,7 @@ def modify(entity: model.Base, auth_user: Optional[model.User]) -> None: if not snapshot.data: return db.session.add(snapshot) + _post_to_webhooks(snapshot) def delete(entity: model.Base, auth_user: Optional[model.User]) -> None: @@ -124,17 +164,24 @@ def delete(entity: model.Base, auth_user: Optional[model.User]) -> None: snapshot_factory = _snapshot_factories[snapshot.resource_type] snapshot.data = snapshot_factory(entity) db.session.add(snapshot) + _post_to_webhooks(snapshot) def merge( - source_entity: model.Base, - target_entity: model.Base, - auth_user: Optional[model.User]) -> None: + source_entity: model.Base, + target_entity: model.Base, + auth_user: Optional[model.User], +) -> None: assert source_entity assert target_entity snapshot = _create( - model.Snapshot.OPERATION_MERGED, source_entity, auth_user) - resource_type, _resource_pkey, resource_name = ( - model.util.get_resource_info(target_entity)) + model.Snapshot.OPERATION_MERGED, source_entity, auth_user + ) + ( + resource_type, + _resource_pkey, + resource_name, + ) = model.util.get_resource_info(target_entity) snapshot.data = [resource_type, resource_name] db.session.add(snapshot) + _post_to_webhooks(snapshot) diff --git a/server/szurubooru/func/tag_categories.py b/server/szurubooru/func/tag_categories.py index bec2f0d..d2c6021 100644 --- a/server/szurubooru/func/tag_categories.py +++ b/server/szurubooru/func/tag_categories.py @@ -1,11 +1,12 @@ import re -from typing import Any, Optional, Dict, List, Callable +from typing import Any, Callable, Dict, List, Optional + import sqlalchemy as sa -from szurubooru import config, db, model, errors, rest -from szurubooru.func import util, serialization, cache +from szurubooru import config, db, errors, model, rest +from szurubooru.func import cache, serialization, util -DEFAULT_CATEGORY_NAME_CACHE_KEY = 'default-tag-category' +DEFAULT_CATEGORY_NAME_CACHE_KEY = "default-tag-category" class TagCategoryNotFoundError(errors.NotFoundError): @@ -29,10 +30,11 @@ class InvalidTagCategoryColorError(errors.ValidationError): def _verify_name_validity(name: str) -> None: - name_regex = config.config['tag_category_name_regex'] + name_regex = config.config["tag_category_name_regex"] if not re.match(name_regex, name): raise InvalidTagCategoryNameError( - 'Name must satisfy regex %r.' % name_regex) + "Name must satisfy regex %r." % name_regex + ) class TagCategorySerializer(serialization.BaseSerializer): @@ -41,11 +43,12 @@ class TagCategorySerializer(serialization.BaseSerializer): def _serializers(self) -> Dict[str, Callable[[], Any]]: return { - 'name': self.serialize_name, - 'version': self.serialize_version, - 'color': self.serialize_color, - 'usages': self.serialize_usages, - 'default': self.serialize_default, + "name": self.serialize_name, + "version": self.serialize_version, + "color": self.serialize_color, + "usages": self.serialize_usages, + "default": self.serialize_default, + "order": self.serialize_order, } def serialize_name(self) -> Any: @@ -63,19 +66,23 @@ class TagCategorySerializer(serialization.BaseSerializer): def serialize_default(self) -> Any: return self.category.default + def serialize_order(self) -> Any: + return self.category.order + def serialize_category( - category: Optional[model.TagCategory], - options: List[str] = []) -> Optional[rest.Response]: + category: Optional[model.TagCategory], options: List[str] = [] +) -> Optional[rest.Response]: if not category: return None return TagCategorySerializer(category).serialize(options) -def create_category(name: str, color: str) -> model.TagCategory: +def create_category(name: str, color: str, order: int) -> model.TagCategory: category = model.TagCategory() update_category_name(category, name) update_category_color(category, color) + update_category_order(category, order) if not get_all_categories(): category.default = True return category @@ -84,18 +91,21 @@ def create_category(name: str, color: str) -> model.TagCategory: def update_category_name(category: model.TagCategory, name: str) -> None: assert category if not name: - raise InvalidTagCategoryNameError('Name cannot be empty.') + raise InvalidTagCategoryNameError("Name cannot be empty.") expr = sa.func.lower(model.TagCategory.name) == name.lower() if category.tag_category_id: expr = expr & ( - model.TagCategory.tag_category_id != category.tag_category_id) + model.TagCategory.tag_category_id != category.tag_category_id + ) already_exists = ( - db.session.query(model.TagCategory).filter(expr).count() > 0) + db.session.query(model.TagCategory).filter(expr).count() > 0 + ) if already_exists: raise TagCategoryAlreadyExistsError( - 'A category with this name already exists.') + "A category with this name already exists." + ) if util.value_exceeds_column_size(name, model.TagCategory.name): - raise InvalidTagCategoryNameError('Name is too long.') + raise InvalidTagCategoryNameError("Name is too long.") _verify_name_validity(name) category.name = name cache.remove(DEFAULT_CATEGORY_NAME_CACHE_KEY) @@ -104,58 +114,66 @@ def update_category_name(category: model.TagCategory, name: str) -> None: def update_category_color(category: model.TagCategory, color: str) -> None: assert category if not color: - raise InvalidTagCategoryColorError('Color cannot be empty.') - if not re.match(r'^#?[0-9a-z]+$', color): - raise InvalidTagCategoryColorError('Invalid color.') + raise InvalidTagCategoryColorError("Color cannot be empty.") + if not re.match(r"^#?[0-9a-z]+$", color): + raise InvalidTagCategoryColorError("Invalid color.") if util.value_exceeds_column_size(color, model.TagCategory.color): - raise InvalidTagCategoryColorError('Color is too long.') + raise InvalidTagCategoryColorError("Color is too long.") category.color = color +def update_category_order(category: model.TagCategory, order: int) -> None: + assert category + category.order = order + + def try_get_category_by_name( - name: str, lock: bool = False) -> Optional[model.TagCategory]: - query = ( - db.session - .query(model.TagCategory) - .filter(sa.func.lower(model.TagCategory.name) == name.lower())) + name: str, lock: bool = False +) -> Optional[model.TagCategory]: + query = db.session.query(model.TagCategory).filter( + sa.func.lower(model.TagCategory.name) == name.lower() + ) if lock: - query = query.with_lockmode('update') + query = query.with_for_update() return query.one_or_none() def get_category_by_name(name: str, lock: bool = False) -> model.TagCategory: category = try_get_category_by_name(name, lock) if not category: - raise TagCategoryNotFoundError('Tag category %r not found.' % name) + raise TagCategoryNotFoundError("Tag category %r not found." % name) return category def get_all_category_names() -> List[str]: - return [row[0] for row in db.session.query(model.TagCategory.name).all()] + return [cat.name for cat in get_all_categories()] def get_all_categories() -> List[model.TagCategory]: - return db.session.query(model.TagCategory).all() + return ( + db.session.query(model.TagCategory) + .order_by(model.TagCategory.order.asc(), model.TagCategory.name.asc()) + .all() + ) def try_get_default_category( - lock: bool = False) -> Optional[model.TagCategory]: - query = ( - db.session - .query(model.TagCategory) - .filter(model.TagCategory.default)) + lock: bool = False, +) -> Optional[model.TagCategory]: + query = db.session.query(model.TagCategory).filter( + model.TagCategory.default + ) if lock: - query = query.with_lockmode('update') + query = query.with_for_update() category = query.first() # if for some reason (e.g. as a result of migration) there's no default # category, get the first record available. if not category: - query = ( - db.session - .query(model.TagCategory) - .order_by(model.TagCategory.tag_category_id.asc())) + query = db.session.query(model.TagCategory).order_by( + model.TagCategory.tag_category_id.asc() + ) if lock: - query = query.with_lockmode('update') + query = query.with_for_update() category = query.first() return category @@ -163,7 +181,7 @@ def try_get_default_category( def get_default_category(lock: bool = False) -> model.TagCategory: category = try_get_default_category(lock) if not category: - raise TagCategoryNotFoundError('No tag category created yet.') + raise TagCategoryNotFoundError("No tag category created yet.") return category @@ -190,9 +208,10 @@ def set_default_category(category: model.TagCategory) -> None: def delete_category(category: model.TagCategory) -> None: assert category if len(get_all_category_names()) == 1: - raise TagCategoryIsInUseError('Cannot delete the last category.') + raise TagCategoryIsInUseError("Cannot delete the last category.") if (category.tag_count or 0) > 0: raise TagCategoryIsInUseError( - 'Tag category has some usages and cannot be deleted. ' + - 'Please remove this category from relevant tags first..') + "Tag category has some usages and cannot be deleted. " + + "Please remove this category from relevant tags first.." + ) db.session.delete(category) diff --git a/server/szurubooru/func/tags.py b/server/szurubooru/func/tags.py index 3384b4e..4981c5d 100644 --- a/server/szurubooru/func/tags.py +++ b/server/szurubooru/func/tags.py @@ -1,11 +1,11 @@ -import json -import os import re -from typing import Any, Optional, Tuple, List, Dict, Callable from datetime import datetime +from typing import Any, Callable, Dict, List, Optional, Tuple + import sqlalchemy as sa -from szurubooru import config, db, model, errors, rest -from szurubooru.func import util, tag_categories, serialization + +from szurubooru import config, db, errors, model, rest +from szurubooru.func import serialization, tag_categories, util class TagNotFoundError(errors.NotFoundError): @@ -38,10 +38,10 @@ class InvalidTagDescriptionError(errors.ValidationError): def _verify_name_validity(name: str) -> None: if util.value_exceeds_column_size(name, model.TagName.name): - raise InvalidTagNameError('Name is too long.') - name_regex = config.config['tag_name_regex'] + raise InvalidTagNameError("Name is too long.") + name_regex = config.config["tag_name_regex"] if not re.match(name_regex, name): - raise InvalidTagNameError('Name must satisfy regex %r.' % name_regex) + raise InvalidTagNameError("Name must satisfy regex %r." % name_regex) def _get_names(tag: model.Tag) -> List[str]: @@ -54,7 +54,8 @@ def _lower_list(names: List[str]) -> List[str]: def _check_name_intersection( - names1: List[str], names2: List[str], case_sensitive: bool) -> bool: + names1: List[str], names2: List[str], case_sensitive: bool +) -> bool: if not case_sensitive: names1 = _lower_list(names1) names2 = _lower_list(names2) @@ -66,17 +67,19 @@ def sort_tags(tags: List[model.Tag]) -> List[model.Tag]: return sorted( tags, key=lambda tag: ( + tag.category.order, default_category_name == tag.category.name, tag.category.name, - tag.names[0].name) + tag.names[0].name, + ), ) def serialize_relation(tag): return { - 'names': [tag_name.name for tag_name in tag.names], - 'category': tag.category.name, - 'usages': tag.post_count, + "names": [tag_name.name for tag_name in tag.names], + "category": tag.category.name, + "usages": tag.post_count, } @@ -86,16 +89,16 @@ class TagSerializer(serialization.BaseSerializer): def _serializers(self) -> Dict[str, Callable[[], Any]]: return { - 'names': self.serialize_names, - 'category': self.serialize_category, - 'version': self.serialize_version, - 'description': self.serialize_description, - 'creationTime': self.serialize_creation_time, - 'lastEditTime': self.serialize_last_edit_time, - 'usages': self.serialize_usages, - 'suggestions': self.serialize_suggestions, - 'implications': self.serialize_implications, - 'metric': self.serialize_metric, + "names": self.serialize_names, + "category": self.serialize_category, + "version": self.serialize_version, + "description": self.serialize_description, + "creationTime": self.serialize_creation_time, + "lastEditTime": self.serialize_last_edit_time, + "usages": self.serialize_usages, + "suggestions": self.serialize_suggestions, + "implications": self.serialize_implications, + "metric": self.serialize_metric, } def serialize_names(self) -> Any: @@ -122,26 +125,29 @@ class TagSerializer(serialization.BaseSerializer): def serialize_suggestions(self) -> Any: return [ serialize_relation(relation) - for relation in sort_tags(self.tag.suggestions)] + for relation in sort_tags(self.tag.suggestions) + ] def serialize_implications(self) -> Any: return [ serialize_relation(relation) - for relation in sort_tags(self.tag.implications)] + for relation in sort_tags(self.tag.implications) + ] def serialize_metric(self) -> Any: if not self.tag.metric: return None else: return { - 'version': self.tag.metric.version, - 'min': self.tag.metric.min, - 'max': self.tag.metric.max, + "version": self.tag.metric.version, + "min": self.tag.metric.min, + "max": self.tag.metric.max, } def serialize_tag( - tag: model.Tag, options: List[str] = []) -> Optional[rest.Response]: + tag: model.Tag, options: List[str] = [] +) -> Optional[rest.Response]: if not tag: return None return TagSerializer(tag).serialize(options) @@ -149,17 +155,17 @@ def serialize_tag( def try_get_tag_by_name(name: str) -> Optional[model.Tag]: return ( - db.session - .query(model.Tag) + db.session.query(model.Tag) .join(model.TagName) .filter(sa.func.lower(model.TagName.name) == name.lower()) - .one_or_none()) + .one_or_none() + ) def get_tag_by_name(name: str) -> model.Tag: tag = try_get_tag_by_name(name) if not tag: - raise TagNotFoundError('Tag %r not found.' % name) + raise TagNotFoundError("Tag %r not found." % name) return tag @@ -173,12 +179,16 @@ def get_tags_by_names(names: List[str]) -> List[model.Tag]: .filter( sa.sql.or_( sa.func.lower(model.TagName.name) == name.lower() - for name in names)) - .all()) + for name in names + ) + ) + .all() + ) def get_or_create_tags_by_names( - names: List[str]) -> Tuple[List[model.Tag], List[model.Tag]]: + names: List[str], +) -> Tuple[List[model.Tag], List[model.Tag]]: names = util.icase_unique(names) existing_tags = get_tags_by_names(names) new_tags = [] @@ -187,7 +197,8 @@ def get_or_create_tags_by_names( found = False for existing_tag in existing_tags: if _check_name_intersection( - _get_names(existing_tag), [name], False): + _get_names(existing_tag), [name], False + ): found = True break if not found: @@ -195,7 +206,8 @@ def get_or_create_tags_by_names( names=[name], category_name=tag_category_name, suggestions=[], - implications=[]) + implications=[], + ) db.session.add(new_tag) new_tags.append(new_tag) return existing_tags, new_tags @@ -207,8 +219,7 @@ def get_tag_siblings(tag: model.Tag) -> List[model.Tag]: pt_alias1 = sa.orm.aliased(model.PostTag) pt_alias2 = sa.orm.aliased(model.PostTag) result = ( - db.session - .query(tag_alias, sa.func.count(pt_alias2.post_id)) + db.session.query(tag_alias, sa.func.count(pt_alias2.post_id)) .join(pt_alias1, pt_alias1.tag_id == tag_alias.tag_id) .join(pt_alias2, pt_alias2.post_id == pt_alias1.post_id) .filter(pt_alias2.tag_id == tag.tag_id) @@ -216,18 +227,23 @@ def get_tag_siblings(tag: model.Tag) -> List[model.Tag]: .group_by(tag_alias.tag_id) .order_by(sa.func.count(pt_alias2.post_id).desc()) .order_by(tag_alias.first_name) - .limit(50)) + .limit(50) + ) return result def delete(source_tag: model.Tag) -> None: assert source_tag db.session.execute( - sa.sql.expression.delete(model.TagSuggestion) - .where(model.TagSuggestion.child_id == source_tag.tag_id)) + sa.sql.expression.delete(model.TagSuggestion).where( + model.TagSuggestion.child_id == source_tag.tag_id + ) + ) db.session.execute( - sa.sql.expression.delete(model.TagImplication) - .where(model.TagImplication.child_id == source_tag.tag_id)) + sa.sql.expression.delete(model.TagImplication).where( + model.TagImplication.child_id == source_tag.tag_id + ) + ) db.session.delete(source_tag) @@ -235,27 +251,27 @@ def merge_tags(source_tag: model.Tag, target_tag: model.Tag) -> None: assert source_tag assert target_tag if source_tag.tag_id == target_tag.tag_id: - raise InvalidTagRelationError('Cannot merge tag with itself.') + raise InvalidTagRelationError("Cannot merge tag with itself.") if source_tag.metric or target_tag.metric: - raise InvalidTagRelationError('Cannot merge tags with metrics.') + raise InvalidTagRelationError("Cannot merge tags with metrics.") def merge_posts(source_tag_id: int, target_tag_id: int) -> None: alias1 = model.PostTag alias2 = sa.orm.util.aliased(model.PostTag) - update_stmt = ( - sa.sql.expression.update(alias1) - .where(alias1.tag_id == source_tag_id)) - update_stmt = ( - update_stmt - .where( - ~sa.exists() - .where(alias1.post_id == alias2.post_id) - .where(alias2.tag_id == target_tag_id))) + update_stmt = sa.sql.expression.update(alias1).where( + alias1.tag_id == source_tag_id + ) + update_stmt = update_stmt.where( + ~sa.exists() + .where(alias1.post_id == alias2.post_id) + .where(alias2.tag_id == target_tag_id) + ) update_stmt = update_stmt.values(tag_id=target_tag_id) db.session.execute(update_stmt) def merge_relations( - table: model.Base, source_tag_id: int, target_tag_id: int) -> None: + table: model.Base, source_tag_id: int, target_tag_id: int + ) -> None: alias1 = table alias2 = sa.orm.util.aliased(table) update_stmt = ( @@ -265,8 +281,10 @@ def merge_tags(source_tag: model.Tag, target_tag: model.Tag) -> None: .where( ~sa.exists() .where(alias2.child_id == alias1.child_id) - .where(alias2.parent_id == target_tag_id)) - .values(parent_id=target_tag_id)) + .where(alias2.parent_id == target_tag_id) + ) + .values(parent_id=target_tag_id) + ) db.session.execute(update_stmt) update_stmt = ( @@ -276,8 +294,10 @@ def merge_tags(source_tag: model.Tag, target_tag: model.Tag) -> None: .where( ~sa.exists() .where(alias2.parent_id == alias1.parent_id) - .where(alias2.child_id == target_tag_id)) - .values(child_id=target_tag_id)) + .where(alias2.child_id == target_tag_id) + ) + .values(child_id=target_tag_id) + ) db.session.execute(update_stmt) def merge_suggestions(source_tag_id: int, target_tag_id: int) -> None: @@ -293,10 +313,11 @@ def merge_tags(source_tag: model.Tag, target_tag: model.Tag) -> None: def create_tag( - names: List[str], - category_name: str, - suggestions: List[str], - implications: List[str]) -> model.Tag: + names: List[str], + category_name: str, + suggestions: List[str], + implications: List[str], +) -> model.Tag: tag = model.Tag() tag.creation_time = datetime.utcnow() update_tag_names(tag, names) @@ -316,7 +337,7 @@ def update_tag_names(tag: model.Tag, names: List[str]) -> None: assert tag names = util.icase_unique([name for name in names if name]) if not len(names): - raise InvalidTagNameError('At least one name must be specified.') + raise InvalidTagNameError("At least one name must be specified.") for name in names: _verify_name_validity(name) @@ -329,7 +350,8 @@ def update_tag_names(tag: model.Tag, names: List[str]) -> None: existing_tags = db.session.query(model.TagName).filter(expr).all() if len(existing_tags): raise TagAlreadyExistsError( - 'One of names is already used by another tag.') + "One of names is already used by another tag." + ) # remove unwanted items for tag_name in tag.names[:]: @@ -351,7 +373,7 @@ def update_tag_names(tag: model.Tag, names: List[str]) -> None: def update_tag_implications(tag: model.Tag, relations: List[str]) -> None: assert tag if _check_name_intersection(_get_names(tag), relations, False): - raise InvalidTagRelationError('Tag cannot imply itself.') + raise InvalidTagRelationError("Tag cannot imply itself.") tag.implications = get_tags_by_names(relations) @@ -359,12 +381,12 @@ def update_tag_implications(tag: model.Tag, relations: List[str]) -> None: def update_tag_suggestions(tag: model.Tag, relations: List[str]) -> None: assert tag if _check_name_intersection(_get_names(tag), relations, False): - raise InvalidTagRelationError('Tag cannot suggest itself.') + raise InvalidTagRelationError("Tag cannot suggest itself.") tag.suggestions = get_tags_by_names(relations) def update_tag_description(tag: model.Tag, description: str) -> None: assert tag if util.value_exceeds_column_size(description, model.Tag.description): - raise InvalidTagDescriptionError('Description is too long.') + raise InvalidTagDescriptionError("Description is too long.") tag.description = description or None diff --git a/server/szurubooru/func/user_tokens.py b/server/szurubooru/func/user_tokens.py index c0f4bad..8d977e4 100644 --- a/server/szurubooru/func/user_tokens.py +++ b/server/szurubooru/func/user_tokens.py @@ -1,8 +1,10 @@ from datetime import datetime -from typing import Any, Optional, List, Dict, Callable -from pyrfc3339 import parser as rfc3339_parser +from typing import Any, Callable, Dict, List, Optional + import pytz -from szurubooru import db, model, rest, errors +from pyrfc3339 import parser as rfc3339_parser + +from szurubooru import db, errors, model, rest from szurubooru.func import auth, serialization, users, util @@ -16,23 +18,22 @@ class InvalidNoteError(errors.ValidationError): class UserTokenSerializer(serialization.BaseSerializer): def __init__( - self, - user_token: model.UserToken, - auth_user: model.User) -> None: + self, user_token: model.UserToken, auth_user: model.User + ) -> None: self.user_token = user_token self.auth_user = auth_user def _serializers(self) -> Dict[str, Callable[[], Any]]: return { - 'user': self.serialize_user, - 'token': self.serialize_token, - 'note': self.serialize_note, - 'enabled': self.serialize_enabled, - 'expirationTime': self.serialize_expiration_time, - 'creationTime': self.serialize_creation_time, - 'lastEditTime': self.serialize_last_edit_time, - 'lastUsageTime': self.serialize_last_usage_time, - 'version': self.serialize_version, + "user": self.serialize_user, + "token": self.serialize_token, + "note": self.serialize_note, + "enabled": self.serialize_enabled, + "expirationTime": self.serialize_expiration_time, + "creationTime": self.serialize_creation_time, + "lastEditTime": self.serialize_last_edit_time, + "lastUsageTime": self.serialize_last_usage_time, + "version": self.serialize_version, } def serialize_user(self) -> Any: @@ -64,31 +65,31 @@ class UserTokenSerializer(serialization.BaseSerializer): def serialize_user_token( - user_token: Optional[model.UserToken], - auth_user: model.User, - options: List[str] = []) -> Optional[rest.Response]: + user_token: Optional[model.UserToken], + auth_user: model.User, + options: List[str] = [], +) -> Optional[rest.Response]: if not user_token: return None return UserTokenSerializer(user_token, auth_user).serialize(options) -def get_by_user_and_token( - user: model.User, token: str) -> model.UserToken: +def get_by_user_and_token(user: model.User, token: str) -> model.UserToken: return ( - db.session - .query(model.UserToken) + db.session.query(model.UserToken) .filter(model.UserToken.user_id == user.user_id) .filter(model.UserToken.token == token) - .one_or_none()) + .one_or_none() + ) def get_user_tokens(user: model.User) -> List[model.UserToken]: assert user return ( - db.session - .query(model.UserToken) + db.session.query(model.UserToken) .filter(model.UserToken.user_id == user.user_id) - .all()) + .all() + ) def create_user_token(user: model.User, enabled: bool) -> model.UserToken: @@ -103,7 +104,8 @@ def create_user_token(user: model.User, enabled: bool) -> model.UserToken: def update_user_token_enabled( - user_token: model.UserToken, enabled: bool) -> None: + user_token: model.UserToken, enabled: bool +) -> None: assert user_token user_token.enabled = enabled update_user_token_edit_time(user_token) @@ -115,28 +117,30 @@ def update_user_token_edit_time(user_token: model.UserToken) -> None: def update_user_token_expiration_time( - user_token: model.UserToken, expiration_time_str: str) -> None: + user_token: model.UserToken, expiration_time_str: str +) -> None: assert user_token try: expiration_time = rfc3339_parser.parse(expiration_time_str, utc=True) expiration_time = expiration_time.astimezone(pytz.UTC) if expiration_time < datetime.utcnow().replace(tzinfo=pytz.UTC): raise InvalidExpirationError( - 'Expiration cannot happen in the past') + "Expiration cannot happen in the past" + ) user_token.expiration_time = expiration_time update_user_token_edit_time(user_token) except ValueError: raise InvalidExpirationError( - 'Expiration is in an invalid format {}'.format( - expiration_time_str)) + "Expiration is in an invalid format {}".format(expiration_time_str) + ) def update_user_token_note(user_token: model.UserToken, note: str) -> None: assert user_token - note = note.strip() if note is not None else '' + note = note.strip() if note is not None else "" note = None if len(note) == 0 else note if util.value_exceeds_column_size(note, model.UserToken.note): - raise InvalidNoteError('Note is too long.') + raise InvalidNoteError("Note is too long.") user_token.note = note update_user_token_edit_time(user_token) diff --git a/server/szurubooru/func/users.py b/server/szurubooru/func/users.py index e5946dc..5cbe3cc 100644 --- a/server/szurubooru/func/users.py +++ b/server/szurubooru/func/users.py @@ -1,9 +1,11 @@ -from datetime import datetime -from typing import Any, Optional, Union, List, Dict, Callable import re +from datetime import datetime +from typing import Any, Callable, Dict, List, Optional, Union + import sqlalchemy as sa -from szurubooru import config, db, model, errors, rest -from szurubooru.func import auth, util, serialization, files, images + +from szurubooru import config, db, errors, model, rest +from szurubooru.func import auth, files, images, serialization, util class UserNotFoundError(errors.NotFoundError): @@ -35,36 +37,41 @@ class InvalidAvatarError(errors.ValidationError): def get_avatar_path(user_name: str) -> str: - return 'avatars/' + user_name.lower() + '.png' + return "avatars/" + user_name.lower() + ".png" def get_avatar_url(user: model.User) -> str: assert user if user.avatar_style == user.AVATAR_GRAVATAR: assert user.email or user.name - return 'https://gravatar.com/avatar/%s?d=retro&s=%d' % ( + return "https://gravatar.com/avatar/%s?d=retro&s=%d" % ( util.get_md5((user.email or user.name).lower()), - config.config['thumbnails']['avatar_width']) + config.config["thumbnails"]["avatar_width"], + ) assert user.name - return '%s/avatars/%s.png' % ( - config.config['data_url'].rstrip('/'), user.name.lower()) + return "%s/avatars/%s.png" % ( + config.config["data_url"].rstrip("/"), + user.name.lower(), + ) def get_email( - user: model.User, - auth_user: model.User, - force_show_email: bool) -> Union[bool, str]: + user: model.User, auth_user: model.User, force_show_email: bool +) -> Union[bool, str]: assert user assert auth_user - if not force_show_email \ - and auth_user.user_id != user.user_id \ - and not auth.has_privilege(auth_user, 'users:edit:any:email'): + if ( + not force_show_email + and auth_user.user_id != user.user_id + and not auth.has_privilege(auth_user, "users:edit:any:email") + ): return False return user.email def get_liked_post_count( - user: model.User, auth_user: model.User) -> Union[bool, int]: + user: model.User, auth_user: model.User +) -> Union[bool, int]: assert user assert auth_user if auth_user.user_id != user.user_id: @@ -73,7 +80,8 @@ def get_liked_post_count( def get_disliked_post_count( - user: model.User, auth_user: model.User) -> Union[bool, int]: + user: model.User, auth_user: model.User +) -> Union[bool, int]: assert user assert auth_user if auth_user.user_id != user.user_id: @@ -83,29 +91,30 @@ def get_disliked_post_count( class UserSerializer(serialization.BaseSerializer): def __init__( - self, - user: model.User, - auth_user: model.User, - force_show_email: bool = False) -> None: + self, + user: model.User, + auth_user: model.User, + force_show_email: bool = False, + ) -> None: self.user = user self.auth_user = auth_user self.force_show_email = force_show_email def _serializers(self) -> Dict[str, Callable[[], Any]]: return { - 'name': self.serialize_name, - 'creationTime': self.serialize_creation_time, - 'lastLoginTime': self.serialize_last_login_time, - 'version': self.serialize_version, - 'rank': self.serialize_rank, - 'avatarStyle': self.serialize_avatar_style, - 'avatarUrl': self.serialize_avatar_url, - 'commentCount': self.serialize_comment_count, - 'uploadedPostCount': self.serialize_uploaded_post_count, - 'favoritePostCount': self.serialize_favorite_post_count, - 'likedPostCount': self.serialize_liked_post_count, - 'dislikedPostCount': self.serialize_disliked_post_count, - 'email': self.serialize_email, + "name": self.serialize_name, + "creationTime": self.serialize_creation_time, + "lastLoginTime": self.serialize_last_login_time, + "version": self.serialize_version, + "rank": self.serialize_rank, + "avatarStyle": self.serialize_avatar_style, + "avatarUrl": self.serialize_avatar_url, + "commentCount": self.serialize_comment_count, + "uploadedPostCount": self.serialize_uploaded_post_count, + "favoritePostCount": self.serialize_favorite_post_count, + "likedPostCount": self.serialize_liked_post_count, + "dislikedPostCount": self.serialize_disliked_post_count, + "email": self.serialize_email, } def serialize_name(self) -> Any: @@ -149,20 +158,22 @@ class UserSerializer(serialization.BaseSerializer): def serialize_user( - user: Optional[model.User], - auth_user: model.User, - options: List[str] = [], - force_show_email: bool = False) -> Optional[rest.Response]: + user: Optional[model.User], + auth_user: model.User, + options: List[str] = [], + force_show_email: bool = False, +) -> Optional[rest.Response]: if not user: return None return UserSerializer(user, auth_user, force_show_email).serialize(options) def serialize_micro_user( - user: Optional[model.User], - auth_user: model.User) -> Optional[rest.Response]: + user: Optional[model.User], auth_user: model.User +) -> Optional[rest.Response]: return serialize_user( - user, auth_user=auth_user, options=['name', 'avatarUrl']) + user, auth_user=auth_user, options=["name", "avatarUrl"] + ) def get_user_count() -> int: @@ -171,33 +182,34 @@ def get_user_count() -> int: def try_get_user_by_name(name: str) -> Optional[model.User]: return ( - db.session - .query(model.User) + db.session.query(model.User) .filter(sa.func.lower(model.User.name) == sa.func.lower(name)) - .one_or_none()) + .one_or_none() + ) def get_user_by_name(name: str) -> model.User: user = try_get_user_by_name(name) if not user: - raise UserNotFoundError('User %r not found.' % name) + raise UserNotFoundError("User %r not found." % name) return user def try_get_user_by_name_or_email(name_or_email: str) -> Optional[model.User]: return ( - db.session - .query(model.User) + db.session.query(model.User) .filter( - (sa.func.lower(model.User.name) == sa.func.lower(name_or_email)) | - (sa.func.lower(model.User.email) == sa.func.lower(name_or_email))) - .one_or_none()) + (sa.func.lower(model.User.name) == sa.func.lower(name_or_email)) + | (sa.func.lower(model.User.email) == sa.func.lower(name_or_email)) + ) + .one_or_none() + ) def get_user_by_name_or_email(name_or_email: str) -> model.User: user = try_get_user_by_name_or_email(name_or_email) if not user: - raise UserNotFoundError('User %r not found.' % name_or_email) + raise UserNotFoundError("User %r not found." % name_or_email) return user @@ -207,7 +219,7 @@ def create_user(name: str, password: str, email: str) -> model.User: update_user_password(user, password) update_user_email(user, email) if get_user_count() > 0: - user.rank = util.flip(auth.RANK_MAP)[config.config['default_rank']] + user.rank = util.flip(auth.RANK_MAP)[config.config["default_rank"]] else: user.rank = model.User.RANK_ADMINISTRATOR user.creation_time = datetime.utcnow() @@ -218,17 +230,18 @@ def create_user(name: str, password: str, email: str) -> model.User: def update_user_name(user: model.User, name: str) -> None: assert user if not name: - raise InvalidUserNameError('Name cannot be empty.') + raise InvalidUserNameError("Name cannot be empty.") if util.value_exceeds_column_size(name, model.User.name): - raise InvalidUserNameError('User name is too long.') + raise InvalidUserNameError("User name is too long.") name = name.strip() - name_regex = config.config['user_name_regex'] + name_regex = config.config["user_name_regex"] if not re.match(name_regex, name): raise InvalidUserNameError( - 'User name %r must satisfy regex %r.' % (name, name_regex)) + "User name %r must satisfy regex %r." % (name, name_regex) + ) other_user = try_get_user_by_name(name) if other_user and other_user.user_id != user.user_id: - raise UserAlreadyExistsError('User %r already exists.' % name) + raise UserAlreadyExistsError("User %r already exists." % name) if user.name and files.has(get_avatar_path(user.name)): files.move(get_avatar_path(user.name), get_avatar_path(name)) user.name = name @@ -237,14 +250,16 @@ def update_user_name(user: model.User, name: str) -> None: def update_user_password(user: model.User, password: str) -> None: assert user if not password: - raise InvalidPasswordError('Password cannot be empty.') - password_regex = config.config['password_regex'] + raise InvalidPasswordError("Password cannot be empty.") + password_regex = config.config["password_regex"] if not re.match(password_regex, password): raise InvalidPasswordError( - 'Password must satisfy regex %r.' % password_regex) + "Password must satisfy regex %r." % password_regex + ) user.password_salt = auth.create_password() password_hash, revision = auth.get_password_hash( - user.password_salt, password) + user.password_salt, password + ) user.password_hash = password_hash user.password_revision = revision @@ -253,53 +268,56 @@ def update_user_email(user: model.User, email: str) -> None: assert user email = email.strip() if util.value_exceeds_column_size(email, model.User.email): - raise InvalidEmailError('Email is too long.') + raise InvalidEmailError("Email is too long.") if not util.is_valid_email(email): - raise InvalidEmailError('E-mail is invalid.') + raise InvalidEmailError("E-mail is invalid.") user.email = email or None def update_user_rank( - user: model.User, rank: str, auth_user: model.User) -> None: + user: model.User, rank: str, auth_user: model.User +) -> None: assert user if not rank: - raise InvalidRankError('Rank cannot be empty.') + raise InvalidRankError("Rank cannot be empty.") rank = util.flip(auth.RANK_MAP).get(rank.strip(), None) all_ranks = list(auth.RANK_MAP.values()) if not rank: - raise InvalidRankError( - 'Rank can be either of %r.' % all_ranks) + raise InvalidRankError("Rank can be either of %r." % all_ranks) if rank in (model.User.RANK_ANONYMOUS, model.User.RANK_NOBODY): - raise InvalidRankError('Rank %r cannot be used.' % auth.RANK_MAP[rank]) - if all_ranks.index(auth_user.rank) \ - < all_ranks.index(rank) and get_user_count() > 0: - raise errors.AuthError('Trying to set higher rank than your own.') + raise InvalidRankError("Rank %r cannot be used." % auth.RANK_MAP[rank]) + if ( + all_ranks.index(auth_user.rank) < all_ranks.index(rank) + and get_user_count() > 0 + ): + raise errors.AuthError("Trying to set higher rank than your own.") user.rank = rank def update_user_avatar( - user: model.User, - avatar_style: str, - avatar_content: Optional[bytes] = None) -> None: + user: model.User, avatar_style: str, avatar_content: Optional[bytes] = None +) -> None: assert user - if avatar_style == 'gravatar': + if avatar_style == "gravatar": user.avatar_style = user.AVATAR_GRAVATAR - elif avatar_style == 'manual': + elif avatar_style == "manual": user.avatar_style = user.AVATAR_MANUAL - avatar_path = 'avatars/' + user.name.lower() + '.png' + avatar_path = "avatars/" + user.name.lower() + ".png" if not avatar_content: if files.has(avatar_path): return - raise InvalidAvatarError('Avatar content missing.') + raise InvalidAvatarError("Avatar content missing.") image = images.Image(avatar_content) image.resize_fill( - int(config.config['thumbnails']['avatar_width']), - int(config.config['thumbnails']['avatar_height'])) + int(config.config["thumbnails"]["avatar_width"]), + int(config.config["thumbnails"]["avatar_height"]), + ) files.save(avatar_path, image.to_png()) else: raise InvalidAvatarError( - 'Avatar style %r is invalid. Valid avatar styles: %r.' % ( - avatar_style, ['gravatar', 'manual'])) + "Avatar style %r is invalid. Valid avatar styles: %r." + % (avatar_style, ["gravatar", "manual"]) + ) def bump_user_login_time(user: model.User) -> None: @@ -312,7 +330,8 @@ def reset_user_password(user: model.User) -> str: password = auth.create_password() user.password_salt = auth.create_password() password_hash, revision = auth.get_password_hash( - user.password_salt, password) + user.password_salt, password + ) user.password_hash = password_hash user.password_revision = revision return password diff --git a/server/szurubooru/func/util.py b/server/szurubooru/func/util.py index 4638d4b..eacdc2a 100644 --- a/server/szurubooru/func/util.py +++ b/server/szurubooru/func/util.py @@ -1,29 +1,32 @@ -import os import hashlib +import os import re import tempfile -from typing import Any, Optional, Union, Tuple, List, Dict, Generator, TypeVar -from datetime import datetime, timedelta from contextlib import contextmanager -from szurubooru import errors +from datetime import datetime, timedelta +from typing import Any, Dict, Generator, List, Optional, Tuple, TypeVar, Union +from szurubooru import errors -T = TypeVar('T') +T = TypeVar("T") def snake_case_to_lower_camel_case(text: str) -> str: - components = text.split('_') - return components[0].lower() + \ - ''.join(word[0].upper() + word[1:].lower() for word in components[1:]) + components = text.split("_") + return components[0].lower() + "".join( + word[0].upper() + word[1:].lower() for word in components[1:] + ) def snake_case_to_upper_train_case(text: str) -> str: - return '-'.join( - word[0].upper() + word[1:].lower() for word in text.split('_')) + return "-".join( + word[0].upper() + word[1:].lower() for word in text.split("_") + ) def snake_case_to_lower_camel_case_keys( - source: Dict[str, Any]) -> Dict[str, Any]: + source: Dict[str, Any] +) -> Dict[str, Any]: target = {} for key, value in source.items(): target[snake_case_to_lower_camel_case(key)] = value @@ -35,7 +38,7 @@ def create_temp_file(**kwargs: Any) -> Generator: (descriptor, path) = tempfile.mkstemp(**kwargs) os.close(descriptor) try: - with open(path, 'r+b') as handle: + with open(path, "r+b") as handle: yield handle finally: os.remove(path) @@ -65,7 +68,7 @@ def flatten_list(source: List[List[T]]) -> List[T]: def get_md5(source: Union[str, bytes]) -> str: if not isinstance(source, bytes): - source = source.encode('utf-8') + source = source.encode("utf-8") md5 = hashlib.md5() md5.update(source) return md5.hexdigest() @@ -73,7 +76,7 @@ def get_md5(source: Union[str, bytes]) -> str: def get_sha1(source: Union[str, bytes]) -> str: if not isinstance(source, bytes): - source = source.encode('utf-8') + source = source.encode("utf-8") sha1 = hashlib.sha1() sha1.update(source) return sha1.hexdigest() @@ -84,12 +87,13 @@ def flip(source: Dict[Any, Any]) -> Dict[Any, Any]: def is_valid_email(email: Optional[str]) -> bool: - ''' Return whether given email address is valid or empty. ''' - return not email or re.match(r'^[^@]*@[^@]*\.[^@]*$', email) is not None + """ Return whether given email address is valid or empty. """ + return not email or re.match(r"^[^@]*@[^@]*\.[^@]*$", email) is not None + +class dotdict(dict): + """ dot.notation access to dictionary attributes. """ -class dotdict(dict): # pylint: disable=invalid-name - ''' dot.notation access to dictionary attributes. ''' def __getattr__(self, attr: str) -> Any: return self.get(attr) @@ -98,51 +102,54 @@ class dotdict(dict): # pylint: disable=invalid-name def parse_time_range(value: str) -> Tuple[datetime, datetime]: - ''' Return tuple containing min/max time for given text representation. ''' + """ Return tuple containing min/max time for given text representation. """ one_day = timedelta(days=1) one_second = timedelta(seconds=1) almost_one_day = one_day - one_second value = value.lower() if not value: - raise errors.ValidationError('Empty date format.') + raise errors.ValidationError("Empty date format.") - if value == 'today': + if value == "today": now = datetime.utcnow() return ( datetime(now.year, now.month, now.day, 0, 0, 0), - datetime(now.year, now.month, now.day, 0, 0, 0) + almost_one_day + datetime(now.year, now.month, now.day, 0, 0, 0) + almost_one_day, ) - if value == 'yesterday': + if value == "yesterday": now = datetime.utcnow() return ( datetime(now.year, now.month, now.day, 0, 0, 0) - one_day, - datetime(now.year, now.month, now.day, 0, 0, 0) - one_second) + datetime(now.year, now.month, now.day, 0, 0, 0) - one_second, + ) - match = re.match(r'^(\d{4})$', value) + match = re.match(r"^(\d{4})$", value) if match: year = int(match.group(1)) return (datetime(year, 1, 1), datetime(year + 1, 1, 1) - one_second) - match = re.match(r'^(\d{4})-(\d{1,2})$', value) + match = re.match(r"^(\d{4})-(\d{1,2})$", value) if match: year = int(match.group(1)) month = int(match.group(2)) return ( datetime(year, month, 1), - datetime(year, month + 1, 1) - one_second) + datetime(year, month + 1, 1) - one_second, + ) - match = re.match(r'^(\d{4})-(\d{1,2})-(\d{1,2})$', value) + match = re.match(r"^(\d{4})-(\d{1,2})-(\d{1,2})$", value) if match: year = int(match.group(1)) month = int(match.group(2)) day = int(match.group(3)) return ( datetime(year, month, day), - datetime(year, month, day + 1) - one_second) + datetime(year, month, day + 1) - one_second, + ) - raise errors.ValidationError('Invalid date format: %r.' % value) + raise errors.ValidationError("Invalid date format: %r." % value) def icase_unique(source: List[str]) -> List[str]: @@ -172,4 +179,4 @@ def get_column_size(column: Any) -> Optional[int]: def chunks(source_list: List[Any], part_size: int) -> Generator: for i in range(0, len(source_list), part_size): - yield source_list[i:i + part_size] + yield source_list[i : i + part_size] diff --git a/server/szurubooru/func/versions.py b/server/szurubooru/func/versions.py index 6e5a367..790b05b 100644 --- a/server/szurubooru/func/versions.py +++ b/server/szurubooru/func/versions.py @@ -1,16 +1,16 @@ -from szurubooru import errors, rest, model +from szurubooru import errors, model, rest def verify_version( - entity: model.Base, - context: rest.Context, - field_name: str = 'version') -> None: + entity: model.Base, context: rest.Context, field_name: str = "version" +) -> None: actual_version = context.get_param_as_int(field_name) expected_version = entity.version if actual_version != expected_version: raise errors.IntegrityError( - 'Someone else modified this in the meantime. ' + - 'Please try again.') + "Someone else modified this in the meantime. " + + "Please try again." + ) def bump_version(entity: model.Base) -> None: |