diff --git a/src/nagini_translation/analyzer.py b/src/nagini_translation/analyzer.py index aae9e0b35..459dc2484 100644 --- a/src/nagini_translation/analyzer.py +++ b/src/nagini_translation/analyzer.py @@ -428,8 +428,17 @@ def find_or_create_class(self, name: str, module=None) -> PythonClass: name = aliases.get(name, name) if self.current_class and name in self.current_class.type_vars: return self.current_class.type_vars[name] - if not module: + if module and module != self.module: + superscope = module + else: + superscope = self.current_class or self.module module = self.module + class_scope = superscope + while isinstance(class_scope, PythonClass): + if class_scope.name == name: + return class_scope + class_scope = class_scope.superscope + # Check all imported modules for the class. for visible_module in module.get_included_modules((), True): if name in visible_module.classes: @@ -438,10 +447,10 @@ def find_or_create_class(self, name: str, module=None) -> PythonClass: else: # Class doesn't exist yet, create it. superclass = self.global_module.classes[OBJECT_TYPE] if name != OBJECT_TYPE else None - cls = self.node_factory.create_python_class(name, module, + cls = self.node_factory.create_python_class(name, superscope, self.node_factory, superclass=superclass) - module.classes[name] = cls + superscope.classes[name] = cls return cls def find_or_create_target_class(self, node: ast.AST) -> PythonClass: @@ -538,13 +547,15 @@ def _visit_ADT(self, cls: PythonClass, actual_bases: List[ast.AST], return actual_bases def visit_ClassDef(self, node: ast.ClassDef) -> None: - if self.current_function or self.current_class: + if self.current_function: raise InvalidProgramException(node, 'nested.class.declaration') name = node.name - self.define_new(self.module, name, node) + container = self.module if self.current_class is None else self.current_class + self.define_new(container, name, node) cls = self.find_or_create_class(name) cls.defined = True cls.node = node + old_class = self.current_class self.current_class = cls actual_bases = [] current_index = 0 @@ -591,7 +602,7 @@ def visit_ClassDef(self, node: ast.ClassDef) -> None: for member in node.body: self.visit(member, node) - self.current_class = None + self.current_class = old_class def _is_illegal_magic_method_name(self, name: str) -> bool: """ @@ -1228,6 +1239,8 @@ def visit_Attribute(self, node: ast.Attribute) -> None: self.track_access(node, real_target) else: receiver = self.typeof(node.value) + if isinstance(receiver, PythonType) and node.attr in receiver.python_class.classes: + return if (isinstance(receiver, UnionType) and not isinstance(receiver, OptionalType)): for type in receiver.get_types() - {None}: @@ -1311,7 +1324,10 @@ def convert_type(self, mypy_type, node, bound_type_vars: Dict[str, PythonType] = msg = f'Type could not be fully inferred (this usually means that a type argument is unknown)' raise InvalidProgramException(node, 'partial.type', message=msg) else: - msg = 'Unsupported type: {}'.format(mypy_type.__class__.__name__) + if mypy_type is None: + msg = 'Internal error: Could not determine type.' + else: + msg = 'Unsupported type: {}'.format(mypy_type.__class__.__name__) raise UnsupportedException(node, desc=msg) return result @@ -1324,11 +1340,36 @@ def _convert_normal_type(self, mypy_type) -> PythonType: name = 'list' if prefix.endswith('.' + name): prefix = prefix[:-(len(name) + 1)] - target_module = self.module + best_module_fit = None for module in self.modules.values(): - if module.type_prefix == prefix: + m_name = module.full_module_name or module.type_prefix + if m_name == prefix or module.type_prefix == prefix: target_module = module break + if prefix.startswith(m_name): + if best_module_fit is None or len(m_name) > len(best_module_fit.full_module_name): + best_module_fit = module + else: + if prefix in IGNORED_IMPORTS: + target_module = self.module.global_module + else: + if best_module_fit: + best_fit_name = best_module_fit.full_module_name or best_module_fit.type_prefix + remaining_prefix = prefix[len(best_fit_name) + 1:] + remaining_parts = remaining_prefix.split('.') + best_fit_container = best_module_fit + while remaining_parts: + if remaining_parts[0] in best_fit_container.classes: + best_fit_container = best_fit_container.classes[remaining_parts[0]] + if len(remaining_parts) == 1: + break + else: + remaining_parts.pop(0) + else: + break + if best_fit_container and name in best_fit_container.classes: + return best_fit_container.classes[name] + raise Exception("Internal error: Could not find module for type.") result = self.find_or_create_class(name, module=target_module) return result @@ -1408,8 +1449,8 @@ def typeof(self, node: ast.AST) -> PythonType: if node.id in self.module.classes: return self.module.classes[node.id] context = [] - if self.current_class is not None: - context.append(self.current_class.name) + if self.current_class: + context.extend(self.current_class.full_name) if self.current_function is not None: context.append(self.current_function.name) context.extend(self.current_scopes) @@ -1420,6 +1461,8 @@ def typeof(self, node: ast.AST) -> PythonType: return self.convert_type(type, node) elif isinstance(node, ast.Attribute): receiver = self.typeof(node.value) + if isinstance(receiver, PythonType) and node.attr in receiver.python_class.classes: + return receiver.python_class.classes[node.attr] if isinstance(receiver, OptionalType): receiver = receiver.optional_type if isinstance(receiver, UnionType) and not isinstance(receiver, OptionalType): @@ -1434,15 +1477,15 @@ def typeof(self, node: ast.AST) -> PythonType: return UnionType(list(set_of_types)) if len(set_of_types) > 1 else set_of_types.pop() contexts = [] if isinstance(receiver, OptionalType): - contexts.append([receiver.optional_type.name]) + contexts.append(receiver.optional_type.python_class.full_name) rec_super = receiver.optional_type.superclass module = receiver.optional_type.module else: - contexts.append([receiver.name]) + contexts.append(receiver.python_class.full_name) rec_super = receiver.superclass module = receiver.module while rec_super is not None: - contexts.append([rec_super.name]) + contexts.append(rec_super.python_class.full_name) rec_super = rec_super.superclass bound_type_vars = None if isinstance(receiver, GenericType) or (isinstance(receiver, OptionalType) and isinstance(receiver.optional_type, GenericType)): @@ -1464,8 +1507,8 @@ def typeof(self, node: ast.AST) -> PythonType: cls = self.module.global_module.classes['type'] return GenericType(cls, [self.current_class]) context = [] - if self.current_class is not None: - context.append(self.current_class.name) + if self.current_class: + context.extend(self.current_class.full_name) context.append(self.current_function.name) context.extend(self.current_scopes) type, _ = self.module.get_type(context, node.arg) diff --git a/src/nagini_translation/lib/constants.py b/src/nagini_translation/lib/constants.py index 6110f8e81..8eaa53003 100644 --- a/src/nagini_translation/lib/constants.py +++ b/src/nagini_translation/lib/constants.py @@ -30,6 +30,7 @@ EXTENDABLE_BUILTINS = [ 'object', 'Exception', + 'BaseLock', 'Lock', 'int' ] diff --git a/src/nagini_translation/lib/program_nodes.py b/src/nagini_translation/lib/program_nodes.py index a503ebf1b..ff70b4c50 100644 --- a/src/nagini_translation/lib/program_nodes.py +++ b/src/nagini_translation/lib/program_nodes.py @@ -174,11 +174,20 @@ def full_module_name(self) -> str: return self.types.module_name return self.type_prefix + @property + def full_name(self) -> List[str]: + if self.type_prefix is None: + return [] # ???? + return self.type_prefix.split(".") + def get_relative_import_name(self, name: str, level: int) -> str: module_name = name if level > 0: current_module_name = self.full_module_name - module_name_to_add = current_module_name.split(".")[:-level] + actual_level = level if not (self.module.file.endswith('__init__.py') or self.module.file.endswith('__init__.pyi')) else level - 1 + module_name_to_add = current_module_name.split(".") + if actual_level != 0: + module_name_to_add = module_name_to_add[:-actual_level] if module_name is not None: module_name_to_add.append(module_name) module_name = ".".join(module_name_to_add) @@ -222,6 +231,14 @@ def process(self, translator: 'Translator') -> None: def scope_prefix(self) -> List[str]: return [] + @property + def all_classes(self) -> OrderedDict[str, 'PythonClass']: + res = OrderedDict() + for cls_name, cls in self.classes.items(): + if cls_name == cls.name: + res.update(cls.all_classes) + return res + def get_func_or_method(self, name: str) -> 'PythonMethod': for module in [self] + self.from_imports + [self.global_module]: if not isinstance(module, PythonModule): @@ -245,6 +262,11 @@ def get_type(self, prefixes: List[str], name: str, """ if self in previous: return None, None + + local_type, local_alts = self.types.get_type(prefixes, name) + if local_type is not None: + return local_type, local_alts + actual_prefix = self.type_prefix.split('.') if self.type_prefix else [] actual_prefix.extend(prefixes) local_type, local_alts = self.types.get_type(actual_prefix, name) @@ -408,6 +430,7 @@ def __init__(self, name: str, superscope: PythonScope, self.predicates = OrderedDict() self.fields = OrderedDict() self.static_fields = OrderedDict() + self.classes = OrderedDict() self.type = None # infer, domain type self.interface = interface self.defined = False @@ -423,6 +446,14 @@ def get_bound_type_vars(self) -> Dict['TypeVar', 'PythonType']: return self.superclass.get_bound_type_vars() return {} + @property + def all_classes(self) -> OrderedDict[str, 'PythonClass']: + res = OrderedDict() + res[".".join(self.full_name)] = self + for cls_name, cls in self.classes.items(): + res.update(cls.all_classes) + return res + @property def is_defining_adt(self) -> bool: """ @@ -670,6 +701,9 @@ def process(self, sil_name: str, translator: 'Translator') -> None: of them. """ self.sil_name = sil_name + for name, cls in self.classes.items(): + cls_name = self.name + '_' + name + cls.process(self.get_fresh_name(cls_name), translator) for name, function in self.functions.items(): func_name = self.name + '_' + name function.process(self.get_fresh_name(func_name), translator) @@ -729,7 +763,7 @@ def get_contents(self, only_top: bool) -> Dict: used by get_target). If 'only_top' is true, returns only top level elements that can be accessed without a receiver. """ - dicts = [self.static_methods, self.static_fields, self.type_vars] + dicts = [self.static_methods, self.static_fields, self.type_vars, self.classes] if not only_top: dicts.extend([self.functions, self.fields, self.methods, self.predicates]) @@ -767,6 +801,13 @@ def try_unbox(self) -> 'PythonClass': def python_class(self) -> 'PythonClass': return self + @property + def full_name(self) -> List[str]: + result = [] + result.extend(self.superscope.full_name) + result.append(self.name) + return result + class GenericType(PythonType): """ diff --git a/src/nagini_translation/lib/typeinfo.py b/src/nagini_translation/lib/typeinfo.py index 3c83bb61d..21390619e 100644 --- a/src/nagini_translation/lib/typeinfo.py +++ b/src/nagini_translation/lib/typeinfo.py @@ -233,6 +233,9 @@ def type_of(self, node): key = (node.name,) if key in self.all_types: return self.all_types[key] + full_key = tuple(self.prefix) + key + if full_key in self.all_types: + return self.all_types[full_key] elif isinstance(node, mypy.nodes.CallExpr): if isinstance(node.callee, mypy.nodes.NameExpr) and node.callee.name == 'Result': key = tuple(self.prefix) diff --git a/src/nagini_translation/resources/builtins.json b/src/nagini_translation/resources/builtins.json index 72744e482..84c7ede0a 100644 --- a/src/nagini_translation/resources/builtins.json +++ b/src/nagini_translation/resources/builtins.json @@ -909,6 +909,7 @@ }, "extends": "object" }, +"BaseLock": {}, "global": { "functions": { "max": { diff --git a/src/nagini_translation/translators/common.py b/src/nagini_translation/translators/common.py index 3a9a7febd..33f52bbdf 100644 --- a/src/nagini_translation/translators/common.py +++ b/src/nagini_translation/translators/common.py @@ -307,8 +307,9 @@ def set_global_defined(self, declaration: PythonNode, module: PythonModule, pos = self.to_position(node, ctx) info = self.no_info(ctx) module_set = module.names_var[1] - decl_id = self.viper.IntLit(self._get_string_value(declaration.name), pos, - info) + string_value = self._get_string_value(declaration.name) + decl_id = self.viper.IntLit(string_value, pos, info) + print("setting defined: " + declaration.name + ", " + str(string_value)) return self._set_global_defined(decl_id, module_set, pos, info) def _set_global_defined(self, decl_int: Expr, module_var: Expr, pos: Position, @@ -904,8 +905,10 @@ def is_valid_super_call(self, node: ast.Call, container) -> bool: def get_target(self, node: ast.AST, ctx: Context) -> PythonModule: container = ctx.actual_function if ctx.actual_function else ctx.module containers = [ctx] - if ctx.current_class: - containers.append(ctx.current_class) + current_class = ctx.current_class + while current_class: + containers.insert(1, current_class) + current_class = current_class.superscope if isinstance(current_class.superscope, PythonClass) else None if isinstance(container, (PythonMethod, PythonIOOperation)): containers.append(container) containers.extend(container.module.get_included_modules()) diff --git a/src/nagini_translation/translators/expression.py b/src/nagini_translation/translators/expression.py index 79cf61d01..fc396fd60 100644 --- a/src/nagini_translation/translators/expression.py +++ b/src/nagini_translation/translators/expression.py @@ -878,7 +878,7 @@ def translate_Attribute(self, node: ast.Attribute, else: # If the receiver is an ADT, attribute access is translated as deconstruction recv_type = self.get_type(node.value, ctx) - if isinstance(recv_type.python_class, PythonClass) and recv_type.python_class.is_adt: + if isinstance(recv_type, PythonType) and recv_type.python_class.is_adt: return self.translate_adt_decons(recv_type.python_class, node, position, ctx) stmt, receiver = self.translate_expr(node.value, ctx, diff --git a/src/nagini_translation/translators/program.py b/src/nagini_translation/translators/program.py index a66c4d2ba..b36329e5f 100644 --- a/src/nagini_translation/translators/program.py +++ b/src/nagini_translation/translators/program.py @@ -1208,8 +1208,8 @@ def translate_program(self, modules: List[PythonModule], sil_progs: Program, for module in modules: ctx.module = module containers = [module] - for class_name, cls in module.classes.items(): - if class_name in PRIMITIVES or class_name != cls.name: + for class_name, cls in module.all_classes.items(): + if class_name in PRIMITIVES: # Skip primitives or entries for type variables. continue if cls.is_adt: @@ -1273,8 +1273,8 @@ def translate_program(self, modules: List[PythonModule], sil_progs: Program, for pred in module.predicates.values(): predicates.append(self.translate_predicate(pred, ctx)) self.track_dependencies(selected_names, selected, pred, ctx) - for class_name, cls in module.classes.items(): - if class_name in PRIMITIVES or class_name != cls.name: + for class_name, cls in module.all_classes.items(): + if class_name in PRIMITIVES: # Skip primitives and type variable entries. continue if cls.is_adt and cls.is_defining_adt: diff --git a/src/nagini_translation/translators/statement.py b/src/nagini_translation/translators/statement.py index f0d3c6b5c..fca48dab8 100644 --- a/src/nagini_translation/translators/statement.py +++ b/src/nagini_translation/translators/statement.py @@ -308,7 +308,8 @@ def translate_stmt_ClassDef(self, node: ast.ClassDef, ctx: Context) -> List[Stmt """ assert self.is_main_method(ctx) # static field definitions - cls = ctx.module.classes[node.name] + container = ctx.current_class or ctx.module + cls = container.classes[node.name] stmts = [] pos = self.to_position(node, ctx) info = self.no_info(ctx)