diff --git a/studio/backend/models/inference.py b/studio/backend/models/inference.py index cf08ecbc12..983b674453 100644 --- a/studio/backend/models/inference.py +++ b/studio/backend/models/inference.py @@ -444,3 +444,111 @@ class ChatCompletion(BaseModel): model: str = "default" choices: list[CompletionChoice] usage: CompletionUsage = Field(default_factory = CompletionUsage) + + +# ===================================================================== +# OpenAI Responses API Models (/v1/responses) +# ===================================================================== + + +# ── Request models ────────────────────────────────────────────── + + +class ResponsesInputTextPart(BaseModel): + """Text content part in a Responses API message (type=input_text).""" + type: Literal["input_text"] + text: str + + +class ResponsesInputImagePart(BaseModel): + """Image content part in a Responses API message (type=input_image).""" + type: Literal["input_image"] + image_url: str = Field(..., description = "data:image/png;base64,... or https://...") + detail: Optional[Literal["auto", "low", "high"]] = "auto" + + +ResponsesContentPart = Union[ResponsesInputTextPart, ResponsesInputImagePart] + + +class ResponsesInputMessage(BaseModel): + """A single message in the Responses API input array.""" + role: Literal["system", "user", "assistant", "developer"] + content: Union[str, list[ResponsesContentPart]] + + +class ResponsesRequest(BaseModel): + """OpenAI Responses API request.""" + model: str = Field("default", description = "Model identifier") + input: Union[str, list[ResponsesInputMessage]] = Field( + default = [], + description = "Input text or message list", + ) + instructions: Optional[str] = Field( + None, description = "System / developer instructions" + ) + temperature: Optional[float] = Field(None, ge = 0.0, le = 2.0) + top_p: Optional[float] = Field(None, ge = 0.0, le = 1.0) + max_output_tokens: Optional[int] = Field(None, ge = 1) + stream: bool = Field(False, description = "Whether to stream the response via SSE") + + # Accepted but ignored -- keeps SDK clients from failing on unsupported fields + tools: Optional[list] = None + tool_choice: Optional[Any] = None + previous_response_id: Optional[str] = None + store: Optional[bool] = None + metadata: Optional[dict] = None + truncation: Optional[Any] = None + user: Optional[str] = None + text: Optional[Any] = None + reasoning: Optional[Any] = None + + model_config = {"extra": "allow"} + + +# ── Response models ───────────────────────────────────────────── + + +class ResponsesOutputTextContent(BaseModel): + """A text content block inside an output message.""" + type: Literal["output_text"] = "output_text" + text: str + annotations: list = Field(default_factory = list) + + +class ResponsesOutputMessage(BaseModel): + """An output message in the Responses API response.""" + type: Literal["message"] = "message" + id: str = Field(default_factory = lambda: f"msg_{uuid.uuid4().hex[:12]}") + status: Literal["completed", "in_progress"] = "completed" + role: Literal["assistant"] = "assistant" + content: list[ResponsesOutputTextContent] = Field(default_factory = list) + + +class ResponsesUsage(BaseModel): + """Token usage for a Responses API response (input_tokens, not prompt_tokens).""" + input_tokens: int = 0 + output_tokens: int = 0 + total_tokens: int = 0 + + +class ResponsesResponse(BaseModel): + """Top-level Responses API response object.""" + id: str = Field(default_factory = lambda: f"resp_{uuid.uuid4().hex[:12]}") + object: Literal["response"] = "response" + created_at: int = Field(default_factory = lambda: int(time.time())) + status: Literal["completed", "in_progress", "failed"] = "completed" + model: str = "default" + output: list[ResponsesOutputMessage] = Field(default_factory = list) + usage: ResponsesUsage = Field(default_factory = ResponsesUsage) + error: Optional[Any] = None + incomplete_details: Optional[Any] = None + instructions: Optional[str] = None + metadata: dict = Field(default_factory = dict) + temperature: Optional[float] = None + top_p: Optional[float] = None + max_output_tokens: Optional[int] = None + previous_response_id: Optional[str] = None + text: Optional[Any] = None + tool_choice: Optional[Any] = None + tools: list = Field(default_factory = list) + truncation: Optional[Any] = None diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index eef44ba9c8..393c9d8289 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -85,6 +85,17 @@ from models.inference import ( CompletionUsage, ValidateModelRequest, ValidateModelResponse, + TextContentPart, + ImageContentPart, + ImageUrl, + ResponsesRequest, + ResponsesInputMessage, + ResponsesInputTextPart, + ResponsesInputImagePart, + ResponsesOutputTextContent, + ResponsesOutputMessage, + ResponsesUsage, + ResponsesResponse, ) from auth.authentication import get_current_subject @@ -1915,25 +1926,244 @@ async def openai_embeddings( # ===================================================================== -from pydantic import BaseModel -from typing import Union +def _normalise_responses_input(payload: ResponsesRequest) -> list[ChatMessage]: + """Convert a ResponsesRequest into a list of ChatMessage for the completions backend.""" + messages: list[ChatMessage] = [] + + # System / developer instructions + if payload.instructions: + messages.append(ChatMessage(role = "system", content = payload.instructions)) + + # Simple string input + if isinstance(payload.input, str): + if payload.input: + messages.append(ChatMessage(role = "user", content = payload.input)) + return messages + + # List of ResponsesInputMessage + for msg in payload.input: + role = "system" if msg.role == "developer" else msg.role + + if isinstance(msg.content, str): + messages.append(ChatMessage(role = role, content = msg.content)) + else: + # Convert Responses content parts -> Chat content parts + parts = [] + for part in msg.content: + if isinstance(part, ResponsesInputTextPart): + parts.append(TextContentPart(type = "text", text = part.text)) + elif isinstance(part, ResponsesInputImagePart): + parts.append(ImageContentPart( + type = "image_url", + image_url = ImageUrl(url = part.image_url, detail = part.detail), + )) + messages.append(ChatMessage(role = role, content = parts if parts else "")) + + return messages -class _ResponsesInputMessage(BaseModel): - role: str - content: str +def _build_chat_request(payload: ResponsesRequest, messages: list[ChatMessage], stream: bool) -> ChatCompletionRequest: + """Build a ChatCompletionRequest from a ResponsesRequest.""" + chat_kwargs = dict( + model = payload.model, + messages = messages, + stream = stream, + ) + if payload.temperature is not None: + chat_kwargs["temperature"] = payload.temperature + if payload.top_p is not None: + chat_kwargs["top_p"] = payload.top_p + if payload.max_output_tokens is not None: + chat_kwargs["max_tokens"] = payload.max_output_tokens + return ChatCompletionRequest(**chat_kwargs) -class ResponsesRequest(BaseModel): - """Minimal OpenAI Responses API request.""" +async def _responses_non_streaming( + payload: ResponsesRequest, + messages: list[ChatMessage], + request: Request, +) -> JSONResponse: + """Handle a non-streaming Responses API call.""" + chat_req = _build_chat_request(payload, messages, stream = False) + result = await openai_chat_completions(chat_req, request) - model: str = "default" - input: Union[str, list[_ResponsesInputMessage]] = [] - instructions: Optional[str] = None - temperature: Optional[float] = None - top_p: Optional[float] = None - max_output_tokens: Optional[int] = None - stream: bool = False + # openai_chat_completions returns a JSONResponse for non-streaming + if isinstance(result, JSONResponse): + body = json.loads(result.body.decode()) + elif isinstance(result, Response): + body = json.loads(result.body.decode()) + else: + body = result + + # Extract content and usage from the Chat Completions response + choices = body.get("choices", []) + text = "" + if choices: + msg = choices[0].get("message", {}) + text = msg.get("content", "") or "" + + usage_data = body.get("usage", {}) + input_tokens = usage_data.get("prompt_tokens", 0) + output_tokens = usage_data.get("completion_tokens", 0) + + resp_id = f"resp_{uuid.uuid4().hex[:12]}" + msg_id = f"msg_{uuid.uuid4().hex[:12]}" + + response = ResponsesResponse( + id = resp_id, + created_at = int(time.time()), + status = "completed", + model = body.get("model", payload.model), + output = [ + ResponsesOutputMessage( + id = msg_id, + status = "completed", + role = "assistant", + content = [ + ResponsesOutputTextContent(text = text), + ], + ), + ], + usage = ResponsesUsage( + input_tokens = input_tokens, + output_tokens = output_tokens, + total_tokens = input_tokens + output_tokens, + ), + temperature = payload.temperature, + top_p = payload.top_p, + max_output_tokens = payload.max_output_tokens, + instructions = payload.instructions, + ) + return JSONResponse(content = response.model_dump()) + + +async def _responses_stream( + payload: ResponsesRequest, + messages: list[ChatMessage], + request: Request, +): + """Handle a streaming Responses API call, emitting named SSE events.""" + resp_id = f"resp_{uuid.uuid4().hex[:12]}" + msg_id = f"msg_{uuid.uuid4().hex[:12]}" + item_id = f"item_{uuid.uuid4().hex[:12]}" + created_at = int(time.time()) + + chat_req = _build_chat_request(payload, messages, stream = True) + result = await openai_chat_completions(chat_req, request) + + async def event_generator(): + full_text = "" + input_tokens = 0 + output_tokens = 0 + + # ── Preamble events ── + yield f"event: response.created\ndata: {json.dumps({'type': 'response.created', 'response': {'id': resp_id, 'object': 'response', 'created_at': created_at, 'status': 'in_progress', 'model': payload.model, 'output': [], 'usage': {'input_tokens': 0, 'output_tokens': 0, 'total_tokens': 0}}})}\n\n" + + # output_item.added + output_item = { + "type": "message", + "id": msg_id, + "status": "in_progress", + "role": "assistant", + "content": [], + } + yield f"event: response.output_item.added\ndata: {json.dumps({'type': 'response.output_item.added', 'output_index': 0, 'item': output_item})}\n\n" + + # content_part.added + content_part = {"type": "output_text", "text": "", "annotations": []} + yield f"event: response.content_part.added\ndata: {json.dumps({'type': 'response.content_part.added', 'item_id': msg_id, 'output_index': 0, 'content_index': 0, 'part': content_part})}\n\n" + + # ── Stream delta events from the inner chat completions stream ── + if isinstance(result, StreamingResponse): + async for raw_chunk in result.body_iterator: + if isinstance(raw_chunk, bytes): + raw_chunk = raw_chunk.decode("utf-8", errors = "replace") + + for line in raw_chunk.split("\n"): + line = line.strip() + if not line.startswith("data: "): + continue + data_str = line[6:] + if data_str == "[DONE]": + continue + try: + chunk_data = json.loads(data_str) + except json.JSONDecodeError: + continue + + choices = chunk_data.get("choices", []) + if not choices: + # Check for usage in final chunk + usage = chunk_data.get("usage") + if usage: + input_tokens = usage.get("prompt_tokens", input_tokens) + output_tokens = usage.get("completion_tokens", output_tokens) + continue + + delta = choices[0].get("delta", {}) + content = delta.get("content") + if content: + full_text += content + delta_event = { + "type": "response.output_text.delta", + "item_id": msg_id, + "output_index": 0, + "content_index": 0, + "delta": content, + } + yield f"event: response.output_text.delta\ndata: {json.dumps(delta_event)}\n\n" + + # Check for usage in chunk + usage = chunk_data.get("usage") + if usage: + input_tokens = usage.get("prompt_tokens", input_tokens) + output_tokens = usage.get("completion_tokens", output_tokens) + + # ── Closing events ── + # output_text.done + yield f"event: response.output_text.done\ndata: {json.dumps({'type': 'response.output_text.done', 'item_id': msg_id, 'output_index': 0, 'content_index': 0, 'text': full_text})}\n\n" + + # content_part.done + yield f"event: response.content_part.done\ndata: {json.dumps({'type': 'response.content_part.done', 'item_id': msg_id, 'output_index': 0, 'content_index': 0, 'part': {'type': 'output_text', 'text': full_text, 'annotations': []}})}\n\n" + + # output_item.done + yield f"event: response.output_item.done\ndata: {json.dumps({'type': 'response.output_item.done', 'output_index': 0, 'item': {'type': 'message', 'id': msg_id, 'status': 'completed', 'role': 'assistant', 'content': [{'type': 'output_text', 'text': full_text, 'annotations': []}]}})}\n\n" + + # response.completed + total_tokens = input_tokens + output_tokens + completed_response = { + "type": "response.completed", + "response": { + "id": resp_id, + "object": "response", + "created_at": created_at, + "status": "completed", + "model": payload.model, + "output": [{ + "type": "message", + "id": msg_id, + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": full_text, "annotations": []}], + }], + "usage": { + "input_tokens": input_tokens, + "output_tokens": output_tokens, + "total_tokens": total_tokens, + }, + }, + } + yield f"event: response.completed\ndata: {json.dumps(completed_response)}\n\n" + + return StreamingResponse( + event_generator(), + media_type = "text/event-stream", + headers = { + "Cache-Control": "no-cache", + "Connection": "keep-alive", + "X-Accel-Buffering": "no", + }, + ) @router.post("/responses") @@ -1945,39 +2175,15 @@ async def openai_responses( """ OpenAI Responses API endpoint. - Converts the Responses-format request into a ChatCompletionRequest - and delegates to the existing chat completions handler. Works with - both GGUF and non-GGUF backends. + Accepts the Responses-format request, converts it to a + ChatCompletionRequest internally, and returns a response + matching the OpenAI Responses API schema (output array, + input_tokens/output_tokens, named SSE events for streaming). """ - # Build messages list from the Responses API input format - messages: list[ChatMessage] = [] - - # System message from instructions - if payload.instructions: - messages.append(ChatMessage(role = "system", content = payload.instructions)) - - # Convert input to messages - if isinstance(payload.input, str): - messages.append(ChatMessage(role = "user", content = payload.input)) - else: - for msg in payload.input: - messages.append(ChatMessage(role = msg.role, content = msg.content)) - + messages = _normalise_responses_input(payload) if not messages: raise HTTPException(status_code = 400, detail = "No input provided.") - # Build a ChatCompletionRequest and delegate - chat_kwargs = dict( - model = payload.model, - messages = messages, - stream = payload.stream, - ) - if payload.temperature is not None: - chat_kwargs["temperature"] = payload.temperature - if payload.top_p is not None: - chat_kwargs["top_p"] = payload.top_p - if payload.max_output_tokens is not None: - chat_kwargs["max_tokens"] = payload.max_output_tokens - - chat_request = ChatCompletionRequest(**chat_kwargs) - return await openai_chat_completions(chat_request, request) + if payload.stream: + return await _responses_stream(payload, messages, request) + return await _responses_non_streaming(payload, messages, request) diff --git a/studio/backend/tests/test_responses_api.py b/studio/backend/tests/test_responses_api.py new file mode 100644 index 0000000000..1905e58de2 --- /dev/null +++ b/studio/backend/tests/test_responses_api.py @@ -0,0 +1,318 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. + +""" +Tests for the OpenAI Responses API schemas and input normalisation. +These tests do NOT require a running server or GPU -- they validate +the Pydantic models and the _normalise_responses_input helper. +""" + +import sys +import os +import json +import re + +# Ensure backend is on path +_backend = os.path.join(os.path.dirname(__file__), "..") +sys.path.insert(0, _backend) + +from models.inference import ( + ResponsesRequest, + ResponsesInputMessage, + ResponsesInputTextPart, + ResponsesInputImagePart, + ResponsesOutputTextContent, + ResponsesOutputMessage, + ResponsesUsage, + ResponsesResponse, + ChatMessage, + TextContentPart, + ImageContentPart, + ImageUrl, + ChatCompletionRequest, +) + + +# ── _normalise_responses_input: copied from routes/inference.py ── +# We cannot import routes.inference directly because routes/__init__.py +# pulls in heavy dependencies (structlog/twisted/torch). This is a +# direct copy of the function for testing purposes. + +def _normalise_responses_input(payload: ResponsesRequest) -> list: + """Convert a ResponsesRequest into a list of ChatMessage for the completions backend.""" + messages = [] + + # System / developer instructions + if payload.instructions: + messages.append(ChatMessage(role = "system", content = payload.instructions)) + + # Simple string input + if isinstance(payload.input, str): + if payload.input: + messages.append(ChatMessage(role = "user", content = payload.input)) + return messages + + # List of ResponsesInputMessage + for msg in payload.input: + role = "system" if msg.role == "developer" else msg.role + + if isinstance(msg.content, str): + messages.append(ChatMessage(role = role, content = msg.content)) + else: + # Convert Responses content parts -> Chat content parts + parts = [] + for part in msg.content: + if isinstance(part, ResponsesInputTextPart): + parts.append(TextContentPart(type = "text", text = part.text)) + elif isinstance(part, ResponsesInputImagePart): + parts.append(ImageContentPart( + type = "image_url", + image_url = ImageUrl(url = part.image_url, detail = part.detail), + )) + messages.append(ChatMessage(role = role, content = parts if parts else "")) + + return messages + + +# ===================================================================== +# Schema validation tests +# ===================================================================== + + +class TestResponsesRequest: + """Validate ResponsesRequest accepts the shapes the OpenAI SDK sends.""" + + def test_minimal_string_input(self): + req = ResponsesRequest(input = "Hello") + assert req.input == "Hello" + assert req.stream is False + assert req.model == "default" + + def test_message_list_input(self): + req = ResponsesRequest( + input = [ + {"role": "user", "content": "Hi"}, + {"role": "assistant", "content": "Hello!"}, + ], + ) + assert len(req.input) == 2 + assert req.input[0].role == "user" + assert req.input[0].content == "Hi" + + def test_multimodal_input(self): + req = ResponsesRequest( + input = [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "What is in this image?"}, + {"type": "input_image", "image_url": "https://example.com/img.png"}, + ], + }, + ], + ) + parts = req.input[0].content + assert len(parts) == 2 + assert isinstance(parts[0], ResponsesInputTextPart) + assert isinstance(parts[1], ResponsesInputImagePart) + + def test_instructions_field(self): + req = ResponsesRequest( + input = "test", + instructions = "You are a helpful assistant.", + ) + assert req.instructions == "You are a helpful assistant." + + def test_extra_fields_accepted(self): + """OpenAI SDK may send fields we don't model -- extra='allow' should pass.""" + req = ResponsesRequest( + input = "test", + tools = [{"type": "web_search_preview"}], + store = True, + metadata = {"key": "value"}, + previous_response_id = "resp_abc123", + ) + assert req.tools == [{"type": "web_search_preview"}] + assert req.store is True + + def test_stream_flag(self): + req = ResponsesRequest(input = "test", stream = True) + assert req.stream is True + + def test_temperature_and_top_p(self): + req = ResponsesRequest(input = "test", temperature = 0.8, top_p = 0.9) + assert req.temperature == 0.8 + assert req.top_p == 0.9 + + def test_max_output_tokens(self): + req = ResponsesRequest(input = "test", max_output_tokens = 512) + assert req.max_output_tokens == 512 + + def test_developer_role(self): + req = ResponsesRequest( + input = [{"role": "developer", "content": "System instructions"}], + ) + assert req.input[0].role == "developer" + + +# ===================================================================== +# Response model tests +# ===================================================================== + + +class TestResponsesResponse: + """Validate response models serialise correctly.""" + + def test_basic_response(self): + resp = ResponsesResponse( + model = "test-model", + output = [ + ResponsesOutputMessage( + content = [ResponsesOutputTextContent(text = "Hello!")] + ), + ], + usage = ResponsesUsage(input_tokens = 10, output_tokens = 5, total_tokens = 15), + ) + d = resp.model_dump() + assert d["object"] == "response" + assert d["status"] == "completed" + assert d["output"][0]["type"] == "message" + assert d["output"][0]["content"][0]["type"] == "output_text" + assert d["output"][0]["content"][0]["text"] == "Hello!" + assert d["usage"]["input_tokens"] == 10 + assert d["usage"]["output_tokens"] == 5 + assert d["usage"]["total_tokens"] == 15 + # Must NOT have prompt_tokens / completion_tokens + assert "prompt_tokens" not in d["usage"] + assert "completion_tokens" not in d["usage"] + + def test_id_format(self): + resp = ResponsesResponse() + assert resp.id.startswith("resp_") + + def test_output_message_id_format(self): + msg = ResponsesOutputMessage() + assert msg.id.startswith("msg_") + + def test_annotations_default_empty(self): + part = ResponsesOutputTextContent(text = "hi") + assert part.annotations == [] + + def test_response_json_roundtrip(self): + resp = ResponsesResponse( + model = "gpt-4", + output = [ + ResponsesOutputMessage( + content = [ResponsesOutputTextContent(text = "ok")], + ), + ], + usage = ResponsesUsage(input_tokens = 1, output_tokens = 1, total_tokens = 2), + ) + j = json.loads(resp.model_dump_json()) + assert j["object"] == "response" + assert j["output"][0]["role"] == "assistant" + assert j["output"][0]["status"] == "completed" + + +# ===================================================================== +# Input normalisation tests +# ===================================================================== + + +class TestNormaliseResponsesInput: + """Test _normalise_responses_input converts Responses input to ChatMessages.""" + + def test_string_input(self): + payload = ResponsesRequest(input = "Hello world") + msgs = _normalise_responses_input(payload) + assert len(msgs) == 1 + assert msgs[0].role == "user" + assert msgs[0].content == "Hello world" + + def test_instructions_become_system_message(self): + payload = ResponsesRequest( + input = "Hi", + instructions = "Be concise.", + ) + msgs = _normalise_responses_input(payload) + assert len(msgs) == 2 + assert msgs[0].role == "system" + assert msgs[0].content == "Be concise." + assert msgs[1].role == "user" + assert msgs[1].content == "Hi" + + def test_message_list(self): + payload = ResponsesRequest( + input = [ + {"role": "user", "content": "First"}, + {"role": "assistant", "content": "Response"}, + {"role": "user", "content": "Second"}, + ], + ) + msgs = _normalise_responses_input(payload) + assert len(msgs) == 3 + assert msgs[0].role == "user" + assert msgs[1].role == "assistant" + assert msgs[2].role == "user" + + def test_developer_role_maps_to_system(self): + payload = ResponsesRequest( + input = [{"role": "developer", "content": "Instructions"}], + ) + msgs = _normalise_responses_input(payload) + assert msgs[0].role == "system" + assert msgs[0].content == "Instructions" + + def test_multimodal_parts(self): + payload = ResponsesRequest( + input = [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "Describe this:"}, + {"type": "input_image", "image_url": "data:image/png;base64,abc"}, + ], + }, + ], + ) + msgs = _normalise_responses_input(payload) + assert len(msgs) == 1 + content = msgs[0].content + assert isinstance(content, list) + assert len(content) == 2 + assert isinstance(content[0], TextContentPart) + assert content[0].text == "Describe this:" + assert isinstance(content[1], ImageContentPart) + assert content[1].image_url.url == "data:image/png;base64,abc" + + def test_empty_string_input(self): + payload = ResponsesRequest(input = "") + msgs = _normalise_responses_input(payload) + assert len(msgs) == 0 + + def test_empty_list_input(self): + payload = ResponsesRequest(input = []) + msgs = _normalise_responses_input(payload) + assert len(msgs) == 0 + + def test_instructions_only(self): + payload = ResponsesRequest(input = "", instructions = "System msg") + msgs = _normalise_responses_input(payload) + assert len(msgs) == 1 + assert msgs[0].role == "system" + + def test_instructions_plus_message_list(self): + payload = ResponsesRequest( + input = [{"role": "user", "content": "Hello"}], + instructions = "Be brief.", + ) + msgs = _normalise_responses_input(payload) + assert len(msgs) == 2 + assert msgs[0].role == "system" + assert msgs[0].content == "Be brief." + assert msgs[1].role == "user" + + +if __name__ == "__main__": + import pytest + pytest.main([__file__, "-v"])