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
13 changes: 9 additions & 4 deletions IPython/core/completer.py
Original file line number Diff line number Diff line change
Expand Up @@ -222,7 +222,6 @@
from IPython.utils.PyColorize import theme_table
from IPython.utils.decorators import sphinx_options
from IPython.utils.dir2 import dir2, get_real_method
from IPython.utils.docs import GENERATING_DOCUMENTATION
from IPython.utils.path import ensure_dir_exists
from IPython.utils.process import arg_split
from traitlets import (
Expand Down Expand Up @@ -1313,14 +1312,19 @@ def _evaluate_expr(self, expr):
),
)
done = True
except Exception as e:
if self.debug:
print("Evaluation exception", e)
except (SyntaxError, TypeError):
# TypeError can show up with something like `+ d`
# where `d` is a dictionary.

# trim the expression to remove any invalid prefix
# e.g. user starts `(d[`, so we get `expr = '(d'`,
# where parenthesis is not closed.
# TODO: make this faster by reusing parts of the computation?
expr = self._trim_expr(expr)
except Exception as e:
if self.debug:
print("Evaluation exception", e)
done = True
return obj

@property
Expand All @@ -1331,6 +1335,7 @@ def _auto_import(self):
self._auto_import_func = import_item(self.auto_import_method)
return self._auto_import_func


def get__all__entries(obj):
"""returns the strings in the __all__ attribute"""
try:
Expand Down
73 changes: 56 additions & 17 deletions IPython/core/guarded_eval.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
from copy import copy
from inspect import isclass, signature, Signature
from inspect import isclass, signature, Signature, getmodule
from typing import (
Annotated,
AnyStr,
Expand Down Expand Up @@ -122,14 +122,24 @@ def can_call(self, func):
def _get_external(module_name: str, access_path: Sequence[str]):
"""Get value from external module given a dotted access path.

Only gets value if the module is already imported.

Raises:
* `KeyError` if module is removed not found, and
* `AttributeError` if access path does not match an exported object
"""
member_type = sys.modules[module_name]
for attr in access_path:
member_type = getattr(member_type, attr)
return member_type
try:
member_type = sys.modules[module_name]
# standard module
for attr in access_path:
member_type = getattr(member_type, attr)
return member_type
except (KeyError, AttributeError):
# handle modules in namespace packages
module_path = ".".join([module_name, *access_path])
if module_path in sys.modules:
return sys.modules[module_path]
raise


def _has_original_dunder_external(
Expand All @@ -139,18 +149,26 @@ def _has_original_dunder_external(
method_name: str,
):
if module_name not in sys.modules:
# LBYLB as it is faster
return False
full_module_path = ".".join([module_name, *access_path])
if full_module_path not in sys.modules:
# LBYLB as it is faster
return False
try:
member_type = _get_external(module_name, access_path)
value_type = type(value)
if type(value) == member_type:
return True
if isinstance(member_type, ModuleType):
value_module = getmodule(value_type)
if not value_module or not value_module.__name__:
return False
if value_module.__name__.startswith(member_type.__name__):
return True
if method_name == "__getattribute__":
# we have to short-circuit here due to an unresolved issue in
# `isinstance` implementation: https://bugs.python.org/issue32683
return False
if isinstance(value, member_type):
if not isinstance(member_type, ModuleType) and isinstance(value, member_type):
method = getattr(value_type, method_name, None)
member_method = getattr(member_type, method_name, None)
if member_method == method:
Expand Down Expand Up @@ -185,35 +203,47 @@ def _has_original_dunder(
return False


def _coerce_path_to_tuples(
allow_list: set[tuple[str, ...] | str],
) -> set[tuple[str, ...]]:
"""Replace dotted paths on the provided allow-list with tuples."""
return {
path if isinstance(path, tuple) else tuple(path.split("."))
for path in allow_list
}


@undoc
@dataclass
class SelectivePolicy(EvaluationPolicy):
allowed_getitem: set[InstancesHaveGetItem] = field(default_factory=set)
allowed_getitem_external: set[tuple[str, ...]] = field(default_factory=set)
allowed_getitem_external: set[tuple[str, ...] | str] = field(default_factory=set)

allowed_getattr: set[MayHaveGetattr] = field(default_factory=set)
allowed_getattr_external: set[tuple[str, ...]] = field(default_factory=set)
allowed_getattr_external: set[tuple[str, ...] | str] = field(default_factory=set)

allowed_operations: set = field(default_factory=set)
allowed_operations_external: set[tuple[str, ...]] = field(default_factory=set)
allowed_operations_external: set[tuple[str, ...] | str] = field(default_factory=set)

_operation_methods_cache: dict[str, set[Callable]] = field(
default_factory=dict, init=False
)

def can_get_attr(self, value, attr):
allowed_getattr_external = _coerce_path_to_tuples(self.allowed_getattr_external)

has_original_attribute = _has_original_dunder(
value,
allowed_types=self.allowed_getattr,
allowed_methods=self._getattribute_methods,
allowed_external=self.allowed_getattr_external,
allowed_external=allowed_getattr_external,
method_name="__getattribute__",
)
has_original_attr = _has_original_dunder(
value,
allowed_types=self.allowed_getattr,
allowed_methods=self._getattr_methods,
allowed_external=self.allowed_getattr_external,
allowed_external=allowed_getattr_external,
method_name="__getattr__",
)

Expand Down Expand Up @@ -245,7 +275,7 @@ def can_get_attr(self, value, attr):
return True # pragma: no cover

# Properties in subclasses of allowed types may be ok if not changed
for module_name, *access_path in self.allowed_getattr_external:
for module_name, *access_path in allowed_getattr_external:
try:
external_class = _get_external(module_name, access_path)
external_class_attr_val = getattr(external_class, attr)
Expand All @@ -257,15 +287,19 @@ def can_get_attr(self, value, attr):

def can_get_item(self, value, item):
"""Allow accessing `__getiitem__` of allow-listed instances unless it was not modified."""
allowed_getitem_external = _coerce_path_to_tuples(self.allowed_getitem_external)
return _has_original_dunder(
value,
allowed_types=self.allowed_getitem,
allowed_methods=self._getitem_methods,
allowed_external=self.allowed_getitem_external,
allowed_external=allowed_getitem_external,
method_name="__getitem__",
)

def can_operate(self, dunders: tuple[str, ...], a, b=None):
allowed_operations_external = _coerce_path_to_tuples(
self.allowed_operations_external
)
objects = [a]
if b is not None:
objects.append(b)
Expand All @@ -275,7 +309,7 @@ def can_operate(self, dunders: tuple[str, ...], a, b=None):
obj,
allowed_types=self.allowed_operations,
allowed_methods=self._operator_dunder_methods(dunder),
allowed_external=self.allowed_operations_external,
allowed_external=allowed_operations_external,
method_name=dunder,
)
for dunder in dunders
Expand Down Expand Up @@ -586,7 +620,12 @@ def eval_node(node: Union[ast.AST, None], context: EvaluationContext):
dunders = _find_dunder(node.op, UNARY_OP_DUNDERS)
if dunders:
if policy.can_operate(dunders, value):
return getattr(value, dunders[0])()
try:
return getattr(value, dunders[0])()
except AttributeError:
raise TypeError(
f"bad operand type for unary {node.op}: {type(value)}"
)
else:
raise GuardRejection(
f"Operation (`{dunders}`) for",
Expand Down
53 changes: 53 additions & 0 deletions tests/test_completer.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
import pytest
import sys
import textwrap
import types
import unittest
import random

Expand Down Expand Up @@ -1403,6 +1404,58 @@ def test_completion_no_autoimport(self):
_, matches = complete(line_buffer="math.")
self.assertNotIn(".pi", matches)

def test_completion_allow_custom_getattr_per_module(self):
factory_code = textwrap.dedent(
"""
class ListFactory:
def __getattr__(self, attr):
return []
"""
)

safe_lib = types.ModuleType("my.safe.lib")
sys.modules["my.safe.lib"] = safe_lib
exec(factory_code, safe_lib.__dict__)

unsafe_lib = types.ModuleType("my.unsafe.lib")
sys.modules["my.unsafe.lib"] = unsafe_lib
exec(factory_code, unsafe_lib.__dict__)

ip = get_ipython()
ip.user_ns["safe_list_factory"] = safe_lib.ListFactory()
ip.user_ns["unsafe_list_factory"] = unsafe_lib.ListFactory()
complete = ip.Completer.complete
with (
evaluation_policy("limited", allowed_getattr_external={"my.safe.lib"}),
jedi_status(False),
):
_, matches = complete(line_buffer="safe_list_factory.example.")
self.assertIn(".append", matches)
# this also checks against https://github.com/ipython/ipython/issues/14916
# because removing "un" would cause this test to incorrectly pass
_, matches = complete(line_buffer="unsafe_list_factory.example.")
self.assertNotIn(".append", matches)

sys.modules["my"] = types.ModuleType("my")

with (
evaluation_policy("limited", allowed_getattr_external={"my"}),
jedi_status(False),
):
_, matches = complete(line_buffer="safe_list_factory.example.")
self.assertIn(".append", matches)
_, matches = complete(line_buffer="unsafe_list_factory.example.")
self.assertIn(".append", matches)

with (
evaluation_policy("limited"),
jedi_status(False),
):
_, matches = complete(line_buffer="safe_list_factory.example.")
self.assertNotIn(".append", matches)
_, matches = complete(line_buffer="unsafe_list_factory.example.")
self.assertNotIn(".append", matches)

def test_dict_key_completion_bytes(self):
"""Test handling of bytes in dict key completion"""
ip = get_ipython()
Expand Down