better error handling for incompatible args

This commit is contained in:
Jesse Bannon 2023-11-15 22:47:57 -08:00
parent 8fe4460ea5
commit 96182af2c1
5 changed files with 32 additions and 12 deletions

View file

@ -19,6 +19,7 @@ from ytdl_sub.script.types.variable import Variable
from ytdl_sub.script.utils.exceptions import IncompatibleFunctionArguments
from ytdl_sub.script.utils.exceptions import InvalidSyntaxException
from ytdl_sub.script.utils.exceptions import UnreachableSyntaxException
from ytdl_sub.script.utils.exceptions import UserException
from ytdl_sub.script.utils.parser_exception_formatter import ParserExceptionFormatter
from ytdl_sub.utils.exceptions import StringFormattingException
from ytdl_sub.validators.string_formatter_validators import is_valid_source_variable_name
@ -112,7 +113,7 @@ class _Parser:
try:
self._syntax_tree = self._parse()
except InvalidSyntaxException as exc:
except UserException as exc:
raise ParserExceptionFormatter(
self._text, self._error_highlight_pos, self._pos, exc
).highlight() from exc
@ -329,9 +330,9 @@ class _Parser:
if ch == ")":
try:
return Function.from_name_and_args(name=function_name, args=function_args)
except IncompatibleFunctionArguments as exc:
except IncompatibleFunctionArguments:
self._set_highlight_position(function_start_pos)
raise InvalidSyntaxException(exc) from exc
raise
if ch != "(":
function_name += ch

View file

@ -219,7 +219,7 @@ class BuiltInFunction(Function):
received_args_str = f"({', '.join([name for name in received_type_names])})"
return f"Expected {self.input_spec.expected_args_str()}.\nReceived {received_args_str}"
return f"Expected {self.input_spec.expected_args_str()}\nReceived {received_args_str}"
def validate_args(self) -> "BuiltInFunction":
if not self.input_spec.is_compatible(input_args=self.args):

View file

@ -1,11 +1,17 @@
from abc import ABC
from ytdl_sub.utils.exceptions import ValidationException
class InvalidSyntaxException(ValidationException):
class UserException(ValidationException, ABC):
"""It's the user's fault!"""
class InvalidSyntaxException(UserException):
"""Syntax is incorrect"""
class IncompatibleFunctionArguments(ValidationException):
class IncompatibleFunctionArguments(UserException):
"""Function has invalid arguments"""

View file

@ -1,11 +1,15 @@
import sys
from typing import List
from typing import TypeVar
from ytdl_sub.script.utils.exceptions import InvalidSyntaxException
from ytdl_sub.script.utils.exceptions import UserException
TUserException = TypeVar("TUserException", bound=UserException)
class ParserExceptionFormatter:
def __init__(self, text: str, start: int, end: int, exception: InvalidSyntaxException):
def __init__(self, text: str, start: int, end: int, exception: TUserException):
self._text = text
self._start = start
self._end = end
@ -65,10 +69,10 @@ class ParserExceptionFormatter:
return "\n" + "\n".join(to_return)
def highlight(self) -> InvalidSyntaxException:
def highlight(self) -> TUserException:
if self.is_multi_line:
invalid_syntax = self.exception_text_lines(border_lines=3)
else:
invalid_syntax = self.exception_text(border=20)
return InvalidSyntaxException(f"{invalid_syntax}\n{str(self._exception)}")
return self._exception.__class__(f"{invalid_syntax}\n{str(self._exception)}")

View file

@ -15,9 +15,14 @@ from ytdl_sub.script.types.resolvable import Boolean
from ytdl_sub.script.types.resolvable import Float
from ytdl_sub.script.types.resolvable import Integer
from ytdl_sub.script.types.resolvable import String
from ytdl_sub.script.utils.exceptions import IncompatibleFunctionArguments
from ytdl_sub.script.utils.exceptions import InvalidSyntaxException
def _incompatible_arguments_match(expected: str, recieved: str) -> str:
return re.escape(f"Expected ({expected})\nReceived ({recieved})")
class TestFunction:
@pytest.mark.parametrize(
"function_str, expected_output",
@ -66,9 +71,13 @@ class TestFunction:
"key"
)
}"""
assert Script({"func": function_str}).resolve() == {
"func": String("winner"),
}
with pytest.raises(
IncompatibleFunctionArguments,
match=_incompatible_arguments_match(
expected="Map, Hashable, Optional", recieved="%if(...)->Union[Map, Array], String"
),
):
Script({"func": function_str}).resolve()
@pytest.mark.parametrize(
"function_str", ["{%array_at({'a': 'dict?'}, 1)}" "{%array_extend('not', 'array')}"]