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.

26 changes: 25 additions & 1 deletion fastapi_startkit/src/fastapi_startkit/masoniteorm/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,30 @@
from .migrations.Migration import Migration
from .migrations.Migrator import Migrator
from .models import Model
from .models.fields import CreatedAtField, DateTimeField, Field, ModelField, UpdatedAtField
from .providers import DatabaseProvider
from .relationships import BelongsTo, BelongsToMany, HasMany, HasManyThrough, HasOne, HasOneThrough, MorphTo

__all__ = ["DatabaseProvider", "PostgresConfig", "MySQLConfig", "SQLiteConfig", "Model", "DB", "Migration", "Migrator"]
__all__ = [
"DatabaseProvider",
"PostgresConfig",
"MySQLConfig",
"SQLiteConfig",
"Model",
"DB",
"Migration",
"Migrator",
"ModelField",
"DateTimeField",
"CreatedAtField",
"UpdatedAtField",
"Field",
# Relationships
"HasOne",
"BelongsTo",
"HasMany",
"HasManyThrough",
"BelongsToMany",
"HasOneThrough",
"MorphTo",
]
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
"""__MIGRATION_NAME__ Migration."""

from fastapi_startkit.masoniteorm.migrations import Migration
from fastapi_startkit.masoniteorm import Migration


class __MIGRATION_NAME__(Migration):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -105,4 +105,7 @@ def get_dirty(self) -> dict:
}

def get_attributes_for_insert(self) -> dict:
return {**self._attributes, **self._dirty_attributes}
# _dirty_attributes already went through set_attribute (casts applied on assignment).
# _attributes is set raw via new_model_instance, so apply set casts here.
casted = {k: self.caster.set(k, v) for k, v in self._attributes.items()}
return {**casted, **self._dirty_attributes}
97 changes: 92 additions & 5 deletions fastapi_startkit/src/fastapi_startkit/masoniteorm/models/caster.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,66 @@ def set(self, value):
return str(value)


class TimeCast(BaseCast):
"""Casts a value to datetime.time; stored as HH:MM:SS string"""

def get(self, value):
if not value:
return None
if isinstance(value, datetime.time):
return value
return datetime.time.fromisoformat(str(value))

def set(self, value):
if not value:
return None
if isinstance(value, datetime.time):
return value.strftime("%H:%M:%S")
return str(value)


class TimeDeltaCast(BaseCast):
"""Casts a value to datetime.timedelta; stored as total seconds"""

def get(self, value):
if value is None:
return None
if isinstance(value, datetime.timedelta):
return value
return datetime.timedelta(seconds=float(value))

def set(self, value):
if value is None:
return None
if isinstance(value, datetime.timedelta):
return value.total_seconds()
return float(value)


@dataclass
class ModelCast(BaseCast):
model_class: type = field(default=None)

def get(self, value):
if value is None:
return None
if isinstance(value, self.model_class):
return value
data = json.loads(value) if isinstance(value, str) else value
return self.model_class(**data)

def set(self, value) -> Optional[str]:
if value is None:
return None
if isinstance(value, self.model_class):
if hasattr(value, "model_dump_json"):
return value.model_dump_json()
return json.dumps(value.__dict__)
if isinstance(value, dict):
return json.dumps(value)
return value


class Caster:
casts = {}

Expand All @@ -121,6 +181,8 @@ class Caster:
"float": FloatCast,
"date": DateCast,
"decimal": DecimalCast,
"time": TimeCast,
"timedelta": TimeDeltaCast,
}

IGNORE_CASTS = ["caster", "db_manager"]
Expand Down Expand Up @@ -156,21 +218,26 @@ def build_casts(cls, model):
annotations = {
k: v for k, v in annotations.items() if k not in cls.IGNORE_CASTS
}
from .fields import FieldDescriptor
from .fields import ModelField, FieldDescriptor

# 1. Collect all potential fields (annotations + descriptors)
all_field_names = set(annotations.keys())
descriptors = {}
for name, attr in cls.__dict__.items():
if isinstance(attr, FieldDescriptor):
for name, attr in model.__dict__.items():
if isinstance(attr, (FieldDescriptor, ModelField)):
all_field_names.add(name)
descriptors[name] = attr

casts = {}
for field_name in all_field_names:
# 2. Get Type Hint and FieldInfo
typ = annotations.get(field_name) or "str"
descriptor = descriptors.get(field_name, None)

# AttributeField: use the type annotation as the model class
if isinstance(descriptor, ModelField):
casts[field_name] = ModelCast(model_class=typ)
continue

field_info = (
descriptor.field_info
if isinstance(descriptor, FieldDescriptor)
Expand Down Expand Up @@ -204,12 +271,29 @@ def normalize_type(t):
or t is Carbon
):
return "date"
if t is datetime.time:
return "time"
if t is datetime.timedelta:
return "timedelta"
if isinstance(t, type):
if issubclass(t, Enum) or hasattr(t, "get") or hasattr(t, "set"):
return t

return "str"

@staticmethod
def _apply_default(cast: "BaseCast"):
"""Return the Field default/default_factory value, or None if none is set."""
from pydantic_core import PydanticUndefined

if cast.config is None:
return None
if cast.config.default is not PydanticUndefined:
return cast.config.default
if cast.config.default_factory is not None:
return cast.config.default_factory()
return None

def get(self, attribute: str, value: Any) -> Any:
if attribute not in self.casts:
return value
Expand All @@ -220,7 +304,10 @@ def get(self, attribute: str, value: Any) -> Any:
return str(value) if value is not None else None

if isinstance(cast, BaseCast):
return cast.get(value)
result = cast.get(value)
if result is None:
result = self._apply_default(cast)
return result

if isinstance(cast, type):
if issubclass(cast, Enum):
Expand Down
13 changes: 13 additions & 0 deletions fastapi_startkit/src/fastapi_startkit/masoniteorm/models/fields.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,19 @@ def Field(*args, **kwargs) -> Any:
return FieldDescriptor(BaseField(*args, **kwargs))


class ModelField:
def __set_name__(self, owner, name):
self.name = name

def __get__(self, instance, owner):
if instance is None:
return self
return instance.get_attribute(self.name)

def __set__(self, instance, value):
instance.set_attribute(self.name, value)


class DateTimeField:
def __init__(self, fmt: str = "YYYY-MM-DD HH:mm:ss", tz: str = "UTC"):
self.format = fmt
Expand Down
3 changes: 3 additions & 0 deletions fastapi_startkit/src/fastapi_startkit/storage/__init__.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,6 @@
from .storage import Storage
from .config import S3Config, LocalDiskConfig, PublicDiskConfig
from .drivers.fake import FakeDriver
from .providers.provider import StorageProvider

__all__ = ["Storage", "S3Config", "LocalDiskConfig", "PublicDiskConfig", "FakeDriver", "StorageProvider"]
9 changes: 9 additions & 0 deletions fastapi_startkit/tests/masoniteorm/fixtures/casts.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,9 @@
from typing import Optional

from pydantic import BaseModel

class Address(BaseModel):
address: Optional[str] = None
city: Optional[str] = None
state: Optional[str] = None
country: Optional[str] = None
5 changes: 5 additions & 0 deletions fastapi_startkit/tests/masoniteorm/fixtures/migration.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,11 @@ async def migrate(schema: Schema) -> None:
table.string("email").unique()
table.boolean("is_admin").default(False)
table.timestamp("email_verified_at").nullable()
table.date("date_of_birth").nullable()
table.decimal("session_duration").nullable()
table.string("punch_in_time").nullable()
table.json("preferences").nullable()
table.text("address").nullable()
table.timestamps()

async with await schema.create_table_if_not_exists("profiles") as table:
Expand Down
14 changes: 11 additions & 3 deletions fastapi_startkit/tests/masoniteorm/fixtures/model.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,8 @@
from datetime import datetime, timedelta, time, date

from fastapi_startkit.carbon.carbon import Carbon
from fastapi_startkit.masoniteorm.models.fields import Field, DateTimeField
from tests.masoniteorm.fixtures.casts import Address
from fastapi_startkit.masoniteorm import ModelField, Field
from fastapi_startkit.masoniteorm.relationships import (
HasOne,
BelongsTo,
Expand All @@ -9,15 +12,20 @@
HasOneThrough,
MorphTo,
)
from fastapi_startkit.masoniteorm.models.model import Model
from fastapi_startkit.masoniteorm import Model


class User(Model):
id: int
name: str
email: str
email_verified_at: Carbon = DateTimeField(fmt="%Y-%m-%d %H:%M:%S", tz="UTC")
email_verified_at: datetime
date_of_birth: date
session_duration: timedelta
punch_in_time: time = Field(default=time(12, 0, 0))
is_admin: bool
preferences: dict
address: Address = ModelField()

profile: "Profile" = HasOne("Profile", "user_id", "id")
articles: "Articles" = HasMany("Articles", "id", "user_id")
Expand Down
18 changes: 13 additions & 5 deletions fastapi_startkit/tests/masoniteorm/fixtures/seeder.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,19 @@
from .model import User, Profile, Articles, Logo, Country, Port, IncomingShipment, Like, Product
from .model import Articles, Country, IncomingShipment, Like, Logo, Port, Product, Profile, User


async def seeder():
user = await User.query().create(
{"email": "admin@admin.com", "name": "Joe", "is_admin": True}
{
"email": "admin@admin.com",
"name": "Joe",
"is_admin": True,
"email_verified_at": "2024-01-15 08:00:00",
"date_of_birth": "1990-06-15",
"session_duration": 3600.0,
"punch_in_time": "09:00:00",
"preferences": {"theme": "dark", "language": "en"},
"address": {"address": "123 Main St", "city": "Sydney", "state": "NSW", "country": "Australia"},
}
)
await Profile.create({"name": "Joe Profile", "user_id": user.id})
article = await Articles.create(
Expand All @@ -13,9 +23,7 @@ async def seeder():
"published_date": "2020-01-01 00:00:00",
}
)
await Logo.create(
{"article_id": article.id, "published_date": "2020-01-01 00:00:00"}
)
await Logo.create({"article_id": article.id, "published_date": "2020-01-01 00:00:00"})
product = await Product.create({"name": "Widget"})
await Like.create({"likeable_type": "article", "likeable_id": article.id})
await Like.create({"likeable_type": "product", "likeable_id": product.id})
Expand Down
Loading
Loading