Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion generate_schemas.js
Original file line number Diff line number Diff line change
Expand Up @@ -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` +
Expand Down
16 changes: 15 additions & 1 deletion mypy.ini
Original file line number Diff line number Diff line change
@@ -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
disallow_untyped_decorators = False
24 changes: 24 additions & 0 deletions ruff.toml
Original file line number Diff line number Diff line change
@@ -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"
83 changes: 54 additions & 29 deletions src/opengeodeweb_microservice/database/connection.py
Original file line number Diff line number Diff line change
@@ -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()
14 changes: 7 additions & 7 deletions src/opengeodeweb_microservice/database/data.py
Original file line number Diff line number Diff line change
@@ -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):
Expand All @@ -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)
Expand Down
15 changes: 10 additions & 5 deletions src/opengeodeweb_microservice/database/data_types.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from typing import Literal, get_args, cast
from typing import Literal, cast, get_args

GeodePointMeshType = Literal[
"PointSet2D",
Expand Down Expand Up @@ -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"]
Expand Down
55 changes: 30 additions & 25 deletions src/opengeodeweb_microservice/schemas.py
Original file line number Diff line number Diff line change
@@ -1,68 +1,70 @@
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)
)
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]
Expand All @@ -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
Expand Down
37 changes: 16 additions & 21 deletions tests/conftest.py
Original file line number Diff line number Diff line change
@@ -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)
Expand All @@ -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()
Loading
Loading