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
Original file line number Diff line number Diff line change
Expand Up @@ -6,21 +6,21 @@
@dataclass
class SQLiteConfig:
driver: str = "sqlite"
url: Optional[str] = env("DB_URL", None)
url: Optional[str] = env("DB_URL", None, cast=False)
database: str = env("DB_DATABASE", "database.sqlite")
options: Optional[Dict[str, Any]] = None


@dataclass
class MySQLConfig:
driver: str = "mysql"
url: Optional[str] = env("DB_URL", None)
url: Optional[str] = env("DB_URL", None, cast=False)
host: str = env("DB_HOST", "127.0.0.1")
port: int = env("DB_PORT", 3306)
database: str = env("DB_DATABASE", "inertia")
username: str = env("DB_USERNAME", "root")
password: str = env("DB_PASSWORD", "")
unix_socket: str = env("DB_SOCKET", "")
password: str = env("DB_PASSWORD", "", cast=False)
unix_socket: str = env("DB_SOCKET", "", cast=False)
charset: str = env("DB_CHARSET", "utf8mb4")
collation: str = env("DB_COLLATION", "utf8mb4_unicode_ci")
options: Optional[Dict[str, Any]] = None
Expand All @@ -29,12 +29,12 @@ class MySQLConfig:
@dataclass
class PostgresConfig:
driver: str = "postgres"
url: Optional[str] = env("DB_URL", None)
url: Optional[str] = env("DB_URL", None, cast=False)
host: str = env("DB_HOST", "127.0.0.1")
port: int = env("DB_PORT", 5432)
database: str = env("DB_DATABASE", "inertia")
username: str = env("DB_USERNAME", "postgres")
password: str = env("DB_PASSWORD", "")
password: str = env("DB_PASSWORD", "", cast=False)
charset: str = env("DB_CHARSET", "utf8")
sslmode: str = env("DB_SSLMODE", "prefer")
options: Optional[Dict[str, Any]] = None
Original file line number Diff line number Diff line change
@@ -1,15 +1,15 @@
from dataclasses import dataclass, field
from typing import Dict, Any
from typing import Any

from fastapi_startkit.environment import env
from fastapi_startkit.masoniteorm import SQLiteConfig
from fastapi_startkit.masoniteorm.config.config import MySQLConfig, PostgresConfig, SQLiteConfig


@dataclass
class DatabaseConfig:
default: str = field(default_factory=lambda: env("DB_CONNECTION", "pgsql"))

connections: Dict[str, Dict[str, Any]] = field(
connections: dict[str, SQLiteConfig | MySQLConfig | PostgresConfig | dict[str, Any]] = field(
default_factory=lambda: {
"sqlite": SQLiteConfig(
driver="sqlite",
Expand All @@ -19,6 +19,6 @@ class DatabaseConfig:
}
)

migrations: Dict[str, Dict[str, Any]] = field(
migrations: dict[str, str] = field(
default_factory=lambda: {"table": "migrations", "directory": "databases/migrations"}
)
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,8 @@
if TYPE_CHECKING:
from typing import Self

from fastapi_startkit.masoniteorm.query.grammars.BaseGrammar import BaseGrammar


class Transaction:
def __init__(self, owner: Connection):
Expand Down Expand Up @@ -100,11 +102,13 @@ def query(self) -> QueryBuilder:
async def get_connection(self) -> AsyncConnection:
return self.connection or await self.engine.connect()

def get_query_grammar(cls):
pass
@classmethod
def get_query_grammar(cls) -> type[BaseGrammar] | None:
return None

def get_post_processor(self):
pass
@classmethod
def get_post_processor(cls) -> type | None:
return None

async def begin_transaction(self) -> None:
connection = self.connection
Expand Down Expand Up @@ -199,7 +203,7 @@ async def delete(self, query: str, bindings: list | None = None) -> int:

async def select(self, query: str, bindings: list | None = None) -> list[dict]:
result = await self.run(query, bindings)
return result.mappings().all()
return [dict(row) for row in result.mappings().all()]

async def select_one(self, query: str, bindings: list | None = None) -> dict | None:
result = await self.run(query, bindings)
Expand Down
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from typing import Any
from typing import TYPE_CHECKING, Any, cast

from sqlalchemy import StaticPool
from sqlalchemy.pool import NullPool
Expand All @@ -12,6 +12,9 @@
)
from fastapi_startkit.masoniteorm.connections.mysql_connection import MySQLConnection

if TYPE_CHECKING:
from fastapi_startkit.application import Application


class ConnectionFactory:
DRIVER_URLS = {
Expand Down Expand Up @@ -64,7 +67,7 @@ def create_engine(cls, cfg: dict) -> AsyncEngine:
kwargs: dict[str, Any] = {"echo": True}
from fastapi_startkit.application import app

if app().is_testing():
if cast("Application", app()).is_testing():
kwargs["poolclass"] = NullPool
elif cfg["driver"] == "sqlite":
kwargs["connect_args"] = {"check_same_thread": False}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -8,13 +8,13 @@ class MySQLConnection(Connection):
"""Async MySQL connection backed by aiomysql via SQLAlchemy."""

@classmethod
def get_query_grammar(cls):
def get_query_grammar(cls) -> type[MySQLGrammar]:
return MySQLGrammar

@classmethod
def get_default_platform(cls):
return MySQLPlatform

@classmethod
def get_post_processor(cls):
def get_post_processor(cls) -> type[MySQLPostProcessor]:
return MySQLPostProcessor
Original file line number Diff line number Diff line change
Expand Up @@ -13,13 +13,13 @@ async def insert_get_id(self, query: str, bindings: list | None = None) -> int |
return row[0] if row is not None else None

@classmethod
def get_query_grammar(cls):
def get_query_grammar(cls) -> type[PostgresGrammar]:
return PostgresGrammar

@classmethod
def get_default_platform(cls):
return PostgresPlatform

@classmethod
def get_post_processor(cls):
def get_post_processor(cls) -> type[PostgresPostProcessor]:
return PostgresPostProcessor
Original file line number Diff line number Diff line change
Expand Up @@ -6,13 +6,13 @@

class SQliteConnection(Connection):
@classmethod
def get_query_grammar(cls):
def get_query_grammar(cls) -> type[SQLiteGrammar]:
return SQLiteGrammar

@classmethod
def get_default_platform(cls):
return SQLitePlatform

@classmethod
def get_post_processor(cls):
def get_post_processor(cls) -> type[SQLitePostProcessor]:
return SQLitePostProcessor
Original file line number Diff line number Diff line change
@@ -1,3 +1,9 @@
from typing import TYPE_CHECKING, cast

if TYPE_CHECKING:
from fastapi_startkit.application import Application


class DatabaseTransaction:
async def asyncStartTestRun(self):
from fastapi_startkit.masoniteorm.models import Model
Expand All @@ -24,7 +30,7 @@ async def migrate_database():
from fastapi_startkit.masoniteorm.migrations import Migrator
from fastapi_startkit.application import app as get_app

migration_dir = get_app().use_base_path("databases/migrations")
migration_dir = str(cast("Application", get_app()).use_base_path("databases/migrations"))
migrator = Migrator(migration_directory=migration_dir)
await migrator.fresh(ignore_fk=True)
RefreshDatabase.migrated = True
Original file line number Diff line number Diff line change
@@ -0,0 +1,21 @@
from fastapi_startkit.masoniteorm.connections.connection import Connection
from fastapi_startkit.masoniteorm.models import Model

from ..fixtures.model import User
from .test_case import TestCase


class TestConnectionSelect(TestCase):
async def test_select_returns_plain_dicts(self):
await User.create({"email": "select@example.com", "name": "Select", "is_admin": False})

rows = await Model.db_manager.connection(None).select(
"SELECT email, name FROM users WHERE email = ?", ["select@example.com"]
)

assert rows == [{"email": "select@example.com", "name": "Select"}]
assert all(type(row) is dict for row in rows)

def test_base_connection_has_no_grammar_or_processor(self):
assert Connection.get_query_grammar() is None
assert Connection.get_post_processor() is None
Original file line number Diff line number Diff line change
@@ -1,4 +1,7 @@
from fastapi_startkit.masoniteorm.testing.transaction import DatabaseTransaction
from pathlib import Path
from unittest.mock import AsyncMock, MagicMock, patch

from fastapi_startkit.masoniteorm.testing.transaction import DatabaseTransaction, RefreshDatabase

from ..fixtures.model import User
from .test_case import TestCase
Expand All @@ -14,3 +17,25 @@ async def test_start_and_stop_roll_back_writes(self):
finally:
await harness.asyncStopTestRun()
assert await User.where("email", "harness@example.com").first() is None


class TestRefreshDatabaseMigrate(TestCase):
async def test_migrate_database_runs_fresh_once_with_string_directory(self):
application = MagicMock()
application.use_base_path.return_value = Path("/project/databases/migrations")
migrator = MagicMock()
migrator.return_value.fresh = AsyncMock()

with (
patch.object(RefreshDatabase, "migrated", False),
patch("fastapi_startkit.application.app", return_value=application),
patch("fastapi_startkit.masoniteorm.migrations.Migrator", migrator),
):
await RefreshDatabase.migrate_database()
await RefreshDatabase.migrate_database()

assert RefreshDatabase.migrated is True

application.use_base_path.assert_called_once_with("databases/migrations")
migrator.assert_called_once_with(migration_directory="/project/databases/migrations")
migrator.return_value.fresh.assert_awaited_once_with(ignore_fk=True)
Loading