aboutsummaryrefslogtreecommitdiff
path: root/server/szurubooru/func
diff options
context:
space:
mode:
authorrr-2016-04-24 16:34:06 +0200
committerrr-2016-04-24 16:37:11 +0200
commit52f4018bee838868943f87528e188a4e35302799 (patch)
tree48ff092298139aeacc3526573346eda1c222f161 /server/szurubooru/func
parent8fb536c8f0aa4e4a398f1d1db05151a3dd897c6f (diff)
server/comments+posts: add rating
Diffstat (limited to 'server/szurubooru/func')
-rw-r--r--server/szurubooru/func/comments.py7
-rw-r--r--server/szurubooru/func/posts.py5
-rw-r--r--server/szurubooru/func/scores.py53
-rw-r--r--server/szurubooru/func/snapshots.py36
-rw-r--r--server/szurubooru/func/util.py24
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)

© 2015 - 2026 Jakob L. Kreuze