diff --git a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py index 6a0dee03..97eb5845 100644 --- a/src/google/adk/plugins/bigquery_agent_analytics_plugin.py +++ b/src/google/adk/plugins/bigquery_agent_analytics_plugin.py @@ -1484,9 +1484,36 @@ class HybridContentParser: ) async def _parse_content_object( - self, content: types.Content | types.Part + self, + content: types.Content | types.Part, + *, + trace_id: Optional[str] = None, + span_id: Optional[str] = None, + parse_uid: str = "", + content_ordinal: int = 0, ) -> tuple[str, list[dict[str, Any]], bool]: - """Parses a Content or Part object into summary text and content parts.""" + """Parses a Content or Part object into summary text and content parts. + + Args: + content: The Content or Part to parse. + trace_id: Trace id of the calling event. GCS object paths are built + from this argument rather than the instance field, because the + parser is shared across concurrent events and an await inside this + method can resume under another event's identity. Falls back to + the instance field. + span_id: Span id of the calling event, with the same rationale. + parse_uid: Unique per parse() call. Disambiguates object names across + concurrent events. Generated here when not supplied. + content_ordinal: Index of this Content within the calling request. The + part index restarts at zero for each Content, so without this two + messages in one request collide on the same object name. + + Returns: + A tuple of (summary_text, content_parts, is_truncated). + """ + trace_id = trace_id if trace_id is not None else self.trace_id + span_id = span_id if span_id is not None else self.span_id + parse_uid = parse_uid or uuid.uuid4().hex content_parts = [] is_truncated = False summary_text = [] @@ -1518,7 +1545,10 @@ class HybridContentParser: elif hasattr(part, "inline_data") and part.inline_data: if self.offloader: ext = mimetypes.guess_extension(part.inline_data.mime_type) or ".bin" - path = f"{datetime.now().date()}/{self.trace_id}/{self.span_id}_p{idx}{ext}" + path = ( + f"{datetime.now().date()}/{trace_id}/{span_id}_{parse_uid}" + f"_c{content_ordinal}_p{idx}{ext}" + ) try: uri = await self.offloader.upload_content( part.inline_data.data, part.inline_data.mime_type, path @@ -1557,7 +1587,10 @@ class HybridContentParser: if self.offloader and (exceeds_inline_byte_limit or exceeds_char_limit): # Text is too big, treat as file - path = f"{datetime.now().date()}/{self.trace_id}/{self.span_id}_p{idx}.txt" + path = ( + f"{datetime.now().date()}/{trace_id}/{span_id}_{parse_uid}" + f"_c{content_ordinal}_p{idx}.txt" + ) try: uri = await self.offloader.upload_content( part.text, "text/plain", path @@ -1605,8 +1638,32 @@ class HybridContentParser: return summary_str, content_parts, is_truncated - async def parse(self, content: Any) -> tuple[Any, list[dict[str, Any]], bool]: - """Parses content into JSON payload and content parts, potentially offloading to GCS.""" + async def parse( + self, + content: Any, + *, + trace_id: Optional[str] = None, + span_id: Optional[str] = None, + ) -> tuple[Any, list[dict[str, Any]], bool]: + """Parses content into JSON payload and content parts, potentially offloading to GCS. + + Args: + content: The content to parse. + trace_id: Trace id of the calling event, used to build GCS object + paths. Pass it per call: the parser instance is shared across + concurrent events, so a path built from the mutable instance field + can pick up another event's identity across an await. Falls back + to the instance field. + span_id: Span id of the calling event, with the same rationale. + + Returns: + A tuple of (json_payload, content_parts, is_truncated). + """ + trace_id = trace_id if trace_id is not None else self.trace_id + span_id = span_id if span_id is not None else self.span_id + # Unique per parse() call, so two events offloading at the same time + # cannot produce the same object name. + parse_uid = uuid.uuid4().hex json_payload = {} content_parts = [] is_truncated = False @@ -1622,9 +1679,15 @@ class HybridContentParser: if isinstance(content.contents, list) else [content.contents] ) - for c in contents: + for content_idx, c in enumerate(contents): role = getattr(c, "role", "unknown") - summary, parts, trunc = await self._parse_content_object(c) + summary, parts, trunc = await self._parse_content_object( + c, + trace_id=trace_id, + span_id=span_id, + parse_uid=parse_uid, + content_ordinal=content_idx, + ) if trunc: is_truncated = True content_parts.extend(parts) @@ -1642,14 +1705,25 @@ class HybridContentParser: is_truncated = True json_payload["system_prompt"] = truncated_si else: - summary, parts, trunc = await self._parse_content_object(si) + summary, parts, trunc = await self._parse_content_object( + si, + trace_id=trace_id, + span_id=span_id, + parse_uid=parse_uid, + content_ordinal=len(contents), + ) if trunc: is_truncated = True content_parts.extend(parts) json_payload["system_prompt"] = summary elif isinstance(content, (types.Content, types.Part)): - summary, parts, trunc = await self._parse_content_object(content) + summary, parts, trunc = await self._parse_content_object( + content, + trace_id=trace_id, + span_id=span_id, + parse_uid=parse_uid, + ) return {"text_summary": summary}, parts, trunc elif isinstance(content, (dict, list)): @@ -3020,11 +3094,13 @@ class BigQueryAgentAnalyticsPlugin(BasePlugin): logger.warning("Parser not initialized; skipping event %s.", event_type) return - # Update parser's trace/span IDs for GCS pathing (reuse instance) - self.parser.trace_id = trace_id or "no_trace" - self.parser.span_id = span_id or "no_span" + # Pass the ids per call rather than assigning them to the shared parser: + # two events in flight at once would otherwise overwrite each other's + # identity between the assignment and the offload that follows an await. content_json, content_parts, parser_truncated = await self.parser.parse( - raw_content + raw_content, + trace_id=trace_id or "no_trace", + span_id=span_id or "no_span", ) is_truncated = is_truncated or parser_truncated diff --git a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py index aaaf51ea..4bef3c34 100644 --- a/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py +++ b/tests/unittests/plugins/test_bigquery_agent_analytics_plugin.py @@ -2919,27 +2919,34 @@ class TestParserReuse: assert bq_plugin_inst.parser is parser_after_init @pytest.mark.asyncio - async def test_parser_trace_id_updated_per_call( + async def test_parser_ids_are_not_mutated_per_call( self, bq_plugin_inst, mock_write_client, invocation_context, dummy_arrow_schema, ): - """trace_id and span_id on the parser should update per _log_event.""" + """_log_event passes the ids per call instead of writing them on the + + shared parser. + """ parser = bq_plugin_inst.parser original_trace_id = parser.trace_id + original_span_id = parser.span_id - bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context) - await bq_plugin_inst.on_user_message_callback( - invocation_context=invocation_context, - user_message=types.Content(parts=[types.Part(text="Test")]), - ) - await asyncio.sleep(0.01) + with mock.patch.object(parser, "parse", wraps=parser.parse) as mock_parse: + bigquery_agent_analytics_plugin.TraceManager.push_span(invocation_context) + await bq_plugin_inst.on_user_message_callback( + invocation_context=invocation_context, + user_message=types.Content(parts=[types.Part(text="Test")]), + ) + await asyncio.sleep(0.01) - # After logging, trace_id/span_id should have been updated - # (they're derived from TraceManager, not the initial empty strings) - assert parser.span_id != "" + assert parser.trace_id == original_trace_id + assert parser.span_id == original_span_id + _, kwargs = mock_parse.call_args + assert kwargs["span_id"] != "" + assert kwargs["span_id"] != original_span_id @pytest.mark.asyncio async def test_parser_not_recreated_with_constructor( @@ -7833,6 +7840,82 @@ class TestExternalUriSanitization: assert removed +# ================================================================ +# TEST CLASS: GCS offload path identity +# ================================================================ +class TestOffloadPathIdentity: + """Tests that offload paths come from the call, not shared parser state.""" + + @pytest.mark.asyncio + async def test_concurrent_parses_keep_their_own_identity(self): + """Two parses in flight at once do not write under each other's prefix.""" + paths = [] + first_upload_started = asyncio.Event() + + async def upload_content(data, mime_type, path): + paths.append(path) + if len(paths) == 1: + # Hold the first upload open until the second has begun, so both + # parses are suspended inside _parse_content_object at once. + first_upload_started.set() + await asyncio.sleep(0.05) + return f"gs://bucket/{path}" + + offloader = mock.MagicMock() + offloader.upload_content = upload_content + parser = bigquery_agent_analytics_plugin.HybridContentParser( + offloader=offloader, + trace_id="", + span_id="", + max_length=10, + ) + content = types.Content(parts=[types.Part(text="X" * 200)]) + + async def parse_as(trace_id, span_id): + return await parser.parse(content, trace_id=trace_id, span_id=span_id) + + task_a = asyncio.create_task(parse_as("trace-a", "span-a")) + await asyncio.wait_for(first_upload_started.wait(), timeout=5) + task_b = asyncio.create_task(parse_as("trace-b", "span-b")) + await asyncio.gather(task_a, task_b) + + assert len(paths) == 2 + assert sum("trace-a/span-a" in p for p in paths) == 1 + assert sum("trace-b/span-b" in p for p in paths) == 1 + assert parser.trace_id == "" + assert parser.span_id == "" + + @pytest.mark.asyncio + async def test_same_part_index_in_two_messages_does_not_collide(self): + """Two messages in one request get distinct object names.""" + paths = [] + + async def upload_content(data, mime_type, path): + paths.append(path) + return f"gs://bucket/{path}" + + offloader = mock.MagicMock() + offloader.upload_content = upload_content + parser = bigquery_agent_analytics_plugin.HybridContentParser( + offloader=offloader, + trace_id="t", + span_id="s", + max_length=10, + ) + llm_request = llm_request_lib.LlmRequest( + model="gemini-pro", + contents=[ + types.Content(parts=[types.Part(text="A" * 200)]), + types.Content(parts=[types.Part(text="B" * 200)]), + ], + ) + + await parser.parse(llm_request, trace_id="t1", span_id="s1") + + assert len(paths) == 2 + assert paths[0] != paths[1] + + # ================================================================ # TEST CLASS: AGENT_RESPONSE logging (Issue #87) # ================================================================