FIX: don't error when adding to canvas (#4055)

This commit is contained in:
Eric Hare 2024-10-07 13:25:17 -07:00 • committed by GitHub
commit 3ea7be12e9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194

View file

@ -21,13 +21,14 @@ class LangChainHubPromptComponent(Component):
name="langchain_api_key", name="langchain_api_key",
display_name="Your LangChain API Key", display_name="Your LangChain API Key",
info="The LangChain API Key to use.", info="The LangChain API Key to use.",
required=True,
), ),
StrInput( StrInput(
name="langchain_hub_prompt", name="langchain_hub_prompt",
display_name="LangChain Hub Prompt", display_name="LangChain Hub Prompt",
info="The LangChain Hub prompt to use.", info="The LangChain Hub prompt to use, i.e., 'efriis/my-first-prompt'",
value="efriis/my-first-prompt",
refresh_button=True, refresh_button=True,
required=True,
), ),
] ]
@ -36,63 +37,70 @@ class LangChainHubPromptComponent(Component):
] ]
def update_build_config(self, build_config: dict, field_value: str, field_name: str | None = None): def update_build_config(self, build_config: dict, field_value: str, field_name: str | None = None):
if field_name == "langchain_hub_prompt": # If the field is not langchain_hub_prompt or the value is empty, return the build config as is
template = self._fetch_langchain_hub_template() if field_name != "langchain_hub_prompt" or not field_value:
return build_config
# Get the template's messages # Fetch the template
if hasattr(template, "messages"): template = self._fetch_langchain_hub_template()
template_messages = template.messages
else:
template_messages = [HumanMessagePromptTemplate(prompt=template)]
# Extract the messages from the prompt data # Get the template's messages
prompt_template = [] if hasattr(template, "messages"):
for message_data in template_messages: template_messages = template.messages
prompt_template.append(message_data.prompt) else:
template_messages = [HumanMessagePromptTemplate(prompt=template)]
# Regular expression to find all instances of {<string>} # Extract the messages from the prompt data
pattern = r"\{(.*?)\}" prompt_template = []
for message_data in template_messages:
prompt_template.append(message_data.prompt)
# Get all the custom fields # Regular expression to find all instances of {<string>}
custom_fields: list[str] = [] pattern = r"\{(.*?)\}"
full_template = ""
for message in prompt_template:
# Find all matches
matches = re.findall(pattern, message.template)
custom_fields = custom_fields + matches
# Create a string version of the full template # Get all the custom fields
full_template = full_template + "\n" + message.template custom_fields: list[str] = []
full_template = ""
for message in prompt_template:
# Find all matches
matches = re.findall(pattern, message.template)
custom_fields = custom_fields + matches
# No need to reprocess if we have them already # Create a string version of the full template
if all("param_" + custom_field in build_config for custom_field in custom_fields): full_template = full_template + "\n" + message.template
return build_config
# Easter egg: Show template in info popup # No need to reprocess if we have them already
build_config["langchain_hub_prompt"]["info"] = full_template if all("param_" + custom_field in build_config for custom_field in custom_fields):
return build_config
# Remove old parameter inputs if any # Easter egg: Show template in info popup
for key, _ in build_config.copy().items(): build_config["langchain_hub_prompt"]["info"] = full_template
if key.startswith("param_"):
del build_config[key]
# Now create inputs for each # Remove old parameter inputs if any
for custom_field in custom_fields: for key, _ in build_config.copy().items():
new_parameter = DefaultPromptField( if key.startswith("param_"):
name=f"param_{custom_field}", del build_config[key]
display_name=custom_field,
info="Fill in the value for {" + custom_field + "}",
).to_dict()
build_config[f"param_{custom_field}"] = new_parameter # Now create inputs for each
for custom_field in custom_fields:
new_parameter = DefaultPromptField(
name=f"param_{custom_field}",
display_name=custom_field,
info="Fill in the value for {" + custom_field + "}",
).to_dict()
# Add the new parameter to the build config
build_config[f"param_{custom_field}"] = new_parameter
return build_config return build_config
async def build_prompt( async def build_prompt(
self, self,
) -> Message: ) -> Message:
# Get the parameters that # Fetch the template
template = self._fetch_langchain_hub_template() # TODO: doing this twice template = self._fetch_langchain_hub_template()
# Get the parameters from the attributes
original_params = {k[6:] if k.startswith("param_") else k: v for k, v in self._attributes.items()} original_params = {k[6:] if k.startswith("param_") else k: v for k, v in self._attributes.items()}
prompt_value = template.invoke(original_params) prompt_value = template.invoke(original_params)
@ -111,6 +119,7 @@ class LangChainHubPromptComponent(Component):
# Check if the api key is provided # Check if the api key is provided
if not self.langchain_api_key: if not self.langchain_api_key:
msg = "Please provide a LangChain API Key" msg = "Please provide a LangChain API Key"
raise ValueError(msg) raise ValueError(msg)
# Pull the prompt from LangChain Hub # Pull the prompt from LangChain Hub