Files
mlx-knife/tests_2.0/test_server_api.py.disabled
T
The BROKE Cluster Team 57bf6d86be 2.0.0-beta.3: Feature Complete - Full 1.1.1 Parity Achieved
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.
2025-09-14 18:04:18 +02:00

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