This commit is contained in:
parent
3dde59d8be
commit
8c8fda3da6
|
|
@ -13,11 +13,11 @@ from application.states.models import (
|
||||||
from application.states import ConversationState, AuthState
|
from application.states import ConversationState, AuthState
|
||||||
|
|
||||||
|
|
||||||
def conversation_history_item_showing(
|
def conversation_history_item(
|
||||||
item: Tuple[str, Conversation],
|
item: Tuple[str, Conversation],
|
||||||
) -> rx.Component:
|
) -> rx.Component:
|
||||||
"""
|
"""
|
||||||
会话历史项展示
|
会话历史项
|
||||||
:param item: 会话历史项
|
:param item: 会话历史项
|
||||||
:return: Component
|
:return: Component
|
||||||
"""
|
"""
|
||||||
|
|
@ -25,7 +25,8 @@ def conversation_history_item_showing(
|
||||||
|
|
||||||
# 高亮:若为当前会话或显示更多则高亮
|
# 高亮:若为当前会话或显示更多则高亮
|
||||||
highlight: bool = (
|
highlight: bool = (
|
||||||
conversation.id == ConversationState.is_conversation_history_item_more_button_shown
|
conversation.id
|
||||||
|
== ConversationState.is_conversation_history_item_more_button_shown
|
||||||
) | (conversation.id == ConversationState.conversation_id)
|
) | (conversation.id == ConversationState.conversation_id)
|
||||||
|
|
||||||
return rx.list.item(
|
return rx.list.item(
|
||||||
|
|
@ -172,12 +173,12 @@ def conversation_history_item_showing(
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def conversation_history_items_showing(
|
def conversation_history(
|
||||||
is_conversation_history_shown: bool,
|
is_shown: bool,
|
||||||
) -> rx.Component:
|
) -> rx.Component:
|
||||||
"""
|
"""
|
||||||
会话历史列表展示
|
会话历史
|
||||||
:param is_conversation_history_shown: 会话历史展示状态
|
:param is_shown: 展示 / 隐藏
|
||||||
:return: Component
|
:return: Component
|
||||||
"""
|
"""
|
||||||
return rx.box(
|
return rx.box(
|
||||||
|
|
@ -195,8 +196,8 @@ def conversation_history_items_showing(
|
||||||
),
|
),
|
||||||
rx.auto_scroll(
|
rx.auto_scroll(
|
||||||
rx.foreach(
|
rx.foreach(
|
||||||
ConversationState.conversations,
|
ConversationState.conversation_history,
|
||||||
conversation_history_item_showing,
|
conversation_history_item,
|
||||||
),
|
),
|
||||||
flex="1",
|
flex="1",
|
||||||
align_items="stretch",
|
align_items="stretch",
|
||||||
|
|
@ -217,9 +218,9 @@ def conversation_history_items_showing(
|
||||||
backdrop_filter="blur(50px)",
|
backdrop_filter="blur(50px)",
|
||||||
gap="12px",
|
gap="12px",
|
||||||
),
|
),
|
||||||
width=rx.cond(is_conversation_history_shown, "25%", "0px"),
|
width=rx.cond(is_shown, "25%", "0px"),
|
||||||
min_width=rx.cond(is_conversation_history_shown, "240px", "0px"),
|
min_width=rx.cond(is_shown, "240px", "0px"),
|
||||||
max_width=rx.cond(is_conversation_history_shown, "380px", "0px"),
|
max_width=rx.cond(is_shown, "380px", "0px"),
|
||||||
height="100%",
|
height="100%",
|
||||||
transition="all 0.18s ease-in-out",
|
transition="all 0.18s ease-in-out",
|
||||||
overflow="hidden",
|
overflow="hidden",
|
||||||
|
|
@ -864,21 +865,22 @@ def conversation_history_collapse_button(
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def conversation() -> rx.Component:
|
def conversation_page() -> rx.Component:
|
||||||
"""
|
"""
|
||||||
会话页面:布局参考 MetaChat,左侧为折叠面板,右侧为工作区。其中,工作区包含展示区和操作区。若对话历史为空则展示欢迎文案,否则展示对话历史
|
会话页面:布局参考 MetaChat,从左到右分别为会话历史和工作区。其中,工作区从上到下分别为对话历史和用户提示词输入框。
|
||||||
"""
|
"""
|
||||||
# 会话历史展示状态
|
# 会话历史展示状态
|
||||||
is_conversation_history_shown = ConversationState.is_conversation_history_shown
|
is_shown = ConversationState.is_conversation_history_shown
|
||||||
|
|
||||||
return rx.box(
|
return rx.box(
|
||||||
rx.hstack(
|
rx.hstack(
|
||||||
# 会话历史列表展示
|
# 会话历史
|
||||||
conversation_history_items_showing(is_conversation_history_shown),
|
conversation_history(is_shown),
|
||||||
|
# 工作区
|
||||||
rx.box(
|
rx.box(
|
||||||
|
# 对话历史
|
||||||
rx.vstack(
|
rx.vstack(
|
||||||
# 对话列表展示
|
dialog_history(),
|
||||||
dialog_items_showing(),
|
|
||||||
width="100%",
|
width="100%",
|
||||||
height="100%",
|
height="100%",
|
||||||
gap="8px",
|
gap="8px",
|
||||||
|
|
@ -903,7 +905,7 @@ def conversation() -> rx.Component:
|
||||||
transition="all 0.18s ease-in-out",
|
transition="all 0.18s ease-in-out",
|
||||||
),
|
),
|
||||||
# 会话历史折叠面板按钮
|
# 会话历史折叠面板按钮
|
||||||
conversation_history_collapse_button(is_conversation_history_shown),
|
conversation_history_collapse_button(is_shown),
|
||||||
# 挂载事件:恢复会话状态
|
# 挂载事件:恢复会话状态
|
||||||
on_mount=AuthState.resume_conversation_state,
|
on_mount=AuthState.resume_conversation_state,
|
||||||
position="relative",
|
position="relative",
|
||||||
|
|
|
||||||
|
|
@ -28,10 +28,10 @@ from application.states.models import (
|
||||||
Message,
|
Message,
|
||||||
MessageType,
|
MessageType,
|
||||||
Work,
|
Work,
|
||||||
|
WorkType,
|
||||||
ConversationHistoryItem,
|
ConversationHistoryItem,
|
||||||
usage_to_dict,
|
usage_validate_python,
|
||||||
)
|
)
|
||||||
from application.tasks import run_stream_events
|
|
||||||
|
|
||||||
|
|
||||||
def format_conversation_created_at(created_at: datetime) -> str:
|
def format_conversation_created_at(created_at: datetime) -> str:
|
||||||
|
|
@ -69,9 +69,6 @@ class ConversationState(rx.State):
|
||||||
# 会话历史项显示悬停气泡,True 表示显示,False 表示隐藏
|
# 会话历史项显示悬停气泡,True 表示显示,False 表示隐藏
|
||||||
is_conversation_history_item_popover_shown: bool = False
|
is_conversation_history_item_popover_shown: bool = False
|
||||||
|
|
||||||
# 当前工作实例
|
|
||||||
work: Work | None = None
|
|
||||||
|
|
||||||
# 当前数据库状态(私有变量)
|
# 当前数据库状态(私有变量)
|
||||||
# 私有变量:reflex 约定以 _ 开头的变量为私有变量,后端不序列化,前端不可使用
|
# 私有变量:reflex 约定以 _ 开头的变量为私有变量,后端不序列化,前端不可使用
|
||||||
_db_state: DatabaseState | None = None
|
_db_state: DatabaseState | None = None
|
||||||
|
|
@ -117,7 +114,7 @@ class ConversationState(rx.State):
|
||||||
self.is_conversation_history_shown = not self.is_conversation_history_shown
|
self.is_conversation_history_shown = not self.is_conversation_history_shown
|
||||||
|
|
||||||
@rx.var
|
@rx.var
|
||||||
def conversation_history_items(self) -> list[ConversationHistoryItem]:
|
def conversation_history(self) -> list[ConversationHistoryItem]:
|
||||||
"""
|
"""
|
||||||
获取会话历史项列表(按照会话创建日期时间降序排序)
|
获取会话历史项列表(按照会话创建日期时间降序排序)
|
||||||
:return: 会话历史项列表
|
:return: 会话历史项列表
|
||||||
|
|
@ -240,69 +237,88 @@ class ConversationState(rx.State):
|
||||||
conversation.is_running = True
|
conversation.is_running = True
|
||||||
# 构建用户提示词消息实例
|
# 构建用户提示词消息实例
|
||||||
message = Message(
|
message = Message(
|
||||||
type=MessageType.USER_PROMPT, content=conversation.user_prompt
|
type=MessageType.USER_PROMPT,
|
||||||
|
content=(user_prompt := conversation.user_prompt),
|
||||||
)
|
)
|
||||||
|
# 添加至消息字典
|
||||||
|
conversation.messages[message.id] = message
|
||||||
# 清空用户提示词
|
# 清空用户提示词
|
||||||
conversation.user_prompt = ""
|
conversation.user_prompt = ""
|
||||||
# 添加至消息字典
|
|
||||||
user_prompt = conversation.messages.setdefault(message.id, message)
|
|
||||||
yield # 通知前端更新渲染
|
yield # 通知前端更新渲染
|
||||||
|
|
||||||
# 获取数据库状态
|
|
||||||
db_state = await self.get_db_state()
|
|
||||||
|
|
||||||
# 初始化片段索引映射为消息实例唯一标识字典
|
# 初始化片段索引映射为消息实例唯一标识字典
|
||||||
index_map_to_message_id: dict[int, str] = {}
|
index_map_to_message_id: dict[int, str] = {}
|
||||||
# 获取运行流式输出事件
|
# 初始化工具调用唯一标识集合
|
||||||
async for event in run_stream_events(
|
tool_call_ids: set[str] = set()
|
||||||
work=self.work,
|
try:
|
||||||
user_prompt=user_prompt,
|
# 获取数据库状态
|
||||||
message_history=await db_state.get_message_history(
|
db_state = await self.get_db_state()
|
||||||
|
# 获取消息历史列表
|
||||||
|
message_history = await db_state.get_message_history(
|
||||||
conversation_id=self.actived_conversation_id
|
conversation_id=self.actived_conversation_id
|
||||||
),
|
)
|
||||||
):
|
usage = usage_validate_python(conversation.usage)
|
||||||
|
# 匹配工作类型
|
||||||
|
match conversation.work:
|
||||||
|
# 预定航班
|
||||||
|
case WorkType.BOOK_FLIGHT:
|
||||||
|
from application.workshop.book_flight import run_stream_events
|
||||||
|
|
||||||
|
stream_events = run_stream_events(
|
||||||
|
work=conversation.work,
|
||||||
|
user_prompt=user_prompt,
|
||||||
|
message_history=message_history,
|
||||||
|
usage=usage,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 非结构化对话
|
||||||
|
case _:
|
||||||
|
from application.workshop.unstructured_dialogue import (
|
||||||
|
run_stream_events,
|
||||||
|
)
|
||||||
|
|
||||||
|
stream_events = run_stream_events(
|
||||||
|
user_prompt=user_prompt,
|
||||||
|
message_history=message_history,
|
||||||
|
usage=usage,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 获取运行流式输出事件
|
||||||
|
async for event in stream_events:
|
||||||
match event:
|
match event:
|
||||||
# ========== 开始事件 ==========
|
# ========== 开始事件 ==========
|
||||||
case PartStartEvent(
|
case PartStartEvent(
|
||||||
index=index,
|
index=index,
|
||||||
part=part,
|
part=part,
|
||||||
previous_part_kind=previous_part_kind,
|
|
||||||
):
|
):
|
||||||
match part:
|
match part:
|
||||||
# 思考分片开始事件
|
# 思考分片开始事件
|
||||||
case ThinkingPart(content=content):
|
case ThinkingPart(content=content):
|
||||||
message = Message(type=MessageType.THINKING)
|
|
||||||
# 将消息实例唯一标识与片段索引映射
|
|
||||||
message_id_map_to_index[index] = message.id
|
|
||||||
# 若上一分片种类为空则将正在思考设置为是
|
|
||||||
if not previous_part_kind:
|
|
||||||
# 构建消息实例
|
# 构建消息实例
|
||||||
message.is_thinking = True
|
message = Message(
|
||||||
message.content = content
|
type=MessageType.THINKING,
|
||||||
|
title="正在思考",
|
||||||
# 工具检索分片开始事件
|
content=content,
|
||||||
case ToolSearchCallPart(tool_call_id=tool_call_id):
|
is_running=True,
|
||||||
# 创建工具调用唯一标识与片段索引映射
|
|
||||||
tool_call_ids[tool_call_id] = index
|
|
||||||
|
|
||||||
dialog.thoughts[index] = Thought(
|
|
||||||
type="tool_search",
|
|
||||||
content="正在生成检索关键词",
|
|
||||||
)
|
|
||||||
|
|
||||||
# 能力加载分片开始事件
|
|
||||||
case LoadCapabilityCallPart(tool_call_id=tool_call_id):
|
|
||||||
tool_call_ids[tool_call_id] = index
|
|
||||||
|
|
||||||
dialog.thoughts[index] = Thought(
|
|
||||||
type="capability_load",
|
|
||||||
content="正在生成加载参数",
|
|
||||||
)
|
)
|
||||||
|
# 添加至消息字典
|
||||||
|
conversation.messages[message.id] = message
|
||||||
|
# 将消息实例唯一标识与片段索引映射
|
||||||
|
index_map_to_message_id[index] = message.id
|
||||||
|
|
||||||
# 工具调用分片开始事件
|
# 工具调用分片开始事件
|
||||||
case ToolCallPart(tool_call_id=tool_call_id):
|
case ToolCallPart(
|
||||||
tool_call_ids[tool_call_id] = index
|
tool_name=tool_name, tool_call_id=tool_call_id
|
||||||
|
):
|
||||||
|
# 构建消息实例
|
||||||
|
message = Message(
|
||||||
|
type=MessageType.TOOL_CALL,
|
||||||
|
title=tool_name,
|
||||||
|
content=",",
|
||||||
|
is_running=True,
|
||||||
|
)
|
||||||
|
tool_call_ids.add(tool_call_id)
|
||||||
|
# 添加至消息字典
|
||||||
dialog.thoughts[index] = Thought(
|
dialog.thoughts[index] = Thought(
|
||||||
type="tool_call",
|
type="tool_call",
|
||||||
content="正在生成调用参数",
|
content="正在生成调用参数",
|
||||||
|
|
@ -310,7 +326,13 @@ class ConversationState(rx.State):
|
||||||
|
|
||||||
# 文本分片开始事件
|
# 文本分片开始事件
|
||||||
case TextPart(content=content):
|
case TextPart(content=content):
|
||||||
dialog.result_output = content
|
# 构建消息实例
|
||||||
|
message = Message(type=MessageType.RESULT_OUTPUT)
|
||||||
|
message.content = content
|
||||||
|
# 添加至消息字典
|
||||||
|
conversation.messages[message.id] = message
|
||||||
|
# 将消息实例唯一标识与片段索引映射
|
||||||
|
index_map_to_message_id[index] = message.id
|
||||||
|
|
||||||
# ========== 增量事件 ==========
|
# ========== 增量事件 ==========
|
||||||
case PartDeltaEvent(index=index, delta=delta):
|
case PartDeltaEvent(index=index, delta=delta):
|
||||||
|
|
@ -319,26 +341,35 @@ class ConversationState(rx.State):
|
||||||
case ThinkingPartDelta(
|
case ThinkingPartDelta(
|
||||||
content_delta=content_delta,
|
content_delta=content_delta,
|
||||||
):
|
):
|
||||||
dialog.thoughts[index].content += content_delta or ""
|
# 若内容增量不为空则增量更新
|
||||||
|
if content_delta:
|
||||||
|
conversation.messages[
|
||||||
|
index_map_to_message_id[index]
|
||||||
|
].content += content_delta
|
||||||
|
|
||||||
# 文本分片增量事件
|
# 文本分片增量事件
|
||||||
case TextPartDelta(
|
case TextPartDelta(
|
||||||
content_delta=content_delta,
|
content_delta=content_delta,
|
||||||
):
|
):
|
||||||
dialog.result_output += content_delta
|
# 增量更新
|
||||||
|
conversation.messages[
|
||||||
|
index_map_to_message_id[index]
|
||||||
|
].content += content_delta
|
||||||
|
|
||||||
# ========== 结束事件 ==========
|
# ========== 结束事件 ==========
|
||||||
case PartEndEvent(
|
case PartEndEvent(
|
||||||
index=index,
|
index=index,
|
||||||
part=part,
|
part=part,
|
||||||
next_part_kind=next_part_kind,
|
|
||||||
):
|
):
|
||||||
match part:
|
match part:
|
||||||
# 思考分片结束事件
|
# 思考分片结束事件
|
||||||
case ThinkingPart(content=content):
|
case ThinkingPart(content=content):
|
||||||
# 若下一分片种类为文本则将思考状态设置为思考完成
|
# 获取消息实例
|
||||||
if next_part_kind == "text":
|
message = conversation.messages[
|
||||||
dialog.is_thinking = False
|
index_map_to_message_id[index]
|
||||||
|
]
|
||||||
|
message.is_running = False
|
||||||
|
message.title = "思考完成"
|
||||||
|
|
||||||
# ========== 函数工具调用事件 ==========
|
# ========== 函数工具调用事件 ==========
|
||||||
case FunctionToolCallEvent(tool_call_id=tool_call_id, part=part):
|
case FunctionToolCallEvent(tool_call_id=tool_call_id, part=part):
|
||||||
|
|
@ -403,17 +434,21 @@ class ConversationState(rx.State):
|
||||||
# 创建结果记录
|
# 创建结果记录
|
||||||
await db_state.create_result_record(
|
await db_state.create_result_record(
|
||||||
conversation_id=self.conversation_id,
|
conversation_id=self.conversation_id,
|
||||||
dialog_id=dialog.id,
|
_id=dialog.id,
|
||||||
new_messages=result.new_messages(),
|
new_messages=result.new_messages(),
|
||||||
)
|
)
|
||||||
|
|
||||||
# 强制更新会话并推送前端
|
# 强制更新会话并推送前端
|
||||||
self.conversations[self.conversation_id] = conversation
|
self.conversations[self.actived_conversation_id] = conversation
|
||||||
yield
|
yield
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
...
|
||||||
|
finally:
|
||||||
# 将正在运行设置为否
|
# 将正在运行设置为否
|
||||||
conversation.is_running = False
|
conversation.is_running = False
|
||||||
self.conversations[self.conversation_id] = conversation
|
|
||||||
|
self.conversations[self.actived_conversation_id] = conversation
|
||||||
yield
|
yield
|
||||||
|
|
||||||
@rx.event
|
@rx.event
|
||||||
|
|
@ -430,10 +465,12 @@ class ConversationState(rx.State):
|
||||||
self.user_prompt = f"帮我找一班从 {conversation.task.deps.origin} 到 {conversation.task.deps.destination} 在 {conversation.task.deps.date} 的航班"
|
self.user_prompt = f"帮我找一班从 {conversation.task.deps.origin} 到 {conversation.task.deps.destination} 在 {conversation.task.deps.date} 的航班"
|
||||||
|
|
||||||
@rx.event
|
@rx.event
|
||||||
def toggle_collapse(self, dialog_id: str) -> None:
|
def toggle_message_collapse(self, message_id: str) -> None:
|
||||||
"""
|
"""
|
||||||
展开/折叠思考折叠面板
|
展开/折叠消息组件
|
||||||
"""
|
"""
|
||||||
# 指定运行
|
# 当前会话
|
||||||
dialog = self.conversations[self.conversation_id].dialogs[dialog_id]
|
conversation = self.conversations.get(self.actived_conversation_id)
|
||||||
dialog.is_expanded = not dialog.is_expanded
|
if not conversation:
|
||||||
|
return
|
||||||
|
conversation.messages[message_id].is_expanded ^= True
|
||||||
|
|
|
||||||
|
|
@ -14,11 +14,9 @@ from sqlmodel import Field, JSON, SQLModel, select, update
|
||||||
|
|
||||||
from application.states.models import (
|
from application.states.models import (
|
||||||
Conversation,
|
Conversation,
|
||||||
TaskType,
|
|
||||||
TaskStatus,
|
TaskStatus,
|
||||||
RunStatus,
|
RunStatus,
|
||||||
MessageType,
|
MessageType,
|
||||||
Task,
|
|
||||||
Run,
|
Run,
|
||||||
Message,
|
Message,
|
||||||
deps_to_object,
|
deps_to_object,
|
||||||
|
|
@ -109,7 +107,7 @@ class MessageRecord(SQLModel, table=True, table_name="message"):
|
||||||
content: str = Field(default="", description="消息内容")
|
content: str = Field(default="", description="消息内容")
|
||||||
|
|
||||||
|
|
||||||
class RunRecord(SQLModel, table=True, table_name="run_result"):
|
class RunRecord(SQLModel, table=True, table_name="run"):
|
||||||
"""
|
"""
|
||||||
运行记录
|
运行记录
|
||||||
"""
|
"""
|
||||||
|
|
@ -333,9 +331,9 @@ class DatabaseState(rx.State):
|
||||||
)
|
)
|
||||||
for run in result.all():
|
for run in result.all():
|
||||||
message_history.extend(
|
message_history.extend(
|
||||||
ModelMessagesTypeAdapter.validate_python(
|
ModelMessagesTypeAdapter.validate_json(
|
||||||
run.new_messages
|
run.new_messages
|
||||||
) # 将 run.new_messages 由 Python 类型反序列化为 List[ModelMessage]
|
) # 将 run.new_messages 由 JSON 字符串反序列化为 List[ModelMessage]
|
||||||
)
|
)
|
||||||
return message_history
|
return message_history
|
||||||
|
|
||||||
|
|
@ -357,9 +355,9 @@ class DatabaseState(rx.State):
|
||||||
RunRecord(
|
RunRecord(
|
||||||
id=id,
|
id=id,
|
||||||
conversation_id=conversation_id,
|
conversation_id=conversation_id,
|
||||||
new_messages=ModelMessagesTypeAdapter.dump_python(
|
new_messages=ModelMessagesTypeAdapter.dump_json(
|
||||||
new_messages
|
new_messages
|
||||||
), # 将 messages 由 List[ModelMessage] 序列化为 Python 类型
|
).decode(), # 将 messages 由 List[ModelMessage] 序列化为 JSON 字符串
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
await session.commit()
|
await session.commit()
|
||||||
|
|
|
||||||
|
|
@ -19,7 +19,9 @@ class MessageType(StrEnum):
|
||||||
|
|
||||||
USER_PROMPT = "user_prompt"
|
USER_PROMPT = "user_prompt"
|
||||||
THINKING = "thinking"
|
THINKING = "thinking"
|
||||||
AGENT_RUN_RESULT_OUTPUT = "agent_run_result_output"
|
TOOL_CALL = "tool_call"
|
||||||
|
RESULT_OUTPUT = "result_output"
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class Message(BaseModel):
|
class Message(BaseModel):
|
||||||
|
|
@ -31,8 +33,8 @@ class Message(BaseModel):
|
||||||
type: MessageType = Field(..., description="消息类型")
|
type: MessageType = Field(..., description="消息类型")
|
||||||
title: str = Field(default="", description="消息标题")
|
title: str = Field(default="", description="消息标题")
|
||||||
content: str = Field(default="", description="消息内容")
|
content: str = Field(default="", description="消息内容")
|
||||||
is_thinking: bool = Field(
|
is_running: bool = Field(
|
||||||
default=False, description="正在思考,True 表示正在思考,False 表示思考完成"
|
default=False, description="正在运行,True 表示正在运行,False 表示运行完成"
|
||||||
)
|
)
|
||||||
is_expanded: bool = Field(
|
is_expanded: bool = Field(
|
||||||
default=False, description="展开组件,True 表示展开,False 表示折叠"
|
default=False, description="展开组件,True 表示展开,False 表示折叠"
|
||||||
|
|
@ -109,11 +111,11 @@ class Conversation(BaseModel):
|
||||||
|
|
||||||
id: str = Field(..., description="会话唯一标识")
|
id: str = Field(..., description="会话唯一标识")
|
||||||
description: str = Field(..., description="会话描述")
|
description: str = Field(..., description="会话描述")
|
||||||
|
work: Work | None = Field(default=None, description="绑定工作实例")
|
||||||
user_prompt: str = Field(default="", description="用户提示词")
|
user_prompt: str = Field(default="", description="用户提示词")
|
||||||
usage: dict[str, Any] = Field(..., description="会话使用量")
|
usage: dict[str, Any] = Field(..., description="会话使用量")
|
||||||
messages: dict[str, Message] = Field(..., description="消息字典")
|
messages: dict[str, Message] = Field(..., description="消息字典")
|
||||||
created_at: datetime = Field(..., description="会话创建日期时间")
|
created_at: datetime = Field(..., description="会话创建日期时间")
|
||||||
|
|
||||||
is_running: bool = Field(
|
is_running: bool = Field(
|
||||||
default=False, description="正在运行,True 表示正在运行, False 表示运行结束"
|
default=False, description="正在运行,True 表示正在运行, False 表示运行结束"
|
||||||
)
|
)
|
||||||
|
|
|
||||||
|
|
@ -1,63 +0,0 @@
|
||||||
# -*- coding: utf-8 -*-
|
|
||||||
"""
|
|
||||||
任务模块
|
|
||||||
"""
|
|
||||||
from typing import AsyncGenerator, List
|
|
||||||
|
|
||||||
from pydantic_ai import Agent, ModelMessage
|
|
||||||
from pydantic_ai.messages import AgentStreamEvent
|
|
||||||
from pydantic_ai.run import AgentRunResultEvent
|
|
||||||
|
|
||||||
from application.states.models import TaskNodeResultEvent, TaskType, Task, Message
|
|
||||||
from application.tasks.models import DEEPSEEK_V4_FLASH_MODEL
|
|
||||||
|
|
||||||
|
|
||||||
instruction = """
|
|
||||||
# 角色
|
|
||||||
专业友好AI助手,结构化解答各类问题。
|
|
||||||
|
|
||||||
# 输出硬性规则
|
|
||||||
1. 全文强制标准Markdown,禁止纯文本;不要额外说明排版格式,直接输出内容;
|
|
||||||
2. 层级使用 `#/##/###`,列表用 `-` 无序列表或数字有序列表;
|
|
||||||
3. 代码块用 ```语言名``` 包裹;
|
|
||||||
4. 重点内容标注 **粗体**/*斜体*;
|
|
||||||
5. 思考、工具日志仅输出文本,适配前端折叠面板,禁止输出HTML标签;
|
|
||||||
6. 内容分点拆分,排版整洁适配前端Markdown渲染。
|
|
||||||
|
|
||||||
# 行文要求
|
|
||||||
语言通俗,逻辑完整简洁,无多余废话。
|
|
||||||
"""
|
|
||||||
|
|
||||||
|
|
||||||
async def run_stream_events(
|
|
||||||
work: Work | None,
|
|
||||||
user_prompt: Message,
|
|
||||||
message_history: List[ModelMessage],
|
|
||||||
) -> AsyncGenerator[AgentStreamEvent | AgentRunResultEvent | TaskNodeResultEvent, None]:
|
|
||||||
"""
|
|
||||||
以流式事件模式运行
|
|
||||||
"""
|
|
||||||
if task:
|
|
||||||
match task.type:
|
|
||||||
case TaskType.BOOK_FLIGHT:
|
|
||||||
from application.tasks.book_flight import run_stream_events
|
|
||||||
|
|
||||||
async for event in run_stream_events(
|
|
||||||
task=task,
|
|
||||||
user_prompt=user_prompt,
|
|
||||||
message_history=message_history,
|
|
||||||
):
|
|
||||||
print(event)
|
|
||||||
yield event
|
|
||||||
|
|
||||||
else:
|
|
||||||
agent = Agent(
|
|
||||||
model=DEEPSEEK_V4_FLASH_MODEL,
|
|
||||||
instructions=instruction,
|
|
||||||
)
|
|
||||||
async with agent.run_stream_events(
|
|
||||||
user_prompt=user_prompt,
|
|
||||||
message_history=message_history,
|
|
||||||
) as events:
|
|
||||||
async for event in events:
|
|
||||||
yield event
|
|
||||||
|
|
@ -0,0 +1 @@
|
||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
|
@ -1,6 +1,6 @@
|
||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
"""
|
"""
|
||||||
预定航班任务
|
预定航班
|
||||||
"""
|
"""
|
||||||
import datetime
|
import datetime
|
||||||
from typing import AsyncGenerator, Literal
|
from typing import AsyncGenerator, Literal
|
||||||
|
|
@ -181,9 +181,11 @@ def init_task() -> Task:
|
||||||
|
|
||||||
|
|
||||||
async def run_stream_events(
|
async def run_stream_events(
|
||||||
task: Task | None,
|
usage: RunUsage,
|
||||||
user_prompt: str,
|
user_prompt: str,
|
||||||
message_history: list[ModelMessage],
|
message_history: list[ModelMessage],
|
||||||
|
work: WorkType | None = None,
|
||||||
|
|
||||||
) -> AsyncGenerator:
|
) -> AsyncGenerator:
|
||||||
result = None
|
result = None
|
||||||
while True:
|
while True:
|
||||||
|
|
@ -0,0 +1,54 @@
|
||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
"""
|
||||||
|
非结构化对话
|
||||||
|
"""
|
||||||
|
from typing import AsyncGenerator
|
||||||
|
|
||||||
|
from pydantic_ai import Agent, ModelMessage, RunUsage
|
||||||
|
from pydantic_ai.messages import AgentStreamEvent
|
||||||
|
from pydantic_ai.run import AgentRunResultEvent
|
||||||
|
|
||||||
|
from application.workshop.models import DEEPSEEK_V4_FLASH_MODEL
|
||||||
|
|
||||||
|
|
||||||
|
instruction = """
|
||||||
|
# 角色
|
||||||
|
专业友好AI助手,结构化解答各类问题。
|
||||||
|
|
||||||
|
# 输出硬性规则
|
||||||
|
1. 全文强制标准Markdown,禁止纯文本;不要额外说明排版格式,直接输出内容;
|
||||||
|
2. 层级使用 `#/##/###`,列表用 `-` 无序列表或数字有序列表;
|
||||||
|
3. 代码块用 ```语言名``` 包裹;
|
||||||
|
4. 重点内容标注 **粗体**/*斜体*;
|
||||||
|
5. 思考、工具日志仅输出文本,适配前端折叠面板,禁止输出HTML标签;
|
||||||
|
6. 内容分点拆分,排版整洁适配前端Markdown渲染。
|
||||||
|
|
||||||
|
# 行文要求
|
||||||
|
语言通俗,逻辑完整简洁,无多余废话。
|
||||||
|
"""
|
||||||
|
|
||||||
|
agent = Agent(
|
||||||
|
model=DEEPSEEK_V4_FLASH_MODEL,
|
||||||
|
instructions=instruction,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
async def run_stream_events(
|
||||||
|
user_prompt: str,
|
||||||
|
message_history: list[ModelMessage],
|
||||||
|
usage: RunUsage,
|
||||||
|
) -> AsyncGenerator[AgentStreamEvent | AgentRunResultEvent, None]:
|
||||||
|
"""
|
||||||
|
运行并流式输出事件
|
||||||
|
:param user_prompt: 用户提示词
|
||||||
|
:param message_history: 消息历史
|
||||||
|
:param usage: 使用量
|
||||||
|
:return: AsyncGenerator
|
||||||
|
"""
|
||||||
|
async with agent.run_stream_events(
|
||||||
|
user_prompt=user_prompt,
|
||||||
|
message_history=message_history,
|
||||||
|
usage=usage,
|
||||||
|
) as events:
|
||||||
|
async for event in events:
|
||||||
|
yield event
|
||||||
Loading…
Reference in New Issue