diff options
| author | rr- | 2016-04-24 16:34:06 +0200 |
|---|---|---|
| committer | rr- | 2016-04-24 16:37:11 +0200 |
| commit | 52f4018bee838868943f87528e188a4e35302799 (patch) | |
| tree | 48ff092298139aeacc3526573346eda1c222f161 /server/szurubooru/func | |
| parent | 8fb536c8f0aa4e4a398f1d1db05151a3dd897c6f (diff) | |
server/comments+posts: add rating
Diffstat (limited to 'server/szurubooru/func')
| -rw-r--r-- | server/szurubooru/func/comments.py | 7 | ||||
| -rw-r--r-- | server/szurubooru/func/posts.py | 5 | ||||
| -rw-r--r-- | server/szurubooru/func/scores.py | 53 | ||||
| -rw-r--r-- | server/szurubooru/func/snapshots.py | 36 | ||||
| -rw-r--r-- | server/szurubooru/func/util.py | 24 |
5 files changed, 92 insertions, 33 deletions
diff --git a/server/szurubooru/func/comments.py b/server/szurubooru/func/comments.py index 576cca2..edd7887 100644 --- a/server/szurubooru/func/comments.py +++ b/server/szurubooru/func/comments.py @@ -1,12 +1,12 @@ import datetime from szurubooru import db, errors -from szurubooru.func import users, posts +from szurubooru.func import users, posts, scores class CommentNotFoundError(errors.NotFoundError): pass class EmptyCommentTextError(errors.ValidationError): pass def serialize_comment(comment, authenticated_user): - return { + ret = { 'id': comment.comment_id, 'user': users.serialize_user(comment.user, authenticated_user), 'post': posts.serialize_post(comment.post, authenticated_user), @@ -14,6 +14,9 @@ def serialize_comment(comment, authenticated_user): 'creationTime': comment.creation_time, 'lastEditTime': comment.last_edit_time, } + if authenticated_user: + ret['ownScore'] = scores.get_score(comment, authenticated_user) + return ret def try_get_comment_by_id(comment_id): return db.session \ diff --git a/server/szurubooru/func/posts.py b/server/szurubooru/func/posts.py index 5d947b2..f9ca093 100644 --- a/server/szurubooru/func/posts.py +++ b/server/szurubooru/func/posts.py @@ -1,7 +1,7 @@ import datetime import sqlalchemy from szurubooru import db, errors -from szurubooru.func import users, snapshots +from szurubooru.func import users, snapshots, scores class PostNotFoundError(errors.NotFoundError): pass class PostAlreadyFeaturedError(errors.ValidationError): pass @@ -36,7 +36,8 @@ def serialize_post(post, authenticated_user): for rel in post.favorited_by], } - # TODO: fetch own score if needed + if authenticated_user: + ret['ownScore'] = scores.get_score(post, authenticated_user) return ret diff --git a/server/szurubooru/func/scores.py b/server/szurubooru/func/scores.py new file mode 100644 index 0000000..62ee773 --- /dev/null +++ b/server/szurubooru/func/scores.py @@ -0,0 +1,53 @@ +import datetime +from szurubooru import db, errors +from szurubooru.func import util + +class InvalidScoreError(errors.ValidationError): pass + +def _get_table_info(entity): + resource_type, _, _ = util.get_resource_info(entity) + if resource_type == 'post': + return db.PostScore, lambda table: table.post_id + elif resource_type == 'comment': + return db.CommentScore, lambda table: table.comment_id + else: + assert False + +def _get_score_entity(entity, user): + table, get_column = _get_table_info(entity) + return db.session \ + .query(table) \ + .filter(get_column(table) == get_column(entity)) \ + .filter(table.user_id == user.user_id) \ + .one_or_none() + +def delete_score(entity, user): + score_entity = _get_score_entity(entity, user) + if score_entity: + db.session.delete(score_entity) + +def get_score(entity, user): + score_entity = _get_score_entity(entity, user) + if score_entity: + return score_entity.score + else: + return 0 + +def set_score(entity, user, score): + if not score: + delete_score(entity, user) + return + if score not in (-1, 1): + raise InvalidScoreError( + 'Score %r is invalid. Valid scores: %r.' % (score, (-1, 1))) + score_entity = _get_score_entity(entity, user) + if score_entity: + score_entity.score = score + else: + table, get_column = _get_table_info(entity) + score_entity = table() + setattr(score_entity, get_column(table).name, get_column(entity)) + score_entity.score = score + score_entity.user = user + score_entity.time = datetime.datetime.now() + db.session.add(score_entity) diff --git a/server/szurubooru/func/snapshots.py b/server/szurubooru/func/snapshots.py index 7399bdb..6543e0f 100644 --- a/server/szurubooru/func/snapshots.py +++ b/server/szurubooru/func/snapshots.py @@ -1,6 +1,6 @@ import datetime -from sqlalchemy.inspection import inspect from szurubooru import db +from szurubooru.func import util def get_tag_snapshot(tag): return { @@ -33,33 +33,11 @@ def get_tag_category_snapshot(category): # pylint: disable=invalid-name serializers = { - 'tag': ( - get_tag_snapshot, - lambda tag: tag.first_name), - 'tag_category': ( - get_tag_category_snapshot, - lambda category: category.name), - 'post': ( - get_post_snapshot, - lambda post: post.post_id), + 'tag': get_tag_snapshot, + 'tag_category': get_tag_category_snapshot, + 'post': get_post_snapshot, } -def get_resource_info(entity): - resource_type = entity.__table__.name - assert resource_type in serializers - - primary_key = inspect(entity).identity - assert primary_key is not None - assert len(primary_key) == 1 - - resource_repr = serializers[resource_type][1](entity) - assert resource_repr - - resource_id = primary_key[0] - assert resource_id - - return (resource_type, resource_id, resource_repr) - def get_previous_snapshot(snapshot): return db.session \ .query(db.Snapshot) \ @@ -71,7 +49,7 @@ def get_previous_snapshot(snapshot): .first() def get_snapshots(entity): - resource_type, resource_id, _ = get_resource_info(entity) + resource_type, resource_id, _ = util.get_resource_info(entity) return db.session \ .query(db.Snapshot) \ .filter(db.Snapshot.resource_type == resource_type) \ @@ -103,7 +81,7 @@ def get_serialized_history(entity): return ret def save(operation, entity, auth_user): - resource_type, resource_id, resource_repr = get_resource_info(entity) + resource_type, resource_id, resource_repr = util.get_resource_info(entity) now = datetime.datetime.now() snapshot = db.Snapshot() @@ -112,7 +90,7 @@ def save(operation, entity, auth_user): snapshot.resource_type = resource_type snapshot.resource_id = resource_id snapshot.resource_repr = resource_repr - snapshot.data = serializers[resource_type][0](entity) + snapshot.data = serializers[resource_type](entity) snapshot.user = auth_user earlier_snapshots = get_snapshots(entity) diff --git a/server/szurubooru/func/util.py b/server/szurubooru/func/util.py index 6c14d96..83d3931 100644 --- a/server/szurubooru/func/util.py +++ b/server/szurubooru/func/util.py @@ -1,7 +1,31 @@ import datetime import re +from sqlalchemy.inspection import inspect from szurubooru.errors import ValidationError +def get_resource_info(entity): + serializers = { + 'tag': lambda tag: tag.first_name, + 'tag_category': lambda category: category.name, + 'comment': lambda comment: comment.comment_id, + 'post': lambda post: post.post_id, + } + + resource_type = entity.__table__.name + assert resource_type in serializers + + primary_key = inspect(entity).identity + assert primary_key is not None + assert len(primary_key) == 1 + + resource_repr = serializers[resource_type](entity) + assert resource_repr + + resource_id = primary_key[0] + assert resource_id + + return (resource_type, resource_id, resource_repr) + def is_valid_email(email): ''' Return whether given email address is valid or empty. ''' return not email or re.match(r'^[^@]*@[^@]*\.[^@]*$', email) |