diff --git a/headroom/proxy/handlers/gemini.py b/headroom/proxy/handlers/gemini.py index efb9cfe08..dba8bbbb1 100644 --- a/headroom/proxy/handlers/gemini.py +++ b/headroom/proxy/handlers/gemini.py @@ -56,6 +56,24 @@ class GeminiHandlerMixin: return ANTIGRAVITY_DAILY_API_URL return getattr(self, "CLOUDCODE_API_URL", DEFAULT_CLOUDCODE_API_URL).rstrip("/") + @staticmethod + def _dict_parts(content: Any) -> list[dict]: + """Return the dict entries of a Gemini content's ``parts``. + + ``parts`` is request-controlled. ``.get("parts", [])`` only falls back + when the key is absent, so a present-but-null ``parts`` returns ``None`` + (crashing ``for part in parts``), and a list carrying a bare string — + which a client that treats ``parts`` as a string array can send — + crashes ``part.get(...)`` / ``key in part`` semantics downstream. + Coerce to a clean list of dict parts so every caller can iterate safely. + """ + if not isinstance(content, dict): + return [] + parts = content.get("parts") + if not isinstance(parts, list): + return [] + return [part for part in parts if isinstance(part, dict)] + def _has_non_text_parts(self, content: dict) -> bool: """Check if a Gemini content entry has non-text parts. @@ -73,8 +91,7 @@ class GeminiHandlerMixin: Returns: True if any part contains non-text data. """ - parts = content.get("parts", []) - for part in parts: + for part in self._dict_parts(content): if any( key in part for key in ( @@ -110,7 +127,7 @@ class GeminiHandlerMixin: opt_iter = iter(optimized_contents) result: list[dict] = [] for idx, content in enumerate(original_contents): - had_text = any("text" in p for p in content.get("parts", [])) + had_text = any("text" in p for p in self._dict_parts(content)) if idx in preserved_indices: result.append(preserved_contents[idx]) if had_text: @@ -155,8 +172,8 @@ class GeminiHandlerMixin: # Add system instruction as system message if system_instruction: - parts = system_instruction.get("parts", []) - text_parts = [p.get("text", "") for p in parts if "text" in p] + sys_parts = self._dict_parts(system_instruction) + text_parts = [p.get("text", "") for p in sys_parts if "text" in p] if text_parts: messages.append({"role": "system", "content": "\n".join(text_parts)}) @@ -166,12 +183,12 @@ class GeminiHandlerMixin: if self._has_non_text_parts(content): preserved_indices.add(idx) - role = content.get("role", "user") + role = content.get("role", "user") if isinstance(content, dict) else "user" # Map Gemini roles to OpenAI roles if role == "model": role = "assistant" - parts = content.get("parts", []) + parts = self._dict_parts(content) text_parts = [p.get("text", "") for p in parts if "text" in p] if text_parts: diff --git a/tests/test_gemini_function_response_waste.py b/tests/test_gemini_function_response_waste.py index 638db4f8a..20161e66c 100644 --- a/tests/test_gemini_function_response_waste.py +++ b/tests/test_gemini_function_response_waste.py @@ -124,6 +124,45 @@ class TestFunctionResponseConversion: assert "loop" in text +class TestMalformedPartsToleration: + """A request-controlled `parts` that is null or carries non-dict elements + must not crash the compression-path conversion helpers.""" + + def test_null_parts_does_not_crash(self, proxy): + contents = [ + {"role": "user", "parts": None}, + {"role": "user", "parts": [{"text": "real"}]}, + ] + messages, preserved = proxy._gemini_contents_to_messages(contents) + assert messages == [{"role": "user", "content": "real"}] + assert preserved == set() + + def test_string_part_elements_do_not_crash(self, proxy): + # A client that treats `parts` as a string array sends bare strings; + # they carry no `text` key, so they contribute nothing but must not + # crash `.get`. + contents = [{"role": "user", "parts": ["bare string", {"text": "kept"}]}] + messages, _ = proxy._gemini_contents_to_messages(contents) + assert messages == [{"role": "user", "content": "kept"}] + + def test_null_part_element_is_skipped(self, proxy): + contents = [{"role": "user", "parts": [None, {"text": "kept"}]}] + messages, _ = proxy._gemini_contents_to_messages(contents) + assert messages == [{"role": "user", "content": "kept"}] + + def test_has_non_text_parts_tolerates_null_parts(self, proxy): + assert proxy._has_non_text_parts({"role": "user", "parts": None}) is False + assert proxy._has_non_text_parts({"role": "user", "parts": ["str"]}) is False + assert proxy._has_non_text_parts({"parts": [{"inlineData": {"data": "x"}}]}) is True + + def test_non_dict_content_entry_is_tolerated(self, proxy): + # A non-dict entry in contents[] is treated as an empty user turn rather + # than crashing content.get / the parts iteration. + contents = ["not a dict", {"role": "user", "parts": [{"text": "kept"}]}] + messages, _ = proxy._gemini_contents_to_messages(contents) + assert messages == [{"role": "user", "content": "kept"}] + + class TestFunctionResponseWasteParsing: def test_function_response_payload_reaches_waste_signals(self, proxy, tokenizer): contents = [