diff --git a/midas/generator/generator.py b/midas/generator/generator.py index b31164a..52c2d76 100644 --- a/midas/generator/generator.py +++ b/midas/generator/generator.py @@ -78,6 +78,9 @@ class Generator(p.Stmt.Visitor[ast.stmt], p.Expr.Visitor[ast.expr]): self._typed_ast = typed_ast body: list[ast.stmt] = self._visit_body(typed_ast.stmts, can_be_empty=True) predicates: list[ast.stmt] = self._constraint_generator.get_definitions() + assertion_definitions: list[ast.stmt] = list( + typed_ast.assertions.definitions.values() + ) body = predicates + body @@ -87,6 +90,8 @@ class Generator(p.Stmt.Visitor[ast.stmt], p.Expr.Visitor[ast.expr]): if self.define_is_column: body = [self._is_column_definition()] + body + body = assertion_definitions + body + module = ast.Module(body=body, type_ignores=[]) module = ast.fix_missing_locations(module) return module diff --git a/midas/generator/stubs.py b/midas/generator/stubs.py index 7ddcb06..d42f3db 100644 --- a/midas/generator/stubs.py +++ b/midas/generator/stubs.py @@ -273,14 +273,11 @@ class StubsGenerator: ), ) - case ColumnType(type=inner): + case ColumnType(): self.import_pandas = True - return ast.Subscript( - value=ast.Attribute( - value=ast.Name(id="pd"), - attr="Series", - ), - slice=self.dump_type(inner), + return ast.Attribute( + value=ast.Name(id="pd"), + attr="Series", ) case DataFrameType():