millforge 0.1.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.
- millforge/__init__.py +1174 -0
- millforge/_forge/LICENSE +21 -0
- millforge/_forge/PROVENANCE.json +295 -0
- millforge/_forge/UPDATE_POLICY.md +24 -0
- millforge/_forge/__init__.py +14 -0
- millforge/_forge/adapter.py +2232 -0
- millforge/_forge/base_runner.py +121 -0
- millforge/_forge/clients/__init__.py +10 -0
- millforge/_forge/clients/base.py +200 -0
- millforge/_forge/context/__init__.py +23 -0
- millforge/_forge/context/manager.py +178 -0
- millforge/_forge/context/strategies.py +335 -0
- millforge/_forge/core/__init__.py +16 -0
- millforge/_forge/core/inference.py +433 -0
- millforge/_forge/core/messages.py +119 -0
- millforge/_forge/core/runner.py +479 -0
- millforge/_forge/core/steps.py +108 -0
- millforge/_forge/core/workflow.py +400 -0
- millforge/_forge/errors.py +222 -0
- millforge/_forge/guardrails/__init__.py +21 -0
- millforge/_forge/guardrails/error_tracker.py +71 -0
- millforge/_forge/guardrails/guardrails.py +194 -0
- millforge/_forge/guardrails/nudge.py +47 -0
- millforge/_forge/guardrails/response_validator.py +119 -0
- millforge/_forge/guardrails/step_enforcer.py +183 -0
- millforge/_forge/prompts/__init__.py +16 -0
- millforge/_forge/prompts/nudges.py +95 -0
- millforge/_forge/prompts/templates.py +285 -0
- millforge/_version.py +3 -0
- millforge/artifacts.py +570 -0
- millforge/base/__init__.py +97 -0
- millforge/base/composition.py +402 -0
- millforge/base/context.py +285 -0
- millforge/base/harness.py +138 -0
- millforge/base/identity.py +465 -0
- millforge/base/options.py +34 -0
- millforge/base/platform.py +17 -0
- millforge/base/prompt.py +317 -0
- millforge/base/runner.py +546 -0
- millforge/compiled_plan.py +970 -0
- millforge/compiler/__init__.py +231 -0
- millforge/compiler/artifact_validation.py +257 -0
- millforge/compiler/canonicalization.py +169 -0
- millforge/compiler/capabilities.py +66 -0
- millforge/compiler/catalogs.py +500 -0
- millforge/compiler/diagnostics.py +491 -0
- millforge/compiler/graph.py +678 -0
- millforge/compiler/lowering.py +198 -0
- millforge/compiler/output.py +692 -0
- millforge/compiler/parsing.py +1424 -0
- millforge/compiler/requests.py +1180 -0
- millforge/compiler/schema_validation.py +272 -0
- millforge/compiler/semantic.py +490 -0
- millforge/compiler/service.py +448 -0
- millforge/compiler/source.py +375 -0
- millforge/compiler/validators.py +184 -0
- millforge/connectors/__init__.py +95 -0
- millforge/connectors/admission.py +801 -0
- millforge/connectors/broker.py +202 -0
- millforge/connectors/contracts.py +1159 -0
- millforge/connectors/diagnostics.py +189 -0
- millforge/connectors/fake.py +66 -0
- millforge/connectors/runtime.py +236 -0
- millforge/contracts.py +2860 -0
- millforge/custom_tools/__init__.py +67 -0
- millforge/custom_tools/compiler.py +724 -0
- millforge/custom_tools/contracts.py +1093 -0
- millforge/custom_tools/diagnostics.py +205 -0
- millforge/eval_artifacts.py +952 -0
- millforge/eval_boundary.py +2435 -0
- millforge/eval_fixtures/__init__.py +1 -0
- millforge/eval_fixtures/default_pack/__init__.py +1 -0
- millforge/eval_fixtures/default_pack/fixtures/fixture.08a.bug_diagnosis.traceback.v1.json +52 -0
- millforge/eval_fixtures/default_pack/fixtures/fixture.08a.direct_edit.import_sort.v1.json +52 -0
- millforge/eval_fixtures/default_pack/fixtures/fixture.08a.evidence_discipline.no_source_change.v1.json +51 -0
- millforge/eval_fixtures/default_pack/fixtures/fixture.08a.false_closure.visible_green.v1.json +52 -0
- millforge/eval_fixtures/default_pack/fixtures/fixture.08a.multi_file.api_contract.v1.json +54 -0
- millforge/eval_fixtures/default_pack/fixtures/fixture.08a.recovery.malformed_artifact.v1.json +54 -0
- millforge/eval_fixtures/default_pack/manifest.json +12 -0
- millforge/eval_modes.py +1282 -0
- millforge/eval_presets.py +1398 -0
- millforge/eval_reports.py +2517 -0
- millforge/eval_suite.py +2429 -0
- millforge/eval_trials.py +2632 -0
- millforge/eval_workflow.py +794 -0
- millforge/exceptions.py +122 -0
- millforge/model_backend.py +2098 -0
- millforge/protocols.py +340 -0
- millforge/py.typed +0 -0
- millforge/runtime.py +1791 -0
- millforge/testing/__init__.py +1089 -0
- millforge/tools/__init__.py +83 -0
- millforge/tools/builtin_runtime.py +1339 -0
- millforge/tools/builtins.py +773 -0
- millforge/tools/execution.py +1545 -0
- millforge/tools/path_policy.py +155 -0
- millforge/tools/pi_compat/PI_LICENSE +21 -0
- millforge/tools/pi_compat/PROVENANCE.json +55 -0
- millforge/tools/pi_compat/UPDATE_POLICY.md +36 -0
- millforge/tools/pi_compat/__init__.py +34 -0
- millforge/tools/pi_compat/contracts.py +49 -0
- millforge/tools/pi_compat/editing.py +390 -0
- millforge/tools/pi_compat/mutations.py +57 -0
- millforge/tools/pi_compat/operations.py +401 -0
- millforge/tools/pi_compat/paths.py +155 -0
- millforge/tools/pi_compat/process.py +1375 -0
- millforge/tools/pi_compat/search.py +738 -0
- millforge/tools/pi_compat/truncation.py +267 -0
- millforge/tools/pi_compat_catalog.py +396 -0
- millforge/tools/pi_compat_runtime.py +460 -0
- millforge/tools/registry.py +553 -0
- millforge/tools/results.py +533 -0
- millforge-0.1.0.dist-info/METADATA +844 -0
- millforge-0.1.0.dist-info/RECORD +116 -0
- millforge-0.1.0.dist-info/WHEEL +4 -0
- millforge-0.1.0.dist-info/licenses/LICENSE +201 -0
|
@@ -0,0 +1,479 @@
|
|
|
1
|
+
"""WorkflowRunner — the agentic tool-calling loop."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import inspect
|
|
7
|
+
import json
|
|
8
|
+
from collections.abc import Awaitable, Callable
|
|
9
|
+
from typing import Any
|
|
10
|
+
|
|
11
|
+
from millforge._forge.clients.base import LLMClient, StreamChunk
|
|
12
|
+
from millforge._forge.context.manager import ContextManager
|
|
13
|
+
from millforge._forge.core.inference import (
|
|
14
|
+
_NUDGE_KIND_TO_TYPE,
|
|
15
|
+
_build_tool_call_infos,
|
|
16
|
+
run_inference,
|
|
17
|
+
)
|
|
18
|
+
from millforge._forge.core.messages import (
|
|
19
|
+
Message,
|
|
20
|
+
MessageMeta,
|
|
21
|
+
MessageRole,
|
|
22
|
+
MessageType,
|
|
23
|
+
)
|
|
24
|
+
from millforge._forge.core.workflow import TextResponse, ToolCall, Workflow
|
|
25
|
+
from millforge._forge.errors import (
|
|
26
|
+
MaxIterationsError,
|
|
27
|
+
NonRetryableToolError,
|
|
28
|
+
PrerequisiteError,
|
|
29
|
+
StepEnforcementError,
|
|
30
|
+
ToolExecutionError,
|
|
31
|
+
ToolResolutionError,
|
|
32
|
+
WorkflowCancelledError,
|
|
33
|
+
)
|
|
34
|
+
from millforge._forge.guardrails import ErrorTracker, ResponseValidator, StepEnforcer
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
class WorkflowRunner:
|
|
38
|
+
"""Executes a Workflow against an LLMClient with context management.
|
|
39
|
+
|
|
40
|
+
1. Builds the initial message list (system prompt + user input)
|
|
41
|
+
2. Sends messages to the LLM via the client (streaming or batch)
|
|
42
|
+
3. Inspects the response — if TextResponse (malformed/refusal), retries with nudge
|
|
43
|
+
4. Validates and executes returned tool calls (batch-aware)
|
|
44
|
+
5. Manages context budget via ContextManager
|
|
45
|
+
6. Enforces required steps via StepEnforcer
|
|
46
|
+
7. Terminates on terminal tool or max iterations
|
|
47
|
+
|
|
48
|
+
Retry logic lives here, not on the client.
|
|
49
|
+
"""
|
|
50
|
+
|
|
51
|
+
def __init__(
|
|
52
|
+
self,
|
|
53
|
+
client: LLMClient,
|
|
54
|
+
context_manager: ContextManager,
|
|
55
|
+
max_iterations: int = 10,
|
|
56
|
+
max_retries_per_step: int = 3,
|
|
57
|
+
max_tool_errors: int = 2,
|
|
58
|
+
stream: bool = False,
|
|
59
|
+
on_chunk: Callable[[StreamChunk], Awaitable[None]] | None = None,
|
|
60
|
+
on_message: Callable[[Message], None] | None = None,
|
|
61
|
+
rescue_enabled: bool = True,
|
|
62
|
+
retry_nudge: Callable[[str], str] | str | None = None,
|
|
63
|
+
max_premature_attempts: int = 3,
|
|
64
|
+
max_prereq_violations: int = 2,
|
|
65
|
+
tool_call_invoker: Callable[[ToolCall], Awaitable[Any] | Any] | None = None,
|
|
66
|
+
):
|
|
67
|
+
"""
|
|
68
|
+
Args:
|
|
69
|
+
client: The LLM client to send messages through.
|
|
70
|
+
context_manager: Manages context budget and triggers compaction.
|
|
71
|
+
max_iterations: Hard ceiling on total LLM round trips. Retries
|
|
72
|
+
consume iterations.
|
|
73
|
+
max_retries_per_step: Consecutive formatting failures before
|
|
74
|
+
raising ToolCallError. Resets on any valid ToolCall.
|
|
75
|
+
max_tool_errors: Consecutive tool execution errors before raising
|
|
76
|
+
ToolExecutionError. Errors are fed back to the model for
|
|
77
|
+
self-correction. Resets on successful execution.
|
|
78
|
+
stream: If True, uses send_stream(). Streaming is a side channel
|
|
79
|
+
— the runner still waits for the FINAL chunk before acting.
|
|
80
|
+
on_chunk: Async callback for each StreamChunk (awaited per chunk).
|
|
81
|
+
Ignored if stream=False.
|
|
82
|
+
on_message: Callback fired when a Message is appended to history.
|
|
83
|
+
Does not affect runner behavior.
|
|
84
|
+
rescue_enabled: If False, skip rescue_tool_call() — TextResponse
|
|
85
|
+
goes straight to retry nudge (or failure if retries=0).
|
|
86
|
+
retry_nudge: Custom nudge for bare text responses. Pass a string
|
|
87
|
+
for a static message, or a callable ``(raw_response) -> str``
|
|
88
|
+
for dynamic nudges. If None, uses the default.
|
|
89
|
+
max_premature_attempts: Premature terminal attempts allowed before
|
|
90
|
+
raising StepEnforcementError. Defaults to upstream behavior.
|
|
91
|
+
max_prereq_violations: Consecutive prerequisite violations allowed
|
|
92
|
+
before raising PrerequisiteError. Defaults to upstream behavior.
|
|
93
|
+
"""
|
|
94
|
+
self.client = client
|
|
95
|
+
self.context_manager = context_manager
|
|
96
|
+
self.max_iterations = max_iterations
|
|
97
|
+
self.max_retries_per_step = max_retries_per_step
|
|
98
|
+
self.max_tool_errors = max_tool_errors
|
|
99
|
+
self.stream = stream
|
|
100
|
+
self.on_chunk = on_chunk
|
|
101
|
+
self.on_message = on_message
|
|
102
|
+
self.rescue_enabled = rescue_enabled
|
|
103
|
+
self.max_premature_attempts = max_premature_attempts
|
|
104
|
+
self.max_prereq_violations = max_prereq_violations
|
|
105
|
+
self.tool_call_invoker = tool_call_invoker
|
|
106
|
+
if isinstance(retry_nudge, str):
|
|
107
|
+
self._retry_nudge_fn: Callable[[str], str] | None = (
|
|
108
|
+
lambda _raw, _msg=retry_nudge: _msg
|
|
109
|
+
)
|
|
110
|
+
else:
|
|
111
|
+
self._retry_nudge_fn = retry_nudge
|
|
112
|
+
|
|
113
|
+
async def run(
|
|
114
|
+
self,
|
|
115
|
+
workflow: Workflow,
|
|
116
|
+
user_message: str,
|
|
117
|
+
prompt_vars: dict[str, str] | None = None,
|
|
118
|
+
initial_messages: list[Message] | None = None,
|
|
119
|
+
cancel_event: asyncio.Event | None = None,
|
|
120
|
+
) -> Any:
|
|
121
|
+
"""Execute the workflow and return the terminal tool's result.
|
|
122
|
+
|
|
123
|
+
Args:
|
|
124
|
+
workflow: The workflow to execute.
|
|
125
|
+
user_message: The user's input message.
|
|
126
|
+
prompt_vars: Variables for the system prompt template.
|
|
127
|
+
initial_messages: If provided, seeds the conversation with these
|
|
128
|
+
messages instead of building a fresh system prompt + user
|
|
129
|
+
input. The on_message callback fires only for NEW messages
|
|
130
|
+
created during this run, not the replayed history. The caller
|
|
131
|
+
must include the system prompt and new user message in the
|
|
132
|
+
seed.
|
|
133
|
+
cancel_event: If provided and set, the runner will raise
|
|
134
|
+
WorkflowCancelledError at the start of the next iteration.
|
|
135
|
+
Checked once per loop, before the inference call.
|
|
136
|
+
|
|
137
|
+
Raises:
|
|
138
|
+
MaxIterationsError: If max_iterations exceeded without terminal tool.
|
|
139
|
+
ToolCallError: If max_retries_per_step exhausted on a single step.
|
|
140
|
+
ToolExecutionError: If a tool callable raised and the model failed
|
|
141
|
+
to self-correct after max_tool_errors consecutive attempts.
|
|
142
|
+
WorkflowCancelledError: If cancel_event was set during execution.
|
|
143
|
+
"""
|
|
144
|
+
# Step 1 — Build initial messages
|
|
145
|
+
if initial_messages is not None:
|
|
146
|
+
messages: list[Message] = list(initial_messages)
|
|
147
|
+
|
|
148
|
+
def _emit(msg: Message) -> None:
|
|
149
|
+
messages.append(msg)
|
|
150
|
+
if self.on_message is not None:
|
|
151
|
+
self.on_message(msg)
|
|
152
|
+
else:
|
|
153
|
+
rendered_prompt = workflow.build_system_prompt(**(prompt_vars or {}))
|
|
154
|
+
messages: list[Message] = []
|
|
155
|
+
|
|
156
|
+
def _emit(msg: Message) -> None:
|
|
157
|
+
messages.append(msg)
|
|
158
|
+
if self.on_message is not None:
|
|
159
|
+
self.on_message(msg)
|
|
160
|
+
|
|
161
|
+
_emit(
|
|
162
|
+
Message(
|
|
163
|
+
MessageRole.SYSTEM,
|
|
164
|
+
rendered_prompt,
|
|
165
|
+
MessageMeta(MessageType.SYSTEM_PROMPT),
|
|
166
|
+
)
|
|
167
|
+
)
|
|
168
|
+
_emit(
|
|
169
|
+
Message(
|
|
170
|
+
MessageRole.USER, user_message, MessageMeta(MessageType.USER_INPUT)
|
|
171
|
+
)
|
|
172
|
+
)
|
|
173
|
+
|
|
174
|
+
# Step 2 — Initialize guardrail middleware
|
|
175
|
+
tool_names = list(workflow.tools.keys())
|
|
176
|
+
validator = ResponseValidator(
|
|
177
|
+
tool_names,
|
|
178
|
+
rescue_enabled=self.rescue_enabled,
|
|
179
|
+
retry_nudge_fn=self._retry_nudge_fn,
|
|
180
|
+
)
|
|
181
|
+
tool_prerequisites = {
|
|
182
|
+
name: td.prerequisites
|
|
183
|
+
for name, td in workflow.tools.items()
|
|
184
|
+
if td.prerequisites
|
|
185
|
+
}
|
|
186
|
+
step_enforcer = StepEnforcer(
|
|
187
|
+
required_steps=workflow.required_steps,
|
|
188
|
+
terminal_tools=workflow.terminal_tools,
|
|
189
|
+
tool_prerequisites=tool_prerequisites,
|
|
190
|
+
max_premature_attempts=self.max_premature_attempts,
|
|
191
|
+
max_prereq_violations=self.max_prereq_violations,
|
|
192
|
+
)
|
|
193
|
+
error_tracker = ErrorTracker(
|
|
194
|
+
max_retries=self.max_retries_per_step,
|
|
195
|
+
max_tool_errors=self.max_tool_errors,
|
|
196
|
+
)
|
|
197
|
+
|
|
198
|
+
# Step 3 — Main loop (one LLM call per iteration, retries consume iterations)
|
|
199
|
+
tool_specs = workflow.get_tool_specs()
|
|
200
|
+
tool_call_counter = 0
|
|
201
|
+
iteration = 0
|
|
202
|
+
|
|
203
|
+
while iteration < self.max_iterations:
|
|
204
|
+
# 3.0 — Check for cancellation
|
|
205
|
+
if cancel_event is not None and cancel_event.is_set():
|
|
206
|
+
raise WorkflowCancelledError(
|
|
207
|
+
messages=messages,
|
|
208
|
+
completed_steps=step_enforcer.completed_steps,
|
|
209
|
+
iteration=iteration,
|
|
210
|
+
)
|
|
211
|
+
|
|
212
|
+
# 3a — Inference: compact, fold, serialize, send, validate, retry
|
|
213
|
+
result = await run_inference(
|
|
214
|
+
messages=messages,
|
|
215
|
+
client=self.client,
|
|
216
|
+
context_manager=self.context_manager,
|
|
217
|
+
validator=validator,
|
|
218
|
+
error_tracker=error_tracker,
|
|
219
|
+
tool_specs=tool_specs,
|
|
220
|
+
tool_call_counter=tool_call_counter,
|
|
221
|
+
step_index=iteration,
|
|
222
|
+
step_hint=step_enforcer.summary_hint(),
|
|
223
|
+
max_attempts=self.max_iterations - iteration,
|
|
224
|
+
stream=self.stream,
|
|
225
|
+
on_chunk=self.on_chunk,
|
|
226
|
+
)
|
|
227
|
+
# max_attempts exhausted — iteration budget spent
|
|
228
|
+
if result is None:
|
|
229
|
+
break
|
|
230
|
+
# Retries consume iterations (preserves pre-extraction semantics)
|
|
231
|
+
iteration += result.attempts
|
|
232
|
+
# Emit new messages from retries (assistant text, nudges)
|
|
233
|
+
for msg in result.new_messages:
|
|
234
|
+
if self.on_message is not None:
|
|
235
|
+
self.on_message(msg)
|
|
236
|
+
tool_call_counter = result.tool_call_counter
|
|
237
|
+
|
|
238
|
+
# Intentional text response — emit and continue the loop.
|
|
239
|
+
# The model chose text over tools; consume an iteration.
|
|
240
|
+
if isinstance(result.response, TextResponse):
|
|
241
|
+
_emit(
|
|
242
|
+
Message(
|
|
243
|
+
MessageRole.ASSISTANT,
|
|
244
|
+
result.response.content,
|
|
245
|
+
MessageMeta(MessageType.TEXT_RESPONSE, step_index=iteration),
|
|
246
|
+
)
|
|
247
|
+
)
|
|
248
|
+
continue
|
|
249
|
+
|
|
250
|
+
tool_calls = result.response
|
|
251
|
+
|
|
252
|
+
# 3b — Check for premature terminal
|
|
253
|
+
step_check = step_enforcer.check(tool_calls)
|
|
254
|
+
|
|
255
|
+
if step_check.needs_nudge:
|
|
256
|
+
if step_enforcer.premature_exhausted:
|
|
257
|
+
attempted = next(
|
|
258
|
+
tc.tool
|
|
259
|
+
for tc in tool_calls
|
|
260
|
+
if tc.tool in workflow.terminal_tools
|
|
261
|
+
)
|
|
262
|
+
raise StepEnforcementError(
|
|
263
|
+
terminal_tool=attempted,
|
|
264
|
+
attempts=step_enforcer.premature_attempts,
|
|
265
|
+
pending_steps=step_enforcer.pending(),
|
|
266
|
+
)
|
|
267
|
+
if tool_calls[0].reasoning:
|
|
268
|
+
_emit(
|
|
269
|
+
Message(
|
|
270
|
+
MessageRole.ASSISTANT,
|
|
271
|
+
tool_calls[0].reasoning,
|
|
272
|
+
MessageMeta(MessageType.REASONING, step_index=iteration),
|
|
273
|
+
)
|
|
274
|
+
)
|
|
275
|
+
tc_infos, tool_call_counter = _build_tool_call_infos(
|
|
276
|
+
tool_calls, tool_call_counter
|
|
277
|
+
)
|
|
278
|
+
_emit(
|
|
279
|
+
Message(
|
|
280
|
+
MessageRole.ASSISTANT,
|
|
281
|
+
"",
|
|
282
|
+
MessageMeta(MessageType.TOOL_CALL, step_index=iteration),
|
|
283
|
+
tool_calls=tc_infos,
|
|
284
|
+
reasoning_content=tool_calls[0].reasoning_content,
|
|
285
|
+
)
|
|
286
|
+
)
|
|
287
|
+
nudge = step_check.nudge
|
|
288
|
+
nudge_type = _NUDGE_KIND_TO_TYPE[nudge.kind]
|
|
289
|
+
# Surface premature-terminal violation as a tool error result.
|
|
290
|
+
# See prereq path below for rationale.
|
|
291
|
+
for tc_info in tc_infos:
|
|
292
|
+
_emit(
|
|
293
|
+
Message(
|
|
294
|
+
MessageRole.TOOL,
|
|
295
|
+
f"[StepEnforcementError] {nudge.content}",
|
|
296
|
+
MessageMeta(nudge_type, step_index=iteration),
|
|
297
|
+
tool_name=tc_info.name,
|
|
298
|
+
tool_call_id=tc_info.call_id,
|
|
299
|
+
)
|
|
300
|
+
)
|
|
301
|
+
continue
|
|
302
|
+
|
|
303
|
+
# 3b.2 — Check prerequisites
|
|
304
|
+
prereq_check = step_enforcer.check_prerequisites(tool_calls)
|
|
305
|
+
|
|
306
|
+
if prereq_check.needs_nudge:
|
|
307
|
+
if step_enforcer.prereq_exhausted:
|
|
308
|
+
# Find the first violating tool for the error
|
|
309
|
+
for tc in tool_calls:
|
|
310
|
+
prereqs = tool_prerequisites.get(tc.tool)
|
|
311
|
+
if prereqs:
|
|
312
|
+
result = step_enforcer._tracker.check_prerequisites(
|
|
313
|
+
tc.tool,
|
|
314
|
+
tc.args,
|
|
315
|
+
prereqs,
|
|
316
|
+
)
|
|
317
|
+
if not result.satisfied:
|
|
318
|
+
raise PrerequisiteError(
|
|
319
|
+
tool_name=tc.tool,
|
|
320
|
+
violations=step_enforcer.prereq_violations,
|
|
321
|
+
missing_prereqs=result.missing,
|
|
322
|
+
)
|
|
323
|
+
if tool_calls[0].reasoning:
|
|
324
|
+
_emit(
|
|
325
|
+
Message(
|
|
326
|
+
MessageRole.ASSISTANT,
|
|
327
|
+
tool_calls[0].reasoning,
|
|
328
|
+
MessageMeta(MessageType.REASONING, step_index=iteration),
|
|
329
|
+
)
|
|
330
|
+
)
|
|
331
|
+
tc_infos, tool_call_counter = _build_tool_call_infos(
|
|
332
|
+
tool_calls, tool_call_counter
|
|
333
|
+
)
|
|
334
|
+
_emit(
|
|
335
|
+
Message(
|
|
336
|
+
MessageRole.ASSISTANT,
|
|
337
|
+
"",
|
|
338
|
+
MessageMeta(MessageType.TOOL_CALL, step_index=iteration),
|
|
339
|
+
tool_calls=tc_infos,
|
|
340
|
+
reasoning_content=tool_calls[0].reasoning_content,
|
|
341
|
+
)
|
|
342
|
+
)
|
|
343
|
+
nudge = prereq_check.nudge
|
|
344
|
+
nudge_type = _NUDGE_KIND_TO_TYPE[nudge.kind]
|
|
345
|
+
# Surface the prereq violation as a tool error result rather
|
|
346
|
+
# than a trailing user nudge. Models are pretrained on the
|
|
347
|
+
# "tool failed → try something else" shape; the user-nudge
|
|
348
|
+
# shape was getting muddied by _merge_consecutive folding it
|
|
349
|
+
# into the original user message, hiding the correction signal.
|
|
350
|
+
# Pair one tool-error result with each tool_call in the batch
|
|
351
|
+
# so the message structure stays consistent.
|
|
352
|
+
for tc_info in tc_infos:
|
|
353
|
+
_emit(
|
|
354
|
+
Message(
|
|
355
|
+
MessageRole.TOOL,
|
|
356
|
+
f"[PrerequisiteError] {nudge.content}",
|
|
357
|
+
MessageMeta(nudge_type, step_index=iteration),
|
|
358
|
+
tool_name=tc_info.name,
|
|
359
|
+
tool_call_id=tc_info.call_id,
|
|
360
|
+
)
|
|
361
|
+
)
|
|
362
|
+
continue
|
|
363
|
+
|
|
364
|
+
# 3c — Execute all tool calls in the batch
|
|
365
|
+
tc_infos, tool_call_counter = _build_tool_call_infos(
|
|
366
|
+
tool_calls, tool_call_counter
|
|
367
|
+
)
|
|
368
|
+
call_ids = [tc.call_id for tc in tc_infos]
|
|
369
|
+
|
|
370
|
+
# Emit reasoning (from first call) and assistant message
|
|
371
|
+
if tool_calls[0].reasoning:
|
|
372
|
+
_emit(
|
|
373
|
+
Message(
|
|
374
|
+
MessageRole.ASSISTANT,
|
|
375
|
+
tool_calls[0].reasoning,
|
|
376
|
+
MessageMeta(MessageType.REASONING, step_index=iteration),
|
|
377
|
+
)
|
|
378
|
+
)
|
|
379
|
+
_emit(
|
|
380
|
+
Message(
|
|
381
|
+
MessageRole.ASSISTANT,
|
|
382
|
+
"",
|
|
383
|
+
MessageMeta(MessageType.TOOL_CALL, step_index=iteration),
|
|
384
|
+
tool_calls=tc_infos,
|
|
385
|
+
reasoning_content=tool_calls[0].reasoning_content,
|
|
386
|
+
)
|
|
387
|
+
)
|
|
388
|
+
|
|
389
|
+
# Execute each tool and emit results
|
|
390
|
+
batch_had_error = False
|
|
391
|
+
last_error: tuple[str, Exception] | None = None
|
|
392
|
+
terminal_result = None
|
|
393
|
+
for i, tc in enumerate(tool_calls):
|
|
394
|
+
tc_id = call_ids[i]
|
|
395
|
+
fn = workflow.get_callable(tc.tool)
|
|
396
|
+
try:
|
|
397
|
+
if self.tool_call_invoker is not None:
|
|
398
|
+
result_val = self.tool_call_invoker(tc)
|
|
399
|
+
if inspect.isawaitable(result_val):
|
|
400
|
+
result_val = await result_val
|
|
401
|
+
elif inspect.iscoroutinefunction(fn):
|
|
402
|
+
result_val = await fn(**tc.args)
|
|
403
|
+
else:
|
|
404
|
+
result_val = fn(**tc.args)
|
|
405
|
+
except ToolResolutionError as exc:
|
|
406
|
+
_emit(
|
|
407
|
+
Message(
|
|
408
|
+
MessageRole.TOOL,
|
|
409
|
+
f"[ToolResolutionError] {exc}",
|
|
410
|
+
MessageMeta(MessageType.TOOL_RESULT, step_index=iteration),
|
|
411
|
+
tool_name=tc.tool,
|
|
412
|
+
tool_call_id=tc_id,
|
|
413
|
+
)
|
|
414
|
+
)
|
|
415
|
+
if tc.tool in workflow.terminal_tools:
|
|
416
|
+
terminal_result = exc
|
|
417
|
+
continue
|
|
418
|
+
except NonRetryableToolError:
|
|
419
|
+
raise
|
|
420
|
+
except Exception as exc:
|
|
421
|
+
batch_had_error = True
|
|
422
|
+
last_error = (tc.tool, exc)
|
|
423
|
+
_emit(
|
|
424
|
+
Message(
|
|
425
|
+
MessageRole.TOOL,
|
|
426
|
+
f"[ToolError] {type(exc).__name__}: {exc}",
|
|
427
|
+
MessageMeta(MessageType.TOOL_RESULT, step_index=iteration),
|
|
428
|
+
tool_name=tc.tool,
|
|
429
|
+
tool_call_id=tc_id,
|
|
430
|
+
)
|
|
431
|
+
)
|
|
432
|
+
if tc.tool in workflow.terminal_tools:
|
|
433
|
+
terminal_result = exc
|
|
434
|
+
continue
|
|
435
|
+
|
|
436
|
+
# Success
|
|
437
|
+
step_enforcer.record(tc.tool, tc.args)
|
|
438
|
+
result_str = (
|
|
439
|
+
result_val
|
|
440
|
+
if isinstance(result_val, str)
|
|
441
|
+
else json.dumps(result_val)
|
|
442
|
+
)
|
|
443
|
+
_emit(
|
|
444
|
+
Message(
|
|
445
|
+
MessageRole.TOOL,
|
|
446
|
+
result_str,
|
|
447
|
+
MessageMeta(MessageType.TOOL_RESULT, step_index=iteration),
|
|
448
|
+
tool_name=tc.tool,
|
|
449
|
+
tool_call_id=tc_id,
|
|
450
|
+
)
|
|
451
|
+
)
|
|
452
|
+
|
|
453
|
+
if tc.tool in workflow.terminal_tools:
|
|
454
|
+
terminal_result = result_val
|
|
455
|
+
|
|
456
|
+
# 3d — Post-batch bookkeeping
|
|
457
|
+
if batch_had_error:
|
|
458
|
+
error_tracker.record_result(success=False)
|
|
459
|
+
if error_tracker.tool_errors_exhausted:
|
|
460
|
+
assert last_error is not None
|
|
461
|
+
raise ToolExecutionError(
|
|
462
|
+
last_error[0],
|
|
463
|
+
cause=last_error[1],
|
|
464
|
+
)
|
|
465
|
+
else:
|
|
466
|
+
error_tracker.reset_errors()
|
|
467
|
+
step_enforcer.reset_premature()
|
|
468
|
+
step_enforcer.reset_prereq_violations()
|
|
469
|
+
|
|
470
|
+
# 3e — If terminal tool was in the batch and succeeded, return
|
|
471
|
+
if terminal_result is not None and not isinstance(
|
|
472
|
+
terminal_result, Exception
|
|
473
|
+
):
|
|
474
|
+
return terminal_result
|
|
475
|
+
|
|
476
|
+
# Step 4 — Max iterations exceeded
|
|
477
|
+
raise MaxIterationsError(
|
|
478
|
+
self.max_iterations, step_enforcer.completed_steps, step_enforcer.pending()
|
|
479
|
+
)
|
|
@@ -0,0 +1,108 @@
|
|
|
1
|
+
"""Required-step tracking and prerequisite enforcement."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from dataclasses import dataclass, field
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
@dataclass
|
|
10
|
+
class PrerequisiteCheck:
|
|
11
|
+
"""Result of checking prerequisites for a tool call.
|
|
12
|
+
|
|
13
|
+
If ``satisfied`` is False, ``missing`` lists the prerequisite tool names
|
|
14
|
+
that have not been called (or not called with matching args).
|
|
15
|
+
"""
|
|
16
|
+
|
|
17
|
+
satisfied: bool
|
|
18
|
+
missing: list[str]
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
@dataclass
|
|
22
|
+
class StepTracker:
|
|
23
|
+
"""Tracks which required steps have been completed and which tools
|
|
24
|
+
have been executed (with args) for prerequisite enforcement.
|
|
25
|
+
|
|
26
|
+
Lives on the WorkflowRunner, outside the message history.
|
|
27
|
+
Compaction cannot invalidate step completion. See P0-1.
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
required_steps: list[str]
|
|
31
|
+
completed_steps: dict[str, None] = field(default_factory=dict)
|
|
32
|
+
executed_tools: dict[str, list[dict[str, Any]]] = field(default_factory=dict)
|
|
33
|
+
|
|
34
|
+
def record(self, tool_name: str, args: dict[str, Any] | None = None) -> None:
|
|
35
|
+
"""Record a successful tool execution.
|
|
36
|
+
|
|
37
|
+
Args:
|
|
38
|
+
tool_name: The tool that was executed.
|
|
39
|
+
args: The arguments the tool was called with. Stored for
|
|
40
|
+
arg-matched prerequisite checking.
|
|
41
|
+
"""
|
|
42
|
+
self.completed_steps[tool_name] = None
|
|
43
|
+
self.executed_tools.setdefault(tool_name, []).append(args or {})
|
|
44
|
+
|
|
45
|
+
def is_satisfied(self) -> bool:
|
|
46
|
+
"""True if all required steps have been called."""
|
|
47
|
+
return all(s in self.completed_steps for s in self.required_steps)
|
|
48
|
+
|
|
49
|
+
def pending(self) -> list[str]:
|
|
50
|
+
"""Return required steps not yet completed, preserving original order."""
|
|
51
|
+
return [s for s in self.required_steps if s not in self.completed_steps]
|
|
52
|
+
|
|
53
|
+
def check_prerequisites(
|
|
54
|
+
self,
|
|
55
|
+
tool_name: str,
|
|
56
|
+
args: dict[str, Any],
|
|
57
|
+
prerequisites: list[str | dict[str, str]],
|
|
58
|
+
) -> PrerequisiteCheck:
|
|
59
|
+
"""Check whether prerequisites are satisfied for a tool call.
|
|
60
|
+
|
|
61
|
+
Args:
|
|
62
|
+
tool_name: The tool about to be called (for error context).
|
|
63
|
+
args: The arguments the tool is being called with.
|
|
64
|
+
prerequisites: The prerequisite definitions from ToolDef.
|
|
65
|
+
|
|
66
|
+
Returns:
|
|
67
|
+
PrerequisiteCheck with satisfied=True if all prereqs are met,
|
|
68
|
+
or satisfied=False with the list of unsatisfied prereq tool names.
|
|
69
|
+
"""
|
|
70
|
+
missing: list[str] = []
|
|
71
|
+
for prereq in prerequisites:
|
|
72
|
+
if isinstance(prereq, str):
|
|
73
|
+
# Name-only: any prior call to this tool satisfies it
|
|
74
|
+
if prereq not in self.executed_tools:
|
|
75
|
+
missing.append(prereq)
|
|
76
|
+
elif isinstance(prereq, dict):
|
|
77
|
+
# Arg-matched: a prior call with the configured prior/current
|
|
78
|
+
# argument values satisfies it. ``match_arg`` is the legacy
|
|
79
|
+
# same-name shorthand retained for private Forge compatibility.
|
|
80
|
+
prereq_tool = prereq["tool"]
|
|
81
|
+
prerequisite_arg = prereq.get(
|
|
82
|
+
"prerequisite_arg", prereq.get("match_arg", "")
|
|
83
|
+
)
|
|
84
|
+
current_arg = prereq.get("current_arg", prereq.get("match_arg", ""))
|
|
85
|
+
# Defensive: malformed (non-dict) args can't satisfy an
|
|
86
|
+
# arg-match. ResponseValidator rejects them before dispatch, but
|
|
87
|
+
# a granular caller may reach here directly — treat as
|
|
88
|
+
# unsatisfied rather than crashing on ``.get``.
|
|
89
|
+
if not isinstance(args, dict):
|
|
90
|
+
missing.append(prereq_tool)
|
|
91
|
+
continue
|
|
92
|
+
required_value = args.get(current_arg)
|
|
93
|
+
if prereq_tool not in self.executed_tools:
|
|
94
|
+
missing.append(prereq_tool)
|
|
95
|
+
continue
|
|
96
|
+
if not any(
|
|
97
|
+
call.get(prerequisite_arg) == required_value
|
|
98
|
+
for call in self.executed_tools[prereq_tool]
|
|
99
|
+
):
|
|
100
|
+
missing.append(prereq_tool)
|
|
101
|
+
|
|
102
|
+
return PrerequisiteCheck(satisfied=len(missing) == 0, missing=missing)
|
|
103
|
+
|
|
104
|
+
def summary_hint(self) -> str:
|
|
105
|
+
"""Human-readable hint for injection into compacted summaries."""
|
|
106
|
+
if not self.completed_steps:
|
|
107
|
+
return "[No steps completed yet]"
|
|
108
|
+
return f"[Steps completed: {', '.join(self.completed_steps)}]"
|