diff --git a/langextract/providers/gemini.py b/langextract/providers/gemini.py index 18416820..8e7352f9 100644 --- a/langextract/providers/gemini.py +++ b/langextract/providers/gemini.py @@ -348,6 +348,34 @@ def _is_retryable_error(self, error: Exception) -> bool: return bool(_RETRYABLE_MESSAGE_RE.search(str(error))) + @staticmethod + def _describe_missing_gemini_text(response: Any) -> str: + """Builds a diagnostic message for a Gemini response with no text. + + `response.text` returns None (not an exception) when the prompt was + blocked before generation, when the candidate stopped for a non-STOP + reason (safety, recitation, etc.) with no content, or when the response + contains only non-text parts (e.g. a function call). Without this, + callers previously got ScoredOutput(score=1.0, output=None) reported as + a successful empty extraction, silently discarding the actual reason. + """ + prompt_feedback = getattr(response, 'prompt_feedback', None) + block_reason = getattr(prompt_feedback, 'block_reason', None) + if block_reason: + block_message = getattr(prompt_feedback, 'block_reason_message', None) + detail = f': {block_message}' if block_message else '' + return f'Gemini blocked the prompt ({block_reason}){detail}' + + candidates = getattr(response, 'candidates', None) + if candidates: + finish_reason = getattr(candidates[0], 'finish_reason', None) + if finish_reason and finish_reason != 'STOP': + return ( + f'Gemini response contained no text (finish_reason={finish_reason})' + ) + + return 'Gemini response contained no text content.' + def _process_single_prompt( self, prompt: str, config: dict ) -> core_types.ScoredOutput: @@ -368,8 +396,15 @@ def _process_single_prompt( response = self._client.models.generate_content( model=self.model_id, contents=prompt, config=call_config ) - return core_types.ScoredOutput(score=1.0, output=response.text) + output_text = response.text + if output_text is None: + raise exceptions.InferenceRuntimeError( + self._describe_missing_gemini_text(response) + ) + return core_types.ScoredOutput(score=1.0, output=output_text) + except exceptions.InferenceRuntimeError: + raise except Exception as e: if attempt < self.max_retries and self._is_retryable_error(e): # Cap after jitter so the named maximum applies to the real sleep. diff --git a/tests/provider_schema_test.py b/tests/provider_schema_test.py index 0ad06f16..506a0c69 100644 --- a/tests/provider_schema_test.py +++ b/tests/provider_schema_test.py @@ -740,5 +740,91 @@ def test_apply_output_schema_rejects_constructor_gemini_schema(self): model.apply_output_schema(self.output_schema) +class GeminiRefusalHandlingTest(absltest.TestCase): + """Tests Gemini safety-block and empty-response handling. + + response.text returns None (not an exception) when a prompt is blocked + before generation or a candidate stops for a non-STOP reason with no + content -- mirrors the OpenAI refusal-handling gap fixed for #491, but + for the Gemini realtime path, which was left unaddressed. + """ + + def setUp(self): + super().setUp() + + patcher = mock.patch("google.genai.Client", autospec=True) + self.addCleanup(patcher.stop) + + mock_client_cls = patcher.start() + self.mock_client = mock_client_cls.return_value + + self.model = gemini.GeminiLanguageModel( + model_id="gemini-3.5-flash", + api_key="test_key", + ) + + def test_process_single_prompt_returns_text(self): + response = mock.Mock(text="Hello") + self.mock_client.models.generate_content.return_value = response + + result = self.model._process_single_prompt("prompt", {}) + + self.assertEqual(result.output, "Hello") + self.assertEqual(result.score, 1.0) + + def test_process_single_prompt_raises_on_blocked_prompt(self): + prompt_feedback = mock.Mock( + block_reason="SAFETY", block_reason_message="blocked content" + ) + response = mock.Mock( + text=None, prompt_feedback=prompt_feedback, candidates=[] + ) + self.mock_client.models.generate_content.return_value = response + + with self.assertRaisesRegex( + exceptions.InferenceRuntimeError, + "Gemini blocked the prompt \\(SAFETY\\): blocked content", + ): + self.model._process_single_prompt("prompt", {}) + + def test_process_single_prompt_raises_on_non_stop_finish_reason(self): + candidate = mock.Mock(finish_reason="SAFETY") + response = mock.Mock( + text=None, prompt_feedback=None, candidates=[candidate] + ) + self.mock_client.models.generate_content.return_value = response + + with self.assertRaisesRegex( + exceptions.InferenceRuntimeError, + "no text \\(finish_reason=SAFETY\\)", + ): + self.model._process_single_prompt("prompt", {}) + + def test_process_single_prompt_raises_when_no_diagnostic_available(self): + response = mock.Mock(text=None, prompt_feedback=None, candidates=[]) + + self.mock_client.models.generate_content.return_value = response + + with self.assertRaisesRegex( + exceptions.InferenceRuntimeError, + "contained no text content", + ): + self.model._process_single_prompt("prompt", {}) + + def test_process_single_prompt_does_not_retry_on_refusal(self): + prompt_feedback = mock.Mock( + block_reason="SAFETY", block_reason_message=None + ) + response = mock.Mock( + text=None, prompt_feedback=prompt_feedback, candidates=[] + ) + self.mock_client.models.generate_content.return_value = response + + with self.assertRaises(exceptions.InferenceRuntimeError): + self.model._process_single_prompt("prompt", {}) + + self.mock_client.models.generate_content.assert_called_once() + + if __name__ == "__main__": absltest.main()