Skip to content

Commit 5b46641

Browse files
jsondaicopybara-github
authored andcommitted
fix: GenAI Client(evals) - Accept OpenAI-style and Message items in conversation history
PiperOrigin-RevId: 995377725
1 parent f6bf219 commit 5b46641

3 files changed

Lines changed: 171 additions & 34 deletions

File tree

‎agentplatform/_genai/_evals_data_converters.py‎

Lines changed: 35 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -57,6 +57,20 @@ def _create_placeholder_response_candidate(
5757
)
5858

5959

60+
def _openai_message_to_eval_message(
61+
turn_id: int, message: dict[str, Any]
62+
) -> types.evals.Message:
63+
"""Converts an OpenAI chat message into a conversation history message."""
64+
role = message.get("role", "user")
65+
return types.evals.Message(
66+
turn_id=str(turn_id),
67+
content=genai_types.Content(
68+
parts=[genai_types.Part(text=message.get("content", ""))], role=role
69+
),
70+
author=role,
71+
)
72+
73+
6074
class _GeminiEvalDataConverter(_evals_utils.EvalDataConverter):
6175
"""Converter for dataset in the Gemini format."""
6276

@@ -199,9 +213,11 @@ def convert(self, raw_data: list[dict[str, Any]]) -> types.EvaluationDataset:
199213
if not prompt_data:
200214
prompt_data = item.pop("source", None)
201215

202-
conversation_history_data = item.pop("conversation_history", None)
216+
history_column = "conversation_history"
217+
conversation_history_data = item.pop(history_column, None)
203218
if conversation_history_data is None:
204-
conversation_history_data = item.pop("history", None)
219+
history_column = "history"
220+
conversation_history_data = item.pop(history_column, None)
205221
response_data = item.pop("response", None)
206222
reference_data = item.pop("reference", None)
207223
system_instruction_data = item.pop("instruction", None)
@@ -241,6 +257,16 @@ def convert(self, raw_data: list[dict[str, Any]]) -> types.EvaluationDataset:
241257
content=content,
242258
)
243259
)
260+
elif isinstance(content, types.evals.Message):
261+
conversation_history.append(content)
262+
elif (
263+
isinstance(content, dict)
264+
and isinstance(content.get("content"), str)
265+
and isinstance(content.get("role", "user"), str)
266+
):
267+
conversation_history.append(
268+
_openai_message_to_eval_message(turn_id, content)
269+
)
244270
elif isinstance(content, dict):
245271
try:
246272
validated_content = genai_types.Content.model_validate(
@@ -254,18 +280,20 @@ def convert(self, raw_data: list[dict[str, Any]]) -> types.EvaluationDataset:
254280
)
255281
except ValidationError as e:
256282
logger.warning(
257-
"Item at index %s in 'history' column for case "
283+
"Item at index %s in '%s' column for case"
258284
" %s is a dict but could not be validated as"
259285
" genai_types.Content: %s",
260286
turn_id,
287+
history_column,
261288
eval_case_id,
262289
e,
263290
)
264291
else:
265292
logger.warning(
266-
"Invalid type in 'history' column for case %s at index %s. "
267-
"Expected genai_types.Content or dict, but got %s. "
268-
"Skipping this history item.",
293+
"Invalid type in '%s' column for case %s at index %s. "
294+
"Expected genai_types.Content, types.evals.Message or "
295+
"dict, but got %s. Skipping this history item.",
296+
history_column,
269297
eval_case_id,
270298
turn_id,
271299
type(content),
@@ -492,17 +520,7 @@ def _parse_messages(self, messages: list[dict[str, Any]]) -> tuple[
492520
messages = messages[1:]
493521

494522
for turn_id, msg in enumerate(messages):
495-
role = msg.get("role", "user")
496-
content = msg.get("content", "")
497-
conversation_history.append(
498-
types.evals.Message(
499-
turn_id=str(turn_id),
500-
content=genai_types.Content(
501-
parts=[genai_types.Part(text=content)], role=role
502-
),
503-
author=role,
504-
)
505-
)
523+
conversation_history.append(_openai_message_to_eval_message(turn_id, msg))
506524

507525
if conversation_history:
508526
last_message = conversation_history.pop()

‎tests/unit/agentplatform/genai/test_evals.py‎

Lines changed: 101 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -45,6 +45,9 @@
4545
types as agentplatform_genai_types,
4646
)
4747
from agentplatform._genai.types import common as common_types
48+
from vertexai._genai import (
49+
_evals_data_converters as vertexai_evals_data_converters,
50+
)
4851
from google.genai import client
4952
from google.genai import errors as genai_errors
5053
from google.genai import types as genai_types
@@ -61,6 +64,12 @@
6164

6265
pytestmark = pytest.mark.usefixtures("google_auth_mock")
6366

67+
_CONVERTER_MODULES = pytest.mark.parametrize(
68+
"converters",
69+
[_evals_data_converters, vertexai_evals_data_converters],
70+
ids=["agent_platform", "vertexai"],
71+
)
72+
6473

6574
class TestDropEmptyColumns:
6675
"""Unit tests for the _drop_empty_columns function."""
@@ -5778,6 +5787,84 @@ def test_convert_with_conversation_history_column_name(self):
57785787
eval_case.conversation_history[1].content.parts[0].text == "Old model msg"
57795788
)
57805789

5790+
@_CONVERTER_MODULES
5791+
def test_convert_openai_style_history(self, converters):
5792+
result_dataset = converters._FlattenEvalDataConverter().convert(
5793+
[
5794+
{
5795+
"prompt": "Code word?",
5796+
"response": "BLUE",
5797+
"conversation_history": [
5798+
{"content": "My code word is BLUE."},
5799+
{"role": "assistant", "content": "Noted."},
5800+
],
5801+
}
5802+
]
5803+
)
5804+
5805+
assert result_dataset.eval_cases[0].conversation_history == [
5806+
converters.types.evals.Message(
5807+
turn_id="0",
5808+
content=genai_types.Content(
5809+
parts=[genai_types.Part(text="My code word is BLUE.")], role="user"
5810+
),
5811+
author="user",
5812+
),
5813+
converters.types.evals.Message(
5814+
turn_id="1",
5815+
content=genai_types.Content(
5816+
parts=[genai_types.Part(text="Noted.")], role="assistant"
5817+
),
5818+
author="assistant",
5819+
),
5820+
]
5821+
5822+
@_CONVERTER_MODULES
5823+
def test_convert_message_history_items(self, converters):
5824+
history = [
5825+
converters.types.evals.Message(
5826+
turn_id="turn-0",
5827+
content=genai_types.Content(
5828+
parts=[genai_types.Part(text="My code word is BLUE.")], role="user"
5829+
),
5830+
author="user",
5831+
)
5832+
]
5833+
5834+
result_dataset = converters._FlattenEvalDataConverter().convert(
5835+
[{"prompt": "Code word?", "response": "BLUE", "history": history}]
5836+
)
5837+
5838+
assert result_dataset.eval_cases[0].conversation_history == history
5839+
5840+
@_CONVERTER_MODULES
5841+
@pytest.mark.parametrize("column", ["conversation_history", "history"])
5842+
@pytest.mark.parametrize(
5843+
"item,expected_warning",
5844+
[
5845+
(
5846+
{"role": "user", "text": "Hi"},
5847+
"Item at index 0 in '{column}' column for case eval_case_0 is a dict",
5848+
),
5849+
(
5850+
{"role": 1, "content": "Hi"},
5851+
"Item at index 0 in '{column}' column for case eval_case_0 is a dict",
5852+
),
5853+
(42, "Invalid type in '{column}' column for case eval_case_0 at index 0."),
5854+
],
5855+
ids=["invalid_dict_item", "non_str_role", "invalid_item_type"],
5856+
)
5857+
def test_convert_invalid_history_item_logs_warning(
5858+
self, converters, column, item, expected_warning, caplog
5859+
):
5860+
with caplog.at_level("WARNING", logger=converters.logger.name):
5861+
result_dataset = converters._FlattenEvalDataConverter().convert(
5862+
[{"prompt": "Hello", "response": "Hi", column: [item]}]
5863+
)
5864+
5865+
assert result_dataset.eval_cases[0].conversation_history == []
5866+
assert expected_warning.format(column=column) in caplog.text
5867+
57815868
def test_convert_missing_response_raises_value_error(self):
57825869
raw_data_df = pd.DataFrame({"prompt": ["Hello"]}) # Missing response
57835870
raw_data = raw_data_df.to_dict(orient="records")
@@ -6032,6 +6119,20 @@ def test_convert_skips_missing_request_or_response(self):
60326119
result_dataset = self.converter.convert(raw_data)
60336120
assert len(result_dataset.eval_cases) == 0
60346121

6122+
@pytest.mark.parametrize(
6123+
"message,role,text",
6124+
[({"role": "assistant", "content": "Hi"}, "assistant", "Hi"), ({}, "user", "")],
6125+
ids=["role_and_content", "defaults"],
6126+
)
6127+
def test_openai_message_to_eval_message(self, message, role, text):
6128+
assert _evals_data_converters._openai_message_to_eval_message(
6129+
1, message
6130+
) == agentplatform_genai_types.evals.Message(
6131+
turn_id="1",
6132+
content=genai_types.Content(parts=[genai_types.Part(text=text)], role=role),
6133+
author=role,
6134+
)
6135+
60356136

60366137
class TestObservabilityDataConverter:
60376138
"""Unit tests for the ObservabilityDataConverter class."""

‎vertexai/_genai/_evals_data_converters.py‎

Lines changed: 35 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -57,6 +57,20 @@ def _create_placeholder_response_candidate(
5757
)
5858

5959

60+
def _openai_message_to_eval_message(
61+
turn_id: int, message: dict[str, Any]
62+
) -> types.evals.Message:
63+
"""Converts an OpenAI chat message into a conversation history message."""
64+
role = message.get("role", "user")
65+
return types.evals.Message(
66+
turn_id=str(turn_id),
67+
content=genai_types.Content(
68+
parts=[genai_types.Part(text=message.get("content", ""))], role=role
69+
),
70+
author=role,
71+
)
72+
73+
6074
class _GeminiEvalDataConverter(_evals_utils.EvalDataConverter):
6175
"""Converter for dataset in the Gemini format."""
6276

@@ -199,9 +213,11 @@ def convert(self, raw_data: list[dict[str, Any]]) -> types.EvaluationDataset:
199213
if not prompt_data:
200214
prompt_data = item.pop("source", None)
201215

202-
conversation_history_data = item.pop("conversation_history", None)
216+
history_column = "conversation_history"
217+
conversation_history_data = item.pop(history_column, None)
203218
if conversation_history_data is None:
204-
conversation_history_data = item.pop("history", None)
219+
history_column = "history"
220+
conversation_history_data = item.pop(history_column, None)
205221
response_data = item.pop("response", None)
206222
reference_data = item.pop("reference", None)
207223
system_instruction_data = item.pop("instruction", None)
@@ -241,6 +257,16 @@ def convert(self, raw_data: list[dict[str, Any]]) -> types.EvaluationDataset:
241257
content=content,
242258
)
243259
)
260+
elif isinstance(content, types.evals.Message):
261+
conversation_history.append(content)
262+
elif (
263+
isinstance(content, dict)
264+
and isinstance(content.get("content"), str)
265+
and isinstance(content.get("role", "user"), str)
266+
):
267+
conversation_history.append(
268+
_openai_message_to_eval_message(turn_id, content)
269+
)
244270
elif isinstance(content, dict):
245271
try:
246272
validated_content = genai_types.Content.model_validate(
@@ -254,18 +280,20 @@ def convert(self, raw_data: list[dict[str, Any]]) -> types.EvaluationDataset:
254280
)
255281
except ValidationError as e:
256282
logger.warning(
257-
"Item at index %s in 'history' column for case "
283+
"Item at index %s in '%s' column for case"
258284
" %s is a dict but could not be validated as"
259285
" genai_types.Content: %s",
260286
turn_id,
287+
history_column,
261288
eval_case_id,
262289
e,
263290
)
264291
else:
265292
logger.warning(
266-
"Invalid type in 'history' column for case %s at index %s. "
267-
"Expected genai_types.Content or dict, but got %s. "
268-
"Skipping this history item.",
293+
"Invalid type in '%s' column for case %s at index %s. "
294+
"Expected genai_types.Content, types.evals.Message or "
295+
"dict, but got %s. Skipping this history item.",
296+
history_column,
269297
eval_case_id,
270298
turn_id,
271299
type(content),
@@ -492,17 +520,7 @@ def _parse_messages(self, messages: list[dict[str, Any]]) -> tuple[
492520
messages = messages[1:]
493521

494522
for turn_id, msg in enumerate(messages):
495-
role = msg.get("role", "user")
496-
content = msg.get("content", "")
497-
conversation_history.append(
498-
types.evals.Message(
499-
turn_id=str(turn_id),
500-
content=genai_types.Content(
501-
parts=[genai_types.Part(text=content)], role=role
502-
),
503-
author=role,
504-
)
505-
)
523+
conversation_history.append(_openai_message_to_eval_message(turn_id, msg))
506524

507525
if conversation_history:
508526
last_message = conversation_history.pop()

0 commit comments

Comments
 (0)