diff --git a/config.example.json b/config.example.json index 9279e34..b97e6a3 100644 --- a/config.example.json +++ b/config.example.json @@ -1,7 +1,6 @@ { "host": "0.0.0.0", "port": 8080, - "dev_mode": false, "logging": { "level": "info" } diff --git a/derivation.nix b/derivation.nix index 342c32e..d27b324 100644 --- a/derivation.nix +++ b/derivation.nix @@ -3,18 +3,7 @@ with python313Packages; buildPythonApplication { pname = "cereal"; version = "0.0.1"; - propagatedBuildInputs = [ - flask - requests - waitress - sqlalchemy - argon2-cffi - pyjwt - cryptography - jinja2 - flask-talisman - flask-cors - ]; + propagatedBuildInputs = [ flask requests waitress sqlalchemy argon2-cffi pyjwt cryptography jinja2]; src = ./.; pyproject = true; build-system = [setuptools]; diff --git a/shell.nix b/shell.nix index 11f109f..e5e7b1b 100644 --- a/shell.nix +++ b/shell.nix @@ -9,8 +9,6 @@ let pyjwt cryptography jinja2 - flask-talisman - flask-cors ]); in with pkgs; diff --git a/src/api/__init__.py b/src/api/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/src/config/__init__.py b/src/config/__init__.py index 1448c56..c2c9171 100644 --- a/src/config/__init__.py +++ b/src/config/__init__.py @@ -1,6 +1,7 @@ import json from dataclasses import dataclass from typing import Any + from .logging import Logging from .email import Email from .auth import Auth @@ -12,7 +13,6 @@ from .parse import parse_nested_config class Config: host: str port: int | None - dev_mode: bool logging: Logging email: Email database: Database @@ -25,17 +25,9 @@ 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 ee15f85..8851384 100644 --- a/src/infra/__init__.py +++ b/src/infra/__init__.py @@ -1,4 +1,4 @@ -from .db import Base, CRUDRepoImpl, UpdateMissingEntryError, get_database +from .db import Base, CRUD, CRUDRepo, UpdateMissingEntryError, get_database from .users import UserModel, UserRepoImpl from .publication import ( PublicationEntryModel, @@ -11,6 +11,7 @@ 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 ad9ff98..6359df4 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,7 +22,25 @@ class UpdateMissingEntryError(Exception): super().__init__(f'Attempted to update non-existing {model_name} {id}') -class CRUDRepoImpl[T: Base, U, C](CRUD[U, C]): +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]): def __init__( self, db: Engine, diff --git a/src/infra/db_tests.py b/src/infra/db_tests.py index 01ecf4f..e764924 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 CRUDRepoImpl, Base +from src.infra import CRUDRepo, Base import pytest class CRUDRepoHelper[A: Base, B, C]: @pytest.fixture - def repo(self, db: Engine) -> CRUDRepoImpl[A, B, C]: + def repo(self, db: Engine) -> CRUDRepo[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: CRUDRepoImpl[A, B, C], dto: C) -> B: + def item(self, repo: CRUDRepo[A, B, C], dto: C) -> B: return repo.create(dto, None) @pytest.fixture @@ -31,28 +31,26 @@ 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: CRUDRepoImpl[A, B, C], dto: C): + def test_create_returns_none_if_invalid(self, repo: CRUDRepo[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: CRUDRepoImpl[A, B, C], dto: C): + def test_create_returns_a_value_if_valid(self, repo: CRUDRepo[A, B, C], dto: C): assert repo.create(dto, lambda _: None) is not None - def test_get_returns_none_when_missing(self, repo: CRUDRepoImpl[A, B, C]): + def test_get_returns_none_when_missing(self, repo: CRUDRepo[A, B, C]): assert repo.get(42) is None - def test_get_returns_a_value_if_it_exists(self, repo: CRUDRepoImpl[A, B, C], item: B): + def test_get_returns_a_value_if_it_exists(self, repo: CRUDRepo[A, B, C], item: B): assert repo.get(1) == item - def test_update( - self, repo: CRUDRepoImpl[A, B, C], transform: Callable[[B], B], item: B, expected_transformed_value: B - ): + def test_update(self, repo: CRUDRepo[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: CRUDRepoImpl[A, B, C], item: B): + def test_delete(self, repo: CRUDRepo[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 c696be3..6c8131d 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, CRUDRepoImpl +from src.infra.db import Base, CRUDRepo 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, CRUDRepoImpl[SubscriptionModel, Subscription, SubscriptionCreateParams]): +class SubscriptionRepoImpl(SubscriptionRepo, CRUDRepo[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 a4c7a14..aa17d43 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, CRUDRepoImpl +from src.infra.db import Base, CRUDRepo 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(CRUDRepoImpl[UserModel, User, UserDTO], UserRepo): +class UserRepoImpl(CRUDRepo[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 b9f0d53..f9fb8fd 100644 --- a/src/main.py +++ b/src/main.py @@ -9,9 +9,6 @@ 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 @@ -36,8 +33,6 @@ 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 deleted file mode 100644 index dceec86..0000000 --- a/src/services/repo.py +++ /dev/null @@ -1,20 +0,0 @@ -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 720f747..856ad6c 100644 --- a/src/services/subscription.py +++ b/src/services/subscription.py @@ -1,4 +1,4 @@ -from src.services.repo import CRUD +from src.infra.db 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 cd6bc74..6f8f40c 100644 --- a/src/services/users/repo.py +++ b/src/services/users/repo.py @@ -1,4 +1,4 @@ -from src.services.repo import CRUD +from src.infra.db import CRUD import abc from src.config.email import EmailAddress