custom functions

This commit is contained in:
Jesse Bannon 2023-11-09 17:02:03 -08:00
parent 1724893b27
commit 99d38b55fb
10 changed files with 317 additions and 44 deletions

View file

@ -12,6 +12,7 @@ from ytdl_sub.script.types.resolvable import Boolean
from ytdl_sub.script.types.resolvable import Float from ytdl_sub.script.types.resolvable import Float
from ytdl_sub.script.types.resolvable import Integer from ytdl_sub.script.types.resolvable import Integer
from ytdl_sub.script.types.resolvable import String from ytdl_sub.script.types.resolvable import String
from ytdl_sub.script.types.variable import FunctionArgument
from ytdl_sub.script.types.variable import Variable from ytdl_sub.script.types.variable import Variable
from ytdl_sub.utils.exceptions import StringFormattingException from ytdl_sub.utils.exceptions import StringFormattingException
from ytdl_sub.validators.string_formatter_validators import is_valid_source_variable_name from ytdl_sub.validators.string_formatter_validators import is_valid_source_variable_name
@ -71,6 +72,30 @@ class _Parser:
assert is_valid_source_variable_name(var_name, raise_exception=False) assert is_valid_source_variable_name(var_name, raise_exception=False)
return Variable(var_name) return Variable(var_name)
def _parse_function_argument(self) -> FunctionArgument:
"""
Begin parsing function args after the first ``$``, i.e. ``$1``
"""
var_name = ""
while ch := self._read(increment_pos=False):
if ch.isspace() and not var_name:
self._pos += 1
continue
if ch in ["}", ",", ")", "]"] or ch.isspace():
break
is_numeric = ch.isnumeric()
if not is_numeric:
raise StringFormattingException("invalid function var name")
var_name += ch
self._pos += 1
if not var_name:
raise StringFormattingException("invalid var name")
return FunctionArgument(name=f"${var_name}")
def _parse_numeric(self) -> Integer | Float: def _parse_numeric(self) -> Integer | Float:
numeric_string = "" numeric_string = ""
while ch := self._read(increment_pos=False): while ch := self._read(increment_pos=False):
@ -125,6 +150,9 @@ class _Parser:
if self._read(increment_pos=False) == "{": if self._read(increment_pos=False) == "{":
self._pos += 1 self._pos += 1
return self._parse_map() return self._parse_map()
if self._read(increment_pos=False) == "$":
self._pos += 1
return self._parse_function_argument()
if self._read(increment_pos=False).isascii() and self._read(increment_pos=False).islower(): if self._read(increment_pos=False).isascii() and self._read(increment_pos=False).islower():
return self._parse_variable() return self._parse_variable()
raise StringFormattingException( raise StringFormattingException(
@ -167,7 +195,7 @@ class _Parser:
while ch := self._read(): while ch := self._read():
if ch == ")": if ch == ")":
return Function(name=function_name, args=function_args) return Function.from_name_and_args(name=function_name, args=function_args)
if ch != "(": if ch != "(":
function_name += ch function_name += ch

View file

@ -0,0 +1,36 @@
from typing import Dict, Optional
from ytdl_sub.script.parser import parse
from ytdl_sub.script.syntax_tree import SyntaxTree
from ytdl_sub.script.types.resolvable import Resolvable
from ytdl_sub.script.types.variable import Variable
class Script:
@classmethod
def _is_function(cls, override_name: str):
return override_name.startswith("%")
@classmethod
def _function_name(self, function_key: str) -> str:
return function_key[1:]
def __init__(self, overrides: Dict[str, str]):
self._functions: Dict[str, SyntaxTree] = {
self._function_name(function_key): parse(function_value)
for function_key, function_value in overrides.items()
if self._is_function(function_key)
}
self._variables: Dict[str, SyntaxTree] = {
override_name: parse(override_value)
for override_name, override_value in overrides.items()
if not self._is_function(override_name)
}
def resolve(self, pre_resolved_variables: Optional[Dict[Variable, Resolvable]] = None) -> Dict[str, Resolvable]:
return SyntaxTree.resolve_overrides(
parsed_overrides=self._variables,
custom_functions=self._functions,
pre_resolved_variables=pre_resolved_variables,
)

View file

@ -1,5 +1,5 @@
from dataclasses import dataclass from dataclasses import dataclass
from typing import Dict from typing import Dict, Optional
from typing import List from typing import List
from typing import Set from typing import Set
@ -7,6 +7,7 @@ from ytdl_sub.script.types.function import Function
from ytdl_sub.script.types.resolvable import ArgumentType from ytdl_sub.script.types.resolvable import ArgumentType
from ytdl_sub.script.types.resolvable import Resolvable from ytdl_sub.script.types.resolvable import Resolvable
from ytdl_sub.script.types.resolvable import String from ytdl_sub.script.types.resolvable import String
from ytdl_sub.script.types.variable import FunctionArgument
from ytdl_sub.script.types.variable import Variable from ytdl_sub.script.types.variable import Variable
from ytdl_sub.script.types.variable_dependency import VariableDependency from ytdl_sub.script.types.variable_dependency import VariableDependency
from ytdl_sub.utils.exceptions import StringFormattingException from ytdl_sub.utils.exceptions import StringFormattingException
@ -27,16 +28,40 @@ class SyntaxTree(VariableDependency):
for token in self.ast: for token in self.ast:
if isinstance(token, Variable): if isinstance(token, Variable):
variables.add(token) variables.add(token)
elif isinstance(token, Function): elif isinstance(token, VariableDependency):
variables.update(token.variables) variables.update(token.variables)
return variables return variables
def resolve(self, resolved_variables: Dict[Variable, Resolvable]) -> Resolvable: @property
def function_arguments(self) -> Set[FunctionArgument]:
"""
Returns
-------
All function arguments used within the SyntaxTree
"""
function_arguments: Set[FunctionArgument] = set()
for token in self.ast:
if isinstance(token, FunctionArgument):
function_arguments.add(token)
elif isinstance(token, VariableDependency):
function_arguments.update(token.function_arguments)
return function_arguments
def resolve(
self,
resolved_variables: Dict[Variable, Resolvable],
custom_functions: Dict[str, "VariableDependency"],
) -> Resolvable:
resolved: List[Resolvable] = [] resolved: List[Resolvable] = []
for token in self.ast: for token in self.ast:
resolved.append( resolved.append(
self._resolve_argument_type(resolved_variables=resolved_variables, arg=token) self._resolve_argument_type(
arg=token,
resolved_variables=resolved_variables,
custom_functions=custom_functions,
)
) )
# If only one resolvable resides in the AST, return as that # If only one resolvable resides in the AST, return as that
@ -47,13 +72,15 @@ class SyntaxTree(VariableDependency):
return String("".join([str(res) for res in resolved])) return String("".join([str(res) for res in resolved]))
@classmethod @classmethod
def resolve_overrides(cls, parsed_overrides: Dict[str, "SyntaxTree"]) -> Dict[str, Resolvable]: def resolve_overrides(
cls, parsed_overrides: Dict[str, "SyntaxTree"], custom_functions: Dict[str, "SyntaxTree"], pre_resolved_variables: Optional[Dict[Variable, Resolvable]]
) -> Dict[str, Resolvable]:
overrides: Dict[Variable, "SyntaxTree"] = { overrides: Dict[Variable, "SyntaxTree"] = {
Variable(name): ast for name, ast in parsed_overrides.items() Variable(name): ast for name, ast in parsed_overrides.items()
} }
unresolved_variables: List[Variable] = list(overrides.keys()) unresolved_variables: List[Variable] = list(overrides.keys())
resolved_variables: Dict[Variable, Resolvable] = {} resolved_variables: Dict[Variable, Resolvable] = pre_resolved_variables if pre_resolved_variables else {}
while unresolved_variables: while unresolved_variables:
unresolved_count: int = len(unresolved_variables) unresolved_count: int = len(unresolved_variables)
@ -63,7 +90,8 @@ class SyntaxTree(VariableDependency):
resolved_variables=resolved_variables resolved_variables=resolved_variables
): ):
resolved_variables[variable] = overrides[variable].resolve( resolved_variables[variable] = overrides[variable].resolve(
resolved_variables=resolved_variables resolved_variables=resolved_variables,
custom_functions=custom_functions,
) )
unresolved_variables.remove(variable) unresolved_variables.remove(variable)

View file

@ -5,6 +5,7 @@ from typing import Set
from ytdl_sub.script.types.resolvable import ArgumentType from ytdl_sub.script.types.resolvable import ArgumentType
from ytdl_sub.script.types.resolvable import Resolvable from ytdl_sub.script.types.resolvable import Resolvable
from ytdl_sub.script.types.variable import FunctionArgument
from ytdl_sub.script.types.variable import Variable from ytdl_sub.script.types.variable import Variable
from ytdl_sub.script.types.variable_dependency import VariableDependency from ytdl_sub.script.types.variable_dependency import VariableDependency
@ -20,12 +21,38 @@ class UnresolvedArray(Array, VariableDependency, ArgumentType):
@property @property
def variables(self) -> Set[Variable]: def variables(self) -> Set[Variable]:
return {value for value in self.value if isinstance(value, Variable)} variables: Set[Variable] = set()
for arg in self.value:
if isinstance(arg, Variable):
variables.add(arg)
elif isinstance(arg, VariableDependency):
variables.update(arg.variables)
def resolve(self, resolved_variables: Dict[Variable, Resolvable]) -> Resolvable: return variables
@property
def function_arguments(self) -> Set[FunctionArgument]:
variables: Set[FunctionArgument] = set()
for arg in self.value:
if isinstance(arg, FunctionArgument):
variables.add(arg)
elif isinstance(arg, VariableDependency):
variables.update(arg.function_arguments)
return variables
def resolve(
self,
resolved_variables: Dict[Variable, Resolvable],
custom_functions: Dict[str, "VariableDependency"],
) -> Resolvable:
return ResolvedArray( return ResolvedArray(
[ [
self._resolve_argument_type(resolved_variables=resolved_variables, arg=arg) self._resolve_argument_type(
arg=arg,
resolved_variables=resolved_variables,
custom_functions=custom_functions,
)
for arg in self.value for arg in self.value
] ]
) )

View file

@ -1,5 +1,7 @@
import copy
import functools import functools
import inspect import inspect
from abc import ABC
from dataclasses import dataclass from dataclasses import dataclass
from inspect import FullArgSpec from inspect import FullArgSpec
from typing import Callable from typing import Callable
@ -17,6 +19,7 @@ from ytdl_sub.script.types.resolvable import Resolvable
from ytdl_sub.script.types.resolvable import Resolvable_0 from ytdl_sub.script.types.resolvable import Resolvable_0
from ytdl_sub.script.types.resolvable import Resolvable_1 from ytdl_sub.script.types.resolvable import Resolvable_1
from ytdl_sub.script.types.resolvable import Resolvable_2 from ytdl_sub.script.types.resolvable import Resolvable_2
from ytdl_sub.script.types.variable import FunctionArgument
from ytdl_sub.script.types.variable import Variable from ytdl_sub.script.types.variable import Variable
from ytdl_sub.script.types.variable_dependency import VariableDependency from ytdl_sub.script.types.variable_dependency import VariableDependency
from ytdl_sub.utils.exceptions import StringFormattingException from ytdl_sub.utils.exceptions import StringFormattingException
@ -122,21 +125,87 @@ class FunctionInputSpec:
@dataclass(frozen=True) @dataclass(frozen=True)
class Function(VariableDependency, ArgumentType): class Function(VariableDependency, ArgumentType, ABC):
name: str name: str
args: List[ArgumentType] args: List[ArgumentType]
def __post_init__(self): @property
if not self.input_spec.is_compatible(input_args=self.args): def variables(self) -> Set[Variable]:
raise StringFormattingException( """
f"Invalid arguments passed to function {self.name}.\n" Returns
f"{self._expected_received_error_msg()}" -------
) All variables used within the function
"""
variables: Set[Variable] = set()
for arg in self.args:
if isinstance(arg, Variable):
variables.add(arg)
elif isinstance(arg, VariableDependency):
variables.update(arg.variables)
return variables
@property
def function_arguments(self) -> Set[FunctionArgument]:
"""
Returns
-------
All function arguments used within the function
"""
function_arguments: Set[FunctionArgument] = set()
for arg in self.args:
if isinstance(arg, FunctionArgument):
function_arguments.add(arg)
elif isinstance(arg, VariableDependency):
function_arguments.update(arg.function_arguments)
return function_arguments
@classmethod
def from_name_and_args(cls, name: str, args: List[ArgumentType]) -> "Function":
if hasattr(Functions, name) or hasattr(Functions, name + "_"):
return BuiltInFunction(name=name, args=args)
return CustomFunction(name=name, args=args)
class CustomFunction(Function):
def resolve(
self,
resolved_variables: Dict[Variable, Resolvable],
custom_functions: Dict[str, "VariableDependency"],
) -> Resolvable:
resolved_args: List[Resolvable] = [
self._resolve_argument_type(
arg=arg, resolved_variables=resolved_variables, custom_functions=custom_functions
)
for arg in self.args
]
if self.name in custom_functions:
if len(self.args) != len(custom_functions[self.name].function_arguments):
raise StringFormattingException("Custom function arg length does not equal")
resolved_variables_with_args = copy.deepcopy(resolved_variables)
for i, arg in enumerate(resolved_args):
function_arg = FunctionArgument(name=f"${i+1}") # Function args are 1-based
if function_arg in resolved_variables_with_args:
raise StringFormattingException("nested custom functions???")
resolved_variables_with_args[function_arg] = arg
return custom_functions[self.name].resolve(
resolved_variables=resolved_variables_with_args,
custom_functions=custom_functions,
)
else:
raise StringFormattingException(f"Custom function {self.name} does not exist")
class BuiltInFunction(Function):
def _expected_received_error_msg(self) -> str: def _expected_received_error_msg(self) -> str:
received_type_names: List[str] = [] received_type_names: List[str] = []
for arg in self.args: for arg in self.args:
if isinstance(arg, Function): if isinstance(arg, BuiltInFunction):
received_type_names.append(f"%{arg.name}(...)->{arg.output_type.__name__}") received_type_names.append(f"%{arg.name}(...)->{arg.output_type.__name__}")
else: else:
received_type_names.append(arg.__class__.__name__) received_type_names.append(arg.__class__.__name__)
@ -145,6 +214,13 @@ class Function(VariableDependency, ArgumentType):
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 __post_init__(self):
if not self.input_spec.is_compatible(input_args=self.args):
raise StringFormattingException(
f"Invalid arguments passed to function {self.name}.\n"
f"{self._expected_received_error_msg()}"
)
@property @property
def callable(self) -> Callable[..., Resolvable]: def callable(self) -> Callable[..., Resolvable]:
if hasattr(Functions, self.name): if hasattr(Functions, self.name):
@ -181,25 +257,15 @@ class Function(VariableDependency, ArgumentType):
return output_type return output_type
@property def resolve(
def variables(self) -> Set[Variable]: self,
""" resolved_variables: Dict[Variable, Resolvable],
Returns custom_functions: Dict[str, "VariableDependency"],
------- ) -> Resolvable:
All variables used within the function resolved_args: List[Resolvable] = [
""" self._resolve_argument_type(
variables: Set[Variable] = set() arg=arg, resolved_variables=resolved_variables, custom_functions=custom_functions
for arg in self.args: )
if isinstance(arg, Variable):
variables.add(arg)
elif isinstance(arg, Function):
variables.update(arg.variables)
return variables
def resolve(self, resolved_variables: Dict[Variable, Resolvable]) -> Resolvable:
resolved_args = [
self._resolve_argument_type(resolved_variables=resolved_variables, arg=arg)
for arg in self.args for arg in self.args
] ]

View file

@ -6,6 +6,7 @@ from typing import Set
from ytdl_sub.script.types.resolvable import ArgumentType from ytdl_sub.script.types.resolvable import ArgumentType
from ytdl_sub.script.types.resolvable import Hashable from ytdl_sub.script.types.resolvable import Hashable
from ytdl_sub.script.types.resolvable import Resolvable from ytdl_sub.script.types.resolvable import Resolvable
from ytdl_sub.script.types.variable import FunctionArgument
from ytdl_sub.script.types.variable import Variable from ytdl_sub.script.types.variable import Variable
from ytdl_sub.script.types.variable_dependency import VariableDependency from ytdl_sub.script.types.variable_dependency import VariableDependency
from ytdl_sub.utils.exceptions import StringFormattingException from ytdl_sub.utils.exceptions import StringFormattingException
@ -26,21 +27,47 @@ class UnresolvedMap(Map, VariableDependency, ArgumentType):
for key, value in self.value.items(): for key, value in self.value.items():
if isinstance(key, Variable): if isinstance(key, Variable):
output.add(key) output.add(key)
elif isinstance(key, VariableDependency):
output.update(key.variables)
if isinstance(value, Variable): if isinstance(value, Variable):
output.add(key) output.add(key)
elif isinstance(value, VariableDependency):
output.update(value.variables)
return output return output
def resolve(self, resolved_variables: Dict[Variable, Resolvable]) -> Resolvable: @property
def function_arguments(self) -> Set[FunctionArgument]:
output: Set[FunctionArgument] = set()
for key, value in self.value.items():
if isinstance(key, FunctionArgument):
output.add(key)
elif isinstance(key, VariableDependency):
output.update(key.function_arguments)
if isinstance(value, FunctionArgument):
output.add(key)
elif isinstance(value, VariableDependency):
output.update(value.function_arguments)
return output
def resolve(
self,
resolved_variables: Dict[Variable, Resolvable],
custom_functions: Dict[str, VariableDependency],
) -> Resolvable:
output: Dict[Hashable, Resolvable] = {} output: Dict[Hashable, Resolvable] = {}
for key, value in self.value.items(): for key, value in self.value.items():
resolved_key = self._resolve_argument_type( resolved_key = self._resolve_argument_type(
resolved_variables=resolved_variables, arg=key arg=key, resolved_variables=resolved_variables, custom_functions=custom_functions
) )
if not isinstance(resolved_key, Hashable): if not isinstance(resolved_key, Hashable):
raise StringFormattingException("key is not hashable") raise StringFormattingException("key is not hashable")
output[resolved_key] = self._resolve_argument_type( output[resolved_key] = self._resolve_argument_type(
resolved_variables=resolved_variables, arg=value arg=value, resolved_variables=resolved_variables, custom_functions=custom_functions
) )
return ResolvedMap(output) return ResolvedMap(output)

View file

@ -6,3 +6,8 @@ from ytdl_sub.script.types.resolvable import ArgumentType
@dataclass(frozen=True) @dataclass(frozen=True)
class Variable(ArgumentType): class Variable(ArgumentType):
name: str name: str
@dataclass(frozen=True)
class FunctionArgument(Variable):
pass

View file

@ -7,6 +7,7 @@ from typing import final
from ytdl_sub.script.types.resolvable import ArgumentType from ytdl_sub.script.types.resolvable import ArgumentType
from ytdl_sub.script.types.resolvable import Resolvable from ytdl_sub.script.types.resolvable import Resolvable
from ytdl_sub.script.types.variable import FunctionArgument
from ytdl_sub.script.types.variable import Variable from ytdl_sub.script.types.variable import Variable
from ytdl_sub.utils.exceptions import StringFormattingException from ytdl_sub.utils.exceptions import StringFormattingException
@ -18,12 +19,24 @@ class VariableDependency(ABC):
def variables(self) -> Set[Variable]: def variables(self) -> Set[Variable]:
raise NotImplemented() raise NotImplemented()
@property
@abstractmethod @abstractmethod
def resolve(self, resolved_variables: Dict[Variable, Resolvable]) -> Resolvable: def function_arguments(self) -> Set[FunctionArgument]:
raise NotImplemented()
@abstractmethod
def resolve(
self,
resolved_variables: Dict[Variable, Resolvable],
custom_functions: Dict[str, "VariableDependency"],
) -> Resolvable:
raise NotImplemented() raise NotImplemented()
def _resolve_argument_type( def _resolve_argument_type(
self, resolved_variables: Dict[Variable, Resolvable], arg: ArgumentType self,
arg: ArgumentType,
resolved_variables: Dict[Variable, Resolvable],
custom_functions: Dict[str, "VariableDependency"],
) -> Resolvable: ) -> Resolvable:
if isinstance(arg, Resolvable): if isinstance(arg, Resolvable):
return arg return arg
@ -32,7 +45,9 @@ class VariableDependency(ABC):
raise StringFormattingException("should never reach@") raise StringFormattingException("should never reach@")
return resolved_variables[arg] return resolved_variables[arg]
if isinstance(arg, VariableDependency): if isinstance(arg, VariableDependency):
return arg.resolve(resolved_variables) return arg.resolve(
resolved_variables=resolved_variables, custom_functions=custom_functions
)
assert False, "never reach here" assert False, "never reach here"

View file

@ -47,6 +47,10 @@ class TestParser:
parsed_extend.resolve({}) parsed_extend.resolve({})
assert False assert False
def test_function_argument(self):
parsed = parse("hello {%map([$1, $2])}")
assert False
def test_conditional(self): def test_conditional(self):
parsed = parse("hello {%if(True, 'hi', 3.4)}") parsed = parse("hello {%if(True, 'hi', 3.4)}")
assert parsed == SyntaxTree( assert parsed == SyntaxTree(

View file

@ -0,0 +1,37 @@
from typing import Dict
import pytest
from ytdl_sub.script.script import Script
from ytdl_sub.script.syntax_tree import SyntaxTree
from ytdl_sub.script.types.function import Function
from ytdl_sub.script.types.resolvable import String
from ytdl_sub.script.types.variable import Variable
from ytdl_sub.utils.exceptions import StringFormattingException
class TestSyntaxTree:
def test_simple(self):
script = Script(
{
"a": "a",
"b": "{b_}",
"b_": "b",
}
)
def test_custom_function(self):
script = Script(
{
"%custom_func": "return {[$1, $2]}",
"aa": "a",
"bb": "b",
"cc": "{%custom_func(aa, bb)}",
}
)
out = script.resolve()
assert False