Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 9 additions & 2 deletions src/mistral_common/protocol/instruct/validator.py
Original file line number Diff line number Diff line change
Expand Up @@ -293,9 +293,16 @@ def _validate_last_message(self, message: UATS, continue_final_message: bool) ->
if continue_final_message:
raise InvalidMessageStructureException("Cannot continue final message in finetuning mode")
else:
bad_assistant = isinstance(message, AssistantMessage) and not message.prefix and not continue_final_message
bad_role = message.role not in {Roles.user, Roles.tool}
if bad_assistant and bad_role:
# A non-user/tool trailing message is only valid when it is an
# assistant message to be continued (prefix) or
# ``continue_final_message`` is set. The previous condition ANDed
# ``bad_role`` with an assistant-only flag, so a trailing message of
# another role (e.g. system) was never rejected here.
valid_trailing_assistant = isinstance(message, AssistantMessage) and (
message.prefix or continue_final_message
)
Comment on lines +302 to +304
if bad_role and not valid_trailing_assistant:
raise InvalidMessageStructureException(
f"Expected last role User or Tool (or Assistant with prefix or continue_final_message set to True) "
f"for serving but got {last_message_role}"
Expand Down
16 changes: 16 additions & 0 deletions tests/validation/test_chat_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -147,6 +147,22 @@ def test_ends_with_assistant(self, validator: MistralRequestValidator) -> None:
continue_final_message=False,
)

def test_ends_with_system(self, validator: MistralRequestValidator) -> None:
with pytest.raises(
InvalidMessageStructureException,
match=(
r"Expected last role User or Tool \(or Assistant with prefix or continue_final_message set to "
r"True\) for serving but got system"
),
Comment on lines +153 to +156
):
validator.validate_messages(
messages=[
UserMessage(content="foo"),
SystemMessage(content="foo"),
],
continue_final_message=False,
)

def test_assistant_prefix(self, validator: MistralRequestValidator) -> None:
validator.validate_messages(
messages=[
Expand Down