[diffusion] UX: replace deprecated ORJSONResponse with orjson_response (#21755)

Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Mick
2026-03-31 21:41:33 +08:00
committed by GitHub
co-authored by Claude Opus 4.6
parent 20d07c4384
commit 7790645b82
3 changed files with 15 additions and 15 deletions
@@ -9,7 +9,6 @@ from typing import TYPE_CHECKING
import torch
from fastapi import APIRouter, FastAPI, Request
from fastapi.responses import ORJSONResponse
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
from sglang.multimodal_gen.runtime.entrypoints.openai import image_api, video_api
@@ -25,6 +24,7 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client
from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.srt.utils.json_response import orjson_response
from sglang.version import __version__
if TYPE_CHECKING:
@@ -235,7 +235,7 @@ vertex_router = APIRouter()
@vertex_router.post(VERTEX_ROUTE)
async def vertex_generate(vertex_req: VertexGenerateReqInput):
if not vertex_req.instances:
return ORJSONResponse({"predictions": []})
return orjson_response({"predictions": []})
server_args = get_global_server_args()
params = vertex_req.parameters or {}
@@ -263,7 +263,7 @@ async def vertex_generate(vertex_req: VertexGenerateReqInput):
results = await asyncio.gather(*futures)
return ORJSONResponse({"predictions": results})
return orjson_response({"predictions": results})
def create_app(server_args: ServerArgs):
@@ -2,7 +2,6 @@ import time
from typing import Any, List, Optional, Union
from fastapi import APIRouter, Body, HTTPException
from fastapi.responses import ORJSONResponse
from pydantic import BaseModel, Field
from sglang.multimodal_gen.registry import get_model_info
@@ -17,6 +16,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBa
from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client
from sglang.multimodal_gen.runtime.server_args import get_global_server_args
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.srt.utils.json_response import orjson_response
router = APIRouter(prefix="/v1")
logger = init_logger(__name__)
@@ -173,7 +173,7 @@ async def list_loras():
raise HTTPException(status_code=500, detail=str(e))
@router.get("/models", response_class=ORJSONResponse)
@router.get("/models")
async def available_models():
"""Show available models. OpenAI-compatible endpoint with extended diffusion info."""
server_args = get_global_server_args()
@@ -206,7 +206,7 @@ async def available_models():
return {"object": "list", "data": [model_card.model_dump()]}
@router.get("/models/{model:path}", response_class=ORJSONResponse)
@router.get("/models/{model:path}")
async def retrieve_model(model: str):
"""Retrieve a model instance. OpenAI-compatible endpoint with extended diffusion info."""
server_args = get_global_server_args()
@@ -214,9 +214,8 @@ async def retrieve_model(model: str):
raise HTTPException(status_code=500, detail="Server args not initialized")
if model != server_args.model_path:
return ORJSONResponse(
status_code=404,
content={
return orjson_response(
{
"error": {
"message": f"The model '{model}' does not exist",
"type": "invalid_request_error",
@@ -224,6 +223,7 @@ async def retrieve_model(model: str):
"code": "model_not_found",
}
},
status_code=404,
)
model_info = get_model_info(
@@ -1,13 +1,13 @@
"""Weight update API for the diffusion engine."""
from fastapi import APIRouter, Request
from fastapi.responses import ORJSONResponse
from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
GetWeightsChecksumReqInput,
UpdateWeightFromDiskReqInput,
)
from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client
from sglang.srt.utils.json_response import orjson_response
router = APIRouter()
@@ -18,7 +18,7 @@ async def update_weights_from_disk(request: Request):
body = await request.json()
model_path = body.get("model_path")
if not model_path:
return ORJSONResponse(
return orjson_response(
{"success": False, "message": "model_path is required"},
status_code=400,
)
@@ -32,7 +32,7 @@ async def update_weights_from_disk(request: Request):
try:
response = await async_scheduler_client.forward(req)
except Exception as e:
return ORJSONResponse(
return orjson_response(
{"success": False, "message": str(e)},
status_code=500,
)
@@ -40,7 +40,7 @@ async def update_weights_from_disk(request: Request):
result = response.output
success = result.get("success", False)
message = result.get("message", "Unknown status")
return ORJSONResponse(
return orjson_response(
{"success": success, "message": message},
status_code=200 if success else 400,
)
@@ -57,6 +57,6 @@ async def get_weights_checksum(request: Request):
try:
response = await async_scheduler_client.forward(req)
except Exception as e:
return ORJSONResponse({"error": str(e)}, status_code=500)
return orjson_response({"error": str(e)}, status_code=500)
return ORJSONResponse(response.output, status_code=200)
return orjson_response(response.output, status_code=200)