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