summaryrefslogtreecommitdiff
path: root/server
diff options
context:
space:
mode:
authorReAnzu <anzu@reanzu.com>2018-02-25 04:44:02 -0600
committerrr- <rr-@sakuya.pl>2018-03-25 22:23:29 +0200
commit2a69f0193f9e5aa451e0dbcb8c89070f3e49ffa5 (patch)
treebda640614ae711e8fef0e0463f118f85ef83ac0e /server
parente35e70992736340bff3995ca6f425baf2f431e2f (diff)
server/auth: add token authentication
* Users are only authenticated against their password on login, and to retrieve a token * Passwords are wiped from the GUI frontend and cookies after login and token retrieval * Tokens are revoked at the end of the session/logout * If the user chooses the "remember me" option, the token is stored in the cookie * Tokens correctly delete themselves on logout * Tokens can expire at user-specified date * Tokens have their last usage time * Tokens can have user defined descriptions * Users can manage login tokens in their account settings
Diffstat (limited to 'server')
-rw-r--r--server/requirements.txt4
-rw-r--r--server/szurubooru/api/__init__.py1
-rw-r--r--server/szurubooru/api/user_token_api.py83
-rw-r--r--server/szurubooru/func/auth.py28
-rw-r--r--server/szurubooru/func/user_tokens.py146
-rw-r--r--server/szurubooru/func/users.py4
-rw-r--r--server/szurubooru/func/util.py6
-rw-r--r--server/szurubooru/middleware/authenticator.py59
-rw-r--r--server/szurubooru/migrations/versions/a39c7f98a7fa_add_user_token_table.py39
-rw-r--r--server/szurubooru/model/__init__.py2
-rw-r--r--server/szurubooru/model/user.py22
-rw-r--r--server/szurubooru/rest/__init__.py1
-rw-r--r--server/szurubooru/tests/api/test_user_token_creating.py29
-rw-r--r--server/szurubooru/tests/api/test_user_token_deleting.py30
-rw-r--r--server/szurubooru/tests/api/test_user_token_retrieving.py31
-rw-r--r--server/szurubooru/tests/api/test_user_token_updating.py42
-rw-r--r--server/szurubooru/tests/conftest.py25
-rw-r--r--server/szurubooru/tests/func/test_auth.py24
-rw-r--r--server/szurubooru/tests/func/test_user_tokens.py155
-rw-r--r--server/szurubooru/tests/middleware/__init__.py0
-rw-r--r--server/szurubooru/tests/middleware/test_authenticator.py93
-rw-r--r--server/szurubooru/tests/model/test_user_token.py14
22 files changed, 812 insertions, 26 deletions
diff --git a/server/requirements.txt b/server/requirements.txt
index 7cc4786..b11c3b5 100644
--- a/server/requirements.txt
+++ b/server/requirements.txt
@@ -11,4 +11,6 @@ scipy>=0.18.1
elasticsearch>=5.0.0
elasticsearch-dsl>=5.0.0
scikit-image>=0.12
-pynacl>=1.2.1 \ No newline at end of file
+pynacl>=1.2.1
+pytz>=2018.3
+pyRFC3339>=1.0 \ No newline at end of file
diff --git a/server/szurubooru/api/__init__.py b/server/szurubooru/api/__init__.py
index 2a2d5af..0d7f75f 100644
--- a/server/szurubooru/api/__init__.py
+++ b/server/szurubooru/api/__init__.py
@@ -1,5 +1,6 @@
import szurubooru.api.info_api
import szurubooru.api.user_api
+import szurubooru.api.user_token_api
import szurubooru.api.post_api
import szurubooru.api.tag_api
import szurubooru.api.tag_category_api
diff --git a/server/szurubooru/api/user_token_api.py b/server/szurubooru/api/user_token_api.py
new file mode 100644
index 0000000..7739823
--- /dev/null
+++ b/server/szurubooru/api/user_token_api.py
@@ -0,0 +1,83 @@
+from typing import Dict
+from szurubooru import model, rest
+from szurubooru.func import auth, users, user_tokens, serialization, versions
+
+
+def _serialize(
+ ctx: rest.Context, user_token: model.UserToken) -> rest.Response:
+ return user_tokens.serialize_user_token(
+ user_token,
+ ctx.user,
+ options=serialization.get_serialization_options(ctx))
+
+
+@rest.routes.get('/user-tokens/(?P<user_name>[^/]+)/?')
+def get_user_tokens(
+ ctx: rest.Context, params: Dict[str, str] = {}) -> rest.Response:
+ user = users.get_user_by_name(params['user_name'])
+ infix = 'self' if ctx.user.user_id == user.user_id else 'any'
+ auth.verify_privilege(ctx.user, 'user_tokens:list:%s' % infix)
+ user_token_list = user_tokens.get_user_tokens(user)
+ return {
+ 'results': [_serialize(ctx, token) for token in user_token_list]
+ }
+
+
+@rest.routes.post('/user-token/(?P<user_name>[^/]+)/?')
+def create_user_token(
+ ctx: rest.Context, params: Dict[str, str] = {}) -> rest.Response:
+ user = users.get_user_by_name(params['user_name'])
+ infix = 'self' if ctx.user.user_id == user.user_id else 'any'
+ auth.verify_privilege(ctx.user, 'user_tokens:create:%s' % infix)
+ enabled = ctx.get_param_as_bool('enabled', True)
+ user_token = user_tokens.create_user_token(user, enabled)
+ if ctx.has_param('note'):
+ note = ctx.get_param_as_string('note')
+ user_tokens.update_user_token_note(user_token, note)
+ if ctx.has_param('expirationTime'):
+ expiration_time = ctx.get_param_as_string('expirationTime')
+ user_tokens.update_user_token_expiration_time(
+ user_token, expiration_time)
+ ctx.session.add(user_token)
+ ctx.session.commit()
+ return _serialize(ctx, user_token)
+
+
+@rest.routes.put('/user-token/(?P<user_name>[^/]+)/(?P<user_token>[^/]+)/?')
+def update_user_token(
+ ctx: rest.Context, params: Dict[str, str] = {}) -> rest.Response:
+ user = users.get_user_by_name(params['user_name'])
+ infix = 'self' if ctx.user.user_id == user.user_id else 'any'
+ auth.verify_privilege(ctx.user, 'user_tokens:edit:%s' % infix)
+ user_token = user_tokens.get_by_user_and_token(user, params['user_token'])
+ versions.verify_version(user_token, ctx)
+ versions.bump_version(user_token)
+ if ctx.has_param('enabled'):
+ auth.verify_privilege(ctx.user, 'user_tokens:edit:%s' % infix)
+ user_tokens.update_user_token_enabled(
+ user_token, ctx.get_param_as_bool('enabled'))
+ if ctx.has_param('note'):
+ auth.verify_privilege(ctx.user, 'user_tokens:edit:%s' % infix)
+ note = ctx.get_param_as_string('note')
+ user_tokens.update_user_token_note(user_token, note)
+ if ctx.has_param('expirationTime'):
+ auth.verify_privilege(ctx.user, 'user_tokens:edit:%s' % infix)
+ expiration_time = ctx.get_param_as_string('expirationTime')
+ user_tokens.update_user_token_expiration_time(
+ user_token, expiration_time)
+ user_tokens.update_user_token_edit_time(user_token)
+ ctx.session.commit()
+ return _serialize(ctx, user_token)
+
+
+@rest.routes.delete('/user-token/(?P<user_name>[^/]+)/(?P<user_token>[^/]+)/?')
+def delete_user_token(
+ ctx: rest.Context, params: Dict[str, str]) -> rest.Response:
+ user = users.get_user_by_name(params['user_name'])
+ infix = 'self' if ctx.user.user_id == user.user_id else 'any'
+ auth.verify_privilege(ctx.user, 'user_tokens:delete:%s' % infix)
+ user_token = user_tokens.get_by_user_and_token(user, params['user_token'])
+ if user_token is not None:
+ ctx.session.delete(user_token)
+ ctx.session.commit()
+ return {}
diff --git a/server/szurubooru/func/auth.py b/server/szurubooru/func/auth.py
index c9740fe..65be79a 100644
--- a/server/szurubooru/func/auth.py
+++ b/server/szurubooru/func/auth.py
@@ -1,10 +1,12 @@
-from typing import Tuple
+from typing import Tuple, Optional
import hashlib
import random
+import uuid
from collections import OrderedDict
+from datetime import datetime
from nacl import pwhash
from nacl.exceptions import InvalidkeyError
-from szurubooru import config, model, errors, db
+from szurubooru import config, db, model, errors
from szurubooru.func import util
@@ -26,7 +28,8 @@ def get_password_hash(salt: str, password: str) -> Tuple[str, int]:
).decode('utf8'), 3
-def get_sha256_legacy_password_hash(salt: str, password: str) -> Tuple[str, int]:
+def get_sha256_legacy_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'))
@@ -78,6 +81,21 @@ def is_valid_password(user: model.User, password: str) -> bool:
return False
+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()):
+ return False
+ return True
+
+
def has_privilege(user: model.User, privilege_name: str) -> bool:
assert user
all_ranks = list(RANK_MAP.keys())
@@ -102,3 +120,7 @@ def generate_authentication_token(user: model.User) -> str:
digest.update(config.config['secret'].encode('utf8'))
digest.update(user.password_salt.encode('utf8'))
return digest.hexdigest()
+
+
+def generate_authorization_token() -> str:
+ return uuid.uuid4().__str__()
diff --git a/server/szurubooru/func/user_tokens.py b/server/szurubooru/func/user_tokens.py
new file mode 100644
index 0000000..c0f4bad
--- /dev/null
+++ b/server/szurubooru/func/user_tokens.py
@@ -0,0 +1,146 @@
+from datetime import datetime
+from typing import Any, Optional, List, Dict, Callable
+from pyrfc3339 import parser as rfc3339_parser
+import pytz
+from szurubooru import db, model, rest, errors
+from szurubooru.func import auth, serialization, users, util
+
+
+class InvalidExpirationError(errors.ValidationError):
+ pass
+
+
+class InvalidNoteError(errors.ValidationError):
+ pass
+
+
+class UserTokenSerializer(serialization.BaseSerializer):
+ def __init__(
+ 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,
+ }
+
+ def serialize_user(self) -> Any:
+ return users.serialize_micro_user(self.user_token.user, self.auth_user)
+
+ def serialize_creation_time(self) -> Any:
+ return self.user_token.creation_time
+
+ def serialize_last_edit_time(self) -> Any:
+ return self.user_token.last_edit_time
+
+ def serialize_last_usage_time(self) -> Any:
+ return self.user_token.last_usage_time
+
+ def serialize_token(self) -> Any:
+ return self.user_token.token
+
+ def serialize_note(self) -> Any:
+ return self.user_token.note
+
+ def serialize_enabled(self) -> Any:
+ return self.user_token.enabled
+
+ def serialize_expiration_time(self) -> Any:
+ return self.user_token.expiration_time
+
+ def serialize_version(self) -> Any:
+ return self.user_token.version
+
+
+def serialize_user_token(
+ 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:
+ return (
+ db.session
+ .query(model.UserToken)
+ .filter(model.UserToken.user_id == user.user_id)
+ .filter(model.UserToken.token == token)
+ .one_or_none())
+
+
+def get_user_tokens(user: model.User) -> List[model.UserToken]:
+ assert user
+ return (
+ db.session
+ .query(model.UserToken)
+ .filter(model.UserToken.user_id == user.user_id)
+ .all())
+
+
+def create_user_token(user: model.User, enabled: bool) -> model.UserToken:
+ assert user
+ user_token = model.UserToken()
+ user_token.user = user
+ user_token.token = auth.generate_authorization_token()
+ user_token.enabled = enabled
+ user_token.creation_time = datetime.utcnow()
+ user_token.last_usage_time = datetime.utcnow()
+ return user_token
+
+
+def update_user_token_enabled(
+ user_token: model.UserToken, enabled: bool) -> None:
+ assert user_token
+ user_token.enabled = enabled
+ update_user_token_edit_time(user_token)
+
+
+def update_user_token_edit_time(user_token: model.UserToken) -> None:
+ assert user_token
+ user_token.last_edit_time = datetime.utcnow()
+
+
+def update_user_token_expiration_time(
+ 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')
+ 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))
+
+
+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 = None if len(note) == 0 else note
+ if util.value_exceeds_column_size(note, model.UserToken.note):
+ raise InvalidNoteError('Note is too long.')
+ user_token.note = note
+ update_user_token_edit_time(user_token)
+
+
+def bump_usage_time(user_token: model.UserToken) -> None:
+ assert user_token
+ user_token.last_usage_time = datetime.utcnow()
diff --git a/server/szurubooru/func/users.py b/server/szurubooru/func/users.py
index 012debc..e5946dc 100644
--- a/server/szurubooru/func/users.py
+++ b/server/szurubooru/func/users.py
@@ -1,6 +1,6 @@
-import re
-from typing import Any, Optional, Union, List, Dict, Callable
from datetime import datetime
+from typing import Any, Optional, Union, List, Dict, Callable
+import re
import sqlalchemy as sa
from szurubooru import config, db, model, errors, rest
from szurubooru.func import auth, util, serialization, files, images
diff --git a/server/szurubooru/func/util.py b/server/szurubooru/func/util.py
index ba2d4dc..5e82286 100644
--- a/server/szurubooru/func/util.py
+++ b/server/szurubooru/func/util.py
@@ -160,6 +160,12 @@ def value_exceeds_column_size(value: Optional[str], column: Any) -> bool:
return len(value) > max_length
+def get_column_size(column: Any) -> Optional[int]:
+ if not column:
+ return None
+ return column.property.columns[0].type.length
+
+
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]
diff --git a/server/szurubooru/middleware/authenticator.py b/server/szurubooru/middleware/authenticator.py
index 644fe3b..4340ec9 100644
--- a/server/szurubooru/middleware/authenticator.py
+++ b/server/szurubooru/middleware/authenticator.py
@@ -1,11 +1,11 @@
import base64
-from typing import Optional
-from szurubooru import db, model, errors, rest
-from szurubooru.func import auth, users
+from typing import Optional, Tuple
+from szurubooru import model, errors, rest
+from szurubooru.func import auth, users, user_tokens
from szurubooru.rest.errors import HttpBadRequest
-def _authenticate(username: str, password: str) -> model.User:
+def _authenticate_basic_auth(username: str, password: str) -> model.User:
''' Try to authenticate user. Throw AuthError for invalid users. '''
user = users.get_user_by_name(username)
if not auth.is_valid_password(user, password):
@@ -13,34 +13,61 @@ def _authenticate(username: str, password: str) -> model.User:
return user
-def _get_user(ctx: rest.Context) -> Optional[model.User]:
+def _authenticate_token(
+ username: str, token: str) -> Tuple[model.User, model.UserToken]:
+ ''' Try to authenticate user. Throw AuthError for invalid users. '''
+ user = users.get_user_by_name(username)
+ user_token = user_tokens.get_by_user_and_token(user, token)
+ if not auth.is_valid_token(user_token):
+ raise errors.AuthError('Invalid token.')
+ return user, user_token
+
+
+def _get_user(ctx: rest.Context, bump_login: bool) -> Optional[model.User]:
if not ctx.has_header('Authorization'):
return None
+ auth_token = None
+
try:
auth_type, credentials = ctx.get_header('Authorization').split(' ', 1)
- if auth_type.lower() != 'basic':
+ if auth_type.lower() == 'basic':
+ username, password = base64.decodebytes(
+ credentials.encode('ascii')).decode('utf8').split(':', 1)
+ auth_user = _authenticate_basic_auth(username, password)
+ elif auth_type.lower() == 'token':
+ username, token = base64.decodebytes(
+ credentials.encode('ascii')).decode('utf8').split(':', 1)
+ auth_user, auth_token = _authenticate_token(username, token)
+ else:
raise HttpBadRequest(
'ValidationError',
- 'Only basic HTTP authentication is supported.')
- username, password = base64.decodebytes(
- credentials.encode('ascii')).decode('utf8').split(':', 1)
- return _authenticate(username, password)
+ 'Only basic or token HTTP authentication is supported.')
except ValueError as err:
msg = (
- 'Basic authentication header value are not properly formed. '
+ 'Authorization header values are not properly formed. '
'Supplied header {0}. Got error: {1}')
raise HttpBadRequest(
'ValidationError',
msg.format(ctx.get_header('Authorization'), str(err)))
+ if bump_login and auth_user.user_id:
+ users.bump_user_login_time(auth_user)
+ if auth_token is not None:
+ user_tokens.bump_usage_time(auth_token)
+ ctx.session.commit()
+
+ return auth_user
+
-@rest.middleware.pre_hook
def process_request(ctx: rest.Context) -> None:
''' Bind the user to request. Update last login time if needed. '''
- auth_user = _get_user(ctx)
+ bump_login = ctx.get_param_as_bool('bump-login', default=False)
+ auth_user = _get_user(ctx, bump_login)
if auth_user:
ctx.user = auth_user
- if ctx.get_param_as_bool('bump-login', default=False) and ctx.user.user_id:
- users.bump_user_login_time(ctx.user)
- ctx.session.commit()
+
+
+@rest.middleware.pre_hook
+def process_request_hook(ctx: rest.Context) -> None:
+ process_request(ctx)
diff --git a/server/szurubooru/migrations/versions/a39c7f98a7fa_add_user_token_table.py b/server/szurubooru/migrations/versions/a39c7f98a7fa_add_user_token_table.py
new file mode 100644
index 0000000..899eaa7
--- /dev/null
+++ b/server/szurubooru/migrations/versions/a39c7f98a7fa_add_user_token_table.py
@@ -0,0 +1,39 @@
+'''
+Added a user_token table for API authorization
+
+Revision ID: a39c7f98a7fa
+Created at: 2018-02-25 01:31:27.345595
+'''
+
+import sqlalchemy as sa
+from alembic import op
+
+
+revision = 'a39c7f98a7fa'
+down_revision = '9ef1a1643c2a'
+branch_labels = None
+depends_on = None
+
+
+def upgrade():
+ op.create_table(
+ 'user_token',
+ sa.Column('id', sa.Integer(), nullable=False),
+ sa.Column('user_id', sa.Integer(), nullable=False),
+ sa.Column('token', sa.Unicode(length=36), nullable=False),
+ sa.Column('note', sa.Unicode(length=128), nullable=True),
+ sa.Column('enabled', sa.Boolean(), nullable=False),
+ sa.Column('expiration_time', sa.DateTime(), nullable=True),
+ sa.Column('creation_time', sa.DateTime(), nullable=False),
+ sa.Column('last_edit_time', sa.DateTime(), nullable=True),
+ sa.Column('last_usage_time', sa.DateTime(), nullable=True),
+ sa.Column('version', sa.Integer(), nullable=False),
+ sa.ForeignKeyConstraint(['user_id'], ['user.id'], ondelete='CASCADE'),
+ sa.PrimaryKeyConstraint('id'))
+ op.create_index(
+ op.f('ix_user_token_user_id'), 'user_token', ['user_id'], unique=False)
+
+
+def downgrade():
+ op.drop_index(op.f('ix_user_token_user_id'), table_name='user_token')
+ op.drop_table('user_token')
diff --git a/server/szurubooru/model/__init__.py b/server/szurubooru/model/__init__.py
index ad2231c..4892b97 100644
--- a/server/szurubooru/model/__init__.py
+++ b/server/szurubooru/model/__init__.py
@@ -1,5 +1,5 @@
from szurubooru.model.base import Base
-from szurubooru.model.user import User
+from szurubooru.model.user import User, UserToken
from szurubooru.model.tag_category import TagCategory
from szurubooru.model.tag import Tag, TagName, TagSuggestion, TagImplication
from szurubooru.model.post import (
diff --git a/server/szurubooru/model/user.py b/server/szurubooru/model/user.py
index 39c5a91..2d599e8 100644
--- a/server/szurubooru/model/user.py
+++ b/server/szurubooru/model/user.py
@@ -86,3 +86,25 @@ class User(Base):
'version_id_col': version,
'version_id_generator': False,
}
+
+
+class UserToken(Base):
+ __tablename__ = 'user_token'
+
+ user_token_id = sa.Column('id', sa.Integer, primary_key=True)
+ user_id = sa.Column(
+ 'user_id',
+ sa.Integer,
+ sa.ForeignKey('user.id', ondelete='CASCADE'),
+ nullable=False,
+ index=True)
+ token = sa.Column('token', sa.Unicode(36), nullable=False)
+ note = sa.Column('note', sa.Unicode(128), nullable=True)
+ enabled = sa.Column('enabled', sa.Boolean, nullable=False, default=True)
+ expiration_time = sa.Column('expiration_time', sa.DateTime, nullable=True)
+ creation_time = sa.Column('creation_time', sa.DateTime, nullable=False)
+ last_edit_time = sa.Column('last_edit_time', sa.DateTime)
+ last_usage_time = sa.Column('last_usage_time', sa.DateTime)
+ version = sa.Column('version', sa.Integer, default=1, nullable=False)
+
+ user = sa.orm.relationship('User')
diff --git a/server/szurubooru/rest/__init__.py b/server/szurubooru/rest/__init__.py
index 14a3e30..d6b3ef2 100644
--- a/server/szurubooru/rest/__init__.py
+++ b/server/szurubooru/rest/__init__.py
@@ -1,2 +1,3 @@
from szurubooru.rest.app import application
from szurubooru.rest.context import Context, Response
+import szurubooru.rest.routes
diff --git a/server/szurubooru/tests/api/test_user_token_creating.py b/server/szurubooru/tests/api/test_user_token_creating.py
new file mode 100644
index 0000000..f550f63
--- /dev/null
+++ b/server/szurubooru/tests/api/test_user_token_creating.py
@@ -0,0 +1,29 @@
+from unittest.mock import patch
+import pytest
+from szurubooru import api
+from szurubooru.func import user_tokens, users
+
+
+@pytest.fixture(autouse=True)
+def inject_config(config_injector):
+ config_injector({'privileges': {'user_tokens:create:self': 'regular'}})
+
+
+def test_creating_user_token(
+ user_token_factory, context_factory, fake_datetime):
+ user_token = user_token_factory()
+ with patch('szurubooru.func.user_tokens.create_user_token'), \
+ patch('szurubooru.func.user_tokens.serialize_user_token'), \
+ patch('szurubooru.func.users.get_user_by_name'), \
+ fake_datetime('1969-02-12'):
+ users.get_user_by_name.return_value = user_token.user
+ user_tokens.serialize_user_token.return_value = 'serialized user token'
+ user_tokens.create_user_token.return_value = user_token
+ result = api.user_token_api.create_user_token(
+ context_factory(user=user_token.user),
+ {
+ 'user_name': user_token.user.name
+ })
+ assert result == 'serialized user token'
+ user_tokens.create_user_token.assert_called_once_with(
+ user_token.user, True)
diff --git a/server/szurubooru/tests/api/test_user_token_deleting.py b/server/szurubooru/tests/api/test_user_token_deleting.py
new file mode 100644
index 0000000..8534152
--- /dev/null
+++ b/server/szurubooru/tests/api/test_user_token_deleting.py
@@ -0,0 +1,30 @@
+from unittest.mock import patch
+import pytest
+from szurubooru import api, db
+from szurubooru.func import user_tokens, users
+
+
+@pytest.fixture(autouse=True)
+def inject_config(config_injector):
+ config_injector({'privileges': {'user_tokens:delete:self': 'regular'}})
+
+
+def test_deleting_user_token(
+ user_token_factory, context_factory, fake_datetime):
+ user_token = user_token_factory()
+ db.session.add(user_token)
+ db.session.commit()
+ with patch('szurubooru.func.user_tokens.get_by_user_and_token'), \
+ patch('szurubooru.func.users.get_user_by_name'), \
+ fake_datetime('1969-02-12'):
+ users.get_user_by_name.return_value = user_token.user
+ user_tokens.get_by_user_and_token.return_value = user_token
+ result = api.user_token_api.delete_user_token(
+ context_factory(user=user_token.user),
+ {
+ 'user_name': user_token.user.name,
+ 'user_token': user_token.token
+ })
+ assert result == {}
+ user_tokens.get_by_user_and_token.assert_called_once_with(
+ user_token.user, user_token.token)
diff --git a/server/szurubooru/tests/api/test_user_token_retrieving.py b/server/szurubooru/tests/api/test_user_token_retrieving.py
new file mode 100644
index 0000000..01b2534
--- /dev/null
+++ b/server/szurubooru/tests/api/test_user_token_retrieving.py
@@ -0,0 +1,31 @@
+from unittest.mock import patch
+import pytest
+from szurubooru import api
+from szurubooru.func import user_tokens, users
+
+
+@pytest.fixture(autouse=True)
+def inject_config(config_injector):
+ config_injector({'privileges': {'user_tokens:list:self': 'regular'}})
+
+
+def test_retrieving_user_tokens(
+ user_token_factory, context_factory, fake_datetime):
+ user_token1 = user_token_factory()
+ user_token2 = user_token_factory(user=user_token1.user)
+ user_token3 = user_token_factory(user=user_token1.user)
+ with patch('szurubooru.func.user_tokens.get_user_tokens'), \
+ patch('szurubooru.func.user_tokens.serialize_user_token'), \
+ patch('szurubooru.func.users.get_user_by_name'), \
+ fake_datetime('1969-02-12'):
+ users.get_user_by_name.return_value = user_token1.user
+ user_tokens.serialize_user_token.return_value = 'serialized user token'
+ user_tokens.get_user_tokens.return_value = [user_token1, user_token2,
+ user_token3]
+ result = api.user_token_api.get_user_tokens(
+ context_factory(user=user_token1.user),
+ {
+ 'user_name': user_token1.user.name
+ })
+ assert result == {'results': ['serialized user token'] * 3}
+ user_tokens.get_user_tokens.assert_called_once_with(user_token1.user)
diff --git a/server/szurubooru/tests/api/test_user_token_updating.py b/server/szurubooru/tests/api/test_user_token_updating.py
new file mode 100644
index 0000000..bf725a3
--- /dev/null
+++ b/server/szurubooru/tests/api/test_user_token_updating.py
@@ -0,0 +1,42 @@
+from unittest.mock import patch
+import pytest
+from szurubooru import api, db
+from szurubooru.func import user_tokens, users
+
+
+@pytest.fixture(autouse=True)
+def inject_config(config_injector):
+ config_injector({'privileges': {'user_tokens:edit:self': 'regular'}})
+
+
+def test_edit_user_token(user_token_factory, context_factory, fake_datetime):
+ user_token = user_token_factory()
+ db.session.add(user_token)
+ db.session.commit()
+ with patch('szurubooru.func.user_tokens.get_by_user_and_token'), \
+ patch('szurubooru.func.user_tokens.update_user_token_enabled'), \
+ patch('szurubooru.func.user_tokens.update_user_token_edit_time'), \
+ patch('szurubooru.func.user_tokens.serialize_user_token'), \
+ patch('szurubooru.func.users.get_user_by_name'), \
+ fake_datetime('1969-02-12'):
+ users.get_user_by_name.return_value = user_token.user
+ user_tokens.serialize_user_token.return_value = 'serialized user token'
+ user_tokens.get_by_user_and_token.return_value = user_token
+ result = api.user_token_api.update_user_token(
+ context_factory(
+ params={
+ 'version': user_token.version,
+ 'enabled': False,
+ },
+ user=user_token.user),
+ {
+ 'user_name': user_token.user.name,
+ 'user_token': user_token.token
+ })
+ assert result == 'serialized user token'
+ user_tokens.get_by_user_and_token.assert_called_once_with(
+ user_token.user, user_token.token)
+ user_tokens.update_user_token_enabled.assert_called_once_with(
+ user_token, False)
+ user_tokens.update_user_token_edit_time.assert_called_once_with(
+ user_token)
diff --git a/server/szurubooru/tests/conftest.py b/server/szurubooru/tests/conftest.py
index db7806e..f6eeee1 100644
--- a/server/szurubooru/tests/conftest.py
+++ b/server/szurubooru/tests/conftest.py
@@ -93,11 +93,11 @@ def session(query_logger): # pylint: disable=unused-argument
@pytest.fixture
def context_factory(session):
- def factory(params=None, files=None, user=None):
+ def factory(params=None, files=None, user=None, headers=None):
ctx = rest.Context(
method=None,
url=None,
- headers={},
+ headers=headers or {},
params=params or {},
files=files or {})
ctx.session = session
@@ -134,6 +134,27 @@ def user_factory():
@pytest.fixture
+def user_token_factory(user_factory):
+ def factory(
+ user=None,
+ token=None,
+ expiration_time=None,
+ enabled=None,
+ creation_time=None):
+ if user is None:
+ user = user_factory()
+ db.session.add(user)
+ user_token = model.UserToken()
+ user_token.user = user
+ user_token.token = token or 'dummy'
+ user_token.expiration_time = expiration_time
+ user_token.enabled = enabled if enabled is not None else True
+ user_token.creation_time = creation_time or datetime(1997, 1, 1)
+ return user_token
+ return factory
+
+
+@pytest.fixture
def tag_category_factory():
def factory(name=None, color='dummy', default=False):
category = model.TagCategory()
diff --git a/server/szurubooru/tests/func/test_auth.py b/server/szurubooru/tests/func/test_auth.py
index 5d0955a..6dc79bb 100644
--- a/server/szurubooru/tests/func/test_auth.py
+++ b/server/szurubooru/tests/func/test_auth.py
@@ -1,5 +1,6 @@
-from szurubooru.func import auth
+from datetime import datetime, timedelta
import pytest
+from szurubooru.func import auth
@pytest.fixture(autouse=True)
@@ -41,3 +42,24 @@ def test_is_valid_password_auto_upgrades_user_password_hash(user_factory):
assert result is True
assert user.password_hash != hash
assert user.password_revision > revision
+
+
+def test_is_valid_token(user_token_factory):
+ user_token = user_token_factory()
+ assert auth.is_valid_token(user_token)
+
+
+def test_expired_token_is_invalid(user_token_factory):
+ past_expiration = (datetime.utcnow() - timedelta(minutes=30))
+ user_token = user_token_factory(expiration_time=past_expiration)
+ assert not auth.is_valid_token(user_token)
+
+
+def test_disabled_token_is_invalid(user_token_factory):
+ user_token = user_token_factory(enabled=False)
+ assert not auth.is_valid_token(user_token)
+
+
+def test_generate_authorization_token():
+ result = auth.generate_authorization_token()
+ assert result != auth.generate_authorization_token()
diff --git a/server/szurubooru/tests/func/test_user_tokens.py b/server/szurubooru/tests/func/test_user_tokens.py
new file mode 100644
index 0000000..8c3577c
--- /dev/null
+++ b/server/szurubooru/tests/func/test_user_tokens.py
@@ -0,0 +1,155 @@
+from datetime import datetime, timedelta
+from unittest.mock import patch
+import pytest
+import pytz
+import random
+import string
+from szurubooru import db, model
+from szurubooru.func import user_tokens, users, auth, util
+
+
+def test_serialize_user_token(user_token_factory):
+ user_token = user_token_factory()
+ db.session.add(user_token)
+ db.session.flush()
+ with patch('szurubooru.func.users.get_avatar_url'):
+ users.get_avatar_url.return_value = 'https://example.com/avatar.png'
+ result = user_tokens.serialize_user_token(user_token, user_token.user)
+ assert result == {
+ 'creationTime': datetime(1997, 1, 1, 0, 0),
+ 'enabled': True,
+ 'expirationTime': None,
+ 'lastEditTime': None,
+ 'lastUsageTime': None,
+ 'note': None,
+ 'token': 'dummy',
+ 'user': {
+ 'avatarUrl': 'https://example.com/avatar.png',
+ 'name': user_token.user.name},
+ 'version': 1
+ }
+
+
+def test_serialize_user_token_none():
+ result = user_tokens.serialize_user_token(None, None)
+ assert result is None
+
+
+def test_get_by_user_and_token(user_token_factory):
+ user_token = user_token_factory()
+ db.session.add(user_token)
+ db.session.flush()
+ db.session.commit()
+ result = user_tokens.get_by_user_and_token(
+ user_token.user, user_token.token)
+ assert result == user_token
+
+
+def test_get_user_tokens(user_token_factory):
+ user_token1 = user_token_factory()
+ user_token2 = user_token_factory(user=user_token1.user)
+ db.session.add(user_token1)
+ db.session.add(user_token2)
+ db.session.flush()
+ db.session.commit()
+ result = user_tokens.get_user_tokens(user_token1.user)
+ assert result == [user_token1, user_token2]
+
+
+def test_create_user_token(user_factory):
+ user = user_factory()
+ db.session.add(user)
+ db.session.flush()
+ db.session.commit()
+ with patch('szurubooru.func.auth.generate_authorization_token'):
+ auth.generate_authorization_token.return_value = 'test'
+ result = user_tokens.create_user_token(user, True)
+ assert result.token == 'test'
+ assert result.user == user
+
+
+def test_update_user_token_enabled(user_token_factory):
+ user_token = user_token_factory()
+ user_tokens.update_user_token_enabled(user_token, False)
+ assert user_token.enabled is False
+ assert user_token.last_edit_time is not None
+
+
+def test_update_user_token_edit_time(user_token_factory):
+ user_token = user_token_factory()
+ assert user_token.last_edit_time is None
+ user_tokens.update_user_token_edit_time(user_token)
+ assert user_token.last_edit_time is not None
+
+
+def test_update_user_token_note(user_token_factory):
+ user_token = user_token_factory()
+ assert user_token.note is None
+ user_tokens.update_user_token_note(user_token, ' Test Note ')
+ assert user_token.note == 'Test Note'
+ assert user_token.last_edit_time is not None
+
+
+def test_update_user_token_note_input_too_long(user_token_factory):
+ user_token = user_token_factory()
+ assert user_token.note is None
+ note_max_length = util.get_column_size(model.UserToken.note) + 1
+ note = ''.join(
+ random.choice(string.ascii_letters) for _ in range(note_max_length))
+ with pytest.raises(user_tokens.InvalidNoteError):
+ user_tokens.update_user_token_note(user_token, note)
+
+
+def test_update_user_token_expiration_time(user_token_factory):
+ user_token = user_token_factory()
+ assert user_token.expiration_time is None
+ expiration_time_str = (
+ (datetime.utcnow() + timedelta(days=1))
+ .replace(tzinfo=pytz.utc)
+ ).isoformat()
+ user_tokens.update_user_token_expiration_time(
+ user_token, expiration_time_str)
+ assert user_token.expiration_time.isoformat() == expiration_time_str
+ assert user_token.last_edit_time is not None
+
+
+def test_update_user_token_expiration_time_in_past(user_token_factory):
+ user_token = user_token_factory()
+ assert user_token.expiration_time is None
+ expiration_time_str = (
+ (datetime.utcnow() - timedelta(days=1))
+ .replace(tzinfo=pytz.utc)
+ ).isoformat()
+ with pytest.raises(
+ user_tokens.InvalidExpirationError,
+ match='Expiration cannot happen in the past'):
+ user_tokens.update_user_token_expiration_time(
+ user_token, expiration_time_str)
+
+
+@pytest.mark.parametrize('expiration_time_str', [
+ datetime.utcnow().isoformat(),
+ (datetime.utcnow() - timedelta(days=1)).ctime(),
+ '1970/01/01 00:00:01.0000Z',
+ '70/01/01 00:00:01.0000Z',
+ ''.join(random.choice(string.ascii_letters) for _ in range(15)),
+ ''.join(random.choice(string.digits) for _ in range(8))
+])
+def test_update_user_token_expiration_time_invalid_format(
+ expiration_time_str, user_token_factory):
+ user_token = user_token_factory()
+ assert user_token.expiration_time is None
+
+ with pytest.raises(
+ user_tokens.InvalidExpirationError,
+ match='Expiration is in an invalid format %s'
+ % expiration_time_str):
+ user_tokens.update_user_token_expiration_time(
+ user_token, expiration_time_str)
+
+
+def test_bump_usage_time(user_token_factory, fake_datetime):
+ user_token = user_token_factory()
+ with fake_datetime('1997-01-01'):
+ user_tokens.bump_usage_time(user_token)
+ assert user_token.last_usage_time == datetime(1997, 1, 1)
diff --git a/server/szurubooru/tests/middleware/__init__.py b/server/szurubooru/tests/middleware/__init__.py
new file mode 100644
index 0000000..e69de29
--- /dev/null
+++ b/server/szurubooru/tests/middleware/__init__.py
diff --git a/server/szurubooru/tests/middleware/test_authenticator.py b/server/szurubooru/tests/middleware/test_authenticator.py
new file mode 100644
index 0000000..be21a93
--- /dev/null
+++ b/server/szurubooru/tests/middleware/test_authenticator.py
@@ -0,0 +1,93 @@
+from unittest.mock import patch
+import pytest
+from szurubooru import db
+from szurubooru.func import auth, users, user_tokens
+from szurubooru.middleware import authenticator
+from szurubooru.rest import errors
+
+
+def test_process_request_no_header(context_factory):
+ ctx = context_factory()
+ authenticator.process_request(ctx)
+ assert ctx.user.name is None
+
+
+def test_process_request_bump_login(context_factory, user_factory):
+ user = user_factory()
+ db.session.add(user)
+ db.session.flush()
+ ctx = context_factory(
+ headers={
+ 'Authorization': 'Basic dGVzdFVzZXI6dGVzdFRva2Vu'
+ },
+ params={
+ 'bump-login': 'true'
+ })
+ with patch('szurubooru.func.auth.is_valid_password'), \
+ patch('szurubooru.func.users.get_user_by_name'):
+ users.get_user_by_name.return_value = user
+ auth.is_valid_password.return_value = True
+ authenticator.process_request(ctx)
+ assert user.last_login_time is not None
+
+
+def test_process_request_bump_login_with_token(
+ context_factory, user_token_factory):
+ user_token = user_token_factory()
+ db.session.add(user_token)
+ db.session.flush()
+ ctx = context_factory(
+ headers={
+ 'Authorization': 'Token dGVzdFVzZXI6dGVzdFRva2Vu'
+ },
+ params={
+ 'bump-login': 'true'
+ })
+ with patch('szurubooru.func.auth.is_valid_token'), \
+ patch('szurubooru.func.users.get_user_by_name'), \
+ patch('szurubooru.func.user_tokens.get_by_user_and_token'):
+ users.get_user_by_name.return_value = user_token.user
+ user_tokens.get_by_user_and_token.return_value = user_token
+ auth.is_valid_token.return_value = True
+ authenticator.process_request(ctx)
+ assert user_token.user.last_login_time is not None
+ assert user_token.last_usage_time is not None
+
+
+def test_process_request_basic_auth_valid(context_factory, user_factory):
+ user = user_factory()
+ ctx = context_factory(
+ headers={
+ 'Authorization': 'Basic dGVzdFVzZXI6dGVzdFBhc3N3b3Jk'
+ })
+ with patch('szurubooru.func.auth.is_valid_password'), \
+ patch('szurubooru.func.users.get_user_by_name'):
+ users.get_user_by_name.return_value = user
+ auth.is_valid_password.return_value = True
+ authenticator.process_request(ctx)
+ assert ctx.user == user
+
+
+def test_process_request_token_auth_valid(context_factory, user_token_factory):
+ user_token = user_token_factory()
+ ctx = context_factory(
+ headers={
+ 'Authorization': 'Token dGVzdFVzZXI6dGVzdFRva2Vu'
+ })
+ with patch('szurubooru.func.auth.is_valid_token'), \
+ patch('szurubooru.func.users.get_user_by_name'), \
+ patch('szurubooru.func.user_tokens.get_by_user_and_token'):
+ users.get_user_by_name.return_value = user_token.user
+ user_tokens.get_by_user_and_token.return_value = user_token
+ auth.is_valid_token.return_value = True
+ authenticator.process_request(ctx)
+ assert ctx.user == user_token.user
+
+
+def test_process_request_bad_header(context_factory):
+ ctx = context_factory(
+ headers={
+ 'Authorization': 'Secret SuperSecretValue'
+ })
+ with pytest.raises(errors.HttpBadRequest):
+ authenticator.process_request(ctx)
diff --git a/server/szurubooru/tests/model/test_user_token.py b/server/szurubooru/tests/model/test_user_token.py
new file mode 100644
index 0000000..0280082
--- /dev/null
+++ b/server/szurubooru/tests/model/test_user_token.py
@@ -0,0 +1,14 @@
+from datetime import datetime
+from szurubooru import db
+
+
+def test_saving_user_token(user_token_factory):
+ user_token = user_token_factory()
+ db.session.add(user_token)
+ db.session.flush()
+ db.session.refresh(user_token)
+ assert not db.session.dirty
+ assert user_token.user is not None
+ assert user_token.token == 'dummy'
+ assert user_token.enabled is True
+ assert user_token.creation_time == datetime(1997, 1, 1)