ai-push-hooks 0.2.1 → 0.3.1
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.
- package/CHANGELOG.md +73 -1
- package/README.md +80 -525
- package/SECURITY.md +102 -14
- package/ai-push-hooks.toml +9 -2
- package/bin/ai-push-hooks.js +6 -6
- package/package.json +3 -2
- package/pyproject.toml +1 -1
- package/src/ai_push_hooks/artifacts.py +67 -13
- package/src/ai_push_hooks/config.py +575 -22
- package/src/ai_push_hooks/engine.py +116 -7
- package/src/ai_push_hooks/executors/apply.py +75 -36
- package/src/ai_push_hooks/executors/ask.py +224 -0
- package/src/ai_push_hooks/executors/exec.py +17 -801
- package/src/ai_push_hooks/executors/runner_workflow.py +478 -0
- package/src/ai_push_hooks/executors/runners/__init__.py +78 -0
- package/src/ai_push_hooks/executors/runners/claude.py +286 -0
- package/src/ai_push_hooks/executors/runners/codex.py +254 -0
- package/src/ai_push_hooks/executors/runners/command.py +178 -0
- package/src/ai_push_hooks/executors/runners/contracts.py +597 -0
- package/src/ai_push_hooks/executors/runners/opencode.py +528 -0
- package/src/ai_push_hooks/executors/runners/opencode_support.py +276 -0
- package/src/ai_push_hooks/executors/runners/process.py +464 -0
- package/src/ai_push_hooks/executors/runners/registry.py +117 -0
- package/src/ai_push_hooks/executors/step_commands.py +478 -0
- package/src/ai_push_hooks/git_utils.py +834 -0
- package/src/ai_push_hooks/hook.py +1 -1
- package/src/ai_push_hooks/modules/beads.py +1 -1
- package/src/ai_push_hooks/modules/docs.py +129 -89
- package/src/ai_push_hooks/modules/pr.py +1 -1
- package/src/ai_push_hooks/plugin_loader.py +422 -0
- package/src/ai_push_hooks/plugins.py +134 -0
- package/src/ai_push_hooks/prompts_builtin.py +9 -2
- package/src/ai_push_hooks/types.py +407 -75
- package/vendor/README.md +15 -0
- package/vendor/requirements.txt +1 -0
- package/vendor/tomli-2.4.0-py3-none-any.whl +0 -0
- package/src/ai_push_hooks/executors/llm.py +0 -624
|
@@ -8,10 +8,26 @@ from .artifacts import ArtifactStore
|
|
|
8
8
|
from .config import resolve_prompt_text
|
|
9
9
|
from .executors.apply import run_apply_step
|
|
10
10
|
from .executors.assertions import ASSERTION_HANDLERS
|
|
11
|
-
from .executors.exec import EXEC_HANDLERS
|
|
12
|
-
from .
|
|
11
|
+
from .executors.exec import EXEC_HANDLERS
|
|
12
|
+
from .git_utils import env_bool
|
|
13
|
+
from .executors.ask import run_ask_step
|
|
14
|
+
from .executors.step_commands import execute_step_command
|
|
13
15
|
from .modules import COLLECTORS
|
|
14
|
-
from .
|
|
16
|
+
from .plugin_loader import PluginDispatcher
|
|
17
|
+
from .plugins import (
|
|
18
|
+
validate_assert_result,
|
|
19
|
+
validate_collector_result,
|
|
20
|
+
validate_exec_result,
|
|
21
|
+
)
|
|
22
|
+
from .types import (
|
|
23
|
+
CollectorResult,
|
|
24
|
+
HookError,
|
|
25
|
+
ModuleRuntimeState,
|
|
26
|
+
RuntimeContext,
|
|
27
|
+
StepConfig,
|
|
28
|
+
StepResult,
|
|
29
|
+
WorkflowRunResult,
|
|
30
|
+
)
|
|
15
31
|
|
|
16
32
|
CollectorHandler = Callable[[RuntimeContext, ModuleRuntimeState], CollectorResult]
|
|
17
33
|
ExecHandler = Callable[[RuntimeContext, ModuleRuntimeState, StepConfig, list[pathlib.Path]], dict[str, Any]]
|
|
@@ -26,7 +42,7 @@ class WorkflowEngine:
|
|
|
26
42
|
collectors: dict[str, CollectorHandler] | None = None,
|
|
27
43
|
exec_handlers: dict[str, ExecHandler] | None = None,
|
|
28
44
|
assertion_handlers: dict[str, AssertionHandler] | None = None,
|
|
29
|
-
|
|
45
|
+
ask_executor: Callable[[RuntimeContext, StepConfig, str, list[pathlib.Path], str], Any] = run_ask_step,
|
|
30
46
|
apply_executor: Callable[[RuntimeContext, ModuleRuntimeState, StepConfig, str, list[pathlib.Path], str], dict[str, object]] = run_apply_step,
|
|
31
47
|
) -> None:
|
|
32
48
|
self.context = context
|
|
@@ -34,11 +50,16 @@ class WorkflowEngine:
|
|
|
34
50
|
self.collectors = collectors or COLLECTORS
|
|
35
51
|
self.exec_handlers = exec_handlers or EXEC_HANDLERS
|
|
36
52
|
self.assertion_handlers = assertion_handlers or ASSERTION_HANDLERS
|
|
37
|
-
self.
|
|
53
|
+
self.ask_executor = ask_executor
|
|
38
54
|
self.apply_executor = apply_executor
|
|
55
|
+
self._plugin_dispatcher: PluginDispatcher | None = None
|
|
39
56
|
|
|
40
57
|
def run(self) -> WorkflowRunResult:
|
|
41
58
|
self.artifacts.prepare()
|
|
59
|
+
# A loader/cache is deliberately scoped to one workflow run. In
|
|
60
|
+
# particular, a second run must not reuse a module snapshot from the
|
|
61
|
+
# first run.
|
|
62
|
+
self._plugin_dispatcher = PluginDispatcher()
|
|
42
63
|
states = [
|
|
43
64
|
ModuleRuntimeState(module=self.context.config.modules[module_id])
|
|
44
65
|
for module_id in self.context.config.workflow.modules
|
|
@@ -110,14 +131,20 @@ class WorkflowEngine:
|
|
|
110
131
|
return StepResult(status="skipped", artifacts={"result.json": path}, metadata={})
|
|
111
132
|
|
|
112
133
|
if step.type == "collect":
|
|
134
|
+
if step.python:
|
|
135
|
+
input_paths = self._resolve_plugin_inputs(state, step)
|
|
136
|
+
result = validate_collector_result(
|
|
137
|
+
self._dispatch_plugin(state, step, input_paths)
|
|
138
|
+
)
|
|
139
|
+
return self._persist_plugin_collect(state, step, result)
|
|
113
140
|
return self._run_collect(state, step)
|
|
114
141
|
|
|
115
142
|
input_paths = [self.artifacts.resolve_input(state, reference) for reference in step.inputs]
|
|
116
143
|
stage_name = f"{state.module.id}.{step.id}"
|
|
117
144
|
|
|
118
|
-
if step.type == "
|
|
145
|
+
if step.type == "ask":
|
|
119
146
|
prompt = resolve_prompt_text(self.context.repo_root, step)
|
|
120
|
-
payload = self.
|
|
147
|
+
payload = self.ask_executor(self.context, step, prompt, input_paths, stage_name)
|
|
121
148
|
artifact_name = step.output or "result.json"
|
|
122
149
|
if isinstance(payload, (dict, list)) or artifact_name.endswith(".json"):
|
|
123
150
|
path = self.artifacts.write_json(state, state.step_index, step.id, artifact_name, payload)
|
|
@@ -132,6 +159,22 @@ class WorkflowEngine:
|
|
|
132
159
|
return StepResult(artifacts={"result.json": path})
|
|
133
160
|
|
|
134
161
|
if step.type == "exec":
|
|
162
|
+
if step.python:
|
|
163
|
+
plugin_inputs = dict(zip(step.inputs, input_paths))
|
|
164
|
+
payload = validate_exec_result(
|
|
165
|
+
self._dispatch_plugin(state, step, plugin_inputs)
|
|
166
|
+
)
|
|
167
|
+
path = self._persist_plugin_result(state, step, payload)
|
|
168
|
+
return StepResult(artifacts={"result.json": path})
|
|
169
|
+
if step.command:
|
|
170
|
+
persisted = execute_step_command(
|
|
171
|
+
self.context,
|
|
172
|
+
state,
|
|
173
|
+
step,
|
|
174
|
+
dict(zip(step.inputs, input_paths)),
|
|
175
|
+
artifacts=self.artifacts,
|
|
176
|
+
)
|
|
177
|
+
return StepResult(artifacts=dict(persisted.artifacts))
|
|
135
178
|
handler = self.exec_handlers.get(step.executor or "")
|
|
136
179
|
if handler is None:
|
|
137
180
|
raise HookError(f"Unknown exec handler: {step.executor}")
|
|
@@ -140,6 +183,24 @@ class WorkflowEngine:
|
|
|
140
183
|
return StepResult(artifacts={"result.json": path})
|
|
141
184
|
|
|
142
185
|
if step.type == "assert":
|
|
186
|
+
if step.python:
|
|
187
|
+
plugin_inputs = dict(zip(step.inputs, input_paths))
|
|
188
|
+
payload = validate_assert_result(
|
|
189
|
+
self._dispatch_plugin(state, step, plugin_inputs)
|
|
190
|
+
)
|
|
191
|
+
path = self._persist_plugin_result(state, step, payload)
|
|
192
|
+
if not payload["ok"]:
|
|
193
|
+
raise HookError(payload.get("message", "assertion failed"))
|
|
194
|
+
return StepResult(artifacts={"result.json": path})
|
|
195
|
+
if step.command:
|
|
196
|
+
persisted = execute_step_command(
|
|
197
|
+
self.context,
|
|
198
|
+
state,
|
|
199
|
+
step,
|
|
200
|
+
dict(zip(step.inputs, input_paths)),
|
|
201
|
+
artifacts=self.artifacts,
|
|
202
|
+
)
|
|
203
|
+
return StepResult(artifacts=dict(persisted.artifacts))
|
|
143
204
|
handler = self.assertion_handlers.get(step.assertion or "")
|
|
144
205
|
if handler is None:
|
|
145
206
|
raise HookError(f"Unknown assertion handler: {step.assertion}")
|
|
@@ -151,6 +212,54 @@ class WorkflowEngine:
|
|
|
151
212
|
|
|
152
213
|
raise HookError(f"Unsupported step type: {step.type}")
|
|
153
214
|
|
|
215
|
+
def _resolve_plugin_inputs(
|
|
216
|
+
self, state: ModuleRuntimeState, step: StepConfig
|
|
217
|
+
) -> dict[str, pathlib.Path]:
|
|
218
|
+
"""Resolve declared inputs in declaration order for a callback."""
|
|
219
|
+
|
|
220
|
+
return {
|
|
221
|
+
reference: self.artifacts.resolve_input(state, reference)
|
|
222
|
+
for reference in step.inputs
|
|
223
|
+
}
|
|
224
|
+
|
|
225
|
+
def _dispatch_plugin(
|
|
226
|
+
self,
|
|
227
|
+
state: ModuleRuntimeState,
|
|
228
|
+
step: StepConfig,
|
|
229
|
+
input_paths: dict[str, pathlib.Path],
|
|
230
|
+
) -> Any:
|
|
231
|
+
dispatcher = self._plugin_dispatcher
|
|
232
|
+
if dispatcher is None: # pragma: no cover - only direct private calls
|
|
233
|
+
dispatcher = PluginDispatcher()
|
|
234
|
+
return dispatcher.dispatch(self.context, state, step, input_paths)
|
|
235
|
+
|
|
236
|
+
def _persist_plugin_collect(
|
|
237
|
+
self, state: ModuleRuntimeState, step: StepConfig, result: CollectorResult
|
|
238
|
+
) -> StepResult:
|
|
239
|
+
# Serialize and enforce both limits before the first write/register.
|
|
240
|
+
serialized = self.artifacts.serialize_plugin_artifacts(result.artifacts)
|
|
241
|
+
artifacts: dict[str, pathlib.Path] = {}
|
|
242
|
+
for artifact_name, content in serialized.items():
|
|
243
|
+
artifacts[artifact_name] = self.artifacts.write_bytes(
|
|
244
|
+
state, state.step_index, step.id, artifact_name, content
|
|
245
|
+
)
|
|
246
|
+
metadata = dict(result.metadata)
|
|
247
|
+
if result.skip_module:
|
|
248
|
+
metadata["skip_module"] = True
|
|
249
|
+
metadata["skip_reason"] = result.skip_reason
|
|
250
|
+
return StepResult(artifacts=artifacts, metadata=metadata)
|
|
251
|
+
|
|
252
|
+
def _persist_plugin_result(
|
|
253
|
+
self, state: ModuleRuntimeState, step: StepConfig, payload: dict[str, Any]
|
|
254
|
+
) -> pathlib.Path:
|
|
255
|
+
# Use the same bounded serializer as collector artifacts. Validation
|
|
256
|
+
# happens first, and serialization happens before the result is written
|
|
257
|
+
# or registered.
|
|
258
|
+
serialized = self.artifacts.serialize_plugin_artifacts({"result.json": payload})
|
|
259
|
+
return self.artifacts.write_bytes(
|
|
260
|
+
state, state.step_index, step.id, "result.json", serialized["result.json"]
|
|
261
|
+
)
|
|
262
|
+
|
|
154
263
|
def _run_collect(self, state: ModuleRuntimeState, step: StepConfig) -> StepResult:
|
|
155
264
|
handler = self.collectors.get(step.collector or "")
|
|
156
265
|
if handler is None:
|
|
@@ -9,6 +9,7 @@ import tempfile
|
|
|
9
9
|
from dataclasses import dataclass
|
|
10
10
|
from typing import Any
|
|
11
11
|
|
|
12
|
+
from ..config import resolve_runner_profile
|
|
12
13
|
from ..paths import (
|
|
13
14
|
atomic_write_bytes,
|
|
14
15
|
ensure_private_directory,
|
|
@@ -20,14 +21,15 @@ from ..paths import (
|
|
|
20
21
|
sanitize_file_mode,
|
|
21
22
|
)
|
|
22
23
|
from ..types import HookError, ModuleRuntimeState, RuntimeContext, StepConfig
|
|
23
|
-
from
|
|
24
|
+
from ..git_utils import (
|
|
24
25
|
list_repo_changes,
|
|
25
26
|
path_matches,
|
|
26
27
|
resolve_git_common_dir,
|
|
27
28
|
resolve_git_dir,
|
|
28
29
|
run_command,
|
|
29
30
|
)
|
|
30
|
-
from .
|
|
31
|
+
from .runner_workflow import run_runner_once
|
|
32
|
+
from .runners.opencode_support import validate_hook_owned_artifacts
|
|
31
33
|
|
|
32
34
|
METADATA_MAX_FILES = 20_000
|
|
33
35
|
METADATA_MAX_BYTES = 64 * 1024 * 1024
|
|
@@ -354,6 +356,8 @@ def _copy_checkout_to_staging(
|
|
|
354
356
|
context: RuntimeContext,
|
|
355
357
|
staging_root: pathlib.Path,
|
|
356
358
|
allow_paths: tuple[str, ...],
|
|
359
|
+
*,
|
|
360
|
+
project_access: str = "artifacts",
|
|
357
361
|
) -> dict[str, DestinationState]:
|
|
358
362
|
repo_root = context.repo_root
|
|
359
363
|
git_roots = (
|
|
@@ -364,16 +368,25 @@ def _copy_checkout_to_staging(
|
|
|
364
368
|
copied_bytes = 0
|
|
365
369
|
baselines: dict[str, DestinationState] = {}
|
|
366
370
|
resolved_repo_root = repo_root.resolve(strict=True)
|
|
367
|
-
|
|
368
|
-
|
|
369
|
-
|
|
370
|
-
|
|
371
|
-
|
|
372
|
-
|
|
373
|
-
|
|
374
|
-
|
|
371
|
+
checkout_paths = _tracked_and_unignored_paths(repo_root)
|
|
372
|
+
if project_access == "project":
|
|
373
|
+
projection_candidates = checkout_paths
|
|
374
|
+
else:
|
|
375
|
+
projection_candidates = {
|
|
376
|
+
relative_path
|
|
377
|
+
for relative_path in checkout_paths
|
|
378
|
+
if not _is_protected_path(relative_path)
|
|
379
|
+
and any(path_matches(relative_path, pattern) for pattern in allow_paths)
|
|
380
|
+
}
|
|
381
|
+
ignored_paths = _ignored_changed_paths(repo_root, projection_candidates)
|
|
382
|
+
projection_paths = projection_candidates - ignored_paths
|
|
383
|
+
for relative_path in sorted(projection_paths):
|
|
384
|
+
if _is_protected_path(relative_path):
|
|
385
|
+
continue
|
|
375
386
|
source = _repo_path_from_git(repo_root, relative_path)
|
|
376
387
|
if path_has_symlink(repo_root, source):
|
|
388
|
+
if project_access == "project":
|
|
389
|
+
continue
|
|
377
390
|
raise HookError(
|
|
378
391
|
f"Allowed checkout path is or traverses a symlink or reparse point: {relative_path}"
|
|
379
392
|
)
|
|
@@ -384,7 +397,9 @@ def _copy_checkout_to_staging(
|
|
|
384
397
|
continue
|
|
385
398
|
if not source.exists():
|
|
386
399
|
continue
|
|
387
|
-
if not source.
|
|
400
|
+
if not stat.S_ISREG(source.lstat().st_mode):
|
|
401
|
+
if project_access == "project":
|
|
402
|
+
continue
|
|
388
403
|
raise HookError(f"Allowed checkout path is not a regular file: {relative_path}")
|
|
389
404
|
copied_files += 1
|
|
390
405
|
if copied_files > STAGING_MAX_FILES:
|
|
@@ -449,8 +464,11 @@ def _changed_staging_paths(
|
|
|
449
464
|
unexpected = sorted(
|
|
450
465
|
path
|
|
451
466
|
for path in all_paths
|
|
452
|
-
if
|
|
453
|
-
|
|
467
|
+
if before.get(path) != after.get(path)
|
|
468
|
+
and (
|
|
469
|
+
_is_protected_path(path)
|
|
470
|
+
or not any(path_matches(path, pattern) for pattern in allow_paths)
|
|
471
|
+
)
|
|
454
472
|
)
|
|
455
473
|
if unexpected:
|
|
456
474
|
raise HookError("Apply staging workspace contains paths outside allowlist: " + ", ".join(unexpected))
|
|
@@ -732,14 +750,33 @@ def _verify_post_propagation_security_state(
|
|
|
732
750
|
)
|
|
733
751
|
|
|
734
752
|
|
|
735
|
-
def _apply_prompt(
|
|
753
|
+
def _apply_prompt(
|
|
754
|
+
prompt: str,
|
|
755
|
+
allow_paths: tuple[str, ...],
|
|
756
|
+
*,
|
|
757
|
+
project_access: str = "artifacts",
|
|
758
|
+
) -> str:
|
|
736
759
|
rendered_paths = "\n".join(f"- {pattern}" for pattern in allow_paths)
|
|
760
|
+
if project_access == "project":
|
|
761
|
+
projection_text = (
|
|
762
|
+
"This workspace contains the eligible readable project projection and may "
|
|
763
|
+
"include readable repository files beyond the allowlist."
|
|
764
|
+
)
|
|
765
|
+
else:
|
|
766
|
+
projection_text = (
|
|
767
|
+
"This workspace contains only eligible readable files selected by the "
|
|
768
|
+
"allowlist."
|
|
769
|
+
)
|
|
737
770
|
return (
|
|
738
771
|
prompt.rstrip()
|
|
739
|
-
+ "\n\nMANDATORY STAGING
|
|
740
|
-
+
|
|
772
|
+
+ "\n\nMANDATORY STAGING BOUNDARY:\n"
|
|
773
|
+
+ projection_text
|
|
774
|
+
+ "\nOnly changes to paths matching the allowlist below propagate to the real "
|
|
775
|
+
"checkout. Any other staging change fails the apply step.\n"
|
|
776
|
+
+ "Allowlisted paths:\n"
|
|
741
777
|
+ rendered_paths
|
|
742
|
-
+ "\
|
|
778
|
+
+ "\nTool availability is controlled by runner and user policy; this instruction "
|
|
779
|
+
+ "does not impose a universal tool ban. Do not create symlinks.\n"
|
|
743
780
|
)
|
|
744
781
|
|
|
745
782
|
|
|
@@ -778,11 +815,12 @@ def run_apply_step(
|
|
|
778
815
|
input_paths: list[pathlib.Path],
|
|
779
816
|
stage_name: str,
|
|
780
817
|
) -> dict[str, object]:
|
|
781
|
-
validated_inputs =
|
|
818
|
+
validated_inputs = validate_hook_owned_artifacts(context, input_paths)
|
|
782
819
|
for input_path in validated_inputs:
|
|
783
820
|
if input_path.name.endswith("issues.json"):
|
|
784
821
|
issues = json.loads(input_path.read_text(encoding="utf-8"))
|
|
785
822
|
if isinstance(issues, list) and not issues:
|
|
823
|
+
# Legacy compatibility shortcut; keep until an explicit replacement exists.
|
|
786
824
|
return {"changed": False, "changed_files": [], "skipped": True}
|
|
787
825
|
|
|
788
826
|
_assert_apply_targets_checked_out_head(context)
|
|
@@ -798,17 +836,25 @@ def run_apply_step(
|
|
|
798
836
|
propagated_expected: dict[str, StagedFile | None] = {}
|
|
799
837
|
with tempfile.TemporaryDirectory(prefix="ai-push-hooks-apply-") as temporary_directory:
|
|
800
838
|
staging_root = pathlib.Path(temporary_directory).resolve(strict=True)
|
|
801
|
-
|
|
839
|
+
profile = resolve_runner_profile(context.config, step)
|
|
840
|
+
destination_baselines = _copy_checkout_to_staging(
|
|
841
|
+
context,
|
|
842
|
+
staging_root,
|
|
843
|
+
step.allow_paths,
|
|
844
|
+
project_access=profile.project_access,
|
|
845
|
+
)
|
|
802
846
|
staged_before = _inventory_staging(staging_root)
|
|
803
847
|
try:
|
|
804
|
-
result =
|
|
848
|
+
result = run_runner_once(
|
|
805
849
|
context,
|
|
806
|
-
|
|
807
|
-
|
|
808
|
-
|
|
809
|
-
|
|
810
|
-
|
|
811
|
-
|
|
850
|
+
step,
|
|
851
|
+
_apply_prompt(
|
|
852
|
+
prompt,
|
|
853
|
+
step.allow_paths,
|
|
854
|
+
project_access=profile.project_access,
|
|
855
|
+
),
|
|
856
|
+
validated_inputs,
|
|
857
|
+
stage_name,
|
|
812
858
|
working_directory=staging_root,
|
|
813
859
|
)
|
|
814
860
|
except Exception as exc: # noqa: BLE001
|
|
@@ -823,13 +869,6 @@ def run_apply_step(
|
|
|
823
869
|
if call_error is None:
|
|
824
870
|
call_error = exc
|
|
825
871
|
|
|
826
|
-
if result is not None:
|
|
827
|
-
try:
|
|
828
|
-
finalize_opencode_session(context, stage_name, result.session_id)
|
|
829
|
-
except Exception as exc: # noqa: BLE001
|
|
830
|
-
if call_error is None:
|
|
831
|
-
call_error = exc
|
|
832
|
-
|
|
833
872
|
_verify_pre_propagation_security_state(
|
|
834
873
|
context,
|
|
835
874
|
baseline,
|
|
@@ -840,9 +879,9 @@ def run_apply_step(
|
|
|
840
879
|
if call_error is not None:
|
|
841
880
|
raise HookError(f"Apply step failed in isolated staging: {call_error}") from call_error
|
|
842
881
|
if result is None:
|
|
843
|
-
raise HookError("Apply step failed without
|
|
844
|
-
if result.
|
|
845
|
-
details = result.stderr.strip() or result.stdout.strip() or f"exit code {result.
|
|
882
|
+
raise HookError("Apply step failed without a runner result")
|
|
883
|
+
if result.returncode != 0:
|
|
884
|
+
details = result.stderr.strip() or result.stdout.strip() or f"exit code {result.returncode}"
|
|
846
885
|
raise HookError(f"Apply step failed in isolated staging: {details}")
|
|
847
886
|
propagated_expected = _propagate_staging_changes(
|
|
848
887
|
context,
|
|
@@ -0,0 +1,224 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import json
|
|
4
|
+
import os
|
|
5
|
+
import pathlib
|
|
6
|
+
import tempfile
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
from ..types import HookError, RuntimeContext, StepConfig
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def extract_json_array(text: str) -> list[Any]:
|
|
13
|
+
start = text.find("[")
|
|
14
|
+
end = text.rfind("]")
|
|
15
|
+
if start < 0 or end < start:
|
|
16
|
+
raise HookError("Could not find JSON array in model output")
|
|
17
|
+
try:
|
|
18
|
+
payload = json.loads(text[start : end + 1])
|
|
19
|
+
except json.JSONDecodeError as exc:
|
|
20
|
+
raise HookError(f"Failed to parse JSON array from model output: {exc}") from exc
|
|
21
|
+
if not isinstance(payload, list):
|
|
22
|
+
raise HookError("Model output JSON is not an array")
|
|
23
|
+
return payload
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def extract_json_object(text: str) -> dict[str, Any]:
|
|
27
|
+
start = text.find("{")
|
|
28
|
+
end = text.rfind("}")
|
|
29
|
+
if start < 0 or end < start:
|
|
30
|
+
raise HookError("Could not find JSON object in model output")
|
|
31
|
+
try:
|
|
32
|
+
payload = json.loads(text[start : end + 1])
|
|
33
|
+
except json.JSONDecodeError as exc:
|
|
34
|
+
raise HookError(f"Failed to parse JSON object from model output: {exc}") from exc
|
|
35
|
+
if not isinstance(payload, dict):
|
|
36
|
+
raise HookError("Model output JSON is not an object")
|
|
37
|
+
return payload
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def validate_schema(schema: str | None, payload: Any) -> Any:
|
|
41
|
+
if schema is None:
|
|
42
|
+
return payload
|
|
43
|
+
if schema == "string_array":
|
|
44
|
+
if not isinstance(payload, list) or not all(isinstance(item, str) for item in payload):
|
|
45
|
+
raise HookError("Expected schema string_array")
|
|
46
|
+
return payload
|
|
47
|
+
if schema == "docs_issue_array":
|
|
48
|
+
if not isinstance(payload, list):
|
|
49
|
+
raise HookError("Expected schema docs_issue_array")
|
|
50
|
+
for item in payload:
|
|
51
|
+
if not isinstance(item, dict):
|
|
52
|
+
raise HookError("docs_issue_array items must be objects")
|
|
53
|
+
if not str(item.get("file", "")).strip() or not str(item.get("description", "")).strip():
|
|
54
|
+
raise HookError("docs_issue_array items require file and description")
|
|
55
|
+
return payload
|
|
56
|
+
if schema == "beads_alignment_result":
|
|
57
|
+
if not isinstance(payload, dict):
|
|
58
|
+
raise HookError("Expected schema beads_alignment_result")
|
|
59
|
+
commands = payload.get("commands", [])
|
|
60
|
+
if commands is not None and (
|
|
61
|
+
not isinstance(commands, list) or not all(isinstance(item, str) for item in commands)
|
|
62
|
+
):
|
|
63
|
+
raise HookError("beads_alignment_result.commands must be an array of strings")
|
|
64
|
+
return payload
|
|
65
|
+
if schema == "pr_create_payload":
|
|
66
|
+
if not isinstance(payload, dict):
|
|
67
|
+
raise HookError("Expected schema pr_create_payload")
|
|
68
|
+
return payload
|
|
69
|
+
raise HookError(f"Unsupported schema: {schema}")
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def _safe_invalid_output(invocation: Any, output: str) -> str:
|
|
73
|
+
from .runners.contracts import request_sensitive_diagnostics
|
|
74
|
+
|
|
75
|
+
request = invocation.request
|
|
76
|
+
return request_sensitive_diagnostics(request, output, max_chars=400, env=os.environ)
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
def run_ask_step(
|
|
80
|
+
context: RuntimeContext,
|
|
81
|
+
step: StepConfig,
|
|
82
|
+
prompt: str,
|
|
83
|
+
input_paths: list[pathlib.Path],
|
|
84
|
+
stage_name: str,
|
|
85
|
+
) -> Any:
|
|
86
|
+
from .runner_workflow import _finalize_invocation, _invoke_runner, _named_error
|
|
87
|
+
from ..config import resolve_runner_profile
|
|
88
|
+
|
|
89
|
+
total_attempts = context.config.llm.json_max_retries + 1
|
|
90
|
+
prompt_text = prompt
|
|
91
|
+
last_error = ""
|
|
92
|
+
last_output = ""
|
|
93
|
+
wants_json = bool(step.schema)
|
|
94
|
+
expects_json_array = step.schema in {"string_array", "docs_issue_array"}
|
|
95
|
+
session_id: str | None = None
|
|
96
|
+
resume_session = False
|
|
97
|
+
selected_profile = step.runner or context.config.llm.runner
|
|
98
|
+
try:
|
|
99
|
+
profile = resolve_runner_profile(context.config, step)
|
|
100
|
+
except Exception as exc: # noqa: BLE001
|
|
101
|
+
raise _named_error(selected_profile, "unknown", stage_name, exc) from exc
|
|
102
|
+
retained_invocation = None
|
|
103
|
+
|
|
104
|
+
# Artifact-only analysis must retain its historical empty scratch cwd;
|
|
105
|
+
# project-aware analysis gets the actual checkout root so the adapter can
|
|
106
|
+
# provide the selected runner's project-read behavior.
|
|
107
|
+
if profile.project_access == "project":
|
|
108
|
+
working_directory = context.repo_root.resolve(strict=True)
|
|
109
|
+
temporary_directory = None
|
|
110
|
+
else:
|
|
111
|
+
temporary_directory = tempfile.TemporaryDirectory(prefix="ai-push-hooks-ask-")
|
|
112
|
+
working_directory = pathlib.Path(temporary_directory.name).resolve(strict=True)
|
|
113
|
+
|
|
114
|
+
try:
|
|
115
|
+
for attempt in range(1, total_attempts + 1):
|
|
116
|
+
invocation = _invoke_runner(
|
|
117
|
+
context,
|
|
118
|
+
step,
|
|
119
|
+
prompt_text,
|
|
120
|
+
input_paths,
|
|
121
|
+
stage_name,
|
|
122
|
+
working_directory=working_directory,
|
|
123
|
+
session_id=session_id,
|
|
124
|
+
resume_session=resume_session,
|
|
125
|
+
attempt=attempt,
|
|
126
|
+
total_attempts=total_attempts,
|
|
127
|
+
prior_invocation=retained_invocation,
|
|
128
|
+
)
|
|
129
|
+
retained_invocation = None
|
|
130
|
+
result = invocation.result
|
|
131
|
+
try:
|
|
132
|
+
if not wants_json:
|
|
133
|
+
payload = result.final_text
|
|
134
|
+
else:
|
|
135
|
+
if expects_json_array:
|
|
136
|
+
payload = extract_json_array(result.final_text)
|
|
137
|
+
else:
|
|
138
|
+
payload = extract_json_object(result.final_text)
|
|
139
|
+
payload = validate_schema(step.schema, payload)
|
|
140
|
+
except HookError as exc:
|
|
141
|
+
last_error = str(exc)
|
|
142
|
+
last_output = result.final_text
|
|
143
|
+
if attempt >= total_attempts:
|
|
144
|
+
_finalize_invocation(context, invocation, failed=True)
|
|
145
|
+
safe_error = _safe_invalid_output(invocation, last_error)
|
|
146
|
+
raise HookError(
|
|
147
|
+
f"Runner profile `{profile.name}` ({profile.type}) failed at stage "
|
|
148
|
+
f"`{stage_name}`: invalid JSON: {safe_error}. "
|
|
149
|
+
f"{_safe_invalid_output(invocation, last_output)}"
|
|
150
|
+
) from exc
|
|
151
|
+
|
|
152
|
+
snippet = last_output[: context.config.llm.invalid_json_feedback_max_chars]
|
|
153
|
+
suffix = (
|
|
154
|
+
"Return ONLY valid JSON array."
|
|
155
|
+
if expects_json_array
|
|
156
|
+
else "Return ONLY valid JSON object."
|
|
157
|
+
)
|
|
158
|
+
prompt_text = (
|
|
159
|
+
prompt
|
|
160
|
+
+ "\n\nIMPORTANT: Your previous response was invalid JSON and could not be parsed.\n"
|
|
161
|
+
+ f"Parse error: {last_error}\n"
|
|
162
|
+
+ suffix
|
|
163
|
+
+ "\nPrevious invalid output:\n```text\n"
|
|
164
|
+
+ snippet
|
|
165
|
+
+ "\n```"
|
|
166
|
+
)
|
|
167
|
+
|
|
168
|
+
session = result.session
|
|
169
|
+
can_resume = bool(
|
|
170
|
+
getattr(getattr(invocation.runner, "capabilities", None), "supports_resume", False)
|
|
171
|
+
and session is not None
|
|
172
|
+
and session.session_id
|
|
173
|
+
and session.resumable
|
|
174
|
+
)
|
|
175
|
+
if context.config.llm.json_retry_new_session or not can_resume:
|
|
176
|
+
if not session or not session.session_id:
|
|
177
|
+
retry_reason = "session absent"
|
|
178
|
+
elif context.config.llm.json_retry_new_session:
|
|
179
|
+
retry_reason = "fresh session configured"
|
|
180
|
+
else:
|
|
181
|
+
retry_reason = "runner does not support resume"
|
|
182
|
+
retry_message = "Retrying with a fresh runner invocation."
|
|
183
|
+
if retry_reason == "runner does not support resume":
|
|
184
|
+
retry_message = (
|
|
185
|
+
"Retrying with a fresh runner invocation; unsupported session reuse."
|
|
186
|
+
)
|
|
187
|
+
elif retry_reason == "session absent":
|
|
188
|
+
retry_message = (
|
|
189
|
+
"Retrying with a fresh runner invocation; no reusable session was captured."
|
|
190
|
+
)
|
|
191
|
+
context.logger.status(
|
|
192
|
+
"llm.retry_fresh_session",
|
|
193
|
+
retry_message,
|
|
194
|
+
stage_name=stage_name,
|
|
195
|
+
runner_profile=profile.name,
|
|
196
|
+
runner_type=profile.type,
|
|
197
|
+
reason=retry_reason,
|
|
198
|
+
)
|
|
199
|
+
_finalize_invocation(context, invocation, failed=True)
|
|
200
|
+
session_id = None
|
|
201
|
+
resume_session = False
|
|
202
|
+
else:
|
|
203
|
+
# Keep the exact captured session ID. Do not invent a
|
|
204
|
+
# provider-specific resume command for completion output.
|
|
205
|
+
session_id = session.session_id
|
|
206
|
+
resume_session = True
|
|
207
|
+
retained_invocation = invocation
|
|
208
|
+
# The invocation itself completed, although its response
|
|
209
|
+
# failed downstream validation; report truthful metadata
|
|
210
|
+
# before reusing the retained session.
|
|
211
|
+
from .runner_workflow import _completion
|
|
212
|
+
|
|
213
|
+
_completion(context, invocation, failed=True)
|
|
214
|
+
continue
|
|
215
|
+
|
|
216
|
+
_finalize_invocation(context, invocation, failed=False)
|
|
217
|
+
return payload
|
|
218
|
+
raise HookError(
|
|
219
|
+
f"Runner profile `{profile.name}` ({profile.type}) failed at stage `{stage_name}`: "
|
|
220
|
+
"model did not return a valid result"
|
|
221
|
+
) # pragma: no cover
|
|
222
|
+
finally:
|
|
223
|
+
if temporary_directory is not None:
|
|
224
|
+
temporary_directory.cleanup()
|