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 stubs/peewee/METADATA.toml
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
version = "4.0.8"
version = "4.1.2"
upstream-repository = "https://github.com/coleifer/peewee"
# We're not providing stubs for all playhouse modules right now
# https://github.com/python/typeshed/pull/11731#issuecomment-2065729058
Expand Down
108 changes: 65 additions & 43 deletions stubs/peewee/peewee.pyi
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@ _Model: TypeAlias = Model
_M = TypeVar("_M", bound=Model, default=Model)
# __get__/__set__ value type. Bare Field defaults to Field[Any].
_V = TypeVar("_V", default=Any)
_DatabaseType: TypeAlias = Database | DatabaseProxy

# Common field kwargs, Unpack-ed into the field __new__ overloads.
@type_check_only
Expand Down Expand Up @@ -80,7 +81,7 @@ SNAKE_CASE_STEP2: Final[re.Pattern[str]]
IDENTIFIER_RE: Final[re.Pattern[str]]

def make_identifier(s: str) -> str: ...
def chunked(it, n) -> Generator[list[Incomplete]]: ...
def chunked(it: Iterable[_T], n: int) -> Generator[list[_T]]: ...

class _callable_context_manager:
def __call__(self, fn): ...
Expand Down Expand Up @@ -227,7 +228,7 @@ class BaseTable(Source):
class _BoundTableContext(_callable_context_manager):
table: Incomplete
database: Incomplete
def __init__(self, table, database) -> None: ...
def __init__(self, table, database: _DatabaseType) -> None: ...
def __enter__(self): ...
def __exit__(
self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: TracebackType | None
Expand All @@ -241,8 +242,8 @@ class Table(_HashableSource, BaseTable): # type: ignore[misc]
self, name, columns=None, primary_key=None, schema: str | None = None, alias=None, _model=None, _database=None
) -> None: ...
def clone(self) -> Table: ...
def bind(self, database=None) -> Self: ...
def bind_ctx(self, database=None) -> _BoundTableContext: ...
def bind(self, database: _DatabaseType | None = None) -> Self: ...
def bind_ctx(self, database: _DatabaseType | None = None) -> _BoundTableContext: ...
def select(self, *columns) -> Select: ...
def insert(self, insert=None, columns=None, **kwargs) -> Insert: ...
def replace(self, insert=None, columns=None, **kwargs): ...
Expand Down Expand Up @@ -438,7 +439,7 @@ class SQL(ColumnBase):
def __init__(self, sql, params=None) -> None: ...
def __sql__(self, ctx): ...

def Check(constraint, name=None) -> Node: ...
def Check(constraint: str, name: str | None = None) -> SQL | NodeList: ...
def Default(value) -> SQL: ...

class Function(ColumnBase):
Expand Down Expand Up @@ -565,17 +566,17 @@ class OnConflict(Node):
class BaseQuery(Node):
default_row_type: Incomplete
def __init__(self, _database=None, **kwargs) -> None: ...
def bind(self, database=None) -> Self: ...
def bind(self, database: _DatabaseType | None = None) -> Self: ...
def clone(self) -> Self: ...
def dicts(self, as_dict: bool = True) -> Self: ...
def tuples(self, as_tuple: bool = True) -> Self: ...
def namedtuples(self, as_namedtuple: bool = True) -> Self: ...
def objects(self, constructor=None) -> Self: ...
def __sql__(self, ctx) -> None: ...
def sql(self) -> tuple[str, list[Any]]: ... # Returns (sql, params), params are bound query values
def execute(self, database=None): ...
async def aexecute(self, database=None): ...
def iterator(self, database=None): ...
def execute(self, database: _DatabaseType | None = None): ...
async def aexecute(self, database: _DatabaseType | None = None): ...
def iterator(self, database: _DatabaseType | None = None): ...
def __iter__(self): ...
def __getitem__(self, value): ...
def __len__(self) -> int: ...
Expand Down Expand Up @@ -612,20 +613,20 @@ class SelectQuery(Query):
def select_from(self, *columns) -> Select: ...

class SelectBase(_HashableSource, Source, SelectQuery): # type: ignore[misc]
def peek(self, database=None, n: int = 1): ...
def first(self, database=None, n: int = 1): ...
def scalar(self, database=None, as_tuple: bool = False, as_dict: bool = False): ...
def scalars(self, database=None) -> Generator[Incomplete]: ...
def count(self, database=None, clear_limit: bool = False) -> int: ...
def exists(self, database=None) -> bool: ...
def get(self, database=None): ...
def peek(self, database: _DatabaseType | None = None, n: int = 1): ...
def first(self, database: _DatabaseType | None = None, n: int = 1): ...
def scalar(self, database: _DatabaseType | None = None, as_tuple: bool = False, as_dict: bool = False): ...
def scalars(self, database: _DatabaseType | None = None) -> Generator[Incomplete]: ...
def count(self, database: _DatabaseType | None = None, clear_limit: bool = False) -> int: ...
def exists(self, database: _DatabaseType | None = None) -> bool: ...
def get(self, database: _DatabaseType | None = None): ...

class CompoundSelectQuery(SelectBase):
lhs: Incomplete
op: Incomplete
rhs: Incomplete
def __init__(self, lhs, op, rhs) -> None: ...
def exists(self, database=None) -> bool: ...
def exists(self, database: _DatabaseType | None = None) -> bool: ...
def __sql__(self, ctx): ...

class Select(SelectBase):
Expand Down Expand Up @@ -672,8 +673,8 @@ class _WriteQuery(Query):
def cte(self, name, recursive: bool = False, columns=None, materialized=None) -> CTE: ...
def returning(self, *returning) -> Self: ...
def apply_returning(self, ctx): ...
def execute_returning(self, database): ...
def handle_result(self, database, cursor): ...
def execute_returning(self, database: _DatabaseType): ...
def handle_result(self, database: _DatabaseType, cursor): ...
def __sql__(self, ctx): ...

class Update(_WriteQuery):
Expand All @@ -697,7 +698,7 @@ class Insert(_WriteQuery):
def get_default_data(self): ...
def get_default_columns(self) -> list[Incomplete] | None: ...
def __sql__(self, ctx): ...
def handle_result(self, database, cursor): ...
def handle_result(self, database: _DatabaseType, cursor): ...

class Delete(_WriteQuery):
def __sql__(self, ctx): ...
Expand Down Expand Up @@ -753,12 +754,17 @@ class ColumnMetadata(NamedTuple):
primary_key: Incomplete
table: Incomplete
default: Incomplete
full_type: str | None = None
identity: bool = False

class ForeignKeyMetadata(NamedTuple):
column: Incomplete
dest_table: Incomplete
dest_column: Incomplete
table: Incomplete
name: str | None = None
on_delete: str | None = None
on_update: str | None = None

class ViewMetadata(NamedTuple):
name: Incomplete
Expand Down Expand Up @@ -814,9 +820,10 @@ class Database(_callable_context_manager):
autoconnect: Incomplete
thread_safe: Incomplete
connect_params: Incomplete
def __deepcopy__(self, memo: Any) -> Self: ...
def __init__(
self,
database,
database: str | None,
thread_safe: bool = True,
autorollback: bool = False,
field_types=None,
Expand All @@ -827,7 +834,7 @@ class Database(_callable_context_manager):
) -> None: ...
database: Incomplete
deferred: Incomplete
def init(self, database, **kwargs) -> None: ...
def init(self, database: str | None, **kwargs) -> None: ...
def __enter__(self) -> Self: ...
def __exit__(
self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: TracebackType | None
Expand Down Expand Up @@ -895,10 +902,10 @@ class SqliteDatabase(Database):
truncate_table: bool
nulls_ordering: bool
def __init__(
self, database, pragmas=None, regexp_function: bool = False, rank_functions: bool = False, *args, **kwargs
self, database: str | None, pragmas=None, regexp_function: bool = False, rank_functions: bool = False, *args, **kwargs
) -> None: ...
returning_clause: Incomplete
def init(self, database, pragmas=None, timeout: int = 5, returning_clause=None, **kwargs) -> None: ...
def init(self, database: str | None, pragmas=None, timeout: int = 5, returning_clause=None, **kwargs) -> None: ...
def pragma(self, key, value=..., permanent: bool = False, schema: str | None = None): ...
cache_size: Incomplete
foreign_keys: Incomplete
Expand Down Expand Up @@ -1006,7 +1013,7 @@ class PostgresqlDatabase(Database):
psycopg3_adapter: Incomplete
def init(
self,
database,
database: str | None,
register_unicode: bool = True,
encoding=None,
isolation_level=None,
Expand Down Expand Up @@ -1050,7 +1057,8 @@ class MySQLDatabase(Database):
safe_create_index: bool
safe_drop_index: bool
sql_mode: str
def init(self, database, **kwargs) -> None: ...
mariadb: bool
def init(self, database: str | None, mariadb: bool | None = None, **kwargs) -> None: ...
def is_connection_usable(self) -> bool: ...
def default_values_insert(self, ctx): ...
def begin(self, isolation_level: str | None = None) -> None: ...
Expand Down Expand Up @@ -1495,6 +1503,7 @@ class TimestampField(BigIntegerField[_V]):
resolution: Incomplete
ticks_to_microsecond: Incomplete
utc: Incomplete
formats: Incomplete

@overload
def __new__(cls, *args: Any, null: Literal[True], **kwargs: Unpack[_FieldKwargs]) -> TimestampField[datetime | None]: ...
Expand Down Expand Up @@ -1677,7 +1686,7 @@ class _SortedFieldList:
class SchemaManager:
model: Incomplete
context_options: Incomplete
def __init__(self, model, database=None, **context_options) -> None: ...
def __init__(self, model, database: _DatabaseType | None = None, **context_options) -> None: ...

@property
def database(self): ...
Expand Down Expand Up @@ -1731,7 +1740,7 @@ class Metadata:
def __init__(
self,
model,
database=None,
database: _DatabaseType | None = None,
table_name=None,
indexes=None,
primary_key=None,
Expand Down Expand Up @@ -1778,7 +1787,7 @@ class Metadata:
def get_primary_keys(self): ...
def get_default_dict(self): ...
def fields_to_index(self) -> list[Incomplete]: ...
def set_database(self, database) -> None: ...
def set_database(self, database: _DatabaseType) -> None: ...
def set_table_name(self, table_name) -> None: ...

class SubclassAwareMetadata(Metadata):
Expand Down Expand Up @@ -1806,7 +1815,7 @@ class _BoundModelsContext(_callable_context_manager):
database: Incomplete
bind_refs: Incomplete
bind_backrefs: Incomplete
def __init__(self, models, database, bind_refs, bind_backrefs) -> None: ...
def __init__(self, models, database: _DatabaseType, bind_refs, bind_backrefs) -> None: ...
def __enter__(self): ...
def __exit__(
self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: TracebackType | None
Expand All @@ -1820,7 +1829,7 @@ class Model(metaclass=ModelBase):
@classmethod
def validate_model(cls) -> None: ...
@classmethod
def alias(cls, alias=None) -> ModelAlias: ...
def alias(cls, alias=None) -> ModelAlias[Self]: ...
@classmethod
def select(cls, *fields) -> ModelSelect[Self]: ...
@classmethod
Expand All @@ -1846,7 +1855,7 @@ class Model(metaclass=ModelBase):
@classmethod
def bulk_update(cls, model_list, fields, batch_size=None): ...
@classmethod
def noop(cls) -> NoopModelSelect: ...
def noop(cls) -> NoopModelSelect[Self]: ...
@classmethod
def get(cls, *query, **filters) -> Self: ...
@classmethod
Expand Down Expand Up @@ -1875,9 +1884,9 @@ class Model(metaclass=ModelBase):
def __ne__(self, other) -> Expression | bool: ... # type: ignore[override]
def __sql__(self, ctx): ...
@classmethod
def bind(cls, database, bind_refs: bool = True, bind_backrefs: bool = True, _exclude=None) -> bool: ...
def bind(cls, database: _DatabaseType, bind_refs: bool = True, bind_backrefs: bool = True, _exclude=None) -> bool: ...
@classmethod
def bind_ctx(cls, database, bind_refs: bool = True, bind_backrefs: bool = True) -> _BoundModelsContext: ...
def bind_ctx(cls, database: _DatabaseType, bind_refs: bool = True, bind_backrefs: bool = True) -> _BoundModelsContext: ...
@classmethod
def table_exists(cls): ...
@classmethod
Expand All @@ -1891,12 +1900,12 @@ class Model(metaclass=ModelBase):
@classmethod
def add_index(cls, *fields, **kwargs) -> None: ...

class ModelAlias(Node):
def __init__(self, model, alias=None) -> None: ...
class ModelAlias(Node, Generic[_M]):
def __init__(self, model: type[_M], alias=None) -> None: ...
def __getattr__(self, attr: str): ...
def __setattr__(self, attr: str, value) -> None: ...
def get_field_aliases(self) -> list[Incomplete]: ...
def select(self, *selection) -> ModelSelect: ...
def select(self, *selection) -> ModelSelect[_M]: ...
def __call__(self, **kwargs): ...
def __sql__(self, ctx): ...

Expand Down Expand Up @@ -1936,8 +1945,10 @@ class BaseModelSelect(_ModelQueryHelper):
__sub__ = except_
def __iter__(self): ...
def prefetch(self, *subqueries): ...
def get(self, database=None): ...
def get_or_none(self, database=None): ...
def with_related(self, *loads: Load | ForeignKeyField[Any] | BackrefAccessor) -> Self: ...
def iterator(self, database: _DatabaseType | None = ...) -> Iterator[Any]: ...
def get(self, database: _DatabaseType | None = None): ...
def get_or_none(self, database: _DatabaseType | None = None): ...
def group_by(self, *columns) -> Self: ...

class ModelCompoundSelectQuery(BaseModelSelect, CompoundSelectQuery): # type: ignore[misc]
Expand All @@ -1948,8 +1959,8 @@ class ModelSelect(BaseModelSelect, Select, Generic[_M]): # type: ignore[misc]
model: type[_M]
def __init__(self, model, fields_or_models, is_default: bool = False) -> None: ...
def __iter__(self) -> Iterator[_M]: ...
def get(self, database=None) -> _M: ...
def get_or_none(self, database=None) -> _M | None: ...
def get(self, database: _DatabaseType | None = None) -> _M: ...
def get_or_none(self, database: _DatabaseType | None = None) -> _M | None: ...
def clone(self) -> Self: ...
def select(self, *fields_or_models) -> ModelSelect[_M]: ...
def select_extend(self, *columns) -> Self: ...
Expand All @@ -1963,7 +1974,7 @@ class ModelSelect(BaseModelSelect, Select, Generic[_M]): # type: ignore[misc]
def create_table(self, name, safe: bool = True, **meta): ...
def __sql_selection__(self, ctx, is_subquery: bool = False): ...

class NoopModelSelect(ModelSelect):
class NoopModelSelect(ModelSelect[_M]):
def __sql__(self, ctx): ...

class _ModelWriteQueryHelper(_ModelQueryHelper):
Expand All @@ -1982,7 +1993,7 @@ class ModelInsert(_ModelWriteQueryHelper, Insert): # type: ignore[misc]

class ModelDelete(_ModelWriteQueryHelper, Delete): ... # type: ignore[misc]

class ManyToManyQuery(ModelSelect):
class ManyToManyQuery(ModelSelect[_M]):
def __init__(self, instance, accessor, rel, *args, **kwargs) -> None: ...
def add(self, value, clear_existing: bool = False) -> None: ...
def remove(self, value): ...
Expand Down Expand Up @@ -2053,6 +2064,16 @@ class PrefetchQuery(_PrefetchQuery):

def prefetch(sq, *subqueries): ...

class Load(Node):
def __init__(
self,
rel: ForeignKeyField[Any] | BackrefAccessor,
query: ModelSelect[Any] | None = ...,
strategy: int = ...,
per_parent: int | None = ...,
) -> None: ...
def then(self, *children: Load | ForeignKeyField[Any] | BackrefAccessor) -> Self: ...

__all__ = [
"AnyField",
"AsIs",
Expand Down Expand Up @@ -2104,6 +2125,7 @@ __all__ = [
"IPField",
"JOIN",
"JSONField",
"Load",
"ManyToManyField",
"Model",
"ModelIndex",
Expand Down