fix: function to try to avoid input keys erros with memory

This commit is contained in:
Gabriel Almeida 2023-04-04 21:46:05 -03:00
commit fe95790331

View file

@ -72,26 +72,41 @@ def process_graph(data_graph: Dict[str, Any]):
return {"result": str(result), "thought": thought.strip()} return {"result": str(result), "thought": thought.strip()}
def fix_memory_inputs_for_intermediate_steps(langchain_object): def fix_memory_inputs(langchain_object):
""" """
Fix memory inputs by replacing the memory key with the input key. Fix memory inputs by replacing the memory key with the input key.
""" """
langchain_object.return_intermediate_steps = True # Possible memory keys
langchain_object.memory.memory_key # "chat_history", "history"
input_key = [ # if memory_key is "chat_history" and input_keys has "history"
key # we need to replace "chat_history" with "history"
for key in langchain_object.input_keys mem_key_dict = {
if key != langchain_object.memory.memory_key "chat_history": "history",
][0] "history": "chat_history",
# get output_key }
output_key = [ memory_key = langchain_object.memory.memory_key
key possible_new_mem_key = mem_key_dict.get(memory_key)
for key in langchain_object.output_keys if possible_new_mem_key is not None:
if key != langchain_object.memory.memory_key # get input_key
][0] input_key = [
# set input_key and output_key in memory key
langchain_object.memory.input_key = input_key for key in langchain_object.input_keys
langchain_object.memory.output_key = output_key if key not in [memory_key, possible_new_mem_key]
][0]
# get output_key
output_key = [
key
for key in langchain_object.output_keys
if key not in [memory_key, possible_new_mem_key]
][0]
# set input_key and output_key in memory
langchain_object.memory.input_key = input_key
langchain_object.memory.output_key = output_key
for input_key in langchain_object.input_keys:
if input_key == possible_new_mem_key:
langchain_object.memory.memory_key = possible_new_mem_key
def get_result_and_thought_using_graph(langchain_object, message: str): def get_result_and_thought_using_graph(langchain_object, message: str):
@ -117,17 +132,15 @@ def get_result_and_thought_using_graph(langchain_object, message: str):
# Deactivating until we have a frontend solution # Deactivating until we have a frontend solution
# to display intermediate steps # to display intermediate steps
langchain_object.return_intermediate_steps = False langchain_object.return_intermediate_steps = False
if langchain_object.return_intermediate_steps:
fix_memory_inputs_for_intermediate_steps(langchain_object) fix_memory_inputs(langchain_object)
try: try:
output = langchain_object(chat_input) output = langchain_object(chat_input)
except ValueError as exc: except ValueError as exc:
# make the error message more informative # make the error message more informative
logger.debug(f"Error: {str(exc)}") logger.debug(f"Error: {str(exc)}")
if hasattr(langchain_object, "memory"): output = langchain_object.run(chat_input)
langchain_object.memory.memory_key = memory_key
output = langchain_object(chat_input)
intermediate_steps = ( intermediate_steps = (
output.get("intermediate_steps", []) if isinstance(output, dict) else [] output.get("intermediate_steps", []) if isinstance(output, dict) else []