fix(api): prevent upload filename collisions, reject unsupported image URLs

Three fixes in the API upload handling:

1. Multipart uploads now prefix filenames with a UUID to prevent
   overwrites when two requests upload files with the same name.
2. JSON image_url content blocks with remote HTTPS URLs now return
   a 400 error instead of silently dropping the image.
3. Model validation runs for both JSON and multipart requests,
   fixing an inconsistency where multipart bypassed the check.
This commit is contained in:
Mohamed Elkholy 2026-04-15 13:58:48 -04:00 committed by Xubin Ren
parent e1fdca7d40
commit 54b48a7431

View File

@ -29,6 +29,7 @@ _DATA_URL_RE = re.compile(r"^data:([^;]+);base64,(.+)$", re.DOTALL)
class _FileSizeExceeded(Exception): class _FileSizeExceeded(Exception):
"""Raised when an uploaded file exceeds the size limit.""" """Raised when an uploaded file exceeds the size limit."""
API_SESSION_KEY = "api:default" API_SESSION_KEY = "api:default"
API_CHAT_ID = "default" API_CHAT_ID = "default"
@ -37,6 +38,7 @@ API_CHAT_ID = "default"
# Response helpers # Response helpers
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
def _error_json(status: int, message: str, err_type: str = "invalid_request_error") -> web.Response: def _error_json(status: int, message: str, err_type: str = "invalid_request_error") -> web.Response:
return web.json_response( return web.json_response(
{"error": {"message": message, "type": err_type, "code": status}}, {"error": {"message": message, "type": err_type, "code": status}},
@ -74,6 +76,7 @@ def _response_text(value: Any) -> str:
# Upload helpers # Upload helpers
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
def _save_base64_data_url(data_url: str, media_dir: Path) -> str | None: def _save_base64_data_url(data_url: str, media_dir: Path) -> str | None:
"""Decode a data:...;base64,... URL and save to disk.""" """Decode a data:...;base64,... URL and save to disk."""
m = _DATA_URL_RE.match(data_url) m = _DATA_URL_RE.match(data_url)
@ -85,9 +88,7 @@ def _save_base64_data_url(data_url: str, media_dir: Path) -> str | None:
except Exception: except Exception:
return None return None
if len(raw) > MAX_FILE_SIZE: if len(raw) > MAX_FILE_SIZE:
raise _FileSizeExceeded( raise _FileSizeExceeded(f"File exceeds {MAX_FILE_SIZE // (1024 * 1024)}MB limit")
f"File exceeds {MAX_FILE_SIZE // (1024 * 1024)}MB limit"
)
ext = mimetypes.guess_extension(mime_type) or ".bin" ext = mimetypes.guess_extension(mime_type) or ".bin"
filename = f"{uuid.uuid4().hex[:12]}{ext}" filename = f"{uuid.uuid4().hex[:12]}{ext}"
dest = media_dir / safe_filename(filename) dest = media_dir / safe_filename(filename)
@ -121,6 +122,11 @@ def _parse_json_content(body: dict) -> tuple[str, list[str]]:
saved = _save_base64_data_url(url, media_dir) saved = _save_base64_data_url(url, media_dir)
if saved: if saved:
media_paths.append(saved) media_paths.append(saved)
elif url:
raise ValueError(
"Remote image URLs are not supported. "
"Use base64 data URLs or upload files via multipart/form-data."
)
text = " ".join(text_parts) text = " ".join(text_parts)
elif isinstance(user_content, str): elif isinstance(user_content, str):
text = user_content text = user_content
@ -130,12 +136,13 @@ def _parse_json_content(body: dict) -> tuple[str, list[str]]:
return text, media_paths return text, media_paths
async def _parse_multipart(request: web.Request) -> tuple[str, list[str], str | None]: async def _parse_multipart(request: web.Request) -> tuple[str, list[str], str | None, str | None]:
"""Parse multipart/form-data. Returns (text, media_paths, session_id).""" """Parse multipart/form-data. Returns (text, media_paths, session_id, model)."""
media_dir = get_media_dir("api") media_dir = get_media_dir("api")
reader = await request.multipart() reader = await request.multipart()
text = "" text = ""
session_id = None session_id = None
model = None
media_paths: list[str] = [] media_paths: list[str] = []
while True: while True:
@ -146,11 +153,16 @@ async def _parse_multipart(request: web.Request) -> tuple[str, list[str], str |
text = (await part.read()).decode("utf-8") text = (await part.read()).decode("utf-8")
elif part.name == "session_id": elif part.name == "session_id":
session_id = (await part.read()).decode("utf-8").strip() session_id = (await part.read()).decode("utf-8").strip()
elif part.name == "model":
model = (await part.read()).decode("utf-8").strip()
elif part.name == "files": elif part.name == "files":
raw = await part.read() raw = await part.read()
if len(raw) > MAX_FILE_SIZE: if len(raw) > MAX_FILE_SIZE:
raise _FileSizeExceeded(f"File '{part.filename}' exceeds {MAX_FILE_SIZE // (1024*1024)}MB limit") raise _FileSizeExceeded(
filename = safe_filename(part.filename or f"{uuid.uuid4().hex[:12]}.bin") f"File '{part.filename}' exceeds {MAX_FILE_SIZE // (1024 * 1024)}MB limit"
)
base = safe_filename(part.filename or "upload.bin")
filename = f"{uuid.uuid4().hex[:12]}_{base}"
dest = media_dir / filename dest = media_dir / filename
dest.write_bytes(raw) dest.write_bytes(raw)
media_paths.append(str(dest)) media_paths.append(str(dest))
@ -158,13 +170,14 @@ async def _parse_multipart(request: web.Request) -> tuple[str, list[str], str |
if not text: if not text:
text = "请分析上传的文件" text = "请分析上传的文件"
return text, media_paths, session_id return text, media_paths, session_id, model
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Route handlers # Route handlers
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
async def handle_chat_completions(request: web.Request) -> web.Response: async def handle_chat_completions(request: web.Request) -> web.Response:
"""POST /v1/chat/completions — supports JSON and multipart/form-data.""" """POST /v1/chat/completions — supports JSON and multipart/form-data."""
content_type = request.content_type or "" content_type = request.content_type or ""
@ -177,16 +190,17 @@ async def handle_chat_completions(request: web.Request) -> web.Response:
try: try:
if content_type.startswith("multipart/"): if content_type.startswith("multipart/"):
text, media_paths, session_id = await _parse_multipart(request) text, media_paths, session_id, requested_model = await _parse_multipart(request)
else: else:
try: try:
body = await request.json() body = await request.json()
except Exception: except Exception:
return _error_json(400, "Invalid JSON body") return _error_json(400, "Invalid JSON body")
if body.get("stream", False): if body.get("stream", False):
return _error_json(400, "stream=true is not supported yet. Set stream=false or omit it.") return _error_json(
if (requested_model := body.get("model")) and requested_model != model_name: 400, "stream=true is not supported yet. Set stream=false or omit it."
return _error_json(400, f"Only configured model '{model_name}' is available") )
requested_model = body.get("model")
text, media_paths = _parse_json_content(body) text, media_paths = _parse_json_content(body)
session_id = body.get("session_id") session_id = body.get("session_id")
except ValueError as e: except ValueError as e:
@ -197,11 +211,16 @@ async def handle_chat_completions(request: web.Request) -> web.Response:
logger.exception("Error parsing upload") logger.exception("Error parsing upload")
return _error_json(413, "File too large or invalid upload") return _error_json(413, "File too large or invalid upload")
if requested_model and requested_model != model_name:
return _error_json(400, f"Only configured model '{model_name}' is available")
session_key = f"api:{session_id}" if session_id else API_SESSION_KEY session_key = f"api:{session_id}" if session_id else API_SESSION_KEY
session_locks: dict[str, asyncio.Lock] = request.app["session_locks"] session_locks: dict[str, asyncio.Lock] = request.app["session_locks"]
session_lock = session_locks.setdefault(session_key, asyncio.Lock()) session_lock = session_locks.setdefault(session_key, asyncio.Lock())
logger.info("API request session_key={} media={} text={}", session_key, len(media_paths), text[:80]) logger.info(
"API request session_key={} media={} text={}", session_key, len(media_paths), text[:80]
)
_FALLBACK = EMPTY_FINAL_RESPONSE_MESSAGE _FALLBACK = EMPTY_FINAL_RESPONSE_MESSAGE
@ -252,17 +271,19 @@ async def handle_chat_completions(request: web.Request) -> web.Response:
async def handle_models(request: web.Request) -> web.Response: async def handle_models(request: web.Request) -> web.Response:
"""GET /v1/models""" """GET /v1/models"""
model_name = request.app.get("model_name", "nanobot") model_name = request.app.get("model_name", "nanobot")
return web.json_response({ return web.json_response(
"object": "list", {
"data": [ "object": "list",
{ "data": [
"id": model_name, {
"object": "model", "id": model_name,
"created": 0, "object": "model",
"owned_by": "nanobot", "created": 0,
} "owned_by": "nanobot",
], }
}) ],
}
)
async def handle_health(request: web.Request) -> web.Response: async def handle_health(request: web.Request) -> web.Response:
@ -274,7 +295,10 @@ async def handle_health(request: web.Request) -> web.Response:
# App factory # App factory
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
def create_app(agent_loop, model_name: str = "nanobot", request_timeout: float = 120.0) -> web.Application:
def create_app(
agent_loop, model_name: str = "nanobot", request_timeout: float = 120.0
) -> web.Application:
"""Create the aiohttp application. """Create the aiohttp application.
Args: Args: