From 8c18b7af6897bd842727f06a8bbcc0622be93e03 Mon Sep 17 00:00:00 2001 From: Kushida Date: Tue, 14 Jul 2026 23:07:41 +0300 Subject: [PATCH] openai_server: fix path traversal in web static serving (relative_to, not startswith) + empty-prompt validation, with tests (#212) --- c/openai_server.py | 10 ++++++-- c/tests/test_openai_server.py | 43 ++++++++++++++++++++++++++++++++++- 2 files changed, 50 insertions(+), 3 deletions(-) diff --git a/c/openai_server.py b/c/openai_server.py index c88888d..f31e8f7 100644 --- a/c/openai_server.py +++ b/c/openai_server.py @@ -731,12 +731,16 @@ class APIHandler(BaseHTTPRequestHandler): Read-only, no auth (same trust level as /health), traversal-safe.""" if path.startswith("/v1/") or path == "/health": return False - base = self.WEB_DIST + base = self.WEB_DIST.resolve() if not base.is_dir(): return False rel = unquote(path).lstrip("/") or "index.html" 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 target = base / "index.html" if not target.is_file(): @@ -1045,6 +1049,8 @@ class APIHandler(BaseHTTPRequestHandler): prompt = body.get("prompt") if not isinstance(prompt, str): 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) diff --git a/c/tests/test_openai_server.py b/c/tests/test_openai_server.py index fd184d6..0064070 100644 --- a/c/tests/test_openai_server.py +++ b/c/tests/test_openai_server.py @@ -2,13 +2,15 @@ import io import json import math import socket +import tempfile import threading import unittest from unittest.mock import patch from urllib.error import HTTPError 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, read_engine_turn, render_chat, serve) @@ -490,6 +492,12 @@ class HTTPTest(unittest.TestCase): self.assertEqual(body["choices"][0]["text"], "Héllo") 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): with self.assertRaises(HTTPError) as caught: self.request("/v1/chat/completions", { @@ -499,6 +507,39 @@ class HTTPTest(unittest.TestCase): 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): def setUp(self): self.engine = BlockingEngine()