Код: Выделить всё
main.pyКод: Выделить всё
import asyncio
from langchain_community.chat_models import ChatOpenAI
from langchain.agents import AgentExecutor, create_openai_tools_agent
from langchain.tools import tool
import gradio as gr
from go2_robot_sdk.scripts.webrtc_driver import Go2Connection
langchain_llm_client = ChatOpenAI(
model='gpt-4o',
temperature=0.,
api_key=OPENAI_API_KEY,
streaming=True,
max_tokens=None,
)
conn = Go2Connection(
robot_ip=ROBOT_IP,
robot_num="0",
token="",
on_open=lambda: logger.info("Data channel opened"),
on_validated=lambda _: logger.info("Robot validated"),
on_message=lambda msg, _, __: logger.info(f"Message received: {msg}")
)
async def connect_to_robot() -> None:
await conn.connect()
logger.info("Connected to robot")
while conn.robot_validation == "PENDING":
await asyncio.sleep(0.1)
if conn.robot_validation != "SUCCESS":
logger.info("Failed to validate robot connection")
return
logger.info("Robot connection validated successfully")
@tool
async def go2_scrape():
conn.data_channel.send(
gen_command(ROBOT_CMD['Scrape'])
)
logger.info(f"Scrape command sent")
async def agent_completion(
agent_executor,
message: str,
tools: List,
) -> AsyncGenerator:
tool_names = [tool.name for tool in tools]
async for event in agent_executor.astream_events(
{
"input": message,
"tools": tools,
"tool_names": tool_names,
"agent_scratchpad": lambda x: format_to_openai_tool_messages(x["intermediate_steps"]),
},
version='v2'
):
kind = event['event']
if kind == "on_chain_start":
if (
event["name"] == "Agent"
):
yield(
f"\n### 執行代理: `{event['name']}`,代理輸入: `{event['data'].get('input')}`\n"
)
elif kind == "on_chat_model_stream":
content = event["data"]["chunk"].content
if content:
yield content
elif kind == "on_tool_start":
yield(
f"\n### 執行任務: `{event['name']}`,任務輸入: `{event['data'].get('input')}`\n"
)
elif kind == "on_tool_end":
yield(
f"\n### 任務完成: `{event['name']}`,任務結果: \n"
)
if isinstance(event['data'].get('output'), AsyncGenerator):
async for event_chunk in event['data'].get('output'):
yield event_chunk
else:
yield(
f"`{event['data'].get('output')}`\n"
)
elif kind == "on_chain_end":
if (
event["name"] == "Agent"
):
yield(
f"\n### 代理完成: `{event['name']}`,代理結果: \n"
)
yield(
f"{event['data'].get('output')['output']}\n"
)
def user_text_input(message: str, history: List[List[str]]) -> List[List[str]]:
history.append([message, ""])
return history
async def bot_response(history: List[List[str]]):
global snapshot_session
user_message = history[-1][0]
responses = agent_completion(agent_executor, user_message, tools)
async for response in responses:
history[-1][1] += response
yield history
def clear_history() -> List[List[str]]:
return []
if __name__ == '__main__':
asyncio.run(connect_to_robot())
with gr.Blocks() as demo:
tools = [
go2_scrape,
]
agent = create_openai_tools_agent(langchain_llm_client, tools, AGENT_PROMPT)
agent_executor = AgentExecutor(
agent=agent,
tools=tools,
verbose=False,
return_intermediate_steps=True
)
with gr.Row():
with gr.Column(scale=6):
chatbot = gr.Chatbot(label='Chatbot')
with gr.Row():
with gr.Column():
text_input = gr.Textbox(label="Type your message here", placeholder="Type your message here...")
clear = gr.Button("Clean Conversation History")
text_input.submit(user_text_input, [text_input, chatbot], [chatbot], queue=False).then(
bot_response, chatbot, chatbot
).then(
lambda x: gr.update(value=''), None, [text_input]
)
Код: Выделить всё
go2_ros_sdk/scripts/webrtc_driver.pyКод: Выделить всё
import base64
import hashlib
import json
import logging
import time
import aiohttp
from aiortc import RTCPeerConnection, RTCSessionDescription
from aiortc.contrib.media import MediaBlackhole
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
class Go2Connection():
def __init__(self, robot_ip, robot_num, token="", on_validated=None, on_message=None, on_open=None):
self.pc = RTCPeerConnection()
self.robot_ip = robot_ip
self.robot_num = str(robot_num)
self.token = token
self.robot_validation = "PENDING"
self.on_validated = on_validated
self.on_message = on_message
self.on_open = on_open
self.audio_track = MediaBlackhole()
self.video_track = MediaBlackhole()
self.data_channel = self.pc.createDataChannel("data", id=0)
self.data_channel.on("open", self.on_data_channel_open)
self.data_channel.on("message", self.on_data_channel_message)
self.pc.on("track", self.on_track)
self.pc.on("connectionstatechange", self.on_connection_state_change)
def on_connection_state_change(self):
logger.info(f"Connection state is {self.pc.connectionState}")
def on_track(self, track):
logger.info(f"Receiving {track.kind}")
async def generate_offer(self):
await self.audio_track.start()
await self.video_track.start()
offer = await self.pc.createOffer()
await self.pc.setLocalDescription(offer)
return offer.sdp
async def set_answer(self, sdp):
answer = RTCSessionDescription(sdp, type="answer")
await self.pc.setRemoteDescription(answer)
def on_data_channel_open(self):
logger.info("Data channel is open")
if self.on_open:
self.on_open()
def on_data_channel_message(self, msg):
logger.info(f"Received message: {msg}")
if self.data_channel.readyState != "open":
self.data_channel._setReadyState("open")
try:
if isinstance(msg, str):
msgobj = json.loads(msg)
if msgobj.get("type") == "validation":
self.validate_robot_conn(msgobj)
if self.on_message:
self.on_message(msg, msgobj, self.robot_num)
except json.JSONDecodeError:
pass
async def connect(self):
offer = await self.generate_offer()
url = f"http://{self.robot_ip}:8081/offer"
headers = {"Content-Type": "application/json"}
data = {
"sdp": offer,
"id": "STA_localNetwork",
"type": "offer",
"token": "",
}
connected = False
while not connected:
async with aiohttp.ClientSession() as session:
async with session.post(url, json=data, headers=headers) as resp:
if resp.status == 200:
answer_data = await resp.json()
answer_sdp = answer_data.get("sdp")
await self.set_answer(answer_sdp)
connected = True
else:
logger.info(f"Failed to get answer from server: Reason: {resp}")
logger.info("Try to reconnect...")
time.sleep(1)
def validate_robot_conn(self, message):
if message.get("data") == "Validation Ok.":
self.robot_validation = "SUCCESS"
if self.on_validated:
self.on_validated(self.robot_num)
else:
self.publish(
"",
self.encrypt_key(message.get("data")),
"validation",
)
def publish(self, topic, data, msg_type):
if self.data_channel.readyState != "open":
logger.info(f"Data channel is not open. State is {self.data_channel.readyState}")
return
payload = {
"type": msg_type,
"topic": topic,
"data": data,
}
payload_dumped = json.dumps(payload)
logger.info(f"-> Sending message {payload_dumped}")
self.data_channel.send(payload_dumped)
@staticmethod
def hex_to_base64(hex_str):
bytes_array = bytes.fromhex(hex_str)
return base64.b64encode(bytes_array).decode("utf-8")
@staticmethod
def encrypt_key(key):
prefixed_key = f"UnitreeGo2_{key}"
encrypted = Go2Connection.encrypt_by_md5(prefixed_key)
return Go2Connection.hex_to_base64(encrypted)
@staticmethod
def encrypt_by_md5(input_str):
hash_obj = hashlib.md5()
hash_obj.update(input_str.encode("utf-8"))
return hash_obj.hexdigest()
Я попытался использовать WebRTC для управления роботизированной собакой Go2 и инкапсулировал управляющие действия в агент LangChain инструменты, запускаемые агентом LangChain astream_event. Однако я столкнулся со следующими проблемами:
[*]Инструмент работает правильно при первом запуске, но не работает при последующих попытках без ответа от Go2 и без журналов ошибок. .
[*]При использовании data_channel.send ожидаемое сообщение обратного вызова не отображается в терминале при запуске через агент LangChain.
Из-за отсутствия журнала ошибок определить причину сложно. Я подозреваю, что модуль aiortc может застрять в бесконечном цикле при первом триггере astream_event, что не позволяет инструменту работать со второй попытки.
Подробнее здесь: https://stackoverflow.com/questions/787 ... ream-event