diff options
| author | ReAnzu <anzu@reanzu.com> | 2018-02-25 04:44:02 -0600 |
|---|---|---|
| committer | rr- <rr-@sakuya.pl> | 2018-03-25 22:23:29 +0200 |
| commit | 2a69f0193f9e5aa451e0dbcb8c89070f3e49ffa5 (patch) | |
| tree | bda640614ae711e8fef0e0463f118f85ef83ac0e /server | |
| parent | e35e70992736340bff3995ca6f425baf2f431e2f (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')
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) |