Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
103 changes: 98 additions & 5 deletions modules/weko-accounts/tests/test_rest.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,8 @@

from flask import json

from weko_accounts.errors import VersionNotFoundRESTError, UserAllreadyLoggedInError, UserNotFoundError, InvalidPasswordError, DisabledUserError
from weko_accounts.errors import VersionNotFoundRESTError, UserAllreadyLoggedInError, \
InvalidCredentialsError, InvalidLoginRequestError, DisabledUserError


# .tox/c1/bin/pytest --cov=weko_accounts tests/test_rest.py::test_WekoLogin_post -vv -s --cov-branch --cov-report=term --basetemp=/code/modules/weko-accounts/.tox/c1/tmp
Expand Down Expand Up @@ -57,8 +58,9 @@ def test_WekoLogin_post(app, client, users_login):
content_type='application/json',
)
res_data = json.loads(res.get_data())
assert res.status_code == UserNotFoundError.code
assert res_data['message'] == UserNotFoundError.description
assert res.status_code == InvalidCredentialsError.code
assert res_data['message'] == InvalidCredentialsError.description
unknown_user_body = res.get_data()

# Invalid password : 403 error
req_json = {
Expand All @@ -71,8 +73,10 @@ def test_WekoLogin_post(app, client, users_login):
content_type='application/json',
)
res_data = json.loads(res.get_data())
assert res.status_code == InvalidPasswordError.code
assert res_data['message'] == InvalidPasswordError.description
assert res.status_code == InvalidCredentialsError.code
assert res_data['message'] == InvalidCredentialsError.description
# Unknown account and wrong password are indistinguishable
assert res.get_data() == unknown_user_body

# Inactive user : 403 error
req_json = {
Expand Down Expand Up @@ -118,6 +122,95 @@ def test_WekoLogin_post(app, client, users_login):
assert res_data['message'] == UserAllreadyLoggedInError.description


# .tox/c1/bin/pytest --cov=weko_accounts tests/test_rest.py::test_WekoLogin_post_invalid_body -vv -s --cov-branch --cov-report=term --basetemp=/code/modules/weko-accounts/.tox/c1/tmp
def test_WekoLogin_post_invalid_body(app, client, users_login):
"""Malformed request bodies are rejected with 400."""
version = 'v1'
bodies = [
None,
[],
{},
{'email': users_login[5]['email']},
{'password': 'dummy'},
{'email': users_login[5]['email'], 'password': 1},
{'email': ['a'], 'password': 'dummy'},
{'email': '', 'password': ''},
]
for body in bodies:
res = client.post(
f'/{version}/login',
data=json.dumps(body),
content_type='application/json',
)
res_data = json.loads(res.get_data())
assert res.status_code == InvalidLoginRequestError.code
assert res_data['message'] == InvalidLoginRequestError.description

# Body that is not JSON at all
res = client.post(
f'/{version}/login',
data='not json',
content_type='application/json',
)
assert res.status_code == InvalidLoginRequestError.code


# .tox/c1/bin/pytest --cov=weko_accounts tests/test_rest.py::test_WekoAccountsREST_limiter -vv -s --cov-branch --cov-report=term --basetemp=/code/modules/weko-accounts/.tox/c1/tmp
def test_WekoAccountsREST_limiter(instance_path):
"""Only the login API of the REST application is rate limited."""
import os

from flask import Flask
from invenio_db import db as db_
from weko_accounts import WekoAccountsREST
from weko_accounts.utils import login_limiter

app_ = Flask('testapi', instance_path=instance_path)
app_.config.update(
WEKO_API_LIMIT_RATE_DEFAULT=['2 per minute'],
# the teardown of the REST blueprint commits the db session.
# sqlite is not used, because invenio-db registers sqlite settings
# on every engine and breaks the other tests
SQLALCHEMY_DATABASE_URI=os.getenv(
'SQLALCHEMY_DATABASE_URI',
'postgresql+psycopg2://invenio:dbpass123@postgresql:5432/wekotest'),
SQLALCHEMY_TRACK_MODIFICATIONS=False,
)
db_.init_app(app_)
# Flask-Limiter registers the limit every time as_view() applies the
# decorator, so the limits of the apps made by the other tests pile up.
# A real process makes the REST application only once.
login_limiter._dynamic_route_limits.clear()
before = len(app_.before_request_funcs.get(None, []))
ext = WekoAccountsREST(app_)

# the shared limiter, whose default limits apply to every endpoint,
# is not initialized on the REST application
assert app_.extensions.get('limiter') is login_limiter
after = len(app_.before_request_funcs.get(None, []))
assert after == before + 1

# Initializing again on the same app does not register the hook twice
ext.init_limiter(app_)
assert len(app_.before_request_funcs.get(None, [])) == after

@app_.route('/other')
def other():
return 'ok'

login_limiter.reset()
with app_.test_client() as c:
# the login API is limited by WEKO_API_LIMIT_RATE_DEFAULT
codes = [c.post('/v1/login', data='x',
content_type='application/json').status_code
for _ in range(3)]
assert codes[:2] == [InvalidLoginRequestError.code] * 2
assert codes[2] == 429

# other endpoints are not limited
assert all(c.get('/other').status_code == 200 for _ in range(5))


# .tox/c1/bin/pytest --cov=weko_accounts tests/test_rest.py::test_WekoLogout_post -vv -s --cov-branch --cov-report=term --basetemp=/code/modules/weko-accounts/.tox/c1/tmp
def test_WekoLogout_post(app, client, users_login):
"""Test WekoLogout.post method."""
Expand Down
14 changes: 14 additions & 0 deletions modules/weko-accounts/weko_accounts/errors.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,20 @@ class InvalidPasswordError(RESTException):
description = 'Invalid password.'


class InvalidCredentialsError(RESTException):
"""Email or password is incorrect."""

code = 403
description = 'Invalid email or password.'


class InvalidLoginRequestError(RESTException):
"""Login request body is malformed."""

code = 400
description = 'Invalid request.'


class DisabledUserError(RESTException):
"""Account is disabled."""

Expand Down
13 changes: 13 additions & 0 deletions modules/weko-accounts/weko_accounts/ext.py
Original file line number Diff line number Diff line change
Expand Up @@ -171,8 +171,21 @@ def init_app(self, app):
blueprint = create_blueprint(app, app.config['WEKO_ACCOUNTS_REST_ENDPOINTS'])
app.register_blueprint(blueprint)
app.extensions['weko_accounts_rest'] = self
self.init_limiter(app)
self.init_unauthorized_handler(app)

def init_limiter(self, app):
"""Initialize rate limiting of the login API.

Only the login API is limited. The shared limiter of
:class:`WekoAccounts` is not used here, because its default limits
would apply to every endpoint of the REST application.

:param app: An instance of :class:`flask.Flask`.
"""
from .utils import login_limiter
login_limiter.init_app(app)

def init_unauthorized_handler(self, app):
"""Return 401 JSON instead of redirecting to the login screen.

Expand Down
32 changes: 22 additions & 10 deletions modules/weko-accounts/weko_accounts/rest.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,15 +25,16 @@
from flask import Blueprint, current_app, jsonify, request, make_response
from flask_login import login_user, logout_user
from flask_security import current_user
from flask_security.utils import verify_password
from flask_security.utils import hash_password, verify_password

from invenio_accounts.models import User
from invenio_db import db
from invenio_rest import ContentNegotiatedMethodView
from weko_logging.activity_logger import UserActivityLogger

from .errors import VersionNotFoundRESTError, UserAllreadyLoggedInError, UserNotFoundError, InvalidPasswordError, DisabledUserError
from .utils import limiter
from .errors import VersionNotFoundRESTError, UserAllreadyLoggedInError, \
InvalidCredentialsError, InvalidLoginRequestError, DisabledUserError
from .utils import limiter, login_limit_value, login_limiter


def create_blueprint(app, endpoints):
Expand Down Expand Up @@ -91,11 +92,14 @@ class WekoLogin(ContentNegotiatedMethodView):

view_name = '{0}_accounts'

# Flask-Limiter matches limits by the name of the view function, so the
# limit is applied to the function made by as_view(), not to post().
decorators = [login_limiter.limit(login_limit_value)]

def __init__(self, *args, **kwargs):
"""Constructor."""
super(WekoLogin, self).__init__(*args, **kwargs)

@limiter.limit('')
def post(self, **kwargs):
"""
Login as weko user.
Expand All @@ -113,21 +117,29 @@ def post(self, **kwargs):

def post_v1(self, **kwargs):

data = request.get_json()
email = data['email']
password = data['password']
data = request.get_json(silent=True)
if not isinstance(data, dict):
raise InvalidLoginRequestError()
email = data.get('email')
password = data.get('password')
if not isinstance(email, str) or not isinstance(password, str) \
or not email or not password:
raise InvalidLoginRequestError()

# Check if user is already logged in
if current_user.is_authenticated:
raise UserAllreadyLoggedInError()

# Get User
user = User.query.filter_by(email=email).first()
if not user:
raise UserNotFoundError()
if not user or not user.password:
# Spend the same hashing cost as a real check so that the
# response does not depend on whether the account exists.
hash_password(password)
raise InvalidCredentialsError()
# Verify password
if not verify_password(password, user.password):
raise InvalidPasswordError()
raise InvalidCredentialsError()

# Check if user is active
if not user.active:
Expand Down
18 changes: 18 additions & 0 deletions modules/weko-accounts/weko_accounts/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,24 @@ def api_view():

"""

login_limiter = Limiter(
app=None,
key_func=lambda: f"{request.endpoint}_{get_remote_addr()}",
default_limits=[],
)
"""Limiter only for the login API of the REST application.

It has no default limits, so only the views decorated with it are limited.
The shared :data:`limiter` is not initialized on the REST application,
because its default limits would apply to every API endpoint.
"""


def login_limit_value():
"""Return the rate limit of the login API from the configuration."""
return ';'.join(current_app.config.get(
'WEKO_API_LIMIT_RATE_DEFAULT', WEKO_API_LIMIT_RATE_DEFAULT))


def get_remote_addr():
"""
Expand Down
Loading