aboutsummaryrefslogtreecommitdiff
path: root/server/szurubooru/func/posts.py
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/posts.py
parentdd03540398fcab11bfc2c3cebed2fb7a78592c16 (diff)
parent59452711668b9f7a7aeea3c57301b22882987c2b (diff)
Merge remote-tracking branch 'origin/master' into hunternif
# Conflicts: # client/css/post-content-control.styl # client/css/post-list-view.styl # client/html/post_edit_sidebar.tpl # client/js/controllers/post_list_controller.js # client/js/controllers/post_main_controller.js # client/js/controllers/post_upload_controller.js # client/js/controllers/tag_controller.js # client/js/controllers/user_list_controller.js # client/js/controls/expander_control.js # client/js/controls/post_content_control.js # client/js/controls/post_edit_sidebar_control.js # client/js/controls/post_readonly_sidebar_control.js # client/js/controls/tag_input_control.js # client/js/main.js # client/js/models/abstract_list.js # client/js/models/post.js # client/js/models/post_list.js # client/js/models/settings.js # client/js/models/tag.js # client/js/models/tag_list.js # client/js/tags.js # client/js/util/search.js # client/js/util/touch.js # client/js/util/uri.js # client/js/util/views.js # client/js/views/post_main_view.js # client/js/views/post_upload_view.js # client/js/views/posts_header_view.js # client/js/views/posts_page_view.js # client/js/views/settings_view.js # client/js/views/tag_view.js # client/package-lock.json # client/package.json # server/config.yaml.dist # server/szurubooru/api/__init__.py # server/szurubooru/api/post_api.py # server/szurubooru/api/tag_api.py # server/szurubooru/func/posts.py # server/szurubooru/func/tags.py # server/szurubooru/model/__init__.py # server/szurubooru/model/post.py # server/szurubooru/model/tag.py # server/szurubooru/search/configs/__init__.py # server/szurubooru/search/configs/post_search_config.py # server/szurubooru/search/executor.py # server/szurubooru/tests/api/test_post_retrieving.py # server/szurubooru/tests/api/test_post_updating.py # server/szurubooru/tests/api/test_tag_updating.py # server/szurubooru/tests/conftest.py # server/szurubooru/tests/func/test_posts.py # server/szurubooru/tests/func/test_tags.py # server/szurubooru/tests/search/configs/test_post_search_config.py
Diffstat (limited to 'server/szurubooru/func/posts.py')
-rw-r--r--server/szurubooru/func/posts.py668
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 []

© 2015 - 2026 Jakob L. Kreuze