diff --git a/generate_schemas.js b/generate_schemas.js index d5fc3df..952b89b 100755 --- a/generate_schemas.js +++ b/generate_schemas.js @@ -164,11 +164,12 @@ async function appendPythonResponse(filename, jsonData, requestContent) { const routeName = `${filename}_route`; exportedNames.push(routeName); return ( + "from pathlib import Path\n" + schemasImport + "\n" + pythonContent.trimEnd() + `\n\n\n${routeName} = Route(\n` + - ` schema=load_schema(__file__),\n` + + ` schema=load_schema(Path(__file__)),\n` + ` params=${paramsClass},\n` + ` response=${responseClass},\n` + `)\n\n` + diff --git a/mypy.ini b/mypy.ini index 089c049..7f3ca7c 100644 --- a/mypy.ini +++ b/mypy.ini @@ -1,4 +1,18 @@ [mypy] strict = True +enable_error_code = + deprecated, + exhaustive-match, + explicit-override, + ignore-without-code, + mutable-override, + possibly-undefined, + redundant-expr, + redundant-self, + truthy-bool, + truthy-iterable, + unimported-reveal, + unused-awaitable, + unused-ignore files = src/ -disallow_untyped_decorators = False \ No newline at end of file +disallow_untyped_decorators = False diff --git a/ruff.toml b/ruff.toml new file mode 100644 index 0000000..46d218b --- /dev/null +++ b/ruff.toml @@ -0,0 +1,24 @@ +# https://docs.astral.sh/ruff/configuration/ + +target-version = "py312" +line-length = 100 +indent-width = 4 +extend-exclude = ["**/schemas/**"] # generated schemas + +[lint] +select = ["ALL"] +ignore = [ + "CPY001", # ignore copyright notice + "D", # ignore undocumented + "COM812", # missing-trailing-comma +] + +[lint.per-file-ignores] +"tests/**" = [ + "S101", # assert is how pytest works + "PLR2004", # expected values in asserts are not magic numbers +] + +[format] +quote-style = "double" +indent-style = "space" diff --git a/src/opengeodeweb_microservice/database/connection.py b/src/opengeodeweb_microservice/database/connection.py index ba43e76..b9cc7c0 100644 --- a/src/opengeodeweb_microservice/database/connection.py +++ b/src/opengeodeweb_microservice/database/connection.py @@ -1,37 +1,62 @@ """Database connection management""" -from sqlalchemy import create_engine -from sqlalchemy.orm import sessionmaker, scoped_session, Session +import logging +from dataclasses import dataclass +from pathlib import Path + +from sqlalchemy import Engine, create_engine +from sqlalchemy.orm import Session, scoped_session, sessionmaker + from .base import Base -DATABASE_FILENAME = "project.db" - -engine = None -session_factory = None -scoped_session_registry = None - - -def init_database(db_path: str = DATABASE_FILENAME, create_tables: bool = True) -> None: - global engine, session_factory, scoped_session_registry - - if engine is None: - engine = create_engine( - f"sqlite:///{db_path}", - connect_args={"check_same_thread": False}, - ) - print(f"Database engine created for {db_path}", flush=True) - session_factory = sessionmaker(bind=engine) - scoped_session_registry = scoped_session(session_factory) - if create_tables: - Base.metadata.create_all(engine) - print(f"Database tables created for {db_path}", flush=True) - else: - print(f"Database connected (tables not created) for {db_path}", flush=True) +DATABASE_FILENAME = Path("project.db") + +logger = logging.getLogger(__name__) + + +class DatabaseNotInitializedError(RuntimeError): + def __init__(self) -> None: + super().__init__("Database not initialized. Call init_database() first.") + + +@dataclass +class _DatabaseState: + engine: Engine | None = None + scoped_session_registry: scoped_session[Session] | None = None + + +_state = _DatabaseState() + + +def init_database(db_path: Path = DATABASE_FILENAME, *, create_tables: bool = True) -> None: + if _state.engine is not None: + logger.info("Database engine already exists for %s, reusing", db_path) + return + + _state.engine = create_engine( + f"sqlite:///{db_path}", + connect_args={"check_same_thread": False}, + ) + logger.info("Database engine created for %s", db_path) + _state.scoped_session_registry = scoped_session(sessionmaker(bind=_state.engine)) + if create_tables: + Base.metadata.create_all(_state.engine) + logger.info("Database tables created for %s", db_path) else: - print(f"Database engine already exists for {db_path}, reusing", flush=True) + logger.info("Database connected (tables not created) for %s", db_path) + + +def close_database() -> None: + """Release every session and connection so the database file can be replaced.""" + if _state.scoped_session_registry is not None: + _state.scoped_session_registry.remove() + if _state.engine is not None: + _state.engine.dispose() + _state.engine = None + _state.scoped_session_registry = None def get_session() -> Session: - if scoped_session_registry is None: - raise RuntimeError("Database not initialized. Call init_database() first.") - return scoped_session_registry() + if _state.scoped_session_registry is None: + raise DatabaseNotInitializedError + return _state.scoped_session_registry() diff --git a/src/opengeodeweb_microservice/database/data.py b/src/opengeodeweb_microservice/database/data.py index f0249ef..6f577c6 100644 --- a/src/opengeodeweb_microservice/database/data.py +++ b/src/opengeodeweb_microservice/database/data.py @@ -1,9 +1,11 @@ -from sqlalchemy import String, JSON, select +import uuid + +from sqlalchemy import String, select from sqlalchemy.orm import Mapped, mapped_column -from .connection import get_session + from .base import Base -from .data_types import GeodeObjectType, ViewerType, ViewerElementsType -import uuid +from .connection import get_session +from .data_types import GeodeObjectType, ViewerElementsType, ViewerType class Data(Base): @@ -15,9 +17,7 @@ class Data(Base): geode_id: Mapped[str] = mapped_column(String, nullable=False) geode_object: Mapped[GeodeObjectType] = mapped_column(String, nullable=False) viewer_object: Mapped[ViewerType] = mapped_column(String, nullable=False) - viewer_elements_type: Mapped[ViewerElementsType] = mapped_column( - String, nullable=False - ) + viewer_elements_type: Mapped[ViewerElementsType] = mapped_column(String, nullable=False) native_file: Mapped[str | None] = mapped_column(String, nullable=True) viewable_file: Mapped[str | None] = mapped_column(String, nullable=True) light_viewable_file: Mapped[str | None] = mapped_column(String, nullable=True) diff --git a/src/opengeodeweb_microservice/database/data_types.py b/src/opengeodeweb_microservice/database/data_types.py index dcaf1db..574711f 100644 --- a/src/opengeodeweb_microservice/database/data_types.py +++ b/src/opengeodeweb_microservice/database/data_types.py @@ -1,4 +1,4 @@ -from typing import Literal, get_args, cast +from typing import Literal, cast, get_args GeodePointMeshType = Literal[ "PointSet2D", @@ -60,12 +60,17 @@ def _flatten_literal_args(literal: object) -> tuple[str, ...]: GeodeObjectType_values = _flatten_literal_args(GeodeObjectType) -def geode_object_type(value: str) -> GeodeObjectType: - if value not in GeodeObjectType_values: - raise ValueError( +class InvalidGeodeObjectTypeError(ValueError): + def __init__(self, value: str) -> None: + super().__init__( f"Invalid GeodeObjectType: {value!r}. Must be one of {GeodeObjectType_values}" ) - return cast(GeodeObjectType, value) + + +def geode_object_type(value: str) -> GeodeObjectType: + if value not in GeodeObjectType_values: + raise InvalidGeodeObjectTypeError(value) + return cast("GeodeObjectType", value) ViewerType = Literal["mesh", "model"] diff --git a/src/opengeodeweb_microservice/schemas.py b/src/opengeodeweb_microservice/schemas.py index 2d009c0..025bf14 100644 --- a/src/opengeodeweb_microservice/schemas.py +++ b/src/opengeodeweb_microservice/schemas.py @@ -1,37 +1,41 @@ -import os -import glob -import json import dataclasses +import json +import sys +from collections.abc import Sized from dataclasses import dataclass -from typing import Any, Generic, TypeVar +from pathlib import Path +from typing import TYPE_CHECKING, Any from dataclasses_json import DataClassJsonMixin -type SchemaDict = dict[str, str] +if TYPE_CHECKING: + from _typeshed import DataclassInstance -ERROR_SCHEMA_PATH = os.path.join(os.path.dirname(__file__), "error.json") +type SchemaDict = dict[str, str] -ParamsT = TypeVar("ParamsT", bound=DataClassJsonMixin) -ResponseT = TypeVar("ResponseT", bound=DataClassJsonMixin) +ERROR_SCHEMA_PATH = Path(__file__).parent / "error.json" MAX_VALUE_LENGTH = 80 -def _format_value(value: Any, max_length: int) -> str: +def _format_value(value: object, max_length: int) -> str: if dataclasses.is_dataclass(value) and not isinstance(value, type): return format_dataclass(value, max_length) text = repr(value) if len(text) <= max_length: return text details = type(value).__name__ - if hasattr(value, "__len__"): + if isinstance(value, Sized): details += f", len={len(value)}" return f"{text[:max_length]}... ({details})" -def format_dataclass(instance: Any, max_length: int = MAX_VALUE_LENGTH) -> str: - """Like repr(), but every field value longer than max_length is truncated and followed by its type and length.""" +def format_dataclass(instance: "DataclassInstance", max_length: int = MAX_VALUE_LENGTH) -> str: + """Like repr(), but truncate every field value longer than max_length. + + Each truncated value is followed by its type and length. + """ fields = ", ".join( f"{field.name}={_format_value(getattr(instance, field.name), max_length)}" for field in dataclasses.fields(instance) @@ -39,30 +43,28 @@ def format_dataclass(instance: Any, max_length: int = MAX_VALUE_LENGTH) -> str: return f"{type(instance).__name__}({fields})" -def print_dataclass(instance: Any) -> None: - print(format_dataclass(instance), flush=True) +def print_dataclass(instance: "DataclassInstance") -> None: + sys.stdout.write(f"{format_dataclass(instance)}\n") + sys.stdout.flush() -def get_schemas_dict(path: str) -> dict[str, SchemaDict]: +def get_schemas_dict(path: Path) -> dict[str, SchemaDict]: schemas_dict: dict[str, SchemaDict] = {} - for json_file in glob.glob(os.path.join(path, "*.json")): - filename = os.path.basename(json_file) - with open(os.path.join(path, json_file), "r") as file: - file_content = json.load(file) - schemas_dict[os.path.splitext(filename)[0]] = file_content + for json_file in path.glob("*.json"): + with json_file.open() as file: + schemas_dict[json_file.stem] = json.load(file) return schemas_dict -def load_schema(python_file: str) -> dict[str, Any]: +def load_schema(python_file: Path) -> dict[str, Any]: """Load the JSON route schema sitting next to a generated schema module.""" - json_file = os.path.splitext(python_file)[0] + ".json" - with open(json_file, "r") as file: + with python_file.with_suffix(".json").open() as file: schema: dict[str, Any] = json.load(file) return schema @dataclass(frozen=True) -class Route(Generic[ParamsT, ResponseT]): +class Route[ParamsT: DataClassJsonMixin, ResponseT: DataClassJsonMixin]: """Generated binding between a route schema, its request params and its success response.""" schema: dict[str, Any] @@ -72,7 +74,10 @@ class Route(Generic[ParamsT, ResponseT]): @dataclass class BinaryResponse(DataClassJsonMixin): - """Response of a route streaming a file (schema `"response": {"format": "binary"}`) instead of JSON.""" + """Response of a route streaming a file instead of JSON. + + Matches a schema with `"response": {"format": "binary"}`. + """ @dataclass diff --git a/tests/conftest.py b/tests/conftest.py index 4512ae4..f7f44f5 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,10 +1,17 @@ -import os -from typing import Generator +import contextlib +from collections.abc import Generator +from pathlib import Path + import pytest -from opengeodeweb_microservice.database.connection import init_database, get_session + +from opengeodeweb_microservice.database.connection import ( + close_database, + get_session, + init_database, +) from opengeodeweb_microservice.database.data import Data -DB_PATH = os.path.join(os.path.dirname(__file__), "test_project.db") +DB_PATH = Path(__file__).parent / "test_project.db" @pytest.fixture(scope="session", autouse=True) @@ -14,28 +21,16 @@ def setup_database() -> Generator[None, None, None]: _cleanup_database(DB_PATH) -def _cleanup_database(db_path: str) -> None: - try: - session = get_session() - session.close() - except Exception: - pass - - if os.path.exists(db_path): - try: - os.remove(db_path) - except PermissionError: - pass +def _cleanup_database(db_path: Path) -> None: + close_database() + with contextlib.suppress(PermissionError): + db_path.unlink(missing_ok=True) @pytest.fixture(autouse=True) def clean_database() -> Generator[None, None, None]: with get_session() as session: - session = get_session() session.query(Data).delete() session.commit() yield - try: - session.rollback() - except Exception: - pass + session.rollback() diff --git a/tests/test_database.py b/tests/test_database.py index 329334f..b1e279a 100644 --- a/tests/test_database.py +++ b/tests/test_database.py @@ -1,19 +1,20 @@ +import pytest from sqlalchemy import select -from opengeodeweb_microservice.database.data import Data from opengeodeweb_microservice.database.connection import get_session +from opengeodeweb_microservice.database.data import Data GEODE_ID = "01a08187-2c4c-7e64-85c5-52c3439f0626" -def test_data_crud_operations(clean_database: None) -> None: +@pytest.mark.usefixtures("clean_database") +def test_data_crud_operations() -> None: data = Data.create( geode_id=GEODE_ID, geode_object="test_object", viewer_object="test_viewer", viewer_elements_type="test_type", ) - print("id", data.id, flush=True) assert data.id is not None assert isinstance(data.id, str) assert len(data.id) == 32 @@ -28,7 +29,8 @@ def test_data_crud_operations(clean_database: None) -> None: assert non_existent is None -def test_data_with_file_assignments(clean_database: None) -> None: +@pytest.mark.usefixtures("clean_database") +def test_data_with_file_assignments() -> None: data = Data.create( geode_id=GEODE_ID, geode_object="geode_object", @@ -52,7 +54,8 @@ def test_data_with_file_assignments(clean_database: None) -> None: assert retrieved.geode_object == "geode_object" -def test_data_geode_id_is_not_unique(clean_database: None) -> None: +@pytest.mark.usefixtures("clean_database") +def test_data_geode_id_is_not_unique() -> None: first = Data.create( geode_id=GEODE_ID, geode_object="geode_object", diff --git a/tests/test_schemas.py b/tests/test_schemas.py index 570038b..c72759e 100644 --- a/tests/test_schemas.py +++ b/tests/test_schemas.py @@ -1,9 +1,8 @@ import json - -import fastjsonschema # type: ignore - from dataclasses import dataclass +import fastjsonschema + from opengeodeweb_microservice.schemas import ( ERROR_SCHEMA_PATH, ErrorResponse, @@ -12,13 +11,9 @@ def test_error_response_matches_error_schema() -> None: - with open(ERROR_SCHEMA_PATH, "r") as file: + with ERROR_SCHEMA_PATH.open() as file: validate = fastjsonschema.compile(json.load(file)) - validate( - ErrorResponse( - code=500, name="Internal Server Error", description="boom" - ).to_dict() - ) + validate(ErrorResponse(code=500, name="Internal Server Error", description="boom").to_dict()) @dataclass @@ -38,7 +33,5 @@ def test_format_dataclass_truncates_long_values() -> None: _Outer(name="short", content="x" * 1000, inner=_Inner(values=list(range(500)))), max_length=20, ) - assert text.startswith( - "_Outer(name='short', content='xxxxxxxxxxxxxxxxxxx... (str, len=1000)" - ) + assert text.startswith("_Outer(name='short', content='xxxxxxxxxxxxxxxxxxx... (str, len=1000)") assert "inner=_Inner(values=[0, 1, 2, 3, 4, 5, 6... (list, len=500))" in text