Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 11 additions & 1 deletion src/google/adk/utils/streaming_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,10 @@
re.VERBOSE,
)

_JSON_HIGH_SURROGATE_SUFFIX_RE = re.compile(
r'(?<!\\)(?:\\\\)*\\u[dD][89aAbB][0-9a-fA-F]{2}$'
)


def _unescape_json_path_string(s: str) -> str:
def replace(match: re.Match[str]) -> str:
Expand Down Expand Up @@ -782,10 +786,16 @@ def _complete_json(self) -> str:
if self._in_string:
if self._escaped:
return ''
prefix = ''.join(self.accumulated_parts)
# A high surrogate needs the following low surrogate before json.loads
# can decode the pair into a character that can be serialized as UTF-8.
# Escaped backslashes are literal text and must not delay streaming.
if _JSON_HIGH_SURROGATE_SUFFIX_RE.search(prefix):
return ''
suffix = '"' + ''.join(
'}' if op == '{' else ']' for op in reversed(self._stack)
)
return ''.join(self.accumulated_parts) + suffix
return prefix + suffix

last_non_ws = ''
last_part_idx = -1
Expand Down
33 changes: 33 additions & 0 deletions tests/unittests/utils/test_streaming_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -2207,6 +2207,39 @@ def test_tracker_empty_chunk_after_numeric_does_not_raise(self):
diffs_2 = tracker.handle_chunk("")
assert diffs_2 == []

@pytest.mark.parametrize(
"chunks",
[
['{"text": "hello ', r"\ud83d", r"\ude00", ' world"}'],
[r'{"text": "hello \ud83d', r'\ude00 world"}'],
['{"text": "hello ', r"\uD83D", r"\uDE00", ' world"}'],
['{"text": "hello ', r"\ud83d\u", "de00", ' world"}'],
],
)
def test_tracker_preserves_surrogate_pairs_across_chunks(self, chunks):
"""Streamed emoji arguments remain intact and serializable."""
tracker = streaming_utils._JsonPathTracker()

partial_args = [
arg for chunk in chunks for arg in tracker.handle_chunk(chunk)
]

assert "".join(arg.string_value or "" for arg in partial_args) == (
"hello 😀 world"
)
for arg in partial_args:
arg.model_dump_json()

def test_tracker_streams_literal_surrogate_escape_without_waiting(self):
"""An escaped backslash is ordinary text, not a surrogate pair."""
tracker = streaming_utils._JsonPathTracker()

partial_args = tracker.handle_chunk(r'{"text": "hello \\ud83d')

assert len(partial_args) == 1
assert partial_args[0].string_value == r"hello \ud83d"
assert partial_args[0].will_continue is True

def test_tracker_escape_at_chunk_boundary_does_not_corrupt_fast_path(self):
tracker = streaming_utils._JsonPathTracker()
diffs_1 = tracker.handle_chunk('{"text": "line1\\')
Expand Down
Loading