This commit is contained in:
parent
a7fead671a
commit
9e0c58743a
|
|
@ -1,8 +1,8 @@
|
||||||
"""empty message
|
"""empty message
|
||||||
|
|
||||||
Revision ID: 0d3f238840fe
|
Revision ID: d52f1ee0f9c4
|
||||||
Revises:
|
Revises:
|
||||||
Create Date: 2026-08-06 13:35:59.116161
|
Create Date: 2026-08-11 14:59:37.246745
|
||||||
|
|
||||||
"""
|
"""
|
||||||
from typing import Sequence, Union
|
from typing import Sequence, Union
|
||||||
|
|
@ -12,7 +12,7 @@ import sqlalchemy as sa
|
||||||
import sqlmodel
|
import sqlmodel
|
||||||
|
|
||||||
# revision identifiers, used by Alembic.
|
# revision identifiers, used by Alembic.
|
||||||
revision: str = '0d3f238840fe'
|
revision: str = 'd52f1ee0f9c4'
|
||||||
down_revision: Union[str, Sequence[str], None] = None
|
down_revision: Union[str, Sequence[str], None] = None
|
||||||
branch_labels: Union[str, Sequence[str], None] = None
|
branch_labels: Union[str, Sequence[str], None] = None
|
||||||
depends_on: Union[str, Sequence[str], None] = None
|
depends_on: Union[str, Sequence[str], None] = None
|
||||||
|
|
@ -23,11 +23,11 @@ def upgrade() -> None:
|
||||||
# ### commands auto generated by Alembic - please adjust! ###
|
# ### commands auto generated by Alembic - please adjust! ###
|
||||||
op.create_table('captcharecord',
|
op.create_table('captcharecord',
|
||||||
sa.Column('email', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
sa.Column('email', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
sa.Column('captcha', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
sa.Column('code', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
sa.Column('is_valid', sa.Boolean(), nullable=False),
|
sa.Column('is_valid', sa.Boolean(), nullable=False),
|
||||||
sa.Column('is_verified', sa.Boolean(), nullable=False),
|
sa.Column('is_verified', sa.Boolean(), nullable=False),
|
||||||
sa.Column('expired_at', sa.DateTime(), nullable=False),
|
sa.Column('expired_at', sa.DateTime(), nullable=False),
|
||||||
sa.PrimaryKeyConstraint('email', 'captcha', 'expired_at')
|
sa.PrimaryKeyConstraint('email', 'code', 'expired_at')
|
||||||
)
|
)
|
||||||
with op.batch_alter_table('captcharecord', schema=None) as batch_op:
|
with op.batch_alter_table('captcharecord', schema=None) as batch_op:
|
||||||
batch_op.create_index(batch_op.f('ix_captcharecord_is_valid'), ['is_valid'], unique=False)
|
batch_op.create_index(batch_op.f('ix_captcharecord_is_valid'), ['is_valid'], unique=False)
|
||||||
|
|
@ -37,6 +37,7 @@ def upgrade() -> None:
|
||||||
sa.Column('id', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
sa.Column('id', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
sa.Column('user_id', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
sa.Column('user_id', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
sa.Column('description', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
sa.Column('description', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
|
sa.Column('usage', sa.JSON(), nullable=False),
|
||||||
sa.Column('is_deleted', sa.Boolean(), nullable=False),
|
sa.Column('is_deleted', sa.Boolean(), nullable=False),
|
||||||
sa.Column('created_at', sa.DateTime(), nullable=False),
|
sa.Column('created_at', sa.DateTime(), nullable=False),
|
||||||
sa.PrimaryKeyConstraint('id')
|
sa.PrimaryKeyConstraint('id')
|
||||||
|
|
@ -45,24 +46,26 @@ def upgrade() -> None:
|
||||||
batch_op.create_index(batch_op.f('ix_conversationrecord_is_deleted'), ['is_deleted'], unique=False)
|
batch_op.create_index(batch_op.f('ix_conversationrecord_is_deleted'), ['is_deleted'], unique=False)
|
||||||
batch_op.create_index(batch_op.f('ix_conversationrecord_user_id'), ['user_id'], unique=False)
|
batch_op.create_index(batch_op.f('ix_conversationrecord_user_id'), ['user_id'], unique=False)
|
||||||
|
|
||||||
op.create_table('dialogrecord',
|
op.create_table('messagerecord',
|
||||||
sa.Column('id', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
sa.Column('id', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
sa.Column('conversation_id', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
sa.Column('conversation_id', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
sa.Column('user_prompt', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
sa.Column('type', sa.Enum('USER_PROMPT', 'THINKING', 'TOOL_CALL', 'RESULT_OUTPUT', name='messagetype'), nullable=False),
|
||||||
sa.Column('thoughts', sa.JSON(), nullable=False),
|
sa.Column('title', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
sa.Column('result_output', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
sa.Column('content', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
sa.Column('usage', sa.JSON(), nullable=False),
|
|
||||||
sa.PrimaryKeyConstraint('id')
|
sa.PrimaryKeyConstraint('id')
|
||||||
)
|
)
|
||||||
with op.batch_alter_table('dialogrecord', schema=None) as batch_op:
|
with op.batch_alter_table('messagerecord', schema=None) as batch_op:
|
||||||
batch_op.create_index(batch_op.f('ix_dialogrecord_conversation_id'), ['conversation_id'], unique=False)
|
batch_op.create_index(batch_op.f('ix_messagerecord_conversation_id'), ['conversation_id'], unique=False)
|
||||||
|
|
||||||
op.create_table('resultrecord',
|
op.create_table('runrecord',
|
||||||
|
sa.Column('id', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
sa.Column('conversation_id', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
sa.Column('conversation_id', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
sa.Column('dialog_id', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
sa.Column('new_messages', sa.JSON(), nullable=False),
|
||||||
sa.Column('new_messages', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
sa.PrimaryKeyConstraint('id')
|
||||||
sa.PrimaryKeyConstraint('conversation_id', 'dialog_id')
|
|
||||||
)
|
)
|
||||||
|
with op.batch_alter_table('runrecord', schema=None) as batch_op:
|
||||||
|
batch_op.create_index(batch_op.f('ix_runrecord_conversation_id'), ['conversation_id'], unique=False)
|
||||||
|
|
||||||
op.create_table('userrecord',
|
op.create_table('userrecord',
|
||||||
sa.Column('id', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
sa.Column('id', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
sa.Column('email', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
sa.Column('email', sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||||
|
|
@ -81,11 +84,14 @@ def downgrade() -> None:
|
||||||
batch_op.drop_index(batch_op.f('ix_userrecord_email'))
|
batch_op.drop_index(batch_op.f('ix_userrecord_email'))
|
||||||
|
|
||||||
op.drop_table('userrecord')
|
op.drop_table('userrecord')
|
||||||
op.drop_table('resultrecord')
|
with op.batch_alter_table('runrecord', schema=None) as batch_op:
|
||||||
with op.batch_alter_table('dialogrecord', schema=None) as batch_op:
|
batch_op.drop_index(batch_op.f('ix_runrecord_conversation_id'))
|
||||||
batch_op.drop_index(batch_op.f('ix_dialogrecord_conversation_id'))
|
|
||||||
|
|
||||||
op.drop_table('dialogrecord')
|
op.drop_table('runrecord')
|
||||||
|
with op.batch_alter_table('messagerecord', schema=None) as batch_op:
|
||||||
|
batch_op.drop_index(batch_op.f('ix_messagerecord_conversation_id'))
|
||||||
|
|
||||||
|
op.drop_table('messagerecord')
|
||||||
with op.batch_alter_table('conversationrecord', schema=None) as batch_op:
|
with op.batch_alter_table('conversationrecord', schema=None) as batch_op:
|
||||||
batch_op.drop_index(batch_op.f('ix_conversationrecord_user_id'))
|
batch_op.drop_index(batch_op.f('ix_conversationrecord_user_id'))
|
||||||
batch_op.drop_index(batch_op.f('ix_conversationrecord_is_deleted'))
|
batch_op.drop_index(batch_op.f('ix_conversationrecord_is_deleted'))
|
||||||
|
|
@ -3,14 +3,14 @@
|
||||||
会话页面
|
会话页面
|
||||||
"""
|
"""
|
||||||
import reflex as rx
|
import reflex as rx
|
||||||
from typing import Tuple
|
|
||||||
|
|
||||||
|
from application.states import AuthState, ConversationState
|
||||||
from application.states.models import (
|
from application.states.models import (
|
||||||
|
ConversationHistoryItem,
|
||||||
MessageHistoryItem,
|
MessageHistoryItem,
|
||||||
MessageType,
|
MessageType,
|
||||||
ConversationHistoryItem,
|
WorkFlowType,
|
||||||
)
|
)
|
||||||
from application.states import ConversationState, AuthState
|
|
||||||
|
|
||||||
|
|
||||||
def conversation_history_item(
|
def conversation_history_item(
|
||||||
|
|
@ -23,10 +23,9 @@ def conversation_history_item(
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# 高亮
|
# 高亮
|
||||||
is_highlight = item.id in (
|
is_highlight = (
|
||||||
ConversationState.is_conversation_history_item_popover_shown,
|
item.id == ConversationState.is_conversation_history_item_popover_shown
|
||||||
ConversationState.actived_conversation_id,
|
) | (item.id == ConversationState.actived_conversation_id)
|
||||||
)
|
|
||||||
|
|
||||||
return rx.list.item(
|
return rx.list.item(
|
||||||
rx.vstack(
|
rx.vstack(
|
||||||
|
|
@ -227,109 +226,86 @@ def conversation_history(
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def guidance() -> rx.Component:
|
def use_guidance() -> rx.Component:
|
||||||
"""
|
"""
|
||||||
引导
|
使用指引
|
||||||
:return: Component
|
:return: Component
|
||||||
"""
|
"""
|
||||||
return rx.vstack(
|
return rx.vstack(
|
||||||
|
# 品牌图标、名称和介绍
|
||||||
rx.vstack(
|
rx.vstack(
|
||||||
# 品牌图标、名称和介绍
|
rx.hstack(
|
||||||
|
rx.image(
|
||||||
|
src="/logo.png",
|
||||||
|
width="64px",
|
||||||
|
height="64px",
|
||||||
|
),
|
||||||
|
rx.text(
|
||||||
|
"棱镜球",
|
||||||
|
font_size="calc(var(--prismui-font-size-3) * 2)",
|
||||||
|
font_weight="var(--prismui-font-weight-3)",
|
||||||
|
letter_spacing="1px",
|
||||||
|
),
|
||||||
|
align_items="center",
|
||||||
|
gap="8px",
|
||||||
|
),
|
||||||
rx.vstack(
|
rx.vstack(
|
||||||
rx.hstack(
|
rx.text(
|
||||||
rx.image(
|
"产品汪的棱镜球,精准捕捉各类业务需求",
|
||||||
src="/logo.png",
|
|
||||||
width="64px",
|
|
||||||
height="64px",
|
|
||||||
),
|
|
||||||
rx.box(
|
|
||||||
"棱镜球",
|
|
||||||
font_size="calc(var(--prismui-font-size-3) * 2)",
|
|
||||||
font_weight="var(--prismui-font-weight-3)",
|
|
||||||
letter_spacing="1px",
|
|
||||||
),
|
|
||||||
display="flex",
|
|
||||||
align_items="center",
|
|
||||||
gap="8px",
|
|
||||||
),
|
|
||||||
rx.vstack(
|
|
||||||
rx.text(
|
|
||||||
"产品汪的棱镜球,精准捕捉各类业务需求",
|
|
||||||
width="100%",
|
|
||||||
text_align="center",
|
|
||||||
),
|
|
||||||
rx.text(
|
|
||||||
"可帮助检索业务资料,协助生成产品方案、输出 PRD 、流程图和原型图等",
|
|
||||||
width="100%",
|
|
||||||
text_align="center",
|
|
||||||
),
|
|
||||||
width="100%",
|
width="100%",
|
||||||
line_height="1.5",
|
text_align="center",
|
||||||
),
|
),
|
||||||
display="flex",
|
rx.text(
|
||||||
flex_direction="column",
|
"可帮助检索业务资料,协助生成产品方案、输出 PRD 、流程图和原型图等",
|
||||||
width="100%",
|
width="100%",
|
||||||
|
text_align="center",
|
||||||
|
),
|
||||||
|
line_height="1.5",
|
||||||
|
gap="4px",
|
||||||
|
),
|
||||||
|
align_items="center",
|
||||||
|
gap="12px",
|
||||||
|
),
|
||||||
|
# 演示案例
|
||||||
|
rx.vstack(
|
||||||
|
rx.text(
|
||||||
|
"演示案例",
|
||||||
|
margin_bottom="16px",
|
||||||
|
line_height="24px",
|
||||||
|
font_size="var(--prismui-font-size-3)",
|
||||||
|
font_weight="var(--prismui-font-weight-3)",
|
||||||
|
),
|
||||||
|
rx.hstack(
|
||||||
|
rx.box(
|
||||||
|
"预定航班",
|
||||||
|
on_click=lambda: ConversationState.init_work_flow(
|
||||||
|
WorkFlowType.BOOK_FLIGHT
|
||||||
|
),
|
||||||
|
padding="10px 16px",
|
||||||
|
background_color="var(--prismui-background-color-6)",
|
||||||
|
border_radius="var(--prismui-border-radius-9)",
|
||||||
|
color="var(--prismui-color-10)",
|
||||||
|
cursor="pointer",
|
||||||
|
),
|
||||||
|
flex_wrap="wrap",
|
||||||
align_items="center",
|
align_items="center",
|
||||||
gap="12px",
|
gap="12px",
|
||||||
),
|
),
|
||||||
# 演示案例
|
|
||||||
rx.vstack(
|
|
||||||
rx.hstack(
|
|
||||||
rx.text(
|
|
||||||
"演示案例",
|
|
||||||
line_height="24px",
|
|
||||||
font_size="var(--prismui-font-size-3)",
|
|
||||||
font_weight="var(--prismui-font-weight-3)",
|
|
||||||
),
|
|
||||||
display="flex",
|
|
||||||
justify_content="space-between",
|
|
||||||
align_items="center",
|
|
||||||
width="100%",
|
|
||||||
margin_bottom="16px",
|
|
||||||
),
|
|
||||||
rx.hstack(
|
|
||||||
rx.box(
|
|
||||||
"生成产品需求文档",
|
|
||||||
on_click=ConversationState.generate_prd,
|
|
||||||
padding="10px 16px",
|
|
||||||
background_color="var(--prismui-background-color-6)",
|
|
||||||
border_radius="var(--prismui-border-radius-9)",
|
|
||||||
color="var(--prismui-color-10)",
|
|
||||||
cursor="pointer",
|
|
||||||
),
|
|
||||||
display="flex",
|
|
||||||
flex_wrap="wrap",
|
|
||||||
align_items="center",
|
|
||||||
gap="12px",
|
|
||||||
),
|
|
||||||
width="100%",
|
|
||||||
padding="24px",
|
|
||||||
background_color="var(--prismui-background-color-1)",
|
|
||||||
border_radius="var(--prismui-border-radius-7)",
|
|
||||||
),
|
|
||||||
display="flex",
|
|
||||||
flex_direction="column",
|
|
||||||
align_items="center",
|
|
||||||
gap="24px",
|
|
||||||
width="100%",
|
width="100%",
|
||||||
min_height="0",
|
padding="24px",
|
||||||
margin="auto 0",
|
background_color="var(--prismui-background-color-1)",
|
||||||
color="var(--prismui-color-text)",
|
border_radius="var(--prismui-border-radius-7)",
|
||||||
),
|
),
|
||||||
display="flex",
|
|
||||||
flex="1",
|
flex="1",
|
||||||
flex_direction="column",
|
justify_content="center",
|
||||||
justify_content="flex-start",
|
align_items="center",
|
||||||
gap="24px",
|
gap="24px",
|
||||||
width="100%",
|
width="100%",
|
||||||
max_width="1200px",
|
max_width="1200px",
|
||||||
|
height="100%",
|
||||||
padding="0 12px",
|
padding="0 12px",
|
||||||
overflow="auto",
|
margin="auto",
|
||||||
style={
|
color="var(--prismui-color-text)",
|
||||||
"&::-webkit-scrollbar": {
|
|
||||||
"display": "none",
|
|
||||||
},
|
|
||||||
},
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -415,13 +391,29 @@ def thinking(item: MessageHistoryItem) -> rx.Component:
|
||||||
cursor="pointer",
|
cursor="pointer",
|
||||||
),
|
),
|
||||||
# 思考的内容
|
# 思考的内容
|
||||||
rx.text(
|
rx.box(
|
||||||
item.content,
|
rx.box(
|
||||||
line_height="1.6",
|
rx.cond(
|
||||||
font_size="var(--prismui-font-size-1)",
|
item.is_shown,
|
||||||
color="var(--prismui-color-3)",
|
rx.text(
|
||||||
|
item.content,
|
||||||
|
width="100%",
|
||||||
|
margin="4px 0",
|
||||||
|
line_height="1.6",
|
||||||
|
font_size="var(--prismui-font-size-1)",
|
||||||
|
color="var(--prismui-color-3)",
|
||||||
|
),
|
||||||
|
rx.fragment(),
|
||||||
|
),
|
||||||
|
min_height="0",
|
||||||
|
),
|
||||||
|
display="grid",
|
||||||
|
grid_template_rows=rx.cond(item.is_shown, "1fr", "0fr"),
|
||||||
|
width="100%",
|
||||||
|
opacity=rx.cond(item.is_shown, "1", "0"),
|
||||||
|
transition="grid-template-rows 0.18s ease-in-out, opacity 0.18s ease-in-out",
|
||||||
|
overflow="hidden",
|
||||||
),
|
),
|
||||||
margin_bottom="8px",
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -553,28 +545,11 @@ def message_history_item(item: MessageHistoryItem) -> rx.Component:
|
||||||
"""
|
"""
|
||||||
|
|
||||||
return rx.vstack(
|
return rx.vstack(
|
||||||
# 若等待流式输出则显示加载动效,否则根据消息类型渲染组件
|
rx.match(
|
||||||
rx.cond(
|
item.type,
|
||||||
ConversationState.awaiting_stream,
|
(MessageType.USER_PROMPT, user_prompt(item)),
|
||||||
rx.box(
|
(MessageType.THINKING, thinking(item)),
|
||||||
rx.box(
|
(MessageType.RESULT_OUTPUT, result_output(item)),
|
||||||
class_name="loadin",
|
|
||||||
flex_shrink="0",
|
|
||||||
),
|
|
||||||
display="flex",
|
|
||||||
justify_content="center",
|
|
||||||
align_items="center",
|
|
||||||
width="32px",
|
|
||||||
height="18px",
|
|
||||||
margin_bottom="8px",
|
|
||||||
overflow="hidden",
|
|
||||||
),
|
|
||||||
rx.match(
|
|
||||||
item.type,
|
|
||||||
(MessageType.USER_PROMPT, user_prompt(item)),
|
|
||||||
(MessageType.THINKING, thinking(item)),
|
|
||||||
(MessageType.RESULT_OUTPUT, result_output(item)),
|
|
||||||
),
|
|
||||||
),
|
),
|
||||||
padding="0 16px",
|
padding="0 16px",
|
||||||
width="100%",
|
width="100%",
|
||||||
|
|
@ -587,15 +562,36 @@ def message_history() -> rx.Component:
|
||||||
:return: Component
|
:return: Component
|
||||||
"""
|
"""
|
||||||
return rx.box(
|
return rx.box(
|
||||||
# 若消息历史为空则显示引导,否则渲染消息历史项
|
# 若消息历史为空则显示使用指引,否则渲染消息历史项
|
||||||
rx.cond(
|
rx.cond(
|
||||||
ConversationState.message_history.length() == 0,
|
ConversationState.message_history.length() == 0,
|
||||||
guidance(),
|
use_guidance(),
|
||||||
rx.vstack(
|
rx.vstack(
|
||||||
rx.auto_scroll(
|
rx.auto_scroll(
|
||||||
rx.foreach(
|
rx.vstack(
|
||||||
ConversationState.message_history,
|
rx.foreach(
|
||||||
message_history_item,
|
ConversationState.message_history,
|
||||||
|
message_history_item,
|
||||||
|
),
|
||||||
|
# 若等待流式输出则显示加载动效
|
||||||
|
rx.cond(
|
||||||
|
ConversationState.awaiting_stream,
|
||||||
|
rx.box(
|
||||||
|
rx.box(
|
||||||
|
class_name="awaiting_stream",
|
||||||
|
flex_shrink="0",
|
||||||
|
),
|
||||||
|
display="flex",
|
||||||
|
justify_content="flex‑start",
|
||||||
|
align_items="center",
|
||||||
|
width="100%",
|
||||||
|
padding="0 16px",
|
||||||
|
margin_bottom="8px",
|
||||||
|
overflow="hidden",
|
||||||
|
),
|
||||||
|
rx.fragment(),
|
||||||
|
),
|
||||||
|
width="100%",
|
||||||
),
|
),
|
||||||
width="100%",
|
width="100%",
|
||||||
padding="0 12px",
|
padding="0 12px",
|
||||||
|
|
@ -616,6 +612,11 @@ def message_history() -> rx.Component:
|
||||||
overflow_y="auto",
|
overflow_y="auto",
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
|
display="flex",
|
||||||
|
flex_direction="column",
|
||||||
|
width="100%",
|
||||||
|
height="100%",
|
||||||
|
min_height="0",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -136,9 +136,10 @@ class AuthState(rx.State):
|
||||||
发送验证码(后台任务)
|
发送验证码(后台任务)
|
||||||
:return: None
|
:return: None
|
||||||
"""
|
"""
|
||||||
db_state = await self.get_db_state()
|
|
||||||
|
|
||||||
async with self:
|
async with self:
|
||||||
|
# 获取当前数据库状态
|
||||||
|
db_state = await self.get_db_state()
|
||||||
# 创建验证记录
|
# 创建验证记录
|
||||||
captcha_code = await db_state.create_captcha_record(email=self.email)
|
captcha_code = await db_state.create_captcha_record(email=self.email)
|
||||||
|
|
||||||
|
|
@ -238,7 +239,7 @@ class AuthState(rx.State):
|
||||||
|
|
||||||
self.is_logging_in = True
|
self.is_logging_in = True
|
||||||
|
|
||||||
# 获取数据库状态
|
# 获取当前数据库状态
|
||||||
db_state = await self.get_db_state()
|
db_state = await self.get_db_state()
|
||||||
|
|
||||||
# 核验验证码
|
# 核验验证码
|
||||||
|
|
@ -259,7 +260,7 @@ class AuthState(rx.State):
|
||||||
self.user_id = user_id
|
self.user_id = user_id
|
||||||
|
|
||||||
self.email = ""
|
self.email = ""
|
||||||
self.captcha = ""
|
self.captcha_code = ""
|
||||||
self.is_captcha_code_sent = False
|
self.is_captcha_code_sent = False
|
||||||
self.resend_captcha_code_countdown = 0
|
self.resend_captcha_code_countdown = 0
|
||||||
self.is_policies_agreed = True
|
self.is_policies_agreed = True
|
||||||
|
|
@ -292,6 +293,6 @@ class AuthState(rx.State):
|
||||||
conversation_state = await self.get_state(ConversationState)
|
conversation_state = await self.get_state(ConversationState)
|
||||||
conversation_state.user_id = ""
|
conversation_state.user_id = ""
|
||||||
conversation_state.conversations = {}
|
conversation_state.conversations = {}
|
||||||
conversation_state.conversation_id = ""
|
conversation_state.actived_conversation_id = ""
|
||||||
conversation_state.is_conversation_history_shown = False
|
conversation_state.is_conversation_history_shown = False
|
||||||
self.user_id = ""
|
self.user_id = ""
|
||||||
|
|
|
||||||
|
|
@ -301,28 +301,30 @@ class ConversationState(rx.State):
|
||||||
index_map_to_message_id: dict[int, str] = {}
|
index_map_to_message_id: dict[int, str] = {}
|
||||||
# 初始化工具调用唯一标识集合
|
# 初始化工具调用唯一标识集合
|
||||||
tool_call_ids: set[str] = set()
|
tool_call_ids: set[str] = set()
|
||||||
|
|
||||||
|
# 匹配工作流类型
|
||||||
|
match conversation.work_flow:
|
||||||
|
# 预定航班
|
||||||
|
case WorkFlowType.BOOK_FLIGHT:
|
||||||
|
from application.workshop.book_flight import run_stream_events
|
||||||
|
|
||||||
|
# 非结构化对话
|
||||||
|
case _:
|
||||||
|
from application.workshop.unstructured_dialogue import (
|
||||||
|
run_stream_events,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 消息列表
|
||||||
|
messages: list[Message] = []
|
||||||
|
# 运行并流式输出事件
|
||||||
|
stream_events = run_stream_events(
|
||||||
|
user_prompt=user_prompt,
|
||||||
|
message_history=await db_state.get_message_history(
|
||||||
|
conversation_id=self.actived_conversation_id
|
||||||
|
),
|
||||||
|
usage=usage_validate_python(dict(conversation.usage)),
|
||||||
|
)
|
||||||
try:
|
try:
|
||||||
# 匹配工作流类型
|
|
||||||
match conversation.work_flow:
|
|
||||||
# 预定航班
|
|
||||||
case WorkFlowType.BOOK_FLIGHT:
|
|
||||||
from application.workshop.book_flight import run_stream_events
|
|
||||||
|
|
||||||
# 非结构化对话
|
|
||||||
case _:
|
|
||||||
from application.workshop.unstructured_dialogue import (
|
|
||||||
run_stream_events,
|
|
||||||
)
|
|
||||||
|
|
||||||
# 运行并流式输出事件
|
|
||||||
stream_events = run_stream_events(
|
|
||||||
user_prompt=user_prompt,
|
|
||||||
message_history=await db_state.get_message_history(
|
|
||||||
conversation_id=self.actived_conversation_id
|
|
||||||
),
|
|
||||||
usage=usage_validate_python(conversation.usage),
|
|
||||||
)
|
|
||||||
|
|
||||||
# 获取运行流式输出事件
|
# 获取运行流式输出事件
|
||||||
async for event in stream_events:
|
async for event in stream_events:
|
||||||
# 将等待流式输出设置为否
|
# 将等待流式输出设置为否
|
||||||
|
|
@ -395,7 +397,15 @@ class ConversationState(rx.State):
|
||||||
]
|
]
|
||||||
message.is_running = False
|
message.is_running = False
|
||||||
message.title = "思考完成"
|
message.title = "思考完成"
|
||||||
|
messages.append(message)
|
||||||
|
|
||||||
|
case TextPart(content=content):
|
||||||
|
# 获取消息实例
|
||||||
|
message = conversation.messages[
|
||||||
|
index_map_to_message_id[index]
|
||||||
|
]
|
||||||
message.content = content
|
message.content = content
|
||||||
|
messages.append(message)
|
||||||
|
|
||||||
# ========== 智能体运行结果事件 ==========
|
# ========== 智能体运行结果事件 ==========
|
||||||
case AgentRunResultEvent(result=result):
|
case AgentRunResultEvent(result=result):
|
||||||
|
|
@ -409,15 +419,9 @@ class ConversationState(rx.State):
|
||||||
new_messages=result.new_messages(),
|
new_messages=result.new_messages(),
|
||||||
)
|
)
|
||||||
|
|
||||||
if isinstance(message, Message):
|
# 强制更新会话并推送前端
|
||||||
# 强制更新会话并推送前端
|
self.conversations[self.actived_conversation_id] = conversation
|
||||||
self.conversations[self.actived_conversation_id] = conversation
|
yield
|
||||||
yield
|
|
||||||
# 创建对话记录
|
|
||||||
await db_state.create_message_record(
|
|
||||||
conversation_id=conversation.id,
|
|
||||||
message=message,
|
|
||||||
)
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
...
|
...
|
||||||
|
|
@ -426,10 +430,15 @@ class ConversationState(rx.State):
|
||||||
conversation.is_running = False
|
conversation.is_running = False
|
||||||
# 将等待流式输出设置为否
|
# 将等待流式输出设置为否
|
||||||
conversation.awaiting_stream = False
|
conversation.awaiting_stream = False
|
||||||
|
|
||||||
self.conversations[self.actived_conversation_id] = conversation
|
self.conversations[self.actived_conversation_id] = conversation
|
||||||
yield
|
yield
|
||||||
|
|
||||||
|
# 批量创建消息记录
|
||||||
|
await db_state.create_message_records(
|
||||||
|
conversation_id=conversation.id,
|
||||||
|
messages=messages,
|
||||||
|
)
|
||||||
|
|
||||||
@rx.event
|
@rx.event
|
||||||
async def init_work_flow(self, work_flow_type: WorkFlowType) -> None:
|
async def init_work_flow(self, work_flow_type: WorkFlowType) -> None:
|
||||||
"""
|
"""
|
||||||
|
|
@ -441,13 +450,12 @@ class ConversationState(rx.State):
|
||||||
return
|
return
|
||||||
|
|
||||||
match work_flow_type:
|
match work_flow_type:
|
||||||
# 预定航班
|
# 预定航班工作流
|
||||||
case WorkFlowType.BOOK_FLIGHT:
|
case WorkFlowType.BOOK_FLIGHT:
|
||||||
from application.workshop.book_flight import init_work_flow
|
from application.workshop.book_flight import init_work_flow
|
||||||
|
|
||||||
# 初始化预定航班任务
|
# 初始化预定航班任务
|
||||||
conversation.work_flow = init_work_flow()
|
conversation.work_flow = (work_flow := init_work_flow())
|
||||||
conversation.user_prompt = f"帮我找一班从 {conversation.work_flow.deps.origin} 到 {conversation.work_flow.deps.destination} 在 {conversation.work_flow.deps.date} 的航班"
|
|
||||||
|
|
||||||
@rx.event
|
@rx.event
|
||||||
def toggle_message_history_item_shown(self, message_id: str) -> None:
|
def toggle_message_history_item_shown(self, message_id: str) -> None:
|
||||||
|
|
|
||||||
|
|
@ -5,21 +5,21 @@
|
||||||
from datetime import datetime, timedelta
|
from datetime import datetime, timedelta
|
||||||
from random import choices
|
from random import choices
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from pydantic_ai import ModelMessage, ModelMessagesTypeAdapter, RunUsage
|
from pydantic_ai import ModelMessage, ModelMessagesTypeAdapter, RunUsage
|
||||||
from pydantic_ai._uuid import uuid7
|
from pydantic_ai._uuid import uuid7
|
||||||
import reflex as rx
|
import reflex as rx
|
||||||
|
|
||||||
from sqlmodel import Field, JSON, SQLModel, select, update
|
from sqlmodel import Field, JSON, SQLModel, select, update
|
||||||
|
|
||||||
from application.states.models import (
|
from application.states.models import (
|
||||||
Conversation,
|
Conversation,
|
||||||
MessageType,
|
|
||||||
Message,
|
Message,
|
||||||
|
MessageType,
|
||||||
usage_dump_python,
|
usage_dump_python,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class CaptchaRecord(SQLModel, table=True, table_name="captcha_code"):
|
class CaptchaRecord(SQLModel, table=True):
|
||||||
"""
|
"""
|
||||||
验证码记录
|
验证码记录
|
||||||
"""
|
"""
|
||||||
|
|
@ -47,7 +47,7 @@ class CaptchaRecord(SQLModel, table=True, table_name="captcha_code"):
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class UserRecord(SQLModel, table=True, table_name="user"):
|
class UserRecord(SQLModel, table=True):
|
||||||
"""
|
"""
|
||||||
用户记录
|
用户记录
|
||||||
"""
|
"""
|
||||||
|
|
@ -60,7 +60,7 @@ class UserRecord(SQLModel, table=True, table_name="user"):
|
||||||
email: str = Field(..., index=True, description="邮箱")
|
email: str = Field(..., index=True, description="邮箱")
|
||||||
|
|
||||||
|
|
||||||
class ConversationRecord(SQLModel, table=True, table_name="conversation"):
|
class ConversationRecord(SQLModel, table=True):
|
||||||
"""
|
"""
|
||||||
会话记录
|
会话记录
|
||||||
"""
|
"""
|
||||||
|
|
@ -85,7 +85,7 @@ class ConversationRecord(SQLModel, table=True, table_name="conversation"):
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class MessageRecord(SQLModel, table=True, table_name="message"):
|
class MessageRecord(SQLModel, table=True):
|
||||||
"""
|
"""
|
||||||
消息记录
|
消息记录
|
||||||
"""
|
"""
|
||||||
|
|
@ -101,7 +101,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"):
|
class RunRecord(SQLModel, table=True):
|
||||||
"""
|
"""
|
||||||
运行记录
|
运行记录
|
||||||
"""
|
"""
|
||||||
|
|
@ -310,6 +310,29 @@ class DatabaseState(rx.State):
|
||||||
await session.refresh(record)
|
await session.refresh(record)
|
||||||
return {message.id: message}
|
return {message.id: message}
|
||||||
|
|
||||||
|
async def create_message_records(
|
||||||
|
self,
|
||||||
|
conversation_id: str,
|
||||||
|
messages: list[Message],
|
||||||
|
) -> None:
|
||||||
|
"""
|
||||||
|
创建消息记录
|
||||||
|
:param conversation_id: 会话唯一标识
|
||||||
|
:param message: 消息实例
|
||||||
|
:return: 消息实例
|
||||||
|
"""
|
||||||
|
async with rx.asession() as session:
|
||||||
|
for message in messages:
|
||||||
|
record = MessageRecord(
|
||||||
|
id=message.id,
|
||||||
|
conversation_id=conversation_id,
|
||||||
|
type=message.type,
|
||||||
|
title=message.title,
|
||||||
|
content=message.content,
|
||||||
|
)
|
||||||
|
session.add(record)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
async def get_message_history(self, conversation_id: str) -> list[ModelMessage]:
|
async def get_message_history(self, conversation_id: str) -> list[ModelMessage]:
|
||||||
"""
|
"""
|
||||||
获取消息历史列表
|
获取消息历史列表
|
||||||
|
|
|
||||||
|
|
@ -2,9 +2,9 @@
|
||||||
"""
|
"""
|
||||||
面向 reflex.state 的类
|
面向 reflex.state 的类
|
||||||
"""
|
"""
|
||||||
|
from datetime import datetime
|
||||||
from enum import StrEnum
|
from enum import StrEnum
|
||||||
from typing import Any
|
from typing import Any
|
||||||
from datetime import datetime
|
|
||||||
|
|
||||||
from pydantic import BaseModel, Field, TypeAdapter
|
from pydantic import BaseModel, Field, TypeAdapter
|
||||||
from pydantic_ai import RunUsage
|
from pydantic_ai import RunUsage
|
||||||
|
|
@ -40,30 +40,6 @@ class Message(BaseModel):
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class RunStatus(StrEnum):
|
|
||||||
"""
|
|
||||||
运行状态枚举
|
|
||||||
"""
|
|
||||||
|
|
||||||
RUNNING = "running"
|
|
||||||
FINISHED = "finished"
|
|
||||||
|
|
||||||
|
|
||||||
class Run(BaseModel):
|
|
||||||
"""
|
|
||||||
运行类
|
|
||||||
"""
|
|
||||||
|
|
||||||
id: str = Field(default_factory=lambda: str(uuid7()), description="运行唯一标识")
|
|
||||||
messages: dict[str, Message] = Field(default_factory=dict, description="消息字典")
|
|
||||||
status: RunStatus = Field(default=RunStatus.RUNNING, description="运行状态")
|
|
||||||
usage: RunUsage = Field(default=RunUsage(), description="运行使用量")
|
|
||||||
usage_limits: UsageLimits | None = Field(default=None, description="运行使用量限制")
|
|
||||||
is_thinking: bool = Field(
|
|
||||||
default=False, description="正在思考,True 表示正在思考,False 表示未正在思考"
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class WorkFlowType(StrEnum):
|
class WorkFlowType(StrEnum):
|
||||||
"""
|
"""
|
||||||
工作流类型枚举
|
工作流类型枚举
|
||||||
|
|
@ -77,7 +53,7 @@ class Deps(BaseModel):
|
||||||
依赖项类
|
依赖项类
|
||||||
"""
|
"""
|
||||||
|
|
||||||
...
|
pass
|
||||||
|
|
||||||
|
|
||||||
class WorkFlow(BaseModel):
|
class WorkFlow(BaseModel):
|
||||||
|
|
@ -86,7 +62,7 @@ class WorkFlow(BaseModel):
|
||||||
"""
|
"""
|
||||||
|
|
||||||
type: WorkFlowType = Field(..., description="工作流类型")
|
type: WorkFlowType = Field(..., description="工作流类型")
|
||||||
deps: Deps | None = Field(default=None, description="工作流依赖项")
|
deps: Any = Field(default=None, description="工作流依赖项")
|
||||||
usage: RunUsage = Field(default=RunUsage(), description="工作流使用量")
|
usage: RunUsage = Field(default=RunUsage(), description="工作流使用量")
|
||||||
usage_limits: UsageLimits | None = Field(
|
usage_limits: UsageLimits | None = Field(
|
||||||
default=None,
|
default=None,
|
||||||
|
|
@ -109,7 +85,10 @@ class Conversation(BaseModel):
|
||||||
is_running: bool = Field(
|
is_running: bool = Field(
|
||||||
default=False, description="正在运行,True 表示正在运行, False 表示运行结束"
|
default=False, description="正在运行,True 表示正在运行, False 表示运行结束"
|
||||||
)
|
)
|
||||||
awaiting_stream: bool = Field(default=False, description="等待流式输出,True 表示等待流式输出,False 表示已开始流式输出或已完成")
|
awaiting_stream: bool = Field(
|
||||||
|
default=False,
|
||||||
|
description="等待流式输出,True 表示等待流式输出,False 表示已开始流式输出或已完成",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class ConversationHistoryItem(BaseModel):
|
class ConversationHistoryItem(BaseModel):
|
||||||
|
|
@ -139,28 +118,6 @@ class MessageHistoryItem(BaseModel):
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
# Dpes 适配器
|
|
||||||
DepsAdapter = TypeAdapter(Deps)
|
|
||||||
|
|
||||||
|
|
||||||
def deps_to_object(deps: dict[str, Any]) -> Deps | None:
|
|
||||||
"""
|
|
||||||
Deps 转为对象
|
|
||||||
"""
|
|
||||||
if not deps:
|
|
||||||
return None
|
|
||||||
return DepsAdapter.validate_python(deps)
|
|
||||||
|
|
||||||
|
|
||||||
def deps_to_dict(deps: Deps | None) -> dict[str, Any]:
|
|
||||||
"""
|
|
||||||
Deps 转为字典
|
|
||||||
"""
|
|
||||||
if not deps:
|
|
||||||
return {}
|
|
||||||
return DepsAdapter.dump_python(deps)
|
|
||||||
|
|
||||||
|
|
||||||
# Usage 适配器
|
# Usage 适配器
|
||||||
UsageAdapter = TypeAdapter(RunUsage)
|
UsageAdapter = TypeAdapter(RunUsage)
|
||||||
|
|
||||||
|
|
@ -185,21 +142,3 @@ def usage_dump_python(usage: RunUsage) -> dict[str, Any]:
|
||||||
|
|
||||||
# UsageLimits 适配器
|
# UsageLimits 适配器
|
||||||
UsageLimitsAdapter = TypeAdapter(UsageLimits)
|
UsageLimitsAdapter = TypeAdapter(UsageLimits)
|
||||||
|
|
||||||
|
|
||||||
def usage_limits_to_object(usage_limits: dict[str, Any]) -> UsageLimits | None:
|
|
||||||
"""
|
|
||||||
UsageLimits 转为对象
|
|
||||||
"""
|
|
||||||
if not usage_limits:
|
|
||||||
return None
|
|
||||||
return UsageLimitsAdapter.validate_python(usage_limits)
|
|
||||||
|
|
||||||
|
|
||||||
def usage_limits_to_dict(usage_limits: UsageLimits | None) -> dict[str, Any]:
|
|
||||||
"""
|
|
||||||
UsageLimits 转为字典
|
|
||||||
"""
|
|
||||||
if not usage_limits:
|
|
||||||
return {}
|
|
||||||
return UsageLimitsAdapter.dump_python(usage_limits)
|
|
||||||
|
|
|
||||||
|
|
@ -1,105 +1,48 @@
|
||||||
# -*- coding: utf-8 -*-
|
# -*- coding: utf-8 -*-
|
||||||
"""
|
"""
|
||||||
预定航班
|
预定航班(范式)
|
||||||
"""
|
"""
|
||||||
import datetime
|
import datetime
|
||||||
from typing import AsyncGenerator, Literal
|
from typing import AsyncGenerator, Literal
|
||||||
|
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
from pydantic_ai import Agent, ModelMessage, ModelRetry, RunContext, UsageLimits, RunUsage
|
from pydantic_ai import Agent, ModelMessage, ModelRetry, RunContext, UsageLimits, RunUsage, DeferredToolRequests
|
||||||
from pydantic_ai.run import AgentRunResultEvent
|
from pydantic_ai.run import AgentRunResultEvent
|
||||||
|
|
||||||
from application.states.models import (
|
from application.states.models import (
|
||||||
WorkFlow,
|
WorkFlow,
|
||||||
WorkFlowType,
|
WorkFlowType,
|
||||||
usage_limits_to_object,
|
|
||||||
usage_to_dict,
|
|
||||||
usage_to_object,
|
|
||||||
Deps,
|
|
||||||
)
|
)
|
||||||
from application.workshop.models import DEEPSEEK_V4_FLASH_MODEL, MODEL_SETTINGS
|
from application.workshop.models import DEEPSEEK_V4_FLASH_MODEL, MODEL_SETTINGS_DISABLED_THINKING
|
||||||
|
|
||||||
|
|
||||||
class Deps_(BaseModel):
|
class Dependences(BaseModel):
|
||||||
"""
|
"""
|
||||||
依赖项类
|
依赖类
|
||||||
"""
|
"""
|
||||||
|
|
||||||
flight_info: str = Field(..., description="所有航班信息")
|
|
||||||
date: datetime.date = Field(..., description="航班日期")
|
date: datetime.date = Field(..., description="航班日期")
|
||||||
origin: str = Field(..., description="出发机场")
|
origin: str = Field(..., description="航班出发机场")
|
||||||
destination: str = Field(..., description="到达机场")
|
destination: str = Field(..., description="航班到达机场")
|
||||||
|
source_material: str = Field(..., description="航班来源资料")
|
||||||
|
|
||||||
|
class Flight(BaseModel):
|
||||||
class FlightDetail(BaseModel):
|
|
||||||
"""
|
"""
|
||||||
航班详情类
|
航班类
|
||||||
"""
|
"""
|
||||||
|
|
||||||
number: str = Field(description="航班号")
|
number: str = Field(..., description="航班号")
|
||||||
date: datetime.date = Field(description="航班日期")
|
date: datetime.date = Field(..., description="航班日期")
|
||||||
origin: str = Field(description="出发机场")
|
origin: str = Field(..., description="航班出发机场")
|
||||||
destination: str = Field(description="到达机场")
|
destination: str = Field(..., description="航班到达机场")
|
||||||
price: int = Field(description="机票价格")
|
airfare: int = Field(..., description="航班机票价格")
|
||||||
|
|
||||||
|
|
||||||
class NoFlightFound(BaseModel):
|
class FlightNoFound(BaseModel):
|
||||||
"""
|
"""
|
||||||
未查询到航班类
|
未查询到航班类
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
# 航班查询智能体
|
|
||||||
flight_search_agent = Agent[Deps, FlightDetail | NoFlightFound](
|
|
||||||
model=DEEPSEEK_V4_FLASH_MODEL,
|
|
||||||
model_settings=MODEL_SETTINGS,
|
|
||||||
deps_type=Deps,
|
|
||||||
output_type=FlightDetail | NoFlightFound,
|
|
||||||
system_prompt="你的工作是在给定日期、出发机场和到达机场为用户找到最便宜的航班",
|
|
||||||
)
|
|
||||||
|
|
||||||
# 所有航班详情提取智能体
|
|
||||||
flight_details_extraction_agent = Agent(
|
|
||||||
model=DEEPSEEK_V4_FLASH_MODEL,
|
|
||||||
model_settings=MODEL_SETTINGS,
|
|
||||||
output_type=list[FlightDetail],
|
|
||||||
system_prompt="从给定文本中提取所有航班详细,包括航班号、航班日期、出发机场、到达机场和机票价格",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@flight_search_agent.tool
|
|
||||||
async def extract_flight_details(ctx: RunContext[Deps]) -> list[FlightDetail]:
|
|
||||||
"""
|
|
||||||
工具:提取所有航班详情
|
|
||||||
"""
|
|
||||||
result = await flight_details_extraction_agent.run(
|
|
||||||
ctx.deps.flight_info, usage=ctx.usage
|
|
||||||
)
|
|
||||||
return result.output
|
|
||||||
|
|
||||||
|
|
||||||
@flight_search_agent.output_validator
|
|
||||||
async def validate_output(
|
|
||||||
ctx: RunContext[Deps], output: FlightDetail | NoFlightFound
|
|
||||||
) -> FlightDetail | NoFlightFound:
|
|
||||||
"""
|
|
||||||
输出校验:航班详情
|
|
||||||
"""
|
|
||||||
if isinstance(output, NoFlightFound):
|
|
||||||
return output
|
|
||||||
|
|
||||||
errors = []
|
|
||||||
if output.date != ctx.deps.date:
|
|
||||||
errors.append(f"航班日期应为 {ctx.deps.date}, 不是 {output.date}")
|
|
||||||
if output.origin != ctx.deps.origin:
|
|
||||||
errors.append(f"出发机场应为 {ctx.deps.origin}, 不是 {output.origin}")
|
|
||||||
if output.destination != ctx.deps.destination:
|
|
||||||
errors.append(f"到达机场应为 {ctx.deps.destination}, 不是 {output.destination}")
|
|
||||||
if errors:
|
|
||||||
raise ModelRetry("\n".join(errors))
|
|
||||||
return output
|
|
||||||
|
|
||||||
|
|
||||||
class SeatPreference(BaseModel):
|
class SeatPreference(BaseModel):
|
||||||
"""
|
"""
|
||||||
座位偏好类
|
座位偏好类
|
||||||
|
|
@ -109,21 +52,77 @@ class SeatPreference(BaseModel):
|
||||||
column: Literal["A", "B", "C", "D", "E", "F"] = Field(description="座位列")
|
column: Literal["A", "B", "C", "D", "E", "F"] = Field(description="座位列")
|
||||||
|
|
||||||
|
|
||||||
class NoSeatExtracted(BaseModel):
|
class SeatPreferenceNoExtracted(BaseModel):
|
||||||
"""
|
"""
|
||||||
未提取到座位偏好类
|
未提取到座位偏好类
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
# 提取座位偏好智能体(无依赖项)
|
# 航班查询智能体
|
||||||
seat_preference_extraction_agent = Agent[object, SeatPreference | NoSeatExtracted](
|
flight_inquiry_agent = Agent[Dependences, Flight | FlightNoFound | DeferredToolRequests](
|
||||||
model=DEEPSEEK_V4_FLASH_MODEL,
|
model=DEEPSEEK_V4_FLASH_MODEL,
|
||||||
model_settings=MODEL_SETTINGS,
|
model_settings=MODEL_SETTINGS_DISABLED_THINKING,
|
||||||
output_type=SeatPreference | NoSeatExtracted,
|
deps_type=Dependences,
|
||||||
system_prompt="提取用户的座位偏好。座位规则说明:A 座、F 座为靠窗座位;第 1 排是前排座位,腿部空间更大;14 排、20 排同样拥有加宽腿部空间",
|
output_type=Flight | FlightNoFound | DeferredToolRequests,
|
||||||
|
system_prompt="你的任务是根据给定的日期、出发机场和到达机场帮助用户找到最便宜的航班。",
|
||||||
)
|
)
|
||||||
|
|
||||||
flight_info = """
|
# 航班信息提取智能体
|
||||||
|
flights_extraction_agent = Agent(
|
||||||
|
model=DEEPSEEK_V4_FLASH_MODEL,
|
||||||
|
model_settings=MODEL_SETTINGS_DISABLED_THINKING,
|
||||||
|
output_type=list[Flight],
|
||||||
|
system_prompt="你的任务是根据给定的航班来源资料提取出所有航班信息,包括航班号、航班日期、航班出发机场、航班到达机场和航班机票价格。禁止编造。")
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
@flight_inquiry_agent.tool
|
||||||
|
async def extract_flight_details(ctx: RunContext[Dependences]) -> list[Flight]:
|
||||||
|
"""
|
||||||
|
提取所有航班信息
|
||||||
|
"""
|
||||||
|
result = await flights_extraction_agent.run(
|
||||||
|
ctx.deps.source_material, usage=ctx.usage
|
||||||
|
)
|
||||||
|
return result.output
|
||||||
|
|
||||||
|
|
||||||
|
@flight_inquiry_agent.output_validator
|
||||||
|
async def validate_output(
|
||||||
|
ctx: RunContext[Dependences], output: Flight | FlightNoFound | DeferredToolRequests
|
||||||
|
) -> Flight | FlightNoFound | DeferredToolRequests:
|
||||||
|
"""
|
||||||
|
输出校验:航班信息
|
||||||
|
"""
|
||||||
|
# 不校验未查询到航班
|
||||||
|
if isinstance(output, FlightNoFound):
|
||||||
|
return output
|
||||||
|
|
||||||
|
# 不校验延迟工具请求
|
||||||
|
if isinstance(output, DeferredToolRequests):
|
||||||
|
return output
|
||||||
|
|
||||||
|
errors = []
|
||||||
|
if output.date != ctx.deps.date:
|
||||||
|
errors.append(f"航班日期应为 {ctx.deps.date}, 不是 {output.date}")
|
||||||
|
if output.origin != ctx.deps.origin:
|
||||||
|
errors.append(f"航班出发机场应为 {ctx.deps.origin}, 不是 {output.origin}")
|
||||||
|
if output.destination != ctx.deps.destination:
|
||||||
|
errors.append(f"航班到达机场应为 {ctx.deps.destination}, 不是 {output.destination}")
|
||||||
|
if errors:
|
||||||
|
raise ModelRetry("\n".join(errors))
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
# 座位偏好提取智能体
|
||||||
|
seat_preference_extraction_agent = Agent[object, SeatPreference | SeatPreferenceNoExtracted](
|
||||||
|
model=DEEPSEEK_V4_FLASH_MODEL,
|
||||||
|
model_settings=MODEL_SETTINGS_DISABLED_THINKING,
|
||||||
|
output_type=SeatPreference | SeatPreferenceNoExtracted,
|
||||||
|
system_prompt="你的任务是根据用户回答提取座位偏好。座位规则说明:A 座、F 座为靠窗座位;第 1 排是前排座位,腿部空间更大;14 排、20 排同样拥有加宽腿部空间",
|
||||||
|
)
|
||||||
|
|
||||||
|
source_material = """
|
||||||
1. Flight SFO-AK123
|
1. Flight SFO-AK123
|
||||||
- Price: $350
|
- Price: $350
|
||||||
- Origin: San Francisco International Airport (SFO)
|
- Origin: San Francisco International Airport (SFO)
|
||||||
|
|
@ -170,7 +169,7 @@ flight_info = """
|
||||||
def init_work_flow() -> WorkFlow:
|
def init_work_flow() -> WorkFlow:
|
||||||
return WorkFlow(
|
return WorkFlow(
|
||||||
type=WorkFlowType.BOOK_FLIGHT,
|
type=WorkFlowType.BOOK_FLIGHT,
|
||||||
deps=Deps(
|
deps=Deps_(
|
||||||
flight_info=flight_info,
|
flight_info=flight_info,
|
||||||
date=datetime.date(2025, 1, 10),
|
date=datetime.date(2025, 1, 10),
|
||||||
origin="SFO",
|
origin="SFO",
|
||||||
|
|
@ -197,8 +196,7 @@ async def run_stream_events(
|
||||||
user_prompt=user_prompt,
|
user_prompt=user_prompt,
|
||||||
deps=task.deps,
|
deps=task.deps,
|
||||||
message_history=message_history,
|
message_history=message_history,
|
||||||
usage=usage_to_object(dict(task.usage)),
|
usage=usage,
|
||||||
usage_limits=usage_limits_to_object(dict(task.usage_limits)),
|
|
||||||
) as events:
|
) as events:
|
||||||
async for event in events:
|
async for event in events:
|
||||||
if not isinstance(event, AgentRunResultEvent):
|
if not isinstance(event, AgentRunResultEvent):
|
||||||
|
|
|
||||||
|
|
@ -14,4 +14,4 @@ DEEPSEEK_V4_FLASH_MODEL = OpenAIChatModel(
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
MODEL_SETTINGS = ModelSettings(extra_body={"thinking": {"type": "disabled"}})
|
MODEL_SETTINGS_DISABLED_THINKING = ModelSettings(extra_body={"thinking": {"type": "disabled"}}) # 禁用思考模式
|
||||||
|
|
|
||||||
|
|
@ -61,13 +61,13 @@ pre, pre code {
|
||||||
font-family: Consolas, "Microsoft YaHei", sans-serif !important;
|
font-family: Consolas, "Microsoft YaHei", sans-serif !important;
|
||||||
}
|
}
|
||||||
|
|
||||||
.loadin {
|
.awaiting_stream {
|
||||||
width: 6px;
|
width: 6px;
|
||||||
aspect-ratio: 1;
|
aspect-ratio: 1;
|
||||||
border-radius: 50%;
|
border-radius: 50%;
|
||||||
animation: loadin 0.9s infinite ease-in-out;
|
animation: awaiting_stream 0.9s infinite ease-in-out;
|
||||||
}
|
}
|
||||||
@keyframes loadin {
|
@keyframes awaiting_stream {
|
||||||
0% {
|
0% {
|
||||||
box-shadow: 12px 0 var(--prismui-color-5), -12px 0 var(--prismui-color-9);
|
box-shadow: 12px 0 var(--prismui-color-5), -12px 0 var(--prismui-color-9);
|
||||||
background: var(--prismui-color-5);
|
background: var(--prismui-color-5);
|
||||||
|
|
|
||||||
Binary file not shown.
|
|
@ -446,7 +446,7 @@
|
||||||
|
|
||||||
"@types/unist": ["@types/unist@3.0.3", "https://registry.npmmirror.com/@types/unist/-/unist-3.0.3.tgz", {}, "sha512-ko/gIFJRv177XgZsZcBwnqJN5x/Gien8qNOn0D5bQU/zAzVf9Zt3BlcUiLqhV9y4ARk0GbT3tnUiPNgnTXzc/Q=="],
|
"@types/unist": ["@types/unist@3.0.3", "https://registry.npmmirror.com/@types/unist/-/unist-3.0.3.tgz", {}, "sha512-ko/gIFJRv177XgZsZcBwnqJN5x/Gien8qNOn0D5bQU/zAzVf9Zt3BlcUiLqhV9y4ARk0GbT3tnUiPNgnTXzc/Q=="],
|
||||||
|
|
||||||
"@ungap/structured-clone": ["@ungap/structured-clone@1.3.2", "https://registry.npmmirror.com/@ungap/structured-clone/-/structured-clone-1.3.2.tgz", {}, "sha512-5jsZFwgR5rTdKwidH9Qmat75RKwqfpKlWWB1frDkljN127mwqBu8K0PYo7/hFpF03IEJpfVPpCQDY/eDx3iHvA=="],
|
"@ungap/structured-clone": ["@ungap/structured-clone@1.3.3", "https://registry.npmmirror.com/@ungap/structured-clone/-/structured-clone-1.3.3.tgz", {}, "sha512-60YRaenCQcVjYEKOcG824+DRGGIQ3VKErcBoAEDJZz5bKIs2ZG+X/H9Nk+Q6EVkwJk5QNApxbrc5QtBSwtrXAg=="],
|
||||||
|
|
||||||
"arg": ["arg@5.0.2", "https://registry.npmmirror.com/arg/-/arg-5.0.2.tgz", {}, "sha512-PYjyFOLKQ9y57JvQ6QLo8dAgNqswh8M1RMJYdQduT6xbWSgK36P/Z/v+p888pM69jMMfS8Xd8F6I1kQ/I9HUGg=="],
|
"arg": ["arg@5.0.2", "https://registry.npmmirror.com/arg/-/arg-5.0.2.tgz", {}, "sha512-PYjyFOLKQ9y57JvQ6QLo8dAgNqswh8M1RMJYdQduT6xbWSgK36P/Z/v+p888pM69jMMfS8Xd8F6I1kQ/I9HUGg=="],
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue