4 Commits
7 changed files with 84 additions and 41 deletions
+1
View File
@@ -135,6 +135,7 @@ class ExtensionType:
class FunctionType: class FunctionType:
pos_args: list[Argument] pos_args: list[Argument]
args: list[Argument]
kw_args: list[Argument] kw_args: list[Argument]
returns: Type returns: Type
+1
View File
@@ -293,6 +293,7 @@ class ExtensionType(Type):
@dataclass(frozen=True) @dataclass(frozen=True)
class FunctionType(Type): class FunctionType(Type):
pos_args: list[Argument] pos_args: list[Argument]
args: list[Argument]
kw_args: list[Argument] kw_args: list[Argument]
returns: Type returns: Type
+10
View File
@@ -297,6 +297,14 @@ class MidasAstPrinter(
self._mark_last() self._mark_last()
self._print_function_arg(arg) self._print_function_arg(arg)
self._write_line("args")
with self._child_level():
for i, arg in enumerate(type.args):
self._idx = i
if i == len(type.args) - 1:
self._mark_last()
self._print_function_arg(arg)
self._write_line("kw_args") self._write_line("kw_args")
with self._child_level(): with self._child_level():
for i, arg in enumerate(type.kw_args): for i, arg in enumerate(type.kw_args):
@@ -447,11 +455,13 @@ class MidasPrinter(m.Expr.Visitor[str], m.Stmt.Visitor[str], m.Type.Visitor[str]
def visit_function_type(self, type: m.FunctionType) -> str: def visit_function_type(self, type: m.FunctionType) -> str:
pos_args: list[str] = [self._print_arg(arg) for arg in type.pos_args] pos_args: list[str] = [self._print_arg(arg) for arg in type.pos_args]
mixed_args: list[str] = [self._print_arg(arg) for arg in type.args]
kw_args: list[str] = [self._print_arg(arg) for arg in type.kw_args] kw_args: list[str] = [self._print_arg(arg) for arg in type.kw_args]
args: list[str] = pos_args args: list[str] = pos_args
if len(pos_args) != 0: if len(pos_args) != 0:
args.append("/") args.append("/")
args += mixed_args
if len(kw_args) != 0: if len(kw_args) != 0:
args.append("*") args.append("*")
args += kw_args args += kw_args
+17 -17
View File
@@ -43,6 +43,8 @@ class MidasTyper(m.Stmt.Visitor[None], m.Expr.Visitor[None], m.Type.Visitor[Type
tokens: list[Token] = lexer.process() tokens: list[Token] = lexer.process()
parser: MidasParser = MidasParser(tokens) parser: MidasParser = MidasParser(tokens)
stmts: list[m.Stmt] = parser.parse() stmts: list[m.Stmt] = parser.parse()
for error in parser.errors:
self.reporter.error(error.token.get_location(), error.message)
self.resolve(stmts) self.resolve(stmts)
def get_type(self, name: str) -> Type: def get_type(self, name: str) -> Type:
@@ -164,25 +166,23 @@ class MidasTyper(m.Stmt.Visitor[None], m.Expr.Visitor[None], m.Type.Visitor[Type
) )
def visit_function_type(self, type: m.FunctionType) -> Type: def visit_function_type(self, type: m.FunctionType) -> Type:
n_pos_args: int = len(type.pos_args)
n_args: int = len(type.args)
def process_arg(arg: m.FunctionType.Argument, i: int) -> Function.Argument:
return Function.Argument(
pos=i,
name=arg.name.lexeme if arg.name is not None else str(i),
type=arg.type.accept(self),
required=arg.required,
)
return Function( return Function(
pos_args=[ pos_args=[process_arg(arg, i) for i, arg in enumerate(type.pos_args)],
Function.Argument( args=[process_arg(arg, i + n_pos_args) for i, arg in enumerate(type.args)],
pos=i,
name=arg.name.lexeme if arg.name is not None else str(i),
type=arg.type.accept(self),
required=arg.required,
)
for i, arg in enumerate(type.pos_args)
],
args=[],
kw_args=[ kw_args=[
Function.Argument( process_arg(arg, i + n_pos_args + n_args)
pos=i, for i, arg in enumerate(type.kw_args)
name=arg.name.lexeme if arg.name is not None else str(i),
type=arg.type.accept(self),
required=arg.required,
)
for i, arg in enumerate(type.kw_args, start=len(type.pos_args))
], ],
returns=type.returns.accept(self), returns=type.returns.accept(self),
) )
+8 -2
View File
@@ -336,7 +336,7 @@ class PythonTyper(
if not self._is_binary_function(function): if not self._is_binary_function(function):
self.reporter.error( self.reporter.error(
expr.location, expr.location,
f"Wrong definition of binary operation. Expected function with 2 positional-only parameters, got {function}", f"Wrong definition of binary operation. Expected function with 1 positional-only parameters, got {function}",
) )
return UnknownType() return UnknownType()
@@ -481,7 +481,13 @@ class PythonTyper(
return self.types.apply_generic(list_type, [UnknownType()]) return self.types.apply_generic(list_type, [UnknownType()])
def visit_base_type(self, node: p.BaseType) -> Type: def visit_base_type(self, node: p.BaseType) -> Type:
base: Type = self.types.get_type(node.base) base: Type
try:
base = self.types.get_type(node.base)
except NameError:
self.reporter.warning(node.location, f"Unknown type '{node.base}'")
return UnknownType()
if node.param is not None: if node.param is not None:
param: Type = node.param.accept(self) param: Type = node.param.accept(self)
return self.types.apply_generic(base, [param]) return self.types.apply_generic(base, [param])
+34 -10
View File
@@ -499,19 +499,36 @@ class MidasParser(Parser):
TokenType.LEFT_PAREN, "Expected '(' before function parameters" TokenType.LEFT_PAREN, "Expected '(' before function parameters"
) )
pos_args: list[FunctionType.Argument] = [] pos_args: list[FunctionType.Argument] = []
args: list[FunctionType.Argument] = []
kw_args: list[FunctionType.Argument] = [] kw_args: list[FunctionType.Argument] = []
positional: bool = True args_first_tokens: list[Token] = []
section: int = 0
while not self.is_at_end() and not self.check(TokenType.RIGHT_PAREN): while not self.is_at_end() and not self.check(TokenType.RIGHT_PAREN):
if positional and ( match section:
self.match(TokenType.STAR) or self.match(TokenType.SLASH) case 0 if self.match(TokenType.SLASH):
): pos_args = args
positional = False args = []
else: args_first_tokens = []
section = 1
case 0 | 1 if self.match(TokenType.STAR):
section = 2
case _:
# Record first token of mixed argument for errors if unnamed
if section != 2:
args_first_tokens.append(self.peek())
name: Optional[Token] = None name: Optional[Token] = None
if self.check_identifier() and self.check_next(TokenType.COLON): if section == 2:
name = self.consume_identifier("Expected keyword argument name")
self.consume(
TokenType.COLON, "Expected ':' after argument name"
)
elif self.check_identifier() and self.check_next(TokenType.COLON):
name = self.advance() name = self.advance()
self.advance() self.advance()
type: Type = self.type_expr() type: Type = self.type_expr()
optional: bool = self.match(TokenType.QMARK) optional: bool = self.match(TokenType.QMARK)
arg = FunctionType.Argument( arg = FunctionType.Argument(
@@ -520,13 +537,19 @@ class MidasParser(Parser):
type=type, type=type,
required=not optional, required=not optional,
) )
if positional: if section == 2:
pos_args.append(arg)
else:
kw_args.append(arg) kw_args.append(arg)
else:
args.append(arg)
if not self.match(TokenType.COMMA): if not self.match(TokenType.COMMA):
break break
for arg, token in zip(args, args_first_tokens):
if arg.name is None:
# Not raised because we can keep parsing
self.error(token, "Unnamed mixed argument")
self.consume(TokenType.RIGHT_PAREN, "Expected ')' after function parameters") self.consume(TokenType.RIGHT_PAREN, "Expected ')' after function parameters")
self.consume(TokenType.ARROW, "Expected '->' before result type") self.consume(TokenType.ARROW, "Expected '->' before result type")
@@ -535,6 +558,7 @@ class MidasParser(Parser):
return FunctionType( return FunctionType(
location=l_paren.location_to(self.previous()), location=l_paren.location_to(self.previous()),
pos_args=pos_args, pos_args=pos_args,
args=args,
kw_args=kw_args, kw_args=kw_args,
returns=result, returns=result,
) )
+1
View File
@@ -173,6 +173,7 @@ class MidasAstJsonSerializer(
return { return {
"_type": "FunctionType", "_type": "FunctionType",
"pos_args": [self._serialize_func_arg(arg) for arg in type.pos_args], "pos_args": [self._serialize_func_arg(arg) for arg in type.pos_args],
"args": [self._serialize_func_arg(arg) for arg in type.args],
"kw_args": [self._serialize_func_arg(arg) for arg in type.kw_args], "kw_args": [self._serialize_func_arg(arg) for arg in type.kw_args],
"returns": type.returns.accept(self), "returns": type.returns.accept(self),
} }