mirror of
https://github.com/NCBM/plyngent.git
synced 2026-07-23 14:14:57 +08:00
core/agent: streaming loop with text deltas and tool-call reconstruction
This commit is contained in:
+81
-17
@@ -1,6 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING, cast
|
||||||
|
|
||||||
from msgspec import UNSET
|
from msgspec import UNSET
|
||||||
|
|
||||||
@@ -11,10 +11,12 @@ from plyngent.lmproto.openai_compatible.model import (
|
|||||||
ChatCompletionsParam,
|
ChatCompletionsParam,
|
||||||
ToolChatMessage,
|
ToolChatMessage,
|
||||||
)
|
)
|
||||||
|
from plyngent.typedef import Unset # noqa: TC001
|
||||||
|
|
||||||
from .events import (
|
from .events import (
|
||||||
AgentEvent,
|
AgentEvent,
|
||||||
AssistantMessageEvent,
|
AssistantMessageEvent,
|
||||||
|
ErrorEvent,
|
||||||
MaxRoundsEvent,
|
MaxRoundsEvent,
|
||||||
TextDeltaEvent,
|
TextDeltaEvent,
|
||||||
ToolCallEvent,
|
ToolCallEvent,
|
||||||
@@ -43,7 +45,11 @@ async def _execute_tool_calls(
|
|||||||
for call in tool_calls:
|
for call in tool_calls:
|
||||||
yield ToolCallEvent(tool_call=call)
|
yield ToolCallEvent(tool_call=call)
|
||||||
if isinstance(call, AssistantFunctionToolCall):
|
if isinstance(call, AssistantFunctionToolCall):
|
||||||
result_text = await tools.execute(call.function.name, call.function.arguments)
|
try:
|
||||||
|
result_text = await tools.execute(call.function.name, call.function.arguments)
|
||||||
|
except Exception as exc: # noqa: BLE001
|
||||||
|
result_text = f"error: tool {call.function.name!r} failed: {exc}"
|
||||||
|
yield ErrorEvent(message=result_text)
|
||||||
tool_msg = ToolChatMessage(content=result_text, tool_call_id=call.id)
|
tool_msg = ToolChatMessage(content=result_text, tool_call_id=call.id)
|
||||||
else:
|
else:
|
||||||
tool_msg = ToolChatMessage(
|
tool_msg = ToolChatMessage(
|
||||||
@@ -54,7 +60,61 @@ async def _execute_tool_calls(
|
|||||||
yield ToolResultEvent(message=tool_msg)
|
yield ToolResultEvent(message=tool_msg)
|
||||||
|
|
||||||
|
|
||||||
async def run_chat_loop( # noqa: PLR0913
|
async def _stream_and_build_assistant( # noqa: C901
|
||||||
|
client: ChatClient,
|
||||||
|
param: ChatCompletionsParam,
|
||||||
|
) -> tuple[AssistantChatMessage, list[TextDeltaEvent]]:
|
||||||
|
"""Stream a round, return the assistant message and any text deltas."""
|
||||||
|
raw_lines_attr = getattr(client, "chat_completions_raw_lines", None)
|
||||||
|
if raw_lines_attr is None:
|
||||||
|
response = await client.chat_completions(param, stream=False)
|
||||||
|
if not response.choices:
|
||||||
|
msg = "chat completion response contained no choices"
|
||||||
|
raise RuntimeError(msg)
|
||||||
|
assistant = response.choices[0].message
|
||||||
|
text_events = []
|
||||||
|
if isinstance(assistant.content, str) and assistant.content:
|
||||||
|
text_events.append(TextDeltaEvent(content=assistant.content))
|
||||||
|
return assistant, text_events
|
||||||
|
|
||||||
|
stream_iter = await raw_lines_attr(param) # type: ignore[misc]
|
||||||
|
raw_lines: list[bytes] = []
|
||||||
|
content_parts: list[str] = []
|
||||||
|
finish_reason: str | None = None
|
||||||
|
text_events = []
|
||||||
|
stream_decoder = getattr(client, "stream_decoder", None)
|
||||||
|
|
||||||
|
async for raw in stream_iter:
|
||||||
|
raw_lines.append(raw)
|
||||||
|
if stream_decoder is not None:
|
||||||
|
try:
|
||||||
|
chunk = stream_decoder.decode(raw)
|
||||||
|
except Exception: # noqa: BLE001
|
||||||
|
continue
|
||||||
|
if chunk.choices:
|
||||||
|
delta_text = chunk.choices[0].delta.content
|
||||||
|
if isinstance(delta_text, str) and delta_text:
|
||||||
|
content_parts.append(delta_text)
|
||||||
|
text_events.append(TextDeltaEvent(content=delta_text))
|
||||||
|
fr = chunk.choices[0].finish_reason
|
||||||
|
if isinstance(fr, str):
|
||||||
|
finish_reason = fr
|
||||||
|
|
||||||
|
full_content = "".join(content_parts) or ""
|
||||||
|
tool_calls: list[AnyAssistantToolCall] | Unset = UNSET
|
||||||
|
|
||||||
|
if finish_reason in ("tool_calls", "function_call"):
|
||||||
|
from plyngent.lmproto.openai_compatible.client import merge_stream_tool_calls
|
||||||
|
|
||||||
|
calls = merge_stream_tool_calls(raw_lines)
|
||||||
|
if calls:
|
||||||
|
tool_calls = cast("list[AnyAssistantToolCall]", calls)
|
||||||
|
|
||||||
|
assistant = AssistantChatMessage(content=full_content or None, tool_calls=tool_calls)
|
||||||
|
return assistant, text_events
|
||||||
|
|
||||||
|
|
||||||
|
async def run_chat_loop( # noqa: PLR0913, C901
|
||||||
client: ChatClient,
|
client: ChatClient,
|
||||||
messages: list[AnyChatMessage],
|
messages: list[AnyChatMessage],
|
||||||
*,
|
*,
|
||||||
@@ -63,12 +123,14 @@ async def run_chat_loop( # noqa: PLR0913
|
|||||||
max_rounds: int = DEFAULT_MAX_ROUNDS,
|
max_rounds: int = DEFAULT_MAX_ROUNDS,
|
||||||
temperature: float | None = None,
|
temperature: float | None = None,
|
||||||
on_limit: LimitContinueHook | None = None,
|
on_limit: LimitContinueHook | None = None,
|
||||||
|
stream: bool = True,
|
||||||
) -> AsyncIterator[AgentEvent]:
|
) -> AsyncIterator[AgentEvent]:
|
||||||
"""Multi-round chat/tool loop; mutates ``messages`` in place and yields events.
|
"""Multi-round chat/tool loop; mutates ``messages`` in place and yields events.
|
||||||
|
|
||||||
|
When ``stream=True``, text tokens yield as they arrive; for tool-calling
|
||||||
|
rounds, the full response is reconstructed from stream deltas.
|
||||||
|
|
||||||
Continues until the model returns no tool calls, or ``max_rounds`` is hit.
|
Continues until the model returns no tool calls, or ``max_rounds`` is hit.
|
||||||
If ``on_limit`` is set and returns True when the cap is reached, another
|
|
||||||
batch of ``max_rounds`` is granted and the loop continues.
|
|
||||||
"""
|
"""
|
||||||
tool_items: Sequence[AnyToolItem] | None = None
|
tool_items: Sequence[AnyToolItem] | None = None
|
||||||
if tools is not None and len(tools) > 0:
|
if tools is not None and len(tools) > 0:
|
||||||
@@ -86,17 +148,23 @@ async def run_chat_loop( # noqa: PLR0913
|
|||||||
temperature=temperature if temperature is not None else UNSET,
|
temperature=temperature if temperature is not None else UNSET,
|
||||||
tools=list(tool_items) if tool_items is not None else UNSET,
|
tools=list(tool_items) if tool_items is not None else UNSET,
|
||||||
)
|
)
|
||||||
response = await client.chat_completions(param, stream=False)
|
|
||||||
if not response.choices:
|
if stream:
|
||||||
msg = "chat completion response contained no choices"
|
assistant, text_events = await _stream_and_build_assistant(client, param)
|
||||||
raise RuntimeError(msg)
|
for event in text_events:
|
||||||
assistant = response.choices[0].message
|
yield event
|
||||||
|
else:
|
||||||
|
response = await client.chat_completions(param, stream=False)
|
||||||
|
if not response.choices:
|
||||||
|
msg = "chat completion response contained no choices"
|
||||||
|
raise RuntimeError(msg)
|
||||||
|
assistant = response.choices[0].message
|
||||||
|
if isinstance(assistant.content, str) and assistant.content:
|
||||||
|
yield TextDeltaEvent(content=assistant.content)
|
||||||
|
|
||||||
messages.append(assistant)
|
messages.append(assistant)
|
||||||
yield AssistantMessageEvent(message=assistant)
|
yield AssistantMessageEvent(message=assistant)
|
||||||
|
|
||||||
if isinstance(assistant.content, str) and assistant.content:
|
|
||||||
yield TextDeltaEvent(content=assistant.content)
|
|
||||||
|
|
||||||
tool_calls = assistant.tool_calls
|
tool_calls = assistant.tool_calls
|
||||||
if tool_calls is UNSET or not tool_calls:
|
if tool_calls is UNSET or not tool_calls:
|
||||||
return
|
return
|
||||||
@@ -112,7 +180,3 @@ async def run_chat_loop( # noqa: PLR0913
|
|||||||
continue
|
continue
|
||||||
yield MaxRoundsEvent(rounds=allowance, continued=False)
|
yield MaxRoundsEvent(rounds=allowance, continued=False)
|
||||||
return
|
return
|
||||||
|
|
||||||
|
|
||||||
def collect_assistant_messages(events: Sequence[AgentEvent]) -> list[AssistantChatMessage]:
|
|
||||||
return [e.message for e in events if isinstance(e, AssistantMessageEvent)]
|
|
||||||
|
|||||||
Reference in New Issue
Block a user