Compare commits
3
Commits
fd5399f50a
...
e6375f1aa9
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e6375f1aa9
|
||
|
|
d16e192a3a
|
||
|
|
3f61f84e5a
|
No files matched your search
+1
-1
@@ -1,5 +1,5 @@
|
|||||||
from pathlib import Path
|
|
||||||
import re
|
import re
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
HEADER = '''"""
|
HEADER = '''"""
|
||||||
This file was generated by a script. Any manual changes might be overwritten.
|
This file was generated by a script. Any manual changes might be overwritten.
|
||||||
|
|||||||
@@ -44,7 +44,9 @@ class Function:
|
|||||||
name: str
|
name: str
|
||||||
posonlyargs: list[Argument]
|
posonlyargs: list[Argument]
|
||||||
args: list[Argument]
|
args: list[Argument]
|
||||||
|
sink: Optional[Argument]
|
||||||
kwonlyargs: list[Argument]
|
kwonlyargs: list[Argument]
|
||||||
|
kw_sink: Optional[Argument]
|
||||||
returns: Optional[MidasType]
|
returns: Optional[MidasType]
|
||||||
body: list[Stmt]
|
body: list[Stmt]
|
||||||
|
|
||||||
@@ -53,6 +55,7 @@ class Function:
|
|||||||
location: Optional[Location] = None
|
location: Optional[Location] = None
|
||||||
name: str
|
name: str
|
||||||
type: Optional[MidasType]
|
type: Optional[MidasType]
|
||||||
|
default: Optional[Expr]
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def all_args(self) -> list[Argument]:
|
def all_args(self) -> list[Argument]:
|
||||||
|
|||||||
@@ -117,7 +117,9 @@ class Function(Stmt):
|
|||||||
name: str
|
name: str
|
||||||
posonlyargs: list[Argument]
|
posonlyargs: list[Argument]
|
||||||
args: list[Argument]
|
args: list[Argument]
|
||||||
|
sink: Optional[Argument]
|
||||||
kwonlyargs: list[Argument]
|
kwonlyargs: list[Argument]
|
||||||
|
kw_sink: Optional[Argument]
|
||||||
returns: Optional[MidasType]
|
returns: Optional[MidasType]
|
||||||
body: list[Stmt]
|
body: list[Stmt]
|
||||||
|
|
||||||
@@ -126,6 +128,7 @@ class Function(Stmt):
|
|||||||
location: Optional[Location] = None
|
location: Optional[Location] = None
|
||||||
name: str
|
name: str
|
||||||
type: Optional[MidasType]
|
type: Optional[MidasType]
|
||||||
|
default: Optional[Expr]
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def all_args(self) -> list[Argument]:
|
def all_args(self) -> list[Argument]:
|
||||||
|
|||||||
+114
-8
@@ -1,4 +1,5 @@
|
|||||||
import logging
|
import logging
|
||||||
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
@@ -19,6 +20,13 @@ class ReturnException(Exception):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, kw_only=True)
|
||||||
|
class MappedArgument:
|
||||||
|
expr: p.Expr
|
||||||
|
type: Type
|
||||||
|
argument: Function.Argument
|
||||||
|
|
||||||
|
|
||||||
class Checker(
|
class Checker(
|
||||||
p.Stmt.Visitor[None],
|
p.Stmt.Visitor[None],
|
||||||
p.Expr.Visitor[Type],
|
p.Expr.Visitor[Type],
|
||||||
@@ -126,15 +134,18 @@ class Checker(
|
|||||||
kw_args: list[Function.Argument] = []
|
kw_args: list[Function.Argument] = []
|
||||||
|
|
||||||
def eval_arg_type(arg: p.Function.Argument) -> Type:
|
def eval_arg_type(arg: p.Function.Argument) -> Type:
|
||||||
if arg.type is None:
|
if arg.type is not None:
|
||||||
return UnknownType()
|
return arg.type.accept(self)
|
||||||
return arg.type.accept(self)
|
if arg.default is not None:
|
||||||
|
return arg.default.accept(self)
|
||||||
|
return UnknownType()
|
||||||
|
|
||||||
for arg in stmt.posonlyargs:
|
for arg in stmt.posonlyargs:
|
||||||
pos_args.append(
|
pos_args.append(
|
||||||
Function.Argument(
|
Function.Argument(
|
||||||
name=arg.name,
|
name=arg.name,
|
||||||
type=eval_arg_type(arg),
|
type=eval_arg_type(arg),
|
||||||
|
required=arg.default is None,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
for arg in stmt.args:
|
for arg in stmt.args:
|
||||||
@@ -142,6 +153,7 @@ class Checker(
|
|||||||
Function.Argument(
|
Function.Argument(
|
||||||
name=arg.name,
|
name=arg.name,
|
||||||
type=eval_arg_type(arg),
|
type=eval_arg_type(arg),
|
||||||
|
required=arg.default is None,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
for arg in stmt.kwonlyargs:
|
for arg in stmt.kwonlyargs:
|
||||||
@@ -149,6 +161,7 @@ class Checker(
|
|||||||
Function.Argument(
|
Function.Argument(
|
||||||
name=arg.name,
|
name=arg.name,
|
||||||
type=eval_arg_type(arg),
|
type=eval_arg_type(arg),
|
||||||
|
required=arg.default is None,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -175,7 +188,9 @@ class Checker(
|
|||||||
else:
|
else:
|
||||||
returns = inferred_return
|
returns = inferred_return
|
||||||
|
|
||||||
|
# TODO: handle *args and **kwargs sinks
|
||||||
function: Function = Function(
|
function: Function = Function(
|
||||||
|
name=stmt.name,
|
||||||
pos_args=pos_args,
|
pos_args=pos_args,
|
||||||
args=args,
|
args=args,
|
||||||
kw_args=kw_args,
|
kw_args=kw_args,
|
||||||
@@ -240,11 +255,18 @@ class Checker(
|
|||||||
self.import_midas(path)
|
self.import_midas(path)
|
||||||
return UnknownType()
|
return UnknownType()
|
||||||
callee: Type = self.evaluate(expr.callee)
|
callee: Type = self.evaluate(expr.callee)
|
||||||
arguments: list[Type] = [self.evaluate(arg) for arg in expr.arguments]
|
if not isinstance(callee, Function):
|
||||||
keywords: dict[str, Type] = {
|
self.error(expr.callee.location, "Callee is not a function")
|
||||||
name: self.evaluate(arg) for name, arg in expr.keywords.items()
|
return UnknownType()
|
||||||
}
|
function: Function = callee
|
||||||
return UnknownType()
|
mapped: list[MappedArgument] = self.map_call_arguments(function, expr)
|
||||||
|
for arg in mapped:
|
||||||
|
if arg.type != arg.argument.type:
|
||||||
|
self.error(
|
||||||
|
arg.expr.location,
|
||||||
|
f"Wrong type for argument '{arg.argument.name}', expected {arg.argument.type}, got {arg.type}",
|
||||||
|
)
|
||||||
|
return function.returns
|
||||||
|
|
||||||
def visit_get_expr(self, expr: p.GetExpr) -> Type: ...
|
def visit_get_expr(self, expr: p.GetExpr) -> Type: ...
|
||||||
|
|
||||||
@@ -277,3 +299,87 @@ class Checker(
|
|||||||
def visit_frame_column(self, node: p.FrameColumn) -> Type: ...
|
def visit_frame_column(self, node: p.FrameColumn) -> Type: ...
|
||||||
|
|
||||||
def visit_frame_type(self, node: p.FrameType) -> Type: ...
|
def visit_frame_type(self, node: p.FrameType) -> Type: ...
|
||||||
|
|
||||||
|
def map_call_arguments(
|
||||||
|
self, function: Function, call: p.CallExpr
|
||||||
|
) -> list[MappedArgument]:
|
||||||
|
positional: list[tuple[p.Expr, Type]] = [
|
||||||
|
(arg, self.evaluate(arg)) for arg in call.arguments
|
||||||
|
]
|
||||||
|
keywords: dict[str, tuple[p.Expr, Type]] = {
|
||||||
|
name: (arg, self.evaluate(arg)) for name, arg in call.keywords.items()
|
||||||
|
}
|
||||||
|
set_args: set[str] = set()
|
||||||
|
|
||||||
|
required_positional: set[str] = {
|
||||||
|
arg.name for arg in function.pos_args + function.args if arg.required
|
||||||
|
}
|
||||||
|
required_keyword: set[str] = {
|
||||||
|
arg.name for arg in function.kw_args if arg.required
|
||||||
|
}
|
||||||
|
|
||||||
|
mapped: list[MappedArgument] = []
|
||||||
|
|
||||||
|
pos_params: list[Function.Argument] = list(function.pos_args)
|
||||||
|
mixed_params: list[Function.Argument] = list(function.args)
|
||||||
|
kw_params: dict[str, Function.Argument] = {
|
||||||
|
arg.name: arg for arg in function.kw_args
|
||||||
|
}
|
||||||
|
|
||||||
|
# TODO: handle *args and **kwargs sinks
|
||||||
|
for arg in positional:
|
||||||
|
param: Function.Argument
|
||||||
|
if len(pos_params) != 0:
|
||||||
|
param = pos_params.pop(0)
|
||||||
|
elif len(mixed_params) != 0:
|
||||||
|
param = mixed_params.pop(0)
|
||||||
|
else:
|
||||||
|
self.error(arg[0].location, "Too many positional arguments")
|
||||||
|
break
|
||||||
|
required_positional.discard(param.name)
|
||||||
|
required_keyword.discard(param.name)
|
||||||
|
set_args.add(param.name)
|
||||||
|
mapped.append(
|
||||||
|
MappedArgument(
|
||||||
|
expr=arg[0],
|
||||||
|
type=arg[1],
|
||||||
|
argument=param,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
kw_params.update({arg.name: arg for arg in mixed_params})
|
||||||
|
for name, arg in keywords.items():
|
||||||
|
param: Function.Argument
|
||||||
|
if name not in kw_params:
|
||||||
|
if name in set_args:
|
||||||
|
self.error(
|
||||||
|
arg[0].location, f"Multiple values for argument '{name}'"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.error(arg[0].location, f"Unknown keyword argument '{name}'")
|
||||||
|
continue
|
||||||
|
param = kw_params.pop(name)
|
||||||
|
required_positional.discard(name)
|
||||||
|
required_keyword.discard(name)
|
||||||
|
set_args.add(name)
|
||||||
|
mapped.append(
|
||||||
|
MappedArgument(
|
||||||
|
expr=arg[0],
|
||||||
|
type=arg[1],
|
||||||
|
argument=param,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
if len(required_positional) != 0:
|
||||||
|
self.error(
|
||||||
|
call.location,
|
||||||
|
f"Missing required positional arguments: {required_positional}",
|
||||||
|
)
|
||||||
|
|
||||||
|
if len(required_keyword) != 0:
|
||||||
|
self.error(
|
||||||
|
call.location,
|
||||||
|
f"Missing required keyword arguments: {required_keyword}",
|
||||||
|
)
|
||||||
|
|
||||||
|
return mapped
|
||||||
@@ -26,6 +26,7 @@ class UnitType:
|
|||||||
|
|
||||||
@dataclass(frozen=True, kw_only=True)
|
@dataclass(frozen=True, kw_only=True)
|
||||||
class Function:
|
class Function:
|
||||||
|
name: str
|
||||||
pos_args: list[Argument]
|
pos_args: list[Argument]
|
||||||
args: list[Argument]
|
args: list[Argument]
|
||||||
kw_args: list[Argument]
|
kw_args: list[Argument]
|
||||||
@@ -35,6 +36,7 @@ class Function:
|
|||||||
class Argument:
|
class Argument:
|
||||||
name: str
|
name: str
|
||||||
type: Type
|
type: Type
|
||||||
|
required: bool
|
||||||
|
|
||||||
|
|
||||||
Type = BaseType | SimpleType | UnknownType | UnitType | Function
|
Type = BaseType | SimpleType | UnknownType | UnitType | Function
|
||||||
@@ -4,9 +4,9 @@ from abc import ABC, abstractmethod
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Generic, Optional, Protocol, TextIO, TypeVar
|
from typing import Generic, Optional, Protocol, TextIO, TypeVar
|
||||||
|
|
||||||
from midas.ast.location import Location
|
|
||||||
import midas.ast.midas as m
|
import midas.ast.midas as m
|
||||||
import midas.ast.python as p
|
import midas.ast.python as p
|
||||||
|
from midas.ast.location import Location
|
||||||
|
|
||||||
H = TypeVar("H", bound="Highlighter", contravariant=True)
|
H = TypeVar("H", bound="Highlighter", contravariant=True)
|
||||||
|
|
||||||
|
|||||||
+46
-7
@@ -2,7 +2,6 @@ import ast
|
|||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
from midas.ast.location import Location
|
from midas.ast.location import Location
|
||||||
|
|
||||||
from midas.ast.python import (
|
from midas.ast.python import (
|
||||||
AssignStmt,
|
AssignStmt,
|
||||||
BaseType,
|
BaseType,
|
||||||
@@ -136,14 +135,23 @@ class PythonParser:
|
|||||||
args=ast.arguments(
|
args=ast.arguments(
|
||||||
posonlyargs=posonlyargs,
|
posonlyargs=posonlyargs,
|
||||||
args=args,
|
args=args,
|
||||||
|
vararg=sink,
|
||||||
kwonlyargs=kwonlyargs,
|
kwonlyargs=kwonlyargs,
|
||||||
|
kwarg=kw_sink,
|
||||||
|
defaults=defaults,
|
||||||
|
kw_defaults=kw_defaults,
|
||||||
),
|
),
|
||||||
returns=returns,
|
returns=returns,
|
||||||
body=raw_body,
|
body=raw_body,
|
||||||
):
|
):
|
||||||
|
|
||||||
def parse_args(args_list: list[ast.arg]) -> list[Function.Argument]:
|
def parse_args(
|
||||||
return [self._parse_function_argument(arg) for arg in args_list]
|
args_list: list[ast.arg], defaults: list[Optional[Expr]]
|
||||||
|
) -> list[Function.Argument]:
|
||||||
|
return [
|
||||||
|
self._parse_function_argument(arg, default)
|
||||||
|
for arg, default in zip(args_list, defaults)
|
||||||
|
]
|
||||||
|
|
||||||
body: list[Stmt] = []
|
body: list[Stmt] = []
|
||||||
for stmt in raw_body:
|
for stmt in raw_body:
|
||||||
@@ -152,19 +160,49 @@ class PythonParser:
|
|||||||
body.append(stmts)
|
body.append(stmts)
|
||||||
elif stmts is not None:
|
elif stmts is not None:
|
||||||
body.extend(stmts)
|
body.extend(stmts)
|
||||||
|
|
||||||
|
parsed_defaults: list[Optional[Expr]] = [
|
||||||
|
self.parse_expr(default) for default in defaults
|
||||||
|
]
|
||||||
|
n_posargs: int = len(posonlyargs)
|
||||||
|
n_args: int = len(args)
|
||||||
|
n_all_posargs = n_posargs + n_args
|
||||||
|
parsed_defaults = [
|
||||||
|
None,
|
||||||
|
] * (n_all_posargs - len(defaults)) + parsed_defaults
|
||||||
|
|
||||||
|
posargs_defaults: list[Optional[Expr]] = parsed_defaults[:n_posargs]
|
||||||
|
args_defaults: list[Optional[Expr]] = parsed_defaults[n_posargs:]
|
||||||
|
kwargs_defaults: list[Optional[Expr]] = [
|
||||||
|
self.parse_expr(default) if default is not None else None
|
||||||
|
for default in kw_defaults
|
||||||
|
]
|
||||||
|
|
||||||
return Function(
|
return Function(
|
||||||
location=loc,
|
location=loc,
|
||||||
name=name,
|
name=name,
|
||||||
posonlyargs=parse_args(posonlyargs),
|
posonlyargs=parse_args(posonlyargs, posargs_defaults),
|
||||||
args=parse_args(args),
|
args=parse_args(args, args_defaults),
|
||||||
kwonlyargs=parse_args(kwonlyargs),
|
sink=(
|
||||||
|
self._parse_function_argument(sink, None)
|
||||||
|
if sink is not None
|
||||||
|
else None
|
||||||
|
),
|
||||||
|
kwonlyargs=parse_args(kwonlyargs, kwargs_defaults),
|
||||||
|
kw_sink=(
|
||||||
|
self._parse_function_argument(kw_sink, None)
|
||||||
|
if kw_sink is not None
|
||||||
|
else None
|
||||||
|
),
|
||||||
returns=self._parse_type(returns) if returns is not None else None,
|
returns=self._parse_type(returns) if returns is not None else None,
|
||||||
body=body,
|
body=body,
|
||||||
)
|
)
|
||||||
case _:
|
case _:
|
||||||
print(f"Unsupported function definition: {ast.unparse(node)}")
|
print(f"Unsupported function definition: {ast.unparse(node)}")
|
||||||
|
|
||||||
def _parse_function_argument(self, arg: ast.arg) -> Function.Argument:
|
def _parse_function_argument(
|
||||||
|
self, arg: ast.arg, default: Optional[Expr]
|
||||||
|
) -> Function.Argument:
|
||||||
loc: Location = Location.from_ast(arg)
|
loc: Location = Location.from_ast(arg)
|
||||||
name: str = arg.arg
|
name: str = arg.arg
|
||||||
type: Optional[MidasType] = None
|
type: Optional[MidasType] = None
|
||||||
@@ -174,6 +212,7 @@ class PythonParser:
|
|||||||
location=loc,
|
location=loc,
|
||||||
name=name,
|
name=name,
|
||||||
type=type,
|
type=type,
|
||||||
|
default=default,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _parse_type(
|
def _parse_type(
|
||||||
|
|||||||
Reference in new issue
Block a user