diff --git a/google/genai/chats.py b/google/genai/chats.py index 3ba5f3c95..e904c539c 100644 --- a/google/genai/chats.py +++ b/google/genai/chats.py @@ -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]: @@ -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, @@ -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, @@ -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: @@ -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 @@ -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: @@ -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 @@ -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( @@ -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: @@ -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] = [] @@ -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 @@ -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: @@ -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 diff --git a/google/genai/tests/chats/test_continuation_token.py b/google/genai/tests/chats/test_continuation_token.py index d136b2b8d..4933a2e0d 100644 --- a/google/genai/tests/chats/test_continuation_token.py +++ b/google/genai/tests/chats/test_continuation_token.py @@ -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) == []