From 7790645b82d8d1293229d68fd0b6fdb17e80a4f5 Mon Sep 17 00:00:00 2001 From: Mick Date: Tue, 31 Mar 2026 21:41:33 +0800 Subject: [PATCH] [diffusion] UX: replace deprecated ORJSONResponse with orjson_response (#21755) Co-authored-by: Claude Opus 4.6 --- .../runtime/entrypoints/http_server.py | 6 +++--- .../runtime/entrypoints/openai/common_api.py | 12 ++++++------ .../runtime/entrypoints/post_training/weights_api.py | 12 ++++++------ 3 files changed, 15 insertions(+), 15 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py index d07febdb1..303cea786 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py @@ -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): diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py index 921f64410..328d9f6f1 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/common_api.py @@ -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( diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/weights_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/weights_api.py index 1b9312d8e..7bc0054f7 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/weights_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/weights_api.py @@ -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)