diff --git a/pytest.ini b/pytest.ini new file mode 100644 index 0000000..fb33402 --- /dev/null +++ b/pytest.ini @@ -0,0 +1,4 @@ +[pytest] +testpaths = src +python_files = *_test.py +addopts = -ra -q diff --git a/setup.py b/setup.py index d244a90..2b6aaed 100644 --- a/setup.py +++ b/setup.py @@ -6,5 +6,8 @@ setup( packages=find_packages(), include_package_data=True, package_data={'src': ['templates/**/*.html', 'templates/**/*.txt']}, + exclude_package_data={ + '': ['*_test.py'], + }, scripts=['./src/main.py'], ) diff --git a/shell.nix b/shell.nix index 585c993..e5e7b1b 100644 --- a/shell.nix +++ b/shell.nix @@ -20,6 +20,7 @@ mkShell { python313Packages.ruff python313Packages.python-lsp-server python313Packages.jedi-language-server + python313Packages.pytest ty ]; } diff --git a/src/infra/__init__.py b/src/infra/__init__.py index e69de29..8851384 100644 --- a/src/infra/__init__.py +++ b/src/infra/__init__.py @@ -0,0 +1,27 @@ +from .db import Base, CRUD, CRUDRepo, UpdateMissingEntryError, get_database +from .users import UserModel, UserRepoImpl +from .publication import ( + PublicationEntryModel, + PublicationSequenceModel, + PublicationOrderModel, + PublicationModel, + PublicationRepoImpl, +) +from .subscription import SubscriptionModel, SubscriptionRepoImpl + +__all__ = [ + 'Base', + 'CRUD', + 'CRUDRepo', + 'UpdateMissingEntryError', + 'get_database', + 'UserModel', + 'UserRepoImpl', + 'PublicationEntryModel', + 'PublicationSequenceModel', + 'PublicationOrderModel', + 'PublicationModel', + 'PublicationRepoImpl', + 'SubscriptionModel', + 'SubscriptionRepoImpl', +] diff --git a/src/infra/db_test.py b/src/infra/db_test.py new file mode 100644 index 0000000..e362239 --- /dev/null +++ b/src/infra/db_test.py @@ -0,0 +1,8 @@ +from src.infra import Base +from sqlalchemy import Engine, create_engine + + +def mock_db() -> Engine: + engine = create_engine('sqlite:///:memory:', echo=True) + Base.metadata.create_all(engine) + return engine diff --git a/src/infra/users.py b/src/infra/users.py index a01e4bc..cf4170b 100644 --- a/src/infra/users.py +++ b/src/infra/users.py @@ -18,7 +18,7 @@ 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() + email_confirmed: Mapped[bool] = mapped_column(default=False) password_hash: Mapped[str] = mapped_column(VARCHAR(255)) subscriptions: Mapped[list['SubscriptionModel']] = relationship(back_populates='user') diff --git a/src/infra/users_test.py b/src/infra/users_test.py new file mode 100644 index 0000000..119c135 --- /dev/null +++ b/src/infra/users_test.py @@ -0,0 +1,31 @@ +from src.utils.secret import SecretBox +from src.services.users.data import UserDTO, User +from src.services.users.repo import UserRepo +from src.infra.users import UserRepoImpl +from src.infra.db_test import mock_db +import pytest + +DB = mock_db() + + +@pytest.fixture +def db(): + return DB + + +@pytest.fixture +def repo(db): + return UserRepoImpl(db) + + +@pytest.fixture +def users(repo: UserRepo): + return [ + repo.create(UserDTO('example@example.com', SecretBox('hunter1')), None), + repo.create(UserDTO('example2@example.com', SecretBox('hunter2')), None), + ] + + +def test_getting_a_user_by_email(users: list[User], repo: UserRepo): + assert repo.get_user_by_email('example@example.com') == users[0] + assert repo.get_user_by_email('fred@example.com') is None diff --git a/src/main.py b/src/main.py index aca6e66..f9fb8fd 100644 --- a/src/main.py +++ b/src/main.py @@ -6,8 +6,7 @@ from src.services.email import EmailService, get_email_service from src.services.auth import AuthService from src.services.users.service import UserService from dataclasses import dataclass -from src.infra.db import get_database -from src.infra.users import UserRepoImpl +from src.infra import get_database, UserRepoImpl from waitress import serve from flask import Flask import argparse diff --git a/src/utils/secret.py b/src/utils/secret.py index 1828827..15fdc96 100644 --- a/src/utils/secret.py +++ b/src/utils/secret.py @@ -7,6 +7,12 @@ class SecretBox[T]: def expose_secret(self) -> T: return self._item + def __eq__(self, other): + if not isinstance(other, SecretBox): + return False + + return self._item == other._item + def __str__(self): return ''