Team Ai
Apppublic

Underground-Digital/Workflow-Engine

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
variable_pool.py172 linesDownload Raw Back to entities
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