aboutsummaryrefslogtreecommitdiff
path: root/server/szurubooru/func
diff options
context:
space:
mode:
authorHunternif2021-08-04 01:56:00 +0100
committerHunternif2021-08-04 01:56:00 +0100
commitca861cdc44ca476cec2949f237bbc7503862ad58 (patch)
tree520f918a9a77432bf283b46b4296812eb4332ffe /server/szurubooru/func
parentdd03540398fcab11bfc2c3cebed2fb7a78592c16 (diff)
parent59452711668b9f7a7aeea3c57301b22882987c2b (diff)
Merge remote-tracking branch 'origin/master' into hunternif
# Conflicts: # client/css/post-content-control.styl # client/css/post-list-view.styl # client/html/post_edit_sidebar.tpl # client/js/controllers/post_list_controller.js # client/js/controllers/post_main_controller.js # client/js/controllers/post_upload_controller.js # client/js/controllers/tag_controller.js # client/js/controllers/user_list_controller.js # client/js/controls/expander_control.js # client/js/controls/post_content_control.js # client/js/controls/post_edit_sidebar_control.js # client/js/controls/post_readonly_sidebar_control.js # client/js/controls/tag_input_control.js # client/js/main.js # client/js/models/abstract_list.js # client/js/models/post.js # client/js/models/post_list.js # client/js/models/settings.js # client/js/models/tag.js # client/js/models/tag_list.js # client/js/tags.js # client/js/util/search.js # client/js/util/touch.js # client/js/util/uri.js # client/js/util/views.js # client/js/views/post_main_view.js # client/js/views/post_upload_view.js # client/js/views/posts_header_view.js # client/js/views/posts_page_view.js # client/js/views/settings_view.js # client/js/views/tag_view.js # client/package-lock.json # client/package.json # server/config.yaml.dist # server/szurubooru/api/__init__.py # server/szurubooru/api/post_api.py # server/szurubooru/api/tag_api.py # server/szurubooru/func/posts.py # server/szurubooru/func/tags.py # server/szurubooru/model/__init__.py # server/szurubooru/model/post.py # server/szurubooru/model/tag.py # server/szurubooru/search/configs/__init__.py # server/szurubooru/search/configs/post_search_config.py # server/szurubooru/search/executor.py # server/szurubooru/tests/api/test_post_retrieving.py # server/szurubooru/tests/api/test_post_updating.py # server/szurubooru/tests/api/test_tag_updating.py # server/szurubooru/tests/conftest.py # server/szurubooru/tests/func/test_posts.py # server/szurubooru/tests/func/test_tags.py # server/szurubooru/tests/search/configs/test_post_search_config.py
Diffstat (limited to 'server/szurubooru/func')
-rw-r--r--server/szurubooru/func/auth.py99
-rw-r--r--server/szurubooru/func/cache.py11
-rw-r--r--server/szurubooru/func/comments.py43
-rw-r--r--server/szurubooru/func/diff.py26
-rw-r--r--server/szurubooru/func/favorites.py11
-rw-r--r--server/szurubooru/func/file_uploads.py12
-rw-r--r--server/szurubooru/func/files.py9
-rw-r--r--server/szurubooru/func/image_hash.py414
-rw-r--r--server/szurubooru/func/images.py337
-rw-r--r--server/szurubooru/func/mailer.py14
-rw-r--r--server/szurubooru/func/mime.py65
-rw-r--r--server/szurubooru/func/net.py94
-rw-r--r--server/szurubooru/func/pool_categories.py207
-rw-r--r--server/szurubooru/func/pools.py337
-rw-r--r--server/szurubooru/func/posts.py668
-rw-r--r--server/szurubooru/func/scores.py22
-rw-r--r--server/szurubooru/func/serialization.py12
-rw-r--r--server/szurubooru/func/snapshots.py141
-rw-r--r--server/szurubooru/func/tag_categories.py113
-rw-r--r--server/szurubooru/func/tags.py162
-rw-r--r--server/szurubooru/func/user_tokens.py70
-rw-r--r--server/szurubooru/func/users.py185
-rw-r--r--server/szurubooru/func/util.py69
-rw-r--r--server/szurubooru/func/versions.py12
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:

© 2015 - 2026 Jakob L. Kreuze