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
2 changes: 1 addition & 1 deletion example/config-app/uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

2 changes: 1 addition & 1 deletion example/database-app/uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Original file line number Diff line number Diff line change
Expand Up @@ -743,7 +743,7 @@ def morphs(self, column, nullable=False, indexes=True):
self._last_column = _columns
return self

def to_sql(self):
async def to_sql(self):
"""Compiles the blueprint class into a sql statement.

Returns:
Expand All @@ -754,13 +754,10 @@ def to_sql(self):
elif self._action == "create_table_if_not_exists":
return self.platform().compile_create_sql(self.table, if_not_exists=True)
else:
if not self._dry and self.table.from_table is None:
# get current table schema
table = self.platform().get_current_schema(
if self.table.from_table is None:
self.table.from_table = await self.platform().get_current_schema(
self.connection, self.table.name, schema=self.schema
)
self.table.from_table = table

return self.platform().compile_alter_sql(self.table)

def __enter__(self):
Expand All @@ -781,12 +778,11 @@ def __exit__(self, exc_type, exc_value, exc_traceback):
async def __aenter__(self):
return self

# TODO: review
async def __aexit__(self, exc_type, exc_value, exc_traceback):
if self._dry:
return

sql = self.to_sql()
sql = await self.to_sql()
if isinstance(sql, list):
for q in sql:
await self.connection.statement(q, ())
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -435,10 +435,10 @@ def compile_column_exists(self, table, column):
def compile_get_all_tables(self, database, schema=None):
return f"SELECT table_name FROM information_schema.tables WHERE table_schema = '{database}'"

def get_current_schema(self, connection, table_name, schema=None):
async def get_current_schema(self, connection, table_name, schema=None):
table = Table(table_name)
sql = f"DESCRIBE {table_name}"
result = connection.query(sql, ())
result = await connection.select(sql, ())
reversed_type_map = {v: k for k, v in self.type_map.items()}

for column in result:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -492,7 +492,7 @@ def compile_column_exists(self, table, column):
def compile_get_all_tables(self, database=None, schema=None):
return f"SELECT table_name FROM information_schema.tables WHERE table_schema = '{schema or 'public'}' AND table_catalog = '{database}' AND table_type = 'BASE TABLE'"

def get_current_schema(self, connection, table_name, schema=None):
async def get_current_schema(self, connection, table_name, schema=None):
sql = self.table_information_string().format(
table=table_name, schema=schema or "public"
)
Expand All @@ -501,7 +501,7 @@ def get_current_schema(self, connection, table_name, schema=None):
reversed_type_map.update(self.table_info_map)
table = Table(table_name)

result = connection.query(sql, ())
result = await connection.select(sql, ())
for column in result:
column_type = reversed_type_map.get(column["data_type"].upper())

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -367,13 +367,13 @@ def columnize_names(self, columns):

return names

def get_current_schema(self, connection, table_name, schema=None):
async def get_current_schema(self, connection, table_name, schema=None):
sql = f"PRAGMA table_info({table_name})"

reversed_type_map = {v: k for k, v in self.type_map.items()}
table = Table(table_name)

result = connection.query(sql, ())
result = await connection.select(sql, ())
for column in result:
column_type = self.get_column_type(
reversed_type_map, column["type"].upper()
Expand Down
42 changes: 42 additions & 0 deletions fastapi_startkit/src/fastapi_startkit/masoniteorm/schema/schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,48 @@


class Schema:
_type_hints_map = {
"string": str,
"char": str,
"big_increments": int,
"integer": int,
"big_integer": int,
"tiny_integer": int,
"small_integer": int,
"medium_integer": int,
"integer_unsigned": int,
"big_integer_unsigned": int,
"tiny_integer_unsigned": int,
"small_integer_unsigned": int,
"medium_integer_unsigned": int,
"increments": int,
"uuid": str,
"binary": bytes,
"boolean": bool,
"decimal": float,
"double": float,
"enum": str,
"text": str,
"tiny_text": str,
"float": float,
"geometry": str,
"json": dict,
"jsonb": bytes,
"inet": str,
"cidr": str,
"macaddr": str,
"long_text": str,
"point": str,
"time": str,
"timestamp": str,
"date": str,
"year": str,
"datetime": str,
"tiny_increments": int,
"unsigned": int,
"unsigned_integer": int,
}

def __init__(self, manager) -> None:
self._manager = manager
self._connection = None
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,11 @@
from fastapi_startkit.masoniteorm import Migration


class AddBodyToPostsTable(Migration):
async def up(self):
async with await self.schema.table("posts") as table:
table.text("body").nullable()

async def down(self):
async with await self.schema.table("posts") as table:
table.drop_column("body")
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@ async def test_can_add_columns(self):

self.assertEqual(len(blueprint.table.added_columns), 2)
self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
[
'CREATE TABLE "users" ("name" VARCHAR(255) NOT NULL, "age" INTEGER NOT NULL)'
],
Expand All @@ -32,7 +32,7 @@ async def test_can_add_tiny_text(self):

self.assertEqual(len(blueprint.table.added_columns), 1)
self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
['CREATE TABLE "users" ("description" TEXT NOT NULL)'],
)

Expand All @@ -45,7 +45,7 @@ async def test_can_add_unsigned_decimal(self):

self.assertEqual(len(blueprint.table.added_columns), 1)
self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
['CREATE TABLE "users" ("amount" DECIMAL(19, 4) NOT NULL)'],
)

Expand All @@ -59,7 +59,7 @@ async def test_can_create_table_if_not_exists(self):

self.assertEqual(len(blueprint.table.added_columns), 2)
self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
[
'CREATE TABLE IF NOT EXISTS "users" ("name" VARCHAR(255) NOT NULL, "age" INTEGER NOT NULL)'
],
Expand All @@ -76,7 +76,7 @@ async def test_can_add_columns_with_constraint(self):

self.assertEqual(len(blueprint.table.added_columns), 2)
self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
[
'CREATE TABLE "users" ("name" VARCHAR(255) NOT NULL, "age" INTEGER NOT NULL, UNIQUE(name))'
],
Expand All @@ -90,7 +90,7 @@ async def test_can_have_float_type(self):
blueprint.float("amount")

self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
[
"""CREATE TABLE "users" ("""
"""\"amount" FLOAT(19, 4) NOT NULL)"""
Expand All @@ -109,7 +109,7 @@ async def test_can_add_columns_with_foreign_key_constraint(self):

self.assertEqual(len(blueprint.table.added_columns), 3)
self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
[
'CREATE TABLE "users" '
'("name" VARCHAR(255) NOT NULL, '
Expand All @@ -134,7 +134,7 @@ async def test_can_add_columns_with_foreign_key_constraint_name(self):

self.assertEqual(len(blueprint.table.added_columns), 3)
self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
[
'CREATE TABLE "users" '
'("name" VARCHAR(255) NOT NULL, '
Expand All @@ -154,7 +154,7 @@ async def test_can_use_morphs_for_polymorphism_relationships(self):

self.assertEqual(len(blueprint.table.added_columns), 2)
self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
[
'CREATE TABLE "likes" ("record_id" INTEGER UNSIGNED NOT NULL, "record_type" VARCHAR NOT NULL)',
'CREATE INDEX likes_record_id_index ON "likes"(record_id)',
Expand All @@ -181,7 +181,7 @@ async def test_can_advanced_table_creation(self):

self.assertEqual(len(blueprint.table.added_columns), 11)
self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
[
"""CREATE TABLE "users" ("id" INTEGER NOT NULL, "name" VARCHAR(255) NOT NULL, "gender" VARCHAR(255) CHECK(gender IN ('male', 'female')) NOT NULL, "email" VARCHAR(255) NOT NULL, """
""""password" VARCHAR(255) NOT NULL, "option" VARCHAR(255) NOT NULL DEFAULT 'ADMIN', "admin" INTEGER NOT NULL DEFAULT 0, "remember_token" VARCHAR(255) NULL, """
Expand All @@ -207,7 +207,7 @@ async def test_can_create_indexes(self):

self.assertEqual(len(blueprint.table.added_columns), 0)
self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
[
'CREATE INDEX users_name_index ON "users"(name)',
'CREATE INDEX active_idx ON "users"(active)',
Expand All @@ -229,7 +229,7 @@ async def test_can_create_indexes_on_previous_column(self):

self.assertEqual(len(blueprint.table.added_columns), 2)
self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
[
'ALTER TABLE "users" ADD COLUMN "email" VARCHAR NOT NULL',
'ALTER TABLE "users" ADD COLUMN "active" VARCHAR NOT NULL',
Expand All @@ -250,7 +250,7 @@ async def test_can_have_composite_keys(self):

self.assertEqual(len(blueprint.table.added_columns), 3)
self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
[
'CREATE TABLE "users" '
'("name" VARCHAR(255) NOT NULL, '
Expand All @@ -272,7 +272,7 @@ async def test_can_have_column_primary_key(self):

self.assertEqual(len(blueprint.table.added_columns), 3)
self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
[
'CREATE TABLE "users" '
'("name" VARCHAR(255) NOT NULL, '
Expand Down Expand Up @@ -309,7 +309,7 @@ async def test_can_advanced_table_creation2(self):

self.assertEqual(len(blueprint.table.added_columns), 17)
self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
[
'CREATE TABLE "users" ("id" INTEGER NOT NULL, "name" VARCHAR(255) NOT NULL, "duration" VARCHAR(255) NOT NULL, '
'"url" VARCHAR(255) NOT NULL, "payload" JSON NOT NULL, "birth" VARCHAR(4) NOT NULL, "last_address" VARCHAR(255) NULL, "route_origin" VARCHAR(255) NULL, "mac_address" VARCHAR(255) NULL, '
Expand Down Expand Up @@ -390,7 +390,7 @@ async def test_can_have_unsigned_columns(self):
blueprint.medium_integer("medium_profile_id").unsigned()

self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
[
"""CREATE TABLE "users" ("""
""""profile_id" INTEGER UNSIGNED NOT NULL, """
Expand Down Expand Up @@ -444,7 +444,7 @@ async def test_can_add_enum(self):

self.assertEqual(len(blueprint.table.added_columns), 1)
self.assertEqual(
blueprint.to_sql(),
await blueprint.to_sql(),
[
"CREATE TABLE \"users\" (\"status\" VARCHAR(255) CHECK(status IN ('active', 'inactive')) NOT NULL DEFAULT 'active')"
],
Expand Down
Loading
Loading