Repository navigation
Expand file tree
/
Copy pathgraph.py
More file actions
106 lines (86 loc) · 3.29 KB
/
Copy pathgraph.py
File metadata and controls
106 lines (86 loc) · 3.29 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
from typing import Annotated, Optional, TypedDict
from uuid import uuid4
from langchain_core.messages import (
AIMessageChunk,
AnyMessage,
HumanMessage,
HumanMessageChunk,
)
from langgraph.checkpoint.sqlite import SqliteSaver
from langgraph.graph import START, END, StateGraph, add_messages
from agents.supervisor_agent import supervisor_agent
from agents.code_agent import code_agent
from agents.research_agent import research_agent
from agents.image_agent import image_agent
from utils import pretty_print_messages
import pprint
import sqlite3
import dotenv
dotenv.load_dotenv()
sqlite_connection = sqlite3.connect("checkpoint.sqlite", check_same_thread=False, timeout=30.0)
sqlite_connection.execute("PRAGMA journal_mode=WAL;")
memory = SqliteSaver(sqlite_connection)
class GeneratedImage(TypedDict, total=False):
data: str
mime_type: str
prompt: str
file_path: Optional[str]
url: Optional[str]
filename: Optional[str]
class AgentState(TypedDict):
messages: Annotated[list[AnyMessage], add_messages]
generated_image: Optional[GeneratedImage]
_code_subagent = code_agent()
_research_subagent = research_agent()
def node_code_agent(state: AgentState):
last_msg = state["messages"][-1]
task_content = getattr(last_msg, "content", str(last_msg))
res = _code_subagent.invoke({"messages": [HumanMessage(content=task_content)]})
final_output = res["messages"][-1].content
return {"messages": [HumanMessage(content=f"Code agent output:\n{final_output}")]}
def node_research_agent(state: AgentState):
last_msg = state["messages"][-1]
task_content = getattr(last_msg, "content", str(last_msg))
res = _research_subagent.invoke({"messages": [HumanMessage(content=task_content)]})
final_output = res["messages"][-1].content
return {"messages": [HumanMessage(content=f"Research agent output:\n{final_output}")]}
supervisor = (
StateGraph(AgentState)
.add_node(
supervisor_agent(),
destinations=("research_agent", "code_agent", "image_agent", END),
)
.add_node("research_agent", node_research_agent)
.add_node("code_agent", node_code_agent)
.add_node("image_agent", image_agent)
.add_edge(START, "supervisor")
.add_edge("research_agent", "supervisor")
.add_edge("code_agent", "supervisor")
.add_edge("image_agent", "supervisor")
.compile(checkpointer=memory)
)
def serialise_ai_message_chunk(chunk):
if isinstance(chunk, AIMessageChunk):
return chunk.content
if isinstance(chunk, HumanMessageChunk):
return chunk.content
else:
raise TypeError(
f"Object of type {type(chunk).__name__} is not correctly formatted for serialisation"
)
if __name__ == "__main__":
while True:
input_query = input("Enter prompt: ")
config = {"configurable": {"thread_id": str(uuid4())}}
for chunk in supervisor.stream(
{"messages": [HumanMessage(content=input_query)], "generated_image": None},
config=config,
stream_mode="updates",
):
node_name = list(chunk)[0]
# pprint.pprint(chunk, indent=4)
for msg in chunk[node_name]["messages"]:
print(25 * "+")
print("node name: ", node_name)
print(msg)
print(25 * "-")