diff options
| author | Hunternif | 2021-08-04 01:56:00 +0100 |
|---|---|---|
| committer | Hunternif | 2021-08-04 01:56:00 +0100 |
| commit | ca861cdc44ca476cec2949f237bbc7503862ad58 (patch) | |
| tree | 520f918a9a77432bf283b46b4296812eb4332ffe /server/szurubooru/func/posts.py | |
| parent | dd03540398fcab11bfc2c3cebed2fb7a78592c16 (diff) | |
| parent | 59452711668b9f7a7aeea3c57301b22882987c2b (diff) | |
Merge remote-tracking branch 'origin/master' into hunternif
# Conflicts:
# client/css/post-content-control.styl
# client/css/post-list-view.styl
# client/html/post_edit_sidebar.tpl
# client/js/controllers/post_list_controller.js
# client/js/controllers/post_main_controller.js
# client/js/controllers/post_upload_controller.js
# client/js/controllers/tag_controller.js
# client/js/controllers/user_list_controller.js
# client/js/controls/expander_control.js
# client/js/controls/post_content_control.js
# client/js/controls/post_edit_sidebar_control.js
# client/js/controls/post_readonly_sidebar_control.js
# client/js/controls/tag_input_control.js
# client/js/main.js
# client/js/models/abstract_list.js
# client/js/models/post.js
# client/js/models/post_list.js
# client/js/models/settings.js
# client/js/models/tag.js
# client/js/models/tag_list.js
# client/js/tags.js
# client/js/util/search.js
# client/js/util/touch.js
# client/js/util/uri.js
# client/js/util/views.js
# client/js/views/post_main_view.js
# client/js/views/post_upload_view.js
# client/js/views/posts_header_view.js
# client/js/views/posts_page_view.js
# client/js/views/settings_view.js
# client/js/views/tag_view.js
# client/package-lock.json
# client/package.json
# server/config.yaml.dist
# server/szurubooru/api/__init__.py
# server/szurubooru/api/post_api.py
# server/szurubooru/api/tag_api.py
# server/szurubooru/func/posts.py
# server/szurubooru/func/tags.py
# server/szurubooru/model/__init__.py
# server/szurubooru/model/post.py
# server/szurubooru/model/tag.py
# server/szurubooru/search/configs/__init__.py
# server/szurubooru/search/configs/post_search_config.py
# server/szurubooru/search/executor.py
# server/szurubooru/tests/api/test_post_retrieving.py
# server/szurubooru/tests/api/test_post_updating.py
# server/szurubooru/tests/api/test_tag_updating.py
# server/szurubooru/tests/conftest.py
# server/szurubooru/tests/func/test_posts.py
# server/szurubooru/tests/func/test_tags.py
# server/szurubooru/tests/search/configs/test_post_search_config.py
Diffstat (limited to 'server/szurubooru/func/posts.py')
| -rw-r--r-- | server/szurubooru/func/posts.py | 668 |
1 files changed, 408 insertions, 260 deletions
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 [] |