-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy pathmessage_parser.py
More file actions
64 lines (55 loc) · 2.2 KB
/
Copy pathmessage_parser.py
File metadata and controls
64 lines (55 loc) · 2.2 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
from typing import Iterable, Iterator, Union
from langchain_core.messages import (
AIMessage,
BaseMessage,
HumanMessage,
ToolCall,
ToolMessage,
)
from ai_datastream.messages import ChatMessage, MessageRole, ToolInvocation
class LanggraphMessageParser:
def _parse_user_message(self, message: ChatMessage) -> Union[HumanMessage, None]:
if not message.content:
return None
return HumanMessage(content=message.content)
def _parse_tool_call(self, tool_invocation: ToolInvocation) -> ToolCall:
return ToolCall(
id=tool_invocation.tool_call_id,
name=tool_invocation.tool_name,
args=tool_invocation.args,
)
def _parse_ai_message(self, message: ChatMessage) -> Union[AIMessage, None]:
tool_calls = []
if message.tool_invocations:
tool_calls = [
self._parse_tool_call(tool_invocation)
for tool_invocation in message.tool_invocations
]
if not tool_calls and not message.content:
return None
return AIMessage(
content=message.content or "",
tool_calls=tool_calls,
)
def _parse_tool_message(self, tool_invocation: ToolInvocation) -> ToolMessage:
return ToolMessage(
content=tool_invocation.result,
tool_call_id=tool_invocation.tool_call_id,
)
def parse(self, message: ChatMessage) -> Iterator[BaseMessage]:
if message.role == MessageRole.USER:
user_message = self._parse_user_message(message)
if user_message:
yield user_message
elif message.role == MessageRole.ASSISTANT:
ai_message = self._parse_ai_message(message)
if ai_message:
yield ai_message
if message.tool_invocations:
for tool_invocation in message.tool_invocations:
yield self._parse_tool_message(tool_invocation)
else:
raise ValueError(f"Invalid message role: {message.role}")
def parse_many(self, messages: Iterable[ChatMessage]) -> Iterator[BaseMessage]:
for message in messages:
yield from self.parse(message)