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.
- agentprobe/__init__.py +104 -0
- agentprobe/agents/__init__.py +0 -0
- agentprobe/agents/base.py +32 -0
- agentprobe/agents/rule_based.py +336 -0
- agentprobe/agents/scripted.py +30 -0
- agentprobe/agents/target_agent.py +106 -0
- agentprobe/agreement.py +80 -0
- agentprobe/classifier.py +159 -0
- agentprobe/cli.py +684 -0
- agentprobe/diff.py +150 -0
- agentprobe/domain.py +121 -0
- agentprobe/domains/__init__.py +0 -0
- agentprobe/domains/access_control/__init__.py +0 -0
- agentprobe/domains/access_control/agent.py +90 -0
- agentprobe/domains/access_control/clean.py +154 -0
- agentprobe/domains/access_control/complex_agent.py +123 -0
- agentprobe/domains/access_control/decoy.py +124 -0
- agentprobe/domains/access_control/domain.py +35 -0
- agentprobe/domains/access_control/entities.py +43 -0
- agentprobe/domains/access_control/injector_prompt.py +196 -0
- agentprobe/domains/access_control/rule_based_agent.py +263 -0
- agentprobe/domains/access_control/scenarios.py +17 -0
- agentprobe/domains/access_control/split.py +96 -0
- agentprobe/domains/access_control/tools.py +235 -0
- agentprobe/domains/access_control/trap.py +100 -0
- agentprobe/feedback.py +121 -0
- agentprobe/generic_world.py +99 -0
- agentprobe/injection.py +475 -0
- agentprobe/injector.py +810 -0
- agentprobe/llm.py +123 -0
- agentprobe/playbook.py +211 -0
- agentprobe/quickstart.py +295 -0
- agentprobe/reachability.py +196 -0
- agentprobe/registry.py +313 -0
- agentprobe/report.py +666 -0
- agentprobe/runner.py +317 -0
- agentprobe/scenario.py +75 -0
- agentprobe/scenarios/__init__.py +0 -0
- agentprobe/scenarios/clean.py +194 -0
- agentprobe/scenarios/decoy.py +272 -0
- agentprobe/scenarios/registry.py +16 -0
- agentprobe/scenarios/split.py +203 -0
- agentprobe/scenarios/trap.py +215 -0
- agentprobe/termui.py +154 -0
- agentprobe/tools.py +275 -0
- agentprobe/trajectory.py +107 -0
- agentprobe/triage.py +153 -0
- agentprobe/validate_scenarios.py +489 -0
- agentprobe/world.py +189 -0
- agentprobe_testing-0.5.0.dist-info/METADATA +127 -0
- agentprobe_testing-0.5.0.dist-info/RECORD +55 -0
- agentprobe_testing-0.5.0.dist-info/WHEEL +5 -0
- agentprobe_testing-0.5.0.dist-info/entry_points.txt +4 -0
- agentprobe_testing-0.5.0.dist-info/licenses/LICENSE +109 -0
- 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
|
+
)
|