[diffusion] feat: add cloud storage support for API (#14579)
This commit is contained in:
@@ -154,6 +154,32 @@ sglang serve "${SERVER_ARGS[@]}"
|
|||||||
|
|
||||||
For detailed API usage, including Image, Video Generation and LoRA management, please refer to the [OpenAI API Documentation](openai_api.md).
|
For detailed API usage, including Image, Video Generation and LoRA management, please refer to the [OpenAI API Documentation](openai_api.md).
|
||||||
|
|
||||||
|
### Cloud Storage Support
|
||||||
|
|
||||||
|
SGLang diffusion supports automatically uploading generated images and videos to S3-compatible cloud storage (e.g., AWS S3, MinIO, Alibaba Cloud OSS, Tencent Cloud COS).
|
||||||
|
|
||||||
|
When enabled, the server follows a **Generate -> Upload -> Delete** workflow:
|
||||||
|
1. The artifact is generated to a temporary local file.
|
||||||
|
2. The file is immediately uploaded to the configured S3 bucket in a background thread.
|
||||||
|
3. Upon successful upload, the local file is deleted.
|
||||||
|
4. The API response returns the public URL of the uploaded object.
|
||||||
|
|
||||||
|
#### Configuration
|
||||||
|
|
||||||
|
Cloud storage is enabled via environment variables. Note that `boto3` must be installed separately (`pip install boto3`) to use this feature.
|
||||||
|
|
||||||
|
```bash
|
||||||
|
# Enable S3 storage
|
||||||
|
export SGLANG_CLOUD_STORAGE_TYPE=s3
|
||||||
|
export SGLANG_S3_BUCKET_NAME=my-bucket
|
||||||
|
export SGLANG_S3_ACCESS_KEY_ID=your-access-key
|
||||||
|
export SGLANG_S3_SECRET_ACCESS_KEY=your-secret-key
|
||||||
|
|
||||||
|
# Optional: Custom endpoint for MinIO/OSS/COS
|
||||||
|
export SGLANG_S3_ENDPOINT_URL=https://minio.example.com
|
||||||
|
```
|
||||||
|
|
||||||
|
See [Environment Variables Documentation](environment_variables.md) for more details.
|
||||||
|
|
||||||
## Generate
|
## Generate
|
||||||
|
|
||||||
|
|||||||
@@ -17,3 +17,16 @@ See [cache-dit documentation](cache_dit.md) for details.
|
|||||||
| `SGLANG_CACHE_DIT_SCM_POLICY` | dynamic | SCM caching policy |
|
| `SGLANG_CACHE_DIT_SCM_POLICY` | dynamic | SCM caching policy |
|
||||||
| `SGLANG_CACHE_DIT_SCM_COMPUTE_BINS` | not set | Custom SCM compute bins |
|
| `SGLANG_CACHE_DIT_SCM_COMPUTE_BINS` | not set | Custom SCM compute bins |
|
||||||
| `SGLANG_CACHE_DIT_SCM_CACHE_BINS` | not set | Custom SCM cache bins |
|
| `SGLANG_CACHE_DIT_SCM_CACHE_BINS` | not set | Custom SCM cache bins |
|
||||||
|
|
||||||
|
## Cloud Storage
|
||||||
|
|
||||||
|
These variables configure S3-compatible cloud storage for automatically uploading generated images and videos.
|
||||||
|
|
||||||
|
| Environment Variable | Default | Description |
|
||||||
|
|---------------------------------|---------|--------------------------------------------------------|
|
||||||
|
| `SGLANG_CLOUD_STORAGE_TYPE` | not set | Set to `s3` to enable cloud storage |
|
||||||
|
| `SGLANG_S3_BUCKET_NAME` | not set | The name of the S3 bucket |
|
||||||
|
| `SGLANG_S3_ENDPOINT_URL` | not set | Custom endpoint URL (for MinIO, OSS, etc.) |
|
||||||
|
| `SGLANG_S3_REGION_NAME` | us-east-1 | AWS region name |
|
||||||
|
| `SGLANG_S3_ACCESS_KEY_ID` | not set | AWS Access Key ID |
|
||||||
|
| `SGLANG_S3_SECRET_ACCESS_KEY` | not set | AWS Secret Access Key |
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import (
|
|||||||
ImageResponse,
|
ImageResponse,
|
||||||
ImageResponseData,
|
ImageResponseData,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.entrypoints.openai.storage import cloud_storage
|
||||||
from sglang.multimodal_gen.runtime.entrypoints.openai.stores import IMAGE_STORE
|
from sglang.multimodal_gen.runtime.entrypoints.openai.stores import IMAGE_STORE
|
||||||
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
|
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
|
||||||
_parse_size,
|
_parse_size,
|
||||||
@@ -136,39 +137,64 @@ async def generations(
|
|||||||
)
|
)
|
||||||
save_file_path = save_file_path_list[0]
|
save_file_path = save_file_path_list[0]
|
||||||
|
|
||||||
|
resp_format = (request.response_format or "b64_json").lower()
|
||||||
|
b64_data = None
|
||||||
|
|
||||||
|
# 1. Read content first if needed (while file exists)
|
||||||
|
if resp_format == "b64_json":
|
||||||
|
with open(save_file_path, "rb") as f:
|
||||||
|
b64_data = base64.b64encode(f.read()).decode("utf-8")
|
||||||
|
|
||||||
|
# 2. Upload and Delete local file
|
||||||
|
cloud_url = await cloud_storage.upload_and_cleanup(save_file_path)
|
||||||
|
|
||||||
|
# 3. Update Database
|
||||||
await IMAGE_STORE.upsert(
|
await IMAGE_STORE.upsert(
|
||||||
request_id,
|
request_id,
|
||||||
{
|
{
|
||||||
"id": request_id,
|
"id": request_id,
|
||||||
"created_at": int(time.time()),
|
"created_at": int(time.time()),
|
||||||
"file_path": save_file_path,
|
"file_path": None if cloud_url else save_file_path,
|
||||||
|
"url": cloud_url,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
resp_format = (request.response_format or "b64_json").lower()
|
# 4. Return Response
|
||||||
|
|
||||||
if resp_format == "b64_json":
|
if resp_format == "b64_json":
|
||||||
with open(save_file_path, "rb") as f:
|
|
||||||
b64 = base64.b64encode(f.read()).decode("utf-8")
|
|
||||||
response_kwargs = {
|
response_kwargs = {
|
||||||
"data": [
|
"data": [
|
||||||
ImageResponseData(
|
ImageResponseData(
|
||||||
b64_json=b64,
|
b64_json=b64_data,
|
||||||
|
revised_prompt=request.prompt,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
}
|
||||||
|
elif resp_format == "url":
|
||||||
|
if not cloud_url:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail="response_format='url' requires cloud storage to be configured.",
|
||||||
|
)
|
||||||
|
response_kwargs = {
|
||||||
|
"data": [
|
||||||
|
ImageResponseData(
|
||||||
|
url=cloud_url,
|
||||||
revised_prompt=request.prompt,
|
revised_prompt=request.prompt,
|
||||||
file_path=os.path.abspath(save_file_path),
|
file_path=os.path.abspath(save_file_path),
|
||||||
)
|
)
|
||||||
],
|
],
|
||||||
}
|
}
|
||||||
response_kwargs = add_common_data_to_response(
|
|
||||||
response_kwargs, request_id=request_id, result=result
|
|
||||||
)
|
|
||||||
return ImageResponse(**response_kwargs)
|
|
||||||
else:
|
else:
|
||||||
# Return error, not supported
|
# Return error, not supported
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=400, detail="response_format=url is not supported"
|
status_code=400, detail=f"response_format={resp_format} is not supported"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
response_kwargs = add_common_data_to_response(
|
||||||
|
response_kwargs, request_id=request_id, result=result
|
||||||
|
)
|
||||||
|
return ImageResponse(**response_kwargs)
|
||||||
|
|
||||||
|
|
||||||
@router.post("/edits", response_model=ImageResponse)
|
@router.post("/edits", response_model=ImageResponse)
|
||||||
async def edits(
|
async def edits(
|
||||||
@@ -248,38 +274,54 @@ async def edits(
|
|||||||
async_scheduler_client, batch
|
async_scheduler_client, batch
|
||||||
)
|
)
|
||||||
save_file_path = save_file_path_list[0]
|
save_file_path = save_file_path_list[0]
|
||||||
|
|
||||||
|
resp_format = (response_format or "b64_json").lower()
|
||||||
|
b64_data = None
|
||||||
|
|
||||||
|
# 1. Read content first if needed (while file exists)
|
||||||
|
if resp_format == "b64_json":
|
||||||
|
with open(save_file_path, "rb") as f:
|
||||||
|
b64_data = base64.b64encode(f.read()).decode("utf-8")
|
||||||
|
|
||||||
|
# 2. Upload and Delete local file
|
||||||
|
cloud_url = await cloud_storage.upload_and_cleanup(save_file_path)
|
||||||
|
|
||||||
|
# 3. Update Database
|
||||||
await IMAGE_STORE.upsert(
|
await IMAGE_STORE.upsert(
|
||||||
request_id,
|
request_id,
|
||||||
{
|
{
|
||||||
"id": request_id,
|
"id": request_id,
|
||||||
"created_at": int(time.time()),
|
"created_at": int(time.time()),
|
||||||
"file_path": save_file_path,
|
"file_path": None if cloud_url else save_file_path,
|
||||||
|
"url": cloud_url,
|
||||||
"input_image_paths": input_paths, # Store all input image paths
|
"input_image_paths": input_paths, # Store all input image paths
|
||||||
"num_input_images": len(input_paths),
|
"num_input_images": len(input_paths),
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
# Default to b64_json to align with gpt-image-1 behavior in OpenAI examples
|
# 4. Return Response
|
||||||
if (response_format or "b64_json").lower() == "b64_json":
|
if (response_format or "b64_json").lower() == "b64_json":
|
||||||
response_kwargs = {"data": []}
|
response_kwargs = {"data": []}
|
||||||
for save_file_path in save_file_path_list:
|
for path in save_file_path_list:
|
||||||
with open(save_file_path, "rb") as f:
|
if path == save_file_path and b64_data is not None:
|
||||||
b64 = base64.b64encode(f.read()).decode("utf-8")
|
b64 = b64_data
|
||||||
response_kwargs["data"].append(
|
else:
|
||||||
ImageResponseData(
|
with open(path, "rb") as f:
|
||||||
b64_json=b64,
|
b64 = base64.b64encode(f.read()).decode("utf-8")
|
||||||
revised_prompt=prompt,
|
response_kwargs["data"].append(
|
||||||
file_path=os.path.abspath(save_file_path),
|
ImageResponseData(
|
||||||
)
|
b64_json=b64,
|
||||||
|
revised_prompt=prompt,
|
||||||
|
file_path=os.path.abspath(path),
|
||||||
)
|
)
|
||||||
|
)
|
||||||
if result.peak_memory_mb and result.peak_memory_mb > 0:
|
if result.peak_memory_mb and result.peak_memory_mb > 0:
|
||||||
response_kwargs["peak_memory_mb"] = result.peak_memory_mb
|
response_kwargs["peak_memory_mb"] = result.peak_memory_mb
|
||||||
else:
|
else:
|
||||||
url = f"/v1/images/{request_id}/content"
|
|
||||||
response_kwargs = {
|
response_kwargs = {
|
||||||
"data": [
|
"data": [
|
||||||
ImageResponseData(
|
ImageResponseData(
|
||||||
url=url,
|
url=cloud_url if cloud_url else f"/v1/images/{request_id}/content",
|
||||||
revised_prompt=prompt,
|
revised_prompt=prompt,
|
||||||
file_path=os.path.abspath(save_file_path),
|
file_path=os.path.abspath(save_file_path),
|
||||||
)
|
)
|
||||||
@@ -289,7 +331,6 @@ async def edits(
|
|||||||
response_kwargs = add_common_data_to_response(
|
response_kwargs = add_common_data_to_response(
|
||||||
response_kwargs, request_id=request_id, result=result
|
response_kwargs, request_id=request_id, result=result
|
||||||
)
|
)
|
||||||
|
|
||||||
return ImageResponse(**response_kwargs)
|
return ImageResponse(**response_kwargs)
|
||||||
|
|
||||||
|
|
||||||
@@ -301,6 +342,12 @@ async def download_image_content(
|
|||||||
if not item:
|
if not item:
|
||||||
raise HTTPException(status_code=404, detail="Image not found")
|
raise HTTPException(status_code=404, detail="Image not found")
|
||||||
|
|
||||||
|
if item.get("url"):
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail=f"Image has been uploaded to cloud storage. Please use the cloud URL: {item.get('url')}",
|
||||||
|
)
|
||||||
|
|
||||||
file_path = item.get("file_path")
|
file_path = item.get("file_path")
|
||||||
if not file_path or not os.path.exists(file_path):
|
if not file_path or not os.path.exists(file_path):
|
||||||
raise HTTPException(status_code=404, detail="Image is still being generated")
|
raise HTTPException(status_code=404, detail="Image is still being generated")
|
||||||
|
|||||||
@@ -55,6 +55,7 @@ class VideoResponse(BaseModel):
|
|||||||
size: str = ""
|
size: str = ""
|
||||||
seconds: str = "4"
|
seconds: str = "4"
|
||||||
quality: str = "standard"
|
quality: str = "standard"
|
||||||
|
url: Optional[str] = None
|
||||||
remixed_from_video_id: Optional[str] = None
|
remixed_from_video_id: Optional[str] = None
|
||||||
completed_at: Optional[int] = None
|
completed_at: Optional[int] = None
|
||||||
expires_at: Optional[int] = None
|
expires_at: Optional[int] = None
|
||||||
|
|||||||
@@ -0,0 +1,107 @@
|
|||||||
|
import asyncio
|
||||||
|
import os
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
|
|
||||||
|
logger = init_logger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class CloudStorage:
|
||||||
|
def __init__(self):
|
||||||
|
self.enabled = os.getenv("SGLANG_CLOUD_STORAGE_TYPE", "").lower() == "s3"
|
||||||
|
if not self.enabled:
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
import boto3
|
||||||
|
except ImportError:
|
||||||
|
logger.error(
|
||||||
|
"boto3 is not installed. Please install it with `pip install boto3` to use cloud storage."
|
||||||
|
)
|
||||||
|
self.enabled = False
|
||||||
|
return
|
||||||
|
|
||||||
|
self.bucket_name = os.getenv("SGLANG_S3_BUCKET_NAME")
|
||||||
|
if not self.bucket_name:
|
||||||
|
self.enabled = False
|
||||||
|
return
|
||||||
|
|
||||||
|
endpoint_url = os.getenv("SGLANG_S3_ENDPOINT_URL") or None
|
||||||
|
region_name = os.getenv("SGLANG_S3_REGION_NAME") or None
|
||||||
|
|
||||||
|
self.client = boto3.client(
|
||||||
|
"s3",
|
||||||
|
aws_access_key_id=os.getenv("SGLANG_S3_ACCESS_KEY_ID"),
|
||||||
|
aws_secret_access_key=os.getenv("SGLANG_S3_SECRET_ACCESS_KEY"),
|
||||||
|
endpoint_url=endpoint_url,
|
||||||
|
region_name=region_name,
|
||||||
|
)
|
||||||
|
self.endpoint_url = endpoint_url
|
||||||
|
self.region_name = region_name
|
||||||
|
|
||||||
|
def is_enabled(self) -> bool:
|
||||||
|
return self.enabled
|
||||||
|
|
||||||
|
async def upload_file(self, local_path: str, destination_key: str) -> Optional[str]:
|
||||||
|
if not self.is_enabled():
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _sync_upload():
|
||||||
|
"""Synchronous part of the upload to run in a thread."""
|
||||||
|
ext = os.path.splitext(local_path)[1].lower()
|
||||||
|
content_type = {
|
||||||
|
".png": "image/png",
|
||||||
|
".jpg": "image/jpeg",
|
||||||
|
".jpeg": "image/jpeg",
|
||||||
|
".webp": "image/webp",
|
||||||
|
".mp4": "video/mp4",
|
||||||
|
}.get(ext, "application/octet-stream")
|
||||||
|
|
||||||
|
# Use the client created once in __init__
|
||||||
|
self.client.upload_file(
|
||||||
|
local_path,
|
||||||
|
self.bucket_name,
|
||||||
|
destination_key,
|
||||||
|
ExtraArgs={"ContentType": content_type},
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Offload the blocking I/O call to a thread executor
|
||||||
|
await asyncio.get_running_loop().run_in_executor(None, _sync_upload)
|
||||||
|
except Exception as e:
|
||||||
|
# If upload fails, log the error and return None for fallback
|
||||||
|
logger.error(f"Upload failed for {destination_key}: {e}")
|
||||||
|
return None
|
||||||
|
|
||||||
|
# Simplified URL generation with a default region
|
||||||
|
if self.endpoint_url:
|
||||||
|
url = (
|
||||||
|
f"{self.endpoint_url.rstrip('/')}/{self.bucket_name}/{destination_key}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
region = self.region_name or "us-east-1"
|
||||||
|
url = f"https://{self.bucket_name}.s3.{region}.amazonaws.com/{destination_key}"
|
||||||
|
|
||||||
|
logger.info(f"Uploaded {local_path} to {url}")
|
||||||
|
return url
|
||||||
|
|
||||||
|
async def upload_and_cleanup(self, file_path: str) -> Optional[str]:
|
||||||
|
"""Helper to upload a file and delete the local copy if successful."""
|
||||||
|
if not self.is_enabled():
|
||||||
|
return None
|
||||||
|
|
||||||
|
key = os.path.basename(file_path)
|
||||||
|
url = await self.upload_file(file_path, key)
|
||||||
|
|
||||||
|
if url:
|
||||||
|
try:
|
||||||
|
# pass if removal fails
|
||||||
|
os.remove(file_path)
|
||||||
|
except OSError as e:
|
||||||
|
logger.warning(f"Failed to remove temporary file {file_path}: {e}")
|
||||||
|
return url
|
||||||
|
|
||||||
|
|
||||||
|
# Global instance
|
||||||
|
cloud_storage = CloudStorage()
|
||||||
@@ -27,6 +27,7 @@ from sglang.multimodal_gen.runtime.entrypoints.openai.protocol import (
|
|||||||
VideoListResponse,
|
VideoListResponse,
|
||||||
VideoResponse,
|
VideoResponse,
|
||||||
)
|
)
|
||||||
|
from sglang.multimodal_gen.runtime.entrypoints.openai.storage import cloud_storage
|
||||||
from sglang.multimodal_gen.runtime.entrypoints.openai.stores import VIDEO_STORE
|
from sglang.multimodal_gen.runtime.entrypoints.openai.stores import VIDEO_STORE
|
||||||
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
|
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
|
||||||
_parse_size,
|
_parse_size,
|
||||||
@@ -127,11 +128,19 @@ async def _dispatch_job_async(job_id: str, batch: Req) -> None:
|
|||||||
from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client
|
from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client
|
||||||
|
|
||||||
try:
|
try:
|
||||||
_, result = await process_generation_batch(async_scheduler_client, batch)
|
save_file_path_list, result = await process_generation_batch(
|
||||||
|
async_scheduler_client, batch
|
||||||
|
)
|
||||||
|
save_file_path = save_file_path_list[0]
|
||||||
|
|
||||||
|
cloud_url = await cloud_storage.upload_and_cleanup(save_file_path)
|
||||||
|
|
||||||
update_fields = {
|
update_fields = {
|
||||||
"status": "completed",
|
"status": "completed",
|
||||||
"progress": 100,
|
"progress": 100,
|
||||||
"completed_at": int(time.time()),
|
"completed_at": int(time.time()),
|
||||||
|
"url": cloud_url,
|
||||||
|
"file_path": save_file_path if not cloud_url else None,
|
||||||
}
|
}
|
||||||
update_fields = add_common_data_to_response(
|
update_fields = add_common_data_to_response(
|
||||||
update_fields, request_id=job_id, result=result
|
update_fields, request_id=job_id, result=result
|
||||||
@@ -334,6 +343,12 @@ async def download_video_content(
|
|||||||
if not job:
|
if not job:
|
||||||
raise HTTPException(status_code=404, detail="Video not found")
|
raise HTTPException(status_code=404, detail="Video not found")
|
||||||
|
|
||||||
|
if job.get("url"):
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=400,
|
||||||
|
detail=f"Video has been uploaded to cloud storage. Please use the cloud URL: {job.get('url')}",
|
||||||
|
)
|
||||||
|
|
||||||
file_path = job.get("file_path")
|
file_path = job.get("file_path")
|
||||||
if not file_path or not os.path.exists(file_path):
|
if not file_path or not os.path.exists(file_path):
|
||||||
raise HTTPException(status_code=404, detail="Generation is still in-progress")
|
raise HTTPException(status_code=404, detail="Generation is still in-progress")
|
||||||
|
|||||||
@@ -0,0 +1,232 @@
|
|||||||
|
"""
|
||||||
|
Test suite for S3 CloudStorage integration.
|
||||||
|
|
||||||
|
Tests verify file upload, cleanup, URL generation, and error handling.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import importlib
|
||||||
|
import os
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
import sglang.multimodal_gen.runtime.entrypoints.openai.storage as storage_mod
|
||||||
|
from sglang.multimodal_gen.runtime.entrypoints.openai.storage import CloudStorage
|
||||||
|
|
||||||
|
|
||||||
|
def _create_temp_file(tmp_path, name="test.png", content=b"\x89PNG\r\n\x1a\nfake"):
|
||||||
|
"""Create a temporary test file."""
|
||||||
|
p = tmp_path / name
|
||||||
|
p.write_bytes(content)
|
||||||
|
return str(p)
|
||||||
|
|
||||||
|
|
||||||
|
# UNIT TESTS
|
||||||
|
|
||||||
|
|
||||||
|
def test_upload_file_success(tmp_path):
|
||||||
|
"""Test successful upload with correct URL generation."""
|
||||||
|
file_path = _create_temp_file(tmp_path, "image.png")
|
||||||
|
|
||||||
|
storage_mod.cloud_storage.enabled = True
|
||||||
|
storage_mod.cloud_storage.bucket_name = "my-bucket"
|
||||||
|
storage_mod.cloud_storage.endpoint_url = "https://s3.example.com"
|
||||||
|
storage_mod.cloud_storage.region_name = None
|
||||||
|
|
||||||
|
called = {}
|
||||||
|
|
||||||
|
def fake_upload(local_path, bucket, key, ExtraArgs=None):
|
||||||
|
called["local_path"] = local_path
|
||||||
|
called["bucket"] = bucket
|
||||||
|
called["key"] = key
|
||||||
|
called["extra"] = ExtraArgs
|
||||||
|
|
||||||
|
storage_mod.cloud_storage.client = SimpleNamespace(upload_file=fake_upload)
|
||||||
|
|
||||||
|
url = asyncio.run(storage_mod.cloud_storage.upload_file(file_path, "image.png"))
|
||||||
|
|
||||||
|
assert url == "https://s3.example.com/my-bucket/image.png"
|
||||||
|
assert called["local_path"] == file_path
|
||||||
|
assert called["bucket"] == "my-bucket"
|
||||||
|
assert called["key"] == "image.png"
|
||||||
|
assert called["extra"]["ContentType"] == "image/png"
|
||||||
|
|
||||||
|
|
||||||
|
def test_upload_and_cleanup(tmp_path):
|
||||||
|
"""Test that local file is deleted after successful upload."""
|
||||||
|
file_path = _create_temp_file(tmp_path, "cleanup.png")
|
||||||
|
|
||||||
|
storage_mod.cloud_storage.enabled = True
|
||||||
|
storage_mod.cloud_storage.bucket_name = "my-bucket"
|
||||||
|
storage_mod.cloud_storage.endpoint_url = "https://s3.example.com"
|
||||||
|
storage_mod.cloud_storage.client = SimpleNamespace(
|
||||||
|
upload_file=lambda *args, **kwargs: None
|
||||||
|
)
|
||||||
|
|
||||||
|
assert os.path.exists(file_path)
|
||||||
|
|
||||||
|
url = asyncio.run(storage_mod.cloud_storage.upload_and_cleanup(file_path))
|
||||||
|
|
||||||
|
assert url == "https://s3.example.com/my-bucket/cleanup.png"
|
||||||
|
assert not os.path.exists(file_path)
|
||||||
|
|
||||||
|
|
||||||
|
def test_upload_failure_preserves_file(tmp_path):
|
||||||
|
"""Test that file is preserved when upload fails."""
|
||||||
|
file_path = _create_temp_file(tmp_path, "preserve.png")
|
||||||
|
|
||||||
|
storage_mod.cloud_storage.enabled = True
|
||||||
|
storage_mod.cloud_storage.bucket_name = "my-bucket"
|
||||||
|
storage_mod.cloud_storage.endpoint_url = "https://s3.example.com"
|
||||||
|
|
||||||
|
def fake_upload_raises(*args, **kwargs):
|
||||||
|
raise RuntimeError("simulated failure")
|
||||||
|
|
||||||
|
storage_mod.cloud_storage.client = SimpleNamespace(upload_file=fake_upload_raises)
|
||||||
|
|
||||||
|
result = asyncio.run(storage_mod.cloud_storage.upload_and_cleanup(file_path))
|
||||||
|
|
||||||
|
assert result is None
|
||||||
|
assert os.path.exists(file_path)
|
||||||
|
|
||||||
|
|
||||||
|
def test_disabled_storage_returns_none(tmp_path):
|
||||||
|
"""Test that disabled storage returns None."""
|
||||||
|
file_path = _create_temp_file(tmp_path, "test.png")
|
||||||
|
|
||||||
|
prev_enabled = storage_mod.cloud_storage.enabled
|
||||||
|
storage_mod.cloud_storage.enabled = False
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = asyncio.run(
|
||||||
|
storage_mod.cloud_storage.upload_file(file_path, "test.png")
|
||||||
|
)
|
||||||
|
assert result is None
|
||||||
|
finally:
|
||||||
|
storage_mod.cloud_storage.enabled = prev_enabled
|
||||||
|
|
||||||
|
|
||||||
|
def test_aws_url_with_region(tmp_path):
|
||||||
|
"""Test AWS S3 URL generation with specific region."""
|
||||||
|
file_path = _create_temp_file(tmp_path, "aws.png")
|
||||||
|
|
||||||
|
storage_mod.cloud_storage.enabled = True
|
||||||
|
storage_mod.cloud_storage.bucket_name = "aws-bucket"
|
||||||
|
storage_mod.cloud_storage.endpoint_url = None
|
||||||
|
storage_mod.cloud_storage.region_name = "us-west-2"
|
||||||
|
storage_mod.cloud_storage.client = SimpleNamespace(
|
||||||
|
upload_file=lambda *args, **kwargs: None
|
||||||
|
)
|
||||||
|
|
||||||
|
url = asyncio.run(storage_mod.cloud_storage.upload_file(file_path, "aws.png"))
|
||||||
|
|
||||||
|
assert url == "https://aws-bucket.s3.us-west-2.amazonaws.com/aws.png"
|
||||||
|
|
||||||
|
|
||||||
|
def test_aws_url_default_region(tmp_path):
|
||||||
|
"""Test AWS S3 URL defaults to us-east-1 when region not specified."""
|
||||||
|
file_path = _create_temp_file(tmp_path, "default.png")
|
||||||
|
|
||||||
|
storage_mod.cloud_storage.enabled = True
|
||||||
|
storage_mod.cloud_storage.bucket_name = "default-bucket"
|
||||||
|
storage_mod.cloud_storage.endpoint_url = None
|
||||||
|
storage_mod.cloud_storage.region_name = None
|
||||||
|
storage_mod.cloud_storage.client = SimpleNamespace(
|
||||||
|
upload_file=lambda *args, **kwargs: None
|
||||||
|
)
|
||||||
|
|
||||||
|
url = asyncio.run(storage_mod.cloud_storage.upload_file(file_path, "default.png"))
|
||||||
|
|
||||||
|
assert url == "https://default-bucket.s3.us-east-1.amazonaws.com/default.png"
|
||||||
|
|
||||||
|
|
||||||
|
def test_custom_endpoint_url(tmp_path):
|
||||||
|
"""Test URL generation with custom endpoint (MinIO/OSS/COS)."""
|
||||||
|
file_path = _create_temp_file(tmp_path, "custom.png")
|
||||||
|
|
||||||
|
storage_mod.cloud_storage.enabled = True
|
||||||
|
storage_mod.cloud_storage.bucket_name = "custom-bucket"
|
||||||
|
storage_mod.cloud_storage.endpoint_url = "https://minio.example.com/"
|
||||||
|
storage_mod.cloud_storage.region_name = None
|
||||||
|
storage_mod.cloud_storage.client = SimpleNamespace(
|
||||||
|
upload_file=lambda *args, **kwargs: None
|
||||||
|
)
|
||||||
|
|
||||||
|
url = asyncio.run(storage_mod.cloud_storage.upload_file(file_path, "custom.png"))
|
||||||
|
|
||||||
|
# Verify trailing slash is stripped
|
||||||
|
assert url == "https://minio.example.com/custom-bucket/custom.png"
|
||||||
|
|
||||||
|
|
||||||
|
def test_content_type_detection(tmp_path):
|
||||||
|
"""Test Content-Type header for different file extensions."""
|
||||||
|
storage_mod.cloud_storage.enabled = True
|
||||||
|
storage_mod.cloud_storage.bucket_name = "test-bucket"
|
||||||
|
storage_mod.cloud_storage.endpoint_url = "https://s3.test"
|
||||||
|
|
||||||
|
test_cases = [
|
||||||
|
("image.png", "image/png"),
|
||||||
|
("image.jpg", "image/jpeg"),
|
||||||
|
("image.jpeg", "image/jpeg"),
|
||||||
|
("image.webp", "image/webp"),
|
||||||
|
("video.mp4", "video/mp4"),
|
||||||
|
("file.bin", "application/octet-stream"),
|
||||||
|
]
|
||||||
|
|
||||||
|
for filename, expected_type in test_cases:
|
||||||
|
called = {}
|
||||||
|
|
||||||
|
def fake_upload(local_path, bucket, key, ExtraArgs=None):
|
||||||
|
called["content_type"] = ExtraArgs.get("ContentType")
|
||||||
|
|
||||||
|
storage_mod.cloud_storage.client = SimpleNamespace(upload_file=fake_upload)
|
||||||
|
|
||||||
|
file_path = _create_temp_file(tmp_path, filename)
|
||||||
|
asyncio.run(storage_mod.cloud_storage.upload_file(file_path, filename))
|
||||||
|
|
||||||
|
assert called["content_type"] == expected_type
|
||||||
|
|
||||||
|
|
||||||
|
# requires moto and boto3
|
||||||
|
has_moto = (
|
||||||
|
importlib.util.find_spec("moto") is not None
|
||||||
|
and importlib.util.find_spec("boto3") is not None
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not has_moto, reason="moto/boto3 not installed")
|
||||||
|
def test_integration_with_moto(tmp_path):
|
||||||
|
"""Integration test using moto to mock real S3 service."""
|
||||||
|
import boto3
|
||||||
|
from moto import mock_aws
|
||||||
|
|
||||||
|
os.environ["SGLANG_CLOUD_STORAGE_TYPE"] = "s3"
|
||||||
|
os.environ["SGLANG_S3_BUCKET_NAME"] = "integration-test"
|
||||||
|
os.environ["SGLANG_S3_REGION_NAME"] = "us-east-1"
|
||||||
|
|
||||||
|
with mock_aws():
|
||||||
|
s3 = boto3.client("s3", region_name="us-east-1")
|
||||||
|
s3.create_bucket(Bucket="integration-test")
|
||||||
|
|
||||||
|
storage = CloudStorage()
|
||||||
|
assert storage.is_enabled()
|
||||||
|
|
||||||
|
file_path = _create_temp_file(tmp_path, "integration.png", b"test_data")
|
||||||
|
|
||||||
|
url = asyncio.run(storage.upload_and_cleanup(file_path))
|
||||||
|
|
||||||
|
assert url is not None
|
||||||
|
assert "integration-test" in url
|
||||||
|
assert "integration.png" in url
|
||||||
|
assert not os.path.exists(file_path)
|
||||||
|
|
||||||
|
obj = s3.get_object(Bucket="integration-test", Key="integration.png")
|
||||||
|
assert obj["Body"].read() == b"test_data"
|
||||||
|
|
||||||
|
for key in [
|
||||||
|
"SGLANG_CLOUD_STORAGE_TYPE",
|
||||||
|
"SGLANG_S3_BUCKET_NAME",
|
||||||
|
"SGLANG_S3_REGION_NAME",
|
||||||
|
]:
|
||||||
|
os.environ.pop(key, None)
|
||||||
Reference in New Issue
Block a user