Skip to content
Merged
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
77 changes: 67 additions & 10 deletions google/genai/chats.py
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,29 @@ def _validate_response(response: GenerateContentResponse) -> bool:
return _validate_content(response.candidates[0].content)


def _update_stream_finish_reason(
chunk: GenerateContentResponse,
finish_reason: Optional[types.FinishReason],
hop_continuation_token: Optional[bytes],
enable_continuation: bool,
) -> tuple[Optional[types.FinishReason], Optional[bytes]]:
if not chunk.candidates:
return finish_reason, hop_continuation_token
candidate = chunk.candidates[0]
if candidate.continuation_token:
hop_continuation_token = candidate.continuation_token
if candidate.finish_reason:
finish_reason = candidate.finish_reason
if (
enable_continuation
and hop_continuation_token
and _is_resumable_finish_reason(finish_reason)
):
finish_reason = None
hop_continuation_token = None
return finish_reason, hop_continuation_token


def _extract_curated_history(
comprehensive_history: list[Content],
) -> list[Content]:
Expand Down Expand Up @@ -542,6 +565,10 @@ def send_message_stream(

if disable_afc:
if isinstance(self._modules, Models):
enable_continuation = _should_enable_automatic_continuation(
parsed_config, default_enabled=True
)
hop_continuation_token = None
for chunk in self._generate_content_stream_with_continuation(
contents=contents_to_model, # type: ignore[arg-type]
config=parsed_config,
Expand All @@ -550,8 +577,12 @@ def send_message_stream(
is_valid = False
if chunk.candidates and chunk.candidates[0].content:
model_output.append(chunk.candidates[0].content)
if chunk.candidates and chunk.candidates[0].finish_reason:
finish_reason = chunk.candidates[0].finish_reason
finish_reason, hop_continuation_token = _update_stream_finish_reason(
chunk,
finish_reason,
hop_continuation_token,
enable_continuation,
)
yield chunk
self.record_history(
user_input=user_input,
Expand Down Expand Up @@ -579,6 +610,9 @@ def send_message_stream(
f"AFC is enabled with max remote calls: {remaining_remote_calls_afc}."
)
function_map = _extra_utils.get_function_map(parsed_config)
enable_continuation = _should_enable_automatic_continuation(
parsed_config, default_enabled=True
)
i = 0
if isinstance(self._modules, Models):
while remaining_remote_calls_afc > 0:
Expand All @@ -599,6 +633,7 @@ def send_message_stream(

model_output = []
finish_reason = None
hop_continuation_token = None
is_valid = True
func_response_parts = []
chunk = None
Expand All @@ -622,8 +657,12 @@ def send_message_stream(

if chunk.candidates and chunk.candidates[0].content:
model_output.append(chunk.candidates[0].content)
if chunk.candidates and chunk.candidates[0].finish_reason:
finish_reason = chunk.candidates[0].finish_reason
finish_reason, hop_continuation_token = _update_stream_finish_reason(
chunk,
finish_reason,
hop_continuation_token,
enable_continuation,
)
yield chunk

if is_last_remote_call_afc:
Expand All @@ -643,7 +682,7 @@ def send_message_stream(
self.record_history(
user_input=user_input,
model_output=model_output,
is_valid=is_valid,
is_valid=is_valid and finish_reason is not None,
)
user_input = func_response_content

Expand Down Expand Up @@ -1103,6 +1142,10 @@ async def async_generator(): # type: ignore[no-untyped-def]
if disable_afc:
output_contents = []
finish_reason = None
hop_continuation_token = None
enable_continuation = _should_enable_automatic_continuation(
parsed_config, default_enabled=True
)
is_valid = True
chunk = None
async for chunk in self._generate_content_stream_with_continuation(
Expand All @@ -1113,8 +1156,12 @@ async def async_generator(): # type: ignore[no-untyped-def]
is_valid = False
if chunk.candidates and chunk.candidates[0].content:
output_contents.append(chunk.candidates[0].content)
if chunk.candidates and chunk.candidates[0].finish_reason:
finish_reason = chunk.candidates[0].finish_reason
finish_reason, hop_continuation_token = _update_stream_finish_reason(
chunk,
finish_reason,
hop_continuation_token,
enable_continuation,
)
yield chunk

if not output_contents or finish_reason is None:
Expand Down Expand Up @@ -1212,6 +1259,9 @@ async def async_generator(): # type: ignore[no-untyped-def]
mcp_to_genai_tool_adapters,
is_caller_method_async=True,
)
enable_continuation = _should_enable_automatic_continuation(
final_parsed_config, default_enabled=True
)

i = 0
model_output: list[types.Content] = []
Expand All @@ -1236,6 +1286,7 @@ async def async_generator(): # type: ignore[no-untyped-def]

model_output = []
finish_reason = None
hop_continuation_token = None
is_valid = True
func_response_parts = []
chunk = None
Expand All @@ -1261,8 +1312,14 @@ async def async_generator(): # type: ignore[no-untyped-def]

if chunk.candidates and chunk.candidates[0].content:
model_output.append(chunk.candidates[0].content)
if chunk.candidates and chunk.candidates[0].finish_reason:
finish_reason = chunk.candidates[0].finish_reason
finish_reason, hop_continuation_token = (
_update_stream_finish_reason(
chunk,
finish_reason,
hop_continuation_token,
enable_continuation,
)
)
yield chunk

if is_last_remote_call_afc:
Expand All @@ -1282,7 +1339,7 @@ async def async_generator(): # type: ignore[no-untyped-def]
self.record_history(
user_input=user_input,
model_output=model_output,
is_valid=is_valid,
is_valid=is_valid and finish_reason is not None,
)
user_input = func_response_content

Expand Down
116 changes: 116 additions & 0 deletions google/genai/tests/chats/test_continuation_token.py
Original file line number Diff line number Diff line change
Expand Up @@ -1618,3 +1618,119 @@ async def test_async_chat_automatic_continuation_config_rules(mock_api_client):
)
assert mock_gc.call_count == 1
assert resp.text == 'Async Hop 1. '


def test_chat_send_message_stream_incomplete_continuation_not_recorded(
mock_api_client,
):
"""When a continuation hop is cut off without finish_reason, the turn is excluded from curated history."""
models_module = models.Models(mock_api_client)
chats_module = chats.Chats(modules=models_module)
chat = chats_module.create(model='gemini-2.5-pro')

hop1_chunks = [
types.GenerateContentResponse(
candidates=[
types.Candidate(
content=types.Content(
role='model', parts=[types.Part(text='Hop 1 part. ')]
),
continuation_token=b'tok_hop_1',
)
]
),
types.GenerateContentResponse(
candidates=[
types.Candidate(
content=types.Content(
role='model', parts=[types.Part(text='End of hop 1. ')]
),
finish_reason=types.FinishReason.CONTINUATION,
)
]
),
]
hop2_cutoff_chunks = [
types.GenerateContentResponse(
candidates=[
types.Candidate(
content=types.Content(
role='model',
parts=[types.Part(text='Hop 2 cut off mid-stream')],
),
finish_reason=None,
)
]
)
]

with mock.patch.object(
models.Models,
'generate_content_stream',
side_effect=[iter(hop1_chunks), iter(hop2_cutoff_chunks)],
) as mock_stream:
chunks = list(chat.send_message_stream('Write a long story'))
assert mock_stream.call_count == 2
assert [c.text for c in chunks] == [
'Hop 1 part. ',
'End of hop 1. ',
'Hop 2 cut off mid-stream',
]
assert chat.get_history(curated=True) == []


@pytest.mark.asyncio
async def test_async_chat_send_message_stream_incomplete_continuation_not_recorded(
mock_api_client,
):
"""When an async continuation hop is cut off without finish_reason, the turn is excluded from curated history."""
models_module = models.AsyncModels(mock_api_client)
chats_module = chats.AsyncChats(modules=models_module)
chat = chats_module.create(model='gemini-2.5-pro')

hop1_chunks = [
types.GenerateContentResponse(
candidates=[
types.Candidate(
content=types.Content(
role='model', parts=[types.Part(text='Async Hop 1. ')]
),
continuation_token=b'async_tok_hop_1',
finish_reason=types.FinishReason.CONTINUATION,
)
]
)
]
hop2_cutoff_chunks = [
types.GenerateContentResponse(
candidates=[
types.Candidate(
content=types.Content(
role='model',
parts=[types.Part(text='Async Hop 2 cut off')],
),
finish_reason=None,
)
]
)
]

async def _make_async_stream(chunk_list):
for c in chunk_list:
yield c

with mock.patch.object(
models.AsyncModels,
'generate_content_stream',
new_callable=mock.AsyncMock,
side_effect=[
_make_async_stream(hop1_chunks),
_make_async_stream(hop2_cutoff_chunks),
],
) as mock_async_stream:
chunks = []
async for chunk in await chat.send_message_stream('Write a long story'):
chunks.append(chunk)
assert mock_async_stream.call_count == 2
assert [c.text for c in chunks] == ['Async Hop 1. ', 'Async Hop 2 cut off']
assert chat.get_history(curated=True) == []
Loading