chore: 解耦 Executor/Scheduler/TaskManager,修复 stop 页泄漏,移除 ServerState 全局单例

This commit is contained in:
2026-05-12 13:47:55 +08:00
parent 7440e9c809
commit df0845e916
10 changed files with 336 additions and 509 deletions
+10 -2
View File
@@ -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
+30 -21
View File
@@ -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
+14 -14
View File
@@ -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"]
+19 -22
View File
@@ -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={