From 2b712fa14c19053218617c19212f72fded23a35a Mon Sep 17 00:00:00 2001 From: ammaster10s Date: Wed, 9 Sep 2026 23:22:55 +0700 Subject: [PATCH] fix: enforce paid submission guardrails --- pipeline/governor.py | 12 +- pipeline/session.py | 110 ++++++++++---- tests/test_session_guardrails.py | 247 +++++++++++++++++++++++++++++++ 3 files changed, 333 insertions(+), 36 deletions(-) create mode 100644 tests/test_session_guardrails.py diff --git a/pipeline/governor.py b/pipeline/governor.py index 79f3b8e..0f7d73a 100644 --- a/pipeline/governor.py +++ b/pipeline/governor.py @@ -43,14 +43,10 @@ def project(state, corpus_rows, conflict_rate=None): # a shard that hit the 24h wall is not zero progress -- it was salvaged done[sh["stage"]] = done.get(sh["stage"], 0) + sh.get("salvaged", 0) elif sh.get("status") == "submitted" and sh.get("job"): - # in-flight: count live completionStats as "paid for" so remaining-work - # estimates don't double-count rows the job has already processed - try: - import batch - cs = batch.get_job(sh["job"]).get("completionStats", {}) - done[sh["stage"]] = done.get(sh["stage"], 0) + int(cs.get("successfulCount", 0)) - except Exception: - pass + # Unharvested jobs are absent from the spend ledger. Keep their rows in + # the future-cost estimate until harvest records the actual cost; live + # completionStats must not make reserved work disappear from projection. + continue rem_a = max(0, corpus_rows - done["pass_a"]) rem_b = max(0, corpus_rows - done["pass_b"]) rem_c = max(0, int(corpus_rows * conflict_rate) - done["pass_c"]) diff --git a/pipeline/session.py b/pipeline/session.py index a7c0d7a..b2c579a 100644 --- a/pipeline/session.py +++ b/pipeline/session.py @@ -4,7 +4,7 @@ python3 pipeline/session.py status what is in flight / done / pending python3 pipeline/session.py submit N build+submit up to N shards (governor-gated) python3 pipeline/session.py harvest collect finished jobs -> BQ + ledger - python3 pipeline/session.py run harvest, then submit to keep MAX_INFLIGHT busy + python3 pipeline/session.py run harvest, then submit up to MAX_INFLIGHT new shards State lives in state/shards.json, mirrored to GCS after every transition, so a session that dies mid-run resumes rather than restarts. @@ -20,6 +20,7 @@ # Measured 2026-08-31: Pro and Flash draw from SEPARATE batch throughput pools -- # running Pass B alongside Pass A left Pro at +11.9 rows/min (unchanged). So each # stage gets its own in-flight budget rather than sharing one. +# Legacy name: this caps new submissions per `run`; stage caps below bound live jobs. MAX_INFLIGHT = int(os.environ.get("PW_MAX_INFLIGHT", "3")) INFLIGHT_BY_STAGE = {"pass_a": int(os.environ.get("PW_INFLIGHT_A", "4")), "pass_b": int(os.environ.get("PW_INFLIGHT_B", "6")), @@ -128,42 +129,95 @@ def cmd_harvest(): if changed: save(st); ledger.sync() return st -def cmd_submit(n=None): +def cmd_submit(n=1): + if isinstance(n, bool) or not isinstance(n, int) or n < 0: + raise ValueError("submission limit must be a non-negative integer") + if n == 0: + print("submission limit is 0; nothing to submit") + return True + st = load() gov = governor.project(st, st["corpus"]) action, notes = governor.decide(gov) - if action == "HALT": - print("GOVERNOR HALT:", notes[0]); return + # Fail closed: CUT describes changes an operator must apply before spending. + # Treating it as advisory would submit the uncut workload that exceeded target. + if action not in ("PROCEED", "STRETCH"): + print(f"GOVERNOR {action}: submission blocked") + for note in notes: + print(" " + note) + return False print(f"governor: {action} (projected THB {gov['projected_total_thb']})") - todo = [] + submitted = 0 for stage, cap in INFLIGHT_BY_STAGE.items(): infl = sum(1 for s in st["shards"] if s["status"] == "submitted" and s["stage"] == stage) free = max(0, cap - infl) pend = [s for s in st["shards"] if s["status"] == "pending" and s["stage"] == stage] if pend and free == 0: print(f" {stage}: no slots ({infl}/{cap} in flight)") - todo += pend[:free] - if not todo: - return - for s in todo: - rows = rows_for(s["id"]) - _, kept, missing, attempt_id = batch.build( - rows, s["id"], s["model"], s["thinking"] - ) - if kept == 0: - print(f" {s['id']}: no rows with patches, skipping"); continue - job = batch.submit(s["id"], s["model"], attempt_id) - s["status"] = "submitted"; s["job"] = job["name"]; s["built"] = kept; s["missing"] = missing - s["attempt_id"] = attempt_id - s["submitted_at"] = datetime.datetime.now(datetime.timezone.utc).isoformat() - print(f" {s['id']}: submitted {kept} rows ({missing} missing patches) job={job['name'].split('/')[-1]}") - save(st) + for s in pend: + if submitted == n or free == 0: + break + rows = rows_for(s["id"]) + _, kept, missing, attempt_id = batch.build( + rows, s["id"], s["model"], s["thinking"] + ) + if kept == 0: + print(f" {s['id']}: no rows with patches, skipping") + continue + job = batch.submit(s["id"], s["model"], attempt_id) + s["status"] = "submitted" + s["job"] = job["name"] + s["built"] = kept + s["missing"] = missing + s["attempt_id"] = attempt_id + s["submitted_at"] = datetime.datetime.now(datetime.timezone.utc).isoformat() + job_id = job["name"].split("/")[-1] + print( + f" {s['id']}: submitted {kept} rows " + f"({missing} missing patches) job={job_id}" + ) + save(st) + submitted += 1 + free -= 1 + if submitted == n: + break + return True -if __name__ == "__main__": - cmd = sys.argv[1] if len(sys.argv) > 1 else "status" - if cmd == "status": cmd_status() - elif cmd == "harvest": cmd_harvest(); cmd_status() - elif cmd == "submit": cmd_submit(int(sys.argv[2]) if len(sys.argv)>2 else 1) +def _submit_limit(argv): + if len(argv) > 3: + raise ValueError("usage: session.py submit [non-negative integer]") + try: + limit = int(argv[2]) if len(argv) > 2 else 1 + except ValueError as exc: + raise ValueError("submission limit must be a non-negative integer") from exc + if limit < 0: + raise ValueError("submission limit must be a non-negative integer") + return limit + +def main(argv=None): + argv = sys.argv if argv is None else argv + cmd = argv[1] if len(argv) > 1 else "status" + if cmd == "status": + cmd_status() + elif cmd == "harvest": + cmd_harvest(); cmd_status() + elif cmd == "submit": + try: + limit = _submit_limit(argv) + except ValueError as e: + print(f"submit: {e}", file=sys.stderr) + return 2 + allowed = cmd_submit(limit) + if not allowed: + return 3 elif cmd == "run": - cmd_harvest(); cmd_submit(MAX_INFLIGHT); cmd_status() - else: print(__doc__) + cmd_harvest(); allowed = cmd_submit(MAX_INFLIGHT); cmd_status() + if not allowed: + return 3 + else: + print(__doc__) + return 2 + return 0 + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tests/test_session_guardrails.py b/tests/test_session_guardrails.py new file mode 100644 index 0000000..0ae6142 --- /dev/null +++ b/tests/test_session_guardrails.py @@ -0,0 +1,247 @@ +from pathlib import Path +import sys +import unittest +from unittest import mock + + +ROOT = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(ROOT / "pipeline")) + +import session +import governor + + +def shard(shard_id, stage, status="pending"): + return { + "id": shard_id, + "stage": stage, + "status": status, + "model": "gemini-test", + "thinking": 1024, + } + + +class SubmitGuardrailTests(unittest.TestCase): + def state(self): + return { + "corpus": 10, + "shards": [ + shard("a01", "pass_a"), + shard("a02", "pass_a"), + shard("b01", "pass_b"), + shard("c01", "pass_c"), + ], + } + + def submit_patches(self, state, action="PROCEED"): + return ( + mock.patch.object(session, "load", return_value=state), + mock.patch.object( + session.governor, + "project", + return_value={"projected_total_thb": 9000}, + ), + mock.patch.object( + session.governor, + "decide", + return_value=(action, [f"{action} note"]), + ), + mock.patch.object(session, "rows_for", return_value=[{"cve": "CVE-TEST"}]), + mock.patch.object( + session.batch, + "build", + side_effect=lambda rows, shard_id, model, thinking: ( + "input.jsonl", 1, 0, f"attempt-{shard_id}" + ), + ), + mock.patch.object( + session.batch, + "submit", + side_effect=lambda shard_id, model, attempt_id: { + "name": f"jobs/{shard_id}" + }, + ), + mock.patch.object(session, "save"), + ) + + def test_submit_limit_is_global_across_stage_caps(self): + state = self.state() + patches = self.submit_patches(state) + with patches[0], patches[1], patches[2], patches[3], patches[4] as build, \ + patches[5] as submit, patches[6] as save: + self.assertTrue(session.cmd_submit(2)) + + self.assertEqual(build.call_count, 2) + self.assertEqual(submit.call_count, 2) + self.assertEqual([c.args[0] for c in submit.call_args_list], ["a01", "a02"]) + self.assertEqual(save.call_count, 2) + self.assertEqual(state["shards"][2]["status"], "pending") + + def test_submit_limit_applies_after_full_stage(self): + state = self.state() + state["shards"] = [ + shard(f"a{i}", "pass_a", "submitted") for i in range(4) + ] + [shard("a-pending", "pass_a"), shard("b01", "pass_b")] + patches = self.submit_patches(state) + with patches[0], patches[1], patches[2], patches[3], patches[4], \ + patches[5] as submit, patches[6]: + self.assertTrue(session.cmd_submit(1)) + + submit.assert_called_once() + self.assertEqual(submit.call_args.args[0], "b01") + self.assertEqual(state["shards"][4]["status"], "pending") + + def test_sparse_stage_does_not_consume_later_stage_quota(self): + state = self.state() + state["shards"] = [ + shard("a01", "pass_a"), + shard("b01", "pass_b"), + shard("b02", "pass_b"), + ] + patches = self.submit_patches(state) + with patches[0], patches[1], patches[2], patches[3], patches[4], \ + patches[5] as submit, patches[6]: + self.assertTrue(session.cmd_submit(3)) + + self.assertEqual( + [call.args[0] for call in submit.call_args_list], + ["a01", "b01", "b02"], + ) + + def test_zero_row_build_does_not_consume_paid_submission_quota(self): + state = self.state() + patches = self.submit_patches(state) + with patches[0], patches[1], patches[2], patches[3], \ + patches[4] as build, patches[5] as submit, patches[6]: + build.side_effect = lambda rows, shard_id, model, thinking: ( + "input.jsonl", + 0 if shard_id == "a01" else 1, + 0, + f"attempt-{shard_id}", + ) + self.assertTrue(session.cmd_submit(1)) + + self.assertEqual(build.call_count, 2) + submit.assert_called_once() + self.assertEqual(submit.call_args.args[0], "a02") + + def test_halt_and_cut_both_block_submission(self): + for action in ("HALT", "CUT", "UNKNOWN"): + with self.subTest(action=action): + state = self.state() + patches = self.submit_patches(state, action) + with patches[0], patches[1], patches[2], patches[3], \ + patches[4] as build, patches[5] as submit, \ + patches[6] as save: + self.assertFalse(session.cmd_submit(2)) + + build.assert_not_called() + submit.assert_not_called() + save.assert_not_called() + self.assertTrue(all(s["status"] == "pending" for s in state["shards"])) + + def test_stretch_and_proceed_allow_submission(self): + for action in ("STRETCH", "PROCEED"): + with self.subTest(action=action): + state = self.state() + patches = self.submit_patches(state, action) + with patches[0], patches[1], patches[2], patches[3], \ + patches[4], patches[5] as submit, patches[6]: + self.assertTrue(session.cmd_submit(1)) + submit.assert_called_once() + + def test_zero_limit_is_a_successful_no_op(self): + with mock.patch.object(session, "load") as load: + self.assertTrue(session.cmd_submit(0)) + load.assert_not_called() + + def test_negative_or_non_integer_limits_fail_before_governor(self): + with mock.patch.object(session, "load") as load: + for value in (-1, 1.5, True, "1"): + with self.subTest(value=value), self.assertRaises(ValueError): + session.cmd_submit(value) + load.assert_not_called() + + +class SubmitArgumentTests(unittest.TestCase): + def test_submit_argument_defaults_to_one(self): + self.assertEqual(session._submit_limit(["session.py", "submit"]), 1) + + def test_submit_argument_allows_zero(self): + self.assertEqual(session._submit_limit(["session.py", "submit", "0"]), 0) + + def test_submit_argument_rejects_invalid_or_extra_values(self): + for argv in ( + ["session.py", "submit", "-1"], + ["session.py", "submit", "nope"], + ["session.py", "submit", "1", "extra"], + ): + with self.subTest(argv=argv), self.assertRaises(ValueError): + session._submit_limit(argv) + + def test_main_returns_distinct_usage_and_governor_exit_codes(self): + with mock.patch.object(session, "cmd_submit", return_value=False), \ + mock.patch("builtins.print"): + self.assertEqual(session.main(["session.py", "submit", "1"]), 3) + with mock.patch("builtins.print"): + self.assertEqual(session.main(["session.py", "submit", "invalid"]), 2) + + def test_main_does_not_mislabel_operational_value_errors(self): + with mock.patch.object( + session, "cmd_submit", side_effect=ValueError("build failed")): + with self.assertRaisesRegex(ValueError, "build failed"): + session.main(["session.py", "submit", "1"]) + + def test_run_reports_status_and_propagates_governor_block(self): + with mock.patch.object(session, "cmd_harvest") as harvest, \ + mock.patch.object(session, "cmd_submit", return_value=False) as submit, \ + mock.patch.object(session, "cmd_status") as status: + self.assertEqual(session.main(["session.py", "run"]), 3) + + harvest.assert_called_once_with() + submit.assert_called_once_with(session.MAX_INFLIGHT) + status.assert_called_once_with() + + +class GovernorReservationTests(unittest.TestCase): + def test_unharvested_rows_remain_in_future_cost_projection(self): + state = { + "measured": {"pass_a": 1.0, "pass_b": 1.0, "pass_c": 1.0}, + "shards": [{ + "id": "a01", "stage": "pass_a", "status": "submitted", + "job": "jobs/a01", "built": 80, + }], + } + with mock.patch.object(governor.ledger, "totals", return_value={"thb": 100.0}), \ + mock.patch.object( + session.batch, + "get_job", + return_value={"completionStats": {"successfulCount": 80}}, + ) as get_job: + projected = governor.project(state, 100, conflict_rate=0) + + get_job.assert_not_called() + self.assertEqual(projected["remaining"]["pass_a"], 100) + self.assertEqual(projected["future_thb"], 200.0) + self.assertEqual(projected["projected_total_thb"], 300.0) + + def test_decision_thresholds_are_inclusive_and_fail_safe(self): + cases = ( + ({"spent_thb": governor.HARDSTOP_THB, + "projected_total_thb": governor.HARDSTOP_THB}, "HALT"), + ({"spent_thb": 0, + "projected_total_thb": governor.TARGET_THB + 0.01}, "CUT"), + ({"spent_thb": 0, + "projected_total_thb": governor.TARGET_THB}, "PROCEED"), + ({"spent_thb": 0, + "projected_total_thb": governor.FLOOR_THB}, "PROCEED"), + ({"spent_thb": 0, + "projected_total_thb": governor.FLOOR_THB - 0.01}, "STRETCH"), + ) + for projection, expected in cases: + with self.subTest(expected=expected): + self.assertEqual(governor.decide(projection)[0], expected) + + +if __name__ == "__main__": + unittest.main()