-
Notifications
You must be signed in to change notification settings - Fork 2
Expand file tree
/
Copy pathhardcoded.py
More file actions
303 lines (274 loc) · 11.1 KB
/
Copy pathhardcoded.py
File metadata and controls
303 lines (274 loc) · 11.1 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
import json
import random
from asyncio import sleep
from collections.abc import Callable
from typing import Any, TypedDict, cast, override
import inspect_ai._util.constants
from inspect_ai._util.retry import report_http_retry
from inspect_ai.log._samples import set_active_model_event_call
from inspect_ai.model import (
ChatCompletionChoice,
ChatMessage,
ChatMessageAssistant,
GenerateConfig,
ModelAPI,
ModelCall,
ModelOutput,
ModelUsage,
RetryDecision,
modelapi,
)
from inspect_ai.tool import ToolCall, ToolChoice, ToolInfo
class HardcodedToolCall(TypedDict):
tool_name: str
tool_args: dict[str, Any]
class HardcodedHTTPError(Exception):
"""Simulated HTTP error from the hardcoded provider (429 == rate limit)."""
status_code: int
def __init__(self, status_code: int) -> None:
super().__init__(f"hardcoded: simulated HTTP {status_code}")
self.status_code = status_code
class HardcodedModelAPI(ModelAPI):
tool_calls: list[HardcodedToolCall]
repetitions: int
answer: str
delay: float
concurrency: int
failure_rate: float
input_tokens: int
output_tokens: int
rate_limit_capacity: int | None
rate_limit_status: int
rate_limit_retry_after: float | None
rate_limit_swallowed_retries: int
retry_wait_seconds: float | None
in_flight: int
def __init__(
self,
model_name: str,
base_url: str | None = None,
api_key: str | None = None,
config: GenerateConfig | None = None,
tool_calls: list[HardcodedToolCall] | str | list[str] | None = None,
repetitions: int = 1,
answer: str = "done",
delay: float = 0.0,
concurrency: int = inspect_ai._util.constants.DEFAULT_MAX_CONNECTIONS,
failure_rate: float = 0.0,
input_tokens: int = 100,
output_tokens: int = 50,
rate_limit_capacity: int | None = None,
rate_limit_status: int = 429,
rate_limit_retry_after: float | None = None,
rate_limit_swallowed_retries: int = 0,
retry_wait_seconds: float | None = None,
):
super().__init__(
model_name=model_name,
base_url=base_url,
api_key=api_key,
config=config if config is not None else GenerateConfig(),
)
self.tool_calls = self._parse_tool_calls(tool_calls)
self.repetitions = repetitions
self.answer = answer
self.delay = delay
self.concurrency = concurrency
self.failure_rate = failure_rate
self.input_tokens = input_tokens
self.output_tokens = output_tokens
self.rate_limit_capacity = rate_limit_capacity
self.rate_limit_status = rate_limit_status
self.rate_limit_retry_after = rate_limit_retry_after
self.rate_limit_swallowed_retries = rate_limit_swallowed_retries
self.retry_wait_seconds = retry_wait_seconds
self.in_flight = 0
def _parse_tool_calls(
self, tool_calls: list[HardcodedToolCall] | str | list[str] | None
) -> list[HardcodedToolCall]:
if tool_calls is None:
return []
if isinstance(tool_calls, list) and len(tool_calls) == 0:
return []
items: list[Any]
if isinstance(tool_calls, str):
try:
decoded: Any = json.loads(tool_calls)
except json.JSONDecodeError:
decoded = tool_calls
items = cast(list[Any], decoded) if isinstance(decoded, list) else [decoded]
elif isinstance(tool_calls[0], str):
str_items: list[str] = [s for s in tool_calls if isinstance(s, str)]
try:
decoded = json.loads("[" + ",".join(str_items) + "]")
items = (
cast(list[Any], decoded)
if isinstance(decoded, list)
else list(str_items)
)
except json.JSONDecodeError:
items = list(str_items)
else:
items = list(tool_calls)
if len(items) == 0:
return []
if isinstance(items[0], str):
return [
HardcodedToolCall(tool_name="bash", tool_args={"command": cmd})
for cmd in items
if isinstance(cmd, str)
]
result: list[HardcodedToolCall] = []
for tool_call in items:
if not isinstance(tool_call, dict):
raise ValueError(f"Invalid tool call: {tool_call}")
tc = cast(dict[str, Any], tool_call)
if "tool_name" not in tc or "tool_args" not in tc:
raise ValueError(f"Invalid tool call: {tc}")
tool_name = tc["tool_name"]
tool_args = tc["tool_args"]
if not isinstance(tool_name, str):
raise ValueError(f"Invalid tool_name (must be str): {tc}")
if not isinstance(tool_args, dict):
raise ValueError(f"Invalid tool_args (must be dict): {tc}")
result.append(
HardcodedToolCall(
tool_name=tool_name,
tool_args=cast(dict[str, Any], tool_args),
)
)
return result
@override
def max_connections(self) -> int:
return self.concurrency
@override
async def generate(
self,
input: list[ChatMessage],
tools: list[ToolInfo],
tool_choice: ToolChoice,
config: GenerateConfig,
record_call: Callable[[ModelCall], None] | None = None,
) -> ModelOutput | tuple[ModelOutput | Exception, ModelCall]:
# No lock needed: single event loop, no await between increment and
# check. The finally matters because get_model() memoizes instances.
self.in_flight += 1
try:
return await self._generate(input, tools, record_call)
finally:
self.in_flight -= 1
async def _generate(
self,
input: list[ChatMessage],
tools: list[ToolInfo],
record_call: Callable[[ModelCall], None] | None,
) -> ModelOutput | tuple[ModelOutput | Exception, ModelCall]:
# Registered before anything that can fail, so a refused request still
# shows up in the transcript. This is what the first-party providers
# do (openai, anthropic, google, bedrock, mistral all call this helper
# ahead of the request); inspect-ai stamps the error onto it for us.
model_call = set_active_model_event_call(
request={"hardcoded": "test"}, filter=None
)
if record_call:
record_call(model_call)
# in_flight includes this call, so capacity=0 refuses everything and
# capacity>0 only bites while calls overlap (i.e. delay > 0). It must
# be raised, not returned: real providers let a 429 propagate out of
# generate(), and inspect-ai re-wraps a *returned* exception in a bare
# RuntimeError with no status_code that should_retry cannot classify.
if (
self.rate_limit_capacity is not None
and self.in_flight > self.rate_limit_capacity
):
# ...unless the "SDK" absorbs it. Real clients retry internally
# (openai/anthropic max_retries=2), so the 429 never escapes
# generate(), but every extra HTTP attempt is still reported and
# the adaptive controller still cuts. Must run on this task: the
# _request_had_retry it sets is a ContextVar, lost in a child task.
if not self.rate_limit_swallowed_retries:
raise HardcodedHTTPError(self.rate_limit_status)
for _ in range(self.rate_limit_swallowed_retries):
report_http_retry(
"rate_limit" if self.rate_limit_status == 429 else "transient",
self.rate_limit_retry_after,
)
index = sum(1 for m in input if m.role == "assistant")
next_tool_call_index = (
int(index) % len(self.tool_calls) if self.tool_calls else 0
)
repetition_count = int(index) // len(self.tool_calls) if self.tool_calls else 1
next_tool_call = (
self.tool_calls[next_tool_call_index]
if next_tool_call_index < len(self.tool_calls)
else None
)
if self.delay > 0:
await sleep(self.delay)
if random.random() < self.failure_rate:
model_call.response = {"failure": "test"}
try:
raise Exception("Failure")
except Exception as e:
return e, model_call
message: ChatMessageAssistant
if repetition_count >= self.repetitions:
submit_tool = next((tool for tool in tools if tool.name == "submit"), None)
if submit_tool is None:
message = ChatMessageAssistant(content=self.answer)
else:
message = ChatMessageAssistant(
content="I will now submit my answer.",
tool_calls=[
ToolCall(
id="hardcoded_submit",
function=submit_tool.name,
arguments={"answer": self.answer},
)
],
)
choice = ChatCompletionChoice(
message=message,
stop_reason="stop",
)
else:
assert next_tool_call is not None
tool_name = next_tool_call["tool_name"]
tool_args = next_tool_call["tool_args"]
message = ChatMessageAssistant(
content=f"Executing {tool_name} with args: {tool_args}",
tool_calls=[
ToolCall(
id=f"hardcoded_{index}",
function=tool_name,
arguments=tool_args,
)
],
)
choice = ChatCompletionChoice(message=message)
model_call.response = {"test": "hardcoded"}
return ModelOutput(
model="hardcoded",
choices=[choice],
usage=ModelUsage(
input_tokens=self.input_tokens,
output_tokens=self.output_tokens,
total_tokens=self.input_tokens + self.output_tokens,
),
), model_call
@override
def should_retry(self, ex: Exception) -> bool | RetryDecision:
# Classify on the status code, not the exception type: rate_limit_status
# then yields an identical failure that stays transient, and a future
# inspect-ai wrapper preserving status_code still classifies.
if getattr(ex, "status_code", None) == 429:
return RetryDecision.rate_limit(retry_after=self.rate_limit_retry_after)
return True
@override
def retry_wait(self) -> Callable[[Any], float] | None:
# tenacity accepts a bare callable as `wait`, so no tenacity import.
seconds = self.retry_wait_seconds
return None if seconds is None else lambda _: seconds
@modelapi(name="hardcoded")
def hardcoded() -> type[ModelAPI]:
return HardcodedModelAPI