Compare commits

...
3 Commits
Author SHA1 Message Date
HEL e6375f1aa9 chore: tidy 2026-05-29 17:25:12 +02:00
HEL d16e192a3a feat(checker): map and check function call arguments 2026-05-29 15:49:51 +02:00
HEL 3f61f84e5a feat(parser): parse function param defaults and sinks 2026-05-29 15:47:19 +02:00
7 changed files with 170 additions and 17 deletions

No files matched your search

+1 -1
View File
@@ -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.
+3
View File
@@ -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]:
+3
View File
@@ -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
View File
@@ -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
+2
View File
@@ -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
+1 -1
View File
@@ -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
View File
@@ -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(