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,433 @@
|
|
|
1
|
+
"""Inference loop — compact, fold, serialize, send, validate, retry.
|
|
2
|
+
|
|
3
|
+
Extracted from WorkflowRunner so both the runner and the proxy can share
|
|
4
|
+
the same input-processing and validation logic. This is the "front half"
|
|
5
|
+
of the agentic loop: everything up to and including getting a clean
|
|
6
|
+
response from the LLM. The "back half" (step enforcement, tool execution,
|
|
7
|
+
terminal check) stays in the caller.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
from collections.abc import Awaitable, Callable
|
|
13
|
+
from dataclasses import dataclass, field
|
|
14
|
+
from typing import Any
|
|
15
|
+
|
|
16
|
+
from millforge._forge.clients.base import (
|
|
17
|
+
ChunkType,
|
|
18
|
+
LLMClient,
|
|
19
|
+
RawOpenAIMessages,
|
|
20
|
+
RawOpenAITools,
|
|
21
|
+
StreamChunk,
|
|
22
|
+
TokenUsage,
|
|
23
|
+
)
|
|
24
|
+
from millforge._forge.context.manager import ContextManager
|
|
25
|
+
from millforge._forge.core.messages import (
|
|
26
|
+
Message,
|
|
27
|
+
MessageMeta,
|
|
28
|
+
MessageRole,
|
|
29
|
+
MessageType,
|
|
30
|
+
ToolCallInfo,
|
|
31
|
+
)
|
|
32
|
+
from millforge._forge.core.workflow import LLMResponse, TextResponse, ToolCall, ToolSpec
|
|
33
|
+
from millforge._forge.errors import StreamError, ToolCallError
|
|
34
|
+
from millforge._forge.guardrails import ErrorTracker, ResponseValidator
|
|
35
|
+
from millforge._forge.guardrails.nudge import TOOL_ERROR_KINDS
|
|
36
|
+
|
|
37
|
+
# Maps Nudge.kind → MessageType for message emission.
|
|
38
|
+
_NUDGE_KIND_TO_TYPE: dict[str, MessageType] = {
|
|
39
|
+
"retry": MessageType.RETRY_NUDGE,
|
|
40
|
+
"unknown_tool": MessageType.RETRY_NUDGE,
|
|
41
|
+
"tool_arg_validation": MessageType.RETRY_NUDGE,
|
|
42
|
+
"step": MessageType.STEP_NUDGE,
|
|
43
|
+
"prerequisite": MessageType.PREREQUISITE_NUDGE,
|
|
44
|
+
}
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def _get_usage(client: LLMClient) -> TokenUsage | None:
|
|
48
|
+
"""Extract actual token count from the client."""
|
|
49
|
+
last_usage = getattr(client, "last_usage", None)
|
|
50
|
+
if not isinstance(last_usage, dict):
|
|
51
|
+
return None
|
|
52
|
+
slot_id = getattr(client, "_slot_id", None) or 0
|
|
53
|
+
return last_usage.get(slot_id)
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
def _sync_token_count(client: LLMClient, context_manager: ContextManager) -> None:
|
|
57
|
+
"""Feed actual token count from the client into the context manager."""
|
|
58
|
+
usage = _get_usage(client)
|
|
59
|
+
if usage is not None:
|
|
60
|
+
context_manager.update_token_count(usage.total_tokens)
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
@dataclass
|
|
64
|
+
class InferenceResult:
|
|
65
|
+
"""Result of a single inference call (may include transparent retries).
|
|
66
|
+
|
|
67
|
+
Attributes:
|
|
68
|
+
response: The validated LLM response — tool calls or text.
|
|
69
|
+
new_messages: Messages generated during this call (assistant text from
|
|
70
|
+
failed attempts, nudges, and the final assistant response). The
|
|
71
|
+
caller should append these to their message history.
|
|
72
|
+
usage: Token usage for the final successful attempt.
|
|
73
|
+
tool_call_counter: Updated counter for generating unique call IDs.
|
|
74
|
+
"""
|
|
75
|
+
|
|
76
|
+
response: list[ToolCall] | TextResponse
|
|
77
|
+
new_messages: list[Message] = field(default_factory=list)
|
|
78
|
+
usage: TokenUsage | None = None
|
|
79
|
+
tool_call_counter: int = 0
|
|
80
|
+
attempts: int = 1
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
def fold_and_serialize(
|
|
84
|
+
messages: list[Message],
|
|
85
|
+
api_format: str,
|
|
86
|
+
) -> list[dict[str, Any]]:
|
|
87
|
+
"""Reasoning-fold and serialize forge Messages to API dicts.
|
|
88
|
+
|
|
89
|
+
Folds REASONING messages into the following TOOL_CALL message's content
|
|
90
|
+
field so the wire format has one assistant message with both content and
|
|
91
|
+
tool_calls (valid OpenAI format). Internal Message list stays separate
|
|
92
|
+
for compaction.
|
|
93
|
+
"""
|
|
94
|
+
api_messages: list[dict[str, Any]] = []
|
|
95
|
+
pending_reasoning: str | None = None
|
|
96
|
+
|
|
97
|
+
for m in messages:
|
|
98
|
+
if m.metadata.type == MessageType.REASONING and m.role == MessageRole.ASSISTANT:
|
|
99
|
+
pending_reasoning = m.content
|
|
100
|
+
continue
|
|
101
|
+
d = m.to_api_dict(format=api_format)
|
|
102
|
+
if pending_reasoning is not None and m.tool_calls is not None:
|
|
103
|
+
d["content"] = pending_reasoning
|
|
104
|
+
pending_reasoning = None
|
|
105
|
+
elif pending_reasoning is not None:
|
|
106
|
+
api_messages.append({"role": "assistant", "content": pending_reasoning})
|
|
107
|
+
pending_reasoning = None
|
|
108
|
+
api_messages.append(d)
|
|
109
|
+
|
|
110
|
+
if pending_reasoning is not None:
|
|
111
|
+
api_messages.append({"role": "assistant", "content": pending_reasoning})
|
|
112
|
+
|
|
113
|
+
return api_messages
|
|
114
|
+
|
|
115
|
+
|
|
116
|
+
def _build_tool_call_infos(
|
|
117
|
+
tool_calls: list[ToolCall],
|
|
118
|
+
tool_call_counter: int,
|
|
119
|
+
) -> tuple[list[ToolCallInfo], int]:
|
|
120
|
+
"""Retain provider call IDs and assign IDs only to provider-less calls."""
|
|
121
|
+
tc_infos = []
|
|
122
|
+
for tc in tool_calls:
|
|
123
|
+
tc_id = tc.call_id
|
|
124
|
+
if tc_id is None:
|
|
125
|
+
tc_id = f"call_{tool_call_counter:09d}"
|
|
126
|
+
tool_call_counter += 1
|
|
127
|
+
tc_infos.append(ToolCallInfo(name=tc.tool, args=tc.args, call_id=tc_id))
|
|
128
|
+
return tc_infos, tool_call_counter
|
|
129
|
+
|
|
130
|
+
|
|
131
|
+
async def run_inference(
|
|
132
|
+
messages: list[Message],
|
|
133
|
+
client: LLMClient,
|
|
134
|
+
context_manager: ContextManager,
|
|
135
|
+
validator: ResponseValidator,
|
|
136
|
+
error_tracker: ErrorTracker,
|
|
137
|
+
tool_specs: list[ToolSpec],
|
|
138
|
+
tool_call_counter: int = 0,
|
|
139
|
+
step_index: int = 0,
|
|
140
|
+
step_hint: str = "",
|
|
141
|
+
max_attempts: int | None = None,
|
|
142
|
+
stream: bool = False,
|
|
143
|
+
on_chunk: Callable[[StreamChunk], Awaitable[None]] | None = None,
|
|
144
|
+
sampling: dict[str, Any] | None = None,
|
|
145
|
+
passthrough: dict[str, Any] | None = None,
|
|
146
|
+
inbound_anthropic_body: dict[str, Any] | None = None,
|
|
147
|
+
raw_openai_messages: RawOpenAIMessages | None = None,
|
|
148
|
+
raw_openai_tools: RawOpenAITools | None = None,
|
|
149
|
+
) -> InferenceResult | None:
|
|
150
|
+
"""Send messages to the LLM with compaction, folding, validation, and retry.
|
|
151
|
+
|
|
152
|
+
Retries are handled internally — the caller gets a clean response or an
|
|
153
|
+
exception. On each retry, the model's failed output and a corrective
|
|
154
|
+
nudge are appended to the message list and returned in
|
|
155
|
+
``InferenceResult.new_messages`` so the caller can track history.
|
|
156
|
+
|
|
157
|
+
Args:
|
|
158
|
+
messages: The current conversation history (forge Messages). This
|
|
159
|
+
list is mutated during compaction and retry (messages may be
|
|
160
|
+
removed by compaction and added by retries).
|
|
161
|
+
client: The LLM backend client.
|
|
162
|
+
context_manager: For context budget compaction.
|
|
163
|
+
validator: For rescue parsing, retry nudges, unknown tool checks.
|
|
164
|
+
error_tracker: Tracks consecutive retry budget. The caller owns
|
|
165
|
+
this object and passes it in so budget persists across calls.
|
|
166
|
+
tool_specs: Available tools to send to the LLM.
|
|
167
|
+
tool_call_counter: Current counter for generating unique call IDs.
|
|
168
|
+
The updated value is returned in the result.
|
|
169
|
+
step_index: Current iteration index (for compaction and message metadata).
|
|
170
|
+
step_hint: Hint for compaction summarization.
|
|
171
|
+
max_attempts: Maximum LLM calls this invocation may make (including
|
|
172
|
+
retries). When None, bounded only by max_retries. The runner
|
|
173
|
+
passes remaining iteration budget here so retries don't exceed
|
|
174
|
+
max_iterations.
|
|
175
|
+
stream: If True, use send_stream() instead of send().
|
|
176
|
+
on_chunk: Async callback for streaming chunks.
|
|
177
|
+
|
|
178
|
+
Returns:
|
|
179
|
+
InferenceResult with validated tool calls, new messages, and
|
|
180
|
+
updated tool_call_counter. Returns None if max_attempts is
|
|
181
|
+
exhausted without a valid response (caller should treat this
|
|
182
|
+
as iteration budget spent).
|
|
183
|
+
|
|
184
|
+
Raises:
|
|
185
|
+
ToolCallError: If retry budget (max_retries) is exhausted.
|
|
186
|
+
StreamError: If streaming ends without a FINAL chunk.
|
|
187
|
+
"""
|
|
188
|
+
api_format = getattr(client, "api_format", "ollama")
|
|
189
|
+
new_messages: list[Message] = []
|
|
190
|
+
max_retries = error_tracker.max_retries
|
|
191
|
+
attempt_limit = max(max_retries, error_tracker.max_tool_errors) + 1
|
|
192
|
+
if max_attempts is not None:
|
|
193
|
+
attempt_limit = min(attempt_limit, max_attempts)
|
|
194
|
+
attempts = 0
|
|
195
|
+
|
|
196
|
+
# Path-1 verbatim opt-in: drop on any forge mutation (compaction,
|
|
197
|
+
# context warning, retry) so cache_control is only preserved on the
|
|
198
|
+
# clean first-attempt call. ADR-015.
|
|
199
|
+
verbatim_body = inbound_anthropic_body
|
|
200
|
+
|
|
201
|
+
for _attempt in range(attempt_limit):
|
|
202
|
+
attempts += 1
|
|
203
|
+
|
|
204
|
+
# Compact
|
|
205
|
+
compacted = context_manager.maybe_compact(
|
|
206
|
+
messages,
|
|
207
|
+
step_index=step_index,
|
|
208
|
+
step_hint=step_hint,
|
|
209
|
+
)
|
|
210
|
+
# Update the caller's list in-place if compaction changed it
|
|
211
|
+
if compacted is not messages:
|
|
212
|
+
messages.clear()
|
|
213
|
+
messages.extend(compacted)
|
|
214
|
+
verbatim_body = None # mutation
|
|
215
|
+
|
|
216
|
+
# Check context thresholds — inject warning if crossed
|
|
217
|
+
context_warning = context_manager.check_thresholds(messages)
|
|
218
|
+
if context_warning:
|
|
219
|
+
verbatim_body = None # mutation
|
|
220
|
+
|
|
221
|
+
# Fold and serialize. Proxy callers may supply the client's raw OpenAI
|
|
222
|
+
# transcript; on the clean first attempt (no compaction, no warning) we
|
|
223
|
+
# forward it verbatim so the backend sees the client-authored shape
|
|
224
|
+
# instead of forge's parsed/re-emitted form. Any forge mutation
|
|
225
|
+
# (compaction / context warning / retry) falls back to folding.
|
|
226
|
+
use_raw_messages = (
|
|
227
|
+
raw_openai_messages is not None
|
|
228
|
+
and _attempt == 0
|
|
229
|
+
and compacted is messages
|
|
230
|
+
and not context_warning
|
|
231
|
+
)
|
|
232
|
+
if use_raw_messages:
|
|
233
|
+
api_messages = raw_openai_messages
|
|
234
|
+
else:
|
|
235
|
+
api_messages = fold_and_serialize(messages, api_format)
|
|
236
|
+
|
|
237
|
+
# Inject context warning as transient user message (not persisted
|
|
238
|
+
# in conversation history). Uses "user" role because mid-conversation
|
|
239
|
+
# "system" messages break Jinja chat templates on llama-server.
|
|
240
|
+
# Also emit as a CONTEXT_WARNING message so on_message consumers
|
|
241
|
+
# (TUI, CLI) can display it to the user.
|
|
242
|
+
if context_warning:
|
|
243
|
+
api_messages.append({"role": "user", "content": context_warning})
|
|
244
|
+
new_messages.append(
|
|
245
|
+
Message(
|
|
246
|
+
MessageRole.USER,
|
|
247
|
+
context_warning,
|
|
248
|
+
MessageMeta(MessageType.CONTEXT_WARNING, step_index=step_index),
|
|
249
|
+
)
|
|
250
|
+
)
|
|
251
|
+
|
|
252
|
+
# Forward raw tools only on the clean first attempt — on retries forge
|
|
253
|
+
# has appended nudge/tool-error messages, so the parsed tool_specs path
|
|
254
|
+
# (format_tool) is the correct serialization. Pass the kwarg only when
|
|
255
|
+
# set so non-proxy callers (and their client doubles) keep the original
|
|
256
|
+
# call signature.
|
|
257
|
+
raw_tools_kwarg: dict[str, Any] = {}
|
|
258
|
+
if raw_openai_tools is not None and _attempt == 0:
|
|
259
|
+
raw_tools_kwarg["raw_openai_tools"] = raw_openai_tools
|
|
260
|
+
|
|
261
|
+
# Send
|
|
262
|
+
if stream:
|
|
263
|
+
response = await _send_streaming(
|
|
264
|
+
client,
|
|
265
|
+
api_messages,
|
|
266
|
+
tool_specs,
|
|
267
|
+
on_chunk,
|
|
268
|
+
sampling,
|
|
269
|
+
passthrough,
|
|
270
|
+
inbound_anthropic_body=verbatim_body,
|
|
271
|
+
**raw_tools_kwarg,
|
|
272
|
+
)
|
|
273
|
+
else:
|
|
274
|
+
response = await client.send(
|
|
275
|
+
api_messages,
|
|
276
|
+
tools=tool_specs,
|
|
277
|
+
sampling=sampling,
|
|
278
|
+
passthrough=passthrough,
|
|
279
|
+
inbound_anthropic_body=verbatim_body,
|
|
280
|
+
**raw_tools_kwarg,
|
|
281
|
+
)
|
|
282
|
+
# Subsequent attempts (retries) are mutations regardless of outcome.
|
|
283
|
+
verbatim_body = None
|
|
284
|
+
|
|
285
|
+
# Update context manager with real token count if available.
|
|
286
|
+
_sync_token_count(client, context_manager)
|
|
287
|
+
|
|
288
|
+
# Validate
|
|
289
|
+
validation = validator.validate(response)
|
|
290
|
+
|
|
291
|
+
if not validation.needs_retry:
|
|
292
|
+
error_tracker.reset_retries()
|
|
293
|
+
# Intentional text response or validated tool calls
|
|
294
|
+
validated = validation.tool_calls
|
|
295
|
+
return InferenceResult(
|
|
296
|
+
response=validated,
|
|
297
|
+
new_messages=new_messages,
|
|
298
|
+
usage=_get_usage(client),
|
|
299
|
+
tool_call_counter=tool_call_counter,
|
|
300
|
+
attempts=attempts,
|
|
301
|
+
)
|
|
302
|
+
|
|
303
|
+
# Retry path. Budget depends on nudge kind:
|
|
304
|
+
# - tool_arg_validation: tool-error budget (record_result/max_tool_errors).
|
|
305
|
+
# Same family as FileNotFoundError — the model emitted a tool call
|
|
306
|
+
# with bad inputs; the dispatcher just rejected it before tool body.
|
|
307
|
+
# - everything else (retry, unknown_tool): retry budget (max_retries).
|
|
308
|
+
nudge = validation.nudge
|
|
309
|
+
nudge_type = _NUDGE_KIND_TO_TYPE[nudge.kind]
|
|
310
|
+
is_tool_error = nudge.kind in TOOL_ERROR_KINDS
|
|
311
|
+
if is_tool_error:
|
|
312
|
+
error_tracker.record_result(success=False)
|
|
313
|
+
exhausted = error_tracker.tool_errors_exhausted
|
|
314
|
+
budget_label = f"max_tool_errors={error_tracker.max_tool_errors}"
|
|
315
|
+
else:
|
|
316
|
+
error_tracker.record_retry()
|
|
317
|
+
exhausted = error_tracker.retries_exhausted
|
|
318
|
+
budget_label = f"max_retries={max_retries}"
|
|
319
|
+
if exhausted:
|
|
320
|
+
raw = (
|
|
321
|
+
response.content
|
|
322
|
+
if isinstance(response, TextResponse)
|
|
323
|
+
else str([(tc.tool, tc.args) for tc in response])
|
|
324
|
+
)
|
|
325
|
+
raise ToolCallError(
|
|
326
|
+
f"Exhausted after {budget_label} consecutive failed attempts ({nudge.kind})",
|
|
327
|
+
raw_response=raw,
|
|
328
|
+
)
|
|
329
|
+
|
|
330
|
+
# Emit the assistant's failed output, then the corrective signal.
|
|
331
|
+
# Two shapes:
|
|
332
|
+
# - Bare text (no tool_call to anchor on): assistant(text) + user nudge.
|
|
333
|
+
# - Tool call with a recoverable defect (unknown tool name, malformed
|
|
334
|
+
# args): emit assistant(tc) + one tool-error result per tc, mirroring
|
|
335
|
+
# step/prereq enforcement in runner.py. Tool-error rides the canonical
|
|
336
|
+
# channel the model was pretrained on, surviving heavy-context
|
|
337
|
+
# attention drop-off and Mistral _merge_consecutive folding far
|
|
338
|
+
# better than a trailing user-role nudge.
|
|
339
|
+
|
|
340
|
+
if isinstance(response, TextResponse):
|
|
341
|
+
msg = Message(
|
|
342
|
+
MessageRole.ASSISTANT,
|
|
343
|
+
response.content,
|
|
344
|
+
MessageMeta(MessageType.TEXT_RESPONSE, step_index=step_index),
|
|
345
|
+
)
|
|
346
|
+
messages.append(msg)
|
|
347
|
+
new_messages.append(msg)
|
|
348
|
+
# Bare text: no tool_call to attach to, fall back to user nudge.
|
|
349
|
+
nudge_msg = Message(
|
|
350
|
+
MessageRole.USER,
|
|
351
|
+
nudge.content,
|
|
352
|
+
MessageMeta(nudge_type, step_index=step_index),
|
|
353
|
+
)
|
|
354
|
+
messages.append(nudge_msg)
|
|
355
|
+
new_messages.append(nudge_msg)
|
|
356
|
+
else:
|
|
357
|
+
# Tool call with a recoverable defect (unknown tool name, malformed
|
|
358
|
+
# args). Emit reasoning + tool_call, then one tool-error result per
|
|
359
|
+
# tool_call so the corrective signal rides the canonical channel.
|
|
360
|
+
err_prefix = (
|
|
361
|
+
"[ToolArgValidationError]"
|
|
362
|
+
if nudge.kind == "tool_arg_validation"
|
|
363
|
+
else "[UnknownTool]"
|
|
364
|
+
)
|
|
365
|
+
tool_calls = response
|
|
366
|
+
if tool_calls[0].reasoning:
|
|
367
|
+
reasoning_msg = Message(
|
|
368
|
+
MessageRole.ASSISTANT,
|
|
369
|
+
tool_calls[0].reasoning,
|
|
370
|
+
MessageMeta(MessageType.REASONING, step_index=step_index),
|
|
371
|
+
)
|
|
372
|
+
messages.append(reasoning_msg)
|
|
373
|
+
new_messages.append(reasoning_msg)
|
|
374
|
+
tc_infos, tool_call_counter = _build_tool_call_infos(
|
|
375
|
+
tool_calls, tool_call_counter
|
|
376
|
+
)
|
|
377
|
+
tc_msg = Message(
|
|
378
|
+
MessageRole.ASSISTANT,
|
|
379
|
+
"",
|
|
380
|
+
MessageMeta(MessageType.TOOL_CALL, step_index=step_index),
|
|
381
|
+
tool_calls=tc_infos,
|
|
382
|
+
reasoning_content=tool_calls[0].reasoning_content,
|
|
383
|
+
)
|
|
384
|
+
messages.append(tc_msg)
|
|
385
|
+
new_messages.append(tc_msg)
|
|
386
|
+
for tc_info in tc_infos:
|
|
387
|
+
err_msg = Message(
|
|
388
|
+
MessageRole.TOOL,
|
|
389
|
+
f"{err_prefix} {nudge.content}",
|
|
390
|
+
MessageMeta(nudge_type, step_index=step_index),
|
|
391
|
+
tool_name=tc_info.name,
|
|
392
|
+
tool_call_id=tc_info.call_id,
|
|
393
|
+
)
|
|
394
|
+
messages.append(err_msg)
|
|
395
|
+
new_messages.append(err_msg)
|
|
396
|
+
|
|
397
|
+
# max_attempts exhausted without valid response — signal to caller
|
|
398
|
+
return None
|
|
399
|
+
|
|
400
|
+
|
|
401
|
+
async def _send_streaming(
|
|
402
|
+
client: LLMClient,
|
|
403
|
+
api_messages: list[dict[str, Any]],
|
|
404
|
+
tool_specs: list[ToolSpec],
|
|
405
|
+
on_chunk: Callable[[StreamChunk], Awaitable[None]] | None = None,
|
|
406
|
+
sampling: dict[str, Any] | None = None,
|
|
407
|
+
passthrough: dict[str, Any] | None = None,
|
|
408
|
+
inbound_anthropic_body: dict[str, Any] | None = None,
|
|
409
|
+
raw_openai_tools: RawOpenAITools | None = None,
|
|
410
|
+
) -> LLMResponse:
|
|
411
|
+
"""Send via streaming, forwarding chunks to on_chunk callback."""
|
|
412
|
+
response = None
|
|
413
|
+
raw_tools_kwarg: dict[str, Any] = {}
|
|
414
|
+
if raw_openai_tools is not None:
|
|
415
|
+
raw_tools_kwarg["raw_openai_tools"] = raw_openai_tools
|
|
416
|
+
async for chunk in client.send_stream(
|
|
417
|
+
api_messages,
|
|
418
|
+
tools=tool_specs,
|
|
419
|
+
sampling=sampling,
|
|
420
|
+
passthrough=passthrough,
|
|
421
|
+
inbound_anthropic_body=inbound_anthropic_body,
|
|
422
|
+
**raw_tools_kwarg,
|
|
423
|
+
):
|
|
424
|
+
if on_chunk is not None:
|
|
425
|
+
await on_chunk(chunk)
|
|
426
|
+
if chunk.type == ChunkType.FINAL:
|
|
427
|
+
response = chunk.response
|
|
428
|
+
if response is None:
|
|
429
|
+
raise StreamError(
|
|
430
|
+
"Stream ended without FINAL chunk — the client adapter "
|
|
431
|
+
"may be malformed or the connection was interrupted"
|
|
432
|
+
)
|
|
433
|
+
return response
|
|
@@ -0,0 +1,119 @@
|
|
|
1
|
+
"""Message types and serialization."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import json
|
|
6
|
+
from dataclasses import dataclass, field
|
|
7
|
+
from enum import Enum
|
|
8
|
+
from typing import Any
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class MessageRole(str, Enum):
|
|
12
|
+
"""Conversation message roles."""
|
|
13
|
+
|
|
14
|
+
SYSTEM = "system"
|
|
15
|
+
USER = "user"
|
|
16
|
+
ASSISTANT = "assistant"
|
|
17
|
+
TOOL = "tool"
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class MessageType(str, Enum):
|
|
21
|
+
"""Metadata tag for compaction prioritization."""
|
|
22
|
+
|
|
23
|
+
SYSTEM_PROMPT = "system_prompt"
|
|
24
|
+
USER_INPUT = "user_input"
|
|
25
|
+
TOOL_CALL = "tool_call"
|
|
26
|
+
TOOL_RESULT = "tool_result"
|
|
27
|
+
REASONING = "reasoning"
|
|
28
|
+
TEXT_RESPONSE = "text_response"
|
|
29
|
+
STEP_NUDGE = "step_nudge"
|
|
30
|
+
PREREQUISITE_NUDGE = "prerequisite_nudge"
|
|
31
|
+
RETRY_NUDGE = "retry_nudge"
|
|
32
|
+
CONTEXT_WARNING = "context_warning"
|
|
33
|
+
SUMMARY = "summary"
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
@dataclass(frozen=True)
|
|
37
|
+
class MessageMeta:
|
|
38
|
+
"""Metadata attached to a message. Never sent to the API."""
|
|
39
|
+
|
|
40
|
+
type: MessageType
|
|
41
|
+
step_index: int | None = None
|
|
42
|
+
original_type: MessageType | None = None
|
|
43
|
+
token_estimate: int | None = None
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
@dataclass(frozen=True)
|
|
47
|
+
class ToolCallInfo:
|
|
48
|
+
"""One tool call within an assistant message."""
|
|
49
|
+
|
|
50
|
+
name: str
|
|
51
|
+
args: Any
|
|
52
|
+
call_id: str
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
@dataclass
|
|
56
|
+
class Message:
|
|
57
|
+
"""Internal message representation with typed metadata.
|
|
58
|
+
|
|
59
|
+
For assistant messages with tool calls, ``tool_calls`` holds one or more
|
|
60
|
+
ToolCallInfo entries. For tool-result messages, ``tool_name`` and
|
|
61
|
+
``tool_call_id`` pair back to the originating call.
|
|
62
|
+
"""
|
|
63
|
+
|
|
64
|
+
role: MessageRole
|
|
65
|
+
content: str
|
|
66
|
+
metadata: MessageMeta
|
|
67
|
+
# Tool-result pairing fields
|
|
68
|
+
tool_name: str | None = None
|
|
69
|
+
tool_call_id: str | None = None
|
|
70
|
+
# Assistant tool-call list (length >= 1 when present)
|
|
71
|
+
tool_calls: list[ToolCallInfo] | None = None
|
|
72
|
+
reasoning_content: str | None = field(default=None, repr=False)
|
|
73
|
+
|
|
74
|
+
def to_api_dict(self, format: str = "ollama") -> dict[str, Any]:
|
|
75
|
+
"""Serialize for LLM API. Strips metadata.
|
|
76
|
+
|
|
77
|
+
format="ollama": arguments as dict, no "type" field in tool_calls,
|
|
78
|
+
tool results use "tool_name".
|
|
79
|
+
format="openai": arguments as JSON string, "type": "function" and
|
|
80
|
+
"id" required on tool_calls, tool results use "tool_call_id"
|
|
81
|
+
and "name".
|
|
82
|
+
"""
|
|
83
|
+
if self.tool_calls is not None:
|
|
84
|
+
tc_list: list[dict[str, Any]] = []
|
|
85
|
+
for tc in self.tool_calls:
|
|
86
|
+
args: Any = {} if tc.args is None else tc.args
|
|
87
|
+
tc_entry: dict[str, Any] = {
|
|
88
|
+
"function": {
|
|
89
|
+
"name": tc.name,
|
|
90
|
+
"arguments": (
|
|
91
|
+
args
|
|
92
|
+
if format == "openai" and isinstance(args, str)
|
|
93
|
+
else json.dumps(args)
|
|
94
|
+
if format == "openai"
|
|
95
|
+
else args
|
|
96
|
+
),
|
|
97
|
+
},
|
|
98
|
+
}
|
|
99
|
+
if format == "openai":
|
|
100
|
+
tc_entry["type"] = "function"
|
|
101
|
+
tc_entry["id"] = tc.call_id
|
|
102
|
+
tc_list.append(tc_entry)
|
|
103
|
+
payload = {
|
|
104
|
+
"role": self.role.value,
|
|
105
|
+
"content": self.content,
|
|
106
|
+
"tool_calls": tc_list,
|
|
107
|
+
}
|
|
108
|
+
if self.reasoning_content is not None:
|
|
109
|
+
payload["reasoning_content"] = self.reasoning_content
|
|
110
|
+
return payload
|
|
111
|
+
d: dict[str, Any] = {"role": self.role.value, "content": self.content}
|
|
112
|
+
if self.tool_name is not None:
|
|
113
|
+
if format == "openai":
|
|
114
|
+
d["name"] = self.tool_name
|
|
115
|
+
if self.tool_call_id is not None:
|
|
116
|
+
d["tool_call_id"] = self.tool_call_id
|
|
117
|
+
else:
|
|
118
|
+
d["tool_name"] = self.tool_name
|
|
119
|
+
return d
|