diff --git a/config.example.json b/config.example.json index b97e6a3..9279e34 100644 --- a/config.example.json +++ b/config.example.json @@ -1,6 +1,7 @@ { "host": "0.0.0.0", "port": 8080, + "dev_mode": false, "logging": { "level": "info" } diff --git a/derivation.nix b/derivation.nix index d27b324..342c32e 100644 --- a/derivation.nix +++ b/derivation.nix @@ -3,7 +3,18 @@ with python313Packages; buildPythonApplication { pname = "cereal"; version = "0.0.1"; - propagatedBuildInputs = [ flask requests waitress sqlalchemy argon2-cffi pyjwt cryptography jinja2]; + propagatedBuildInputs = [ + flask + requests + waitress + sqlalchemy + argon2-cffi + pyjwt + cryptography + jinja2 + flask-talisman + flask-cors + ]; src = ./.; pyproject = true; build-system = [setuptools]; diff --git a/shell.nix b/shell.nix index e5e7b1b..11f109f 100644 --- a/shell.nix +++ b/shell.nix @@ -9,6 +9,8 @@ let pyjwt cryptography jinja2 + flask-talisman + flask-cors ]); in with pkgs; diff --git a/src/api/__init__.py b/src/api/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/config/__init__.py b/src/config/__init__.py index c2c9171..1448c56 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -1,7 +1,6 @@ import json from dataclasses import dataclass from typing import Any - from .logging import Logging from .email import Email from .auth import Auth @@ -13,6 +12,7 @@ from .parse import parse_nested_config class Config: host: str port: int | None + dev_mode: bool logging: Logging email: Email database: Database @@ -25,9 +25,17 @@ class Config: db_config = parse_nested_config(config, 'database', Database.from_dict) auth_config = parse_nested_config(config, 'auth', Auth.from_dict) + mode = config.get('dev_mode') + + if mode != True: # noqa: E712 + # Explicitly checking for True and falling back to false for anything else. Only explicitly setting + # dev_mode to true will run in dev mode + mode = False + return Config( host=config.get('host', '0.0.0.0'), port=config.get('port'), + dev_mode=mode, email=email_config, logging=log_config, database=db_config, diff --git a/src/infra/__init__.py b/src/infra/__init__.py index 8851384..ee15f85 100644 --- a/src/infra/__init__.py +++ b/src/infra/__init__.py @@ -1,4 +1,4 @@ -from .db import Base, CRUD, CRUDRepo, UpdateMissingEntryError, get_database +from .db import Base, CRUDRepoImpl, UpdateMissingEntryError, get_database from .users import UserModel, UserRepoImpl from .publication import ( PublicationEntryModel, @@ -11,7 +11,6 @@ from .subscription import SubscriptionModel, SubscriptionRepoImpl __all__ = [ 'Base', - 'CRUD', 'CRUDRepo', 'UpdateMissingEntryError', 'get_database', diff --git a/src/infra/db.py b/src/infra/db.py index 6359df4..ad9ff98 100644 --- a/src/infra/db.py +++ b/src/infra/db.py @@ -1,8 +1,8 @@ +from src.services.repo import CRUD from typing import Callable from sqlalchemy.orm import DeclarativeBase, Session from src.config.database import Database as DatabaseConfig from sqlalchemy import create_engine, Engine -import abc class Base(DeclarativeBase): @@ -22,25 +22,7 @@ class UpdateMissingEntryError(Exception): super().__init__(f'Attempted to update non-existing {model_name} {id}') -class CRUD[T, C](abc.ABC): - @abc.abstractmethod - def create(self, domain_item: C, verify: Callable[[T], None] | None) -> T: - pass - - @abc.abstractmethod - def get(self, item_id: int) -> T | None: - pass - - @abc.abstractmethod - def update(self, item_id: int, transform: Callable[[T], T]) -> T: - pass - - @abc.abstractmethod - def delete(self, item_id: int, verify: Callable[[T], None] | None): - pass - - -class CRUDRepo[T: Base, U, C](CRUD[U, C]): +class CRUDRepoImpl[T: Base, U, C](CRUD[U, C]): def __init__( self, db: Engine, diff --git a/src/infra/db_tests.py b/src/infra/db_tests.py index e764924..01ecf4f 100644 --- a/src/infra/db_tests.py +++ b/src/infra/db_tests.py @@ -1,12 +1,12 @@ from sqlalchemy import Engine from typing import Callable -from src.infra import CRUDRepo, Base +from src.infra import CRUDRepoImpl, Base import pytest class CRUDRepoHelper[A: Base, B, C]: @pytest.fixture - def repo(self, db: Engine) -> CRUDRepo[A, B, C]: + def repo(self, db: Engine) -> CRUDRepoImpl[A, B, C]: import pdb pdb.set_trace() @@ -17,7 +17,7 @@ class CRUDRepoHelper[A: Base, B, C]: raise NotImplementedError() @pytest.fixture - def item(self, repo: CRUDRepo[A, B, C], dto: C) -> B: + def item(self, repo: CRUDRepoImpl[A, B, C], dto: C) -> B: return repo.create(dto, None) @pytest.fixture @@ -31,26 +31,28 @@ class CRUDRepoHelper[A: Base, B, C]: def test_create(self, item: B): assert item is not None - def test_create_returns_none_if_invalid(self, repo: CRUDRepo[A, B, C], dto: C): + def test_create_returns_none_if_invalid(self, repo: CRUDRepoImpl[A, B, C], dto: C): def always_invalid(*_args, **_kw): raise ValueError('test') with pytest.raises(ValueError): repo.create(dto, always_invalid) - def test_create_returns_a_value_if_valid(self, repo: CRUDRepo[A, B, C], dto: C): + def test_create_returns_a_value_if_valid(self, repo: CRUDRepoImpl[A, B, C], dto: C): assert repo.create(dto, lambda _: None) is not None - def test_get_returns_none_when_missing(self, repo: CRUDRepo[A, B, C]): + def test_get_returns_none_when_missing(self, repo: CRUDRepoImpl[A, B, C]): assert repo.get(42) is None - def test_get_returns_a_value_if_it_exists(self, repo: CRUDRepo[A, B, C], item: B): + def test_get_returns_a_value_if_it_exists(self, repo: CRUDRepoImpl[A, B, C], item: B): assert repo.get(1) == item - def test_update(self, repo: CRUDRepo[A, B, C], transform: Callable[[B], B], item: B, expected_transformed_value: B): + def test_update( + self, repo: CRUDRepoImpl[A, B, C], transform: Callable[[B], B], item: B, expected_transformed_value: B + ): assert repo.update(1, transform) == expected_transformed_value - def test_delete(self, repo: CRUDRepo[A, B, C], item: B): + def test_delete(self, repo: CRUDRepoImpl[A, B, C], item: B): assert repo.get(1) is not None repo.delete(1, None) assert repo.get(1) is None diff --git a/src/infra/subscription.py b/src/infra/subscription.py index 6c8131d..c696be3 100644 --- a/src/infra/subscription.py +++ b/src/infra/subscription.py @@ -5,7 +5,7 @@ from typing import TYPE_CHECKING from datetime import datetime, timedelta from src.infra.users import UserModel from sqlalchemy.orm import Mapped, mapped_column, relationship, Session, joinedload -from src.infra.db import Base, CRUDRepo +from src.infra.db import Base, CRUDRepoImpl from sqlalchemy import Engine, ForeignKey, CheckConstraint, DateTime, func, UniqueConstraint, select from src.services.subscription import ( SubscriptionRepo, @@ -62,7 +62,7 @@ def mutate_subscription(db_sub: SubscriptionModel, sub: Subscription): db_sub.sequence_notified = sub.sequence_notified -class SubscriptionRepoImpl(SubscriptionRepo, CRUDRepo[SubscriptionModel, Subscription, SubscriptionCreateParams]): +class SubscriptionRepoImpl(SubscriptionRepo, CRUDRepoImpl[SubscriptionModel, Subscription, SubscriptionCreateParams]): def __init__(self, db: Engine): super().__init__( db, SubscriptionModel, model_to_subscription, mutate_subscription, lambda u: SubscriptionModel(**asdict(u)) diff --git a/src/infra/users.py b/src/infra/users.py index aa17d43..a4c7a14 100644 --- a/src/infra/users.py +++ b/src/infra/users.py @@ -6,7 +6,7 @@ from src.config.email import EmailAddress from sqlalchemy import String, VARCHAR, Engine, select from sqlalchemy.orm import mapped_column, Mapped, Session, relationship from argon2 import PasswordHasher -from src.infra.db import Base, CRUDRepo +from src.infra.db import Base, CRUDRepoImpl from src.services.users.data import UserDTO, User from src.services.users.repo import UserRepo @@ -52,7 +52,7 @@ def make_create_user(hasher: PasswordHasher): return create_user -class UserRepoImpl(CRUDRepo[UserModel, User, UserDTO], UserRepo): +class UserRepoImpl(CRUDRepoImpl[UserModel, User, UserDTO], UserRepo): def __init__(self, db: Engine): hasher = PasswordHasher() super().__init__(db, UserModel, model_to_user, mutate_user, make_create_user(hasher)) diff --git a/src/main.py b/src/main.py index f9fb8fd..b9f0d53 100644 --- a/src/main.py +++ b/src/main.py @@ -9,6 +9,9 @@ from dataclasses import dataclass from src.infra import get_database, UserRepoImpl from waitress import serve from flask import Flask +from flask_cors import CORS +from flask_talisman import Talisman + import argparse import logging @@ -33,6 +36,8 @@ class MyFlask(Flask): def create_app(name: str, config: Config) -> Flask: logging.basicConfig(level=config.logging.level, format=('%(asctime)s %(levelname)s [%(name)s] %(message)s')) app = MyFlask(name) + CORS(app) + Talisman(app) # Configure Services database = get_database(config.database) diff --git a/src/services/repo.py b/src/services/repo.py new file mode 100644 index 0000000..dceec86 --- /dev/null +++ b/src/services/repo.py @@ -0,0 +1,20 @@ +from typing import Callable +import abc + + +class CRUD[T, C](abc.ABC): + @abc.abstractmethod + def create(self, domain_item: C, verify: Callable[[T], None] | None) -> T: + pass + + @abc.abstractmethod + def get(self, item_id: int) -> T | None: + pass + + @abc.abstractmethod + def update(self, item_id: int, transform: Callable[[T], T]) -> T: + pass + + @abc.abstractmethod + def delete(self, item_id: int, verify: Callable[[T], None] | None): + pass diff --git a/src/services/subscription.py b/src/services/subscription.py index 856ad6c..720f747 100644 --- a/src/services/subscription.py +++ b/src/services/subscription.py @@ -1,4 +1,4 @@ -from src.infra.db import CRUD +from src.services.repo import CRUD import abc from datetime import datetime from src.services.publications import ( diff --git a/src/services/users/repo.py b/src/services/users/repo.py index 6f8f40c..cd6bc74 100644 --- a/src/services/users/repo.py +++ b/src/services/users/repo.py @@ -1,4 +1,4 @@ -from src.infra.db import CRUD +from src.services.repo import CRUD import abc from src.config.email import EmailAddress