Fix_vecstore_memory (#582)

This commit is contained in:
Gabriel Luiz Freitas Almeida 2023-07-01 09:59:19 -03:00 • committed by GitHub
commit a0920e7abd
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
6 changed files with 22 additions and 756 deletions

View file

@ -100,6 +100,8 @@ def instantiate_llm(node_type, class_object, params: Dict):
def instantiate_memory(node_type, class_object, params): def instantiate_memory(node_type, class_object, params):
try: try:
if "retriever" in params and hasattr(params["retriever"], "as_retriever"):
params["retriever"] = params["retriever"].as_retriever()
return class_object(**params) return class_object(**params)
# I want to catch a specific attribute error that happens # I want to catch a specific attribute error that happens
# when the object does not have a cursor attribute # when the object does not have a cursor attribute

View file

@ -62,6 +62,10 @@ def update_memory_keys(langchain_object, possible_new_mem_key):
if key not in [langchain_object.memory.memory_key, possible_new_mem_key] if key not in [langchain_object.memory.memory_key, possible_new_mem_key]
][0] ][0]
langchain_object.memory.input_key = input_key keys = [input_key, output_key, possible_new_mem_key]
langchain_object.memory.output_key = output_key attrs = ["input_key", "output_key", "memory_key"]
langchain_object.memory.memory_key = possible_new_mem_key for key, attr in zip(keys, attrs):
try:
setattr(langchain_object.memory, attr, key)
except ValueError as exc:
logger.debug(f"{langchain_object.memory} has no attribute {attr} ({exc})")

View file

@ -51,5 +51,6 @@ async def get_result_and_steps(langchain_object, message: str, **kwargs):
) )
thought = format_actions(intermediate_steps) if intermediate_steps else "" thought = format_actions(intermediate_steps) if intermediate_steps else ""
except Exception as exc: except Exception as exc:
logger.exception(exc)
raise ValueError(f"Error: {str(exc)}") from exc raise ValueError(f"Error: {str(exc)}") from exc
return result, thought return result, thought

View file

@ -37,6 +37,7 @@ class MemoryFrontendNode(FrontendNode):
value="", value="",
) )
) )
if self.template.type_name not in {"VectorStoreRetrieverMemory"}:
self.template.add_field( self.template.add_field(
TemplateField( TemplateField(
field_type="str", field_type="str",

View file

@ -51,7 +51,7 @@ class VectorStoreFrontendNode(FrontendNode):
required=False, required=False,
show=True, show=True,
advanced=False, advanced=False,
value=True, value=False,
display_name="Persist", display_name="Persist",
) )
extra_fields.append(extra_field) extra_fields.append(extra_field)

File diff suppressed because one or more lines are too long