Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 38 additions & 0 deletions mlx_lm/server.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,11 @@
# Copyright © 2023-2024 Apple Inc.

import argparse
import hmac
import ipaddress
import json
import logging
import os
import pickle
import platform
import socket
Expand Down Expand Up @@ -1043,6 +1046,19 @@ def _set_stream_headers(self, status_code: int = 200):
self.send_header("Cache-Control", "no-cache")
self._set_cors_headers()

def _authenticate(self) -> bool:
api_key = getattr(self.response_generator.cli_args, "api_key", None)
if not api_key:
return True
authorization = self.headers.get("Authorization", "")
if hmac.compare_digest(authorization, f"Bearer {api_key}"):
return True
self._set_completion_headers(401)
self.send_header("WWW-Authenticate", "Bearer")
self.end_headers()
self.wfile.write(json.dumps({"error": "Unauthorized"}).encode())
return False

def do_OPTIONS(self):
self._set_completion_headers(204)
self.end_headers()
Expand All @@ -1051,6 +1067,9 @@ def do_POST(self):
"""
Respond to a POST request from a client.
"""
if not self._authenticate():
return

request_factories = {
"/v1/completions": self.handle_text_completions,
"/v1/chat/completions": self.handle_chat_completions,
Expand Down Expand Up @@ -1587,6 +1606,9 @@ def do_GET(self):
"""
Respond to a GET request from a client.
"""
if not self._authenticate():
return

if self.path.startswith("/v1/models"):
self.handle_models_request()
elif self.path == "/health":
Expand Down Expand Up @@ -1672,6 +1694,17 @@ def _run_http_server(
server_class=ThreadingHTTPServer,
handler_class=APIHandler,
):
api_key = getattr(response_generator.cli_args, "api_key", None)
try:
is_loopback = ipaddress.ip_address(host).is_loopback
except ValueError:
is_loopback = host.lower() == "localhost"
if not is_loopback and not api_key:
raise ValueError(
"Binding mlx_lm.server to a non-loopback host requires --api-key "
"or MLX_LM_API_KEY"
)

server_address = (host, port)
infos = socket.getaddrinfo(
*server_address, type=socket.SOCK_STREAM, flags=socket.AI_PASSIVE
Expand Down Expand Up @@ -1744,6 +1777,11 @@ def main():
default="*",
help="Allowed origins (default: *)",
)
parser.add_argument(
"--api-key",
default=os.environ.get("MLX_LM_API_KEY"),
help="Bearer token required for HTTP requests (or set MLX_LM_API_KEY)",
)
parser.add_argument(
"--draft-model",
type=str,
Expand Down
28 changes: 28 additions & 0 deletions tests/test_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
ResponseGenerator,
SamplingArguments,
_make_sampler,
_run_http_server,
)
from mlx_lm.utils import load

Expand Down Expand Up @@ -248,6 +249,21 @@ def test_handle_completions(self):
json.loads(requests.post(url, json=post_data).text)["choices"][0]["text"],
)

def test_api_key_protects_model_selection_endpoints(self):
url = f"http://localhost:{self.port}/v1/models"
cli_args = self.response_generator.cli_args
cli_args.api_key = "test-secret"
try:
unauthorized = requests.get(url)
authorized = requests.get(
url, headers={"Authorization": "Bearer test-secret"}
)
finally:
cli_args.api_key = None

self.assertEqual(unauthorized.status_code, 401)
self.assertEqual(authorized.status_code, 200)

def test_handle_chat_completions(self):
url = f"http://localhost:{self.port}/v1/chat/completions"
chat_post_data = {
Expand Down Expand Up @@ -369,6 +385,18 @@ def test_handle_models(self):
self.assertIn("created", model)


class TestRemoteBindingSecurity(unittest.TestCase):
def test_remote_binding_requires_api_key(self):
response_generator = type(
"ResponseGenerator",
(),
{"cli_args": type("Args", (), {"api_key": None})()},
)()

with self.assertRaisesRegex(ValueError, "non-loopback host requires"):
_run_http_server("0.0.0.0", 8080, response_generator)


class TestServerWithDraftModel(unittest.TestCase):
@classmethod
def setUpClass(cls):
Expand Down