openai_server: fix path traversal in web static serving (relative_to, not startswith) + empty-prompt validation, with tests (#212)
This commit is contained in:
+8
-2
@@ -731,12 +731,16 @@ class APIHandler(BaseHTTPRequestHandler):
|
|||||||
Read-only, no auth (same trust level as /health), traversal-safe."""
|
Read-only, no auth (same trust level as /health), traversal-safe."""
|
||||||
if path.startswith("/v1/") or path == "/health":
|
if path.startswith("/v1/") or path == "/health":
|
||||||
return False
|
return False
|
||||||
base = self.WEB_DIST
|
base = self.WEB_DIST.resolve()
|
||||||
if not base.is_dir():
|
if not base.is_dir():
|
||||||
return False
|
return False
|
||||||
rel = unquote(path).lstrip("/") or "index.html"
|
rel = unquote(path).lstrip("/") or "index.html"
|
||||||
target = (base / rel).resolve()
|
target = (base / rel).resolve()
|
||||||
if not str(target).startswith(str(base)) or not target.is_file():
|
try:
|
||||||
|
target.relative_to(base)
|
||||||
|
except ValueError:
|
||||||
|
target = None
|
||||||
|
if target is None or not target.is_file():
|
||||||
if path == "/" or "." not in rel: # SPA fallback
|
if path == "/" or "." not in rel: # SPA fallback
|
||||||
target = base / "index.html"
|
target = base / "index.html"
|
||||||
if not target.is_file():
|
if not target.is_file():
|
||||||
@@ -1045,6 +1049,8 @@ class APIHandler(BaseHTTPRequestHandler):
|
|||||||
prompt = body.get("prompt")
|
prompt = body.get("prompt")
|
||||||
if not isinstance(prompt, str):
|
if not isinstance(prompt, str):
|
||||||
raise APIError(400, "Colibri currently requires `prompt` to be a string.", "prompt")
|
raise APIError(400, "Colibri currently requires `prompt` to be a string.", "prompt")
|
||||||
|
if not prompt:
|
||||||
|
raise APIError(400, "`prompt` must not be empty.", "prompt")
|
||||||
self.generation(body, prompt, request_id, False)
|
self.generation(body, prompt, request_id, False)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -2,13 +2,15 @@ import io
|
|||||||
import json
|
import json
|
||||||
import math
|
import math
|
||||||
import socket
|
import socket
|
||||||
|
import tempfile
|
||||||
import threading
|
import threading
|
||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
from urllib.error import HTTPError
|
from urllib.error import HTTPError
|
||||||
from urllib.request import Request, urlopen
|
from urllib.request import Request, urlopen
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
from openai_server import (APIError, APIServer, ClientCancelled, END, GenerationScheduler,
|
from openai_server import (APIError, APIHandler, APIServer, ClientCancelled, END, GenerationScheduler,
|
||||||
READY, Engine, generation_options, parse_tool_calls,
|
READY, Engine, generation_options, parse_tool_calls,
|
||||||
read_engine_turn, render_chat, serve)
|
read_engine_turn, render_chat, serve)
|
||||||
|
|
||||||
@@ -490,6 +492,12 @@ class HTTPTest(unittest.TestCase):
|
|||||||
self.assertEqual(body["choices"][0]["text"], "Héllo")
|
self.assertEqual(body["choices"][0]["text"], "Héllo")
|
||||||
self.assertEqual(self.engine.calls[-1][0], "Complete me")
|
self.assertEqual(self.engine.calls[-1][0], "Complete me")
|
||||||
|
|
||||||
|
def test_rejects_empty_legacy_completion(self):
|
||||||
|
with self.assertRaises(HTTPError) as caught:
|
||||||
|
self.request("/v1/completions", {"model": "test-model", "prompt": ""})
|
||||||
|
self.assertEqual(caught.exception.code, 400)
|
||||||
|
self.assertEqual(json.load(caught.exception)["error"]["param"], "prompt")
|
||||||
|
|
||||||
def test_rejects_invalid_stream_options(self):
|
def test_rejects_invalid_stream_options(self):
|
||||||
with self.assertRaises(HTTPError) as caught:
|
with self.assertRaises(HTTPError) as caught:
|
||||||
self.request("/v1/chat/completions", {
|
self.request("/v1/chat/completions", {
|
||||||
@@ -499,6 +507,39 @@ class HTTPTest(unittest.TestCase):
|
|||||||
self.assertEqual(caught.exception.code, 400)
|
self.assertEqual(caught.exception.code, 400)
|
||||||
|
|
||||||
|
|
||||||
|
class StaticServingTest(unittest.TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.tmp = tempfile.TemporaryDirectory()
|
||||||
|
root = Path(self.tmp.name)
|
||||||
|
dist = root / "dist"
|
||||||
|
dist.mkdir()
|
||||||
|
(dist / "index.html").write_text("dashboard", encoding="utf-8")
|
||||||
|
sibling = root / "dist-private"
|
||||||
|
sibling.mkdir()
|
||||||
|
(sibling / "secret.txt").write_text("private", encoding="utf-8")
|
||||||
|
self.web_dist = patch.object(APIHandler, "WEB_DIST", dist)
|
||||||
|
self.web_dist.start()
|
||||||
|
self.server = APIServer(("127.0.0.1", 0), FakeEngine(), "test-model")
|
||||||
|
self.thread = threading.Thread(target=self.server.serve_forever, daemon=True)
|
||||||
|
self.thread.start()
|
||||||
|
self.base = f"http://127.0.0.1:{self.server.server_port}"
|
||||||
|
|
||||||
|
def tearDown(self):
|
||||||
|
self.server.scheduler.close()
|
||||||
|
self.server.shutdown()
|
||||||
|
self.server.server_close()
|
||||||
|
self.thread.join(timeout=2)
|
||||||
|
self.web_dist.stop()
|
||||||
|
self.tmp.cleanup()
|
||||||
|
|
||||||
|
def test_static_root_stays_inside_dist_directory(self):
|
||||||
|
with urlopen(self.base + "/", timeout=2) as response:
|
||||||
|
self.assertEqual(response.read(), b"dashboard")
|
||||||
|
with self.assertRaises(HTTPError) as caught:
|
||||||
|
urlopen(self.base + "/%2e%2e/dist-private/secret.txt", timeout=2)
|
||||||
|
self.assertEqual(caught.exception.code, 404)
|
||||||
|
|
||||||
|
|
||||||
class SchedulerHTTPTest(unittest.TestCase):
|
class SchedulerHTTPTest(unittest.TestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
self.engine = BlockingEngine()
|
self.engine = BlockingEngine()
|
||||||
|
|||||||
Reference in New Issue
Block a user