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:
committed by
Xubin Ren
parent
e1fdca7d40
commit
54b48a7431
+49
-25
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user