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,71 @@
1
+ """Error budget tracking — consecutive retries and tool errors."""
2
+
3
+ from __future__ import annotations
4
+
5
+
6
+ class ErrorTracker:
7
+ """Tracks consecutive retry and tool error counts against limits.
8
+
9
+ Stateful — instantiate per session/task.
10
+
11
+ Args:
12
+ max_retries: Consecutive formatting/validation failures before
13
+ exhaustion.
14
+ max_tool_errors: Consecutive tool execution errors before
15
+ exhaustion. Soft errors (ToolResolutionError equivalent)
16
+ do not count.
17
+ """
18
+
19
+ def __init__(self, max_retries: int = 3, max_tool_errors: int = 2) -> None:
20
+ self.max_retries = max_retries
21
+ self.max_tool_errors = max_tool_errors
22
+ self._consecutive_retries = 0
23
+ self._consecutive_tool_errors = 0
24
+
25
+ def record_retry(self) -> None:
26
+ """Record a validation failure (TextResponse or unknown tool)."""
27
+ self._consecutive_retries += 1
28
+
29
+ def reset_retries(self) -> None:
30
+ """Reset retry counter (call on successful validation)."""
31
+ self._consecutive_retries = 0
32
+
33
+ def record_result(self, success: bool, is_soft_error: bool = False) -> None:
34
+ """Record a tool execution result.
35
+
36
+ Args:
37
+ success: True if the tool executed without error.
38
+ is_soft_error: True if the error is a resolution/soft error
39
+ that should not count toward the error budget (e.g.,
40
+ ToolResolutionError). Ignored when success is True.
41
+ """
42
+ if success:
43
+ # Individual success doesn't reset — only a fully clean batch does.
44
+ # Call reset_errors() after a batch with zero errors.
45
+ return
46
+ if not is_soft_error:
47
+ self._consecutive_tool_errors += 1
48
+
49
+ def reset_errors(self) -> None:
50
+ """Reset tool error counter (call after a fully clean batch)."""
51
+ self._consecutive_tool_errors = 0
52
+
53
+ @property
54
+ def retries_exhausted(self) -> bool:
55
+ """True if consecutive retries exceed the limit."""
56
+ return self._consecutive_retries > self.max_retries
57
+
58
+ @property
59
+ def tool_errors_exhausted(self) -> bool:
60
+ """True if consecutive tool errors exceed the limit."""
61
+ return self._consecutive_tool_errors > self.max_tool_errors
62
+
63
+ @property
64
+ def consecutive_retries(self) -> int:
65
+ """Current consecutive retry count."""
66
+ return self._consecutive_retries
67
+
68
+ @property
69
+ def consecutive_tool_errors(self) -> int:
70
+ """Current consecutive tool error count."""
71
+ return self._consecutive_tool_errors
@@ -0,0 +1,194 @@
1
+ """Guardrails -- bundled middleware for foreign orchestration loops.
2
+
3
+ Two-method API that wraps ResponseValidator, StepEnforcer, and ErrorTracker:
4
+
5
+ result = guardrails.check(response) # before execution
6
+ done = guardrails.record(["tool"]) # after execution
7
+
8
+ For granular control, use the individual components directly.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ from collections.abc import Callable
14
+ from dataclasses import dataclass
15
+ from typing import Literal
16
+
17
+ from millforge._forge.core.workflow import LLMResponse, ToolCall
18
+ from millforge._forge.guardrails.error_tracker import ErrorTracker
19
+ from millforge._forge.guardrails.nudge import (
20
+ TOOL_CHANNEL_KINDS,
21
+ TOOL_ERROR_KINDS,
22
+ Nudge,
23
+ )
24
+ from millforge._forge.guardrails.response_validator import ResponseValidator
25
+ from millforge._forge.guardrails.step_enforcer import StepEnforcer
26
+
27
+
28
+ @dataclass(frozen=True)
29
+ class CheckResult:
30
+ """Result of checking an LLM response against all guardrails.
31
+
32
+ Attributes:
33
+ action: What the caller should do next.
34
+ "execute" -- tool_calls are safe to run.
35
+ "retry" -- model produced unusable output (bad format / bare
36
+ text); inject nudge as a user message and re-prompt.
37
+ "tool_error" -- model called a tool incorrectly (unknown name or
38
+ malformed args); inject nudge as a tool result
39
+ (nudge.role == "tool") and re-prompt.
40
+ "step_blocked" -- model tried to skip required steps; inject nudge.
41
+ "fatal" -- error budget exhausted; stop the workflow.
42
+ tool_calls: Validated tool calls (only set when action == "execute").
43
+ nudge: Corrective message to inject (set when action is "retry",
44
+ "tool_error", or "step_blocked"). Emit it with its own role:
45
+ ``messages.append({"role": nudge.role, "content": nudge.content})``.
46
+ reason: Human-readable explanation (only set when action == "fatal").
47
+ """
48
+
49
+ action: Literal["execute", "retry", "tool_error", "step_blocked", "fatal"]
50
+ tool_calls: list[ToolCall] | None = None
51
+ nudge: Nudge | None = None
52
+ reason: str | None = None
53
+
54
+
55
+ class Guardrails:
56
+ """Bundled guardrail middleware for foreign orchestration loops.
57
+
58
+ Wraps ResponseValidator, StepEnforcer, and ErrorTracker into a
59
+ two-method API. Use ``check()`` after each LLM response and
60
+ ``record()`` after executing tools.
61
+
62
+ Args:
63
+ tool_names: Valid tool names for this workflow.
64
+ required_steps: Tools that must be called before the terminal tool.
65
+ Defaults to no required steps.
66
+ terminal_tool: The tool(s) that can end the workflow. Accepts a
67
+ single name or a frozenset of names.
68
+ max_retries: Consecutive bad responses before ``check()`` returns
69
+ ``"fatal"``. Default 3.
70
+ max_tool_errors: Consecutive tool execution failures before
71
+ exhaustion. Default 2.
72
+ rescue_enabled: Attempt to parse tool calls from plain text
73
+ responses. Default True.
74
+ max_premature_attempts: Premature terminal attempts before
75
+ ``check()`` returns ``"fatal"``. Default 3.
76
+ retry_nudge: Custom nudge for bare text responses. Pass a callable
77
+ ``(raw_response) -> str`` for dynamic nudges. If None, uses
78
+ the default.
79
+ """
80
+
81
+ def __init__(
82
+ self,
83
+ tool_names: list[str],
84
+ terminal_tool: str | frozenset[str],
85
+ required_steps: list[str] | None = None,
86
+ max_retries: int = 3,
87
+ max_tool_errors: int = 2,
88
+ rescue_enabled: bool = True,
89
+ max_premature_attempts: int = 3,
90
+ retry_nudge: Callable[[str], str] | None = None,
91
+ ) -> None:
92
+ self._validator = ResponseValidator(
93
+ tool_names=tool_names,
94
+ rescue_enabled=rescue_enabled,
95
+ retry_nudge_fn=retry_nudge,
96
+ )
97
+ if isinstance(terminal_tool, str):
98
+ terminal_tools = frozenset([terminal_tool])
99
+ else:
100
+ terminal_tools = terminal_tool
101
+ self._enforcer = StepEnforcer(
102
+ required_steps=required_steps or [],
103
+ terminal_tools=terminal_tools,
104
+ max_premature_attempts=max_premature_attempts,
105
+ )
106
+ self._errors = ErrorTracker(
107
+ max_retries=max_retries,
108
+ max_tool_errors=max_tool_errors,
109
+ )
110
+
111
+ def check(
112
+ self,
113
+ response: LLMResponse,
114
+ ) -> CheckResult:
115
+ """Check an LLM response against all guardrails.
116
+
117
+ Call this after each LLM response, before executing any tools.
118
+
119
+ Args:
120
+ response: The LLM response -- either a TextResponse or a
121
+ list of ToolCall objects.
122
+
123
+ Returns:
124
+ CheckResult indicating what the caller should do next.
125
+ """
126
+ # Checkpoint 1: Is this response usable?
127
+ validation = self._validator.validate(response)
128
+
129
+ if validation.needs_retry:
130
+ nudge = validation.nudge
131
+ kind = nudge.kind if nudge is not None else ""
132
+ # Budget: malformed args drain the tool-error budget; everything
133
+ # else (bare text, unknown tool) drains the retry budget — matching
134
+ # run_inference so all three integration modes account identically.
135
+ if kind in TOOL_ERROR_KINDS:
136
+ self._errors.record_result(success=False)
137
+ if self._errors.tool_errors_exhausted:
138
+ return CheckResult(
139
+ action="fatal",
140
+ reason="too many consecutive tool-argument errors",
141
+ )
142
+ else:
143
+ self._errors.record_retry()
144
+ if self._errors.retries_exhausted:
145
+ return CheckResult(
146
+ action="fatal",
147
+ reason="too many consecutive bad responses",
148
+ )
149
+ # Channel: tool-call faults (unknown tool, malformed args) ride the
150
+ # tool-result channel; bare-text failures stay a user-role retry.
151
+ action = "tool_error" if kind in TOOL_CHANNEL_KINDS else "retry"
152
+ return CheckResult(action=action, nudge=nudge)
153
+
154
+ self._errors.reset_retries()
155
+
156
+ # Checkpoint 2: Is the model skipping required steps?
157
+ step_check = self._enforcer.check(validation.tool_calls)
158
+
159
+ if step_check.needs_nudge:
160
+ if self._enforcer.premature_exhausted:
161
+ return CheckResult(
162
+ action="fatal",
163
+ reason="model repeatedly skipped required steps",
164
+ )
165
+ return CheckResult(action="step_blocked", nudge=step_check.nudge)
166
+
167
+ return CheckResult(action="execute", tool_calls=validation.tool_calls)
168
+
169
+ def record(self, executed: list[str | tuple[str, dict]]) -> bool:
170
+ """Record which tools were successfully executed.
171
+
172
+ Call this after executing tools to keep the middleware in sync.
173
+
174
+ Args:
175
+ executed: Names of tools that succeeded, or (name, args) tuples
176
+ for prerequisite tracking.
177
+
178
+ Returns:
179
+ True if the terminal tool was reached and all required
180
+ steps are satisfied (workflow is done).
181
+ """
182
+ for entry in executed:
183
+ if isinstance(entry, tuple):
184
+ name, args = entry
185
+ self._enforcer.record(name, args)
186
+ else:
187
+ self._enforcer.record(entry)
188
+ self._errors.reset_errors()
189
+ self._enforcer.reset_premature()
190
+ return self._enforcer.is_satisfied() and any(
191
+ (entry if isinstance(entry, str) else entry[0])
192
+ in self._enforcer.terminal_tools
193
+ for entry in executed
194
+ )
@@ -0,0 +1,47 @@
1
+ """Lightweight nudge message returned by guardrail components."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import dataclass
6
+
7
+
8
+ @dataclass(frozen=True)
9
+ class Nudge:
10
+ """A message to inject into conversation history.
11
+
12
+ Returned by guardrail components when the model needs correction.
13
+ The consumer maps this to their framework's message format::
14
+
15
+ # OpenAI-style
16
+ messages.append({"role": nudge.role, "content": nudge.content})
17
+
18
+ # LangChain
19
+ messages.append(HumanMessage(content=nudge.content))
20
+
21
+ Attributes:
22
+ role: Message role for injection ("user", "system", or "tool").
23
+ content: The nudge text.
24
+ kind: Identifies what generated the nudge ("retry", "unknown_tool",
25
+ "step"). Useful for logging/metrics and for WorkflowRunner to
26
+ map back to MessageType for compaction prioritization.
27
+ tier: Escalation level for step nudges (0 = N/A, 1-3 = escalating).
28
+ """
29
+
30
+ role: str
31
+ content: str
32
+ kind: str
33
+ tier: int = 0
34
+
35
+
36
+ # Nudge kinds that drain the tool-error BUDGET (record_result, max_tool_errors)
37
+ # rather than the retry budget. Malformed args are conceptually "tool called
38
+ # with bad inputs" — same family as a runtime FileNotFoundError. Shared by
39
+ # run_inference (core) and the Guardrails facade so both account identically.
40
+ TOOL_ERROR_KINDS: frozenset[str] = frozenset({"tool_arg_validation"})
41
+
42
+ # Nudge kinds emitted on the tool-result CHANNEL (role="tool") — the model
43
+ # called a tool incorrectly (bad name or bad args), so the correction rides the
44
+ # canonical tool channel rather than a user nudge. Maps to the Guardrails
45
+ # facade's action="tool_error". Superset of TOOL_ERROR_KINDS: an unknown-tool
46
+ # call rides the tool channel but only drains the retry budget.
47
+ TOOL_CHANNEL_KINDS: frozenset[str] = frozenset({"unknown_tool", "tool_arg_validation"})
@@ -0,0 +1,119 @@
1
+ """Response validation — rescue, retry, and unknown-tool nudges."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from collections.abc import Callable
6
+ from dataclasses import dataclass
7
+
8
+ from millforge._forge.core.workflow import LLMResponse, TextResponse, ToolCall
9
+ from millforge._forge.guardrails.nudge import Nudge
10
+ from millforge._forge.prompts.nudges import (
11
+ retry_nudge,
12
+ tool_arg_validation_nudge,
13
+ unknown_tool_nudge,
14
+ )
15
+ from millforge._forge.prompts.templates import rescue_tool_call
16
+
17
+
18
+ @dataclass
19
+ class ValidationResult:
20
+ """Result of validating an LLM response.
21
+
22
+ Exactly one of ``tool_calls`` or ``nudge`` is set:
23
+ - If ``needs_retry`` is False: ``tool_calls`` contains validated tool calls.
24
+ - If ``needs_retry`` is True: ``nudge`` contains the message to inject.
25
+ """
26
+
27
+ tool_calls: list[ToolCall] | None
28
+ nudge: Nudge | None
29
+ needs_retry: bool
30
+
31
+
32
+ class ResponseValidator:
33
+ """Validates LLM responses: rescues tool calls from text, checks tool names.
34
+
35
+ Stateless — safe to reuse across turns and sessions.
36
+
37
+ Args:
38
+ tool_names: Valid tool names for this workflow.
39
+ rescue_enabled: If True, attempt to parse tool calls from TextResponse
40
+ before generating a retry nudge.
41
+ retry_nudge_fn: Custom nudge function for bare text responses. Takes
42
+ the raw response text and returns the nudge message. If None,
43
+ uses the default retry nudge from ``millforge._forge.prompts.nudges``.
44
+ """
45
+
46
+ def __init__(
47
+ self,
48
+ tool_names: list[str],
49
+ rescue_enabled: bool = True,
50
+ retry_nudge_fn: Callable[[str], str] | None = None,
51
+ ) -> None:
52
+ self.tool_names = tool_names
53
+ self.rescue_enabled = rescue_enabled
54
+ self._retry_nudge_fn = retry_nudge_fn or retry_nudge
55
+
56
+ def validate(
57
+ self,
58
+ response: LLMResponse,
59
+ ) -> ValidationResult:
60
+ """Validate an LLM response.
61
+
62
+ Args:
63
+ response: Either a TextResponse or a list of ToolCall objects.
64
+
65
+ Returns:
66
+ ValidationResult with tool_calls on success, or a Nudge on failure.
67
+ """
68
+ # TextResponse: rescue, then retry nudge
69
+ if isinstance(response, TextResponse):
70
+ if self.rescue_enabled:
71
+ rescued = rescue_tool_call(response.content, self.tool_names)
72
+ if rescued:
73
+ return ValidationResult(
74
+ tool_calls=rescued, nudge=None, needs_retry=False
75
+ )
76
+ return ValidationResult(
77
+ tool_calls=None,
78
+ nudge=Nudge(
79
+ role="user",
80
+ content=self._retry_nudge_fn(response.content),
81
+ kind="retry",
82
+ ),
83
+ needs_retry=True,
84
+ )
85
+
86
+ # list[ToolCall]: check for unknown tools first (cheap, no point
87
+ # validating args of a hallucinated tool).
88
+ tool_calls = response
89
+ unknown = [tc for tc in tool_calls if tc.tool not in self.tool_names]
90
+ if unknown:
91
+ return ValidationResult(
92
+ tool_calls=None,
93
+ nudge=Nudge(
94
+ role="tool",
95
+ content=unknown_tool_nudge(unknown[0].tool, self.tool_names),
96
+ kind="unknown_tool",
97
+ ),
98
+ needs_retry=True,
99
+ )
100
+
101
+ # Args-shape check. ToolCall no longer enforces dict-args at
102
+ # construction (see workflow.py); the structural check lives here so
103
+ # malformed args ride the tool-error channel via inference.py instead
104
+ # of crashing the client parser.
105
+ bad_args = [tc for tc in tool_calls if not isinstance(tc.args, dict)]
106
+ if bad_args:
107
+ return ValidationResult(
108
+ tool_calls=None,
109
+ nudge=Nudge(
110
+ role="tool",
111
+ content=tool_arg_validation_nudge(
112
+ bad_args[0].tool, bad_args[0].args
113
+ ),
114
+ kind="tool_arg_validation",
115
+ ),
116
+ needs_retry=True,
117
+ )
118
+
119
+ return ValidationResult(tool_calls=tool_calls, nudge=None, needs_retry=False)
@@ -0,0 +1,183 @@
1
+ """Step enforcement — required step tracking, premature terminal nudges,
2
+ and prerequisite enforcement."""
3
+
4
+ from __future__ import annotations
5
+
6
+ from dataclasses import dataclass
7
+ from typing import Any
8
+
9
+ from millforge._forge.core.steps import StepTracker
10
+ from millforge._forge.core.workflow import ToolCall
11
+ from millforge._forge.guardrails.nudge import Nudge
12
+ from millforge._forge.prompts.nudges import prerequisite_nudge, step_nudge
13
+
14
+
15
+ @dataclass
16
+ class StepCheck:
17
+ """Result of checking tool calls against step requirements.
18
+
19
+ If ``needs_nudge`` is True, ``nudge`` contains the message to inject.
20
+ """
21
+
22
+ nudge: Nudge | None
23
+ needs_nudge: bool
24
+
25
+
26
+ class StepEnforcer:
27
+ """Tracks required steps and enforces them with escalating nudges.
28
+
29
+ Also enforces tool prerequisites — conditional dependencies between tools.
30
+
31
+ Stateful — instantiate per session/task.
32
+
33
+ Args:
34
+ required_steps: Tool names that must be called before the terminal tool.
35
+ terminal_tools: The tools that can end the workflow.
36
+ tool_prerequisites: Map of tool name to its ToolDef.prerequisites list.
37
+ max_premature_attempts: How many premature terminal attempts before
38
+ the enforcer signals exhaustion (via StepCheck or raising).
39
+ max_prereq_violations: How many consecutive prerequisite violations
40
+ before the enforcer signals exhaustion.
41
+ """
42
+
43
+ def __init__(
44
+ self,
45
+ required_steps: list[str],
46
+ terminal_tools: frozenset[str],
47
+ tool_prerequisites: dict[str, list[str | dict[str, str]]] | None = None,
48
+ max_premature_attempts: int = 3,
49
+ max_prereq_violations: int = 2,
50
+ ) -> None:
51
+ self._tracker = StepTracker(required_steps=required_steps)
52
+ self.terminal_tools = terminal_tools
53
+ self._tool_prerequisites = tool_prerequisites or {}
54
+ self.max_premature_attempts = max_premature_attempts
55
+ self.max_prereq_violations = max_prereq_violations
56
+ self._premature_attempts = 0
57
+ self._consecutive_prereq_violations = 0
58
+
59
+ def check(self, tool_calls: list[ToolCall]) -> StepCheck:
60
+ """Check whether tool calls include a premature terminal call.
61
+
62
+ If a terminal tool is in the batch and required steps aren't
63
+ satisfied, returns a StepCheck with an escalating nudge. The
64
+ escalation tier increments on each premature attempt (1=polite,
65
+ 2=direct, 3=aggressive).
66
+
67
+ Args:
68
+ tool_calls: The tool calls the model wants to execute.
69
+
70
+ Returns:
71
+ StepCheck with nudge if premature, or no nudge if clear to proceed.
72
+ """
73
+ has_terminal = any(tc.tool in self.terminal_tools for tc in tool_calls)
74
+
75
+ if has_terminal and not self._tracker.is_satisfied():
76
+ self._premature_attempts += 1
77
+ tier = min(self._premature_attempts, 3)
78
+ # Find which terminal tool was attempted for the nudge message
79
+ attempted = next(
80
+ tc.tool for tc in tool_calls if tc.tool in self.terminal_tools
81
+ )
82
+ return StepCheck(
83
+ nudge=Nudge(
84
+ role="user",
85
+ content=step_nudge(
86
+ attempted,
87
+ self._tracker.pending(),
88
+ tier=tier,
89
+ ),
90
+ kind="step",
91
+ tier=tier,
92
+ ),
93
+ needs_nudge=True,
94
+ )
95
+
96
+ return StepCheck(nudge=None, needs_nudge=False)
97
+
98
+ def check_prerequisites(self, tool_calls: list[ToolCall]) -> StepCheck:
99
+ """Check whether any tool call has unsatisfied prerequisites.
100
+
101
+ Evaluates against pre-batch state. Any violation in the batch blocks
102
+ the entire batch (whole-batch blocking).
103
+
104
+ Args:
105
+ tool_calls: The tool calls the model wants to execute.
106
+
107
+ Returns:
108
+ StepCheck with nudge if any prereq is unsatisfied.
109
+ """
110
+ for tc in tool_calls:
111
+ prereqs = self._tool_prerequisites.get(tc.tool)
112
+ if not prereqs:
113
+ continue
114
+ result = self._tracker.check_prerequisites(tc.tool, tc.args, prereqs)
115
+ if not result.satisfied:
116
+ self._consecutive_prereq_violations += 1
117
+ return StepCheck(
118
+ nudge=Nudge(
119
+ role="user",
120
+ content=prerequisite_nudge(
121
+ tc.tool,
122
+ result.missing,
123
+ ),
124
+ kind="prerequisite",
125
+ ),
126
+ needs_nudge=True,
127
+ )
128
+
129
+ return StepCheck(nudge=None, needs_nudge=False)
130
+
131
+ def record(self, tool_name: str, args: dict[str, Any] | None = None) -> None:
132
+ """Record a successful tool execution."""
133
+ self._tracker.record(tool_name, args)
134
+
135
+ def is_satisfied(self) -> bool:
136
+ """True if all required steps have been completed."""
137
+ return self._tracker.is_satisfied()
138
+
139
+ def pending(self) -> list[str]:
140
+ """Return required steps not yet completed."""
141
+ return self._tracker.pending()
142
+
143
+ def terminal_reached(self, tool_calls: list[ToolCall]) -> bool:
144
+ """True if a terminal tool is in the batch and steps are satisfied."""
145
+ has_terminal = any(tc.tool in self.terminal_tools for tc in tool_calls)
146
+ return has_terminal and self._tracker.is_satisfied()
147
+
148
+ @property
149
+ def premature_attempts(self) -> int:
150
+ """Number of premature terminal attempts so far."""
151
+ return self._premature_attempts
152
+
153
+ @property
154
+ def premature_exhausted(self) -> bool:
155
+ """True if premature attempts exceed the limit."""
156
+ return self._premature_attempts > self.max_premature_attempts
157
+
158
+ @property
159
+ def prereq_violations(self) -> int:
160
+ """Number of consecutive prerequisite violations."""
161
+ return self._consecutive_prereq_violations
162
+
163
+ @property
164
+ def prereq_exhausted(self) -> bool:
165
+ """True if consecutive prereq violations exceed the limit."""
166
+ return self._consecutive_prereq_violations > self.max_prereq_violations
167
+
168
+ def reset_premature(self) -> None:
169
+ """Reset premature attempt counter (call after a clean batch)."""
170
+ self._premature_attempts = 0
171
+
172
+ def reset_prereq_violations(self) -> None:
173
+ """Reset consecutive prereq violation counter (call after a clean batch)."""
174
+ self._consecutive_prereq_violations = 0
175
+
176
+ @property
177
+ def completed_steps(self) -> dict[str, None]:
178
+ """Steps completed so far (for diagnostics / error reporting)."""
179
+ return self._tracker.completed_steps
180
+
181
+ def summary_hint(self) -> str:
182
+ """Human-readable hint for context compaction."""
183
+ return self._tracker.summary_hint()
@@ -0,0 +1,16 @@
1
+ """Private Forge prompt helpers."""
2
+
3
+ from millforge._forge.prompts.nudges import retry_nudge, step_nudge
4
+ from millforge._forge.prompts.templates import (
5
+ build_tool_prompt,
6
+ extract_tool_call,
7
+ rescue_tool_call,
8
+ )
9
+
10
+ __all__ = [
11
+ "build_tool_prompt",
12
+ "extract_tool_call",
13
+ "rescue_tool_call",
14
+ "retry_nudge",
15
+ "step_nudge",
16
+ ]