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.
Files changed (116) hide show
  1. millforge/__init__.py +1174 -0
  2. millforge/_forge/LICENSE +21 -0
  3. millforge/_forge/PROVENANCE.json +295 -0
  4. millforge/_forge/UPDATE_POLICY.md +24 -0
  5. millforge/_forge/__init__.py +14 -0
  6. millforge/_forge/adapter.py +2232 -0
  7. millforge/_forge/base_runner.py +121 -0
  8. millforge/_forge/clients/__init__.py +10 -0
  9. millforge/_forge/clients/base.py +200 -0
  10. millforge/_forge/context/__init__.py +23 -0
  11. millforge/_forge/context/manager.py +178 -0
  12. millforge/_forge/context/strategies.py +335 -0
  13. millforge/_forge/core/__init__.py +16 -0
  14. millforge/_forge/core/inference.py +433 -0
  15. millforge/_forge/core/messages.py +119 -0
  16. millforge/_forge/core/runner.py +479 -0
  17. millforge/_forge/core/steps.py +108 -0
  18. millforge/_forge/core/workflow.py +400 -0
  19. millforge/_forge/errors.py +222 -0
  20. millforge/_forge/guardrails/__init__.py +21 -0
  21. millforge/_forge/guardrails/error_tracker.py +71 -0
  22. millforge/_forge/guardrails/guardrails.py +194 -0
  23. millforge/_forge/guardrails/nudge.py +47 -0
  24. millforge/_forge/guardrails/response_validator.py +119 -0
  25. millforge/_forge/guardrails/step_enforcer.py +183 -0
  26. millforge/_forge/prompts/__init__.py +16 -0
  27. millforge/_forge/prompts/nudges.py +95 -0
  28. millforge/_forge/prompts/templates.py +285 -0
  29. millforge/_version.py +3 -0
  30. millforge/artifacts.py +570 -0
  31. millforge/base/__init__.py +97 -0
  32. millforge/base/composition.py +402 -0
  33. millforge/base/context.py +285 -0
  34. millforge/base/harness.py +138 -0
  35. millforge/base/identity.py +465 -0
  36. millforge/base/options.py +34 -0
  37. millforge/base/platform.py +17 -0
  38. millforge/base/prompt.py +317 -0
  39. millforge/base/runner.py +546 -0
  40. millforge/compiled_plan.py +970 -0
  41. millforge/compiler/__init__.py +231 -0
  42. millforge/compiler/artifact_validation.py +257 -0
  43. millforge/compiler/canonicalization.py +169 -0
  44. millforge/compiler/capabilities.py +66 -0
  45. millforge/compiler/catalogs.py +500 -0
  46. millforge/compiler/diagnostics.py +491 -0
  47. millforge/compiler/graph.py +678 -0
  48. millforge/compiler/lowering.py +198 -0
  49. millforge/compiler/output.py +692 -0
  50. millforge/compiler/parsing.py +1424 -0
  51. millforge/compiler/requests.py +1180 -0
  52. millforge/compiler/schema_validation.py +272 -0
  53. millforge/compiler/semantic.py +490 -0
  54. millforge/compiler/service.py +448 -0
  55. millforge/compiler/source.py +375 -0
  56. millforge/compiler/validators.py +184 -0
  57. millforge/connectors/__init__.py +95 -0
  58. millforge/connectors/admission.py +801 -0
  59. millforge/connectors/broker.py +202 -0
  60. millforge/connectors/contracts.py +1159 -0
  61. millforge/connectors/diagnostics.py +189 -0
  62. millforge/connectors/fake.py +66 -0
  63. millforge/connectors/runtime.py +236 -0
  64. millforge/contracts.py +2860 -0
  65. millforge/custom_tools/__init__.py +67 -0
  66. millforge/custom_tools/compiler.py +724 -0
  67. millforge/custom_tools/contracts.py +1093 -0
  68. millforge/custom_tools/diagnostics.py +205 -0
  69. millforge/eval_artifacts.py +952 -0
  70. millforge/eval_boundary.py +2435 -0
  71. millforge/eval_fixtures/__init__.py +1 -0
  72. millforge/eval_fixtures/default_pack/__init__.py +1 -0
  73. millforge/eval_fixtures/default_pack/fixtures/fixture.08a.bug_diagnosis.traceback.v1.json +52 -0
  74. millforge/eval_fixtures/default_pack/fixtures/fixture.08a.direct_edit.import_sort.v1.json +52 -0
  75. millforge/eval_fixtures/default_pack/fixtures/fixture.08a.evidence_discipline.no_source_change.v1.json +51 -0
  76. millforge/eval_fixtures/default_pack/fixtures/fixture.08a.false_closure.visible_green.v1.json +52 -0
  77. millforge/eval_fixtures/default_pack/fixtures/fixture.08a.multi_file.api_contract.v1.json +54 -0
  78. millforge/eval_fixtures/default_pack/fixtures/fixture.08a.recovery.malformed_artifact.v1.json +54 -0
  79. millforge/eval_fixtures/default_pack/manifest.json +12 -0
  80. millforge/eval_modes.py +1282 -0
  81. millforge/eval_presets.py +1398 -0
  82. millforge/eval_reports.py +2517 -0
  83. millforge/eval_suite.py +2429 -0
  84. millforge/eval_trials.py +2632 -0
  85. millforge/eval_workflow.py +794 -0
  86. millforge/exceptions.py +122 -0
  87. millforge/model_backend.py +2098 -0
  88. millforge/protocols.py +340 -0
  89. millforge/py.typed +0 -0
  90. millforge/runtime.py +1791 -0
  91. millforge/testing/__init__.py +1089 -0
  92. millforge/tools/__init__.py +83 -0
  93. millforge/tools/builtin_runtime.py +1339 -0
  94. millforge/tools/builtins.py +773 -0
  95. millforge/tools/execution.py +1545 -0
  96. millforge/tools/path_policy.py +155 -0
  97. millforge/tools/pi_compat/PI_LICENSE +21 -0
  98. millforge/tools/pi_compat/PROVENANCE.json +55 -0
  99. millforge/tools/pi_compat/UPDATE_POLICY.md +36 -0
  100. millforge/tools/pi_compat/__init__.py +34 -0
  101. millforge/tools/pi_compat/contracts.py +49 -0
  102. millforge/tools/pi_compat/editing.py +390 -0
  103. millforge/tools/pi_compat/mutations.py +57 -0
  104. millforge/tools/pi_compat/operations.py +401 -0
  105. millforge/tools/pi_compat/paths.py +155 -0
  106. millforge/tools/pi_compat/process.py +1375 -0
  107. millforge/tools/pi_compat/search.py +738 -0
  108. millforge/tools/pi_compat/truncation.py +267 -0
  109. millforge/tools/pi_compat_catalog.py +396 -0
  110. millforge/tools/pi_compat_runtime.py +460 -0
  111. millforge/tools/registry.py +553 -0
  112. millforge/tools/results.py +533 -0
  113. millforge-0.1.0.dist-info/METADATA +844 -0
  114. millforge-0.1.0.dist-info/RECORD +116 -0
  115. millforge-0.1.0.dist-info/WHEEL +4 -0
  116. 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,10 @@
1
+ """Private transport-free Forge client protocol helpers."""
2
+
3
+ from millforge._forge.clients.base import ChunkType, LLMClient, StreamChunk, TokenUsage
4
+
5
+ __all__ = [
6
+ "ChunkType",
7
+ "LLMClient",
8
+ "StreamChunk",
9
+ "TokenUsage",
10
+ ]
@@ -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