chore: 解耦 Executor/Scheduler/TaskManager,修复 stop 页泄漏,移除 ServerState 全局单例
This commit is contained in:
@@ -11,6 +11,14 @@ from astrai.inference.server import app
|
||||
@pytest.fixture
|
||||
def client():
|
||||
"""Provide a test client for the FastAPI app."""
|
||||
app.state.server_config = {
|
||||
"device": "cpu",
|
||||
"dtype": "bfloat16",
|
||||
"param_path": None,
|
||||
"max_batch_size": 1,
|
||||
"_test": True,
|
||||
}
|
||||
app.state.engine = None
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
@@ -39,7 +47,7 @@ def mock_engine():
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def loaded_model(mock_engine, monkeypatch):
|
||||
def loaded_model(client, mock_engine):
|
||||
"""Simulate that the engine is loaded."""
|
||||
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
|
||||
app.state.engine = mock_engine
|
||||
return mock_engine
|
||||
|
||||
@@ -162,23 +162,20 @@ def test_prefix_cache_has_page():
|
||||
|
||||
|
||||
def test_task_table_set_get():
|
||||
pool = PagePool(n_pages=8)
|
||||
table = TaskTable(pool, page_size=64)
|
||||
table = TaskTable(page_size=64)
|
||||
table.set("task1", [0, 1, 2], 128)
|
||||
assert table.get("task1") == [0, 1, 2]
|
||||
assert table.get_cached("task1") == 128
|
||||
|
||||
|
||||
def test_task_table_get_missing():
|
||||
pool = PagePool(n_pages=8)
|
||||
table = TaskTable(pool, page_size=64)
|
||||
table = TaskTable(page_size=64)
|
||||
assert table.get("nonexistent") == []
|
||||
assert table.get_cached("nonexistent") == 0
|
||||
|
||||
|
||||
def test_task_table_pop():
|
||||
pool = PagePool(n_pages=8)
|
||||
table = TaskTable(pool, page_size=64)
|
||||
table = TaskTable(page_size=64)
|
||||
table.set("task1", [0, 1], 64)
|
||||
pages, cached = table.pop("task1")
|
||||
assert pages == [0, 1]
|
||||
@@ -186,26 +183,39 @@ def test_task_table_pop():
|
||||
assert table.get("task1") == []
|
||||
|
||||
|
||||
def test_task_table_extend_allocates_pages():
|
||||
pool = PagePool(n_pages=8)
|
||||
table = TaskTable(pool, page_size=64)
|
||||
table.set("task1", [], 0)
|
||||
ok = table.extend("task1", 200)
|
||||
def test_paged_cache_task_extend_allocates():
|
||||
cache = PagedCache(
|
||||
n_layers=1,
|
||||
n_pages=8,
|
||||
page_size=64,
|
||||
n_kv_heads=2,
|
||||
head_dim=8,
|
||||
device=torch.device("cpu"),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
cache._table.set("task1", [], 0)
|
||||
ok = cache.task_extend("task1", 200)
|
||||
assert ok
|
||||
assert len(table.get("task1")) == 4
|
||||
assert len(cache._table.get("task1")) == 4
|
||||
|
||||
|
||||
def test_task_table_extend_fails_when_pool_full():
|
||||
pool = PagePool(n_pages=2)
|
||||
table = TaskTable(pool, page_size=64)
|
||||
table.set("task1", [pool.alloc(), pool.alloc()], 0)
|
||||
ok = table.extend("task1", 300)
|
||||
def test_paged_cache_task_extend_fails_when_pool_full():
|
||||
cache = PagedCache(
|
||||
n_layers=1,
|
||||
n_pages=2,
|
||||
page_size=64,
|
||||
n_kv_heads=2,
|
||||
head_dim=8,
|
||||
device=torch.device("cpu"),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
cache._table.set("task1", [0, 1], 0)
|
||||
ok = cache.task_extend("task1", 300)
|
||||
assert not ok
|
||||
|
||||
|
||||
def test_task_table_table_tensor():
|
||||
pool = PagePool(n_pages=16)
|
||||
table = TaskTable(pool, page_size=64)
|
||||
table = TaskTable(page_size=64)
|
||||
table.set("a", [0, 1], 0)
|
||||
table.set("b", [2, 3, 4], 0)
|
||||
t = table.table_tensor(["a", "b"], torch.device("cpu"))
|
||||
@@ -215,8 +225,7 @@ def test_task_table_table_tensor():
|
||||
|
||||
|
||||
def test_task_table_table_tensor_empty_input():
|
||||
pool = PagePool(n_pages=4)
|
||||
table = TaskTable(pool, page_size=64)
|
||||
table = TaskTable(page_size=64)
|
||||
t = table.table_tensor([], torch.device("cpu"))
|
||||
assert t.numel() == 0
|
||||
|
||||
|
||||
@@ -1,20 +1,20 @@
|
||||
"""Unit tests for _Result accumulator and InferenceEngine.generate()."""
|
||||
"""Unit tests for GenerateResult accumulator and InferenceEngine.generate()."""
|
||||
|
||||
import threading
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from astrai.inference.engine import _Result
|
||||
from astrai.inference.engine import GenerateResult
|
||||
from astrai.inference.task import STOP
|
||||
|
||||
|
||||
def test_result_append_single():
|
||||
r = _Result(count=1)
|
||||
r = GenerateResult(count=1)
|
||||
r.append("hello", 0)
|
||||
assert r.results[0] == "hello"
|
||||
|
||||
|
||||
def test_result_append_multiple_tasks():
|
||||
r = _Result(count=3)
|
||||
r = GenerateResult(count=3)
|
||||
r.append("a", 0)
|
||||
r.append("b", 1)
|
||||
r.append("c", 2)
|
||||
@@ -24,7 +24,7 @@ def test_result_append_multiple_tasks():
|
||||
|
||||
|
||||
def test_result_stop_marks_complete():
|
||||
r = _Result(count=2)
|
||||
r = GenerateResult(count=2)
|
||||
r.append("text", 0)
|
||||
r.append(STOP, 0)
|
||||
r.append("more", 1)
|
||||
@@ -34,14 +34,14 @@ def test_result_stop_marks_complete():
|
||||
|
||||
|
||||
def test_result_stop_does_not_double_count():
|
||||
r = _Result(count=1)
|
||||
r = GenerateResult(count=1)
|
||||
r.append(STOP, 0)
|
||||
r.append(STOP, 0)
|
||||
assert r._completed == 1
|
||||
|
||||
|
||||
def test_result_pop_all_returns_and_clears():
|
||||
r = _Result(count=2)
|
||||
r = GenerateResult(count=2)
|
||||
r.append("a", 0)
|
||||
r.append("b", 1)
|
||||
out = r.pop_all()
|
||||
@@ -52,7 +52,7 @@ def test_result_pop_all_returns_and_clears():
|
||||
|
||||
|
||||
def test_result_wait_blocks_until_data():
|
||||
r = _Result(count=1)
|
||||
r = GenerateResult(count=1)
|
||||
|
||||
def delayed_append():
|
||||
import time
|
||||
@@ -69,13 +69,13 @@ def test_result_wait_blocks_until_data():
|
||||
|
||||
|
||||
def test_result_wait_timeout():
|
||||
r = _Result(count=1)
|
||||
r = GenerateResult(count=1)
|
||||
ok = r.wait(timeout=0.01)
|
||||
assert not ok
|
||||
|
||||
|
||||
def test_result_wait_completion_non_streaming():
|
||||
r = _Result(count=2)
|
||||
r = GenerateResult(count=2)
|
||||
|
||||
def finish_later():
|
||||
import time
|
||||
@@ -93,7 +93,7 @@ def test_result_wait_completion_non_streaming():
|
||||
|
||||
|
||||
def test_result_get_results():
|
||||
r = _Result(count=2)
|
||||
r = GenerateResult(count=2)
|
||||
r.append("hello", 0)
|
||||
r.append("world", 1)
|
||||
results = r.get_results()
|
||||
@@ -148,9 +148,9 @@ def test_engine_generate_streaming_yields_tokens():
|
||||
gen = eng.generate("hello", stream=True)
|
||||
|
||||
cb = callbacks_saved[0]
|
||||
cb("t1", 0)
|
||||
cb("t2", 0)
|
||||
cb(STOP, 0)
|
||||
cb("t1")
|
||||
cb("t2")
|
||||
cb(STOP)
|
||||
|
||||
tokens = list(gen)
|
||||
assert tokens == ["t1", "t2"]
|
||||
|
||||
@@ -2,10 +2,12 @@
|
||||
|
||||
import pytest
|
||||
|
||||
from astrai.inference.server import app
|
||||
|
||||
def test_health_no_model(client, monkeypatch):
|
||||
|
||||
def test_health_no_model(client):
|
||||
"""GET /health should return 200 even when engine not loaded."""
|
||||
monkeypatch.setattr("astrai.inference.server._state.engine", None)
|
||||
app.state.engine = None
|
||||
response = client.get("/health")
|
||||
assert response.status_code == 200
|
||||
data = response.json()
|
||||
@@ -22,15 +24,14 @@ def test_health_with_model(client, loaded_model):
|
||||
assert data["model_loaded"] is True
|
||||
|
||||
|
||||
def test_chat_completions_non_stream(client, loaded_model, monkeypatch):
|
||||
def test_chat_completions_non_stream(client, loaded_model):
|
||||
"""POST /v1/chat/completions with stream=false returns OpenAI-style JSON."""
|
||||
|
||||
async def async_gen():
|
||||
yield "Assistant reply"
|
||||
|
||||
mock_engine = loaded_model
|
||||
mock_engine.generate_async.return_value = async_gen()
|
||||
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
|
||||
app.state.engine = loaded_model
|
||||
loaded_model.generate_async.return_value = async_gen()
|
||||
response = client.post(
|
||||
"/v1/chat/completions",
|
||||
json={
|
||||
@@ -48,16 +49,15 @@ def test_chat_completions_non_stream(client, loaded_model, monkeypatch):
|
||||
assert "prompt_tokens" in data["usage"]
|
||||
|
||||
|
||||
def test_chat_completions_stream(client, loaded_model, monkeypatch):
|
||||
def test_chat_completions_stream(client, loaded_model):
|
||||
"""POST /v1/chat/completions with stream=true returns SSE stream."""
|
||||
|
||||
async def async_gen():
|
||||
yield "cumulative1"
|
||||
yield "cumulative2"
|
||||
|
||||
mock_engine = loaded_model
|
||||
mock_engine.generate_async.return_value = async_gen()
|
||||
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
|
||||
app.state.engine = loaded_model
|
||||
loaded_model.generate_async.return_value = async_gen()
|
||||
response = client.post(
|
||||
"/v1/chat/completions",
|
||||
json={
|
||||
@@ -77,15 +77,14 @@ def test_chat_completions_stream(client, loaded_model, monkeypatch):
|
||||
assert any("[DONE]" in line for line in lines)
|
||||
|
||||
|
||||
def test_messages_non_stream(client, loaded_model, monkeypatch):
|
||||
def test_messages_non_stream(client, loaded_model):
|
||||
"""POST /v1/messages with stream=false returns Anthropic-style JSON."""
|
||||
|
||||
async def async_gen():
|
||||
yield "Assistant reply"
|
||||
|
||||
mock_engine = loaded_model
|
||||
mock_engine.generate_async.return_value = async_gen()
|
||||
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
|
||||
app.state.engine = loaded_model
|
||||
loaded_model.generate_async.return_value = async_gen()
|
||||
response = client.post(
|
||||
"/v1/messages",
|
||||
json={
|
||||
@@ -105,16 +104,15 @@ def test_messages_non_stream(client, loaded_model, monkeypatch):
|
||||
assert "input_tokens" in data["usage"]
|
||||
|
||||
|
||||
def test_messages_stream(client, loaded_model, monkeypatch):
|
||||
def test_messages_stream(client, loaded_model):
|
||||
"""POST /v1/messages with stream=true returns Anthropic SSE stream."""
|
||||
|
||||
async def async_gen():
|
||||
yield "cumulative1"
|
||||
yield "cumulative2"
|
||||
|
||||
mock_engine = loaded_model
|
||||
mock_engine.generate_async.return_value = async_gen()
|
||||
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
|
||||
app.state.engine = loaded_model
|
||||
loaded_model.generate_async.return_value = async_gen()
|
||||
response = client.post(
|
||||
"/v1/messages",
|
||||
json={
|
||||
@@ -137,15 +135,14 @@ def test_messages_stream(client, loaded_model, monkeypatch):
|
||||
assert "message_stop" in content
|
||||
|
||||
|
||||
def test_messages_with_system(client, loaded_model, monkeypatch):
|
||||
def test_messages_with_system(client, loaded_model):
|
||||
"""POST /v1/messages with system prompt."""
|
||||
|
||||
async def async_gen():
|
||||
yield "Reply"
|
||||
|
||||
mock_engine = loaded_model
|
||||
mock_engine.generate_async.return_value = async_gen()
|
||||
monkeypatch.setattr("astrai.inference.server._state.engine", mock_engine)
|
||||
app.state.engine = loaded_model
|
||||
loaded_model.generate_async.return_value = async_gen()
|
||||
response = client.post(
|
||||
"/v1/messages",
|
||||
json={
|
||||
|
||||
Reference in New Issue
Block a user