Add a converter module to house conversion from different LLM provider formats to Invariant API format.

This commit is contained in:
Hemang
2025-03-10 15:18:13 +01:00
committed by Hemang Sarkar
parent 38504b3e45
commit 7576f245bd
7 changed files with 225 additions and 109 deletions
View File
@@ -0,0 +1,96 @@
"""Converts the request and response formats from Anthropic to Invariant API format."""
def convert_anthropic_to_invariant_message_format(
messages: list[dict], keep_empty_tool_response: bool = False
) -> list[dict]:
"""Converts a list of messages from the Anthropic API to the Invariant API format."""
output = []
role_mapping = {
"system": lambda msg: {"role": "system", "content": msg["content"]},
"user": lambda msg: handle_user_message(msg, keep_empty_tool_response),
"assistant": handle_assistant_message,
}
for message in messages:
handler = role_mapping.get(message["role"])
if handler:
output.extend(handler(message))
return output
def handle_user_message(message, keep_empty_tool_response):
"""Handle the user message from the Anthropic API"""
output = []
content = message["content"]
if isinstance(content, list):
user_content = []
for sub_message in content:
if sub_message["type"] == "tool_result":
if sub_message["content"]:
output.append(
{
"role": "tool",
"content": sub_message["content"],
"tool_id": sub_message["tool_use_id"],
}
)
elif keep_empty_tool_response and any(sub_message.values()):
output.append(
{
"role": "tool",
"content": {"is_error": True}
if sub_message["is_error"]
else {},
"tool_id": sub_message["tool_use_id"],
}
)
elif sub_message["type"] == "text":
user_content.append({"type": "text", "text": sub_message["text"]})
elif sub_message["type"] == "image":
user_content.append(
{
"type": "image_url",
"image_url": {
"url": "data:"
+ sub_message["source"]["media_type"]
+ ";base64,"
+ sub_message["source"]["data"],
},
},
)
if user_content:
output.append({"role": "user", "content": user_content})
else:
output.append({"role": "user", "content": content})
return output
def handle_assistant_message(message):
"""Handle the assistant message from the Anthropic API"""
output = []
if isinstance(message["content"], list):
for sub_message in message["content"]:
if sub_message["type"] == "text":
output.append({"role": "assistant", "content": sub_message.get("text")})
elif sub_message["type"] == "tool_use":
output.append(
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"tool_id": sub_message.get("id"),
"type": "function",
"function": {
"name": sub_message.get("name"),
"arguments": sub_message.get("input"),
},
}
],
}
)
else:
output.append({"role": "assistant", "content": message["content"]})
return output
+111
View File
@@ -0,0 +1,111 @@
"""Converts the request and response formats from Gemini to Invariant API format."""
def convert_request(request: dict) -> list[dict]:
"""Converts the request from Gemini API to Invariant API format."""
openai_messages = []
if "systemInstruction" in request:
system_content = " ".join(
part.get("text", "")
for part in request["systemInstruction"].get("parts", [])
)
openai_messages.append({"role": "system", "content": system_content})
if "contents" in request:
for content in request["contents"]:
role = content.get("role", "")
if role == "user":
message_content = []
for part in content.get("parts", []):
if "text" in part:
message_content.append({"type": "text", "text": part["text"]})
elif "inlineData" in part:
message_content.append(
{
"type": "image",
"image_url": {
"url": f"data:{part['inlineData']['mime_type']};base64,{part['inlineData']['data']}"
},
}
)
elif "functionResponse" in part:
openai_messages.append(
{
"role": "tool",
"tool_name": part["functionResponse"]["name"],
"content": part["functionResponse"]["response"].get(
"result", {}
),
}
)
if message_content:
openai_messages.append(
{
"role": "user",
"content": message_content
if len(message_content) > 1
else message_content[0],
}
)
elif role == "model":
for part in content.get("parts", []):
if "text" in part:
openai_messages.append(
{"role": "assistant", "content": part["text"]}
)
elif "functionCall" in part:
openai_messages.append(
{
"role": "assistant",
"tool_calls": [
{
"type": "function",
"function": {
"name": part["functionCall"]["name"],
"arguments": part["functionCall"].get(
"args", {}
),
},
}
],
}
)
return openai_messages
def convert_response(response: dict) -> list[dict]:
"""Converts the response from Gemini API to Invariant API format."""
openai_messages = []
if "candidates" in response:
for candidate in response["candidates"]:
candidate_content = candidate.get("content", {})
role = candidate_content.get("role", "")
if role == "model":
for part in candidate_content.get("parts", []):
if "text" in part:
openai_messages.append(
{"role": "assistant", "content": part["text"]}
)
elif "functionCall" in part:
openai_messages.append(
{
"role": "assistant",
"tool_calls": [
{
"type": "function",
"function": {
"name": part["functionCall"]["name"],
"arguments": part["functionCall"].get(
"args", {}
),
},
}
],
}
)
return openai_messages
+3 -98
View File
@@ -12,11 +12,11 @@ from common.constants import (
IGNORED_HEADERS,
)
from integrations.explorer import push_trace
from converters.anthropic_to_invariant import convert_anthropic_to_invariant_message_format
from common.authorization import extract_authorization_from_headers
gateway = APIRouter()
MISSING_INVARIANT_AUTH_API_KEY = "Missing invariant authorization header"
MISSING_ANTHROPIC_AUTH_HEADER = "Missing Anthropic authorization header"
FAILED_TO_PUSH_TRACE = "Failed to push trace to the dataset: "
END_REASONS = ["end_turn", "max_tokens", "stop_sequence"]
@@ -92,10 +92,10 @@ async def push_to_explorer(
messages = request_body.get("messages", [])
messages += [merged_response]
transformed_messages = convert_anthropic_to_invariant_message_format(messages)
converted_messages = convert_anthropic_to_invariant_message_format(messages)
_ = await push_trace(
dataset_name=dataset_name,
messages=[transformed_messages],
messages=[converted_messages],
invariant_authorization=invariant_authorization,
)
@@ -222,98 +222,3 @@ def update_merged_response(text_json, merged_response):
merged_response[-1]["content"] += text_json.get("delta").get("partial_json")
elif text_json.get("type") == MESSGAE_DELTA:
merged_response[-1]["stop_reason"] = text_json.get("delta").get("stop_reason")
def convert_anthropic_to_invariant_message_format(
messages: list[dict], keep_empty_tool_response: bool = False
) -> list[dict]:
"""Converts a list of messages from the Anthropic API to the Invariant API format."""
output = []
role_mapping = {
"system": lambda msg: {"role": "system", "content": msg["content"]},
"user": lambda msg: handle_user_message(msg, keep_empty_tool_response),
"assistant": lambda msg: handle_assistant_message(msg),
}
for message in messages:
handler = role_mapping.get(message["role"])
if handler:
output.extend(handler(message))
return output
def handle_user_message(message, keep_empty_tool_response):
"""Handle the user message from the Anthropic API"""
output = []
content = message["content"]
if isinstance(content, list):
user_content = []
for sub_message in content:
if sub_message["type"] == "tool_result":
if sub_message["content"]:
output.append(
{
"role": "tool",
"content": sub_message["content"],
"tool_id": sub_message["tool_use_id"],
}
)
elif keep_empty_tool_response and any(sub_message.values()):
output.append(
{
"role": "tool",
"content": {"is_error": True}
if sub_message["is_error"]
else {},
"tool_id": sub_message["tool_use_id"],
}
)
elif sub_message["type"] == "text":
user_content.append({"type": "text", "text": sub_message["text"]})
elif sub_message["type"] == "image":
user_content.append(
{
"type": "image_url",
"image_url": {
"url": "data:"
+ sub_message["source"]["media_type"]
+ ";base64,"
+ sub_message["source"]["data"],
},
},
)
if user_content:
output.append({"role": "user", "content": user_content})
else:
output.append({"role": "user", "content": content})
return output
def handle_assistant_message(message):
"""Handle the assistant message from the Anthropic API"""
output = []
if isinstance(message["content"], list):
for sub_message in message["content"]:
if sub_message["type"] == "text":
output.append({"role": "assistant", "content": sub_message.get("text")})
elif sub_message["type"] == "tool_use":
output.append(
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"tool_id": sub_message.get("id"),
"type": "function",
"function": {
"name": sub_message.get("name"),
"arguments": sub_message.get("input"),
},
}
],
}
)
else:
output.append({"role": "assistant", "content": message["content"]})
return output
+4 -4
View File
@@ -12,6 +12,7 @@ from common.constants import (
IGNORED_HEADERS,
)
from common.authorization import extract_authorization_from_headers
from converters.gemini_to_invariant import convert_request, convert_response
from integrations.explorer import push_trace
gateway = APIRouter()
@@ -118,12 +119,11 @@ async def push_to_explorer(
invariant_authorization: str,
) -> None:
"""Pushes the full trace to the Invariant Explorer"""
# Combine the messages from the request body and the choices from the Gemini response
messages = request_body.get("messages", [])
messages += [choice["message"] for choice in merged_response.get("choices", [])]
converted_requests = convert_request(request_body)
converted_responses = convert_response(merged_response)
_ = await push_trace(
dataset_name=dataset_name,
messages=[messages],
messages=[converted_requests + converted_responses],
invariant_authorization=invariant_authorization,
)
-1
View File
@@ -16,7 +16,6 @@ from common.authorization import extract_authorization_from_headers
gateway = APIRouter()
MISSING_INVARIANT_AUTH_API_KEY = "Missing invariant api key"
MISSING_AUTH_HEADER = "Missing authorization header"
FINISH_REASON_TO_PUSH_TRACE = ["stop", "length", "content_filter"]
OPENAI_AUTHORIZATION_HEADER = "authorization"
+11 -6
View File
@@ -8,6 +8,7 @@ from playwright.async_api import async_playwright
@pytest.fixture
def gateway_url():
"""Get the gateway URL from the environment variable"""
if "INVARIANT_GATEWAY_API_URL" in os.environ:
return os.environ["INVARIANT_GATEWAY_API_URL"]
raise ValueError("Please set the INVARIANT_GATEWAY_API_URL environment variable")
@@ -15,6 +16,7 @@ def gateway_url():
@pytest.fixture
def explorer_api_url():
"""Get the explorer API URL from the environment variable"""
if "INVARIANT_API_URL" in os.environ:
return os.environ["INVARIANT_API_URL"]
raise ValueError("Please set the INVARIANT_API_URL environment variable")
@@ -22,19 +24,22 @@ def explorer_api_url():
@pytest.fixture
async def playwright(scope="session"):
"""Fixture to create a Playwright instance"""
async with async_playwright() as playwright_instance:
yield playwright_instance
@pytest.fixture
async def browser(playwright, scope="session"):
browser = await playwright.firefox.launch(headless=True)
yield browser
await browser.close()
"""Fixture to create a browser instance"""
firefox_browser = await playwright.firefox.launch(headless=True)
yield firefox_browser
await firefox_browser.close()
@pytest.fixture
async def context(browser):
context = await browser.new_context(ignore_https_errors=True)
yield context
await context.close()
"""Fixture to create a browser context"""
browser_context = await browser.new_context(ignore_https_errors=True)
yield browser_context
await browser_context.close()