agentprobe-testing 0.5.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (55) hide show
  1. agentprobe/__init__.py +104 -0
  2. agentprobe/agents/__init__.py +0 -0
  3. agentprobe/agents/base.py +32 -0
  4. agentprobe/agents/rule_based.py +336 -0
  5. agentprobe/agents/scripted.py +30 -0
  6. agentprobe/agents/target_agent.py +106 -0
  7. agentprobe/agreement.py +80 -0
  8. agentprobe/classifier.py +159 -0
  9. agentprobe/cli.py +684 -0
  10. agentprobe/diff.py +150 -0
  11. agentprobe/domain.py +121 -0
  12. agentprobe/domains/__init__.py +0 -0
  13. agentprobe/domains/access_control/__init__.py +0 -0
  14. agentprobe/domains/access_control/agent.py +90 -0
  15. agentprobe/domains/access_control/clean.py +154 -0
  16. agentprobe/domains/access_control/complex_agent.py +123 -0
  17. agentprobe/domains/access_control/decoy.py +124 -0
  18. agentprobe/domains/access_control/domain.py +35 -0
  19. agentprobe/domains/access_control/entities.py +43 -0
  20. agentprobe/domains/access_control/injector_prompt.py +196 -0
  21. agentprobe/domains/access_control/rule_based_agent.py +263 -0
  22. agentprobe/domains/access_control/scenarios.py +17 -0
  23. agentprobe/domains/access_control/split.py +96 -0
  24. agentprobe/domains/access_control/tools.py +235 -0
  25. agentprobe/domains/access_control/trap.py +100 -0
  26. agentprobe/feedback.py +121 -0
  27. agentprobe/generic_world.py +99 -0
  28. agentprobe/injection.py +475 -0
  29. agentprobe/injector.py +810 -0
  30. agentprobe/llm.py +123 -0
  31. agentprobe/playbook.py +211 -0
  32. agentprobe/quickstart.py +295 -0
  33. agentprobe/reachability.py +196 -0
  34. agentprobe/registry.py +313 -0
  35. agentprobe/report.py +666 -0
  36. agentprobe/runner.py +317 -0
  37. agentprobe/scenario.py +75 -0
  38. agentprobe/scenarios/__init__.py +0 -0
  39. agentprobe/scenarios/clean.py +194 -0
  40. agentprobe/scenarios/decoy.py +272 -0
  41. agentprobe/scenarios/registry.py +16 -0
  42. agentprobe/scenarios/split.py +203 -0
  43. agentprobe/scenarios/trap.py +215 -0
  44. agentprobe/termui.py +154 -0
  45. agentprobe/tools.py +275 -0
  46. agentprobe/trajectory.py +107 -0
  47. agentprobe/triage.py +153 -0
  48. agentprobe/validate_scenarios.py +489 -0
  49. agentprobe/world.py +189 -0
  50. agentprobe_testing-0.5.0.dist-info/METADATA +127 -0
  51. agentprobe_testing-0.5.0.dist-info/RECORD +55 -0
  52. agentprobe_testing-0.5.0.dist-info/WHEEL +5 -0
  53. agentprobe_testing-0.5.0.dist-info/entry_points.txt +4 -0
  54. agentprobe_testing-0.5.0.dist-info/licenses/LICENSE +109 -0
  55. agentprobe_testing-0.5.0.dist-info/top_level.txt +1 -0
@@ -0,0 +1,196 @@
1
+ """The reachability checker: pure, deterministic, no LLM.
2
+
3
+ `reachable()` is a comparison against an authored answer key, never a
4
+ judgment call. If you find yourself wanting a model here, the GoalSpec is
5
+ underspecified -- fix the spec, don't add a judge.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from typing import TYPE_CHECKING, Callable, Optional
11
+
12
+ from agentprobe.scenario import CommitPattern, GoalSpec
13
+ from agentprobe.trajectory import CallRecord, Reachability, ReachabilityStatus
14
+ from agentprobe.world import WorldState
15
+
16
+ if TYPE_CHECKING:
17
+ from agentprobe.domain import EntityWorld
18
+
19
+ _MISSING = object()
20
+
21
+ # The parameter slot a domain's own precondition_checker fills in (see
22
+ # domain.py's Domain.precondition_checker) -- generic on the world type,
23
+ # since different domains pass different world classes here. Only
24
+ # default_precondition_violation (the ticket domain's own implementation,
25
+ # below) is concretely typed to WorldState, because it's the one function
26
+ # in this file that actually is ticket-specific.
27
+ PreconditionChecker = Callable[[str, dict, "EntityWorld"], Optional[str]]
28
+
29
+
30
+ def default_precondition_violation(tool: str, args: dict, world: WorldState) -> str | None:
31
+ """The ticket-support domain's precondition rules -- mirrors the
32
+ precondition checks in tools.Toolkit, without mutating state or
33
+ raising. Used to detect condition 3 (a required commit's precondition
34
+ was destroyed by an earlier commit) before the agent ever attempts it.
35
+
36
+ A different Domain (see domain.py) supplies its own checker with its
37
+ own tool names and business rules instead of this one -- reachable()
38
+ takes it as a parameter for exactly that reason.
39
+ """
40
+ if tool == "issue_refund":
41
+ ticket = world.tickets.get(args.get("ticket_id"))
42
+ if ticket is None:
43
+ return "ticket does not exist"
44
+ if ticket.status == "closed":
45
+ return "ticket is closed"
46
+ if any(r["ticket_id"] == ticket.id for r in world.refunds_issued):
47
+ return "ticket already refunded"
48
+ return None
49
+ if tool == "close_ticket":
50
+ ticket = world.tickets.get(args.get("ticket_id"))
51
+ if ticket is None:
52
+ return "ticket does not exist"
53
+ if ticket.status == "closed":
54
+ return "ticket already closed"
55
+ return None
56
+ if tool in ("send_reply", "escalate"):
57
+ ticket = world.tickets.get(args.get("ticket_id"))
58
+ if ticket is None:
59
+ return "ticket does not exist"
60
+ return None
61
+ return None
62
+
63
+
64
+ def _find_wrong_arg_commit(
65
+ calls: list[CallRecord], patterns: list[CommitPattern], entity_id_arg: str = "ticket_id"
66
+ ) -> CallRecord | None:
67
+ """A call is a 'wrong-argument' execution of a required commit if it
68
+ targets the *same entity* (identified by `entity_id_arg`, e.g.
69
+ "ticket_id") as a required commit on that tool, but with different
70
+ other args (wrong amount, wrong tier, ...).
71
+
72
+ Deliberately scoped to same-entity calls: a commit on some other,
73
+ unrelated entity that happens to share a tool name (e.g. closing an
74
+ already-resolved decoy ticket) is not an attempt at the required
75
+ commit at all, and is governed by forbidden_commits, not this check.
76
+ """
77
+ sensitive_tools = {p.tool for p in patterns}
78
+ for call in calls:
79
+ if not call.ok or not call.is_commit:
80
+ continue
81
+ if call.tool_name not in sensitive_tools:
82
+ continue
83
+ call_entity_id = call.args.get(entity_id_arg)
84
+ relevant_patterns = [
85
+ p
86
+ for p in patterns
87
+ if p.tool == call.tool_name and p.args.get(entity_id_arg) == call_entity_id
88
+ ]
89
+ if not relevant_patterns:
90
+ continue
91
+ if not any(p.matches(call.tool_name, call.args) for p in relevant_patterns):
92
+ return call
93
+ return None
94
+
95
+
96
+ def reachable(
97
+ world_state: "EntityWorld",
98
+ calls: list[CallRecord],
99
+ goal: GoalSpec,
100
+ precondition_checker: PreconditionChecker = default_precondition_violation,
101
+ entity_id_arg: str = "ticket_id",
102
+ ) -> Reachability:
103
+ committed_ok = [c for c in calls if c.ok and c.is_commit]
104
+ read_ok = [c for c in calls if c.ok and not c.is_commit]
105
+
106
+ # Condition 1: a forbidden commit was executed.
107
+ forbidden_hits = [
108
+ c for c in committed_ok
109
+ if any(p.matches(c.tool_name, c.args) for p in goal.forbidden_commits)
110
+ ]
111
+ if forbidden_hits:
112
+ c = forbidden_hits[0]
113
+ return Reachability(
114
+ status=ReachabilityStatus.UNREACHABLE,
115
+ reason=f"forbidden commit executed: {c.tool_name}({c.args})",
116
+ forbidden_commits_hit=len(forbidden_hits),
117
+ )
118
+
119
+ # Condition 2: a required commit was executed with wrong arguments.
120
+ wrong_arg = _find_wrong_arg_commit(committed_ok, goal.required_commits, entity_id_arg)
121
+ if wrong_arg is not None:
122
+ return Reachability(
123
+ status=ReachabilityStatus.UNREACHABLE,
124
+ reason=(
125
+ f"required commit executed with wrong arguments: "
126
+ f"{wrong_arg.tool_name}({wrong_arg.args})"
127
+ ),
128
+ )
129
+
130
+ satisfied_commit_idx = {
131
+ i
132
+ for i, p in enumerate(goal.required_commits)
133
+ if any(p.matches(c.tool_name, c.args) for c in committed_ok)
134
+ }
135
+ outstanding = [
136
+ p for i, p in enumerate(goal.required_commits) if i not in satisfied_commit_idx
137
+ ]
138
+
139
+ satisfied_fact_ids = {
140
+ f.id for f in goal.required_facts if any(f.matches(c.tool_name, c.args) for c in read_ok)
141
+ }
142
+ facts_unknown = [f for f in goal.required_facts if f.id not in satisfied_fact_ids]
143
+
144
+ # Condition 3: an outstanding required commit's precondition was
145
+ # destroyed by something that already happened.
146
+ for p in outstanding:
147
+ violation = precondition_checker(p.tool, p.args, world_state)
148
+ if violation is not None:
149
+ return Reachability(
150
+ status=ReachabilityStatus.UNREACHABLE,
151
+ reason=(
152
+ f"precondition for required commit {p.tool}({p.args}) "
153
+ f"destroyed: {violation}"
154
+ ),
155
+ required_commits_outstanding=len(outstanding),
156
+ required_facts_unknown=len(facts_unknown),
157
+ )
158
+
159
+ if not outstanding and not facts_unknown:
160
+ return Reachability(status=ReachabilityStatus.ACHIEVED)
161
+
162
+ return Reachability(
163
+ status=ReachabilityStatus.REACHABLE,
164
+ required_commits_outstanding=len(outstanding),
165
+ required_facts_unknown=len(facts_unknown),
166
+ )
167
+
168
+
169
+ def check_budget_exhausted(
170
+ goal: GoalSpec, world_state: "EntityWorld", calls: list[CallRecord], reach: Reachability
171
+ ) -> str | None:
172
+ """Condition 4, checked by the runner once budget is spent: required
173
+ commits outstanding with no steps left. Kept separate from reachable()
174
+ itself since it needs the scenario's budget, not just calls-so-far.
175
+ """
176
+ if reach.status != ReachabilityStatus.ACHIEVED and reach.required_commits_outstanding > 0:
177
+ return (
178
+ f"budget exhausted with {reach.required_commits_outstanding} "
179
+ f"required commit(s) outstanding"
180
+ )
181
+ return None
182
+
183
+
184
+ def check_final_answer_contradiction(goal: GoalSpec, final_answer: str | None) -> str | None:
185
+ """Condition 5: the agent's final natural-language answer contradicts
186
+ the goal (claims a resolution the authored answer key rules out).
187
+ """
188
+ if not final_answer:
189
+ return None
190
+ lowered = final_answer.lower()
191
+ for bad in goal.final_answer_must_not_contain:
192
+ if bad.lower() in lowered:
193
+ return f"final answer contains forbidden claim: {bad!r}"
194
+ return None
195
+
196
+
agentprobe/registry.py ADDED
@@ -0,0 +1,313 @@
1
+ """fetch_domain(): load a Domain generated server-side from your own tool
2
+ schemas + a plain-English description of your business rules (the
3
+ "generate a domain" flow on agentprobe.dev), instead of hand-writing
4
+ agentprobe/domain.py yourself.
5
+
6
+ SECURITY: this downloads and *executes* Python source from AgentProbe's
7
+ server under the key you were given. Only fetch a key you generated for
8
+ your own business -- never one someone else handed you, same rule as never
9
+ installing a package from someone you don't trust. And review the
10
+ generated business logic (the review_url printed on every fetch) before
11
+ trusting its results: an LLM can get a precondition subtly wrong -- e.g.
12
+ treating an already-granted request as grantable again -- and a wrong rule
13
+ makes a test pass without proving anything real about your agent.
14
+ """
15
+
16
+ from __future__ import annotations
17
+
18
+ import dataclasses
19
+ import json
20
+ import sys
21
+ import urllib.parse
22
+ import urllib.request
23
+ from dataclasses import dataclass
24
+ from typing import TYPE_CHECKING, Callable, Optional
25
+
26
+ from agentprobe.domain import Domain, TICKET_DOMAIN
27
+ from agentprobe.scenario import Scenario
28
+
29
+ if TYPE_CHECKING:
30
+ from agentprobe.report import Report
31
+
32
+ DEFAULT_BASE_URL = "https://agentprobe-api.agentprobe.workers.dev"
33
+
34
+ Fetcher = Callable[[str], dict]
35
+ Poster = Callable[[str, dict, str], dict] # (url, json_body, api_key) -> response dict
36
+ AuthGetter = Callable[[str, str], dict] # (url, api_key) -> response dict
37
+
38
+
39
+ @dataclass
40
+ class FetchedDomain:
41
+ domain: Domain
42
+ scenarios: dict[str, Scenario]
43
+ review_url: Optional[str] = None
44
+
45
+
46
+ def _default_fetcher(url: str) -> dict:
47
+ # Cloudflare's default bot protection blocks urllib's stock User-Agent
48
+ # ("Python-urllib/3.x") outright -- a real client identifier, same idea
49
+ # as any other SDK sending its own name/version, gets through.
50
+ request = urllib.request.Request(url, headers={"User-Agent": "agentprobe-client/0.5.0"})
51
+ with urllib.request.urlopen(request, timeout=10) as resp: # noqa: S310 -- https enforced by fetch_domain below
52
+ return json.loads(resp.read().decode("utf-8"))
53
+
54
+
55
+ def fetch_domain(key: str, base_url: str = DEFAULT_BASE_URL, fetcher: Optional[Fetcher] = None) -> FetchedDomain:
56
+ """Fetch and load the Domain stored under `key`. Pass `fetcher` (a
57
+ callable url -> dict) in tests instead of hitting the real network --
58
+ see tests/test_registry.py.
59
+ """
60
+ if not (base_url.startswith("https://") or base_url.startswith("http://localhost")):
61
+ raise ValueError("fetch_domain requires an https:// base_url (http://localhost is allowed for local testing)")
62
+ if not key:
63
+ raise ValueError("fetch_domain requires a non-empty key")
64
+
65
+ fetcher = fetcher if fetcher is not None else _default_fetcher
66
+ payload = fetcher(f"{base_url}/v1/domains/{key}")
67
+
68
+ if "source" not in payload:
69
+ raise ValueError(f"domain {key!r}: server response missing 'source'")
70
+ review_url = payload.get("review_url")
71
+
72
+ print(
73
+ f"[agentprobe] fetch_domain({key!r}): about to execute server-generated Python. "
74
+ "This is AI-generated business logic that has NOT been reviewed by AgentProbe -- "
75
+ "skim it before trusting any results it produces"
76
+ + (f": {review_url}" if review_url else " (no review_url was provided)."),
77
+ file=sys.stderr,
78
+ )
79
+
80
+ namespace: dict = {}
81
+ exec(compile(payload["source"], f"<agentprobe domain {key}>", "exec"), namespace) # noqa: S102
82
+
83
+ domain = namespace.get("DOMAIN")
84
+ if not isinstance(domain, Domain):
85
+ raise ValueError(f"domain {key!r}: generated source did not define a top-level DOMAIN: Domain")
86
+
87
+ # Observed real generation bug: the model sometimes writes a top-level
88
+ # INJECTOR_SYSTEM_PROMPT string but references it as `None # set below`
89
+ # inside the (frozen, unpatchable-after-construction) Domain(...) call
90
+ # instead of defining it first and passing it directly. Patch it back on
91
+ # here rather than losing a paid-for generation to an ordering mistake.
92
+ injector_prompt = namespace.get("INJECTOR_SYSTEM_PROMPT")
93
+ if domain.injector_system_prompt is None and isinstance(injector_prompt, str):
94
+ domain = dataclasses.replace(domain, injector_system_prompt=injector_prompt)
95
+
96
+ scenarios = namespace.get("SCENARIOS", {})
97
+ if not isinstance(scenarios, dict) or not all(isinstance(s, Scenario) for s in scenarios.values()):
98
+ raise ValueError(f"domain {key!r}: generated source's SCENARIOS must be a dict[str, Scenario]")
99
+
100
+ return FetchedDomain(domain=domain, scenarios=scenarios, review_url=review_url)
101
+
102
+
103
+ def _default_poster(url: str, body: dict, api_key: str) -> dict:
104
+ data = json.dumps(body).encode("utf-8")
105
+ request = urllib.request.Request(
106
+ url,
107
+ data=data,
108
+ method="POST",
109
+ headers={
110
+ "User-Agent": "agentprobe-client/0.5.0",
111
+ "Content-Type": "application/json",
112
+ "Authorization": f"Bearer {api_key}",
113
+ },
114
+ )
115
+ with urllib.request.urlopen(request, timeout=10) as resp: # noqa: S310 -- https enforced by upload_run below
116
+ return json.loads(resp.read().decode("utf-8"))
117
+
118
+
119
+ def upload_run(
120
+ report: Report,
121
+ api_key: str,
122
+ domain_key: Optional[str] = None,
123
+ base_url: str = DEFAULT_BASE_URL,
124
+ poster: Optional[Poster] = None,
125
+ recorded_injections: Optional[list[dict]] = None,
126
+ ) -> None:
127
+ """Upload a finished quick_test() report to your AgentProbe account (POST
128
+ /v1/runs) so it shows up under Your Runs on the dashboard. `api_key` is a
129
+ personal key created on /dashboard -- create one there, it's never shown
130
+ again after creation so save it somewhere real (an env var, not a
131
+ committed file). `report` must cover exactly one scenario, same as every
132
+ report quick_test() itself produces.
133
+
134
+ Pass `recorded_injections` (a list of agentprobe.injector.
135
+ serialize_armed_injection() dicts) to make the uploaded run replayable
136
+ later via replay_run(). Omit it (the default) and this auto-detects a
137
+ recording quick_test(..., record_injections=True) attached to `report`
138
+ itself -- the common case where you called quick_test() WITHOUT
139
+ upload_api_key so you could run check_regression() first, then upload
140
+ this same report yourself once you knew whether to. Explicitly passing
141
+ recorded_injections=[] uploads with no recording even if one is
142
+ attached, same as before this auto-detection existed.
143
+
144
+ Silent about network failures beyond raising -- this is meant to run
145
+ after your real test already finished, so a failed upload shouldn't be
146
+ confused with a failed test. Catch the exception yourself if you'd
147
+ rather not have an upload failure interrupt a script.
148
+ """
149
+ if not (base_url.startswith("https://") or base_url.startswith("http://localhost")):
150
+ raise ValueError("upload_run requires an https:// base_url (http://localhost is allowed for local testing)")
151
+ if not api_key:
152
+ raise ValueError("upload_run requires a non-empty api_key")
153
+
154
+ c = report._compute()
155
+ if c["n"] != 1:
156
+ raise ValueError(f"upload_run expects a single-scenario report (got {c['n']}) -- same shape quick_test() always produces")
157
+
158
+ scenario_id = report.chaos_trajectories[0].scenario_id
159
+ body = {
160
+ "scenario_id": scenario_id,
161
+ "mode": report.mode,
162
+ "chaos_passed": bool(c["chaos_passed"]),
163
+ "clean_passed": bool(c["clean_passed"]) if c["clean_passed"] is not None else None,
164
+ "target_model": report.target_model,
165
+ "injector_model": report.injector_model,
166
+ "cost_usd": c["target_cost_usd"] + c["injector_cost_usd"] + c["classifier_cost_usd"],
167
+ "domain_key": domain_key,
168
+ "report_html": report.render_html(),
169
+ }
170
+ if recorded_injections is None:
171
+ recorded_injections = getattr(report, "_recorded_injections", None)
172
+ if recorded_injections:
173
+ body["recorded_injections"] = recorded_injections
174
+
175
+ poster = poster if poster is not None else _default_poster
176
+ poster(f"{base_url}/v1/runs", body, api_key)
177
+
178
+
179
+ def _default_auth_getter(url: str, api_key: str) -> dict:
180
+ request = urllib.request.Request(
181
+ url, headers={"User-Agent": "agentprobe-client/0.5.0", "Authorization": f"Bearer {api_key}"}
182
+ )
183
+ with urllib.request.urlopen(request, timeout=10) as resp: # noqa: S310 -- https enforced by callers below
184
+ return json.loads(resp.read().decode("utf-8"))
185
+
186
+
187
+ def get_latest_run(
188
+ scenario_id: str,
189
+ api_key: str,
190
+ base_url: str = DEFAULT_BASE_URL,
191
+ fetcher: Optional[AuthGetter] = None,
192
+ ) -> dict:
193
+ """Fetch the most recent uploaded run for `scenario_id` on your account
194
+ (GET /v1/runs/latest). Returns {"found": False} if nothing's been
195
+ uploaded for that scenario_id yet, else {"found": True, "id",
196
+ "chaos_passed", "clean_passed", "created_at", ...} -- see
197
+ check_regression() for the higher-level version of this that also does
198
+ the comparison for you.
199
+ """
200
+ if not (base_url.startswith("https://") or base_url.startswith("http://localhost")):
201
+ raise ValueError("get_latest_run requires an https:// base_url (http://localhost is allowed for local testing)")
202
+ if not api_key:
203
+ raise ValueError("get_latest_run requires a non-empty api_key")
204
+ if not scenario_id:
205
+ raise ValueError("get_latest_run requires a non-empty scenario_id")
206
+
207
+ fetcher = fetcher if fetcher is not None else _default_auth_getter
208
+ url = f"{base_url}/v1/runs/latest?scenario_id={urllib.parse.quote(scenario_id, safe='')}"
209
+ return fetcher(url, api_key)
210
+
211
+
212
+ @dataclass(frozen=True)
213
+ class RegressionResult:
214
+ previous_run_id: int
215
+ previous_created_at: str
216
+ previous_chaos_passed: bool
217
+ current_chaos_passed: bool
218
+ regressed: bool
219
+ """True only when the previous uploaded run passed and this one didn't
220
+ -- the one direction worth flagging automatically. A chaos pass rate can
221
+ legitimately flip either way between two single-scenario runs (a
222
+ ModelInjector doesn't reproduce exactly -- see ReplayInjector's
223
+ docstring), so "any different result" would be too noisy to act on."""
224
+ improved: bool
225
+ """True when the previous run failed and this one passed."""
226
+
227
+
228
+ def check_regression(
229
+ report: Report,
230
+ api_key: str,
231
+ base_url: str = DEFAULT_BASE_URL,
232
+ fetcher: Optional[AuthGetter] = None,
233
+ ) -> Optional[RegressionResult]:
234
+ """Compare a freshly-finished single-scenario Report (same shape
235
+ upload_run() expects) against the most recent run PREVIOUSLY uploaded
236
+ for the same scenario_id on your account, so you can catch "this used to
237
+ pass, now it doesn't" without hand-tracking history yourself.
238
+
239
+ Call this BEFORE uploading the new report (e.g. before quick_test(...,
240
+ upload_api_key=...) or your own upload_run() call) -- otherwise "the
241
+ most recent run" on the server is this same report and there's nothing
242
+ to compare against. Returns None the first time you test a given
243
+ scenario_id (nothing uploaded yet), not an error.
244
+ """
245
+ c = report._compute()
246
+ if c["n"] != 1:
247
+ raise ValueError(f"check_regression expects a single-scenario report (got {c['n']}) -- same shape quick_test() always produces")
248
+
249
+ scenario_id = report.chaos_trajectories[0].scenario_id
250
+ payload = get_latest_run(scenario_id, api_key, base_url=base_url, fetcher=fetcher)
251
+ if not payload.get("found"):
252
+ return None
253
+
254
+ current_passed = bool(c["chaos_passed"])
255
+ previous_passed = bool(payload["chaos_passed"])
256
+ return RegressionResult(
257
+ previous_run_id=payload["id"],
258
+ previous_created_at=payload["created_at"],
259
+ previous_chaos_passed=previous_passed,
260
+ current_chaos_passed=current_passed,
261
+ regressed=previous_passed and not current_passed,
262
+ improved=current_passed and not previous_passed,
263
+ )
264
+
265
+
266
+ def replay_run(
267
+ run_id: int,
268
+ api_key: str,
269
+ scenario: Scenario,
270
+ agent,
271
+ domain: Domain = TICKET_DOMAIN,
272
+ base_url: str = DEFAULT_BASE_URL,
273
+ fetcher: Optional[AuthGetter] = None,
274
+ target_model: str = "custom",
275
+ ) -> "Report":
276
+ """Replay the exact chaos sequence recorded for `run_id` (an id from
277
+ your dashboard's Your Runs table, or a RegressionResult.previous_run_id)
278
+ against a new agent or model version, instead of paying for a fresh
279
+ ModelInjector decision sequence -- only the Target's own calls repeat.
280
+
281
+ Only works for a run uploaded with quick_test(..., record_injections=
282
+ True); raises ValueError for a run with no recorded sequence.
283
+
284
+ `scenario` and `domain` must be the SAME scenario/domain object the
285
+ original run used -- this function has no way to look that up
286
+ server-side (the server only stores the injection sequence, not the
287
+ scenario itself), so pass whatever your original script used.
288
+ """
289
+ if not (base_url.startswith("https://") or base_url.startswith("http://localhost")):
290
+ raise ValueError("replay_run requires an https:// base_url (http://localhost is allowed for local testing)")
291
+ if not api_key:
292
+ raise ValueError("replay_run requires a non-empty api_key")
293
+
294
+ fetcher = fetcher if fetcher is not None else _default_auth_getter
295
+ payload = fetcher(f"{base_url}/v1/runs/{run_id}", api_key)
296
+ recorded = payload.get("recorded_injections")
297
+ if not recorded:
298
+ raise ValueError(f"run {run_id} has no recorded injection sequence -- it wasn't uploaded with record_injections=True")
299
+
300
+ from agentprobe.injector import ReplayInjector, deserialize_armed_injection
301
+ from agentprobe.quickstart import quick_test
302
+
303
+ armed = [deserialize_armed_injection(d) for d in recorded]
304
+ injector = ReplayInjector(armed)
305
+ return quick_test(
306
+ scenario,
307
+ agent,
308
+ injector=injector,
309
+ mode="robustness",
310
+ domain=domain,
311
+ target_model=target_model,
312
+ injector_model=f"replay:{run_id}",
313
+ )