This commit is contained in:
parent
0280261508
commit
026c573556
|
|
@ -57,7 +57,9 @@ def usage_to_object(usage: dict[str, Any]) -> RunUsage:
|
||||||
"""
|
"""
|
||||||
Usage 转为对象
|
Usage 转为对象
|
||||||
"""
|
"""
|
||||||
return UsageAdapter.validate_python(usage) if usage else RunUsage()
|
if isinstance(usage, dict) and usage:
|
||||||
|
return UsageAdapter.validate_python(usage)
|
||||||
|
return RunUsage()
|
||||||
|
|
||||||
|
|
||||||
def usage_to_dict(usage: RunUsage) -> dict[str, Any]:
|
def usage_to_dict(usage: RunUsage) -> dict[str, Any]:
|
||||||
|
|
@ -75,11 +77,9 @@ def usage_limits_to_object(usage_limits: dict[str, Any]) -> UsageLimits:
|
||||||
"""
|
"""
|
||||||
UsageLimits 转为对象
|
UsageLimits 转为对象
|
||||||
"""
|
"""
|
||||||
return (
|
if isinstance(usage_limits, dict) and usage_limits:
|
||||||
UsageLimitsAdapter.validate_python(usage_limits)
|
return UsageLimitsAdapter.validate_python(usage_limits)
|
||||||
if usage_limits
|
return UsageLimits(request_limit=5)
|
||||||
else UsageLimits()
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def usage_limits_to_dict(usage_limits: UsageLimits) -> dict[str, Any]:
|
def usage_limits_to_dict(usage_limits: UsageLimits) -> dict[str, Any]:
|
||||||
|
|
@ -104,7 +104,7 @@ class Task(BaseModel):
|
||||||
|
|
||||||
id: str = Field(default_factory=lambda: str(uuid7()), description="任务唯一标识")
|
id: str = Field(default_factory=lambda: str(uuid7()), description="任务唯一标识")
|
||||||
type: TaskType = Field(..., description="任务类型")
|
type: TaskType = Field(..., description="任务类型")
|
||||||
node: str | None = Field(default="", description="任务节点")
|
node: str = Field(default="", description="任务节点")
|
||||||
deps: Any = Field(default=None, description="任务依赖项")
|
deps: Any = Field(default=None, description="任务依赖项")
|
||||||
usage: dict[str, Any] = Field(default_factory=dict, description="任务使用量")
|
usage: dict[str, Any] = Field(default_factory=dict, description="任务使用量")
|
||||||
usage_limits: dict[str, Any] = Field(
|
usage_limits: dict[str, Any] = Field(
|
||||||
|
|
|
||||||
|
|
@ -13,12 +13,11 @@ from application.states.models import (
|
||||||
Task,
|
Task,
|
||||||
TaskNodeResultEvent,
|
TaskNodeResultEvent,
|
||||||
TaskType,
|
TaskType,
|
||||||
usage_limits_to_dict,
|
|
||||||
usage_limits_to_object,
|
usage_limits_to_object,
|
||||||
usage_to_dict,
|
usage_to_dict,
|
||||||
usage_to_object,
|
usage_to_object,
|
||||||
)
|
)
|
||||||
from application.tasks.models import DEEPSEEK_V4_FLASH_MODEL
|
from application.tasks.models import DEEPSEEK_V4_FLASH_MODEL, MODEL_SETTINGS
|
||||||
|
|
||||||
|
|
||||||
class Deps(BaseModel):
|
class Deps(BaseModel):
|
||||||
|
|
@ -53,6 +52,7 @@ class NoFlightFound(BaseModel):
|
||||||
# 航班查询智能体
|
# 航班查询智能体
|
||||||
flight_search_agent = Agent[Deps, FlightDetail | NoFlightFound](
|
flight_search_agent = Agent[Deps, FlightDetail | NoFlightFound](
|
||||||
model=DEEPSEEK_V4_FLASH_MODEL,
|
model=DEEPSEEK_V4_FLASH_MODEL,
|
||||||
|
model_settings=MODEL_SETTINGS,
|
||||||
deps_type=Deps,
|
deps_type=Deps,
|
||||||
output_type=FlightDetail | NoFlightFound,
|
output_type=FlightDetail | NoFlightFound,
|
||||||
system_prompt="你的工作是在给定日期、出发机场和到达机场为用户找到最便宜的航班",
|
system_prompt="你的工作是在给定日期、出发机场和到达机场为用户找到最便宜的航班",
|
||||||
|
|
@ -61,6 +61,7 @@ flight_search_agent = Agent[Deps, FlightDetail | NoFlightFound](
|
||||||
# 所有航班详情提取智能体
|
# 所有航班详情提取智能体
|
||||||
flight_details_extraction_agent = Agent(
|
flight_details_extraction_agent = Agent(
|
||||||
model=DEEPSEEK_V4_FLASH_MODEL,
|
model=DEEPSEEK_V4_FLASH_MODEL,
|
||||||
|
model_settings=MODEL_SETTINGS,
|
||||||
output_type=list[FlightDetail],
|
output_type=list[FlightDetail],
|
||||||
system_prompt="从给定文本中提取所有航班详细,包括航班号、航班日期、出发机场、到达机场和机票价格",
|
system_prompt="从给定文本中提取所有航班详细,包括航班号、航班日期、出发机场、到达机场和机票价格",
|
||||||
)
|
)
|
||||||
|
|
@ -117,6 +118,7 @@ class NoSeatExtracted(BaseModel):
|
||||||
# 提取座位偏好智能体(无依赖项)
|
# 提取座位偏好智能体(无依赖项)
|
||||||
seat_preference_extraction_agent = Agent[object, SeatPreference | NoSeatExtracted](
|
seat_preference_extraction_agent = Agent[object, SeatPreference | NoSeatExtracted](
|
||||||
model=DEEPSEEK_V4_FLASH_MODEL,
|
model=DEEPSEEK_V4_FLASH_MODEL,
|
||||||
|
model_settings=MODEL_SETTINGS,
|
||||||
output_type=SeatPreference | NoSeatExtracted,
|
output_type=SeatPreference | NoSeatExtracted,
|
||||||
system_prompt="提取用户的座位偏好。座位规则说明:A 座、F 座为靠窗座位;第 1 排是前排座位,腿部空间更大;14 排、20 排同样拥有加宽腿部空间",
|
system_prompt="提取用户的座位偏好。座位规则说明:A 座、F 座为靠窗座位;第 1 排是前排座位,腿部空间更大;14 排、20 排同样拥有加宽腿部空间",
|
||||||
)
|
)
|
||||||
|
|
@ -175,7 +177,6 @@ def init_task() -> Task:
|
||||||
origin="SFO",
|
origin="SFO",
|
||||||
destination="ANC",
|
destination="ANC",
|
||||||
),
|
),
|
||||||
usage_limits=usage_limits_to_dict(UsageLimits(request_limit=5)),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -196,8 +197,8 @@ 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(task.usage),
|
usage=usage_to_object(dict(task.usage)),
|
||||||
usage_limits=usage_limits_to_object(task.usage_limits),
|
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):
|
||||||
|
|
@ -241,8 +242,8 @@ async def run_stream_events(
|
||||||
async with seat_preference_extraction_agent.run_stream_events(
|
async with seat_preference_extraction_agent.run_stream_events(
|
||||||
user_prompt=user_prompt,
|
user_prompt=user_prompt,
|
||||||
message_history=message_history,
|
message_history=message_history,
|
||||||
usage=usage_to_object(task.usage),
|
usage=usage_to_object(dict(task.usage)),
|
||||||
usage_limits=usage_limits_to_object(task.usage_limits),
|
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):
|
||||||
|
|
|
||||||
|
|
@ -2,10 +2,10 @@
|
||||||
"""
|
"""
|
||||||
智能体相关模块
|
智能体相关模块
|
||||||
"""
|
"""
|
||||||
from pydantic_ai.models.openai import OpenAIChatModel
|
from pydantic_ai import ModelSettings
|
||||||
|
from pydantic_ai.models.openai import OpenAIChatModel, OpenAIChatModelSettings
|
||||||
from pydantic_ai.providers.openai import OpenAIProvider
|
from pydantic_ai.providers.openai import OpenAIProvider
|
||||||
|
|
||||||
|
|
||||||
DEEPSEEK_V4_FLASH_MODEL = OpenAIChatModel(
|
DEEPSEEK_V4_FLASH_MODEL = OpenAIChatModel(
|
||||||
model_name="deepseek-v4-flash",
|
model_name="deepseek-v4-flash",
|
||||||
provider=OpenAIProvider(
|
provider=OpenAIProvider(
|
||||||
|
|
@ -13,3 +13,5 @@ DEEPSEEK_V4_FLASH_MODEL = OpenAIChatModel(
|
||||||
api_key="sk-D9Y1mCe8VlvNqLuSC4mAjqEwxJ2nW4C0h8a7EPn8kg9RLsHq",
|
api_key="sk-D9Y1mCe8VlvNqLuSC4mAjqEwxJ2nW4C0h8a7EPn8kg9RLsHq",
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
MODEL_SETTINGS = ModelSettings(extra_body={"thinking": {"type": "disabled"}})
|
||||||
|
|
|
||||||
Binary file not shown.
Loading…
Reference in New Issue