Skip to content
Open
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
90 changes: 50 additions & 40 deletions auth_backend/auth_plugins/email.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,10 @@
import hashlib
import logging
import re
from typing import Annotated, Self

from annotated_types import MinLen
from email_validator import EmailNotValidError, validate_email
from event_schema.auth import UserLogin
from fastapi import Depends, Header, HTTPException, Request
from fastapi.background import BackgroundTasks
Expand All @@ -26,59 +28,59 @@
logger = logging.getLogger(__name__)


def check_email(v):
restricted: set[str] = {
'"',
'#',
'&',
"'",
'(',
')',
'*',
',',
'/',
';',
'<',
'>',
'?',
'[',
'\\',
']',
'^',
'`',
'{',
'|',
'}',
'~',
'\n',
'\r',
}
if "@" not in v:
raise ValueError()
if set(v) & restricted:
raise ValueError()
return v
def check_email(v, validate: bool):
if not validate:
return v

if not isinstance(v, str):
raise ValueError("Email must be a string")
if not v or v != v.strip() or any(char.isspace() for char in v):
raise ValueError("Email must not contain leading, trailing, or internal spaces")

if not v.isascii() or v.count("@") != 1:
raise ValueError("Invalid email address")

local_part, domain_part = v.split("@", 1)
if not local_part or not domain_part:
raise ValueError("Invalid email address")
if local_part.startswith(".") or local_part.endswith(".") or ".." in local_part:
raise ValueError("Invalid email address")
if domain_part.startswith(".") or domain_part.endswith(".") or "." not in domain_part:
raise ValueError("Invalid email address")
if any(part == "" for part in domain_part.split(".")):
raise ValueError("Invalid email address")

if not re.fullmatch(r"[A-Za-z0-9.!#$%&'*+/=?^_`{|}~-]+", local_part):
raise ValueError("Invalid email address")
if not re.fullmatch(r"[A-Za-z0-9.-]+", domain_part.split(".")[0]) or not re.fullmatch(
r"[A-Za-z]{2,}", domain_part.split(".")[-1]
):
raise ValueError("Invalid email address")

try:
validated = validate_email(v, check_deliverability=False, allow_smtputf8=False)
except EmailNotValidError as exc:
raise ValueError("Invalid email address") from exc

return validated.normalized


class EmailLogin(Base):
email: Annotated[str, MinLen(1)]
password: Annotated[str, MinLen(1)]
scopes: list[Scope] | None = None
session_name: str | None = None
email_validator = field_validator("email")(check_email)


class EmailRegister(Base):
email: Annotated[str, MinLen(1)]
password: Annotated[str, MinLen(1)]
email_validator = field_validator("email")(check_email)
email_validator = field_validator("email")(lambda v: check_email(v, validate=True))


class EmailChange(Base):
email: Annotated[str, MinLen(1)]

email_validator = field_validator("email")(check_email)


class ResetPassword(Base):
password: Annotated[str, MinLen(1)]
Expand All @@ -95,8 +97,6 @@ def check_passwords_dont_match(self) -> Self:
class RequestResetForgottenPassword(Base):
email: Annotated[str, MinLen(1)]

email_validator = field_validator("email")(check_email)


class ResetForgottenPassword(Base):
new_password: Annotated[str, MinLen(1)]
Expand Down Expand Up @@ -335,7 +335,17 @@ async def _request_reset_email(
"Registration wasn't completed. Try to registrate again and do not forget to approve your email",
"Регистрация не была завершена. Попробуйте зарегистрироваться снова и не забудьте подтвердить почту",
)
if auth_params["email"].value == scheme.email:
same_email: AuthMethod | None = (
AuthMethod.query(session=txn)
.filter(
AuthMethod.user_id == user_session.user_id,
AuthMethod.auth_method == cls.get_name(),
AuthMethod.param == "email",
func.lower(AuthMethod.value) == scheme.email.lower(),
)
.one_or_none()
)
if same_email:
raise HTTPException(
status_code=401,
detail=StatusResponseModel(
Expand Down Expand Up @@ -487,7 +497,7 @@ async def _request_reset_forgotten_password(
.filter(
AuthMethod.auth_method == Email.get_name(),
AuthMethod.param == "email",
AuthMethod.value == schema.email,
func.lower(AuthMethod.value) == schema.email.lower(),
)
.one_or_none()
)
Expand Down
1 change: 1 addition & 0 deletions requirements.txt
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
fastapi
fastapi-sqlalchemy
psycopg2-binary
psycopg[binary]
pydantic
uvicorn
alembic
Expand Down
6 changes: 3 additions & 3 deletions tests/test_routes/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,7 +57,7 @@ def dbsession():

@pytest.fixture()
def user_id(client_auth: TestClient, dbsession):
time = datetime.datetime.utcnow()
time = datetime.datetime.now(datetime.UTC).strftime("%Y%m%d%H%M%S%f")
body = {"email": f"user{time}@example.com", "password": "string"}
client_auth.post("/email/registration", json=body)
db_user: AuthMethod = (
Expand All @@ -77,7 +77,7 @@ def user_id(client_auth: TestClient, dbsession):
@pytest.fixture()
def user(client_auth: TestClient, dbsession):
url = "/email/login"
time = datetime.datetime.utcnow()
time = datetime.datetime.now(datetime.UTC).strftime("%Y%m%d%H%M%S%f")
body = {"email": f"user{time}@example.com", "password": "string", "scopes": []}
response = client_auth.post("/email/registration", json=body)
db_user: AuthMethod = (
Expand Down Expand Up @@ -128,7 +128,7 @@ def group(dbsession, parent_id):
_ids: list[int] = []

def _group(client: TestClient):
time = datetime.datetime.utcnow()
time = datetime.datetime.now(datetime.UTC).strftime("%Y%m%d%H%M%S%f")
body = {"name": f"group{time}", "parent_id": parent_id, "scopes": []}
response = client.post(url="/group", json=body)
nonlocal _ids
Expand Down
22 changes: 21 additions & 1 deletion tests/test_routes/test_change_email.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@ def test_main_scenario(client_auth: TestClient, dbsession: Session, user):
.one()
.value
)
tmp_email = f"changed{datetime.datetime.utcnow()}@mail.com"
tmp_email = f"changed{datetime.datetime.now(datetime.UTC).strftime('%Y%m%d%H%M%S%f')}@mail.com"
response = client_auth.post(f"{url}/request", json={"email": tmp_email}, headers={"Authorization": login["token"]})
assert response.status_code == status.HTTP_200_OK

Expand Down Expand Up @@ -70,6 +70,26 @@ def test_invalid_jsons(client_auth: TestClient, dbsession: Session, user):
assert response.status_code == status.HTTP_403_FORBIDDEN


def test_legacy_email_is_accepted(client_auth: TestClient, user):
login = user["login_json"]
response = client_auth.post(
f"{url}/request",
json={"email": "legacy@localhost"},
headers={"Authorization": login["token"]},
)
assert response.status_code == status.HTTP_200_OK


def test_email_comparison_ignores_case(client_auth: TestClient, user):
body, login = user["body"], user["login_json"]
response = client_auth.post(
f"{url}/request",
json={"email": body["email"].upper()},
headers={"Authorization": login["token"]},
)
assert response.status_code == status.HTTP_401_UNAUTHORIZED


def test_expired_token(client_auth: TestClient, dbsession: Session, user):
user_id, body, login = user["user_id"], user["body"], user["login_json"]
response = client_auth.post("/logout", headers={"Authorization": login['token']})
Expand Down
38 changes: 38 additions & 0 deletions tests/test_routes/test_change_password.py
Original file line number Diff line number Diff line change
Expand Up @@ -140,6 +140,25 @@ def test_no_token(client_auth: TestClient, dbsession: Session, user_id: int):
assert response.status_code == status.HTTP_200_OK


def test_legacy_email_case_insensitive_lookup(client_auth: TestClient, dbsession: Session, user_id: int):
token = (
dbsession.query(AuthMethod)
.filter(
AuthMethod.user_id == user_id, AuthMethod.param == "confirmation_token", AuthMethod.auth_method == "email"
)
.one()
)
response = client_auth.get(f"/email/approve?token={token.value}")
assert response.status_code == status.HTTP_200_OK

auth_params = Email.get_auth_method_params(user_id, session=dbsession)
auth_params["email"].value = "legacy@localhost"
dbsession.flush()

response = client_auth.post(f"{url}/restore", json={"email": "LEGACY@LOCALHOST"})
assert response.status_code == status.HTTP_200_OK


def test_with_token(client_auth: TestClient, dbsession: Session, user):
user_id, body, response = user["user_id"], user["body"], user["login_json"]
auth_token = response["token"]
Expand Down Expand Up @@ -221,3 +240,22 @@ def test_no_token_two_requests(client_auth: TestClient, dbsession: Session, user
)
assert reset_token_2
assert reset_token_1 != reset_token_2


def test_legacy_email_case_insensitive_lookup(client_auth: TestClient, dbsession: Session, user_id: int):
token = (
dbsession.query(AuthMethod)
.filter(
AuthMethod.user_id == user_id, AuthMethod.param == "confirmation_token", AuthMethod.auth_method == "email"
)
.one()
)
response = client_auth.get(f"/email/approve?token={token.value}")
assert response.status_code == status.HTTP_200_OK

auth_params = Email.get_auth_method_params(user_id, session=dbsession)
auth_params["email"].value = "legacy@localhost"
dbsession.flush()

response = client_auth.post(f"{url}/restore", json={"email": "LEGACY@LOCALHOST"})
assert response.status_code == status.HTTP_200_OK
5 changes: 3 additions & 2 deletions tests/test_routes/test_login.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,7 @@
def test_invalid_email(client: TestClient):
body = {"email": "some_string", "password": "string"}
response = client.post(url, json=body)
assert response.status_code == status.HTTP_422_UNPROCESSABLE_ENTITY
assert response.status_code == status.HTTP_401_UNAUTHORIZED


def test_main_scenario(client_auth: TestClient, dbsession: Session, user):
Expand All @@ -28,7 +28,8 @@ def test_main_scenario(client_auth: TestClient, dbsession: Session, user):


def test_incorrect_data(client_auth: TestClient, dbsession: Session):
body1 = {"email": f"user{datetime.datetime.utcnow()}@example.com", "password": "string", "scopes": []}
time = datetime.datetime.now(datetime.UTC).strftime("%Y%m%d%H%M%S%f")
body1 = {"email": f"user{time}@example.com", "password": "string", "scopes": []}
body2 = {"email": "wrong@example.com", "password": "string", "scopes": []}
body3 = {"email": "some@example.com", "password": "strong", "scopes": []}
body4 = {"email": "wrong@example.com", "password": "strong", "scopes": []}
Expand Down
5 changes: 3 additions & 2 deletions tests/test_routes/test_logout.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from datetime import datetime
import datetime

from fastapi.testclient import TestClient
from sqlalchemy.orm import Session
Expand All @@ -10,7 +10,8 @@


def test_main_scenario(client_auth: TestClient, dbsession: Session):
body = {"email": f"user{datetime.utcnow()}@example.com", "password": "string", "scopes": []}
time = datetime.datetime.now(datetime.UTC).strftime("%Y%m%d%H%M%S%f")
body = {"email": f"user{time}@example.com", "password": "string", "scopes": []}
user_response = client_auth.post("/email/registration", json=body)
query = (
dbsession.query(AuthMethod)
Expand Down
5 changes: 3 additions & 2 deletions tests/test_routes/test_oidc.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from datetime import datetime
import datetime

from fastapi.testclient import TestClient
from sqlalchemy.orm import Session
Expand Down Expand Up @@ -39,7 +39,8 @@ def test_jwks(client_auth: TestClient):

def test_token_from_token_ok(client_auth: TestClient, dbsession: Session):
# Подготовка к тесту
body = {"email": f"user{datetime.utcnow()}@example.com", "password": "string", "scopes": []}
time = datetime.datetime.now(datetime.UTC).strftime("%Y%m%d%H%M%S%f")
body = {"email": f"user{time}@example.com", "password": "string", "scopes": []}
user_response = client_auth.post("/email/registration", json=body)
query = (
dbsession.query(AuthMethod)
Expand Down
Loading
Loading