* feat(factory.py): add support for merging consecutive messages of one role when separated with empty message of another role

This commit is contained in:
nkvch
2024-05-07 12:51:30 +02:00
parent 30003afbf8
commit 7d7b59ff78
+36 -26
View File
@@ -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(