Compare commits

...

3 commits

Author SHA1 Message Date
Campbell Alden
07aac38dbc Configure CORS and security middlewares 2026-09-05 23:53:40 +09:00
Campbell Alden
13779d5a74 Add a dev mode configuration flag 2026-09-05 23:38:55 +09:00
Campbell Alden
c3cd7c2796 Restructure where the CRUD base class lives for clarity and to avoid circular deps 2026-09-05 23:38:46 +09:00
14 changed files with 69 additions and 39 deletions

View file

@ -1,6 +1,7 @@
{
"host": "0.0.0.0",
"port": 8080,
"dev_mode": false,
"logging": {
"level": "info"
}

View file

@ -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];

View file

@ -9,6 +9,8 @@ let
pyjwt
cryptography
jinja2
flask-talisman
flask-cors
]);
in
with pkgs;

0
src/api/__init__.py Normal file
View file

View file

@ -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,

View file

@ -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',

View file

@ -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,

View file

@ -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

View file

@ -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))

View file

@ -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))

View file

@ -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)

20
src/services/repo.py Normal file
View file

@ -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

View file

@ -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 (

View file

@ -1,4 +1,4 @@
from src.infra.db import CRUD
from src.services.repo import CRUD
import abc
from src.config.email import EmailAddress