初始化demo
This commit is contained in:
@@ -0,0 +1,91 @@
|
||||
import asyncio
|
||||
from dataclasses import dataclass
|
||||
from typing import AsyncIterator, Callable, Optional
|
||||
|
||||
from .agent_factory import create_agent
|
||||
from .persistence import Session, SessionStore, get_session_store
|
||||
|
||||
|
||||
@dataclass
|
||||
class StreamEvent:
|
||||
"""Represents a streaming event from the agent."""
|
||||
|
||||
kind: str
|
||||
text: str = ""
|
||||
tool_name: str = ""
|
||||
|
||||
|
||||
async def stream_turn(agent, user_input: str) -> AsyncIterator[StreamEvent]:
|
||||
"""Streams the agent's response to a user input in real-time.
|
||||
|
||||
Yields token/tool events as they arrive, then a final "done" event
|
||||
"""
|
||||
last_tool_id = None
|
||||
|
||||
async for event in agent.stream_async(user_input):
|
||||
if event.get("data"):
|
||||
yield StreamEvent(kind="token", text=event["data"])
|
||||
elif event.get("current_tool_use"):
|
||||
tu = event["current_tool_use"]
|
||||
tid = tu.get("toolUsedId") if isinstance(tu, dict) else None
|
||||
# De-dupe: Strands re-emits current_tool_use on every token while the
|
||||
# tool call is being assembled; surface one frame per distinct call.
|
||||
|
||||
if tid and tid != last_tool_id:
|
||||
last_tool_id = tid
|
||||
yield StreamEvent(kind="tool", tool_name=tu.get("toolName", ""))
|
||||
yield StreamEvent(kind="done")
|
||||
|
||||
|
||||
async def run_turn(
|
||||
user_id: str,
|
||||
session_id_str,
|
||||
user_input: str,
|
||||
*,
|
||||
store: Optional[SessionStore] = None,
|
||||
on_event: Optional[Callable[[StreamEvent], None]] = None,
|
||||
system_prompt: Optional[str] = None,
|
||||
) -> str:
|
||||
"""Runs one full turn end-to-end with history persistence.
|
||||
|
||||
1. Load prior history for (user_id, session_id).
|
||||
2. Build a fresh agent seeded with that history.
|
||||
3. Stream the turn, forwarding events to "on_event" if provided.
|
||||
4. Persist the agent's updated messages back to the store.
|
||||
|
||||
Returns the full assistant reply text.
|
||||
"""
|
||||
store = store or get_session_store()
|
||||
|
||||
# 1. Load prior history
|
||||
session = store.load(user_id, session_id)
|
||||
|
||||
# 2. SEED a fresh agent with the prior messages
|
||||
agent = create_agent(
|
||||
system_prompt=system_prompt,
|
||||
messages=session.messages
|
||||
)
|
||||
|
||||
# 3. Stream the turn
|
||||
reply_parts: list[str] = []
|
||||
async for event in stream_turn(agent, user_input):
|
||||
if event.kind == "token":
|
||||
reply_parts.append(event.text)
|
||||
if on_event is not None:
|
||||
on_event(event)
|
||||
|
||||
# 4. PERSIST updated history. Strands has appended this turn's user message,
|
||||
# any tool-use/tool-result messages, and the assistant reply to
|
||||
# agent.messages - persisting that list is the whole history contract.
|
||||
store.save(Session(session_id=session_id, user_id=user_id, messages=list(agent.messages)))
|
||||
return "".join(reply_parts)
|
||||
|
||||
|
||||
def run_turn_sync(
|
||||
user_id: str,
|
||||
session_id: str,
|
||||
user_input: str,
|
||||
**kwargs
|
||||
) -> str:
|
||||
"""Blocking convenience wrapper around :func: run_turn."""
|
||||
return asyncio.run(run_turn(user_id, session_id, user_input, **kwargs))
|
||||
Reference in New Issue
Block a user