[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 import torch
from fastapi import APIRouter, FastAPI, Request 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.configs.sample.sampling_params import SamplingParams
from sglang.multimodal_gen.runtime.entrypoints.openai import image_api, video_api 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.scheduler_client import async_scheduler_client
from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args 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.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.srt.utils.json_response import orjson_response
from sglang.version import __version__ from sglang.version import __version__
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -235,7 +235,7 @@ vertex_router = APIRouter()
@vertex_router.post(VERTEX_ROUTE) @vertex_router.post(VERTEX_ROUTE)
async def vertex_generate(vertex_req: VertexGenerateReqInput): async def vertex_generate(vertex_req: VertexGenerateReqInput):
if not vertex_req.instances: if not vertex_req.instances:
return ORJSONResponse({"predictions": []}) return orjson_response({"predictions": []})
server_args = get_global_server_args() server_args = get_global_server_args()
params = vertex_req.parameters or {} params = vertex_req.parameters or {}
@@ -263,7 +263,7 @@ async def vertex_generate(vertex_req: VertexGenerateReqInput):
results = await asyncio.gather(*futures) results = await asyncio.gather(*futures)
return ORJSONResponse({"predictions": results}) return orjson_response({"predictions": results})
def create_app(server_args: ServerArgs): def create_app(server_args: ServerArgs):
@@ -2,7 +2,6 @@ import time
from typing import Any, List, Optional, Union from typing import Any, List, Optional, Union
from fastapi import APIRouter, Body, HTTPException from fastapi import APIRouter, Body, HTTPException
from fastapi.responses import ORJSONResponse
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from sglang.multimodal_gen.registry import get_model_info 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.scheduler_client import async_scheduler_client
from sglang.multimodal_gen.runtime.server_args import get_global_server_args 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.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.srt.utils.json_response import orjson_response
router = APIRouter(prefix="/v1") router = APIRouter(prefix="/v1")
logger = init_logger(__name__) logger = init_logger(__name__)
@@ -173,7 +173,7 @@ async def list_loras():
raise HTTPException(status_code=500, detail=str(e)) raise HTTPException(status_code=500, detail=str(e))
@router.get("/models", response_class=ORJSONResponse) @router.get("/models")
async def available_models(): async def available_models():
"""Show available models. OpenAI-compatible endpoint with extended diffusion info.""" """Show available models. OpenAI-compatible endpoint with extended diffusion info."""
server_args = get_global_server_args() server_args = get_global_server_args()
@@ -206,7 +206,7 @@ async def available_models():
return {"object": "list", "data": [model_card.model_dump()]} 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): async def retrieve_model(model: str):
"""Retrieve a model instance. OpenAI-compatible endpoint with extended diffusion info.""" """Retrieve a model instance. OpenAI-compatible endpoint with extended diffusion info."""
server_args = get_global_server_args() 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") raise HTTPException(status_code=500, detail="Server args not initialized")
if model != server_args.model_path: if model != server_args.model_path:
return ORJSONResponse( return orjson_response(
status_code=404, {
content={
"error": { "error": {
"message": f"The model '{model}' does not exist", "message": f"The model '{model}' does not exist",
"type": "invalid_request_error", "type": "invalid_request_error",
@@ -224,6 +223,7 @@ async def retrieve_model(model: str):
"code": "model_not_found", "code": "model_not_found",
} }
}, },
status_code=404,
) )
model_info = get_model_info( model_info = get_model_info(
@@ -1,13 +1,13 @@
"""Weight update API for the diffusion engine.""" """Weight update API for the diffusion engine."""
from fastapi import APIRouter, Request from fastapi import APIRouter, Request
from fastapi.responses import ORJSONResponse
from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import ( from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
GetWeightsChecksumReqInput, GetWeightsChecksumReqInput,
UpdateWeightFromDiskReqInput, UpdateWeightFromDiskReqInput,
) )
from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client
from sglang.srt.utils.json_response import orjson_response
router = APIRouter() router = APIRouter()
@@ -18,7 +18,7 @@ async def update_weights_from_disk(request: Request):
body = await request.json() body = await request.json()
model_path = body.get("model_path") model_path = body.get("model_path")
if not model_path: if not model_path:
return ORJSONResponse( return orjson_response(
{"success": False, "message": "model_path is required"}, {"success": False, "message": "model_path is required"},
status_code=400, status_code=400,
) )
@@ -32,7 +32,7 @@ async def update_weights_from_disk(request: Request):
try: try:
response = await async_scheduler_client.forward(req) response = await async_scheduler_client.forward(req)
except Exception as e: except Exception as e:
return ORJSONResponse( return orjson_response(
{"success": False, "message": str(e)}, {"success": False, "message": str(e)},
status_code=500, status_code=500,
) )
@@ -40,7 +40,7 @@ async def update_weights_from_disk(request: Request):
result = response.output result = response.output
success = result.get("success", False) success = result.get("success", False)
message = result.get("message", "Unknown status") message = result.get("message", "Unknown status")
return ORJSONResponse( return orjson_response(
{"success": success, "message": message}, {"success": success, "message": message},
status_code=200 if success else 400, status_code=200 if success else 400,
) )
@@ -57,6 +57,6 @@ async def get_weights_checksum(request: Request):
try: try:
response = await async_scheduler_client.forward(req) response = await async_scheduler_client.forward(req)
except Exception as e: 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)