From 7d7b59ff78a44c9375167dabb79bb16b56a17c91 Mon Sep 17 00:00:00 2001 From: nkvch Date: Mon, 6 May 2024 16:59:13 +0200 Subject: [PATCH] * feat(factory.py): add support for merging consecutive messages of one role when separated with empty message of another role --- litellm/llms/prompt_templates/factory.py | 62 ++++++++++++++---------- 1 file changed, 36 insertions(+), 26 deletions(-) diff --git a/litellm/llms/prompt_templates/factory.py b/litellm/llms/prompt_templates/factory.py index 0820303685..6e8589d581 100644 --- a/litellm/llms/prompt_templates/factory.py +++ b/litellm/llms/prompt_templates/factory.py @@ -1,27 +1,23 @@ -from enum import Enum -import requests, traceback -import json, re, xml.etree.ElementTree as ET -from jinja2 import Template, exceptions, meta, BaseLoader -from jinja2.sandbox import ImmutableSandboxedEnvironment -from typing import ( - Any, - List, - Mapping, - MutableMapping, - Optional, - Sequence, -) -import litellm -from litellm.types.completion import ( - ChatCompletionUserMessageParam, - ChatCompletionSystemMessageParam, - ChatCompletionMessageParam, - ChatCompletionFunctionMessageParam, - ChatCompletionMessageToolCallParam, - ChatCompletionToolMessageParam, -) -from litellm.types.llms.anthropic import * +import json +import re +import traceback import uuid +import xml.etree.ElementTree as ET +from enum import Enum +from typing import Any, List, Mapping, MutableMapping, Optional, Sequence + +import requests +from jinja2 import BaseLoader, Template, exceptions, meta +from jinja2.sandbox import ImmutableSandboxedEnvironment + +import litellm +from litellm.types.completion import (ChatCompletionFunctionMessageParam, + ChatCompletionMessageParam, + ChatCompletionMessageToolCallParam, + ChatCompletionSystemMessageParam, + ChatCompletionToolMessageParam, + ChatCompletionUserMessageParam) +from litellm.types.llms.anthropic import * def default_pt(messages): @@ -603,9 +599,10 @@ def construct_tool_use_system_prompt( def convert_url_to_base64(url): - import requests import base64 + import requests + for _ in range(3): try: response = requests.get(url) @@ -984,6 +981,7 @@ def anthropic_messages_pt(messages: list): new_messages = [] msg_i = 0 tool_use_param = False + merge_with_previous = False while msg_i < len(messages): user_content = [] init_msg_i = msg_i @@ -1016,7 +1014,13 @@ def anthropic_messages_pt(messages: list): msg_i += 1 if user_content: - new_messages.append({"role": "user", "content": user_content}) + if merge_with_previous: + new_messages[-1]["content"].extend(user_content) + merge_with_previous = False + else: + new_messages.append({"role": "user", "content": user_content}) + else: + merge_with_previous = True assistant_content = [] ## MERGE CONSECUTIVE ASSISTANT CONTENT ## @@ -1044,7 +1048,13 @@ def anthropic_messages_pt(messages: list): msg_i += 1 if assistant_content: - new_messages.append({"role": "assistant", "content": assistant_content}) + if merge_with_previous: + new_messages[-1]["content"].extend(assistant_content) + merge_with_previous = False + else: + new_messages.append({"role": "assistant", "content": assistant_content}) + else: + merge_with_previous = True if msg_i == init_msg_i: # prevent infinite loops raise Exception(