[router]Replace requests lib with openai in e2e_response_api (#13293)
This commit is contained in:
@@ -11,6 +11,8 @@ import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
import openai
|
||||
|
||||
# Add e2e_response_api directory for imports
|
||||
_TEST_DIR = Path(__file__).parent.parent
|
||||
sys.path.insert(0, str(_TEST_DIR))
|
||||
@@ -51,6 +53,7 @@ class TestGrpcBackend(StateManagementTests, MCPTests, StructuredOutputBaseTest):
|
||||
)
|
||||
|
||||
cls.base_url = cls.cluster["base_url"]
|
||||
cls.client = openai.Client(api_key=cls.api_key, base_url=cls.base_url + "/v1")
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
@@ -64,8 +67,7 @@ class TestGrpcBackend(StateManagementTests, MCPTests, StructuredOutputBaseTest):
|
||||
|
||||
def test_structured_output_json_schema(self):
|
||||
"""Override with simpler schema for Llama model (complex schemas not well supported)."""
|
||||
data = {
|
||||
"model": self.model,
|
||||
params = {
|
||||
"input": [
|
||||
{
|
||||
"role": "system",
|
||||
@@ -89,28 +91,26 @@ class TestGrpcBackend(StateManagementTests, MCPTests, StructuredOutputBaseTest):
|
||||
},
|
||||
}
|
||||
|
||||
create_resp = self.make_request("/v1/responses", "POST", data)
|
||||
self.assertEqual(create_resp.status_code, 200)
|
||||
|
||||
create_data = create_resp.json()
|
||||
self.assertIn("id", create_data)
|
||||
self.assertIn("output", create_data)
|
||||
self.assertIn("text", create_data)
|
||||
create_resp = self.create_response(**params)
|
||||
self.assertIsNone(create_resp.error)
|
||||
self.assertIsNotNone(create_resp.id)
|
||||
self.assertIsNotNone(create_resp.output)
|
||||
self.assertIsNotNone(create_resp.text)
|
||||
|
||||
# Verify text format was echoed back correctly
|
||||
self.assertIn("format", create_data["text"])
|
||||
self.assertEqual(create_data["text"]["format"]["type"], "json_schema")
|
||||
self.assertEqual(create_data["text"]["format"]["name"], "math_answer")
|
||||
self.assertIn("schema", create_data["text"]["format"])
|
||||
self.assertIsNotNone(create_resp.text.format)
|
||||
self.assertEqual(create_resp.text.format.type, "json_schema")
|
||||
self.assertEqual(create_resp.text.format.name, "math_answer")
|
||||
self.assertIsNotNone(create_resp.text.format.schema_)
|
||||
|
||||
# Find the message output
|
||||
output_text = next(
|
||||
(
|
||||
content.get("text", "")
|
||||
for item in create_data.get("output", [])
|
||||
if item.get("type") == "message"
|
||||
for content in item.get("content", [])
|
||||
if content.get("type") == "output_text"
|
||||
content.text
|
||||
for item in create_resp.output
|
||||
if item.type == "message"
|
||||
for content in item.content
|
||||
if content.type == "output_text"
|
||||
),
|
||||
None,
|
||||
)
|
||||
@@ -154,6 +154,7 @@ class TestGrpcHarmonyBackend(
|
||||
)
|
||||
|
||||
cls.base_url = cls.cluster["base_url"]
|
||||
cls.client = openai.Client(api_key=cls.api_key, base_url=cls.base_url + "/v1")
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
|
||||
@@ -13,6 +13,8 @@ import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
import openai
|
||||
|
||||
# Add e2e_response_api directory for imports
|
||||
_TEST_DIR = Path(__file__).parent.parent
|
||||
sys.path.insert(0, str(_TEST_DIR))
|
||||
@@ -52,6 +54,7 @@ class TestOpenaiBackend(
|
||||
)
|
||||
|
||||
cls.base_url = cls.cluster["base_url"]
|
||||
cls.client = openai.Client(api_key=cls.api_key, base_url=cls.base_url + "/v1")
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
@@ -93,6 +96,7 @@ class TestXaiBackend(StateManagementTests):
|
||||
)
|
||||
|
||||
cls.base_url = cls.cluster["base_url"]
|
||||
cls.client = openai.Client(api_key=cls.api_key, base_url=cls.base_url + "/v1")
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
|
||||
@@ -5,14 +5,16 @@ This module provides base test classes that can be reused across different backe
|
||||
(OpenAI, XAI, gRPC) with common test logic.
|
||||
"""
|
||||
|
||||
import json
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
import time
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
from typing import Optional, Union
|
||||
|
||||
import requests
|
||||
import openai
|
||||
from openai.types import conversations, responses
|
||||
|
||||
# Add current directory for local imports
|
||||
_TEST_DIR = Path(__file__).parent
|
||||
@@ -28,44 +30,11 @@ class ResponseAPIBaseTest(CustomTestCase):
|
||||
base_url: str = None
|
||||
api_key: str = None
|
||||
model: str = None
|
||||
|
||||
def make_request(
|
||||
self,
|
||||
endpoint: str,
|
||||
method: str = "POST",
|
||||
json_data: Optional[dict] = None,
|
||||
params: Optional[dict] = None,
|
||||
) -> requests.Response:
|
||||
"""
|
||||
Make HTTP request to router.
|
||||
|
||||
Args:
|
||||
endpoint: Endpoint path (e.g., "/v1/responses")
|
||||
method: HTTP method (GET, POST, DELETE)
|
||||
json_data: JSON body for POST requests
|
||||
params: Query parameters
|
||||
|
||||
Returns:
|
||||
requests.Response object
|
||||
"""
|
||||
url = f"{self.base_url}{endpoint}"
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if self.api_key:
|
||||
headers["Authorization"] = f"Bearer {self.api_key}"
|
||||
|
||||
if method == "POST":
|
||||
resp = requests.post(url, json=json_data, headers=headers, params=params)
|
||||
elif method == "GET":
|
||||
resp = requests.get(url, headers=headers, params=params)
|
||||
elif method == "DELETE":
|
||||
resp = requests.delete(url, headers=headers, params=params)
|
||||
else:
|
||||
raise ValueError(f"Unsupported method: {method}")
|
||||
return resp
|
||||
client: openai.OpenAI = None
|
||||
|
||||
def create_response(
|
||||
self,
|
||||
input_text: str,
|
||||
input: Union[str, responses.ResponseInputParam],
|
||||
instructions: Optional[str] = None,
|
||||
stream: bool = False,
|
||||
max_output_tokens: Optional[int] = None,
|
||||
@@ -75,12 +44,12 @@ class ResponseAPIBaseTest(CustomTestCase):
|
||||
tools: Optional[list] = None,
|
||||
background: bool = False,
|
||||
**kwargs,
|
||||
) -> requests.Response:
|
||||
) -> responses.Response | openai.Stream[responses.ResponseStreamEvent]:
|
||||
"""
|
||||
Create a response via POST /v1/responses.
|
||||
|
||||
Args:
|
||||
input_text: User input
|
||||
input: User input
|
||||
instructions: Optional system instructions
|
||||
stream: Whether to stream response
|
||||
max_output_tokens: Optional max tokens to generate
|
||||
@@ -92,178 +61,128 @@ class ResponseAPIBaseTest(CustomTestCase):
|
||||
**kwargs: Additional request parameters
|
||||
|
||||
Returns:
|
||||
requests.Response object
|
||||
Response object for non-stream request
|
||||
ResponseStreamEvent for stream request
|
||||
"""
|
||||
data = {
|
||||
params = {
|
||||
"model": self.model,
|
||||
"input": input_text,
|
||||
"input": input,
|
||||
"stream": stream,
|
||||
**kwargs,
|
||||
}
|
||||
|
||||
if instructions:
|
||||
data["instructions"] = instructions
|
||||
params["instructions"] = instructions
|
||||
|
||||
if max_output_tokens is not None:
|
||||
data["max_output_tokens"] = max_output_tokens
|
||||
params["max_output_tokens"] = max_output_tokens
|
||||
|
||||
if temperature is not None:
|
||||
data["temperature"] = temperature
|
||||
params["temperature"] = temperature
|
||||
|
||||
if previous_response_id:
|
||||
data["previous_response_id"] = previous_response_id
|
||||
params["previous_response_id"] = previous_response_id
|
||||
|
||||
if conversation:
|
||||
data["conversation"] = conversation
|
||||
params["conversation"] = conversation
|
||||
|
||||
if tools:
|
||||
data["tools"] = tools
|
||||
params["tools"] = tools
|
||||
|
||||
if background:
|
||||
data["background"] = background
|
||||
params["background"] = background
|
||||
|
||||
if stream:
|
||||
# For streaming, we need to handle SSE
|
||||
return self._create_streaming_response(data)
|
||||
else:
|
||||
return self.make_request("/v1/responses", "POST", data)
|
||||
return self.client.responses.create(**params)
|
||||
|
||||
def _create_streaming_response(self, data: dict) -> requests.Response:
|
||||
"""Handle streaming response creation."""
|
||||
url = f"{self.base_url}/v1/responses"
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if self.api_key:
|
||||
headers["Authorization"] = f"Bearer {self.api_key}"
|
||||
|
||||
# Return response object with stream=True
|
||||
return requests.post(url, json=data, headers=headers, stream=True)
|
||||
|
||||
def get_response(self, response_id: str) -> requests.Response:
|
||||
def get_response(
|
||||
self, response_id: str
|
||||
) -> responses.Response | openai.Stream[responses.ResponseStreamEvent]:
|
||||
"""Get response by ID via GET /v1/responses/{response_id}."""
|
||||
return self.make_request(f"/v1/responses/{response_id}", "GET")
|
||||
return self.client.responses.retrieve(response_id=response_id)
|
||||
|
||||
def delete_response(self, response_id: str) -> requests.Response:
|
||||
def delete_response(self, response_id: str) -> None:
|
||||
"""Delete response by ID via DELETE /v1/responses/{response_id}."""
|
||||
return self.make_request(f"/v1/responses/{response_id}", "DELETE")
|
||||
return self.client.responses.delete(response_id=response_id)
|
||||
|
||||
def cancel_response(self, response_id: str) -> requests.Response:
|
||||
def cancel_response(self, response_id: str) -> responses.Response:
|
||||
"""Cancel response by ID via POST /v1/responses/{response_id}/cancel."""
|
||||
return self.make_request(f"/v1/responses/{response_id}/cancel", "POST", {})
|
||||
return self.client.responses.cancel(response_id=response_id)
|
||||
|
||||
def get_response_input_items(self, response_id: str) -> requests.Response:
|
||||
def get_response_input_items(
|
||||
self, response_id: str
|
||||
) -> openai.pagination.SyncCursorPage[responses.ResponseItem]:
|
||||
"""Get response input items via GET /v1/responses/{response_id}/input_items."""
|
||||
return self.make_request(f"/v1/responses/{response_id}/input_items", "GET")
|
||||
return self.client.responses.input_items.list(response_id=response_id)
|
||||
|
||||
def create_conversation(self, metadata: Optional[dict] = None) -> requests.Response:
|
||||
def create_conversation(
|
||||
self, metadata: Optional[dict] = None
|
||||
) -> conversations.Conversation:
|
||||
"""Create conversation via POST /v1/conversations."""
|
||||
data = {}
|
||||
params = {}
|
||||
if metadata:
|
||||
data["metadata"] = metadata
|
||||
return self.make_request("/v1/conversations", "POST", data)
|
||||
params["metadata"] = metadata
|
||||
return self.client.conversations.create(**params)
|
||||
|
||||
def get_conversation(self, conversation_id: str) -> requests.Response:
|
||||
def get_conversation(self, conversation_id: str) -> conversations.Conversation:
|
||||
"""Get conversation by ID via GET /v1/conversations/{conversation_id}."""
|
||||
return self.make_request(f"/v1/conversations/{conversation_id}", "GET")
|
||||
return self.client.conversations.retrieve(conversation_id=conversation_id)
|
||||
|
||||
def update_conversation(
|
||||
self, conversation_id: str, metadata: dict
|
||||
) -> requests.Response:
|
||||
) -> conversations.Conversation:
|
||||
"""Update conversation via POST /v1/conversations/{conversation_id}."""
|
||||
return self.make_request(
|
||||
f"/v1/conversations/{conversation_id}", "POST", {"metadata": metadata}
|
||||
return self.client.conversations.update(
|
||||
conversation_id=conversation_id, metadata=metadata
|
||||
)
|
||||
|
||||
def delete_conversation(self, conversation_id: str) -> requests.Response:
|
||||
def delete_conversation(
|
||||
self, conversation_id: str
|
||||
) -> conversations.ConversationDeletedResource:
|
||||
"""Delete conversation via DELETE /v1/conversations/{conversation_id}."""
|
||||
return self.make_request(f"/v1/conversations/{conversation_id}", "DELETE")
|
||||
return self.client.conversations.delete(conversation_id=conversation_id)
|
||||
|
||||
def list_conversation_items(
|
||||
self,
|
||||
conversation_id: str,
|
||||
limit: Optional[int] = None,
|
||||
after: Optional[str] = None,
|
||||
before: Optional[str] = None,
|
||||
order: str = "asc",
|
||||
) -> requests.Response:
|
||||
) -> openai.pagination.SyncConversationCursorPage[conversations.ConversationItem]:
|
||||
"""List conversation items via GET /v1/conversations/{conversation_id}/items."""
|
||||
params = {"order": order}
|
||||
params = {"conversation_id": conversation_id, "order": order}
|
||||
if limit:
|
||||
params["limit"] = limit
|
||||
if after:
|
||||
params["after"] = after
|
||||
if before:
|
||||
params["before"] = before
|
||||
return self.make_request(
|
||||
f"/v1/conversations/{conversation_id}/items", "GET", params=params
|
||||
)
|
||||
return self.client.conversations.items.list(**params)
|
||||
|
||||
def create_conversation_items(
|
||||
self, conversation_id: str, items: list
|
||||
) -> requests.Response:
|
||||
) -> conversations.ConversationItemList:
|
||||
"""Create conversation items via POST /v1/conversations/{conversation_id}/items."""
|
||||
return self.make_request(
|
||||
f"/v1/conversations/{conversation_id}/items", "POST", {"items": items}
|
||||
return self.client.conversations.items.create(
|
||||
conversation_id=conversation_id, items=items
|
||||
)
|
||||
|
||||
def get_conversation_item(
|
||||
self, conversation_id: str, item_id: str
|
||||
) -> requests.Response:
|
||||
) -> conversations.ConversationItem:
|
||||
"""Get conversation item via GET /v1/conversations/{conversation_id}/items/{item_id}."""
|
||||
return self.make_request(
|
||||
f"/v1/conversations/{conversation_id}/items/{item_id}", "GET"
|
||||
return self.client.conversations.items.retrieve(
|
||||
conversation_id=conversation_id, item_id=item_id
|
||||
)
|
||||
|
||||
def delete_conversation_item(
|
||||
self, conversation_id: str, item_id: str
|
||||
) -> requests.Response:
|
||||
) -> conversations.Conversation:
|
||||
"""Delete conversation item via DELETE /v1/conversations/{conversation_id}/items/{item_id}."""
|
||||
return self.make_request(
|
||||
f"/v1/conversations/{conversation_id}/items/{item_id}", "DELETE"
|
||||
return self.client.conversations.items.delete(
|
||||
conversation_id=conversation_id, item_id=item_id
|
||||
)
|
||||
|
||||
def parse_sse_events(self, response: requests.Response) -> list:
|
||||
"""
|
||||
Parse Server-Sent Events from streaming response.
|
||||
|
||||
Args:
|
||||
response: requests.Response with stream=True
|
||||
|
||||
Returns:
|
||||
List of event dictionaries with 'event' and 'data' keys
|
||||
"""
|
||||
events = []
|
||||
current_event = None
|
||||
|
||||
for line in response.iter_lines():
|
||||
if not line:
|
||||
# Empty line signals end of event
|
||||
if current_event and current_event.get("data"):
|
||||
events.append(current_event)
|
||||
current_event = None
|
||||
continue
|
||||
|
||||
line = line.decode("utf-8")
|
||||
|
||||
if line.startswith("event:"):
|
||||
current_event = {"event": line[6:].strip()}
|
||||
elif line.startswith("data:"):
|
||||
if current_event is None:
|
||||
current_event = {}
|
||||
data_str = line[5:].strip()
|
||||
try:
|
||||
current_event["data"] = json.loads(data_str)
|
||||
except json.JSONDecodeError:
|
||||
current_event["data"] = data_str
|
||||
|
||||
# Don't forget the last event if stream ends without empty line
|
||||
if current_event and current_event.get("data"):
|
||||
events.append(current_event)
|
||||
|
||||
return events
|
||||
|
||||
def wait_for_background_task(
|
||||
self, response_id: str, timeout: int = 30, poll_interval: float = 0.5
|
||||
) -> dict:
|
||||
) -> responses.Response:
|
||||
"""
|
||||
Wait for background task to complete.
|
||||
|
||||
@@ -283,17 +202,15 @@ class ResponseAPIBaseTest(CustomTestCase):
|
||||
|
||||
while time.time() - start_time < timeout:
|
||||
resp = self.get_response(response_id)
|
||||
self.assertEqual(resp.status_code, 200)
|
||||
self.assertIsNone(resp.error)
|
||||
self.assertEqual(resp.id, response_id)
|
||||
|
||||
data = resp.json()
|
||||
status = data.get("status")
|
||||
status = resp.status
|
||||
|
||||
if status == "completed":
|
||||
return data
|
||||
return resp
|
||||
elif status == "failed":
|
||||
raise AssertionError(
|
||||
f"Background task failed: {data.get('error', 'Unknown error')}"
|
||||
)
|
||||
raise AssertionError(f"Background task failed: {resp.error}")
|
||||
elif status == "cancelled":
|
||||
raise AssertionError("Background task was cancelled")
|
||||
|
||||
@@ -310,31 +227,29 @@ class StateManagementBaseTest(ResponseAPIBaseTest):
|
||||
def test_basic_response_creation(self):
|
||||
"""Test basic response creation without state."""
|
||||
resp = self.create_response("What is 2+2?", max_output_tokens=50)
|
||||
self.assertEqual(resp.status_code, 200)
|
||||
|
||||
data = resp.json()
|
||||
self.assertIn("id", data)
|
||||
self.assertIn("output", data)
|
||||
self.assertEqual(data["status"], "completed")
|
||||
self.assertIn("usage", data)
|
||||
self.assertIsNotNone(resp.id)
|
||||
self.assertIsNone(resp.error)
|
||||
self.assertEqual(resp.status, "completed")
|
||||
self.assertGreater(len(resp.output_text), 0)
|
||||
self.assertGreater(resp.usage.input_tokens, 0)
|
||||
self.assertGreater(resp.usage.output_tokens, 0)
|
||||
self.assertGreater(resp.usage.total_tokens, 0)
|
||||
|
||||
def test_streaming_response(self):
|
||||
"""Test streaming response."""
|
||||
resp = self.create_response("Count to 5", stream=True, max_output_tokens=50)
|
||||
self.assertEqual(resp.status_code, 200)
|
||||
|
||||
events = self.parse_sse_events(resp)
|
||||
self.assertGreater(len(events), 0)
|
||||
|
||||
# Check for response.created event
|
||||
created_events = [e for e in events if e.get("event") == "response.created"]
|
||||
events = [event for event in resp]
|
||||
created_events = [event for event in events if event.type == "response.created"]
|
||||
self.assertGreater(len(created_events), 0)
|
||||
|
||||
# Check for final completed event or in_progress events
|
||||
self.assertTrue(
|
||||
any(
|
||||
e.get("event") in ["response.completed", "response.in_progress"]
|
||||
for e in events
|
||||
event.type in ["response.completed", "response.in_progress"]
|
||||
for event in events
|
||||
)
|
||||
)
|
||||
|
||||
@@ -346,41 +261,40 @@ class ResponseCRUDBaseTest(ResponseAPIBaseTest):
|
||||
"""Test creating response and retrieving it."""
|
||||
# Create response
|
||||
create_resp = self.create_response("Hello, world!")
|
||||
self.assertEqual(create_resp.status_code, 200)
|
||||
|
||||
create_data = create_resp.json()
|
||||
response_id = create_data["id"]
|
||||
self.assertIsNotNone(create_resp.id)
|
||||
self.assertIsNone(create_resp.error)
|
||||
self.assertEqual(create_resp.status, "completed")
|
||||
self.assertGreater(len(create_resp.output_text), 0)
|
||||
response_id = create_resp.id
|
||||
|
||||
# Get response
|
||||
get_resp = self.get_response(response_id)
|
||||
self.assertEqual(get_resp.status_code, 200)
|
||||
self.assertIsNone(get_resp.error)
|
||||
self.assertEqual(get_resp.id, response_id)
|
||||
self.assertEqual(get_resp.status, "completed")
|
||||
|
||||
get_data = get_resp.json()
|
||||
self.assertEqual(get_data["id"], response_id)
|
||||
self.assertEqual(get_data["status"], "completed")
|
||||
|
||||
input_resp = self.get_response_input_items(get_data["id"])
|
||||
self.assertEqual(input_resp.status_code, 200)
|
||||
input_data = input_resp.json()
|
||||
self.assertIn("data", input_data)
|
||||
self.assertGreater(len(input_data["data"]), 0)
|
||||
input_resp = self.get_response_input_items(get_resp.id)
|
||||
self.assertIsNotNone(input_resp.data)
|
||||
self.assertGreater(len(input_resp.data), 0)
|
||||
|
||||
@unittest.skip("TODO: Add delete response feature")
|
||||
def test_delete_response(self):
|
||||
"""Test deleting response."""
|
||||
# Create response
|
||||
create_resp = self.create_response("Test deletion", max_output_tokens=50)
|
||||
self.assertEqual(create_resp.status_code, 200)
|
||||
create_resp = self.create_response("Test deletion")
|
||||
self.assertIsNotNone(create_resp.id)
|
||||
self.assertIsNone(create_resp.error)
|
||||
self.assertEqual(create_resp.status, "completed")
|
||||
self.assertGreater(len(create_resp.output_text), 0)
|
||||
|
||||
response_id = create_resp.json()["id"]
|
||||
response_id = create_resp.id
|
||||
|
||||
# Delete response
|
||||
delete_resp = self.delete_response(response_id)
|
||||
self.assertEqual(delete_resp.status_code, 200)
|
||||
self.delete_response(response_id)
|
||||
|
||||
# Verify it's deleted (should return 404)
|
||||
get_resp = self.get_response(response_id)
|
||||
self.assertEqual(get_resp.status_code, 404)
|
||||
with self.assertRaises(openai.NotFoundError):
|
||||
self.get_response(response_id)
|
||||
|
||||
@unittest.skip("TODO: Add background response feature")
|
||||
def test_background_response(self):
|
||||
@@ -389,15 +303,15 @@ class ResponseCRUDBaseTest(ResponseAPIBaseTest):
|
||||
create_resp = self.create_response(
|
||||
"Write a short story", background=True, max_output_tokens=100
|
||||
)
|
||||
self.assertEqual(create_resp.status_code, 200)
|
||||
self.assertIsNotNone(create_resp.id)
|
||||
self.assertIsNone(create_resp.error)
|
||||
self.assertIn(create_resp.status, ["in_progress", "queued"])
|
||||
|
||||
create_data = create_resp.json()
|
||||
response_id = create_data["id"]
|
||||
self.assertEqual(create_data["status"], "in_progress")
|
||||
response_id = create_resp.id
|
||||
|
||||
# Wait for completion
|
||||
final_data = self.wait_for_background_task(response_id, timeout=60)
|
||||
self.assertEqual(final_data["status"], "completed")
|
||||
self.assertEqual(final_data.status, "completed")
|
||||
|
||||
|
||||
class ConversationCRUDBaseTest(ResponseAPIBaseTest):
|
||||
@@ -407,72 +321,88 @@ class ConversationCRUDBaseTest(ResponseAPIBaseTest):
|
||||
"""Test creating and retrieving conversation."""
|
||||
# Create conversation
|
||||
create_resp = self.create_conversation(metadata={"user": "test_user"})
|
||||
self.assertEqual(create_resp.status_code, 200)
|
||||
self.assertIsNotNone(create_resp.id)
|
||||
self.assertIsNotNone(create_resp.created_at)
|
||||
|
||||
create_data = create_resp.json()
|
||||
conversation_id = create_data["id"]
|
||||
self.assertEqual(create_data["metadata"]["user"], "test_user")
|
||||
create_data = create_resp.metadata
|
||||
self.assertEqual(create_data["user"], "test_user")
|
||||
conversation_id = create_resp.id
|
||||
|
||||
# Get conversation
|
||||
get_resp = self.get_conversation(conversation_id)
|
||||
self.assertEqual(get_resp.status_code, 200)
|
||||
self.assertIsNotNone(get_resp.id)
|
||||
self.assertIsNotNone(get_resp.created_at)
|
||||
|
||||
get_data = get_resp.json()
|
||||
self.assertEqual(get_data["id"], conversation_id)
|
||||
self.assertEqual(get_data["metadata"]["user"], "test_user")
|
||||
get_data = get_resp.metadata
|
||||
self.assertEqual(get_resp.id, conversation_id)
|
||||
self.assertEqual(get_data["user"], "test_user")
|
||||
|
||||
def test_update_conversation(self):
|
||||
"""Test updating conversation metadata."""
|
||||
# Create conversation
|
||||
create_resp = self.create_conversation(metadata={"key1": "value1"})
|
||||
self.assertEqual(create_resp.status_code, 200)
|
||||
conversation_id = create_resp.json()["id"]
|
||||
self.assertIsNotNone(create_resp.id)
|
||||
self.assertIsNotNone(create_resp.created_at)
|
||||
|
||||
create_data = create_resp.metadata
|
||||
self.assertEqual(create_data["key1"], "value1")
|
||||
self.assertNotIn("key2", create_data)
|
||||
conversation_id = create_resp.id
|
||||
|
||||
# Update conversation
|
||||
update_resp = self.update_conversation(
|
||||
conversation_id, metadata={"key1": "value1", "key2": "value2"}
|
||||
)
|
||||
self.assertEqual(update_resp.status_code, 200)
|
||||
self.assertEqual(update_resp.id, conversation_id)
|
||||
update_data = update_resp.metadata
|
||||
self.assertEqual(update_data["key1"], "value1")
|
||||
self.assertEqual(update_data["key2"], "value2")
|
||||
|
||||
# Verify update
|
||||
get_resp = self.get_conversation(conversation_id)
|
||||
get_data = get_resp.json()
|
||||
self.assertEqual(get_data["metadata"]["key2"], "value2")
|
||||
get_data = get_resp.metadata
|
||||
self.assertEqual(get_data["key1"], "value1")
|
||||
self.assertEqual(get_data["key2"], "value2")
|
||||
|
||||
def test_delete_conversation(self):
|
||||
"""Test deleting conversation."""
|
||||
# Create conversation
|
||||
create_resp = self.create_conversation()
|
||||
self.assertEqual(create_resp.status_code, 200)
|
||||
conversation_id = create_resp.json()["id"]
|
||||
self.assertIsNotNone(create_resp.id)
|
||||
self.assertIsNotNone(create_resp.created_at)
|
||||
conversation_id = create_resp.id
|
||||
|
||||
# Delete conversation
|
||||
delete_resp = self.delete_conversation(conversation_id)
|
||||
self.assertEqual(delete_resp.status_code, 200)
|
||||
self.assertIsNotNone(delete_resp.id)
|
||||
self.assertTrue(delete_resp.deleted)
|
||||
|
||||
# Verify deletion
|
||||
get_resp = self.get_conversation(conversation_id)
|
||||
self.assertEqual(get_resp.status_code, 404)
|
||||
with self.assertRaises(openai.NotFoundError):
|
||||
self.get_conversation(conversation_id)
|
||||
|
||||
def test_list_conversation_items(self):
|
||||
"""Test listing conversation items."""
|
||||
# Create conversation
|
||||
conv_resp = self.create_conversation()
|
||||
conversation_id = conv_resp.json()["id"]
|
||||
self.assertIsNotNone(conv_resp.id)
|
||||
conversation_id = conv_resp.id
|
||||
|
||||
# Create response with conversation
|
||||
self.create_response(
|
||||
resp1 = self.create_response(
|
||||
"First message", conversation=conversation_id, max_output_tokens=50
|
||||
)
|
||||
self.create_response(
|
||||
self.assertIsNone(resp1.error)
|
||||
resp2 = self.create_response(
|
||||
"Second message", conversation=conversation_id, max_output_tokens=50
|
||||
)
|
||||
self.assertIsNone(resp2.error)
|
||||
|
||||
# List items
|
||||
list_resp = self.list_conversation_items(conversation_id)
|
||||
self.assertEqual(list_resp.status_code, 200)
|
||||
self.assertIsNotNone(list_resp)
|
||||
self.assertIsNotNone(list_resp.data)
|
||||
|
||||
list_data = list_resp.json()
|
||||
self.assertIn("data", list_data)
|
||||
list_data = list_resp.data
|
||||
# Should have at least 4 items (2 inputs + 2 outputs)
|
||||
self.assertGreaterEqual(len(list_data["data"]), 4)
|
||||
self.assertGreaterEqual(len(list_data), 4)
|
||||
|
||||
@@ -13,49 +13,10 @@ from pathlib import Path
|
||||
_TEST_DIR = Path(__file__).parent
|
||||
sys.path.insert(0, str(_TEST_DIR))
|
||||
|
||||
from util import CustomTestCase
|
||||
|
||||
|
||||
class ResponseAPIBaseTest(CustomTestCase):
|
||||
"""Base class for Response API tests with common utilities."""
|
||||
|
||||
# To be set by subclasses
|
||||
base_url: str = None
|
||||
api_key: str = None
|
||||
model: str = None
|
||||
|
||||
def make_request(
|
||||
self,
|
||||
endpoint: str,
|
||||
method: str = "POST",
|
||||
json_data: dict = None,
|
||||
params: dict = None,
|
||||
):
|
||||
"""
|
||||
Make HTTP request to router.
|
||||
|
||||
This is a minimal implementation - subclasses should import from basic_crud.
|
||||
"""
|
||||
import requests
|
||||
|
||||
url = f"{self.base_url}{endpoint}"
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if self.api_key:
|
||||
headers["Authorization"] = f"Bearer {self.api_key}"
|
||||
|
||||
if method == "POST":
|
||||
resp = requests.post(url, json=json_data, headers=headers, params=params)
|
||||
elif method == "GET":
|
||||
resp = requests.get(url, headers=headers, params=params)
|
||||
elif method == "DELETE":
|
||||
resp = requests.delete(url, headers=headers, params=params)
|
||||
else:
|
||||
raise ValueError(f"Unsupported method: {method}")
|
||||
return resp
|
||||
from basic_crud import ResponseAPIBaseTest
|
||||
|
||||
|
||||
class FunctionCallingBaseTest(ResponseAPIBaseTest):
|
||||
"""Base class for function calling tests."""
|
||||
|
||||
def test_basic_function_call(self):
|
||||
"""
|
||||
@@ -99,54 +60,41 @@ class FunctionCallingBaseTest(ResponseAPIBaseTest):
|
||||
]
|
||||
|
||||
# 2. Prompt the model with tools defined
|
||||
resp = self.make_request(
|
||||
"/v1/responses",
|
||||
"POST",
|
||||
{
|
||||
"model": self.model,
|
||||
"tools": tools,
|
||||
"input": input_list,
|
||||
},
|
||||
)
|
||||
resp = self.create_response(input=input_list, tools=tools)
|
||||
|
||||
# Should successfully make the request
|
||||
self.assertEqual(resp.status_code, 200)
|
||||
|
||||
data = resp.json()
|
||||
self.assertIsNone(resp.error)
|
||||
|
||||
# Basic response structure
|
||||
self.assertIn("id", data)
|
||||
self.assertIn("status", data)
|
||||
self.assertEqual(data["status"], "completed")
|
||||
self.assertIn("output", data)
|
||||
self.assertIsNotNone(resp.id)
|
||||
self.assertEqual(resp.status, "completed")
|
||||
self.assertIsNotNone(resp.output)
|
||||
|
||||
# Verify output array is not empty
|
||||
output = data["output"]
|
||||
output = resp.output
|
||||
self.assertIsInstance(output, list)
|
||||
self.assertGreater(len(output), 0)
|
||||
|
||||
# Check for function_call in output
|
||||
function_calls = [
|
||||
item for item in output if item.get("type") == "function_call"
|
||||
]
|
||||
function_calls = [item for item in output if item.type == "function_call"]
|
||||
self.assertGreater(
|
||||
len(function_calls), 0, "Response should contain at least one function_call"
|
||||
)
|
||||
|
||||
# Verify function_call structure
|
||||
function_call = function_calls[0]
|
||||
self.assertIn("call_id", function_call)
|
||||
self.assertIn("name", function_call)
|
||||
self.assertEqual(function_call["name"], "get_horoscope")
|
||||
self.assertIn("arguments", function_call)
|
||||
self.assertIsNotNone(function_call.call_id)
|
||||
self.assertIsNotNone(function_call.name)
|
||||
self.assertEqual(function_call.name, "get_horoscope")
|
||||
self.assertIsNotNone(function_call.arguments)
|
||||
|
||||
# Parse arguments
|
||||
args = json.loads(function_call["arguments"])
|
||||
args = json.loads(function_call.arguments)
|
||||
self.assertIn("sign", args)
|
||||
self.assertEqual(args["sign"].lower(), "aquarius")
|
||||
|
||||
# 3. Save function call outputs for subsequent requests
|
||||
input_list += output
|
||||
input_list.append(function_call)
|
||||
|
||||
# 4. Execute the function logic for get_horoscope
|
||||
horoscope = f"{args['sign']}: Next Tuesday you will befriend a baby otter."
|
||||
@@ -155,47 +103,38 @@ class FunctionCallingBaseTest(ResponseAPIBaseTest):
|
||||
input_list.append(
|
||||
{
|
||||
"type": "function_call_output",
|
||||
"call_id": function_call["call_id"],
|
||||
"call_id": function_call.call_id,
|
||||
"output": json.dumps({"horoscope": horoscope}),
|
||||
}
|
||||
)
|
||||
|
||||
# 6. Make second request with function output
|
||||
resp2 = self.make_request(
|
||||
"/v1/responses",
|
||||
"POST",
|
||||
{
|
||||
"model": self.model,
|
||||
"instructions": "Respond only with a horoscope generated by a tool.",
|
||||
"tools": tools,
|
||||
"input": input_list,
|
||||
},
|
||||
resp2 = self.create_response(
|
||||
input=input_list,
|
||||
instructions="Respond only with a horoscope generated by a tool.",
|
||||
tools=tools,
|
||||
)
|
||||
data2 = resp2.json()
|
||||
self.assertEqual(data2["status"], "completed")
|
||||
self.assertIsNone(resp2.error)
|
||||
self.assertEqual(resp2.status, "completed")
|
||||
|
||||
# The model should be able to give a response using the function output
|
||||
output2 = data2["output"]
|
||||
output2 = resp2.output
|
||||
self.assertGreater(len(output2), 0)
|
||||
|
||||
# Find message output
|
||||
messages = [item for item in output2 if item.get("type") == "message"]
|
||||
messages = [item for item in output2 if item.type == "message"]
|
||||
self.assertGreater(
|
||||
len(messages), 0, "Response should contain at least one message"
|
||||
)
|
||||
|
||||
# Verify message contains the horoscope
|
||||
message = messages[0]
|
||||
self.assertIn("content", message)
|
||||
content_parts = message["content"]
|
||||
self.assertIsNotNone(message.content)
|
||||
content_parts = message.content
|
||||
self.assertGreater(len(content_parts), 0)
|
||||
|
||||
# Get text from content
|
||||
text_parts = [
|
||||
part.get("text", "")
|
||||
for part in content_parts
|
||||
if part.get("type") == "output_text"
|
||||
]
|
||||
text_parts = [part.text for part in content_parts if part.type == "output_text"]
|
||||
full_text = " ".join(text_parts).lower()
|
||||
|
||||
# Should mention the horoscope or baby otter
|
||||
|
||||
@@ -56,24 +56,19 @@ class MCPTests(ResponseAPIBaseTest):
|
||||
)
|
||||
|
||||
# Should successfully make the request
|
||||
self.assertEqual(resp.status_code, 200)
|
||||
|
||||
data = resp.json()
|
||||
self.assertIsNone(resp.error)
|
||||
|
||||
# Basic response structure
|
||||
self.assertIn("id", data)
|
||||
self.assertIn("status", data)
|
||||
self.assertEqual(data["status"], "completed")
|
||||
self.assertIn("output", data)
|
||||
self.assertIn("model", data)
|
||||
self.assertIsNotNone(resp.id)
|
||||
self.assertEqual(resp.status, "completed")
|
||||
self.assertIsNotNone(resp.model)
|
||||
self.assertIsNotNone(resp.output)
|
||||
|
||||
# Verify output array is not empty
|
||||
output = data["output"]
|
||||
self.assertIsInstance(output, list)
|
||||
self.assertGreater(len(output), 0)
|
||||
self.assertGreater(len(resp.output_text), 0)
|
||||
|
||||
# Check for MCP-specific output types
|
||||
output_types = [item.get("type") for item in output]
|
||||
output_types = [item.type for item in resp.output]
|
||||
|
||||
# Should have mcp_list_tools - tools are listed before calling
|
||||
self.assertIn(
|
||||
@@ -81,40 +76,38 @@ class MCPTests(ResponseAPIBaseTest):
|
||||
)
|
||||
|
||||
# Should have at least one mcp_call
|
||||
mcp_calls = [item for item in output if item.get("type") == "mcp_call"]
|
||||
mcp_calls = [item for item in resp.output if item.type == "mcp_call"]
|
||||
self.assertGreater(
|
||||
len(mcp_calls), 0, "Response should contain at least one mcp_call"
|
||||
)
|
||||
|
||||
# Verify mcp_call structure
|
||||
for mcp_call in mcp_calls:
|
||||
self.assertIn("id", mcp_call)
|
||||
self.assertIn("status", mcp_call)
|
||||
self.assertEqual(mcp_call["status"], "completed")
|
||||
self.assertIn("server_label", mcp_call)
|
||||
self.assertEqual(mcp_call["server_label"], "brave")
|
||||
self.assertIn("name", mcp_call)
|
||||
self.assertIn("arguments", mcp_call)
|
||||
self.assertIn("output", mcp_call)
|
||||
self.assertIsNotNone(mcp_call.id)
|
||||
self.assertEqual(mcp_call.status, "completed")
|
||||
self.assertEqual(mcp_call.server_label, "brave")
|
||||
self.assertIsNotNone(mcp_call.name)
|
||||
self.assertIsNotNone(mcp_call.arguments)
|
||||
self.assertIsNotNone(mcp_call.output)
|
||||
|
||||
# Strict mode: additional validation for HTTP backends
|
||||
if self.mcp_validation_mode == "strict":
|
||||
# Should have final message output
|
||||
messages = [item for item in output if item.get("type") == "message"]
|
||||
messages = [item for item in resp.output if item.type == "message"]
|
||||
self.assertGreater(
|
||||
len(messages), 0, "Response should contain at least one message"
|
||||
)
|
||||
# Verify message structure
|
||||
for msg in messages:
|
||||
self.assertIn("content", msg)
|
||||
self.assertIsInstance(msg["content"], list)
|
||||
self.assertIsNotNone(msg.content)
|
||||
self.assertIsInstance(msg.content, list)
|
||||
|
||||
# Check content has text
|
||||
for content_item in msg["content"]:
|
||||
if content_item.get("type") == "output_text":
|
||||
self.assertIn("text", content_item)
|
||||
self.assertIsInstance(content_item["text"], str)
|
||||
self.assertGreater(len(content_item["text"]), 0)
|
||||
for content_item in msg.content:
|
||||
if content_item.type == "output_text":
|
||||
self.assertIsNotNone(content_item.text)
|
||||
self.assertIsInstance(content_item.text, str)
|
||||
self.assertGreater(len(content_item.text), 0)
|
||||
|
||||
def test_mcp_basic_tool_call_streaming(self):
|
||||
"""Test basic MCP tool call (streaming).
|
||||
@@ -130,13 +123,10 @@ class MCPTests(ResponseAPIBaseTest):
|
||||
)
|
||||
|
||||
# Should successfully make the request
|
||||
self.assertEqual(resp.status_code, 200)
|
||||
|
||||
events = self.parse_sse_events(resp)
|
||||
events = [event for event in resp]
|
||||
self.assertGreater(len(events), 0)
|
||||
|
||||
event_types = [e.get("event") for e in events]
|
||||
|
||||
event_types = [event.type for event in events]
|
||||
# Check for lifecycle events
|
||||
self.assertIn(
|
||||
"response.created", event_types, "Should have response.created event"
|
||||
@@ -185,31 +175,31 @@ class MCPTests(ResponseAPIBaseTest):
|
||||
)
|
||||
|
||||
# Verify final completed event has full response
|
||||
completed_events = [e for e in events if e.get("event") == "response.completed"]
|
||||
completed_events = [e for e in events if e.type == "response.completed"]
|
||||
self.assertEqual(len(completed_events), 1)
|
||||
|
||||
final_response = completed_events[0].get("data", {}).get("response", {})
|
||||
self.assertIn("id", final_response)
|
||||
self.assertEqual(final_response.get("status"), "completed")
|
||||
self.assertIn("output", final_response)
|
||||
final_response = completed_events[0].response
|
||||
self.assertIsNotNone(final_response.id)
|
||||
self.assertEqual(final_response.status, "completed")
|
||||
self.assertIsNotNone(final_response.output)
|
||||
|
||||
# Verify final output contains expected items
|
||||
final_output = final_response.get("output", [])
|
||||
final_output_types = [item.get("type") for item in final_output]
|
||||
final_output = final_response.output
|
||||
final_output_types = [item.type for item in final_output]
|
||||
|
||||
self.assertIn("mcp_list_tools", final_output_types)
|
||||
self.assertIn("mcp_call", final_output_types)
|
||||
|
||||
# Verify mcp_call items in final output
|
||||
mcp_calls = [item for item in final_output if item.get("type") == "mcp_call"]
|
||||
mcp_calls = [item for item in final_output if item.type == "mcp_call"]
|
||||
self.assertGreater(len(mcp_calls), 0)
|
||||
|
||||
for mcp_call in mcp_calls:
|
||||
self.assertEqual(mcp_call.get("status"), "completed")
|
||||
self.assertEqual(mcp_call.get("server_label"), "brave")
|
||||
self.assertIn("name", mcp_call)
|
||||
self.assertIn("arguments", mcp_call)
|
||||
self.assertIn("output", mcp_call)
|
||||
self.assertEqual(mcp_call.status, "completed")
|
||||
self.assertEqual(mcp_call.server_label, "brave")
|
||||
self.assertIsNotNone(mcp_call.name)
|
||||
self.assertIsNotNone(mcp_call.arguments)
|
||||
self.assertIsNotNone(mcp_call.output)
|
||||
|
||||
# Strict mode: additional validation for HTTP backends
|
||||
if self.mcp_validation_mode == "strict":
|
||||
@@ -239,19 +229,17 @@ class MCPTests(ResponseAPIBaseTest):
|
||||
|
||||
# Verify text deltas combine to final message
|
||||
text_deltas = [
|
||||
e.get("data", {}).get("delta", "")
|
||||
for e in events
|
||||
if e.get("event") == "response.output_text.delta"
|
||||
e.delta for e in events if e.type == "response.output_text.delta"
|
||||
]
|
||||
self.assertGreater(len(text_deltas), 0, "Should have text deltas")
|
||||
|
||||
# Get final text from output_text.done event
|
||||
text_done_events = [
|
||||
e for e in events if e.get("event") == "response.output_text.done"
|
||||
e for e in events if e.type == "response.output_text.done"
|
||||
]
|
||||
self.assertGreater(len(text_done_events), 0)
|
||||
|
||||
final_text = text_done_events[0].get("data", {}).get("text", "")
|
||||
final_text = text_done_events[0].text
|
||||
self.assertGreater(len(final_text), 0, "Final text should not be empty")
|
||||
|
||||
def test_mixed_mcp_and_function_tools(self):
|
||||
@@ -264,38 +252,33 @@ class MCPTests(ResponseAPIBaseTest):
|
||||
)
|
||||
|
||||
# Should successfully make the request
|
||||
self.assertEqual(resp.status_code, 200)
|
||||
|
||||
data = resp.json()
|
||||
self.assertIsNone(resp.error)
|
||||
|
||||
# Basic response structure
|
||||
self.assertIn("id", data)
|
||||
self.assertIn("status", data)
|
||||
self.assertIn("output", data)
|
||||
self.assertIsNotNone(resp.id)
|
||||
self.assertIsNotNone(resp.status)
|
||||
self.assertIsNotNone(resp.output)
|
||||
|
||||
# Verify output array is not empty
|
||||
output = data["output"]
|
||||
output = resp.output
|
||||
self.assertIsInstance(output, list)
|
||||
self.assertGreater(len(output), 0)
|
||||
|
||||
# Check for function_call (not mcp_call for get_weather)
|
||||
function_calls = [
|
||||
item for item in output if item.get("type") == "function_call"
|
||||
]
|
||||
function_calls = [item for item in output if item.type == "function_call"]
|
||||
self.assertGreater(
|
||||
len(function_calls), 0, "Response should contain at least one function_call"
|
||||
)
|
||||
|
||||
# Verify function_call structure for get_weather
|
||||
weather_call = function_calls[0]
|
||||
self.assertIn("name", weather_call)
|
||||
self.assertEqual(weather_call["name"], "get_weather")
|
||||
self.assertIn("call_id", weather_call)
|
||||
self.assertIn("arguments", weather_call)
|
||||
self.assertIn("status", weather_call)
|
||||
self.assertEqual(weather_call.name, "get_weather")
|
||||
self.assertIsNotNone(weather_call.call_id)
|
||||
self.assertIsNotNone(weather_call.arguments)
|
||||
self.assertIsNotNone(weather_call.status)
|
||||
|
||||
# Parse and verify arguments
|
||||
args = json.loads(weather_call["arguments"])
|
||||
args = json.loads(weather_call.arguments)
|
||||
self.assertIn("location", args)
|
||||
self.assertIn("seattle", args["location"].lower())
|
||||
|
||||
@@ -309,12 +292,10 @@ class MCPTests(ResponseAPIBaseTest):
|
||||
)
|
||||
|
||||
# Should successfully make the request
|
||||
self.assertEqual(resp.status_code, 200)
|
||||
|
||||
events = self.parse_sse_events(resp)
|
||||
events = [event for event in resp]
|
||||
self.assertGreater(len(events), 0)
|
||||
|
||||
event_types = [e.get("event") for e in events]
|
||||
event_types = [e.type for e in events]
|
||||
|
||||
# Check for lifecycle events
|
||||
self.assertIn(
|
||||
@@ -345,8 +326,8 @@ class MCPTests(ResponseAPIBaseTest):
|
||||
mcp_call_arg_events = [
|
||||
e
|
||||
for e in events
|
||||
if e.get("event") == "response.mcp_call_arguments.delta"
|
||||
and "get_weather" in str(e.get("data", {}))
|
||||
if e.type == "response.mcp_call_arguments.delta"
|
||||
and "get_weather" in str(e.delta)
|
||||
]
|
||||
self.assertEqual(
|
||||
len(mcp_call_arg_events),
|
||||
@@ -356,9 +337,7 @@ class MCPTests(ResponseAPIBaseTest):
|
||||
|
||||
# Verify function_call_arguments.delta event structure
|
||||
func_arg_deltas = [
|
||||
e
|
||||
for e in events
|
||||
if e.get("event") == "response.function_call_arguments.delta"
|
||||
e for e in events if e.type == "response.function_call_arguments.delta"
|
||||
]
|
||||
self.assertGreater(
|
||||
len(func_arg_deltas), 0, "Should have function_call_arguments.delta events"
|
||||
@@ -367,8 +346,7 @@ class MCPTests(ResponseAPIBaseTest):
|
||||
# Check that at least one delta event contains location arguments
|
||||
has_location = False
|
||||
for event in func_arg_deltas:
|
||||
data = event.get("data", {})
|
||||
delta = data.get("delta", "")
|
||||
delta = event.delta
|
||||
if "location" in delta.lower() or "seattle" in delta.lower():
|
||||
has_location = True
|
||||
break
|
||||
|
||||
@@ -7,6 +7,7 @@ These tests should work across all backends (OpenAI, XAI, gRPC).
|
||||
|
||||
import unittest
|
||||
|
||||
import openai
|
||||
from basic_crud import ResponseAPIBaseTest
|
||||
|
||||
|
||||
@@ -19,57 +20,59 @@ class StateManagementTests(ResponseAPIBaseTest):
|
||||
resp1 = self.create_response(
|
||||
"My name is Alice and my friend is Bob. Remember it."
|
||||
)
|
||||
self.assertEqual(resp1.status_code, 200)
|
||||
response1_id = resp1.json()["id"]
|
||||
self.assertIsNone(resp1.error)
|
||||
self.assertEqual(resp1.status, "completed")
|
||||
response1_id = resp1.id
|
||||
|
||||
# Second response referencing first
|
||||
resp2 = self.create_response(
|
||||
"What is my name", previous_response_id=response1_id
|
||||
)
|
||||
self.assertEqual(resp2.status_code, 200)
|
||||
response2_data = resp2.json()
|
||||
self.assertIsNone(resp2.error)
|
||||
self.assertEqual(resp2.status, "completed")
|
||||
|
||||
# The model should remember the name from previous response
|
||||
output_text = self._extract_output_text(response2_data)
|
||||
self.assertIn("Alice", output_text)
|
||||
self.assertIn("Alice", resp2.output_text)
|
||||
|
||||
# Third response referencing second
|
||||
resp3 = self.create_response(
|
||||
"What is my friend name?",
|
||||
previous_response_id=response2_data["id"],
|
||||
previous_response_id=resp2.id,
|
||||
)
|
||||
response3_data = resp3.json()
|
||||
output_text = self._extract_output_text(response3_data)
|
||||
self.assertEqual(resp3.status_code, 200)
|
||||
self.assertIn("Bob", output_text)
|
||||
self.assertIsNone(resp3.error)
|
||||
self.assertEqual(resp3.status, "completed")
|
||||
self.assertIn("Bob", resp3.output_text)
|
||||
|
||||
@unittest.skip("TODO: Add the invalid previous_response_id check")
|
||||
def test_previous_response_id_invalid(self):
|
||||
"""Test using invalid previous_response_id."""
|
||||
resp = self.create_response(
|
||||
"Test", previous_response_id="resp_invalid123", max_output_tokens=50
|
||||
)
|
||||
self.assertIn(resp.status_code, [400, 404])
|
||||
with self.assertRaises(openai.BadRequestError):
|
||||
self.create_response(
|
||||
"Test", previous_response_id="resp_invalid123", max_output_tokens=50
|
||||
)
|
||||
|
||||
def test_conversation_with_multiple_turns(self):
|
||||
"""Test state management using conversation ID."""
|
||||
# Create conversation
|
||||
conv_resp = self.create_conversation(metadata={"topic": "math"})
|
||||
self.assertEqual(conv_resp.status_code, 200)
|
||||
self.assertIsNotNone(conv_resp.id)
|
||||
self.assertIsNotNone(conv_resp.created_at)
|
||||
|
||||
conversation_id = conv_resp.json()["id"]
|
||||
conversation_id = conv_resp.id
|
||||
|
||||
# First response in conversation
|
||||
resp1 = self.create_response("I have 5 apples.", conversation=conversation_id)
|
||||
self.assertEqual(resp1.status_code, 200)
|
||||
self.assertIsNone(resp1.error)
|
||||
self.assertEqual(resp1.status, "completed")
|
||||
|
||||
# Second response in same conversation
|
||||
resp2 = self.create_response(
|
||||
"How many apples do I have?",
|
||||
conversation=conversation_id,
|
||||
)
|
||||
self.assertEqual(resp2.status_code, 200)
|
||||
output_text = self._extract_output_text(resp2.json())
|
||||
self.assertIsNone(resp2.error)
|
||||
self.assertEqual(resp2.status, "completed")
|
||||
output_text = resp2.output_text
|
||||
|
||||
# Should remember "5 apples"
|
||||
self.assertTrue("5" in output_text or "five" in output_text.lower())
|
||||
@@ -79,14 +82,15 @@ class StateManagementTests(ResponseAPIBaseTest):
|
||||
"If I get 3 more, how many total?",
|
||||
conversation=conversation_id,
|
||||
)
|
||||
self.assertEqual(resp3.status_code, 200)
|
||||
output_text = self._extract_output_text(resp3.json())
|
||||
self.assertIsNone(resp3.error)
|
||||
self.assertEqual(resp3.status, "completed")
|
||||
output_text = resp3.output_text
|
||||
|
||||
# Should calculate 5 + 3 = 8
|
||||
self.assertTrue("8" in output_text or "eight" in output_text.lower())
|
||||
list_resp = self.list_conversation_items(conversation_id)
|
||||
self.assertEqual(list_resp.status_code, 200)
|
||||
items = list_resp.json()["data"]
|
||||
self.assertIsNotNone(list_resp.data)
|
||||
items = list_resp.data
|
||||
# Should have at least 6 items (3 inputs + 3 outputs)
|
||||
self.assertGreaterEqual(len(items), 6)
|
||||
|
||||
@@ -96,20 +100,15 @@ class StateManagementTests(ResponseAPIBaseTest):
|
||||
conversation_id = "conv_123"
|
||||
|
||||
resp1 = self.create_response("Test")
|
||||
response1_id = resp1.json()["id"]
|
||||
response1_id = resp1.id
|
||||
|
||||
# Try to use both parameters
|
||||
resp = self.create_response(
|
||||
"This should fail",
|
||||
previous_response_id=response1_id,
|
||||
conversation=conversation_id,
|
||||
)
|
||||
|
||||
# Should return 400 Bad Request
|
||||
self.assertEqual(resp.status_code, 400)
|
||||
error_data = resp.json()
|
||||
self.assertIn("error", error_data)
|
||||
self.assertIn("mutually exclusive", error_data["error"]["message"].lower())
|
||||
with self.assertRaises(openai.BadRequestError):
|
||||
self.create_response(
|
||||
"This should fail",
|
||||
previous_response_id=response1_id,
|
||||
conversation=conversation_id,
|
||||
)
|
||||
|
||||
# Helper methods
|
||||
|
||||
|
||||
@@ -12,42 +12,16 @@ from pathlib import Path
|
||||
_TEST_DIR = Path(__file__).parent
|
||||
sys.path.insert(0, str(_TEST_DIR))
|
||||
|
||||
from util import CustomTestCase
|
||||
from basic_crud import ResponseAPIBaseTest
|
||||
|
||||
|
||||
class StructuredOutputBaseTest(CustomTestCase):
|
||||
"""Base class for structured output tests with common utilities."""
|
||||
|
||||
# To be set by subclasses
|
||||
base_url: str = None
|
||||
api_key: str = None
|
||||
model: str = None
|
||||
|
||||
def make_request(self, endpoint, method="GET", data=None):
|
||||
"""Make HTTP request to the API."""
|
||||
url = f"{self.base_url}{endpoint}"
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
}
|
||||
|
||||
if method == "GET":
|
||||
response = self.session.get(url, headers=headers)
|
||||
elif method == "POST":
|
||||
response = self.session.post(url, headers=headers, json=data)
|
||||
elif method == "DELETE":
|
||||
response = self.session.delete(url, headers=headers)
|
||||
else:
|
||||
raise ValueError(f"Unsupported method: {method}")
|
||||
|
||||
return response
|
||||
class StructuredOutputBaseTest(ResponseAPIBaseTest):
|
||||
|
||||
def test_structured_output_json_schema(self):
|
||||
"""Test structured output with json_schema format."""
|
||||
|
||||
# Create response with structured output
|
||||
data = {
|
||||
"model": self.model,
|
||||
params = {
|
||||
"input": [
|
||||
{
|
||||
"role": "system",
|
||||
@@ -84,29 +58,27 @@ class StructuredOutputBaseTest(CustomTestCase):
|
||||
},
|
||||
}
|
||||
|
||||
create_resp = self.make_request("/v1/responses", "POST", data)
|
||||
self.assertEqual(create_resp.status_code, 200)
|
||||
|
||||
create_data = create_resp.json()
|
||||
self.assertIn("id", create_data)
|
||||
self.assertIn("output", create_data)
|
||||
self.assertIn("text", create_data)
|
||||
create_resp = self.create_response(**params)
|
||||
self.assertIsNone(create_resp.error)
|
||||
self.assertIsNotNone(create_resp.id)
|
||||
self.assertIsNotNone(create_resp.output)
|
||||
self.assertIsNotNone(create_resp.text)
|
||||
|
||||
# Verify text format was echoed back correctly
|
||||
self.assertIn("format", create_data["text"])
|
||||
self.assertEqual(create_data["text"]["format"]["type"], "json_schema")
|
||||
self.assertEqual(create_data["text"]["format"]["name"], "math_reasoning")
|
||||
self.assertIn("schema", create_data["text"]["format"])
|
||||
self.assertEqual(create_data["text"]["format"]["strict"], True)
|
||||
self.assertIsNotNone(create_resp.text.format)
|
||||
self.assertEqual(create_resp.text.format.type, "json_schema")
|
||||
self.assertEqual(create_resp.text.format.name, "math_reasoning")
|
||||
self.assertIsNotNone(create_resp.text.format.schema_)
|
||||
self.assertEqual(create_resp.text.format.strict, True)
|
||||
|
||||
# Find the message output (output[0] may be reasoning, output[1] is message)
|
||||
output_text = next(
|
||||
(
|
||||
content.get("text", "")
|
||||
for item in create_data.get("output", [])
|
||||
if item.get("type") == "message"
|
||||
for content in item.get("content", [])
|
||||
if content.get("type") == "output_text"
|
||||
content.text
|
||||
for item in create_resp.output
|
||||
if item.type == "message"
|
||||
for content in item.content
|
||||
if content.type == "output_text"
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
@@ -12,6 +12,8 @@ import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
import openai
|
||||
|
||||
# Add e2e_response_api directory for imports
|
||||
_TEST_DIR = Path(__file__).parent.parent
|
||||
sys.path.insert(0, str(_TEST_DIR))
|
||||
@@ -39,6 +41,7 @@ class TestOracleStore(ResponseCRUDBaseTest, ConversationCRUDBaseTest):
|
||||
)
|
||||
|
||||
cls.base_url = cls.cluster["base_url"]
|
||||
cls.client = openai.Client(api_key=cls.api_key, base_url=cls.base_url + "/v1")
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
|
||||
Reference in New Issue
Block a user