diff --git a/scripts/agent_gan_inference_demo.py b/scripts/agent_gan_inference_demo.py index 9ac263b..053d202 100644 --- a/scripts/agent_gan_inference_demo.py +++ b/scripts/agent_gan_inference_demo.py @@ -18,7 +18,10 @@ def _agent_cache_gate(warm_delta: dict, actual_delta: dict) -> bool: return ( - warm_delta["remote_hits"] >= 1 + ( + warm_delta.get("remote_hits", 0) >= 1 + or warm_delta.get("local_hits", 0) >= 1 + ) and warm_delta["tokens_reused"] >= 1 and warm_delta["tokens_computed"] == 0 and warm_delta["fallbacks"] == 0 diff --git a/scripts/agent_gan_repl.py b/scripts/agent_gan_repl.py index d374bfe..4fe4636 100644 --- a/scripts/agent_gan_repl.py +++ b/scripts/agent_gan_repl.py @@ -157,6 +157,23 @@ def _stage( return stage +def _gate_failure(name: str, warm: dict, actual: dict) -> RuntimeError: + keys = ( + "local_hits", + "remote_hits", + "remote_jobs", + "tokens_reused", + "tokens_computed", + "fallbacks", + "remote_job_failures", + ) + compact = lambda delta: {key: delta.get(key, 0) for key in keys} + return RuntimeError( + f"{name} KV gate failed: " + f"warm={compact(warm['delta'])} actual={compact(actual['delta'])}", + ) + + def main() -> int: install_signal_protection() parser = argparse.ArgumentParser() @@ -274,7 +291,11 @@ def get_stats(): generator_text, ) if not generator_stage["ok"] and not telemetry_state["degraded"]: - raise RuntimeError("Generator KV gate failed") + raise _gate_failure( + "Generator", + generator_warm, + generator_actual, + ) if remote_run: _telemetry_request( f"{args.dashboard}/v1/network/benchmarks/{run_id}", @@ -337,7 +358,7 @@ def get_stats(): extra_metrics=context_metrics, ) if not critic_stage["ok"] and not telemetry_state["degraded"]: - raise RuntimeError("Critic KV gate failed") + raise _gate_failure("Critic", critic_warm, critic_actual) completed = None if remote_run: completed = _telemetry_request( diff --git a/tests/inference_engine/bridge/test_agent_gan_demo.py b/tests/inference_engine/bridge/test_agent_gan_demo.py index 5b91fd5..749e7a2 100644 --- a/tests/inference_engine/bridge/test_agent_gan_demo.py +++ b/tests/inference_engine/bridge/test_agent_gan_demo.py @@ -23,6 +23,10 @@ def test_agent_gate_accepts_remote_compute_or_exact_remote_cache_hit(): } assert _agent_cache_gate(warm, actual) assert _agent_cache_gate({**warm, "remote_jobs": 0}, actual) + assert _agent_cache_gate( + {**warm, "remote_hits": 0, "local_hits": 1, "remote_jobs": 0}, + actual, + ) assert not _agent_cache_gate({**warm, "remote_hits": 0}, actual) assert not _agent_cache_gate({**warm, "tokens_reused": 0}, actual) assert not _agent_cache_gate({**warm, "fallbacks": 1}, actual) diff --git a/tests/inference_engine/bridge/test_agent_gan_repl.py b/tests/inference_engine/bridge/test_agent_gan_repl.py index 3c9e0f0..a7912f5 100644 --- a/tests/inference_engine/bridge/test_agent_gan_repl.py +++ b/tests/inference_engine/bridge/test_agent_gan_repl.py @@ -5,6 +5,7 @@ from scripts.agent_gan_repl import ( PrefillHeartbeat, TokenPrinter, + _gate_failure, _stage, _telemetry_request, build_critic_messages, @@ -160,6 +161,18 @@ def timeout(*_args, **_kwargs): assert "inference will continue" in output +def test_gate_failure_exposes_reuse_counters(): + error = _gate_failure( + "Generator", + {"delta": {"local_hits": 0, "remote_hits": 0}}, + {"delta": {"local_hits": 0, "fallbacks": 1}}, + ) + message = str(error) + assert "Generator KV gate failed" in message + assert "'remote_hits': 0" in message + assert "'fallbacks': 1" in message + + def test_interactive_prompts_are_deterministic_for_kv_reuse(): generator_a = build_generator_messages("prove RH") generator_b = build_generator_messages("prove RH")