feat(inference): accept OpenAI multimodal content parts (image_url) in /chat/completions

This commit is contained in:
Roland Tannous 2026-02-14 09:06:25 +00:00
commit 9de38cb773
2 changed files with 116 additions and 19 deletions

View file

@ -5,9 +5,9 @@ from __future__ import annotations
import time
import uuid
from typing import Literal, Optional, List
from typing import Annotated, Literal, Optional, List, Union
from pydantic import BaseModel, Field
from pydantic import BaseModel, Discriminator, Field, Tag
class LoadRequest(BaseModel):
@ -25,7 +25,7 @@ class UnloadRequest(BaseModel):
class GenerateRequest(BaseModel):
"""Request for text generation"""
"""Request for text generation (legacy /generate/stream endpoint)"""
messages: List[dict] = Field(..., description="Chat messages in OpenAI format")
system_prompt: str = Field("You are a helpful AI assistant.", description="System prompt")
temperature: float = Field(0.7, ge=0.0, le=2.0, description="Sampling temperature")
@ -64,10 +64,53 @@ class InferenceStatusResponse(BaseModel):
# =====================================================================
# ── Multimodal content parts (OpenAI vision format) ──────────────
class TextContentPart(BaseModel):
"""Text content part in a multimodal message."""
type: Literal["text"]
text: str
class ImageUrl(BaseModel):
"""Image URL object — supports data URIs and remote URLs."""
url: str = Field(..., description="data:image/png;base64,... or https://...")
detail: Optional[Literal["auto", "low", "high"]] = "auto"
class ImageContentPart(BaseModel):
"""Image content part in a multimodal message."""
type: Literal["image_url"]
image_url: ImageUrl
def _content_part_discriminator(v):
if isinstance(v, dict):
return v.get("type")
return getattr(v, "type", None)
ContentPart = Annotated[
Union[
Annotated[TextContentPart, Tag("text")],
Annotated[ImageContentPart, Tag("image_url")],
],
Discriminator(_content_part_discriminator),
]
"""Union type for multimodal content parts, discriminated by the 'type' field."""
# ── Messages ─────────────────────────────────────────────────────
class ChatMessage(BaseModel):
"""A single message in the conversation."""
"""
A single message in the conversation.
``content`` may be a plain string (text-only) or a list of
content parts for multimodal messages (OpenAI vision format).
"""
role: Literal["system", "user", "assistant"] = Field(..., description="Message role")
content: str = Field(..., description="Message content")
content: Union[str, list[ContentPart]] = Field(..., description="Message content (string or multimodal parts)")
class ChatCompletionRequest(BaseModel):

View file

@ -240,6 +240,59 @@ async def get_status():
# =====================================================================
def _extract_content_parts(
messages: list,
) -> tuple[str, list[dict], "Optional[str]"]:
"""
Parse OpenAI-format messages into components the inference backend expects.
Handles both plain-string ``content`` and multimodal content-part arrays
(``[{type: "text", ...}, {type: "image_url", ...}]``).
Returns:
system_prompt: The system message text (or a default).
chat_messages: Non-system messages with content flattened to strings.
image_base64: Base64 data of the *first* image found, or ``None``.
"""
system_prompt = "You are a helpful AI assistant."
chat_messages: list[dict] = []
first_image_b64: Optional[str] = None
for msg in messages:
# ── System messages → extract as system_prompt ────────
if msg.role == "system":
if isinstance(msg.content, str):
system_prompt = msg.content
elif isinstance(msg.content, list):
# Unlikely but handle: join text parts
system_prompt = "\n".join(
p.text for p in msg.content if p.type == "text"
)
continue
# ── User / assistant messages ─────────────────────────
if isinstance(msg.content, str):
# Plain string content — pass through
chat_messages.append({"role": msg.role, "content": msg.content})
elif isinstance(msg.content, list):
# Multimodal content parts
text_parts: list[str] = []
for part in msg.content:
if part.type == "text":
text_parts.append(part.text)
elif part.type == "image_url" and first_image_b64 is None:
url = part.image_url.url
if url.startswith("data:"):
# data:image/png;base64,<DATA> → extract <DATA>
first_image_b64 = url.split(",", 1)[1] if "," in url else None
else:
logger.warning(
f"Remote image URLs not yet supported: {url[:80]}..."
)
combined_text = "\n".join(text_parts) if text_parts else ""
chat_messages.append({"role": msg.role, "content": combined_text})
return system_prompt, chat_messages, first_image_b64
@router.post("/chat/completions")
@ -247,6 +300,9 @@ async def openai_chat_completions(request: ChatCompletionRequest):
"""
OpenAI-compatible chat completions endpoint.
Supports multimodal messages: ``content`` may be a plain string or a
list of content parts (``text`` / ``image_url``).
Streaming (default): returns SSE chunks matching OpenAI's format.
Non-streaming: returns a single ChatCompletion JSON object.
"""
@ -258,15 +314,10 @@ async def openai_chat_completions(request: ChatCompletionRequest):
detail="No model loaded. Call POST /inference/load first.",
)
# ── Extract system prompt from messages ───────────────────────
system_prompt = "You are a helpful AI assistant."
chat_messages: list[dict] = []
for msg in request.messages:
if msg.role == "system":
system_prompt = msg.content
else:
chat_messages.append({"role": msg.role, "content": msg.content})
# ── Parse messages (handles multimodal content parts) ─────
system_prompt, chat_messages, extracted_image_b64 = _extract_content_parts(
request.messages
)
# If no non-system messages were provided, error out
if not chat_messages:
@ -275,9 +326,12 @@ async def openai_chat_completions(request: ChatCompletionRequest):
detail="At least one non-system message is required.",
)
# ── Decode image if provided (vision models) ──────────────────
# ── Decode image (from content parts OR legacy field) ─────
# Content-part images take priority; fall back to legacy field
image_b64 = extracted_image_b64 or request.image_base64
image = None
if request.image_base64:
if image_b64:
try:
import base64
from PIL import Image
@ -287,10 +341,10 @@ async def openai_chat_completions(request: ChatCompletionRequest):
if not model_info.get("is_vision"):
raise HTTPException(
status_code=400,
detail="Image provided but current model is text-only.",
detail="Image provided but current model is text-only. Load a vision model.",
)
image_data = base64.b64decode(request.image_base64)
image_data = base64.b64decode(image_b64)
image = Image.open(BytesIO(image_data))
image = backend.resize_image(image)
@ -299,7 +353,7 @@ async def openai_chat_completions(request: ChatCompletionRequest):
except Exception as e:
raise HTTPException(status_code=400, detail=f"Failed to decode image: {e}")
# ── Shared generation kwargs ──────────────────────────────────
# ── Shared generation kwargs ──────────────────────────────
gen_kwargs = dict(
messages=chat_messages,
system_prompt=system_prompt,