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
4 changes: 2 additions & 2 deletions etils/epy/contextlib.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,10 +76,10 @@ def __contextmanager__(self) -> Iterable[_T]:
def __enter__(self) -> _T:
# object.__setattr__ to support frozen dataclasses
object.__setattr__(self, '_epy_cm', self.__contextmanager__())
return self._epy_cm.__enter__() # pytype: disable=attribute-error
return self._epy_cm.__enter__()

def __exit__(self, exc_type, exc_value, traceback) -> None:
return self._epy_cm.__exit__(exc_type, exc_value, traceback) # pytype: disable=attribute-error
return self._epy_cm.__exit__(exc_type, exc_value, traceback)


# Should use `contextlib.nested` instead if outputs are required?
Expand Down
18 changes: 9 additions & 9 deletions etils/epy/itertools.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@

from typing import Any, Callable, Iterable, Iterator, TypeVar

# from typing_extensions import Unpack, TypeVarTuple # pytype: disable=not-supported-yet # pylint: disable=g-multiple-import
# from typing_extensions import Unpack, TypeVarTuple # pylint: disable=g-multiple-import

# TODO(pytype): Once supported, should replace
Unpack = Any
Expand Down Expand Up @@ -116,9 +116,9 @@ def splitby(
return false_list, true_list


def zip_dict( # pytype: disable=invalid-annotation
*dicts: Unpack[dict[_KeyT, _ValuesT]],
) -> Iterator[_KeyT, tuple[Unpack[_ValuesT]]]:
def zip_dict(
*dicts: Unpack[dict[_KeyT, _ValuesT]], # pyrefly: ignore[bad-specialization]
) -> Iterator[_KeyT, tuple[Unpack[_ValuesT]]]: # pyrefly: ignore[bad-specialization]
"""Iterate over items of dictionaries grouped by their keys.

Example:
Expand All @@ -144,12 +144,12 @@ def zip_dict( # pytype: disable=invalid-annotation
KeyError: If dicts does not contain the same keys.
"""
# Set does not keep order like dict, so only use set to compare keys
all_keys = set(itertools.chain(*dicts))
all_keys = set(itertools.chain(*dicts)) # pyrefly: ignore[bad-argument-type]
d0 = dicts[0]

if len(all_keys) != len(d0):
raise KeyError(f'Missing keys: {all_keys ^ set(d0)}')
if len(all_keys) != len(d0): # pyrefly: ignore[bad-argument-type]
raise KeyError(f'Missing keys: {all_keys ^ set(d0)}') # pyrefly: ignore[bad-argument-type]

for key in d0: # set merge all keys
for key in d0: # set merge all keys # pyrefly: ignore[not-iterable]
# Will raise KeyError if the dict don't have the same keys
yield key, tuple(d[key] for d in dicts)
yield key, tuple(d[key] for d in dicts) # pyrefly: ignore[bad-index, invalid-yield]
2 changes: 1 addition & 1 deletion etils/epy/lazy_imports_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -206,7 +206,7 @@ def get_dataset() -> tf.data.Dataset:
original_import = builtins.__import__
try:
builtins.__import__ = functools.partial(
_lazy_import, # pyrefly: ignore[bad-argument-type]
_lazy_import,
error_callback=error_callback,
success_callback=success_callback,
)
Expand Down
2 changes: 1 addition & 1 deletion etils/epy/lazy_imports_utils_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,7 +101,7 @@ def test_error_callback():
with epy.lazy_imports(
error_callback=error_callback, success_callback=success_callback
):
import doesnotexist # pylint: disable=g-import-not-at-top,unused-import # pytype: disable=import-error
import doesnotexist # pylint: disable=g-import-not-at-top,unused-import # pyrefly: ignore[missing-import]
error_callback.assert_not_called()
success_callback.assert_not_called()
try:
Expand Down
5 changes: 2 additions & 3 deletions etils/epy/py_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -181,8 +181,7 @@ def __init__(self):
if not isinstance(cls, type):
raise TypeError(f'{cls.__name__} is not')

cls.__init__ = _wrap_init(cls.__init__) # pyrefly: ignore[bad-assignment]
# pyrefly: ignore[bad-assignment]
cls.__init__ = _wrap_init(cls.__init__)
cls.__setattr__ = _wrap_setattr(cls.__setattr__)
return cls

Expand All @@ -195,7 +194,7 @@ def normalize_str_to_list(x: Optional[StrOrStrList]) -> list[str]:
elif not isinstance(x, (list, tuple)):
raise TypeError(f'Expected list. Got: {x!r}')
else: # list/tuple
return list(x) # pyrefly: ignore[no-matching-overload] # pyrefly#2607
return list(x) # pyrefly#2607


def wraps_cls(wrapped: type[Any]) -> Callable[[_Cls], _Cls]:
Expand Down
2 changes: 1 addition & 1 deletion etils/epy/reraise_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,7 +98,7 @@ def wrap_error(
# Dynamically create an exception for:
# * Compatibility with caller core (e.g. `except OriginalError`)

class WrappedException(type(e)): # pyrefly: ignore[invalid-inheritance]
class WrappedException(type(e)):
"""Exception proxy with additional message."""

def __init__(self, msg):
Expand Down
12 changes: 6 additions & 6 deletions etils/epy/text_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -314,7 +314,7 @@ def __repr__(self):
},
)
elif _is_attr(obj, force=force):
import attr # pylint: disable=g-import-not-at-top # pytype: disable=import-error
import attr # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import]

all_fields = attr.fields_dict(type(obj))

Expand Down Expand Up @@ -343,7 +343,7 @@ def __repr__(self):
# TODO(epot): When the new fiddle version is release on PyPI, this
# code could be activated (with the matching test).
elif _is_fiddle(obj):
import fiddle # pylint: disable=g-import-not-at-top # pytype: disable=import-error
import fiddle # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import]

cls_name = type(obj).__name__
formatted_fn_or_cls = obj._fn_or_cls_name_repr() # pylint: disable=protected-access
Expand Down Expand Up @@ -404,7 +404,7 @@ def _is_attr(obj: Any, *, force: bool = False) -> bool:
"""Returns `True` if the object is a `attr` dataclass."""
if 'attr' not in sys.modules:
return False
import attr # pylint: disable=g-import-not-at-top # pytype: disable=import-error
import attr # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import]

if not attr.has(type(obj)):
return False
Expand All @@ -422,7 +422,7 @@ def _is_pydantic(obj: Any, *, force: bool = False) -> bool:
if 'pydantic' not in sys.modules:
return False

import pydantic # pylint: disable=g-import-not-at-top # pytype: disable=import-error
import pydantic # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import]

if not isinstance(obj, pydantic.BaseModel):
return False
Expand All @@ -437,7 +437,7 @@ def _is_immutabledict(obj: Any, *, force: bool = False) -> bool:
"""Returns `True` if the object is an `immutabledict`."""
if 'immutabledict' not in sys.modules:
return False
import immutabledict # pylint: disable=g-import-not-at-top # pytype: disable=import-error
import immutabledict # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import]

if not isinstance(obj, immutabledict.immutabledict):
return False
Expand Down Expand Up @@ -477,7 +477,7 @@ def _is_fiddle(obj: Any) -> bool:
"""Returns `True` if the object is a `fiddle` config object."""
if 'fiddle' not in sys.modules:
return False
import fiddle # pylint: disable=g-import-not-at-top # pytype: disable=import-error
import fiddle # pylint: disable=g-import-not-at-top # pyrefly: ignore[missing-import]

return isinstance(obj, fiddle.Config)

Expand Down
4 changes: 2 additions & 2 deletions etils/epy/text_utils_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,7 +93,7 @@ class A:
t: tuple[str, ...] = ()
l: list[int] = dataclasses.field(default_factory=list)
d: dict[str, int] = dataclasses.field(default_factory=dict)
dc: B = dataclasses.field(default_factory=B) # pytype: disable=invalid-annotation,name-error
dc: B = dataclasses.field(default_factory=B)
s: str = 'aaa'

a = A(
Expand All @@ -102,7 +102,7 @@ class A:
d={'aaaaaaaaaaaaaaaaaaaa': 1, 'bbbbbbbbbbbbbbbbbbbb': 1},
)

repr_ = epy.pretty_repr(a) # pytype: disable=wrong-arg-types
repr_ = epy.pretty_repr(a)
assert repr_ == epy.dedent("""
A(
t=(
Expand Down
Loading