Merge pull request #97 from logspace-ai/chatgpt
This commit is contained in:
commit
088e037c70
3 changed files with 30 additions and 19 deletions
|
|
@ -3,7 +3,6 @@
|
||||||
# - Defer prompts building to the last moment or when they have all the tools
|
# - Defer prompts building to the last moment or when they have all the tools
|
||||||
# - Build each inner agent first, then build the outer agent
|
# - Build each inner agent first, then build the outer agent
|
||||||
|
|
||||||
import logging
|
|
||||||
import types
|
import types
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
from typing import Any, Dict, List
|
from typing import Any, Dict, List
|
||||||
|
|
@ -12,8 +11,7 @@ from langflow.graph.constants import DIRECT_TYPES
|
||||||
from langflow.graph.utils import load_file
|
from langflow.graph.utils import load_file
|
||||||
from langflow.interface import loading
|
from langflow.interface import loading
|
||||||
from langflow.interface.listing import ALL_TYPES_DICT
|
from langflow.interface.listing import ALL_TYPES_DICT
|
||||||
|
from langflow.utils.logger import logger
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
class Node:
|
class Node:
|
||||||
|
|
|
||||||
|
|
@ -71,9 +71,10 @@ def load_flow_from_json(path: str):
|
||||||
data_graph = flow_graph["data"]
|
data_graph = flow_graph["data"]
|
||||||
nodes = data_graph["nodes"]
|
nodes = data_graph["nodes"]
|
||||||
# Substitute ZeroShotPrompt with PromptTemplate
|
# Substitute ZeroShotPrompt with PromptTemplate
|
||||||
nodes = replace_zero_shot_prompt_with_prompt_template(nodes)
|
# nodes = replace_zero_shot_prompt_with_prompt_template(nodes)
|
||||||
# Add input variables
|
# Add input variables
|
||||||
nodes = payload.extract_input_variables(nodes)
|
# nodes = payload.extract_input_variables(nodes)
|
||||||
|
|
||||||
# Nodes, edges and root node
|
# Nodes, edges and root node
|
||||||
edges = data_graph["edges"]
|
edges = data_graph["edges"]
|
||||||
graph = Graph(nodes, edges)
|
graph = Graph(nodes, edges)
|
||||||
|
|
|
||||||
|
|
@ -32,7 +32,7 @@ def build_langchain_object(data_graph):
|
||||||
logger.debug("Building langchain object")
|
logger.debug("Building langchain object")
|
||||||
nodes = data_graph["nodes"]
|
nodes = data_graph["nodes"]
|
||||||
# Add input variables
|
# Add input variables
|
||||||
nodes = payload.extract_input_variables(nodes)
|
# nodes = payload.extract_input_variables(nodes)
|
||||||
# Nodes, edges and root node
|
# Nodes, edges and root node
|
||||||
edges = data_graph["edges"]
|
edges = data_graph["edges"]
|
||||||
graph = Graph(nodes, edges)
|
graph = Graph(nodes, edges)
|
||||||
|
|
@ -75,26 +75,38 @@ def process_graph(data_graph: Dict[str, Any]):
|
||||||
|
|
||||||
def get_result_and_thought_using_graph(loaded_langchain, message: str):
|
def get_result_and_thought_using_graph(loaded_langchain, message: str):
|
||||||
"""Get result and thought from extracted json"""
|
"""Get result and thought from extracted json"""
|
||||||
loaded_langchain.verbose = True
|
|
||||||
try:
|
try:
|
||||||
|
loaded_langchain.verbose = True
|
||||||
with io.StringIO() as output_buffer, contextlib.redirect_stdout(output_buffer):
|
with io.StringIO() as output_buffer, contextlib.redirect_stdout(output_buffer):
|
||||||
chat_input = {}
|
chat_input = None
|
||||||
for key in loaded_langchain.input_keys:
|
for key in loaded_langchain.input_keys:
|
||||||
if key == "chat_history":
|
if key == "chat_history" and hasattr(loaded_langchain, "memory"):
|
||||||
if hasattr(loaded_langchain, "memory"):
|
|
||||||
loaded_langchain.memory.memory_key = "chat_history"
|
loaded_langchain.memory.memory_key = "chat_history"
|
||||||
else:
|
else:
|
||||||
chat_input[key] = message
|
chat_input = {key: message}
|
||||||
|
|
||||||
if hasattr(loaded_langchain, "run"):
|
if hasattr(loaded_langchain, "return_intermediate_steps"):
|
||||||
loaded_langchain = loaded_langchain.run
|
# https://github.com/hwchase17/langchain/issues/2068
|
||||||
result = loaded_langchain(**chat_input)
|
loaded_langchain.return_intermediate_steps = False
|
||||||
|
|
||||||
|
try:
|
||||||
|
output = loaded_langchain(chat_input)
|
||||||
|
except ValueError as exc:
|
||||||
|
logger.debug("Error: %s", str(exc))
|
||||||
|
output = loaded_langchain.run(chat_input)
|
||||||
|
|
||||||
|
intermediate_steps = (
|
||||||
|
output.get("intermediate_steps", []) if isinstance(output, dict) else []
|
||||||
|
)
|
||||||
|
|
||||||
result = (
|
result = (
|
||||||
result.get(loaded_langchain.output_keys[0])
|
output.get(loaded_langchain.output_keys[0])
|
||||||
if isinstance(result, dict)
|
if isinstance(output, dict)
|
||||||
else result
|
else output
|
||||||
)
|
)
|
||||||
|
if intermediate_steps:
|
||||||
|
thought = format_intermediate_steps(intermediate_steps)
|
||||||
|
else:
|
||||||
thought = output_buffer.getvalue()
|
thought = output_buffer.getvalue()
|
||||||
|
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue