From 4b7566e0181f135c27eb15900490cfe5e5508b82 Mon Sep 17 00:00:00 2001 From: Bedram Tamang Date: Sun, 27 Sep 2026 02:44:26 -0700 Subject: [PATCH 1/2] Annotate Model/QueryBuilder typing gaps and fix generic binding Parameterize every unannotated columns/values/dict/list parameter on the public Model and QueryBuilder API (first, first_or_fail, find, find_or_fail, get, get_models, where_in, where_not_in, update, create, insert, fill, etc.). The systemic QueryBuilder[Unknown] issue was caused by bare, unparameterized "QueryBuilder" return annotations overriding pyright's body-inferred type. Replace them with QueryBuilder[Self] (Model classmethods) and Self (QueryBuilder instance methods) so chained calls like User.where(...).first(), User.find(1), and User.where_in(...) resolve to QueryBuilder[User] / User | None instead of Unknown. Also tighten a few internal QueryBuilder annotations (_columns, where_in's duck-typed Collection unwrap, select()'s *args) that the wider parameter types would otherwise leave inconsistent, and add a missing columns normalization guard in get_models() matching the existing pattern in first()/get(). No runtime behavior changes. Co-Authored-By: Claude Sonnet 5 --- .../masoniteorm/models/builder.py | 115 +++++++++-------- .../masoniteorm/models/model.py | 122 +++++++++--------- 2 files changed, 121 insertions(+), 116 deletions(-) diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/builder.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/builder.py index e8841043..1d7e2202 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/builder.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/builder.py @@ -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 ( @@ -22,6 +22,7 @@ from fastapi_startkit.masoniteorm.collection import Collection from fastapi_startkit.masoniteorm.connections.connection import Connection from fastapi_startkit.masoniteorm.models.model import Model + from fastapi_startkit.masoniteorm.query.grammars.BaseGrammar import BaseGrammar TModel = TypeVar("TModel", bound="Model") @@ -67,13 +68,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 @@ -91,35 +92,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: @@ -129,14 +130,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: @@ -149,7 +150,7 @@ 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) @@ -157,7 +158,7 @@ async def first_or_fail(self, columns=None) -> "TModel": 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 = [] @@ -170,7 +171,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) @@ -180,15 +183,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 @@ -218,11 +221,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),) @@ -232,67 +235,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 @@ -322,27 +325,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: @@ -351,7 +354,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: @@ -379,7 +382,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() @@ -492,7 +495,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 @@ -538,20 +541,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,) @@ -574,32 +577,32 @@ 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") @@ -607,7 +610,7 @@ def where_has(self, relation: str, callback=None) -> "QueryBuilder": 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") @@ -616,7 +619,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()] diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/model.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/model.py index fea99d10..a7c4d855 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/model.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/model.py @@ -21,6 +21,8 @@ from fastapi_startkit.masoniteorm.observers import ObservesEvents if TYPE_CHECKING: + from collections.abc import Callable, Iterable + from fastapi_startkit.masoniteorm.models.builder import QueryBuilder, WhereGroup @@ -67,7 +69,7 @@ def __init_subclass__(cls, **kwargs): created_at: Carbon = CreatedAtField(fmt="%Y-%m-%d %H:%M:%S", tz="UTC") # pyright: ignore[reportAssignmentType] updated_at: Carbon = UpdatedAtField(fmt="%Y-%m-%d %H:%M:%S", tz="UTC") # pyright: ignore[reportAssignmentType] - def __init__(self, attributes: dict | None = None, **kwargs): + def __init__(self, attributes: dict[str, Any] | None = None, **kwargs: Any): super().__init__(attributes, **kwargs) self.connection = getattr(self.__class__, "__connection__", "default") self._global_scopes = {} @@ -87,14 +89,14 @@ def is_created(self) -> bool: """Returns True if this model has been persisted to the database.""" return self._exists - def all_attributes(self) -> dict: + def all_attributes(self) -> dict[str, Any]: """Returns all model attributes (original + dirty).""" return self.get_attributes() - def get_builder(self): + def get_builder(self) -> "QueryBuilder[Self]": return self.new_query() - def add_relation(self, data: dict): + def add_relation(self, data: dict[str, Any]): self._relationship.update(data) @property @@ -106,7 +108,7 @@ def get_related(self, key: str): return getattr(self.__class__, key) @classmethod - def with_(cls, *eagers) -> "QueryBuilder": + def with_(cls, *eagers) -> "QueryBuilder[Self]": return cls.query().with_(*eagers) @overload @@ -134,135 +136,135 @@ def where(cls, column: str | dict[str, Any] | WhereGroup[Self], *args: Any) -> Q return cls.query().where(column, *args) @classmethod - def or_where(cls, column, *args) -> "QueryBuilder": + def or_where(cls, column: str, *args: Any) -> "QueryBuilder[Self]": return cls.query().or_where(column, *args) @classmethod - def where_null(cls, column: str) -> "QueryBuilder": + def where_null(cls, column: str) -> "QueryBuilder[Self]": return cls.query().where_null(column) @classmethod - def where_not_null(cls, column: str) -> "QueryBuilder": + def where_not_null(cls, column: str) -> "QueryBuilder[Self]": return cls.query().where_not_null(column) @classmethod - def or_where_null(cls, column: str) -> "QueryBuilder": + def or_where_null(cls, column: str) -> "QueryBuilder[Self]": return cls.query().or_where_null(column) @classmethod - def or_where_not_null(cls, column: str) -> "QueryBuilder": + def or_where_not_null(cls, column: str) -> "QueryBuilder[Self]": return cls.query().or_where_not_null(column) @classmethod - def where_raw(cls, expression: str, bindings=()) -> "QueryBuilder": + def where_raw(cls, expression: str, bindings: tuple[Any, ...] = ()) -> "QueryBuilder[Self]": return cls.query().where_raw(expression, bindings) @classmethod - def or_where_raw(cls, expression: str, bindings=()) -> "QueryBuilder": + def or_where_raw(cls, expression: str, bindings: tuple[Any, ...] = ()) -> "QueryBuilder[Self]": return cls.query().or_where_raw(expression, bindings) @classmethod - def where_in(cls, column: str, values) -> "QueryBuilder": + def where_in(cls, column: str, values: Iterable[Any]) -> "QueryBuilder[Self]": return cls.query().where_in(column, values) @classmethod - def where_not_in(cls, column: str, values) -> "QueryBuilder": + def where_not_in(cls, column: str, values: Iterable[Any]) -> "QueryBuilder[Self]": return cls.query().where_not_in(column, values) @classmethod - def select(cls, *args) -> "QueryBuilder": + def select(cls, *args: str | list[str]) -> "QueryBuilder[Self]": return cls.query().select(*args) @classmethod - def limit(cls, limit: int) -> "QueryBuilder": + def limit(cls, limit: int) -> "QueryBuilder[Self]": return cls.query().limit(limit) @classmethod - def offset(cls, offset: int) -> "QueryBuilder": + def offset(cls, offset: int) -> "QueryBuilder[Self]": return cls.query().offset(offset) @classmethod - def order_by(cls, column: str, direction: str = "asc") -> "QueryBuilder": + def order_by(cls, column: str, direction: str = "asc") -> "QueryBuilder[Self]": return cls.query().order_by(column, direction) @classmethod - def order_by_raw(cls, expression: str) -> "QueryBuilder": + def order_by_raw(cls, expression: str) -> "QueryBuilder[Self]": return cls.query().order_by_raw(expression) @classmethod - def latest(cls, column: str = "created_at") -> "QueryBuilder": + def latest(cls, column: str = "created_at") -> "QueryBuilder[Self]": return cls.query().latest(column) @classmethod - def oldest(cls, column: str = "created_at") -> "QueryBuilder": + def oldest(cls, column: str = "created_at") -> "QueryBuilder[Self]": return cls.query().oldest(column) @classmethod - def group_by(cls, column: str) -> "QueryBuilder": + def group_by(cls, column: str) -> "QueryBuilder[Self]": return cls.query().group_by(column) @classmethod - def group_by_raw(cls, expression: str) -> "QueryBuilder": + def group_by_raw(cls, expression: str) -> "QueryBuilder[Self]": return cls.query().group_by_raw(expression) @classmethod - def having(cls, column: str, equality: str, value) -> "QueryBuilder": + def having(cls, column: str, equality: str, value: Any) -> "QueryBuilder[Self]": return cls.query().having(column, equality, value) @classmethod - def between(cls, column: str, low, high) -> "QueryBuilder": + def between(cls, column: str, low: Any, high: Any) -> "QueryBuilder[Self]": return cls.query().between(column, low, high) @classmethod - def not_between(cls, column: str, low, high) -> "QueryBuilder": + def not_between(cls, column: str, low: Any, high: Any) -> "QueryBuilder[Self]": return cls.query().not_between(column, low, high) @classmethod - def distinct(cls) -> "QueryBuilder": + def distinct(cls) -> "QueryBuilder[Self]": return cls.query().distinct() @classmethod - def join(cls, table: str, column1: str, equality: str, column2: str, clause: str = "join") -> "QueryBuilder": + def join(cls, table: str, column1: str, equality: str, column2: str, clause: str = "join") -> "QueryBuilder[Self]": return cls.query().join(table, column1, equality, column2, clause) @classmethod - def left_join(cls, table: str, column1: str, equality: str, column2: str) -> "QueryBuilder": + def left_join(cls, table: str, column1: str, equality: str, column2: str) -> "QueryBuilder[Self]": return cls.query().left_join(table, column1, equality, column2) @classmethod - def right_join(cls, table: str, column1: str, equality: str, column2: str) -> "QueryBuilder": + def right_join(cls, table: str, column1: str, equality: str, column2: str) -> "QueryBuilder[Self]": return cls.query().right_join(table, column1, equality, column2) @classmethod - def where_column(cls, column1: str, column2: str) -> "QueryBuilder": + def where_column(cls, column1: str, column2: str) -> "QueryBuilder[Self]": return cls.query().where_column(column1, column2) @classmethod - def when(cls, condition, callback) -> "QueryBuilder": + def when(cls, condition: Any, callback: Callable[["QueryBuilder[Self]"], Any]) -> "QueryBuilder[Self]": return cls.query().when(condition, callback) @classmethod - def where_exists(cls, builder: "QueryBuilder") -> "QueryBuilder": + def where_exists(cls, builder: "QueryBuilder[Any]") -> "QueryBuilder[Self]": return cls.query().where_exists(builder) @classmethod - def or_where_exists(cls, builder: "QueryBuilder") -> "QueryBuilder": + def or_where_exists(cls, builder: "QueryBuilder[Any]") -> "QueryBuilder[Self]": return cls.query().or_where_exists(builder) @classmethod - def where_has(cls, relation: str, callback=None) -> "QueryBuilder": + def where_has(cls, relation: str, callback: Callable[..., Any] | None = None) -> "QueryBuilder[Self]": return cls.query().where_has(relation, callback) @classmethod - def or_where_has(cls, relation: str, callback=None) -> "QueryBuilder": + def or_where_has(cls, relation: str, callback: Callable[..., Any] | None = None) -> "QueryBuilder[Self]": return cls.query().or_where_has(relation, callback) @classmethod - async def find(cls, primary_key: str | int, columns=None): + async def find(cls, primary_key: str | int, columns: list[str] | str | None = None) -> Self | None: return await cls.query().find(primary_key, columns) @classmethod - async def find_or_fail(cls, primary_key: str | int, columns=None) -> Self: + async def find_or_fail(cls, primary_key: str | int, columns: list[str] | str | None = None) -> Self: """Fetch the record matching ``primary_key``. Raises: @@ -271,11 +273,11 @@ async def find_or_fail(cls, primary_key: str | int, columns=None) -> Self: return await cls.query().find_or_fail(primary_key, columns) @classmethod - async def first_or_fail(cls, columns=None): + async def first_or_fail(cls, columns: list[str] | str | None = None) -> Self: return await cls.query().first_or_fail(columns) @classmethod - async def first(cls, columns=None): + async def first(cls, columns: list[str] | str | None = None) -> Self | None: return await cls.query().first(columns) @classmethod @@ -322,7 +324,7 @@ def set_connection(self, connection: str): def get_connection_name(self): return self.connection - def new_model_instance(self, attributes=None, exists=False): + def new_model_instance(self, attributes: dict[str, Any] | None = None, exists: bool = False) -> Self: if attributes is None: attributes = {} model = self.__class__() @@ -340,22 +342,22 @@ def resolve_db_manager(cls) -> DatabaseManager: def new_query(self) -> "QueryBuilder[Self]": return self.resolve_db_manager().connection(self.connection).query().set_model(self) - def hydrate(self, items): + def hydrate(self, items: Iterable[dict[str, Any]]) -> Collection[Self]: instance = self.new_model_instance() - items = [instance.new_from_builder(item) for item in items] + models = [instance.new_from_builder(item) for item in items] - return instance.new_collection(items) + return instance.new_collection(models) - def new_collection(self, models: list): + def new_collection(self, models: list[Self]) -> Collection[Self]: collection = Collection(items=models) collection.with_relationship_autoloading() return collection - def new_from_builder(self, attributes: dict, connection: str | None = None): - model = self.new_model_instance([], exists=True) + def new_from_builder(self, attributes: dict[str, Any], connection: str | None = None) -> Self: + model = self.new_model_instance({}, exists=True) model.set_raw_attributes(attributes, True) model.set_connection(connection or self.get_connection_name()) @@ -363,7 +365,7 @@ def new_from_builder(self, attributes: dict, connection: str | None = None): return model - def __getattr__(self, attribute): + def __getattr__(self, attribute: str) -> Any: return self.get_attribute(attribute) @classmethod @@ -371,25 +373,25 @@ def query(cls) -> "QueryBuilder[Self]": return cls().new_query() @classmethod - async def first_or_create(cls, search: dict, attributes: dict | None = None) -> "Model": + async def first_or_create(cls, search: dict[str, Any], attributes: dict[str, Any] | None = None) -> Self: return await cls.query().first_or_create(search, attributes) @classmethod - async def update_or_create(cls, search: dict, attributes: dict | None = None) -> "Model": + async def update_or_create(cls, search: dict[str, Any], attributes: dict[str, Any] | None = None) -> Self: return await cls.query().update_or_create(search, attributes) @classmethod - async def create(cls, attributes: dict): + async def create(cls, attributes: dict[str, Any]) -> Self: instance = cls().new_model_instance(attributes) await instance.save() return instance @classmethod - async def insert(cls, values: dict | list) -> int | None: + async def insert(cls, values: dict[str, Any] | list[dict[str, Any]]) -> int | None: return await cls.query().insert(values) - async def update(self, attributes: dict) -> bool: + async def update(self, attributes: dict[str, Any]) -> bool: if not self._exists: return False @@ -418,13 +420,13 @@ async def touch(self) -> bool: return True - def fill(self, attributes: dict) -> "Model": + def fill(self, attributes: dict[str, Any]) -> Self: for key, value in attributes.items(): if key in self.__fillable__: self.set_attribute(key, value) return self - async def save(self, options: dict | None = None): + async def save(self, options: dict[str, Any] | None = None) -> bool: query = self.new_query() self.observe_events(self, "saving") @@ -439,11 +441,11 @@ async def save(self, options: dict | None = None): return saved - def finish_saving(self, options: dict | None = None): + def finish_saving(self, options: dict[str, Any] | None = None) -> None: self.observe_events(self, "saved") self.sync_original() - async def perform_insert(self, query) -> bool: + async def perform_insert(self, query: "QueryBuilder[Self]") -> bool: attributes = self.get_attributes_for_insert() """if the model set auto incrementing, we need to set back the primary key to the inserted id.""" @@ -460,7 +462,7 @@ async def perform_insert(self, query) -> bool: self.observe_events(self, "created") return True - async def perform_update(self, query) -> bool: + async def perform_update(self, query: "QueryBuilder[Self]") -> bool: dirty = self.get_dirty() if not dirty: return True @@ -476,10 +478,10 @@ def sync_original(self): self._dirty_attributes = {} self._original = dict(self._attributes) - def get_attributes(self) -> dict: + def get_attributes(self) -> dict[str, Any]: return {**self._attributes, **self._dirty_attributes} - def serialize(self) -> dict: + def serialize(self) -> dict[str, Any]: return self.get_attributes() def get_table_name(self): From 7fb3c1345f61257a0e8ec6d884d0cada3f937ffc Mon Sep 17 00:00:00 2001 From: Bedram Tamang Date: Sun, 27 Sep 2026 03:00:31 -0700 Subject: [PATCH 2/2] Remove unused BaseGrammar import flagged by ruff grammar/processor on QueryBuilder.__init__ are kept as Any: the base Connection.get_query_grammar()/get_post_processor() type as type[BaseGrammar] | None (only concrete in subclass overrides), and threading that Optional through reintroduces reportOptionalCall on get_grammar()/insert_get_id()/update() plus unrelated cascading diagnostics on _bindings/_limit/_offset with no net typing benefit. Any keeps the diagnostics at the confirmed 0-error baseline, so the now-unused BaseGrammar import is dropped instead of forcing it into use. Co-Authored-By: Claude Sonnet 5 --- .../src/fastapi_startkit/masoniteorm/models/builder.py | 1 - 1 file changed, 1 deletion(-) diff --git a/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/builder.py b/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/builder.py index 1d7e2202..1739497b 100644 --- a/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/builder.py +++ b/fastapi_startkit/src/fastapi_startkit/masoniteorm/models/builder.py @@ -22,7 +22,6 @@ from fastapi_startkit.masoniteorm.collection import Collection from fastapi_startkit.masoniteorm.connections.connection import Connection from fastapi_startkit.masoniteorm.models.model import Model - from fastapi_startkit.masoniteorm.query.grammars.BaseGrammar import BaseGrammar TModel = TypeVar("TModel", bound="Model")