Underground-Digital/Workflow-Engine
0
1import re2from collections import defaultdict3from collections.abc import Mapping, Sequence4from typing import Any, Union5 6from pydantic import BaseModel, Field7 8from core.file import File, FileAttribute, file_manager9from core.variables import Segment, SegmentGroup, Variable10from core.variables.segments import FileSegment11from factories import variable_factory12 13from ..constants import CONVERSATION_VARIABLE_NODE_ID, ENVIRONMENT_VARIABLE_NODE_ID, SYSTEM_VARIABLE_NODE_ID14from ..enums import SystemVariableKey15 16VariableValue = Union[str, int, float, dict, list, File]17 18 19VARIABLE_PATTERN = re.compile(r"\{\{#([a-zA-Z0-9_]{1,50}(?:\.[a-zA-Z_][a-zA-Z0-9_]{0,29}){1,10})#\}\}")20 21 22class VariablePool(BaseModel):23 # Variable dictionary is a dictionary for looking up variables by their selector.24 # The first element of the selector is the node id, it's the first-level key in the dictionary.25 # Other elements of the selector are the keys in the second-level dictionary. To get the key, we hash the26 # elements of the selector except the first one.27 variable_dictionary: dict[str, dict[int, Segment]] = Field(28 description="Variables mapping",29 default=defaultdict(dict),30 )31 # TODO: This user inputs is not used for pool.32 user_inputs: Mapping[str, Any] = Field(33 description="User inputs",34 )35 system_variables: Mapping[SystemVariableKey, Any] = Field(36 description="System variables",37 )38 environment_variables: Sequence[Variable] = Field(39 description="Environment variables.",40 default_factory=list,41 )42 conversation_variables: Sequence[Variable] = Field(43 description="Conversation variables.",44 default_factory=list,45 )46 47 def __init__(48 self,49 *,50 system_variables: Mapping[SystemVariableKey, Any] | None = None,51 user_inputs: Mapping[str, Any] | None = None,52 environment_variables: Sequence[Variable] | None = None,53 conversation_variables: Sequence[Variable] | None = None,54 **kwargs,55 ):56 environment_variables = environment_variables or []57 conversation_variables = conversation_variables or []58 user_inputs = user_inputs or {}59 system_variables = system_variables or {}60 61 super().__init__(62 system_variables=system_variables,63 user_inputs=user_inputs,64 environment_variables=environment_variables,65 conversation_variables=conversation_variables,66 **kwargs,67 )68 69 for key, value in self.system_variables.items():70 self.add((SYSTEM_VARIABLE_NODE_ID, key.value), value)71 # Add environment variables to the variable pool72 for var in self.environment_variables:73 self.add((ENVIRONMENT_VARIABLE_NODE_ID, var.name), var)74 # Add conversation variables to the variable pool75 for var in self.conversation_variables:76 self.add((CONVERSATION_VARIABLE_NODE_ID, var.name), var)77 78 def add(self, selector: Sequence[str], value: Any, /) -> None:79 """80 Adds a variable to the variable pool.81 82 NOTE: You should not add a non-Segment value to the variable pool83 even if it is allowed now.84 85 Args:86 selector (Sequence[str]): The selector for the variable.87 value (VariableValue): The value of the variable.88 89 Raises:90 ValueError: If the selector is invalid.91 92 Returns:93 None94 """95 if len(selector) < 2:96 raise ValueError("Invalid selector")97 98 if isinstance(value, Segment):99 v = value100 else:101 v = variable_factory.build_segment(value)102 103 hash_key = hash(tuple(selector[1:]))104 self.variable_dictionary[selector[0]][hash_key] = v105 106 def get(self, selector: Sequence[str], /) -> Segment | None:107 """108 Retrieves the value from the variable pool based on the given selector.109 110 Args:111 selector (Sequence[str]): The selector used to identify the variable.112 113 Returns:114 Any: The value associated with the given selector.115 116 Raises:117 ValueError: If the selector is invalid.118 """119 if len(selector) < 2:120 return None121 122 hash_key = hash(tuple(selector[1:]))123 value = self.variable_dictionary[selector[0]].get(hash_key)124 125 if value is None:126 selector, attr = selector[:-1], selector[-1]127 # Python support `attr in FileAttribute` after 3.12128 if attr not in {item.value for item in FileAttribute}:129 return None130 value = self.get(selector)131 if not isinstance(value, FileSegment):132 return None133 attr = FileAttribute(attr)134 attr_value = file_manager.get_attr(file=value.value, attr=attr)135 return variable_factory.build_segment(attr_value)136 137 return value138 139 def remove(self, selector: Sequence[str], /):140 """141 Remove variables from the variable pool based on the given selector.142 143 Args:144 selector (Sequence[str]): A sequence of strings representing the selector.145 146 Returns:147 None148 """149 if not selector:150 return151 if len(selector) == 1:152 self.variable_dictionary[selector[0]] = {}153 return154 hash_key = hash(tuple(selector[1:]))155 self.variable_dictionary[selector[0]].pop(hash_key, None)156 157 def convert_template(self, template: str, /):158 parts = VARIABLE_PATTERN.split(template)159 segments = []160 for part in filter(lambda x: x, parts):161 if "." in part and (variable := self.get(part.split("."))):162 segments.append(variable)163 else:164 segments.append(variable_factory.build_segment(part))165 return SegmentGroup(value=segments)166 167 def get_file(self, selector: Sequence[str], /) -> FileSegment | None:168 segment = self.get(selector)169 if isinstance(segment, FileSegment):170 return segment171 return None172 