story: DGR-032 Implement deterministic fake ShardEngine
This commit is contained in:
200
tests/test_fake_shard_engine.py
Normal file
200
tests/test_fake_shard_engine.py
Normal file
@@ -0,0 +1,200 @@
|
||||
"""DGR-032 ``FakeShardEngine`` tests.
|
||||
|
||||
``FakeShardEngine`` obeys the exact same lifecycle contract every
|
||||
``ShardEngine`` implementation must (see ``shard_engine_contract.py``); the
|
||||
tests here additionally cover this story's own scope: head/middle/tail
|
||||
output shape, isolated multi-session state, and the delay/memory-pressure/
|
||||
malformed-output/crash fault-injection knobs. None of this is real-model
|
||||
evidence — see the module docstring in ``fake_shard_engine.py`` and
|
||||
``test_fake_shard_engine_declares_fixture_evidence_class`` below, which pins
|
||||
the marker DGR-036 will rely on to tell a fixture engine apart from a real
|
||||
one when it certifies fixture-vs-real parity.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from meshnet_node.fake_shard_engine import (
|
||||
MALFORMED_TOKEN_ID_FLOOR,
|
||||
TOKEN_ID_VOCAB_SIZE,
|
||||
FakeShardEngine,
|
||||
FakeShardEngineConfig,
|
||||
)
|
||||
from meshnet_node.shard_engine import (
|
||||
BoundaryBundle,
|
||||
DecodeRequest,
|
||||
EngineTensor,
|
||||
LoadRequest,
|
||||
PrefillRequest,
|
||||
TokenOutput,
|
||||
)
|
||||
from meshnet_node.shard_lifecycle import StatusCode
|
||||
|
||||
from shard_engine_contract import assert_shard_engine_contract
|
||||
|
||||
|
||||
def _load(engine: FakeShardEngine, *, shard_start=0, shard_end=3, total_layers=4, recipe=None):
|
||||
return engine.load(
|
||||
LoadRequest(
|
||||
artifact_path="fixture://fake-shard-engine",
|
||||
shard_start=shard_start,
|
||||
shard_end=shard_end,
|
||||
total_layers=total_layers,
|
||||
recipe=recipe or {},
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def test_fake_shard_engine_obeys_the_shared_shard_engine_contract():
|
||||
assert_shard_engine_contract(FakeShardEngine)
|
||||
|
||||
|
||||
def test_fake_shard_engine_declares_fixture_evidence_class():
|
||||
# DGR-036's fixture-vs-real parity check needs a structural way to tell
|
||||
# a fixture engine apart from a real one; this constant is that marker.
|
||||
assert FakeShardEngine.EVIDENCE_CLASS == "fixture"
|
||||
|
||||
|
||||
def test_head_shard_returns_boundary_bundle_with_post_head_residual_point():
|
||||
engine = FakeShardEngine()
|
||||
_load(engine, shard_start=0, shard_end=1, total_layers=4)
|
||||
result = engine.prefill(
|
||||
PrefillRequest(session_id="s", route_epoch=1, position=0, idempotency_step=0, token_ids=(1, 2))
|
||||
)
|
||||
assert result.status.code is StatusCode.OK
|
||||
assert isinstance(result.output, BoundaryBundle)
|
||||
assert result.output.boundary_point == "post_head_residual"
|
||||
|
||||
|
||||
def test_middle_shard_returns_boundary_bundle_and_passes_through_token_sideband():
|
||||
engine = FakeShardEngine()
|
||||
_load(engine, shard_start=1, shard_end=2, total_layers=8)
|
||||
bundle_in = BoundaryBundle(
|
||||
tensors=(EngineTensor(name="hidden_states", shape=(1, 2), dtype="bfloat16", data=b"\x01\x02\x03\x04"),),
|
||||
architecture="dense",
|
||||
boundary_point="pre_tail_residual",
|
||||
token_id_sideband=(7, 8, 9),
|
||||
)
|
||||
result = engine.prefill(
|
||||
PrefillRequest(session_id="s", route_epoch=1, position=0, idempotency_step=0, input=bundle_in)
|
||||
)
|
||||
assert result.status.code is StatusCode.OK
|
||||
assert isinstance(result.output, BoundaryBundle)
|
||||
assert result.output.boundary_point == "post_middle_residual"
|
||||
assert result.output.token_id_sideband == (7, 8, 9)
|
||||
|
||||
|
||||
def test_tail_shard_returns_token_output_within_advertised_vocab():
|
||||
engine = FakeShardEngine()
|
||||
_load(engine, shard_start=0, shard_end=3, total_layers=4)
|
||||
result = engine.prefill(
|
||||
PrefillRequest(session_id="s", route_epoch=1, position=0, idempotency_step=0, token_ids=(1, 2, 3))
|
||||
)
|
||||
assert isinstance(result.output, TokenOutput)
|
||||
assert 0 <= result.output.token_id < TOKEN_ID_VOCAB_SIZE
|
||||
|
||||
|
||||
def test_session_state_is_isolated_between_two_concurrent_sessions():
|
||||
engine = FakeShardEngine()
|
||||
_load(engine)
|
||||
engine.prefill(PrefillRequest(session_id="a", route_epoch=5, position=0, idempotency_step=0, token_ids=(1,)))
|
||||
engine.prefill(PrefillRequest(session_id="b", route_epoch=1, position=0, idempotency_step=0, token_ids=(1,)))
|
||||
|
||||
# A stale epoch against session "a" must not affect session "b" at all.
|
||||
stale = engine.decode(DecodeRequest(session_id="a", route_epoch=4, position=1, idempotency_step=1, token_id=2))
|
||||
assert stale.status.code is StatusCode.FAILED_PRECONDITION
|
||||
|
||||
still_fine = engine.decode(DecodeRequest(session_id="b", route_epoch=1, position=1, idempotency_step=1, token_id=2))
|
||||
assert still_fine.status.code is StatusCode.OK
|
||||
|
||||
engine.cancel("a")
|
||||
after_cancel_b = engine.decode(
|
||||
DecodeRequest(session_id="b", route_epoch=1, position=2, idempotency_step=2, token_id=3)
|
||||
)
|
||||
assert after_cancel_b.status.code is StatusCode.OK, "cancelling session a must not cancel session b"
|
||||
|
||||
|
||||
def test_step_delay_seconds_invokes_the_configured_sleep_hook():
|
||||
calls: list[float] = []
|
||||
engine = FakeShardEngine(FakeShardEngineConfig(step_delay_seconds=0.25, sleep=calls.append))
|
||||
_load(engine)
|
||||
engine.prefill(PrefillRequest(session_id="s", route_epoch=1, position=0, idempotency_step=0, token_ids=(1,)))
|
||||
engine.decode(DecodeRequest(session_id="s", route_epoch=1, position=1, idempotency_step=1, token_id=2))
|
||||
assert calls == [0.25, 0.25]
|
||||
|
||||
|
||||
def test_memory_budget_bytes_trips_deterministic_resource_exhausted():
|
||||
engine = FakeShardEngine(FakeShardEngineConfig(memory_budget_bytes=8))
|
||||
_load(engine)
|
||||
first = engine.prefill(
|
||||
PrefillRequest(session_id="s", route_epoch=1, position=0, idempotency_step=0, token_ids=(1,))
|
||||
)
|
||||
assert first.status.code is StatusCode.OK # 8 bytes used, exactly at budget
|
||||
|
||||
second = engine.decode(DecodeRequest(session_id="s", route_epoch=1, position=1, idempotency_step=1, token_id=2))
|
||||
assert second.status.code is StatusCode.RESOURCE_EXHAUSTED
|
||||
assert second.status.retryable is True
|
||||
assert second.output is None
|
||||
|
||||
|
||||
def test_malformed_output_is_structurally_valid_but_semantically_wrong_for_tail():
|
||||
engine = FakeShardEngine(FakeShardEngineConfig(malformed_output=True))
|
||||
_load(engine, shard_start=0, shard_end=3, total_layers=4)
|
||||
result = engine.prefill(
|
||||
PrefillRequest(session_id="s", route_epoch=1, position=0, idempotency_step=0, token_ids=(1, 2, 3))
|
||||
)
|
||||
assert result.status.code is StatusCode.OK
|
||||
assert isinstance(result.output, TokenOutput)
|
||||
assert result.output.token_id >= MALFORMED_TOKEN_ID_FLOOR
|
||||
|
||||
|
||||
def test_malformed_output_is_structurally_valid_but_semantically_wrong_for_boundary_bundle():
|
||||
engine = FakeShardEngine(FakeShardEngineConfig(malformed_output=True))
|
||||
_load(engine, shard_start=0, shard_end=1, total_layers=4)
|
||||
result = engine.prefill(
|
||||
PrefillRequest(session_id="s", route_epoch=1, position=0, idempotency_step=0, token_ids=(1, 2, 3))
|
||||
)
|
||||
assert result.status.code is StatusCode.OK
|
||||
assert isinstance(result.output, BoundaryBundle)
|
||||
assert result.output.architecture.startswith("malformed:")
|
||||
assert len(result.output.tensors[0].data) == 1
|
||||
|
||||
|
||||
def test_crash_after_calls_raises_instead_of_returning_a_structured_status():
|
||||
engine = FakeShardEngine(FakeShardEngineConfig(crash_after_calls=2))
|
||||
_load(engine)
|
||||
ok = engine.prefill(PrefillRequest(session_id="s", route_epoch=1, position=0, idempotency_step=0, token_ids=(1,)))
|
||||
assert ok.status.code is StatusCode.OK
|
||||
|
||||
with pytest.raises(RuntimeError):
|
||||
engine.decode(DecodeRequest(session_id="s", route_epoch=1, position=1, idempotency_step=1, token_id=2))
|
||||
|
||||
|
||||
def test_crash_exception_factory_is_configurable():
|
||||
class _SimulatedSegfault(Exception):
|
||||
pass
|
||||
|
||||
engine = FakeShardEngine(
|
||||
FakeShardEngineConfig(crash_after_calls=1, crash_exception_factory=_SimulatedSegfault)
|
||||
)
|
||||
_load(engine)
|
||||
with pytest.raises(_SimulatedSegfault):
|
||||
engine.prefill(PrefillRequest(session_id="s", route_epoch=1, position=0, idempotency_step=0, token_ids=(1,)))
|
||||
|
||||
|
||||
def test_config_rejects_invalid_knob_values():
|
||||
with pytest.raises(ValueError):
|
||||
FakeShardEngineConfig(step_delay_seconds=-1.0)
|
||||
with pytest.raises(ValueError):
|
||||
FakeShardEngineConfig(memory_budget_bytes=-1)
|
||||
with pytest.raises(ValueError):
|
||||
FakeShardEngineConfig(crash_after_calls=0)
|
||||
|
||||
|
||||
def test_load_result_and_capabilities_report_recipe_architecture():
|
||||
engine = FakeShardEngine()
|
||||
load_result = _load(engine, recipe={"architecture": "deepseek-v4-flash"})
|
||||
assert load_result.architecture == "deepseek-v4-flash"
|
||||
caps = engine.capabilities()
|
||||
assert caps.architecture == "deepseek-v4-flash"
|
||||
Reference in New Issue
Block a user