Add tool messages to arcade chat conversation history (#49)
* Adds tool calls and tool responses to the conversation history for streaming and non-streaming. * Adds debug flag for arcade chat which displays the tool calls and tool responses that occurred during generate `arcade chat -d`
This commit is contained in:
parent
399ad5f878
commit
18cc192290
2 changed files with 49 additions and 9 deletions
|
|
@ -21,6 +21,8 @@ from arcade.cli.utils import (
|
||||||
create_cli_catalog,
|
create_cli_catalog,
|
||||||
display_eval_results,
|
display_eval_results,
|
||||||
display_streamed_markdown,
|
display_streamed_markdown,
|
||||||
|
display_tool_messages,
|
||||||
|
get_tool_messages,
|
||||||
markdownify_urls,
|
markdownify_urls,
|
||||||
validate_and_get_config,
|
validate_and_get_config,
|
||||||
)
|
)
|
||||||
|
|
@ -165,6 +167,7 @@ def chat(
|
||||||
"--no-tls",
|
"--no-tls",
|
||||||
help="Whether to disable TLS for the connection to the Arcade Engine.",
|
help="Whether to disable TLS for the connection to the Arcade Engine.",
|
||||||
),
|
),
|
||||||
|
debug: bool = typer.Option(False, "--debug", "-d", help="Show debug information"),
|
||||||
) -> None:
|
) -> None:
|
||||||
"""
|
"""
|
||||||
Chat with a language model.
|
Chat with a language model.
|
||||||
|
|
@ -185,7 +188,7 @@ def chat(
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# start messages conversation
|
# start messages conversation
|
||||||
messages: list[dict[str, Any]] = []
|
history: list[dict[str, Any]] = []
|
||||||
|
|
||||||
chat_header = Text.assemble(
|
chat_header = Text.assemble(
|
||||||
"\n",
|
"\n",
|
||||||
|
|
@ -217,30 +220,39 @@ def chat(
|
||||||
# Add the input to history
|
# Add the input to history
|
||||||
readline.add_history(user_input)
|
readline.add_history(user_input)
|
||||||
|
|
||||||
messages.append({"role": "user", "content": user_input})
|
history.append({"role": "user", "content": user_input})
|
||||||
|
|
||||||
|
tool_messages: list[dict] = []
|
||||||
|
|
||||||
if stream:
|
if stream:
|
||||||
# TODO Fix this in the client so users don't deal with these
|
# TODO Fix this in the client so users don't deal with these
|
||||||
# typing issues
|
# typing issues
|
||||||
stream_response = client.chat.completions.create( # type: ignore[call-overload]
|
stream_response = client.chat.completions.create( # type: ignore[call-overload]
|
||||||
model=model,
|
model=model,
|
||||||
messages=messages,
|
messages=history,
|
||||||
tool_choice="generate",
|
tool_choice="generate",
|
||||||
user=user_email,
|
user=user_email,
|
||||||
stream=True,
|
stream=True,
|
||||||
)
|
)
|
||||||
role, message_content = display_streamed_markdown(stream_response, model)
|
role, message_content, tool_messages = display_streamed_markdown(
|
||||||
|
stream_response, model
|
||||||
|
)
|
||||||
|
|
||||||
|
history += tool_messages
|
||||||
else:
|
else:
|
||||||
response = client.chat.completions.create( # type: ignore[call-overload]
|
response = client.chat.completions.create( # type: ignore[call-overload]
|
||||||
model=model,
|
model=model,
|
||||||
messages=messages,
|
messages=history,
|
||||||
tool_choice="generate",
|
tool_choice="generate",
|
||||||
user=user_email,
|
user=user_email,
|
||||||
stream=False,
|
stream=False,
|
||||||
)
|
)
|
||||||
message_content = response.choices[0].message.content or ""
|
message_content = response.choices[0].message.content or ""
|
||||||
role = response.choices[0].message.role
|
|
||||||
|
|
||||||
|
tool_messages = get_tool_messages(response.choices[0])
|
||||||
|
history += tool_messages
|
||||||
|
|
||||||
|
role = response.choices[0].message.role
|
||||||
if role == "assistant":
|
if role == "assistant":
|
||||||
message_content = markdownify_urls(message_content)
|
message_content = markdownify_urls(message_content)
|
||||||
console.print(
|
console.print(
|
||||||
|
|
@ -249,7 +261,10 @@ def chat(
|
||||||
else:
|
else:
|
||||||
console.print(f"\n[bold magenta]{role}:[/bold magenta] {message_content}")
|
console.print(f"\n[bold magenta]{role}:[/bold magenta] {message_content}")
|
||||||
|
|
||||||
messages.append({"role": role, "content": message_content})
|
if debug:
|
||||||
|
display_tool_messages(tool_messages)
|
||||||
|
|
||||||
|
history.append({"role": role, "content": message_content})
|
||||||
|
|
||||||
except KeyboardInterrupt:
|
except KeyboardInterrupt:
|
||||||
console.print("Chat stopped by user.", style="bold blue")
|
console.print("Chat stopped by user.", style="bold blue")
|
||||||
|
|
|
||||||
|
|
@ -55,13 +55,35 @@ def create_cli_catalog(
|
||||||
return catalog
|
return catalog
|
||||||
|
|
||||||
|
|
||||||
def display_streamed_markdown(stream: Stream[ChatCompletionChunk], model: str) -> tuple[str, str]:
|
def display_tool_messages(tool_messages: list[dict]) -> None:
|
||||||
|
for message in tool_messages:
|
||||||
|
if message["role"] == "assistant":
|
||||||
|
for tool_call in message.get("tool_calls", []):
|
||||||
|
console.print(
|
||||||
|
f"[bright_black]Called tool '{tool_call['function']['name']}' with parameters: {tool_call['function']['arguments']}[/bright_black]"
|
||||||
|
)
|
||||||
|
elif message["role"] == "tool":
|
||||||
|
console.print(
|
||||||
|
f"[bright_black]Tool '{message['name']}' returned: {message['content']}[/bright_black]"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def get_tool_messages(choice: dict) -> list[dict]:
|
||||||
|
if hasattr(choice, "tool_messages") and choice.tool_messages:
|
||||||
|
return choice.tool_messages # type: ignore[no-any-return]
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
def display_streamed_markdown(
|
||||||
|
stream: Stream[ChatCompletionChunk], model: str
|
||||||
|
) -> tuple[str, str, list]:
|
||||||
"""
|
"""
|
||||||
Display the streamed markdown chunks as a single line.
|
Display the streamed markdown chunks as a single line.
|
||||||
"""
|
"""
|
||||||
from rich.live import Live
|
from rich.live import Live
|
||||||
|
|
||||||
full_message = ""
|
full_message = ""
|
||||||
|
tool_messages = []
|
||||||
role = ""
|
role = ""
|
||||||
with Live(console=console, refresh_per_second=10) as live:
|
with Live(console=console, refresh_per_second=10) as live:
|
||||||
for chunk in stream:
|
for chunk in stream:
|
||||||
|
|
@ -76,12 +98,15 @@ def display_streamed_markdown(stream: Stream[ChatCompletionChunk], model: str) -
|
||||||
markdown_chunk = Markdown(full_message)
|
markdown_chunk = Markdown(full_message)
|
||||||
live.update(markdown_chunk)
|
live.update(markdown_chunk)
|
||||||
|
|
||||||
|
# Display and get tool messages if they exist
|
||||||
|
tool_messages += get_tool_messages(choice) # type: ignore[arg-type]
|
||||||
|
|
||||||
# Markdownify URLs in the final message if applicable
|
# Markdownify URLs in the final message if applicable
|
||||||
if role == "assistant":
|
if role == "assistant":
|
||||||
full_message = markdownify_urls(full_message)
|
full_message = markdownify_urls(full_message)
|
||||||
live.update(Markdown(full_message))
|
live.update(Markdown(full_message))
|
||||||
|
|
||||||
return role, full_message
|
return role, full_message, tool_messages
|
||||||
|
|
||||||
|
|
||||||
def markdownify_urls(message: str) -> str:
|
def markdownify_urls(message: str) -> str:
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue