mirror of
https://github.com/wassname/flask-security.git
synced 2026-09-09 11:22:35 +08:00
Add the ability to specify additional fields on the user model that can be used for logging in.
This commit is contained in:
@@ -84,7 +84,8 @@ _default_config = {
|
|||||||
'EMAIL_SUBJECT_PASSWORDLESS': 'Login instructions',
|
'EMAIL_SUBJECT_PASSWORDLESS': 'Login instructions',
|
||||||
'EMAIL_SUBJECT_PASSWORD_NOTICE': 'Your password has been reset',
|
'EMAIL_SUBJECT_PASSWORD_NOTICE': 'Your password has been reset',
|
||||||
'EMAIL_SUBJECT_PASSWORD_CHANGE_NOTICE': 'Your password has been changed',
|
'EMAIL_SUBJECT_PASSWORD_CHANGE_NOTICE': 'Your password has been changed',
|
||||||
'EMAIL_SUBJECT_PASSWORD_RESET': 'Password reset instructions'
|
'EMAIL_SUBJECT_PASSWORD_RESET': 'Password reset instructions',
|
||||||
|
'USER_IDENTITY_ATTRIBUTES': ['email']
|
||||||
}
|
}
|
||||||
|
|
||||||
#: Default Flask-Security messages
|
#: Default Flask-Security messages
|
||||||
|
|||||||
+28
-18
@@ -9,6 +9,8 @@
|
|||||||
:license: MIT, see LICENSE for more details.
|
:license: MIT, see LICENSE for more details.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from .utils import get_identity_attributes
|
||||||
|
|
||||||
|
|
||||||
class Datastore(object):
|
class Datastore(object):
|
||||||
def __init__(self, db):
|
def __init__(self, db):
|
||||||
@@ -179,14 +181,14 @@ class SQLAlchemyUserDatastore(SQLAlchemyDatastore, UserDatastore):
|
|||||||
SQLAlchemyDatastore.__init__(self, db)
|
SQLAlchemyDatastore.__init__(self, db)
|
||||||
UserDatastore.__init__(self, user_model, role_model)
|
UserDatastore.__init__(self, user_model, role_model)
|
||||||
|
|
||||||
def get_user(self, id_or_email):
|
def get_user(self, identifier):
|
||||||
returned = None
|
if self._is_numeric(identifier):
|
||||||
if self._is_numeric(id_or_email):
|
return self.user_model.query.get(identifier)
|
||||||
returned = self.user_model.query.get(id_or_email)
|
for attr in get_identity_attributes():
|
||||||
if not returned:
|
query = getattr(self.user_model, attr).ilike(identifier)
|
||||||
returned = self.user_model.query.filter(
|
rv = self.user_model.query.filter(query).first()
|
||||||
self.user_model.email.ilike(id_or_email)).first()
|
if rv is not None:
|
||||||
return returned
|
return rv
|
||||||
|
|
||||||
def _is_numeric(self, value):
|
def _is_numeric(self, value):
|
||||||
try:
|
try:
|
||||||
@@ -210,12 +212,18 @@ class MongoEngineUserDatastore(MongoEngineDatastore, UserDatastore):
|
|||||||
MongoEngineDatastore.__init__(self, db)
|
MongoEngineDatastore.__init__(self, db)
|
||||||
UserDatastore.__init__(self, user_model, role_model)
|
UserDatastore.__init__(self, user_model, role_model)
|
||||||
|
|
||||||
def get_user(self, id_or_email):
|
def get_user(self, identifier):
|
||||||
from mongoengine import ValidationError
|
from mongoengine import ValidationError
|
||||||
try:
|
try:
|
||||||
return self.user_model.objects(id=id_or_email).first()
|
return self.user_model.objects(id=identifier).first()
|
||||||
except ValidationError:
|
except ValidationError:
|
||||||
return self.user_model.objects(email__iexact=id_or_email).first()
|
pass
|
||||||
|
for attr in get_identity_attributes():
|
||||||
|
query_key = '%s__iexact' % attr
|
||||||
|
query = {query_key: identifier}
|
||||||
|
rv = self.user_model.objects(**query).first()
|
||||||
|
if rv is not None:
|
||||||
|
return rv
|
||||||
|
|
||||||
def find_user(self, **kwargs):
|
def find_user(self, **kwargs):
|
||||||
try:
|
try:
|
||||||
@@ -254,16 +262,18 @@ class PeeweeUserDatastore(PeeweeDatastore, UserDatastore):
|
|||||||
UserDatastore.__init__(self, user_model, role_model)
|
UserDatastore.__init__(self, user_model, role_model)
|
||||||
self.UserRole = role_link
|
self.UserRole = role_link
|
||||||
|
|
||||||
def get_user(self, id_or_email):
|
def get_user(self, identifier):
|
||||||
try:
|
try:
|
||||||
return self.user_model.get(self.user_model.id == id_or_email)
|
return self.user_model.get(self.user_model.id == identifier)
|
||||||
except ValueError:
|
except ValueError:
|
||||||
pass
|
pass
|
||||||
try:
|
|
||||||
return self.user_model.get(self.user_model.email ** id_or_email)
|
for attr in get_identity_attributes():
|
||||||
except self.user_model.DoesNotExist:
|
column = getattr(self.user_model, attr)
|
||||||
pass
|
try:
|
||||||
return None
|
return self.user_model.get(column ** identifier)
|
||||||
|
except self.user_model.DoesNotExist:
|
||||||
|
pass
|
||||||
|
|
||||||
def find_user(self, **kwargs):
|
def find_user(self, **kwargs):
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -220,6 +220,7 @@ class LoginForm(Form, NextFormMixin):
|
|||||||
self.password.errors.append(get_message('PASSWORD_NOT_PROVIDED')[0])
|
self.password.errors.append(get_message('PASSWORD_NOT_PROVIDED')[0])
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
self.user = _datastore.get_user(self.email.data)
|
self.user = _datastore.get_user(self.email.data)
|
||||||
|
|
||||||
if self.user is None:
|
if self.user is None:
|
||||||
|
|||||||
@@ -279,6 +279,16 @@ def get_token_status(token, serializer, max_age=None):
|
|||||||
return expired, invalid, user
|
return expired, invalid, user
|
||||||
|
|
||||||
|
|
||||||
|
def get_identity_attributes(app=None):
|
||||||
|
app = app or current_app
|
||||||
|
attrs = app.config['SECURITY_USER_IDENTITY_ATTRIBUTES']
|
||||||
|
try:
|
||||||
|
attrs = [f.strip() for f in attrs.split(',')]
|
||||||
|
except AttributeError:
|
||||||
|
pass
|
||||||
|
return attrs
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def capture_passwordless_login_requests():
|
def capture_passwordless_login_requests():
|
||||||
login_requests = []
|
login_requests = []
|
||||||
|
|||||||
@@ -826,3 +826,14 @@ class ConfirmableExtendFormsTest(SecurityTest):
|
|||||||
def test_send_confirmation(self):
|
def test_send_confirmation(self):
|
||||||
r = self._get('/confirm', follow_redirects=True)
|
r = self._get('/confirm', follow_redirects=True)
|
||||||
self.assertIn("My Send Confirmation Email Address Field", r.data)
|
self.assertIn("My Send Confirmation Email Address Field", r.data)
|
||||||
|
|
||||||
|
|
||||||
|
class AdditionalUserIdentityAttributes(SecurityTest):
|
||||||
|
|
||||||
|
AUTH_CONFIG = {
|
||||||
|
'SECURITY_USER_IDENTITY_ATTRIBUTES': ('email', 'username')
|
||||||
|
}
|
||||||
|
|
||||||
|
def test_authenticate(self):
|
||||||
|
r = self.authenticate(email='matt')
|
||||||
|
self.assertIn('Hello matt@lp.com', r.data)
|
||||||
|
|||||||
@@ -128,17 +128,17 @@ def create_roles():
|
|||||||
|
|
||||||
|
|
||||||
def create_users(count=None):
|
def create_users(count=None):
|
||||||
users = [('matt@lp.com', 'password', ['admin'], True),
|
users = [('matt@lp.com', 'matt', 'password', ['admin'], True),
|
||||||
('joe@lp.com', 'password', ['editor'], True),
|
('joe@lp.com', 'joe', 'password', ['editor'], True),
|
||||||
('dave@lp.com', 'password', ['admin', 'editor'], True),
|
('dave@lp.com', 'dave', 'password', ['admin', 'editor'], True),
|
||||||
('jill@lp.com', 'password', ['author'], True),
|
('jill@lp.com', 'jill', 'password', ['author'], True),
|
||||||
('tiya@lp.com', 'password', [], False)]
|
('tiya@lp.com', 'tiya', 'password', [], False)]
|
||||||
count = count or len(users)
|
count = count or len(users)
|
||||||
|
|
||||||
for u in users[:count]:
|
for u in users[:count]:
|
||||||
pw = encrypt_password(u[1])
|
pw = encrypt_password(u[2])
|
||||||
ds.create_user(email=u[0], password=pw,
|
ds.create_user(email=u[0], username=u[1], password=pw,
|
||||||
roles=u[2], active=u[3])
|
roles=u[3], active=u[4])
|
||||||
ds.commit()
|
ds.commit()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -30,6 +30,7 @@ def create_app(config, **kwargs):
|
|||||||
|
|
||||||
class User(db.Document, UserMixin):
|
class User(db.Document, UserMixin):
|
||||||
email = db.StringField(unique=True, max_length=255)
|
email = db.StringField(unique=True, max_length=255)
|
||||||
|
username = db.StringField(max_length=255)
|
||||||
password = db.StringField(required=True, max_length=255)
|
password = db.StringField(required=True, max_length=255)
|
||||||
last_login_at = db.DateTimeField()
|
last_login_at = db.DateTimeField()
|
||||||
current_login_at = db.DateTimeField()
|
current_login_at = db.DateTimeField()
|
||||||
|
|||||||
@@ -29,6 +29,7 @@ def create_app(config, **kwargs):
|
|||||||
|
|
||||||
class User(db.Model, UserMixin):
|
class User(db.Model, UserMixin):
|
||||||
email = TextField()
|
email = TextField()
|
||||||
|
username = TextField()
|
||||||
password = TextField()
|
password = TextField()
|
||||||
last_login_at = DateTimeField(null=True)
|
last_login_at = DateTimeField(null=True)
|
||||||
current_login_at = DateTimeField(null=True)
|
current_login_at = DateTimeField(null=True)
|
||||||
|
|||||||
@@ -32,6 +32,7 @@ def create_app(config, **kwargs):
|
|||||||
class User(db.Model, UserMixin):
|
class User(db.Model, UserMixin):
|
||||||
id = db.Column(db.Integer, primary_key=True)
|
id = db.Column(db.Integer, primary_key=True)
|
||||||
email = db.Column(db.String(255), unique=True)
|
email = db.Column(db.String(255), unique=True)
|
||||||
|
username = db.Column(db.String(255))
|
||||||
password = db.Column(db.String(255))
|
password = db.Column(db.String(255))
|
||||||
last_login_at = db.Column(db.DateTime())
|
last_login_at = db.Column(db.DateTime())
|
||||||
current_login_at = db.Column(db.DateTime())
|
current_login_at = db.Column(db.DateTime())
|
||||||
|
|||||||
Reference in New Issue
Block a user