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,121 @@
|
|
|
1
|
+
"""Private Forge graph assembly for the public Millforge base runner facade."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from dataclasses import dataclass
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
from typing import Callable
|
|
8
|
+
|
|
9
|
+
from millforge._forge.adapter import ForgeContextFactory, ForgeGuardrailBackend
|
|
10
|
+
from millforge._forge.errors import ForgeError
|
|
11
|
+
from millforge.base.composition import MillforgeBaseComponents
|
|
12
|
+
from millforge.compiled_plan import CompiledHarnessPlan
|
|
13
|
+
from millforge.contracts import (
|
|
14
|
+
ArtifactRef,
|
|
15
|
+
CompiledHarnessRef,
|
|
16
|
+
HarnessExecutionRequest,
|
|
17
|
+
HarnessExecutionResult,
|
|
18
|
+
)
|
|
19
|
+
from millforge.exceptions import BackendTranslationError
|
|
20
|
+
from millforge.protocols import (
|
|
21
|
+
CancellationResolver,
|
|
22
|
+
ModelClient,
|
|
23
|
+
RuntimeArtifactWriter,
|
|
24
|
+
RuntimeClock,
|
|
25
|
+
)
|
|
26
|
+
from millforge.runtime import DefaultHarnessRuntime
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
@dataclass(frozen=True, slots=True)
|
|
30
|
+
class _PinnedCompiledHarnessLoader:
|
|
31
|
+
plan: CompiledHarnessPlan
|
|
32
|
+
|
|
33
|
+
async def load(self, ref: CompiledHarnessRef) -> CompiledHarnessPlan:
|
|
34
|
+
del ref
|
|
35
|
+
return self.plan
|
|
36
|
+
|
|
37
|
+
|
|
38
|
+
class _LazyRuntimeArtifactWriter:
|
|
39
|
+
"""Defer writer construction until runtime run-directory preparation succeeds."""
|
|
40
|
+
|
|
41
|
+
__slots__ = ("_error", "_factory", "_run_directory", "_writer")
|
|
42
|
+
|
|
43
|
+
def __init__(
|
|
44
|
+
self,
|
|
45
|
+
*,
|
|
46
|
+
factory: Callable[[Path], RuntimeArtifactWriter],
|
|
47
|
+
run_directory: Path,
|
|
48
|
+
) -> None:
|
|
49
|
+
self._factory = factory
|
|
50
|
+
self._run_directory = run_directory
|
|
51
|
+
self._writer: RuntimeArtifactWriter | None = None
|
|
52
|
+
self._error: Exception | None = None
|
|
53
|
+
|
|
54
|
+
def _resolve(self) -> RuntimeArtifactWriter:
|
|
55
|
+
if self._error is not None:
|
|
56
|
+
raise self._error
|
|
57
|
+
if self._writer is None:
|
|
58
|
+
try:
|
|
59
|
+
self._writer = self._factory(self._run_directory)
|
|
60
|
+
except Exception as exc:
|
|
61
|
+
self._error = exc
|
|
62
|
+
raise
|
|
63
|
+
return self._writer
|
|
64
|
+
|
|
65
|
+
async def write_terminal_result(self, ref: ArtifactRef, data: object) -> None:
|
|
66
|
+
await self._resolve().write_terminal_result(ref, data)
|
|
67
|
+
|
|
68
|
+
async def write_execution_summary(self, ref: ArtifactRef, data: object) -> None:
|
|
69
|
+
await self._resolve().write_execution_summary(ref, data)
|
|
70
|
+
|
|
71
|
+
async def write_events(self, ref: ArtifactRef, data: object) -> None:
|
|
72
|
+
await self._resolve().write_events(ref, data)
|
|
73
|
+
|
|
74
|
+
async def write_tool_trace(self, ref: ArtifactRef, data: object) -> None:
|
|
75
|
+
await self._resolve().write_tool_trace(ref, data)
|
|
76
|
+
|
|
77
|
+
async def write_metrics(self, ref: ArtifactRef, data: object) -> None:
|
|
78
|
+
await self._resolve().write_metrics(ref, data)
|
|
79
|
+
|
|
80
|
+
async def write_artifact_manifest(self, ref: ArtifactRef, data: object) -> None:
|
|
81
|
+
await self._resolve().write_artifact_manifest(ref, data)
|
|
82
|
+
|
|
83
|
+
async def write_diagnostic(self, ref: ArtifactRef, data: object) -> None:
|
|
84
|
+
await self._resolve().write_diagnostic(ref, data)
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
async def execute_base_invocation(
|
|
88
|
+
*,
|
|
89
|
+
components: MillforgeBaseComponents,
|
|
90
|
+
request: HarnessExecutionRequest,
|
|
91
|
+
model_client: ModelClient,
|
|
92
|
+
clock: RuntimeClock,
|
|
93
|
+
cancellation_resolver: CancellationResolver,
|
|
94
|
+
artifact_writer_factory: Callable[[Path], RuntimeArtifactWriter],
|
|
95
|
+
) -> HarnessExecutionResult:
|
|
96
|
+
"""Build and execute one fresh private adapter graph."""
|
|
97
|
+
try:
|
|
98
|
+
loader = _PinnedCompiledHarnessLoader(components.compiled_plan)
|
|
99
|
+
writer = _LazyRuntimeArtifactWriter(
|
|
100
|
+
factory=artifact_writer_factory,
|
|
101
|
+
run_directory=request.run_directory.path,
|
|
102
|
+
)
|
|
103
|
+
tool_executor = components.tool_executor.fork_for_invocation()
|
|
104
|
+
backend = ForgeGuardrailBackend(
|
|
105
|
+
model_client=model_client,
|
|
106
|
+
tool_executor=tool_executor,
|
|
107
|
+
plan_loader=loader,
|
|
108
|
+
context_factory=ForgeContextFactory(),
|
|
109
|
+
clock=clock,
|
|
110
|
+
cancellation_resolver=cancellation_resolver,
|
|
111
|
+
)
|
|
112
|
+
runtime = DefaultHarnessRuntime(
|
|
113
|
+
backend=backend,
|
|
114
|
+
plan_loader=loader,
|
|
115
|
+
artifact_writer=writer,
|
|
116
|
+
clock=clock,
|
|
117
|
+
cancellation_resolver=cancellation_resolver,
|
|
118
|
+
)
|
|
119
|
+
return await runtime.execute(request)
|
|
120
|
+
except ForgeError:
|
|
121
|
+
raise BackendTranslationError("Private Forge backend failed") from None
|
|
@@ -0,0 +1,200 @@
|
|
|
1
|
+
"""Streaming types and LLM client protocol."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import json
|
|
6
|
+
from collections.abc import AsyncIterator
|
|
7
|
+
from dataclasses import dataclass
|
|
8
|
+
from enum import Enum
|
|
9
|
+
from typing import Any, Protocol, runtime_checkable
|
|
10
|
+
|
|
11
|
+
from millforge._forge.core.workflow import LLMResponse, ToolSpec
|
|
12
|
+
|
|
13
|
+
# Verbatim OpenAI-shape payloads forwarded by the proxy. The proxy hands the
|
|
14
|
+
# client the user's original ``tools`` array so the backend sees the exact
|
|
15
|
+
# schema the client authored, instead of forge's reconstructed ToolSpec.
|
|
16
|
+
RawOpenAITools = list[dict[str, Any]]
|
|
17
|
+
RawOpenAIMessages = list[dict[str, Any]]
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
@dataclass(frozen=True)
|
|
21
|
+
class TokenUsage:
|
|
22
|
+
"""Token counts from a single LLM response.
|
|
23
|
+
|
|
24
|
+
Populated from the server's ``usage`` field when available (e.g.
|
|
25
|
+
llama-server). Backends that don't report usage leave the client's
|
|
26
|
+
``last_usage`` empty and the context manager falls back to heuristic
|
|
27
|
+
estimation.
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
prompt_tokens: int
|
|
31
|
+
completion_tokens: int
|
|
32
|
+
total_tokens: int
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
# Both Ollama and llama-server use the OpenAI tool schema format today.
|
|
36
|
+
# If a backend diverges, move this back into the relevant client module.
|
|
37
|
+
def format_tool(spec: ToolSpec) -> dict[str, Any]:
|
|
38
|
+
"""Convert a ToolSpec into the OpenAI-compatible tool schema."""
|
|
39
|
+
return {
|
|
40
|
+
"type": "function",
|
|
41
|
+
"function": {
|
|
42
|
+
"name": spec.name,
|
|
43
|
+
"description": spec.description,
|
|
44
|
+
"parameters": spec.get_json_schema(),
|
|
45
|
+
},
|
|
46
|
+
}
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def decode_tool_args(raw: Any) -> Any:
|
|
50
|
+
"""Decode a tool-call ``arguments`` payload, fail-loud.
|
|
51
|
+
|
|
52
|
+
JSON-string args are parsed; on malformed JSON the raw string is returned
|
|
53
|
+
unchanged (a non-dict). ``ResponseValidator``'s args-shape check then routes
|
|
54
|
+
it through the tool-error channel instead of crashing the parser or coercing
|
|
55
|
+
to ``{}`` — so a structural arg failure rides the same lane as a runtime
|
|
56
|
+
tool error rather than a trailing retry nudge.
|
|
57
|
+
|
|
58
|
+
Non-string payloads (an already-decoded dict from Ollama / the Anthropic
|
|
59
|
+
SDK, or any other shape) pass through untouched for the validator to judge.
|
|
60
|
+
A missing or empty payload is a no-arg call (``{}``).
|
|
61
|
+
"""
|
|
62
|
+
if raw is None:
|
|
63
|
+
return {}
|
|
64
|
+
if not isinstance(raw, str):
|
|
65
|
+
return raw
|
|
66
|
+
if not raw:
|
|
67
|
+
return {}
|
|
68
|
+
try:
|
|
69
|
+
return json.loads(raw)
|
|
70
|
+
except json.JSONDecodeError:
|
|
71
|
+
return raw
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
class ChunkType(str, Enum):
|
|
75
|
+
"""What kind of partial data a stream chunk carries."""
|
|
76
|
+
|
|
77
|
+
TEXT_DELTA = "text_delta"
|
|
78
|
+
TOOL_CALL_DELTA = "tool_call_delta"
|
|
79
|
+
FINAL = "final"
|
|
80
|
+
RETRY = "retry"
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
@dataclass(frozen=True)
|
|
84
|
+
class StreamChunk:
|
|
85
|
+
"""A single chunk from a streaming LLM response.
|
|
86
|
+
|
|
87
|
+
Consumers (UI, logging) process TEXT_DELTA and TOOL_CALL_DELTA as they
|
|
88
|
+
arrive. The runner ignores all chunks except FINAL, which carries the
|
|
89
|
+
resolved response. On RETRY, consumers should discard the partial output
|
|
90
|
+
from the failed attempt.
|
|
91
|
+
"""
|
|
92
|
+
|
|
93
|
+
type: ChunkType
|
|
94
|
+
content: str = ""
|
|
95
|
+
response: LLMResponse | None = None
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
@runtime_checkable
|
|
99
|
+
class LLMClient(Protocol):
|
|
100
|
+
"""Interface that client adapters implement.
|
|
101
|
+
|
|
102
|
+
The client is responsible for:
|
|
103
|
+
1. Sending messages to the LLM backend
|
|
104
|
+
2. Parsing the response into ToolCall or TextResponse
|
|
105
|
+
3. Handling native FC or prompt-injected calling internally
|
|
106
|
+
4. Optionally streaming partial responses via send_stream()
|
|
107
|
+
|
|
108
|
+
The client does NOT retry. Retry logic lives in the WorkflowRunner.
|
|
109
|
+
"""
|
|
110
|
+
|
|
111
|
+
api_format: str
|
|
112
|
+
"""Wire format for Message.to_api_dict(): 'ollama' or 'openai'."""
|
|
113
|
+
|
|
114
|
+
model: str
|
|
115
|
+
"""The backend model identity, sent verbatim as the wire "model" field
|
|
116
|
+
(the served-model-name, gguf stem, or model tag depending on backend).
|
|
117
|
+
Distinct from any sampling-registry lookup key a client also derives."""
|
|
118
|
+
|
|
119
|
+
async def send(
|
|
120
|
+
self,
|
|
121
|
+
messages: list[dict[str, str]],
|
|
122
|
+
tools: list[ToolSpec] | None = None,
|
|
123
|
+
sampling: dict[str, Any] | None = None,
|
|
124
|
+
passthrough: dict[str, Any] | None = None,
|
|
125
|
+
inbound_anthropic_body: dict[str, Any] | None = None,
|
|
126
|
+
raw_openai_tools: RawOpenAITools | None = None,
|
|
127
|
+
) -> LLMResponse:
|
|
128
|
+
"""Send messages and return a parsed response.
|
|
129
|
+
|
|
130
|
+
Returns list[ToolCall] if the model produced valid tool invocations.
|
|
131
|
+
Returns TextResponse if the model produced text (reasoning, refusal,
|
|
132
|
+
or malformed output that couldn't be parsed as a tool call).
|
|
133
|
+
|
|
134
|
+
The runner inspects the response and decides whether to retry.
|
|
135
|
+
|
|
136
|
+
Args:
|
|
137
|
+
messages: API-format messages to send.
|
|
138
|
+
tools: Tool specs to include with the request.
|
|
139
|
+
sampling: Optional per-call sampling overrides
|
|
140
|
+
(``temperature``, ``top_p``, ``top_k``, ``min_p``,
|
|
141
|
+
``repeat_penalty``, ``presence_penalty``, ``seed``).
|
|
142
|
+
Per-call values win over instance state for this call only;
|
|
143
|
+
the client's instance fields are not mutated.
|
|
144
|
+
passthrough: Optional dict of inbound body fields forge doesn't
|
|
145
|
+
own. The client merges these into the outbound body before
|
|
146
|
+
overlaying its own fields (model, messages, tools, sampling).
|
|
147
|
+
Used by the proxy to preserve user intent (max_tokens, stop,
|
|
148
|
+
tool_choice, etc.) without forge having to enumerate every
|
|
149
|
+
supported field. None = no extras to merge.
|
|
150
|
+
inbound_anthropic_body: Path-1 only — when set, the AnthropicClient
|
|
151
|
+
will send this body verbatim (bypassing its deconstruct/rebuild
|
|
152
|
+
path) to preserve block-level Anthropic fields like
|
|
153
|
+
``cache_control``. The runner clears this kwarg on any
|
|
154
|
+
forge-mutation (retry / compaction / context warning) so
|
|
155
|
+
only the clean first-attempt call rides verbatim. Other
|
|
156
|
+
clients accept and ignore. See ADR-015.
|
|
157
|
+
raw_openai_tools: Proxy-only — the client's verbatim OpenAI
|
|
158
|
+
``tools`` array. When set, LlamafileClient's native path sends
|
|
159
|
+
it as-is instead of re-emitting ``format_tool(spec)``, so the
|
|
160
|
+
backend sees the original schema (no name/schema drift). Other
|
|
161
|
+
clients accept and ignore.
|
|
162
|
+
"""
|
|
163
|
+
...
|
|
164
|
+
|
|
165
|
+
async def send_stream(
|
|
166
|
+
self,
|
|
167
|
+
messages: list[dict[str, str]],
|
|
168
|
+
tools: list[ToolSpec] | None = None,
|
|
169
|
+
sampling: dict[str, Any] | None = None,
|
|
170
|
+
passthrough: dict[str, Any] | None = None,
|
|
171
|
+
inbound_anthropic_body: dict[str, Any] | None = None,
|
|
172
|
+
raw_openai_tools: RawOpenAITools | None = None,
|
|
173
|
+
) -> AsyncIterator[StreamChunk]:
|
|
174
|
+
"""Send messages and yield streaming chunks.
|
|
175
|
+
|
|
176
|
+
Yields TEXT_DELTA or TOOL_CALL_DELTA chunks as they arrive.
|
|
177
|
+
The final chunk has type FINAL and carries the resolved LLMResponse
|
|
178
|
+
(same list[ToolCall] | TextResponse as send() would return).
|
|
179
|
+
|
|
180
|
+
The runner forwards chunks to its on_chunk callback for UI/logging,
|
|
181
|
+
then inspects the FINAL chunk and decides whether to retry.
|
|
182
|
+
|
|
183
|
+
Args:
|
|
184
|
+
messages: API-format messages to send.
|
|
185
|
+
tools: Tool specs to include with the request.
|
|
186
|
+
sampling: Optional per-call sampling overrides (see ``send``).
|
|
187
|
+
Per-call values win over instance state without mutating self.
|
|
188
|
+
passthrough: Optional inbound-body extras dict (see ``send``).
|
|
189
|
+
inbound_anthropic_body: Optional path-1 verbatim body (see ``send``).
|
|
190
|
+
raw_openai_tools: Optional verbatim OpenAI tools array (see ``send``).
|
|
191
|
+
"""
|
|
192
|
+
...
|
|
193
|
+
|
|
194
|
+
async def get_context_length(self) -> int | None:
|
|
195
|
+
"""Query the backend for its configured context window size."""
|
|
196
|
+
...
|
|
197
|
+
|
|
198
|
+
async def aclose(self) -> None:
|
|
199
|
+
"""Release held network resources (e.g. the httpx connection pool)."""
|
|
200
|
+
...
|
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
"""Private Forge context management subset without hardware discovery."""
|
|
2
|
+
|
|
3
|
+
from millforge._forge.context.manager import (
|
|
4
|
+
CompactEvent,
|
|
5
|
+
ContextManager,
|
|
6
|
+
default_context_warning,
|
|
7
|
+
)
|
|
8
|
+
from millforge._forge.context.strategies import (
|
|
9
|
+
CompactStrategy,
|
|
10
|
+
NoCompact,
|
|
11
|
+
SlidingWindowCompact,
|
|
12
|
+
TieredCompact,
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
__all__ = [
|
|
16
|
+
"CompactEvent",
|
|
17
|
+
"CompactStrategy",
|
|
18
|
+
"ContextManager",
|
|
19
|
+
"NoCompact",
|
|
20
|
+
"SlidingWindowCompact",
|
|
21
|
+
"TieredCompact",
|
|
22
|
+
"default_context_warning",
|
|
23
|
+
]
|
|
@@ -0,0 +1,178 @@
|
|
|
1
|
+
"""Context manager for budget tracking and compaction triggering."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import Callable
|
|
6
|
+
from dataclasses import dataclass
|
|
7
|
+
|
|
8
|
+
from millforge._forge.context.strategies import CompactStrategy
|
|
9
|
+
from millforge._forge.core.messages import Message
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
@dataclass(frozen=True)
|
|
13
|
+
class CompactEvent:
|
|
14
|
+
"""Emitted by ContextManager when compaction fires."""
|
|
15
|
+
|
|
16
|
+
step_index: int
|
|
17
|
+
tokens_before: int
|
|
18
|
+
tokens_after: int
|
|
19
|
+
budget_tokens: int
|
|
20
|
+
messages_before: int
|
|
21
|
+
messages_after: int
|
|
22
|
+
phase_reached: int
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
# ── Default context warning ──────────────────────────────────────
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def default_context_warning(tokens: int, budget: int, pct: float) -> str | None:
|
|
29
|
+
"""Default context threshold callback.
|
|
30
|
+
|
|
31
|
+
Returns an escalating warning based on how full the context is.
|
|
32
|
+
"""
|
|
33
|
+
if pct >= 0.80:
|
|
34
|
+
return (
|
|
35
|
+
f"[Context usage: {pct:.0%} ({tokens:,} / {budget:,} tokens). "
|
|
36
|
+
"Context is nearly full. Older tool results and reasoning will be "
|
|
37
|
+
"compacted soon — key information may be lost. Summarize critical "
|
|
38
|
+
"findings now and prioritize completing the current task.]"
|
|
39
|
+
)
|
|
40
|
+
if pct >= 0.65:
|
|
41
|
+
return (
|
|
42
|
+
f"[Context usage: {pct:.0%} ({tokens:,} / {budget:,} tokens). "
|
|
43
|
+
"Context is filling up. When compaction triggers, older tool results "
|
|
44
|
+
"and reasoning will be condensed. Be concise in your responses and "
|
|
45
|
+
"front-load important information.]"
|
|
46
|
+
)
|
|
47
|
+
return (
|
|
48
|
+
f"[Context usage: {pct:.0%} ({tokens:,} / {budget:,} tokens). "
|
|
49
|
+
"Be mindful of context usage.]"
|
|
50
|
+
)
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class ContextManager:
|
|
54
|
+
"""Manages context window budget and triggers compaction."""
|
|
55
|
+
|
|
56
|
+
def __init__(
|
|
57
|
+
self,
|
|
58
|
+
strategy: CompactStrategy,
|
|
59
|
+
budget_tokens: int,
|
|
60
|
+
on_compact: Callable[[CompactEvent], None] | None = None,
|
|
61
|
+
context_thresholds: list[float] | None = None,
|
|
62
|
+
on_context_threshold: Callable[[int, int, float], str | None] | None = None,
|
|
63
|
+
) -> None:
|
|
64
|
+
"""
|
|
65
|
+
Args:
|
|
66
|
+
strategy: Compaction strategy to use. The strategy owns its own
|
|
67
|
+
compaction thresholds (e.g. ``TieredCompact(compact_threshold=0.75)``
|
|
68
|
+
or ``TieredCompact(phase_thresholds=(0.6, 0.75, 0.9))``).
|
|
69
|
+
budget_tokens: Maximum context budget in tokens.
|
|
70
|
+
on_compact: Callback invoked when compaction fires. Receives a
|
|
71
|
+
CompactEvent with before/after token counts, phase reached,
|
|
72
|
+
and which messages were affected. Use for logging, debugging,
|
|
73
|
+
or surfacing compaction to a UI.
|
|
74
|
+
context_thresholds: Sorted list of budget fractions (e.g.
|
|
75
|
+
``[0.5, 0.65, 0.8]``) at which to fire the context
|
|
76
|
+
threshold callback. Each threshold fires at most once per
|
|
77
|
+
session (resets if usage drops below it after compaction).
|
|
78
|
+
Defaults to None (disabled).
|
|
79
|
+
on_context_threshold: Callback invoked when a context threshold
|
|
80
|
+
is crossed. Receives ``(tokens, budget, pct)`` and returns
|
|
81
|
+
an optional string to inject as a transient system message
|
|
82
|
+
before the next inference call. Return None to skip
|
|
83
|
+
injection. Defaults to None (disabled).
|
|
84
|
+
"""
|
|
85
|
+
self.strategy = strategy
|
|
86
|
+
self.budget_tokens = budget_tokens
|
|
87
|
+
self.on_compact = on_compact
|
|
88
|
+
self._context_thresholds = (
|
|
89
|
+
sorted(context_thresholds) if context_thresholds else []
|
|
90
|
+
)
|
|
91
|
+
self._on_context_threshold = on_context_threshold
|
|
92
|
+
self._fired_thresholds: set[float] = set()
|
|
93
|
+
self._last_known_tokens: int | None = None
|
|
94
|
+
|
|
95
|
+
def update_token_count(self, total_tokens: int) -> None:
|
|
96
|
+
"""Record actual token count from the backend.
|
|
97
|
+
|
|
98
|
+
Called after each LLM response when the backend reports usage.
|
|
99
|
+
Subsequent calls to ``estimate_tokens`` will return this value
|
|
100
|
+
until the next update.
|
|
101
|
+
"""
|
|
102
|
+
self._last_known_tokens = total_tokens
|
|
103
|
+
|
|
104
|
+
def estimate_tokens(self, messages: list[Message]) -> int:
|
|
105
|
+
"""Return actual token count if available, else char/4 heuristic."""
|
|
106
|
+
if self._last_known_tokens is not None:
|
|
107
|
+
return self._last_known_tokens
|
|
108
|
+
return (
|
|
109
|
+
sum(
|
|
110
|
+
len(message.content) + len(message.reasoning_content or "")
|
|
111
|
+
for message in messages
|
|
112
|
+
)
|
|
113
|
+
// 4
|
|
114
|
+
)
|
|
115
|
+
|
|
116
|
+
def check_thresholds(self, messages: list[Message]) -> str | None:
|
|
117
|
+
"""Check context thresholds and return an optional warning to inject.
|
|
118
|
+
|
|
119
|
+
Fires the ``on_context_threshold`` callback when usage crosses a
|
|
120
|
+
configured threshold for the first time. Thresholds reset if usage
|
|
121
|
+
drops below them (e.g. after compaction).
|
|
122
|
+
|
|
123
|
+
Returns:
|
|
124
|
+
A string to inject as a transient system message, or None.
|
|
125
|
+
"""
|
|
126
|
+
if not self._context_thresholds or not self._on_context_threshold:
|
|
127
|
+
return None
|
|
128
|
+
|
|
129
|
+
tokens = self.estimate_tokens(messages)
|
|
130
|
+
if self.budget_tokens <= 0:
|
|
131
|
+
return None
|
|
132
|
+
|
|
133
|
+
pct = tokens / self.budget_tokens
|
|
134
|
+
|
|
135
|
+
# Reset thresholds that usage has dropped below (after compaction)
|
|
136
|
+
self._fired_thresholds = {t for t in self._fired_thresholds if pct >= t}
|
|
137
|
+
|
|
138
|
+
# Find the highest unfired threshold that has been crossed
|
|
139
|
+
highest_crossed: float | None = None
|
|
140
|
+
for threshold in self._context_thresholds:
|
|
141
|
+
if pct >= threshold and threshold not in self._fired_thresholds:
|
|
142
|
+
highest_crossed = threshold
|
|
143
|
+
|
|
144
|
+
if highest_crossed is None:
|
|
145
|
+
return None
|
|
146
|
+
|
|
147
|
+
self._fired_thresholds.add(highest_crossed)
|
|
148
|
+
return self._on_context_threshold(tokens, self.budget_tokens, pct)
|
|
149
|
+
|
|
150
|
+
def maybe_compact(
|
|
151
|
+
self,
|
|
152
|
+
messages: list[Message],
|
|
153
|
+
step_index: int = 0,
|
|
154
|
+
step_hint: str = "",
|
|
155
|
+
) -> list[Message]:
|
|
156
|
+
"""Delegate to the strategy, which owns threshold logic."""
|
|
157
|
+
tokens_before = self.estimate_tokens(messages)
|
|
158
|
+
|
|
159
|
+
result, phase = self.strategy.compact(
|
|
160
|
+
messages, self.budget_tokens, step_hint=step_hint
|
|
161
|
+
)
|
|
162
|
+
|
|
163
|
+
if phase == 0:
|
|
164
|
+
return messages
|
|
165
|
+
|
|
166
|
+
if self.on_compact is not None:
|
|
167
|
+
event = CompactEvent(
|
|
168
|
+
step_index=step_index,
|
|
169
|
+
tokens_before=tokens_before,
|
|
170
|
+
tokens_after=self.estimate_tokens(result),
|
|
171
|
+
budget_tokens=self.budget_tokens,
|
|
172
|
+
messages_before=len(messages),
|
|
173
|
+
messages_after=len(result),
|
|
174
|
+
phase_reached=phase,
|
|
175
|
+
)
|
|
176
|
+
self.on_compact(event)
|
|
177
|
+
|
|
178
|
+
return result
|