Compare commits
No commits in common. "0af1b827aa4000bc07c31253ddc76536bcb56028" and "c1a61730d5a52c294d9f5a8ab9d0f9c7fac5963d" have entirely different histories.
0af1b827aa
...
c1a61730d5
18 changed files with 26 additions and 336 deletions
1
.gitignore
vendored
1
.gitignore
vendored
|
|
@ -4,4 +4,3 @@ __pycache__/*
|
|||
config.json
|
||||
result
|
||||
todo.md
|
||||
*.db
|
||||
|
|
|
|||
|
|
@ -6,10 +6,6 @@
|
|||
}
|
||||
"email": {
|
||||
"bird_api_key": "asdf1234",
|
||||
"sender": "noreply@verifiedsenderdomain.com"
|
||||
},
|
||||
"database": {
|
||||
"url": "sqlite://",
|
||||
"echo": false
|
||||
"sender": "me@example.com"
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ with python313Packages;
|
|||
buildPythonApplication {
|
||||
pname = "cereal";
|
||||
version = "0.0.1";
|
||||
propagatedBuildInputs = [ flask requests waitress sqlalchemy argon2-cffi];
|
||||
propagatedBuildInputs = [ flask requests waitress];
|
||||
src = ./.;
|
||||
pyproject = true;
|
||||
build-system = [setuptools];
|
||||
|
|
|
|||
|
|
@ -4,8 +4,6 @@ let
|
|||
requests
|
||||
flask
|
||||
waitress
|
||||
sqlalchemy
|
||||
argon2-cffi
|
||||
]);
|
||||
in
|
||||
with pkgs;
|
||||
|
|
|
|||
|
|
@ -4,8 +4,7 @@ from typing import Any
|
|||
|
||||
from .logging import Logging
|
||||
from .email import Email
|
||||
from .database import Database
|
||||
from .parse import assert_key_of_type, parse_nested_config
|
||||
from .parse import ConfigParseError, assert_key_of_type
|
||||
|
||||
|
||||
@dataclass
|
||||
|
|
@ -14,7 +13,6 @@ class Config:
|
|||
port: int | None
|
||||
logging: Logging
|
||||
email: Email
|
||||
database: Database
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, config: dict[str, Any]) -> 'Config':
|
||||
|
|
@ -22,13 +20,17 @@ class Config:
|
|||
assert_key_of_type(config, 'email', dict)
|
||||
assert_key_of_type(config, 'host', str)
|
||||
assert_key_of_type(config, 'port', int)
|
||||
log_config = parse_nested_config(config, 'logging', Logging.from_dict)
|
||||
email_config = parse_nested_config(config, 'email', Email.from_dict)
|
||||
db_config = parse_nested_config(config, 'database', Database.from_dict)
|
||||
try:
|
||||
log_config = Logging.from_dict(config['logging'])
|
||||
except ConfigParseError as e:
|
||||
raise ConfigParseError(['logging', *e.keypath], e.issue) from e
|
||||
|
||||
return Config(
|
||||
host=config['host'], port=config['port'], email=email_config, logging=log_config, database=db_config
|
||||
)
|
||||
try:
|
||||
email_config = Email.from_dict(config['email'])
|
||||
except ConfigParseError as e:
|
||||
raise ConfigParseError(['email', *e.keypath], e.issue) from e
|
||||
|
||||
return Config(host=config['host'], port=config['port'], email=email_config, logging=log_config)
|
||||
|
||||
|
||||
def parse_config(filename: str) -> Config:
|
||||
|
|
|
|||
|
|
@ -1,16 +0,0 @@
|
|||
from src.config.parse import assert_key_of_type
|
||||
from typing import Any
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
@dataclass
|
||||
class Database:
|
||||
url: str
|
||||
echo: bool
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, config: dict[str, Any]) -> 'Database':
|
||||
assert_key_of_type(config, 'url', str)
|
||||
assert_key_of_type(config, 'echo', bool)
|
||||
|
||||
return Database(url=config['url'], echo=config['echo'])
|
||||
|
|
@ -1,4 +1,3 @@
|
|||
from src.utils.secret import SecretBox
|
||||
from src.config.parse import assert_key_of_type
|
||||
from typing import Any
|
||||
from dataclasses import dataclass
|
||||
|
|
@ -9,11 +8,11 @@ EmailAddress = str
|
|||
|
||||
@dataclass
|
||||
class Email:
|
||||
bird_api_key: SecretBox[str]
|
||||
bird_api_key: str
|
||||
sender: EmailAddress
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, config: dict[str, Any]) -> 'Email':
|
||||
assert_key_of_type(config, 'bird_api_key', str)
|
||||
assert_key_of_type(config, 'sender', str)
|
||||
return Email(bird_api_key=SecretBox(config['bird_api_key']), sender=config['sender'])
|
||||
return Email(bird_api_key=config['bird_api_key'], sender=config['sender'])
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
from typing import Any, TypeVar, Callable
|
||||
from typing import Any
|
||||
|
||||
|
||||
class ConfigParseError(Exception):
|
||||
|
|
@ -18,14 +18,3 @@ def assert_key_of_type(config: dict[str, Any], key: str, kind: Any):
|
|||
|
||||
if not isinstance(config[key], kind):
|
||||
raise ConfigParseError([key], f'type of "{key}" was {type(config[key])}, expected {kind}')
|
||||
|
||||
|
||||
T = TypeVar('T')
|
||||
|
||||
|
||||
def parse_nested_config(config: dict[str, Any], key: str, parse: Callable[[dict[str, Any]], T]) -> T:
|
||||
assert_key_of_type(config, key, dict)
|
||||
try:
|
||||
return parse(config[key])
|
||||
except ConfigParseError as e:
|
||||
raise ConfigParseError([key, *e.keypath], e.issue)
|
||||
|
|
|
|||
|
|
@ -1,15 +0,0 @@
|
|||
from sqlalchemy.orm import DeclarativeBase
|
||||
from src.config.database import Database as DatabaseConfig
|
||||
from sqlalchemy import create_engine
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
pass
|
||||
|
||||
|
||||
def get_database(config: DatabaseConfig):
|
||||
engine = create_engine(config.url, echo=config.echo)
|
||||
# Ensure that all tables exist in the database
|
||||
Base.metadata.create_all(engine)
|
||||
|
||||
return engine
|
||||
|
|
@ -1,84 +0,0 @@
|
|||
import logging
|
||||
from argon2.exceptions import VerifyMismatchError
|
||||
from src.utils.secret import SecretBox
|
||||
from src.config.email import EmailAddress
|
||||
from sqlalchemy import String, Boolean, VARCHAR, Engine, select
|
||||
from sqlalchemy.orm import mapped_column, Mapped, Session
|
||||
from argon2 import PasswordHasher
|
||||
from src.infra.db import Base
|
||||
from src.services.users import UserRepo, UserDTO, User
|
||||
|
||||
|
||||
class UserModel(Base):
|
||||
__tablename__ = 'user'
|
||||
id: Mapped[int] = mapped_column(primary_key=True)
|
||||
email: Mapped[str] = mapped_column(String(254), unique=True)
|
||||
email_confirmed: Mapped[bool] = mapped_column(Boolean())
|
||||
password_hash: Mapped[str] = mapped_column(VARCHAR(255))
|
||||
|
||||
|
||||
def model_to_user(db_user: UserModel) -> User:
|
||||
return User(
|
||||
id=db_user.id,
|
||||
email=db_user.email,
|
||||
email_confirmed=db_user.email_confirmed,
|
||||
password_hash=SecretBox(db_user.password_hash),
|
||||
)
|
||||
|
||||
|
||||
class UserRepoImpl(UserRepo):
|
||||
def __init__(self, db: Engine):
|
||||
self._db = db
|
||||
self._hasher = PasswordHasher()
|
||||
self._logger = logging.getLogger('UserRepoImpl')
|
||||
|
||||
def create_user(self, user: UserDTO) -> User:
|
||||
pw = self._hasher.hash(user.raw_password.expose_secret())
|
||||
db_user = UserModel(email=user.email, password_hash=pw)
|
||||
with Session(self._db) as session:
|
||||
session.add(db_user)
|
||||
session.commit()
|
||||
|
||||
return model_to_user(db_user)
|
||||
|
||||
def get_user_by_id(self, user_id: int) -> User | None:
|
||||
with Session(self._db) as session:
|
||||
user = session.get(UserModel, user_id)
|
||||
|
||||
if user:
|
||||
return model_to_user(user)
|
||||
|
||||
def get_user_by_email(self, email: EmailAddress) -> User | None:
|
||||
with Session(self._db) as session:
|
||||
user = session.scalar(select(UserModel).where(UserModel.email == email))
|
||||
|
||||
if user:
|
||||
return model_to_user(user)
|
||||
|
||||
def update_user(self, user: User):
|
||||
with Session(self._db) as session:
|
||||
db_user = session.get(UserModel, user.id)
|
||||
if db_user is None:
|
||||
self._logger.warning(f'Attempted to update non-existing user {user.id}')
|
||||
return
|
||||
|
||||
# These fields can be set directly
|
||||
db_user.email_confirmed = user.email_confirmed
|
||||
db_user.password_hash = user.password_hash.expose_secret()
|
||||
|
||||
# If the email is being updated then it should not be considered confirmed.
|
||||
if db_user.email != user.email:
|
||||
db_user.email = user.email
|
||||
db_user.email_confirmed = False
|
||||
|
||||
session.commit()
|
||||
|
||||
def auth_as_user(self, user: UserDTO) -> User | None:
|
||||
full_user = self.get_user_by_email(user.email)
|
||||
if full_user is None:
|
||||
return None
|
||||
|
||||
try:
|
||||
self._hasher.verify(full_user.password_hash.expose_secret(), user.raw_password.expose_secret())
|
||||
except VerifyMismatchError:
|
||||
return None
|
||||
15
src/main.py
15
src/main.py
|
|
@ -1,10 +1,5 @@
|
|||
#!/usr/bin/env python
|
||||
|
||||
from src.notifications.service import make_notification_service
|
||||
from src.services.email import get_service as get_email_service
|
||||
from src.services.users import get_service as get_user_service
|
||||
from src.infra.db import get_database
|
||||
from src.infra.users import UserRepoImpl
|
||||
from waitress import serve
|
||||
from flask import Flask
|
||||
import argparse
|
||||
|
|
@ -20,16 +15,6 @@ def create_app(name: str, config: Config) -> Flask:
|
|||
|
||||
app = Flask(name)
|
||||
|
||||
email_service = get_email_service(config.email)
|
||||
notification_service = make_notification_service(email_service)
|
||||
database = get_database(config.database)
|
||||
user_repo = UserRepoImpl(database)
|
||||
user_service = get_user_service(user_repo)
|
||||
|
||||
app.extensions['email'] = email_service
|
||||
app.extensions['notifications'] = notification_service
|
||||
app.extensions['user_service'] = user_service
|
||||
|
||||
if config.host:
|
||||
app.config['SERVER_NAME'] = config.host
|
||||
|
||||
|
|
|
|||
|
|
@ -1,55 +1,12 @@
|
|||
import abc
|
||||
from typing import Any
|
||||
|
||||
from src.services.email import EmailService
|
||||
from src.config import Config
|
||||
|
||||
from .data import Notification
|
||||
|
||||
# MUSTFIX: Needs real users!
|
||||
import abc
|
||||
|
||||
|
||||
class NotificationSender(abc.ABC):
|
||||
class NotificationService(abc.ABC):
|
||||
"""A service that sends notifications to the user"""
|
||||
|
||||
def __init__(self, name: str):
|
||||
self._name = name
|
||||
|
||||
@property
|
||||
def name(self):
|
||||
return self._name
|
||||
|
||||
@abc.abstractmethod
|
||||
def notify(self, user: Any, notification: Notification):
|
||||
"""Send a notification"""
|
||||
pass
|
||||
|
||||
|
||||
class NotificationService:
|
||||
def __init__(self):
|
||||
self._senders = []
|
||||
|
||||
def register_sender(self, sender: NotificationSender):
|
||||
self._senders.append(sender)
|
||||
|
||||
def notify_user(self, user: Any, notification: Notification):
|
||||
for s in self._senders:
|
||||
if s.name in user.notification_methods:
|
||||
s.notify(user, notification)
|
||||
|
||||
|
||||
class EmailNotifier(NotificationSender):
|
||||
def __init__(self, email_service: EmailService):
|
||||
super().__init__('email')
|
||||
self._email_service = email_service
|
||||
|
||||
def notify(self, user: Any, notification: Notification):
|
||||
# TODO: Convert notification and user information into an EmailDTO for the email service
|
||||
pass
|
||||
|
||||
|
||||
def make_notification_service(email_service: EmailService):
|
||||
service = NotificationService()
|
||||
service.register_sender(EmailNotifier(email_service))
|
||||
|
||||
return service
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import requests
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from email.utils import formatdate, make_msgid
|
||||
import abc
|
||||
|
||||
from ..config.email import Email as EmailConfig, EmailAddress
|
||||
|
|
@ -17,6 +17,10 @@ class EmailDTO:
|
|||
text: str
|
||||
# HTML alternate content. (Optional: Textual content is required)
|
||||
html: str | None = None
|
||||
# A unique ID for identifying this Email
|
||||
message_id: str = field(default_factory=make_msgid)
|
||||
# The date that the email was sent
|
||||
date: str = field(default_factory=lambda: formatdate(localtime=True))
|
||||
|
||||
|
||||
class EmailService:
|
||||
|
|
@ -26,20 +30,12 @@ class EmailService:
|
|||
|
||||
|
||||
class BirdEmailServiceImpl(EmailService):
|
||||
BIRD_API_ENDPOINT = 'https://eu1.platform.bird.com/v1/email/messages'
|
||||
|
||||
def __init__(self, config: EmailConfig):
|
||||
self._config = config
|
||||
|
||||
def send_email(self, email: EmailDTO):
|
||||
payload = {'from': email.sender, 'to': email.to, 'subject': email.subject, 'text': email.text}
|
||||
|
||||
if email.html:
|
||||
payload['html'] = email.html
|
||||
|
||||
headers = {'Authorization': f'Bearer {self._config.bird_api_key.expose_secret()}'}
|
||||
response = requests.post(BirdEmailServiceImpl.BIRD_API_ENDPOINT, json=payload, headers=headers)
|
||||
return response.ok
|
||||
# TODO: Implement Bird Specific Email sending and Error handling
|
||||
pass
|
||||
|
||||
|
||||
def get_service(config: EmailConfig) -> EmailService:
|
||||
|
|
|
|||
|
|
@ -1,102 +0,0 @@
|
|||
from email.headerregistry import Address
|
||||
import abc
|
||||
from dataclasses import dataclass
|
||||
from src.utils.secret import SecretBox
|
||||
from src.config.email import EmailAddress
|
||||
|
||||
|
||||
@dataclass
|
||||
class UserProfile:
|
||||
id: int
|
||||
email: EmailAddress
|
||||
|
||||
|
||||
@dataclass
|
||||
class User:
|
||||
id: int
|
||||
email: EmailAddress
|
||||
email_confirmed: bool
|
||||
password_hash: SecretBox[str]
|
||||
|
||||
def to_profile(self) -> UserProfile:
|
||||
return UserProfile(id=self.id, email=self.email)
|
||||
|
||||
|
||||
@dataclass
|
||||
class UserDTO:
|
||||
email: EmailAddress
|
||||
raw_password: SecretBox[str]
|
||||
|
||||
|
||||
class LoginError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class SignupError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class UserRepo(abc.ABC):
|
||||
@abc.abstractmethod
|
||||
def create_user(self, user: UserDTO) -> User:
|
||||
pass
|
||||
|
||||
@abc.abstractmethod
|
||||
def get_user_by_id(self, user_id: int) -> User | None:
|
||||
pass
|
||||
|
||||
@abc.abstractmethod
|
||||
def get_user_by_email(self, email: EmailAddress) -> User | None:
|
||||
pass
|
||||
|
||||
@abc.abstractmethod
|
||||
def update_user(self, user: User):
|
||||
pass
|
||||
|
||||
@abc.abstractmethod
|
||||
def auth_as_user(self, user: UserDTO) -> User | None:
|
||||
pass
|
||||
|
||||
|
||||
def is_valid_email(email: str) -> bool:
|
||||
try:
|
||||
parsed = Address(addr_spec=email)
|
||||
return bool(parsed.username and parsed.domain and '.' in parsed.domain)
|
||||
except (ValueError, TypeError):
|
||||
return False
|
||||
|
||||
|
||||
def is_valid_password(password: str) -> bool:
|
||||
# TODO: Enforce other or saner rules?
|
||||
long_enough = len(password) > 8
|
||||
short_enough = len(password) < 32
|
||||
has_symbol = any([s in password for s in list('@#$%^&*!?/')])
|
||||
return long_enough and short_enough and has_symbol
|
||||
|
||||
|
||||
class UserService:
|
||||
def __init__(self, repo: UserRepo):
|
||||
self._repo = repo
|
||||
|
||||
def login(self, user: UserDTO) -> UserProfile | None:
|
||||
full_user = self._repo.auth_as_user(user)
|
||||
if full_user:
|
||||
return full_user.to_profile()
|
||||
|
||||
def signup(self, user: UserDTO) -> UserProfile:
|
||||
if not is_valid_email(user.email):
|
||||
raise SignupError(f'{user.email} was not an acceptable email address')
|
||||
|
||||
if not is_valid_password(user.raw_password.expose_secret()):
|
||||
raise SignupError('The given password was not acceptable')
|
||||
|
||||
return self._repo.create_user(user).to_profile()
|
||||
|
||||
def get_user_by_id(self, user_id: int) -> UserProfile | None:
|
||||
user = self._repo.get_user_by_id(user_id)
|
||||
if user:
|
||||
return user.to_profile()
|
||||
|
||||
|
||||
def get_service(repo: UserRepo) -> UserService:
|
||||
return UserService(repo)
|
||||
|
|
@ -1,14 +0,0 @@
|
|||
class SecretBox[T]:
|
||||
"""A helper abstraction for making sure secrets aren't leaked accidentally"""
|
||||
|
||||
def __init__(self, item: T):
|
||||
self._item = item
|
||||
|
||||
def expose_secret(self) -> T:
|
||||
return self._item
|
||||
|
||||
def __str__(self):
|
||||
return '<SECRET>'
|
||||
|
||||
def __repr__(self):
|
||||
return f'SecretBox<{type(self._item)}>'
|
||||
Loading…
Add table
Add a link
Reference in a new issue