This repository has been archived on 2026-08-16. You can view files and clone it, but cannot push or open issues or pull requests.
ollie-core/agent/hooks.go

238 lines
7.1 KiB
Go

package agent
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"os"
"os/exec"
"strings"
"syscall"
"time"
olog "ollie/log"
"ollie/paths"
)
// Hook name constants for well-known agent lifecycle events.
const (
HookAgentSpawn = "agentSpawn"
HookPreTurn = "preTurn"
HookPostTurn = "postTurn"
HookPreTool = "preTool"
HookPostTool = "postTool"
HookPreCompact = "preCompact"
HookPostCompact = "postCompact"
HookTurnError = "turnError"
)
const defaultHookTimeout = 60
// hookTimeout is the hook execution timeout in seconds. Overridable in tests.
var hookTimeout = defaultHookTimeout
// Hooks maps hook names to one or more shell commands.
type Hooks map[string][]string
// HookResult holds the outcome of running a hook.
type HookResult struct {
// Ran is true when a hook command was configured and executed.
Ran bool
// Handled is true when all hook commands exited 0 (no warnings, not blocked).
// Use this to check whether a hook fully handled an event (e.g. turnError).
Handled bool
// Blocked is true when the hook wants to prevent the action (exit 2).
// For Stop hooks, Blocked means "don't stop, continue".
Blocked bool
// Context is stdout from the hook, injected into the conversation.
Context string
// Warning, if non-empty, is a message about hook execution problems
// (e.g. timeout, start failure) that should be surfaced to the user.
Warning string
// Total is the number of hook commands configured for this event.
Total int
// Succeeded is the number of hook commands that exited 0.
Succeeded int
// Failed is the number of hook commands that did not succeed (non-zero, timeout, start failure).
Failed int
// FailedCmds identifies which commands failed (truncated to 40 chars each).
FailedCmds []string
}
// Run executes all commands for the named hook in order, sending payload as
// JSON on stdin for each. Returns a combined HookResult.
//
// Exit codes per command:
// - 0: success. Stdout is appended to combined context.
// - 2: block. Stops execution immediately and returns blocked.
// - other: non-blocking warning (stderr logged, execution continues).
func (h Hooks) Run(ctx context.Context, name string, payload any, log *olog.Logger) HookResult {
cmds := h[name]
if len(cmds) == 0 {
return HookResult{}
}
payloadJSON, _ := json.Marshal(payload)
var cwd string
if m, ok := payload.(map[string]string); ok {
cwd = m["cwd"]
}
total := len(cmds)
var contextParts []string
var warnings []string
var failedCmds []string
succeeded := 0
failed := 0
allHandled := true
for _, cmdStr := range cmds {
log.Debug("hook %s: cmd=%q", name, cmdStr)
result := runHookCmd(ctx, name, cmdStr, payloadJSON, cwd, log)
if result.Warning != "" {
warnings = append(warnings, result.Warning)
allHandled = false
failed++
failedCmds = append(failedCmds, truncateCmd(cmdStr))
} else if !result.Ran {
allHandled = false
failed++
failedCmds = append(failedCmds, truncateCmd(cmdStr))
} else if result.Blocked {
// Blocked counts as "ran" for accounting but stops iteration.
succeeded++
return HookResult{
Ran: true, Blocked: true, Context: result.Context,
Warning: strings.Join(warnings, "; "),
Total: total,
Succeeded: succeeded,
Failed: failed,
FailedCmds: failedCmds,
}
} else {
succeeded++
if result.Context != "" {
contextParts = append(contextParts, result.Context)
}
}
}
return HookResult{
Ran: true, Handled: allHandled,
Context: strings.Join(contextParts, "\n"),
Warning: strings.Join(warnings, "; "),
Total: total,
Succeeded: succeeded,
Failed: failed,
FailedCmds: failedCmds,
}
}
func runHookCmd(ctx context.Context, name, cmdStr string, payloadJSON []byte, cwd string, log *olog.Logger) HookResult {
cmd := exec.CommandContext(ctx, "sh", "-c", cmdStr)
cmd.SysProcAttr = &syscall.SysProcAttr{Setpgid: true}
cmd.Stdin = bytes.NewReader(payloadJSON)
if cwd != "" {
expanded := paths.ExpandHome(cwd)
// Only set Dir if the path exists locally (remote sessions
// have a CWD that only exists on the remote host).
if _, err := os.Stat(expanded); err == nil {
cmd.Dir = expanded
}
}
// Inject payload map keys as OLLIE_* environment variables so hook
// commands can use $OLLIE_SESSION_ID, $OLLIE_MODEL, etc. without
// parsing the JSON payload on stdin.
var payloadMap map[string]string
if err := json.Unmarshal(payloadJSON, &payloadMap); err == nil {
env := cmd.Environ()
for k, v := range payloadMap {
env = append(env, "OLLIE_"+strings.ToUpper(k)+"="+v)
}
cmd.Env = env
}
var stdout, stderr bytes.Buffer
cmd.Stdout = &stdout
cmd.Stderr = &stderr
done := make(chan error, 1)
if err := cmd.Start(); err != nil {
log.Debug("hook %s: start error: %v", name, err)
return HookResult{Ran: true, Warning: fmt.Sprintf("hook %s: failed to start: %v", name, err)}
}
go func() { done <- cmd.Wait() }()
timeout := time.After(time.Duration(hookTimeout) * time.Second)
select {
case err := <-done:
exitCode := 0
if err != nil {
var exitErr *exec.ExitError
if errors.As(err, &exitErr) {
exitCode = exitErr.ExitCode()
} else {
log.Debug("hook %s: wait error: %v", name, err)
return HookResult{}
}
}
switch exitCode {
case 0:
out := strings.TrimSpace(stdout.String())
log.Debug("hook %s: exit=0 context_len=%d", name, len(out))
return HookResult{Ran: true, Context: out}
case 2:
msg := strings.TrimSpace(stderr.String())
log.Debug("hook %s: exit=2 (blocked) msg=%q", name, msg)
return HookResult{Ran: true, Blocked: true, Context: msg}
default:
log.Debug("hook %s: exit=%d (non-blocking error) stderr=%q", name, exitCode, stderr.String())
return HookResult{Ran: false, Warning: fmt.Sprintf("hook %s: exit %d", name, exitCode)}
}
case <-ctx.Done():
syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL) //nolint:errcheck
<-done
log.Debug("hook %s: cancelled (context done)", name)
return HookResult{}
case <-timeout:
syscall.Kill(-cmd.Process.Pid, syscall.SIGKILL) //nolint:errcheck
<-done
log.Debug("hook %s: timed out after %ds", name, hookTimeout)
return HookResult{Ran: true, Warning: fmt.Sprintf("hook %s: timed out after %ds", name, hookTimeout)}
}
}
// Summary returns a human-readable summary of hook execution, e.g.
// "(2 of 3 hooks run) (1 of 3 failed: my-script.sh)"
func (r HookResult) Summary() string {
if !r.Ran || r.Total == 0 {
return ""
}
ran := r.Succeeded + r.Failed // commands that were attempted
s := fmt.Sprintf("(%d of %d hooks run)", ran, r.Total)
if r.Failed > 0 {
detail := strings.Join(r.FailedCmds, ", ")
s += fmt.Sprintf(" (%d of %d failed: %s)", r.Failed, r.Total, detail)
}
return s
}
// truncateCmd returns a short identifier for a hook command string.
func truncateCmd(cmd string) string {
cmd = strings.TrimSpace(cmd)
if len(cmd) <= 40 {
return cmd
}
return cmd[:37] + "..."
}
// hooksRan returns a display string for N hooks having run, e.g. "1 hook run".
// Deprecated: prefer HookResult.Summary() for richer output.
func hooksRan(n int) string {
if n == 1 {
return "1 hook run"
}
return fmt.Sprintf("%d hooks run", n)
}