Files
sglang/sgl-router/py_test/e2e_response_api/mixins/basic_crud.py

409 lines
15 KiB
Python

"""
Base test class for Response API e2e tests.
This module provides base test classes that can be reused across different backends
(OpenAI, XAI, gRPC) with common test logic.
"""
from __future__ import annotations
import sys
import time
import unittest
from pathlib import Path
from typing import Optional, Union
import openai
from openai.types import conversations, responses
# Add current directory for local imports
_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
client: openai.OpenAI = None
def create_response(
self,
input: Union[str, responses.ResponseInputParam],
instructions: Optional[str] = None,
stream: bool = False,
max_output_tokens: Optional[int] = None,
temperature: Optional[float] = None,
previous_response_id: Optional[str] = None,
conversation: Optional[str] = None,
tools: Optional[list] = None,
background: bool = False,
**kwargs,
) -> responses.Response | openai.Stream[responses.ResponseStreamEvent]:
"""
Create a response via POST /v1/responses.
Args:
input: User input
instructions: Optional system instructions
stream: Whether to stream response
max_output_tokens: Optional max tokens to generate
temperature: Sampling temperature
previous_response_id: Optional previous response ID for state management
conversation: Optional conversation ID for state management
tools: Optional list of MCP tools
background: Whether to run in background mode
**kwargs: Additional request parameters
Returns:
Response object for non-stream request
ResponseStreamEvent for stream request
"""
params = {
"model": self.model,
"input": input,
"stream": stream,
**kwargs,
}
if instructions:
params["instructions"] = instructions
if max_output_tokens is not None:
params["max_output_tokens"] = max_output_tokens
if temperature is not None:
params["temperature"] = temperature
if previous_response_id:
params["previous_response_id"] = previous_response_id
if conversation:
params["conversation"] = conversation
if tools:
params["tools"] = tools
if background:
params["background"] = background
return self.client.responses.create(**params)
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.client.responses.retrieve(response_id=response_id)
def delete_response(self, response_id: str) -> None:
"""Delete response by ID via DELETE /v1/responses/{response_id}."""
return self.client.responses.delete(response_id=response_id)
def cancel_response(self, response_id: str) -> responses.Response:
"""Cancel response by ID via POST /v1/responses/{response_id}/cancel."""
return self.client.responses.cancel(response_id=response_id)
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.client.responses.input_items.list(response_id=response_id)
def create_conversation(
self, metadata: Optional[dict] = None
) -> conversations.Conversation:
"""Create conversation via POST /v1/conversations."""
params = {}
if metadata:
params["metadata"] = metadata
return self.client.conversations.create(**params)
def get_conversation(self, conversation_id: str) -> conversations.Conversation:
"""Get conversation by ID via GET /v1/conversations/{conversation_id}."""
return self.client.conversations.retrieve(conversation_id=conversation_id)
def update_conversation(
self, conversation_id: str, metadata: dict
) -> conversations.Conversation:
"""Update conversation via POST /v1/conversations/{conversation_id}."""
return self.client.conversations.update(
conversation_id=conversation_id, metadata=metadata
)
def delete_conversation(
self, conversation_id: str
) -> conversations.ConversationDeletedResource:
"""Delete conversation via DELETE /v1/conversations/{conversation_id}."""
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,
order: str = "asc",
) -> openai.pagination.SyncConversationCursorPage[conversations.ConversationItem]:
"""List conversation items via GET /v1/conversations/{conversation_id}/items."""
params = {"conversation_id": conversation_id, "order": order}
if limit:
params["limit"] = limit
if after:
params["after"] = after
return self.client.conversations.items.list(**params)
def create_conversation_items(
self, conversation_id: str, items: list
) -> conversations.ConversationItemList:
"""Create conversation items via POST /v1/conversations/{conversation_id}/items."""
return self.client.conversations.items.create(
conversation_id=conversation_id, items=items
)
def get_conversation_item(
self, conversation_id: str, item_id: str
) -> conversations.ConversationItem:
"""Get conversation item via GET /v1/conversations/{conversation_id}/items/{item_id}."""
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
) -> conversations.Conversation:
"""Delete conversation item via DELETE /v1/conversations/{conversation_id}/items/{item_id}."""
return self.client.conversations.items.delete(
conversation_id=conversation_id, item_id=item_id
)
def wait_for_background_task(
self, response_id: str, timeout: int = 30, poll_interval: float = 0.5
) -> responses.Response:
"""
Wait for background task to complete.
Args:
response_id: Response ID to poll
timeout: Max seconds to wait
poll_interval: Seconds between polls
Returns:
Final response data
Raises:
TimeoutError: If task doesn't complete in time
AssertionError: If task fails
"""
start_time = time.time()
while time.time() - start_time < timeout:
resp = self.get_response(response_id)
self.assertIsNone(resp.error)
self.assertEqual(resp.id, response_id)
status = resp.status
if status == "completed":
return resp
elif status == "failed":
raise AssertionError(f"Background task failed: {resp.error}")
elif status == "cancelled":
raise AssertionError("Background task was cancelled")
time.sleep(poll_interval)
raise TimeoutError(
f"Background task {response_id} did not complete within {timeout}s"
)
class StateManagementBaseTest(ResponseAPIBaseTest):
"""Base class for state management tests (previous_response_id and conversation)."""
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.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)
# Check for response.created event
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(
event.type in ["response.completed", "response.in_progress"]
for event in events
)
)
class ResponseCRUDBaseTest(ResponseAPIBaseTest):
"""Base class for Response API CRUD tests."""
def test_create_and_get_response(self):
"""Test creating response and retrieving it."""
# Create response
create_resp = self.create_response("Hello, world!")
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.assertIsNone(get_resp.error)
self.assertEqual(get_resp.id, response_id)
self.assertEqual(get_resp.status, "completed")
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")
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
# Delete response
self.delete_response(response_id)
# Verify it's deleted (should return 404)
with self.assertRaises(openai.NotFoundError):
self.get_response(response_id)
@unittest.skip("TODO: Add background response feature")
def test_background_response(self):
"""Test background response execution."""
# Create background response
create_resp = self.create_response(
"Write a short story", background=True, max_output_tokens=100
)
self.assertIsNotNone(create_resp.id)
self.assertIsNone(create_resp.error)
self.assertIn(create_resp.status, ["in_progress", "queued"])
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")
class ConversationCRUDBaseTest(ResponseAPIBaseTest):
"""Base class for Conversation API CRUD tests."""
def test_create_and_get_conversation(self):
"""Test creating and retrieving conversation."""
# Create conversation
create_resp = self.create_conversation(metadata={"user": "test_user"})
self.assertIsNotNone(create_resp.id)
self.assertIsNotNone(create_resp.created_at)
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.assertIsNotNone(get_resp.id)
self.assertIsNotNone(get_resp.created_at)
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.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.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.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.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.assertIsNotNone(delete_resp.id)
self.assertTrue(delete_resp.deleted)
# Verify deletion
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()
self.assertIsNotNone(conv_resp.id)
conversation_id = conv_resp.id
# Create response with conversation
resp1 = self.create_response(
"First message", conversation=conversation_id, max_output_tokens=50
)
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.assertIsNotNone(list_resp)
self.assertIsNotNone(list_resp.data)
list_data = list_resp.data
# Should have at least 4 items (2 inputs + 2 outputs)
self.assertGreaterEqual(len(list_data), 4)