feat(types): add ColumnGroupBy
This commit is contained in:
@@ -6,7 +6,15 @@ import midas.ast.python as p
|
|||||||
from midas.ast.location import Location
|
from midas.ast.location import Location
|
||||||
from midas.checker.frame_methods import Call, MethodRegistry
|
from midas.checker.frame_methods import Call, MethodRegistry
|
||||||
from midas.checker.reporter import FileReporter
|
from midas.checker.reporter import FileReporter
|
||||||
from midas.checker.types import ColumnType, DataFrameType, TupleType, Type, UnknownType
|
from midas.checker.types import (
|
||||||
|
ColumnGroupBy,
|
||||||
|
ColumnType,
|
||||||
|
DataFrameType,
|
||||||
|
FrameGroupBy,
|
||||||
|
TupleType,
|
||||||
|
Type,
|
||||||
|
UnknownType,
|
||||||
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from midas.checker.python import PythonTyper, TypedExpr
|
from midas.checker.python import PythonTyper, TypedExpr
|
||||||
@@ -90,6 +98,26 @@ class FrameManager:
|
|||||||
reporter.error(location, f"Invalid index type {index} on {frame}")
|
reporter.error(location, f"Invalid index type {index} on {frame}")
|
||||||
return UnknownType()
|
return UnknownType()
|
||||||
|
|
||||||
|
def groupby_get(
|
||||||
|
self,
|
||||||
|
reporter: FileReporter,
|
||||||
|
location: Location,
|
||||||
|
groupby: FrameGroupBy,
|
||||||
|
index: p.Expr,
|
||||||
|
) -> Type:
|
||||||
|
result: Type = self.get(reporter, location, groupby.frame, index)
|
||||||
|
match result:
|
||||||
|
case ColumnType():
|
||||||
|
result = ColumnGroupBy(column=result)
|
||||||
|
case TupleType(items=columns):
|
||||||
|
result = TupleType(
|
||||||
|
items=tuple(
|
||||||
|
ColumnGroupBy(column=cast(ColumnType, column))
|
||||||
|
for column in columns
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _set_column(
|
def _set_column(
|
||||||
cls, frame: DataFrameType, name: str, column: ColumnType
|
cls, frame: DataFrameType, name: str, column: ColumnType
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ from midas.checker.types import (
|
|||||||
ConstraintType,
|
ConstraintType,
|
||||||
DataFrameType,
|
DataFrameType,
|
||||||
DerivedType,
|
DerivedType,
|
||||||
|
FrameGroupBy,
|
||||||
Function,
|
Function,
|
||||||
GenericType,
|
GenericType,
|
||||||
TopType,
|
TopType,
|
||||||
@@ -770,6 +771,8 @@ class PythonTyper(
|
|||||||
return self._visit_tuple_subscript(unfolded, expr)
|
return self._visit_tuple_subscript(unfolded, expr)
|
||||||
case DataFrameType():
|
case DataFrameType():
|
||||||
return self._visit_frame_subscript(unfolded, expr)
|
return self._visit_frame_subscript(unfolded, expr)
|
||||||
|
case FrameGroupBy():
|
||||||
|
return self._visit_frame_groupby_subscript(unfolded, expr)
|
||||||
|
|
||||||
operation: Optional[Type] = self.types.lookup_member(object, "__getitem__")
|
operation: Optional[Type] = self.types.lookup_member(object, "__getitem__")
|
||||||
if operation is None:
|
if operation is None:
|
||||||
@@ -1095,3 +1098,10 @@ class PythonTyper(
|
|||||||
self, frame: DataFrameType, expr: p.SubscriptExpr
|
self, frame: DataFrameType, expr: p.SubscriptExpr
|
||||||
) -> Type:
|
) -> Type:
|
||||||
return self.frame_mgr.get(self.reporter, expr.location, frame, expr.index)
|
return self.frame_mgr.get(self.reporter, expr.location, frame, expr.index)
|
||||||
|
|
||||||
|
def _visit_frame_groupby_subscript(
|
||||||
|
self, groupby: FrameGroupBy, expr: p.SubscriptExpr
|
||||||
|
) -> Type:
|
||||||
|
return self.frame_mgr.groupby_get(
|
||||||
|
self.reporter, expr.location, groupby, expr.index
|
||||||
|
)
|
||||||
|
|||||||
@@ -195,6 +195,14 @@ class FrameGroupBy:
|
|||||||
return f"FrameGroupBy[{self.frame}]"
|
return f"FrameGroupBy[{self.frame}]"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, kw_only=True)
|
||||||
|
class ColumnGroupBy:
|
||||||
|
column: ColumnType
|
||||||
|
|
||||||
|
def __str__(self) -> str:
|
||||||
|
return f"ColumnGroupBy[{self.column}]"
|
||||||
|
|
||||||
|
|
||||||
def substitute_typevars(type: Type, substitutions: dict[str, Type]) -> Type:
|
def substitute_typevars(type: Type, substitutions: dict[str, Type]) -> Type:
|
||||||
def sub_argument(arg: Function.Argument):
|
def sub_argument(arg: Function.Argument):
|
||||||
return Function.Argument(
|
return Function.Argument(
|
||||||
@@ -318,6 +326,11 @@ def substitute_typevars(type: Type, substitutions: dict[str, Type]) -> Type:
|
|||||||
frame=cast(DataFrameType, substitute_typevars(frame, substitutions))
|
frame=cast(DataFrameType, substitute_typevars(frame, substitutions))
|
||||||
)
|
)
|
||||||
|
|
||||||
|
case ColumnGroupBy(column=column):
|
||||||
|
return ColumnGroupBy(
|
||||||
|
column=cast(ColumnType, substitute_typevars(column, substitutions))
|
||||||
|
)
|
||||||
|
|
||||||
case UnknownType() | UnitType():
|
case UnknownType() | UnitType():
|
||||||
return type
|
return type
|
||||||
|
|
||||||
@@ -397,6 +410,9 @@ def to_annotation(type: Type) -> str:
|
|||||||
case FrameGroupBy():
|
case FrameGroupBy():
|
||||||
return "pd.api.typing.DataFrameGroupBy"
|
return "pd.api.typing.DataFrameGroupBy"
|
||||||
|
|
||||||
|
case ColumnGroupBy():
|
||||||
|
return "pd.api.typing.SeriesGroupBy"
|
||||||
|
|
||||||
case _:
|
case _:
|
||||||
assert_never(type)
|
assert_never(type)
|
||||||
|
|
||||||
@@ -426,4 +442,5 @@ Type = (
|
|||||||
| ColumnType
|
| ColumnType
|
||||||
| DataFrameType
|
| DataFrameType
|
||||||
| FrameGroupBy
|
| FrameGroupBy
|
||||||
|
| ColumnGroupBy
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from midas.checker.registry import TypesRegistry
|
|||||||
from midas.checker.types import (
|
from midas.checker.types import (
|
||||||
AppliedType,
|
AppliedType,
|
||||||
BaseType,
|
BaseType,
|
||||||
|
ColumnGroupBy,
|
||||||
ColumnType,
|
ColumnType,
|
||||||
ComplexType,
|
ComplexType,
|
||||||
ConstraintType,
|
ConstraintType,
|
||||||
@@ -504,6 +505,7 @@ class Generator(p.Stmt.Visitor[ast.stmt], p.Expr.Visitor[ast.expr]):
|
|||||||
| ExtensionType()
|
| ExtensionType()
|
||||||
| GenericType()
|
| GenericType()
|
||||||
| FrameGroupBy()
|
| FrameGroupBy()
|
||||||
|
| ColumnGroupBy()
|
||||||
):
|
):
|
||||||
self.logger.warning(f"Can't make assertion for type {type}")
|
self.logger.warning(f"Can't make assertion for type {type}")
|
||||||
return []
|
return []
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ from midas.checker.registry import Member, TypesRegistry
|
|||||||
from midas.checker.types import (
|
from midas.checker.types import (
|
||||||
AppliedType,
|
AppliedType,
|
||||||
BaseType,
|
BaseType,
|
||||||
|
ColumnGroupBy,
|
||||||
ColumnType,
|
ColumnType,
|
||||||
ComplexType,
|
ComplexType,
|
||||||
ConstraintType,
|
ConstraintType,
|
||||||
@@ -301,6 +302,19 @@ class StubsGenerator:
|
|||||||
attr="DataFrameGroupBy",
|
attr="DataFrameGroupBy",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
case ColumnGroupBy():
|
||||||
|
self.import_pandas = True
|
||||||
|
return ast.Attribute(
|
||||||
|
value=ast.Attribute(
|
||||||
|
value=ast.Attribute(
|
||||||
|
value=ast.Name(id="pd"),
|
||||||
|
attr="api",
|
||||||
|
),
|
||||||
|
attr="typing",
|
||||||
|
),
|
||||||
|
attr="SeriesGroupBy",
|
||||||
|
)
|
||||||
|
|
||||||
case _:
|
case _:
|
||||||
assert_never(type)
|
assert_never(type)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user