|
|
@@ -1,13 +1,13 @@
|
|
|
-import json
|
|
|
+import re
|
|
|
from typing import Any
|
|
|
|
|
|
-from agent_lab.domain.events import ToolCallEvent
|
|
|
+from agent_lab.domain.events import EVENT_BLOCK_END, EVENT_BLOCK_START, ToolCallEvent
|
|
|
from agent_lab.domain.messages import StreamItem, TokenUsage
|
|
|
|
|
|
|
|
|
class ChatCompletionStreamParser:
|
|
|
def __init__(self) -> None:
|
|
|
- self._tool_calls: dict[int, dict[str, Any]] = {}
|
|
|
+ self._text_events = TextEventProtocolParser()
|
|
|
|
|
|
def feed(self, chunk: dict[str, Any]) -> list[StreamItem]:
|
|
|
items: list[StreamItem] = []
|
|
|
@@ -16,13 +16,10 @@ class ChatCompletionStreamParser:
|
|
|
delta = choice.get("delta") or {}
|
|
|
|
|
|
if "content" in delta and delta["content"] is not None:
|
|
|
- items.append(StreamItem.message_delta(delta["content"]))
|
|
|
+ items.extend(self._text_events.feed(delta["content"]))
|
|
|
|
|
|
- for tool_call in delta.get("tool_calls") or []:
|
|
|
- self._accumulate_tool_call(tool_call)
|
|
|
-
|
|
|
- if choice.get("finish_reason") == "tool_calls":
|
|
|
- items.extend(self._drain_tool_events())
|
|
|
+ if choice.get("finish_reason") is not None:
|
|
|
+ items.extend(self.flush())
|
|
|
|
|
|
usage = chunk.get("usage")
|
|
|
if usage is not None:
|
|
|
@@ -40,54 +37,121 @@ class ChatCompletionStreamParser:
|
|
|
|
|
|
return items
|
|
|
|
|
|
- def _accumulate_tool_call(self, tool_call: dict[str, Any]) -> None:
|
|
|
- index = int(tool_call["index"])
|
|
|
- state = self._tool_calls.setdefault(
|
|
|
- index,
|
|
|
- {
|
|
|
- "id": "",
|
|
|
- "name": "",
|
|
|
- "arguments": "",
|
|
|
- },
|
|
|
- )
|
|
|
-
|
|
|
- if tool_call.get("id"):
|
|
|
- state["id"] = tool_call["id"]
|
|
|
-
|
|
|
- function = tool_call.get("function") or {}
|
|
|
- if function.get("name"):
|
|
|
- state["name"] = function["name"]
|
|
|
- if "arguments" in function and function["arguments"] is not None:
|
|
|
- state["arguments"] += function["arguments"]
|
|
|
-
|
|
|
- def _drain_tool_events(self) -> list[StreamItem]:
|
|
|
+ def flush(self) -> list[StreamItem]:
|
|
|
+ return self._text_events.flush()
|
|
|
+
|
|
|
+
|
|
|
+class TextEventProtocolParser:
|
|
|
+ _event_name_pattern = re.compile(r"^[A-Za-z_][A-Za-z0-9_.:-]*$")
|
|
|
+
|
|
|
+ def __init__(self) -> None:
|
|
|
+ self._mode = "message"
|
|
|
+ self._message_buffer = ""
|
|
|
+ self._event_buffer = ""
|
|
|
+ self._event_count = 0
|
|
|
+
|
|
|
+ def feed(self, content: str) -> list[StreamItem]:
|
|
|
+ items: list[StreamItem] = []
|
|
|
+ pending = content
|
|
|
+
|
|
|
+ while pending:
|
|
|
+ if self._mode == "message":
|
|
|
+ pending = self._feed_message(pending, items)
|
|
|
+ continue
|
|
|
+ pending = self._feed_events(pending, items)
|
|
|
+
|
|
|
+ return items
|
|
|
+
|
|
|
+ def flush(self) -> list[StreamItem]:
|
|
|
+ items: list[StreamItem] = []
|
|
|
+ if self._mode == "message":
|
|
|
+ self._emit_message(self._message_buffer, items)
|
|
|
+ self._message_buffer = ""
|
|
|
+ return items
|
|
|
+
|
|
|
+ items.extend(self._emit_event_lines(self._event_buffer))
|
|
|
+ self._event_buffer = ""
|
|
|
+ return items
|
|
|
+
|
|
|
+ def _feed_message(
|
|
|
+ self,
|
|
|
+ content: str,
|
|
|
+ items: list[StreamItem],
|
|
|
+ ) -> str:
|
|
|
+ self._message_buffer += content
|
|
|
+ marker_index = self._message_buffer.find(EVENT_BLOCK_START)
|
|
|
+ if marker_index >= 0:
|
|
|
+ self._emit_message(self._message_buffer[:marker_index], items)
|
|
|
+ remainder = self._message_buffer[
|
|
|
+ marker_index + len(EVENT_BLOCK_START):
|
|
|
+ ]
|
|
|
+ self._message_buffer = ""
|
|
|
+ self._mode = "events"
|
|
|
+ return remainder
|
|
|
+
|
|
|
+ keep = self._partial_marker_suffix_length(self._message_buffer)
|
|
|
+ if keep:
|
|
|
+ self._emit_message(self._message_buffer[:-keep], items)
|
|
|
+ self._message_buffer = self._message_buffer[-keep:]
|
|
|
+ return ""
|
|
|
+
|
|
|
+ self._emit_message(self._message_buffer, items)
|
|
|
+ self._message_buffer = ""
|
|
|
+ return ""
|
|
|
+
|
|
|
+ def _feed_events(
|
|
|
+ self,
|
|
|
+ content: str,
|
|
|
+ items: list[StreamItem],
|
|
|
+ ) -> str:
|
|
|
+ self._event_buffer += content
|
|
|
+ marker_index = self._event_buffer.find(EVENT_BLOCK_END)
|
|
|
+ if marker_index >= 0:
|
|
|
+ event_text = self._event_buffer[:marker_index]
|
|
|
+ items.extend(self._emit_event_lines(event_text))
|
|
|
+ remainder = self._event_buffer[marker_index + len(EVENT_BLOCK_END):]
|
|
|
+ self._event_buffer = ""
|
|
|
+ self._mode = "message"
|
|
|
+ return remainder
|
|
|
+
|
|
|
+ last_newline = self._event_buffer.rfind("\n")
|
|
|
+ if last_newline >= 0:
|
|
|
+ complete_text = self._event_buffer[: last_newline + 1]
|
|
|
+ self._event_buffer = self._event_buffer[last_newline + 1:]
|
|
|
+ items.extend(self._emit_event_lines(complete_text))
|
|
|
+ return ""
|
|
|
+
|
|
|
+ def _emit_message(self, content: str, items: list[StreamItem]) -> None:
|
|
|
+ if content:
|
|
|
+ items.append(StreamItem.message_delta(content))
|
|
|
+
|
|
|
+ def _emit_event_lines(self, text: str) -> list[StreamItem]:
|
|
|
events: list[StreamItem] = []
|
|
|
- for index in sorted(self._tool_calls):
|
|
|
- state = self._tool_calls[index]
|
|
|
- raw_arguments = state["arguments"]
|
|
|
+ for line in text.splitlines():
|
|
|
+ name = self._clean_event_name(line)
|
|
|
+ if not name:
|
|
|
+ continue
|
|
|
+ if not self._event_name_pattern.fullmatch(name):
|
|
|
+ continue
|
|
|
+ self._event_count += 1
|
|
|
events.append(
|
|
|
StreamItem.event(
|
|
|
ToolCallEvent(
|
|
|
- id=state["id"],
|
|
|
- name=state["name"],
|
|
|
- arguments=self._parse_arguments(raw_arguments),
|
|
|
- raw_arguments=raw_arguments,
|
|
|
+ id=f"event_{self._event_count}",
|
|
|
+ name=name,
|
|
|
+ arguments={},
|
|
|
+ raw_arguments="{}",
|
|
|
)
|
|
|
)
|
|
|
)
|
|
|
-
|
|
|
- self._tool_calls.clear()
|
|
|
return events
|
|
|
|
|
|
- def _parse_arguments(self, raw_arguments: str) -> dict[str, Any]:
|
|
|
- if not raw_arguments:
|
|
|
- return {}
|
|
|
-
|
|
|
- try:
|
|
|
- parsed = json.loads(raw_arguments)
|
|
|
- except json.JSONDecodeError:
|
|
|
- return {"raw": raw_arguments}
|
|
|
+ def _clean_event_name(self, line: str) -> str:
|
|
|
+ return line.strip().removeprefix("-").strip().strip("`'\"")
|
|
|
|
|
|
- if isinstance(parsed, dict):
|
|
|
- return parsed
|
|
|
- return {"value": parsed}
|
|
|
+ def _partial_marker_suffix_length(self, content: str) -> int:
|
|
|
+ max_length = min(len(content), len(EVENT_BLOCK_START) - 1)
|
|
|
+ for length in range(max_length, 0, -1):
|
|
|
+ if EVENT_BLOCK_START.startswith(content[-length:]):
|
|
|
+ return length
|
|
|
+ return 0
|