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
114 changes: 58 additions & 56 deletions fastapi_startkit/src/fastapi_startkit/masoniteorm/models/builder.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
import inspect
from collections.abc import Callable
from collections.abc import Callable, Iterable
from typing import TYPE_CHECKING, Any, Generic, Self, TypeVar, overload

from fastapi_startkit.masoniteorm.expressions.expressions import (
Expand Down Expand Up @@ -67,13 +67,13 @@ class QueryBuilder(EagerLoadMixin, SupportMixin, Generic[TModel]):
"!~~*",
]

def __init__(self, connection: "Connection", grammar, processor):
def __init__(self, connection: "Connection", grammar: Any, processor: Any):
super().__init__()
self.connection = connection
self.grammar = grammar
self.processor = processor

self._columns = []
self._columns: list[Any] = []
self._table = ""
self._limit = False
self._offset = False
Expand All @@ -91,35 +91,35 @@ def __init__(self, connection: "Connection", grammar, processor):
self._global_scopes = {}
self._action = "select"

def set_action(self, action: str) -> "QueryBuilder":
def set_action(self, action: str) -> "Self":
self._action = action
return self

def set_model(self, model) -> "QueryBuilder":
def set_model(self, model: "TModel") -> "Self":
self._model = model
self._table = model.get_table_name()
self._global_scopes = model._global_scopes
return self

def with_(self, *eagers) -> "QueryBuilder":
def with_(self, *eagers) -> "Self":
self._eager_relation.register(eagers)
return self

def get_table_name(self) -> str:
return self._table

def table(self, table: str) -> "QueryBuilder":
def table(self, table: str) -> "Self":
self._table = table
return self

def where_in(self, column: str, values) -> "QueryBuilder":
def where_in(self, column: str, values: Iterable[Any]) -> "Self":
if hasattr(values, "_items"):
values = values._items
values = getattr(values, "_items")
values = list(values) if not isinstance(values, list) else values
self._wheres.append(QueryExpression(column, "IN", values))
return self

def select(self, *args) -> "QueryBuilder":
def select(self, *args: "str | list[str]") -> "Self":
for arg in args:
if isinstance(arg, list):
for column in arg:
Expand All @@ -129,14 +129,14 @@ def select(self, *args) -> "QueryBuilder":
self._columns += (SelectExpression(column),)
return self

def limit(self, limit: int) -> "QueryBuilder":
def limit(self, limit: int) -> "Self":
self._limit = limit
return self

async def find(self, primary_key: str | int, columns=None) -> "TModel | None":
async def find(self, primary_key: str | int, columns: "list[str] | str | None" = None) -> "TModel | None":
return await self.where(self._model.__primary_key__, primary_key).first(columns)

async def find_or_fail(self, primary_key: str | int, columns=None) -> "TModel":
async def find_or_fail(self, primary_key: str | int, columns: "list[str] | str | None" = None) -> "TModel":
"""Return the record matching ``primary_key``.

Raises:
Expand All @@ -149,15 +149,15 @@ async def find_or_fail(self, primary_key: str | int, columns=None) -> "TModel":
raise ModelNotFoundException(f"{type(self._model).__name__} with primary key {primary_key!r} not found.")
return result

async def first_or_fail(self, columns=None) -> "TModel":
async def first_or_fail(self, columns: "list[str] | str | None" = None) -> "TModel":
from fastapi_startkit.masoniteorm.exceptions import ModelNotFoundException

result = await self.first(columns)
if result is None:
raise ModelNotFoundException(f"{type(self._model).__name__} not found.")
return result

async def first(self, columns=None) -> "TModel | None":
async def first(self, columns: "list[str] | str | None" = None) -> "TModel | None":
if not columns:
columns = []

Expand All @@ -170,7 +170,9 @@ async def get(self, columns: "list[str] | str | None" = None) -> "Collection[TMo
columns = []
return await self.get_models(columns)

async def get_models(self, columns=None):
async def get_models(self, columns: "list[str] | str | None" = None) -> "Collection[TModel]":
if not columns:
columns = []
self.select(columns)
models = await self.connection.select(self.to_qmark(), list(self.get_bindings()))
collection = self._model.hydrate(models)
Expand All @@ -180,15 +182,15 @@ async def get_models(self, columns=None):

return collection

def get_bindings(self) -> tuple:
def get_bindings(self) -> tuple[Any, ...]:
return self._bindings

def run_scopes(self) -> "QueryBuilder":
def run_scopes(self) -> "Self":
for name, scope in self._global_scopes.get(self._action, {}).items():
scope(self)
return self

def without_global_scopes(self) -> "QueryBuilder":
def without_global_scopes(self) -> "Self":
self._global_scopes = {}
return self

Expand Down Expand Up @@ -218,11 +220,11 @@ def to_sql(self) -> str:
self.run_scopes()
return self.get_grammar().compile(self._action).to_sql()

def offset(self, offset: int) -> "QueryBuilder":
def offset(self, offset: int) -> "Self":
self._offset = offset
return self

def order_by(self, column, direction: str = "asc") -> "QueryBuilder":
def order_by(self, column: "str | QueryBuilder[Any]", direction: str = "asc") -> "Self":
direction = direction.upper()
if isinstance(column, QueryBuilder):
self._order_by += (OrderByExpression(None, direction, builder=column),)
Expand All @@ -232,67 +234,67 @@ def order_by(self, column, direction: str = "asc") -> "QueryBuilder":
self._order_by += (OrderByExpression(col, direction),)
return self

def order_by_raw(self, expression: str) -> "QueryBuilder":
def order_by_raw(self, expression: str) -> "Self":
self._order_by += (OrderByExpression(expression, raw=True),)
return self

def latest(self, column: str = "created_at") -> "QueryBuilder":
def latest(self, column: str = "created_at") -> "Self":
return self.order_by(column, "desc")

def oldest(self, column: str = "created_at") -> "QueryBuilder":
def oldest(self, column: str = "created_at") -> "Self":
return self.order_by(column, "asc")

def group_by(self, column: str) -> "QueryBuilder":
def group_by(self, column: str) -> "Self":
for col in column.split(","):
col = col.strip()
self._group_by += (GroupByExpression(col),)
return self

def group_by_raw(self, expression: str) -> "QueryBuilder":
def group_by_raw(self, expression: str) -> "Self":
self._group_by += (GroupByExpression(expression, raw=True),)
return self

def having(self, column: str, equality: str, value) -> "QueryBuilder":
def having(self, column: str, equality: str, value: Any) -> "Self":
self._having += (HavingExpression(column, equality, value),)
return self

def where_null(self, column: str) -> "QueryBuilder":
def where_null(self, column: str) -> "Self":
self._wheres += (QueryExpression(column, "=", None, "NULL"),)
return self

def where_not_null(self, column: str) -> "QueryBuilder":
def where_not_null(self, column: str) -> "Self":
self._wheres += (QueryExpression(column, "=", None, "NOT NULL"),)
return self

def or_where_null(self, column: str) -> "QueryBuilder":
def or_where_null(self, column: str) -> "Self":
self._wheres += (QueryExpression(column, "=", None, "NULL", keyword="or"),)
return self

def or_where_not_null(self, column: str) -> "QueryBuilder":
def or_where_not_null(self, column: str) -> "Self":
self._wheres += (QueryExpression(column, "=", None, "NOT NULL", keyword="or"),)
return self

def where_not_in(self, column: str, values) -> "QueryBuilder":
def where_not_in(self, column: str, values: Iterable[Any]) -> "Self":
values = list(values) if not isinstance(values, list) else values
self._wheres.append(QueryExpression(column, "NOT IN", values))
return self

def between(self, column: str, low, high) -> "QueryBuilder":
def between(self, column: str, low: Any, high: Any) -> "Self":
self._wheres += (BetweenExpression(column, low, high, "BETWEEN"),)
return self

def not_between(self, column: str, low, high) -> "QueryBuilder":
def not_between(self, column: str, low: Any, high: Any) -> "Self":
self._wheres += (BetweenExpression(column, low, high, "NOT BETWEEN"),)
return self

def left_join(self, table: str, column1: str, equality: str, column2: str) -> "QueryBuilder":
def left_join(self, table: str, column1: str, equality: str, column2: str) -> "Self":
return self.join(table, column1, equality, column2, clause="left")

def right_join(self, table: str, column1: str, equality: str, column2: str) -> "QueryBuilder":
def right_join(self, table: str, column1: str, equality: str, column2: str) -> "Self":
# SQLite doesn't support RIGHT JOIN — use left join as fallback
return self.join(table, column1, equality, column2, clause="right")

def distinct(self) -> "QueryBuilder":
def distinct(self) -> "Self":
self._distinct = True
return self

Expand Down Expand Up @@ -322,27 +324,27 @@ async def min(self, column: str):
async def avg(self, column: str):
return await self.aggregate("AVG", column)

async def delete(self, column=None, value=None):
async def delete(self, column: str | None = None, value: Any = None):
if column is not None:
self.where(column, value)
self.set_action("delete")
sql = self.to_qmark()
return await self.connection.delete(sql, list(self.get_bindings()))

async def create(self, attributes: dict):
async def create(self, attributes: dict[str, Any]) -> "TModel":
model = self._model.new_model_instance(attributes)
await model.save()

return model

async def first_or_create(self, search: dict, attributes: dict | None = None):
async def first_or_create(self, search: dict[str, Any], attributes: dict[str, Any] | None = None) -> "TModel":
instance = await self.where(search).first()
if instance is not None:
return instance

return await self.create({**(attributes or {}), **search})

async def update_or_create(self, search: dict, attributes: dict | None = None):
async def update_or_create(self, search: dict[str, Any], attributes: dict[str, Any] | None = None) -> "TModel":
instance = await self.where(search).first()
if instance is not None:
if attributes:
Expand All @@ -351,7 +353,7 @@ async def update_or_create(self, search: dict, attributes: dict | None = None):

return await self.create({**(attributes or {}), **search})

async def insert(self, values: dict | list) -> int | None:
async def insert(self, values: dict[str, Any] | list[dict[str, Any]]) -> int | None:
self.set_action("bulk_create")

if not values:
Expand Down Expand Up @@ -379,7 +381,7 @@ async def insert_get_id(

return await self.connection.insert_get_id(sql, bindings)

async def update(self, values: dict) -> int:
async def update(self, values: dict[str, Any]) -> int:
updates = [UpdateQueryExpression(col, val) for col, val in values.items()]
grammar = self.grammar()
sql = grammar._compile_update(query=self, values=updates, qmark=True).to_sql()
Expand Down Expand Up @@ -492,7 +494,7 @@ def new(self):
# with .table(...) as usual.
return self.connection.query().table(self._table)

def invalid_operator(self, operator):
def invalid_operator(self, operator: Any) -> bool:
"""Determine whether an operator is not supported by the builder."""
return not isinstance(operator, str) or operator.lower() not in self.operators

Expand Down Expand Up @@ -538,20 +540,20 @@ def where(self, column: "str | dict[str, Any] | WhereGroup[TModel]", *args: Any)
self._wheres += ((QueryExpression(column, operator, value, "value")),)
return self

def or_where(self, column, *args) -> "QueryBuilder":
def or_where(self, column: str, *args: Any) -> "Self":
operator, value = self._extract_operator_value(*args)
self._wheres += ((QueryExpression(column, operator, value, "value", keyword="or")),)
return self

def where_raw(self, expression: str, bindings=()) -> "QueryBuilder":
def where_raw(self, expression: str, bindings: tuple[Any, ...] = ()) -> "Self":
self._wheres += (QueryExpression(expression, "=", None, raw=True, bindings=bindings),)
return self

def or_where_raw(self, expression: str, bindings=()) -> "QueryBuilder":
def or_where_raw(self, expression: str, bindings: tuple[Any, ...] = ()) -> "Self":
self._wheres += (QueryExpression(expression, "=", None, raw=True, keyword="or", bindings=bindings),)
return self

def join(self, table: str, column1: str, equality: str, column2: str, clause: str = "join") -> "QueryBuilder":
def join(self, table: str, column1: str, equality: str, column2: str, clause: str = "join") -> "Self":
join_clause = JoinClause(table, clause=clause)
join_clause.on(column1, equality, column2)
self._joins += (join_clause,)
Expand All @@ -574,40 +576,40 @@ def _normalize_where_column(self, operator: str, column2: str | None):
)
return operator, column2

def where_column(self, column1: str, operator: str, column2: str | None = None) -> "QueryBuilder":
def where_column(self, column1: str, operator: str, column2: str | None = None) -> "Self":
"""Compare two columns (identifiers, never bound values), joined with AND."""
operator, column2 = self._normalize_where_column(operator, column2)
self._wheres += (QueryExpression(column1, operator, column2, "value_equals"),)
return self

def or_where_column(self, column1: str, operator: str, column2: str | None = None) -> "QueryBuilder":
def or_where_column(self, column1: str, operator: str, column2: str | None = None) -> "Self":
"""Compare two columns (identifiers, never bound values), joined with OR."""
operator, column2 = self._normalize_where_column(operator, column2)
self._wheres += (QueryExpression(column1, operator, column2, "value_equals", keyword="or"),)
return self

def when(self, condition, callback) -> "QueryBuilder":
def when(self, condition: Any, callback: Callable[["Self"], Any]) -> "Self":
if condition:
callback(self)
return self

def where_exists(self, builder: "QueryBuilder") -> "QueryBuilder":
def where_exists(self, builder: "QueryBuilder[Any]") -> "Self":
self._wheres += (QueryExpression(None, "EXISTS", SubSelectExpression(builder)),)
return self

def or_where_exists(self, builder: "QueryBuilder") -> "QueryBuilder":
def or_where_exists(self, builder: "QueryBuilder[Any]") -> "Self":
self._wheres += (QueryExpression(None, "EXISTS", SubSelectExpression(builder), keyword="or"),)
return self

def where_has(self, relation: str, callback=None) -> "QueryBuilder":
def where_has(self, relation: str, callback: Callable[..., Any] | None = None) -> "Self":
related = getattr(self._model.__class__, relation)
if callback:
related.query_where_exists(self, callback, method="where_exists")
else:
related.query_has(self, method="where_exists")
return self

def or_where_has(self, relation: str, callback=None) -> "QueryBuilder":
def or_where_has(self, relation: str, callback: Callable[..., Any] | None = None) -> "Self":
related = getattr(self._model.__class__, relation)
if callback:
related.query_where_exists(self, callback, method="or_where_exists")
Expand All @@ -616,7 +618,7 @@ def or_where_has(self, relation: str, callback=None) -> "QueryBuilder":
return self

@classmethod
def clean_bindings(cls, values):
def clean_bindings(cls, values: dict[str, Any] | list[dict[str, Any]]) -> list[Any]:
if isinstance(values, dict):
values = [values]
return [val for row in values for val in row.values()]
Loading
Loading