map working

This commit is contained in:
Jesse Bannon 2023-11-07 22:30:51 -08:00
parent 20ad7b4f7e
commit ea0d836baf
6 changed files with 113 additions and 18 deletions

View file

@ -1,9 +1,12 @@
from typing import Dict
from typing import List
from typing import Optional
from ytdl_sub.script.types.array import Array
from ytdl_sub.script.types.resolvable import Map
from ytdl_sub.script.types.map import Map
from ytdl_sub.script.types.resolvable import Hashable
from ytdl_sub.script.types.resolvable import Resolvable
from ytdl_sub.script.types.resolvable import String
from ytdl_sub.utils.exceptions import StringFormattingException
@ -21,3 +24,11 @@ class MapFunctions:
output[key_value.value[0]] = key_value.value[1]
return Map(output)
@staticmethod
def get(mapping: Map, key: Hashable, default: Optional[Resolvable] = None) -> Resolvable:
if key not in mapping.value:
if default is not None:
return default
raise StringFormattingException("key not found")
return mapping.value[key]

View file

@ -1,3 +1,4 @@
from typing import Dict
from typing import List
from typing import Optional
@ -6,6 +7,7 @@ from ytdl_sub.script.types.array import Array
from ytdl_sub.script.types.array import UnresolvedArray
from ytdl_sub.script.types.function import ArgumentType
from ytdl_sub.script.types.function import Function
from ytdl_sub.script.types.map import UnresolvedMap
from ytdl_sub.script.types.resolvable import Boolean
from ytdl_sub.script.types.resolvable import Float
from ytdl_sub.script.types.resolvable import Integer
@ -120,6 +122,9 @@ class _Parser:
if self._read(increment_pos=False) == "[":
self._pos += 1
return self._parse_array()
if self._read(increment_pos=False) == "{":
self._pos += 1
assert self._parse_map()
if self._read(increment_pos=False).isascii() and self._read(increment_pos=False).islower():
return self._parse_variable()
raise StringFormattingException(
@ -127,7 +132,7 @@ class _Parser:
"string, boolean, or variable without brackets"
)
def _parse_args(self, breaking_char: str = ")") -> List[ArgumentType]:
def _parse_args(self, breaking_chars: str = ")") -> List[ArgumentType]:
"""
Begin parsing function args after the first ``(``, i.e. ``function_name(``
"""
@ -136,7 +141,7 @@ class _Parser:
arguments: List[ArgumentType] = []
while ch := self._read(increment_pos=False):
if ch == breaking_char:
if ch in breaking_chars:
break
if ch.isspace():
@ -182,10 +187,40 @@ class _Parser:
self._pos += 1
return UnresolvedArray(value=function_args)
else:
function_args = self._parse_args(breaking_char="]")
function_args = self._parse_args(breaking_chars="]")
raise StringFormattingException("Invalid function")
def _parse_map(self) -> UnresolvedMap:
"""
Begin parsing a map after reading the first ``{``
"""
output: Dict[ArgumentType, ArgumentType] = {}
key: Optional[ArgumentType] = None
while ch := self._read(increment_pos=False):
if ch == "}":
if key is not None:
raise StringFormattingException("Key with no value")
self._pos += 1
return UnresolvedMap(value=output)
elif key is None:
key_args = self._parse_args(breaking_chars=":")
if len(key_args) != 1:
raise StringFormattingException("Lazy parsing but got mlutiple args")
key = key_args[0]
elif key is not None and ch == ":":
self._pos += 1
value_args = self._parse_args(breaking_chars=",}")
if len(value_args) != 1:
raise StringFormattingException("Lazy parsing, no value")
output[key] = value_args[0]
key = None
else:
raise StringFormattingException("Invalid map")
def _parse(self) -> SyntaxTree:
bracket_counter = 0
literal_str = ""
@ -216,6 +251,9 @@ class _Parser:
elif ch1 == "[":
self._pos += 1
self._ast.append(self._parse_array())
elif ch1 == "{":
self._pos += 1
self._ast.append(self._parse_map())
else:
self._ast.append(self._parse_variable())
else:

View file

@ -0,0 +1,51 @@
from dataclasses import dataclass
from typing import Dict
from typing import List
from typing import Set
from ytdl_sub.script.types.resolvable import ArgumentType
from ytdl_sub.script.types.resolvable import Hashable
from ytdl_sub.script.types.resolvable import Resolvable
from ytdl_sub.script.types.variable import Variable
from ytdl_sub.script.types.variable_dependency import VariableDependency
from ytdl_sub.utils.exceptions import StringFormattingException
@dataclass(frozen=True)
class Map:
value: Dict[Hashable, Resolvable]
@dataclass(frozen=True)
class UnresolvedMap(Map, VariableDependency, ArgumentType):
value: Dict[ArgumentType, ArgumentType]
@property
def variables(self) -> Set[Variable]:
output: Set[Variable] = set()
for key, value in self.value.items():
if isinstance(key, Variable):
output.add(key)
if isinstance(value, Variable):
output.add(key)
return output
def resolve(self, resolved_variables: Dict[Variable, Resolvable]) -> Resolvable:
output: Dict[Hashable, Resolvable] = {}
for key, value in self.value.items():
resolved_key = self._resolve_argument_type(
resolved_variables=resolved_variables, arg=key
)
if not isinstance(resolved_key, Hashable):
raise StringFormattingException("key is not hashable")
output[resolved_key] = self._resolve_argument_type(
resolved_variables=resolved_variables, arg=value
)
return ResolvedMap(output)
@dataclass(frozen=True)
class ResolvedMap(Map, Resolvable):
pass

View file

@ -1,7 +1,6 @@
from abc import ABC
from dataclasses import dataclass
from typing import Any
from typing import Dict
from typing import Generic
from typing import TypeVar
@ -33,13 +32,17 @@ class Resolvable(Resolvable_0, Resolvable_1, Resolvable_2, ABC):
return str(self.value)
class Hashable(Resolvable, ABC):
pass
@dataclass(frozen=True)
class ResolvableT(Resolvable, ABC, Generic[T]):
class ResolvableT(Hashable, ABC, Generic[T]):
value: T
@dataclass(frozen=True)
class Numeric(ResolvableT[NumericT], ABC, Generic[NumericT]):
class Numeric(ResolvableT[NumericT], Hashable, ABC, Generic[NumericT]):
pass
@ -54,18 +57,10 @@ class Float(Numeric[float], ArgumentType):
@dataclass(frozen=True)
class Boolean(ResolvableT[bool], ArgumentType):
class Boolean(ResolvableT[bool], Hashable, ArgumentType):
pass
@dataclass(frozen=True)
class String(ResolvableT[str], ArgumentType):
class String(ResolvableT[str], Hashable, ArgumentType):
pass
@dataclass(frozen=True)
class Map(Resolvable, ArgumentType):
value: Dict[Resolvable, Resolvable]
def __str__(self) -> str:
return f"[{', '.join([val.value for val in self.value])}]"

View file

@ -43,6 +43,7 @@ class TestParser:
parsed_empty = parse("hello {%map()}")
parsed_with_var = parse("hello {%map([variable_name, 'elem2'])}")
parsed_extend = parse("hi {%map([variable_name, 'elem2'], ['elem3', variable_name])}")
parse_raw_map = parse("hello {{'key': 'value'}}")
parsed_extend.resolve({})
assert False

View file

@ -53,6 +53,5 @@ class TestSyntaxTree:
"b": SyntaxTree(ast=[Function(name="capitalize", args=[Variable("b_")])]),
"b_": SyntaxTree(ast=[Variable("b")]),
}
_ = SyntaxTree.resolve_overrides(parsed_overrides=overrides)
with pytest.raises(StringFormattingException):
_ = SyntaxTree.resolve_overrides(parsed_overrides=overrides)