Compare commits

...
2 Commits
Author SHA1 Message Date
HEL 9b59058881 feat(cli): add highlight command 2026-05-22 22:16:05 +02:00
HEL d0c54db33a feat(parser): store locations in parsed nodes 2026-05-22 22:11:44 +02:00
5 changed files with 277 additions and 9 deletions

No files matched your search

+28 -2
View File
@@ -3,13 +3,39 @@ from __future__ import annotations
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
import ast import ast
from dataclasses import dataclass from dataclasses import dataclass
from typing import Generic, Optional, TypeVar from typing import Generic, Optional, Protocol, TypeVar
T = TypeVar("T") T = TypeVar("T")
@dataclass(frozen=True) class HasLocation(Protocol):
lineno: int
col_offset: int
end_lineno: Optional[int]
end_col_offset: Optional[int]
@dataclass(frozen=True, kw_only=True)
class Location:
lineno: int
col_offset: int
end_lineno: Optional[int]
end_col_offset: Optional[int]
@staticmethod
def from_ast(obj: HasLocation) -> Location:
return Location(
lineno=obj.lineno,
col_offset=obj.col_offset,
end_lineno=obj.end_lineno,
end_col_offset=obj.end_col_offset,
)
@dataclass(frozen=True, kw_only=True)
class Expr(ABC): class Expr(ABC):
location: Optional[Location] = None
@abstractmethod @abstractmethod
def accept(self, visitor: Visitor[T]) -> T: ... def accept(self, visitor: Visitor[T]) -> T: ...
+81
View File
@@ -0,0 +1,81 @@
html,
body {
margin: 0;
font-size: 14pt;
}
* {
box-sizing: border-box;
}
#code {
display: flex;
flex-direction: column;
font-family: monospace;
white-space: pre-wrap;
}
.line {
display: flex;
&:nth-child(odd) {
background-color: rgb(247, 247, 247);
}
.no {
width: 4em;
text-align: right;
padding: 0.2em 0.4em;
border-right: solid black 1px;
flex-shrink: 0;
}
.txt {
flex-grow: 1;
padding: 0.2em 0.8em;
}
}
span {
--col: transparent;
--opacity: 0.1;
--border: 0px;
background-color: rgba(var(--col), var(--opacity));
outline: solid rgb(var(--col)) var(--border);
outline-offset: 2px;
border-radius: 2px;
&:hover:not(:has(*:hover)) {
--opacity: 0.8;
--border: 2px;
z-index: 10;
}
&.base-type {
--col: 108, 233, 108;
}
&.param {
--col: 103, 192, 224;
}
&.constraint-type {
--col: 174, 200, 195;
}
&.frame-column {
--col: 216, 231, 81;
}
&.frame-type {
--col: 231, 46, 40;
}
&.function {
--col: 215, 103, 224;
}
&.argument {
--col: 103, 192, 224;
}
}
+115
View File
@@ -0,0 +1,115 @@
from pathlib import Path
from typing import TextIO
from midas.ast.python import (
BaseType,
ConstraintType,
Expr,
FrameColumn,
FrameType,
Function,
FunctionArgument,
)
class PythonHighlighter(Expr.Visitor[None]):
CSS_PATH: Path = Path(__file__).parent / "highlight.css"
def __init__(self, source: str) -> None:
self.source: str = source
self.lines: list[str] = self.source.splitlines()
self.openings: dict[tuple[int, int], list[str]] = {}
self.closings: dict[tuple[int, int], list[str]] = {}
def highlight(self, node: Expr):
node.accept(self)
def dump(self, buf: TextIO):
css: str = self.CSS_PATH.read_text()
css = "\n".join((" " + line).rstrip() for line in css.splitlines())
lines: list[str] = [
"<!DOCTYPE html>",
'<html lang="en">',
"<head>",
' <meta charset="UTF-8">',
' <meta name="viewport" content="width=device-width, initial-scale=1.0">',
" <title>Highlighted file</title>",
" <style>",
css,
" </style>",
"</head>",
"<body>",
' <div id="code">',
]
for l, line in enumerate(self.lines):
lineno: int = l + 1
line_buf: str = (
f'<div class="line" id="l{lineno}"><div class="no">{lineno}</div><div class="txt">'
)
for c, char in enumerate(line):
pos: tuple[int, int] = (lineno, c)
closings: list[str] = self.closings.get(pos, [])
openings: list[str] = self.openings.get(pos, [])
line_buf += "".join(closings + openings)
line_buf += char
line_buf += "</div></div>"
lines.append(" " + line_buf)
lines.extend(
[
" </div>",
"</body>",
"</html>",
]
)
buf.write("\n".join(lines))
def wrap(self, node: Expr, cls: str):
if node.location is None:
return
if node.location.end_lineno is None or node.location.end_col_offset is None:
return
start_pos: tuple[int, int] = (node.location.lineno, node.location.col_offset)
end_pos: tuple[int, int] = (
node.location.end_lineno,
node.location.end_col_offset,
)
opening: str = f'<span class="{cls}" title="{cls}">'
closing: str = "</span>"
self.openings.setdefault(start_pos, []).append(opening)
self.closings.setdefault(end_pos, []).insert(0, closing)
if start_pos[0] != end_pos[0]:
for l in range(start_pos[0], end_pos[0]):
c: int = len(self.lines[l - 1])
self.closings.setdefault((l, c), []).insert(0, closing)
self.openings.setdefault((l + 1, 0), []).append(opening)
def visit_base_type(self, node: BaseType) -> None:
self.wrap(node, "base-type")
if node.param is not None:
self.wrap(node.param, "param")
node.param.accept(self)
def visit_constraint_type(self, node: ConstraintType) -> None:
self.wrap(node, "constraint-type")
node.type.accept(self)
def visit_frame_column(self, node: FrameColumn) -> None:
self.wrap(node, "frame-column")
if node.type is not None:
node.type.accept(self)
def visit_frame_type(self, node: FrameType) -> None:
self.wrap(node, "frame-type")
for column in node.columns:
column.accept(self)
def visit_function(self, node: Function) -> None:
self.wrap(node, "function")
for arg in node.posonlyargs + node.args + node.kwonlyargs:
arg.accept(self)
def visit_function_argument(self, node: FunctionArgument) -> None:
self.wrap(node, "argument")
if node.type is not None:
node.type.accept(self)
+18
View File
@@ -4,6 +4,7 @@ from typing import Optional, TextIO
import click import click
from midas.ast.printer import PythonAstPrinter from midas.ast.printer import PythonAstPrinter
from midas.cli.highlighter import PythonHighlighter
from midas.parser.python import PythonParser from midas.parser.python import PythonParser
@@ -56,3 +57,20 @@ def dump_ast(output: Optional[TextIO], parse: bool, file: TextIO):
click.echo(dump) click.echo(dump)
else: else:
output.write(dump) output.write(dump)
@utils.command()
@click.option("-o", "--output", type=click.File("w"), default="-")
@click.argument("file", type=click.File("r"))
def highlight(output: TextIO, file: TextIO):
source: str = file.read()
tree: ast.Module = ast.parse(source, filename=file.name)
parser = PythonParser()
parser.visit(tree)
highlighter: PythonHighlighter = PythonHighlighter(source)
for _, annotation in parser.annotations:
if annotation is not None:
highlighter.highlight(annotation)
for func in parser.functions:
highlighter.highlight(func)
highlighter.dump(output)
+35 -7
View File
@@ -8,6 +8,7 @@ from midas.ast.python import (
FrameType, FrameType,
Function, Function,
FunctionArgument, FunctionArgument,
Location,
MidasType, MidasType,
) )
@@ -50,6 +51,7 @@ class PythonParser(ast.NodeVisitor):
self.generic_visit(node) self.generic_visit(node)
def _parse_function(self, node: ast.FunctionDef) -> Function: def _parse_function(self, node: ast.FunctionDef) -> Function:
loc: Location = Location.from_ast(node)
match node: match node:
case ast.FunctionDef( case ast.FunctionDef(
name=name, name=name,
@@ -65,6 +67,7 @@ class PythonParser(ast.NodeVisitor):
return [self._parse_function_argument(arg) for arg in args_list] return [self._parse_function_argument(arg) for arg in args_list]
return Function( return Function(
location=loc,
name=name, name=name,
posonlyargs=parse_args(posonlyargs), posonlyargs=parse_args(posonlyargs),
args=parse_args(args), args=parse_args(args),
@@ -73,24 +76,38 @@ class PythonParser(ast.NodeVisitor):
) )
def _parse_function_argument(self, arg: ast.arg) -> FunctionArgument: def _parse_function_argument(self, arg: ast.arg) -> FunctionArgument:
loc: Location = Location.from_ast(arg)
name: str = arg.arg name: str = arg.arg
type: Optional[MidasType] = None type: Optional[MidasType] = None
if arg.annotation is not None: if arg.annotation is not None:
type = self._parse_type(arg.annotation) type = self._parse_type(arg.annotation)
return FunctionArgument(name=name, type=type) return FunctionArgument(
location=loc,
name=name,
type=type,
)
def _parse_type( def _parse_type(
self, type_expr: ast.expr, root: bool = False self, type_expr: ast.expr, root: bool = False
) -> Optional[MidasType]: ) -> Optional[MidasType]:
loc: Location = Location.from_ast(type_expr)
match type_expr: match type_expr:
case ast.Subscript(value=ast.Name(id="Frame"), slice=schema): case ast.Subscript(value=ast.Name(id="Frame"), slice=schema):
return self._parse_frame_type(schema) return self._parse_frame_type(schema)
case ast.Subscript(value=ast.Name(id=name), slice=param): case ast.Subscript(value=ast.Name(id=name), slice=param):
return BaseType(base=name, param=self._parse_type(param)) return BaseType(
location=loc,
base=name,
param=self._parse_type(param),
)
case ast.Name(id=name): case ast.Name(id=name):
return BaseType(base=name, param=None) return BaseType(
location=loc,
base=name,
param=None,
)
case ast.BinOp(left=left_expr, op=ast.Add(), right=right_expr): case ast.BinOp(left=left_expr, op=ast.Add(), right=right_expr):
left = self._parse_type(left_expr) left = self._parse_type(left_expr)
@@ -107,12 +124,17 @@ class PythonParser(ast.NodeVisitor):
) )
ast.copy_location(constraint, type_expr) ast.copy_location(constraint, type_expr)
return ConstraintType( return ConstraintType(
location=loc,
type=left_type, type=left_type,
constraint=constraint, constraint=constraint,
) )
case _: case _:
return ConstraintType(type=left, constraint=right_expr) return ConstraintType(
location=loc,
type=left,
constraint=right_expr,
)
case _: case _:
if root: if root:
@@ -120,6 +142,7 @@ class PythonParser(ast.NodeVisitor):
raise UnsupportedSyntaxError(type_expr) raise UnsupportedSyntaxError(type_expr)
def _parse_frame_type(self, schema: ast.expr) -> FrameType: def _parse_frame_type(self, schema: ast.expr) -> FrameType:
loc: Location = Location.from_ast(schema)
columns: list[FrameColumn] = [] columns: list[FrameColumn] = []
match schema: match schema:
@@ -133,12 +156,17 @@ class PythonParser(ast.NodeVisitor):
case _: case _:
raise UnsupportedSyntaxError(schema) raise UnsupportedSyntaxError(schema)
return FrameType(columns=columns) return FrameType(location=loc, columns=columns)
def _parse_frame_column(self, column: ast.expr) -> FrameColumn: def _parse_frame_column(self, column: ast.expr) -> FrameColumn:
loc: Location = Location.from_ast(column)
match column: match column:
case ast.Name(): case ast.Name():
return FrameColumn(name=None, type=self._parse_type(column)) return FrameColumn(
location=loc,
name=None,
type=self._parse_type(column),
)
case ast.Slice(lower=ast.Name(id=name), upper=type_expr): case ast.Slice(lower=ast.Name(id=name), upper=type_expr):
if name == "_": if name == "_":
@@ -154,7 +182,7 @@ class PythonParser(ast.NodeVisitor):
type = self._parse_type(type_expr) type = self._parse_type(type_expr)
case _: case _:
raise UnsupportedSyntaxError(type_expr) raise UnsupportedSyntaxError(type_expr)
return FrameColumn(name=name, type=type) return FrameColumn(location=loc, name=name, type=type)
case _: case _:
raise UnsupportedSyntaxError(column) raise UnsupportedSyntaxError(column)