diff --git a/tests/functional/syntax/test_concat.py b/tests/functional/syntax/test_concat.py index 612b45622a..3ddd31eb02 100644 --- a/tests/functional/syntax/test_concat.py +++ b/tests/functional/syntax/test_concat.py @@ -117,6 +117,20 @@ def test_block_fail(assert_compile_failed, get_contract, bad_code, exc): assert_compile_failed(lambda: get_contract(bad_code), exc) +def test_concat_type_mismatch_message_uses_readable_type_names(): + code = """ +@external +def foo() -> Bytes[64]: + return concat(123, b"x") + """ + with pytest.raises(TypeMismatch) as e: + compiler.compile_code(code) + + assert e.value.message == ( + "Expected one of Bytes, String, bytesM but literal can only be cast as int104 or uint96." + ) + + valid_list = [ """ @external diff --git a/tests/functional/syntax/test_len.py b/tests/functional/syntax/test_len.py index b8cc61df1d..c2b8721576 100644 --- a/tests/functional/syntax/test_len.py +++ b/tests/functional/syntax/test_len.py @@ -53,3 +53,31 @@ def foo() -> uint256: @pytest.mark.parametrize("good_code", valid_list) def test_list_success(good_code): assert compile_code(good_code) is not None + + +def test_len_type_mismatch_message_uses_readable_type_names(): + code = """ +@external +def foo(inp: int128) -> uint256: + return len(inp) + """ + with pytest.raises(TypeMismatch) as e: + compile_code(code) + + assert e.value.message == ( + "Given reference has type int128, expected one of String, Bytes, DynArray" + ) + + +def test_index_type_mismatch_message_uses_readable_type_names(): + # `IntegerT.any()` renders its class-level `_id` (`integer`), not the + # applied per-instance name (`uint256`) + code = """ +@external +def foo(x: DynArray[uint256, 3]) -> uint256: + return x[b"ab"] + """ + with pytest.raises(TypeMismatch) as e: + compile_code(code) + + assert e.value.message == "Expected integer but literal can only be cast as Bytes[2]." diff --git a/vyper/semantics/analysis/getters.py b/vyper/semantics/analysis/getters.py index c5c9ed1c1e..df66f7ab72 100644 --- a/vyper/semantics/analysis/getters.py +++ b/vyper/semantics/analysis/getters.py @@ -46,12 +46,12 @@ def generate_public_variable_getters(vyper_module: vy_ast.Module) -> None: # type as the next arg arg, annotation = annotation.slice.elements # type: ignore elif annotation.value.get("id") == "DynArray": - arg = vy_ast.Name(id=type_._id) + arg = vy_ast.Name(id=type_.serialization_name) annotation = annotation.slice.elements[0] # type: ignore else: # for other types, build an input arg node from the expected # type and remove the outer `Subscript` from the annotation - arg = vy_ast.Name(id=type_._id) + arg = vy_ast.Name(id=type_.serialization_name) annotation = annotation.value input_nodes.append(vy_ast.arg(arg=f"arg{i}", annotation=arg)) diff --git a/vyper/semantics/types/__init__.py b/vyper/semantics/types/__init__.py index 412534973b..6986f47294 100644 --- a/vyper/semantics/types/__init__.py +++ b/vyper/semantics/types/__init__.py @@ -23,7 +23,14 @@ def _get_primitive_types(): # are in the namespace instead of concrete type objects. res.extend([BytesT, StringT]) - ret = {t._id: t for t in res} + ret = {} + for t in res: + if isinstance(t, type): + # parametrizable bytestring *classes* (e.g. `Bytes` for `Bytes[5]`) + # are registered under their declared name. + ret[t._id] = t + else: + ret[t.serialization_name] = t ret.update(_get_sequence_types()) return ret diff --git a/vyper/semantics/types/base.py b/vyper/semantics/types/base.py index 3ffe5233e6..30848dccb3 100644 --- a/vyper/semantics/types/base.py +++ b/vyper/semantics/types/base.py @@ -24,6 +24,12 @@ class _GenericTypeAcceptor: def __repr__(self): return f"GenericTypeAcceptor({self.type_})" + def __str__(self): + # `.any()` hands us the *class*, whose `_id` names the type family + # (e.g. `String`, `bytesM`, `integer`). `TYPE_T` has no `_id` but is + # never stringified: its call sites reject a non-type arg first. + return self.type_._id + def __init__(self, type_): self.type_ = type_ @@ -82,7 +88,12 @@ class VyperType: typeclass: str = None # type: ignore - _id: str # rename to `_name` + # Name of the type family (e.g. `bool`, `String`, `DynArray`, `integer`, + # `bytesM`). Types whose applied parameters are fused into a single source + # token (`uint256`, `bytes4`) expose the applied name via + # `serialization_name`; parameterized containers (e.g. `String[5]`, + # `DynArray[uint256, 3]`) render the applied name in `__repr__`. + _id: str _type_members: Optional[Dict] = None _valid_literal: Tuple = () _invalid_locations: Tuple = () @@ -135,11 +146,19 @@ def __lt__(self, other): def __repr__(self): # TODO: add `pretty()` to the VyperType API? + return self.serialization_name + + # The name used to refer to this type as a single source token (AST + # `to_dict` "name", namespace key, synthesized getter annotations). + # Defaults to `_id`; parametric primitives that fuse the applied + # parameters into the name (`uint256`, `bytes4`) override it. + @property + def serialization_name(self): return self._id # return a dict suitable for serializing in the AST def to_dict(self): - ret = {"name": self._id} + ret = {"name": self.serialization_name} if self.decl_node is not None: ret["type_decl_node"] = self.decl_node.get_id_dict() if self.typeclass is not None: diff --git a/vyper/semantics/types/primitives.py b/vyper/semantics/types/primitives.py index 12ee47afaf..8176af905b 100644 --- a/vyper/semantics/types/primitives.py +++ b/vyper/semantics/types/primitives.py @@ -57,6 +57,7 @@ def validate_literal(self, node: vy_ast.Constant) -> None: # one-word bytesM with m possible bytes set, e.g. bytes1..bytes32 class BytesM_T(_PrimT): typeclass = "bytes_m" + _id = "bytesM" _valid_literal = (vy_ast.Hex,) @@ -67,7 +68,7 @@ def __init__(self, m): self.m: int = m @property - def _id(self): + def serialization_name(self): return f"bytes{self.m}" @property @@ -266,6 +267,7 @@ class IntegerT(NumericT): """ typeclass = "integer" + _id = "integer" _valid_literal = (vy_ast.Int,) _equality_attrs = ("is_signed", "bits") @@ -277,8 +279,8 @@ def __init__(self, is_signed, bits): self._is_signed = is_signed self._bits = bits - @cached_property - def _id(self): + @property + def serialization_name(self): u = "u" if not self.is_signed else "" return f"{u}int{self.bits}"