mirror of
https://github.com/cloudstack-llc/mlx-knife.git
synced 2026-07-21 01:55:25 -04:00
57bf6d86be
Major Features Added: • Complete run command implementation with interactive/single-shot modes • MLXRunner core engine ported from 1.x with modular architecture • OpenAI-compatible server with SIGINT-robust supervisor mode • Experimental push feature properly isolated behind environment variable Key Improvements: - Full feature parity with 1.1.1 stable releases - Enhanced human output formatting across all commands - Clean separation of stable (184 tests) vs experimental features - Updated demo GIF showcasing improved 2.0 interface Fixes: - Pull operation cache pollution (Issue #30) with preflight access checks - Test stability improvements across all environments Architecture: - Modular runner design with focused helper modules - Thread-safe model loading and memory management - stable testing across Python 3.9-3.13 Ready for use as comprehensive 1.x alternative.
264 lines
8.6 KiB
Plaintext
264 lines
8.6 KiB
Plaintext
"""
|
|
Test server API endpoints for 2.0 implementation.
|
|
"""
|
|
|
|
import json
|
|
import pytest
|
|
from fastapi.testclient import TestClient
|
|
from unittest.mock import Mock, patch
|
|
|
|
from mlxk2.core.server_base import app
|
|
from mlxk2.core.runner import MLXRunner
|
|
|
|
|
|
class MockMLXRunner:
|
|
"""Mock MLXRunner for testing."""
|
|
|
|
def __init__(self, model_path, verbose=False):
|
|
self.model_spec = model_path
|
|
self.verbose = verbose
|
|
self._context_length = 4096
|
|
|
|
def load_model(self):
|
|
pass
|
|
|
|
def cleanup(self):
|
|
pass
|
|
|
|
def _calculate_dynamic_max_tokens(self, server_mode=False):
|
|
if server_mode:
|
|
return self._context_length // 2 # Half context for server
|
|
else:
|
|
return self._context_length # Full context for run
|
|
|
|
def generate_streaming(self, prompt, max_tokens=None, temperature=0.7,
|
|
top_p=0.9, repetition_penalty=1.1, use_chat_template=True,
|
|
use_chat_stop_tokens=False):
|
|
"""Mock streaming generation."""
|
|
yield "Hello"
|
|
yield " "
|
|
yield "world"
|
|
yield "!"
|
|
|
|
def generate_batch(self, prompt, max_tokens=None, temperature=0.7,
|
|
top_p=0.9, repetition_penalty=1.1, use_chat_template=True,
|
|
use_chat_stop_tokens=False):
|
|
"""Mock batch generation."""
|
|
return "Hello world!"
|
|
|
|
def _format_conversation(self, messages):
|
|
"""Mock conversation formatting."""
|
|
formatted_parts = []
|
|
for msg in messages:
|
|
role = msg["role"]
|
|
content = msg["content"]
|
|
if role == "system":
|
|
formatted_parts.append(f"System: {content}")
|
|
elif role == "user":
|
|
formatted_parts.append(f"Human: {content}")
|
|
elif role == "assistant":
|
|
formatted_parts.append(f"Assistant: {content}")
|
|
|
|
return "\n\n".join(formatted_parts) + "\n\nAssistant: "
|
|
|
|
|
|
@pytest.fixture
|
|
def client():
|
|
"""Create test client."""
|
|
with TestClient(app) as client:
|
|
yield client
|
|
|
|
|
|
@pytest.fixture
|
|
def mock_runner():
|
|
"""Create mock runner."""
|
|
return MockMLXRunner("test-model")
|
|
|
|
|
|
def test_health_endpoint(client):
|
|
"""Test health check endpoint."""
|
|
response = client.get("/health")
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert data["status"] == "healthy"
|
|
assert data["service"] == "mlx-knife-server-2.0"
|
|
|
|
|
|
def test_models_endpoint(client):
|
|
"""Test models listing endpoint."""
|
|
# Mock the model cache and health check
|
|
with patch('mlxk2.core.server_base.get_current_model_cache') as mock_cache, \
|
|
patch('mlxk2.core.server_base.cache_dir_to_hf') as mock_cache_to_hf, \
|
|
patch('mlxk2.core.server_base.detect_framework') as mock_framework, \
|
|
patch('mlxk2.core.server_base.is_model_healthy') as mock_healthy:
|
|
|
|
# Setup mocks
|
|
mock_cache_dir = Mock()
|
|
mock_cache_dir.name = "models--test--model"
|
|
mock_cache_dir.iterdir.return_value = [mock_cache_dir]
|
|
|
|
mock_cache.return_value.iterdir.return_value = [mock_cache_dir]
|
|
mock_cache_to_hf.return_value = "test/model"
|
|
mock_framework.return_value = "MLX"
|
|
mock_healthy.return_value = (True, None)
|
|
|
|
# Mock snapshots directory
|
|
mock_snapshots_dir = Mock()
|
|
mock_snapshots_dir.exists.return_value = True
|
|
mock_snapshot = Mock()
|
|
mock_snapshot.is_dir.return_value = True
|
|
mock_snapshots_dir.iterdir.return_value = [mock_snapshot]
|
|
mock_cache_dir.__truediv__.return_value = mock_snapshots_dir
|
|
|
|
response = client.get("/v1/models")
|
|
assert response.status_code == 200
|
|
data = response.json()
|
|
assert "data" in data
|
|
assert data["object"] == "list"
|
|
|
|
|
|
@patch('mlxk2.core.server_base.get_or_load_model')
|
|
def test_completions_endpoint(mock_get_model, client, mock_runner):
|
|
"""Test completions endpoint."""
|
|
mock_get_model.return_value = mock_runner
|
|
|
|
request_data = {
|
|
"model": "test/model",
|
|
"prompt": "Hello",
|
|
"max_tokens": 10,
|
|
"temperature": 0.7
|
|
}
|
|
|
|
response = client.post("/v1/completions", json=request_data)
|
|
assert response.status_code == 200
|
|
|
|
data = response.json()
|
|
assert data["object"] == "text_completion"
|
|
assert "choices" in data
|
|
assert len(data["choices"]) == 1
|
|
assert data["choices"][0]["text"] == "Hello world!"
|
|
|
|
|
|
@patch('mlxk2.core.server_base.get_or_load_model')
|
|
def test_chat_completions_endpoint(mock_get_model, client, mock_runner):
|
|
"""Test chat completions endpoint."""
|
|
mock_get_model.return_value = mock_runner
|
|
|
|
request_data = {
|
|
"model": "test/model",
|
|
"messages": [
|
|
{"role": "user", "content": "Hello"}
|
|
],
|
|
"max_tokens": 10,
|
|
"temperature": 0.7
|
|
}
|
|
|
|
response = client.post("/v1/chat/completions", json=request_data)
|
|
assert response.status_code == 200
|
|
|
|
data = response.json()
|
|
assert data["object"] == "chat.completion"
|
|
assert "choices" in data
|
|
assert len(data["choices"]) == 1
|
|
assert data["choices"][0]["message"]["role"] == "assistant"
|
|
assert data["choices"][0]["message"]["content"] == "Hello world!"
|
|
|
|
|
|
@patch('mlxk2.core.server_base.get_or_load_model')
|
|
def test_streaming_completions(mock_get_model, client, mock_runner):
|
|
"""Test streaming completions."""
|
|
mock_get_model.return_value = mock_runner
|
|
|
|
request_data = {
|
|
"model": "test/model",
|
|
"prompt": "Hello",
|
|
"stream": True,
|
|
"max_tokens": 10
|
|
}
|
|
|
|
response = client.post("/v1/completions", json=request_data)
|
|
assert response.status_code == 200
|
|
assert response.headers["content-type"] == "text/plain; charset=utf-8"
|
|
|
|
|
|
@patch('mlxk2.core.server_base.get_or_load_model')
|
|
def test_streaming_chat_completions(mock_get_model, client, mock_runner):
|
|
"""Test streaming chat completions."""
|
|
mock_get_model.return_value = mock_runner
|
|
|
|
request_data = {
|
|
"model": "test/model",
|
|
"messages": [
|
|
{"role": "user", "content": "Hello"}
|
|
],
|
|
"stream": True,
|
|
"max_tokens": 10
|
|
}
|
|
|
|
response = client.post("/v1/chat/completions", json=request_data)
|
|
assert response.status_code == 200
|
|
assert response.headers["content-type"] == "text/plain; charset=utf-8"
|
|
|
|
|
|
def test_model_hot_swapping(client):
|
|
"""Test that model hot-swapping clears previous models."""
|
|
with patch('mlxk2.core.server_base.resolve_model_for_operation') as mock_resolve, \
|
|
patch('mlxk2.core.server_base.get_current_model_cache') as mock_cache, \
|
|
patch('mlxk2.core.server_base.MLXRunner') as mock_runner_class:
|
|
|
|
# Setup for first model
|
|
mock_resolve.return_value = ("test/model1", None, None)
|
|
mock_cache_dir = Mock()
|
|
mock_cache_dir.__truediv__.return_value.exists.return_value = True
|
|
mock_cache.return_value = mock_cache_dir
|
|
|
|
mock_runner1 = Mock()
|
|
mock_runner1.load_model = Mock()
|
|
mock_runner1.cleanup = Mock()
|
|
mock_runner_class.return_value = mock_runner1
|
|
|
|
# Load first model
|
|
from mlxk2.core.server_base import get_or_load_model
|
|
runner1 = get_or_load_model("test/model1")
|
|
|
|
# Setup for second model
|
|
mock_resolve.return_value = ("test/model2", None, None)
|
|
mock_runner2 = Mock()
|
|
mock_runner2.load_model = Mock()
|
|
mock_runner2.cleanup = Mock()
|
|
mock_runner_class.return_value = mock_runner2
|
|
|
|
# Load second model - should cleanup first
|
|
runner2 = get_or_load_model("test/model2")
|
|
|
|
# Verify cleanup was called on first runner
|
|
mock_runner1.cleanup.assert_called_once()
|
|
|
|
|
|
def test_server_mode_token_limits():
|
|
"""Test that server mode uses half context for DoS protection."""
|
|
runner = MockMLXRunner("test-model")
|
|
|
|
# Server mode should use half context
|
|
server_tokens = runner._calculate_dynamic_max_tokens(server_mode=True)
|
|
assert server_tokens == 2048 # Half of 4096
|
|
|
|
# Run mode should use full context
|
|
run_tokens = runner._calculate_dynamic_max_tokens(server_mode=False)
|
|
assert run_tokens == 4096 # Full context
|
|
|
|
|
|
@patch('mlxk2.core.server_base.get_or_load_model')
|
|
def test_error_handling(mock_get_model, client):
|
|
"""Test error handling in API endpoints."""
|
|
# Test model not found
|
|
mock_get_model.side_effect = Exception("Model not found")
|
|
|
|
request_data = {
|
|
"model": "nonexistent/model",
|
|
"prompt": "Hello"
|
|
}
|
|
|
|
response = client.post("/v1/completions", json=request_data)
|
|
assert response.status_code == 500
|