| 1 | package main |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "encoding/json" |
| 6 | "fmt" |
| 7 | "os" |
| 8 | "path/filepath" |
| 9 | "strings" |
| 10 | "sync" |
| 11 | "testing" |
| 12 | |
| 13 | "reasonix/internal/hook" |
| 14 | "reasonix/internal/provider" |
| 15 | ) |
| 16 | |
| 17 | const tryHookProbeKind = "desktop-try-hook-probe" |
| 18 | |
| 19 | var tryHookProbe = &tryHookProbeProvider{} |
| 20 | |
| 21 | func init() { |
| 22 | provider.Register(tryHookProbeKind, func(provider.Config) (provider.Provider, error) { |
| 23 | return tryHookProbe, nil |
| 24 | }) |
| 25 | } |
| 26 | |
| 27 | // The try run reads the user's own open workspace, so it runs the hooks a chat |
| 28 | // session there would: the project's and the user's. |
| 29 | func TestTrySubagentProfileRunsWorkspaceAndUserHooks(t *testing.T) { |
| 30 | isolateDesktopUserDirs(t) |
| 31 | root := t.TempDir() |
| 32 | if err := os.WriteFile(filepath.Join(root, "marker.txt"), []byte("try hook probe"), 0o644); err != nil { |
| 33 | t.Fatal(err) |
| 34 | } |
| 35 | if err := os.WriteFile(filepath.Join(root, "reasonix.toml"), []byte(` |
| 36 | default_model = "tryer" |
| 37 | |
| 38 | [[providers]] |
| 39 | name = "tryer" |
| 40 | kind = "`+tryHookProbeKind+`" |
| 41 | model = "try-model" |
| 42 | base_url = "http://127.0.0.1:1" |
| 43 | `), 0o644); err != nil { |
| 44 | t.Fatal(err) |
| 45 | } |
| 46 | approveWorkspace(t, root) |
| 47 | scripts := t.TempDir() |
| 48 | projectLog := filepath.Join(scripts, "project.log") |
| 49 | globalLog := filepath.Join(scripts, "global.log") |
| 50 | writeTryHookSettings(t, hook.ProjectSettingsPath(root), writeTryHookScript(t, scripts, "project-log.sh", projectLog, 0)) |
| 51 | writeTryHookSettings(t, hook.GlobalSettingsPath(""), writeTryHookScript(t, scripts, "global-deny.sh", globalLog, 2)) |
| 52 | if err := hook.ApproveProjectHooks(hook.LoadOptions{ProjectRoot: root}); err != nil { |
| 53 | t.Fatal(err) |
| 54 | } |
| 55 | |
| 56 | a := NewApp() |
| 57 | a.tabs = map[string]*WorkspaceTab{"test": {ID: "test", Scope: "project", WorkspaceRoot: root, Ready: true}} |
| 58 | a.activeTabID = "test" |
| 59 | tryHookProbe.reset() |
| 60 | if _, err := a.TrySubagentProfile(SubagentProfileInput{SystemPrompt: "read marker.txt"}, "read marker.txt"); err != nil { |
| 61 | t.Fatalf("TrySubagentProfile: %v", err) |
| 62 | } |
| 63 | if tryHookProbe.readLeaked() { |
| 64 | t.Fatal("try run received marker.txt contents past a PreToolUse deny") |
| 65 | } |
| 66 | if !tryHookProbe.readBlocked() { |
| 67 | t.Fatal("try run's read_file result was not blocked by PreToolUse") |
| 68 | } |
| 69 | project := readTryHookPayload(t, projectLog, "project") |
| 70 | global := readTryHookPayload(t, globalLog, "global") |
| 71 | for _, p := range []tryHookPayload{project, global} { |
| 72 | if p.ToolName != "read_file" || !strings.HasPrefix(p.SessionID, "try-subagent:") || p.SessionID == "try-subagent:" { |
| 73 | t.Fatalf("hook payload = %+v, want read_file under a per-run try-subagent session", p) |
| 74 | } |
| 75 | } |
| 76 | } |
| 77 | |
| 78 | type tryHookPayload struct{ ToolName, SessionID string } |
| 79 | |
| 80 | func writeTryHookScript(t *testing.T, dir, name, logPath string, exit int) string { |
| 81 | t.Helper() |
| 82 | path := filepath.Join(dir, name) |
| 83 | body := fmt.Sprintf("#!/bin/sh\ncat >> '%s'\nexit %d\n", strings.ReplaceAll(logPath, "'", `'\''`), exit) |
| 84 | if err := os.WriteFile(path, []byte(body), 0o755); err != nil { |
| 85 | t.Fatal(err) |
| 86 | } |
| 87 | return path |
| 88 | } |
| 89 | |
| 90 | func writeTryHookSettings(t *testing.T, path, command string) { |
| 91 | t.Helper() |
| 92 | settings, err := json.Marshal(map[string]any{"hooks": map[string]any{"PreToolUse": []any{map[string]string{"match": "read_file", "command": command}}}}) |
| 93 | if err != nil { |
| 94 | t.Fatal(err) |
| 95 | } |
| 96 | if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { |
| 97 | t.Fatal(err) |
| 98 | } |
| 99 | if err := os.WriteFile(path, settings, 0o644); err != nil { |
| 100 | t.Fatal(err) |
| 101 | } |
| 102 | } |
| 103 | |
| 104 | func readTryHookPayload(t *testing.T, path, scope string) tryHookPayload { |
| 105 | t.Helper() |
| 106 | log, err := os.ReadFile(path) |
| 107 | if err != nil { |
| 108 | t.Fatalf("%s PreToolUse hook never ran for the try run's read_file: %v", scope, err) |
| 109 | } |
| 110 | var p tryHookPayload |
| 111 | if err := json.Unmarshal([]byte(strings.TrimSpace(string(log))), &p); err != nil { |
| 112 | t.Fatalf("decode %s hook payload %q: %v", scope, log, err) |
| 113 | } |
| 114 | return p |
| 115 | } |
| 116 | |
| 117 | type tryHookProbeProvider struct { |
| 118 | mu sync.Mutex |
| 119 | calls int |
| 120 | blocked bool |
| 121 | leaked bool |
| 122 | } |
| 123 | |
| 124 | func (p *tryHookProbeProvider) reset() { |
| 125 | p.mu.Lock() |
| 126 | defer p.mu.Unlock() |
| 127 | p.calls, p.blocked, p.leaked = 0, false, false |
| 128 | } |
| 129 | |
| 130 | func (p *tryHookProbeProvider) readBlocked() bool { |
| 131 | p.mu.Lock() |
| 132 | defer p.mu.Unlock() |
| 133 | return p.blocked |
| 134 | } |
| 135 | |
| 136 | func (p *tryHookProbeProvider) readLeaked() bool { |
| 137 | p.mu.Lock() |
| 138 | defer p.mu.Unlock() |
| 139 | return p.leaked |
| 140 | } |
| 141 | |
| 142 | func (p *tryHookProbeProvider) Name() string { return tryHookProbeKind } |
| 143 | |
| 144 | func (p *tryHookProbeProvider) Stream(_ context.Context, req provider.Request) (<-chan provider.Chunk, error) { |
| 145 | p.mu.Lock() |
| 146 | call := p.calls |
| 147 | p.calls++ |
| 148 | for _, msg := range req.Messages { |
| 149 | if msg.Role != provider.RoleTool || msg.Name != "read_file" { |
| 150 | continue |
| 151 | } |
| 152 | if strings.Contains(msg.Content, "blocked:") { |
| 153 | p.blocked = true |
| 154 | } |
| 155 | if strings.Contains(msg.Content, "try hook probe") { |
| 156 | p.leaked = true |
| 157 | } |
| 158 | } |
| 159 | p.mu.Unlock() |
| 160 | |
| 161 | chunks := []provider.Chunk{{Type: provider.ChunkText, Text: "done"}, {Type: provider.ChunkDone}} |
| 162 | if call == 0 { |
| 163 | chunks = []provider.Chunk{{Type: provider.ChunkToolCall, ToolCall: &provider.ToolCall{ID: "try-read", Name: "read_file", Arguments: `{"path":"marker.txt"}`}}} |
| 164 | } |
| 165 | ch := make(chan provider.Chunk, len(chunks)) |
| 166 | for _, chunk := range chunks { |
| 167 | ch <- chunk |
| 168 | } |
| 169 | close(ch) |
| 170 | return ch, nil |
| 171 | } |
| 172 |