@@ -9,13 +9,19 @@ import (
99 "math/rand"
1010 "net/http"
1111 "os"
12+ goruntime "runtime"
1213 "strings"
1314 "sync"
1415 "time"
1516
1617 "github.com/gorilla/websocket"
1718
1819 "neo-code/internal/tools"
20+ "neo-code/internal/tools/bash"
21+ diagnosetool "neo-code/internal/tools/diagnose"
22+ "neo-code/internal/tools/filesystem"
23+ "neo-code/internal/tools/todo"
24+ "neo-code/internal/tools/webfetch"
1925)
2026
2127// Runner 是本地执行守护进程,主动连接云端 Gateway,接收工具执行请求。
@@ -26,6 +32,7 @@ type Runner struct {
2632 capSigner * CapSigner
2733
2834 mu sync.Mutex
35+ writeMu sync.Mutex // 保护 WebSocket 并发写
2936 running bool
3037 cancel context.CancelFunc
3138}
@@ -38,6 +45,7 @@ type Config struct {
3845 Token string
3946 Workdir string
4047 WorkdirAllowlist []string
48+ Shell string // 用于 bash 工具,空值自动检测
4149 HeartbeatInterval time.Duration
4250 ReconnectBackoffMin time.Duration
4351 ReconnectBackoffMax time.Duration
@@ -71,7 +79,38 @@ func New(cfg Config) (*Runner, error) {
7179 logger = log .New (os .Stderr , "runner: " , log .LstdFlags )
7280 }
7381
82+ shell := cfg .Shell
83+ if shell == "" {
84+ if goruntime .GOOS == "windows" {
85+ shell = "cmd"
86+ } else {
87+ shell = "bash"
88+ }
89+ }
90+ workdir := cfg .Workdir
91+ if workdir == "" {
92+ var err error
93+ workdir , err = os .Getwd ()
94+ if err != nil {
95+ return nil , fmt .Errorf ("runner: get workdir: %w" , err )
96+ }
97+ }
98+
7499 toolMgr := tools .NewRegistry ()
100+ toolMgr .Register (filesystem .New (workdir ))
101+ toolMgr .Register (filesystem .NewWrite (workdir ))
102+ toolMgr .Register (filesystem .NewGrep (workdir ))
103+ toolMgr .Register (filesystem .NewGlob (workdir ))
104+ toolMgr .Register (filesystem .NewEdit (workdir ))
105+ toolMgr .Register (filesystem .NewMove (workdir ))
106+ toolMgr .Register (filesystem .NewCopy (workdir ))
107+ toolMgr .Register (filesystem .NewDelete (workdir ))
108+ toolMgr .Register (filesystem .NewCreateDir (workdir ))
109+ toolMgr .Register (filesystem .NewRemoveDir (workdir ))
110+ toolMgr .Register (bash .New (workdir , shell , cfg .RequestTimeout ))
111+ toolMgr .Register (webfetch .New (webfetch.Config {Timeout : cfg .RequestTimeout }))
112+ toolMgr .Register (diagnosetool .New ())
113+ toolMgr .Register (todo .New ())
75114
76115 capSigner := NewCapSigner (cfg .WorkdirAllowlist )
77116
@@ -130,12 +169,12 @@ func (r *Runner) Run(ctx context.Context) error {
130169
131170func (r * Runner ) connectAndServe (ctx context.Context ) error {
132171 url := fmt .Sprintf ("ws://%s/ws" , r .cfg .GatewayAddress )
133- if r .cfg .Token != "" {
134- url += "?token=" + r .cfg .Token
135- }
136172
137173 header := http.Header {}
138174 header .Set ("X-Runner-ID" , r .cfg .RunnerID )
175+ if r .cfg .Token != "" {
176+ header .Set ("Authorization" , "Bearer " + r .cfg .Token )
177+ }
139178
140179 dialer := websocket.Dialer {
141180 HandshakeTimeout : r .cfg .RequestTimeout ,
@@ -149,7 +188,7 @@ func (r *Runner) connectAndServe(ctx context.Context) error {
149188 }
150189 defer conn .Close ()
151190
152- r .logger .Printf ("connected to gateway at %s" , url )
191+ r .logger .Printf ("connected to gateway at %s (runner=%s) " , r . cfg . GatewayAddress , r . cfg . RunnerID )
153192
154193 // 认证
155194 if err := r .sendRequest (conn , "gateway.authenticate" , map [string ]string {
@@ -224,6 +263,24 @@ func (r *Runner) handleToolRequest(ctx context.Context, conn *websocket.Conn, ms
224263
225264 r .logger .Printf ("executing tool: %s (request_id=%s)" , req .ToolName , req .RequestID )
226265
266+ // 验证 capability token 和路径边界
267+ if err := r .capSigner .VerifyToolRequest (req , r .cfg .Workdir ); err != nil {
268+ r .logger .Printf ("tool request denied: %v" , err )
269+ resultParams := map [string ]any {
270+ "request_id" : req .RequestID ,
271+ "session_id" : req .SessionID ,
272+ "run_id" : req .RunID ,
273+ "runner_id" : r .cfg .RunnerID ,
274+ "tool_call_id" : req .ToolCallID ,
275+ "content" : fmt .Sprintf ("tool request denied: %v" , err ),
276+ "is_error" : true ,
277+ }
278+ if sendErr := r .sendRequest (conn , "gateway.executeToolResult" , resultParams ); sendErr != nil {
279+ r .logger .Printf ("failed to send denied result: %v" , sendErr )
280+ }
281+ return
282+ }
283+
227284 // 执行工具
228285 execCtx , cancel := context .WithTimeout (ctx , r .cfg .RequestTimeout )
229286 defer cancel ()
@@ -268,6 +325,8 @@ func (r *Runner) handlePing(conn *websocket.Conn, msg map[string]any) {
268325 "result" : "pong" ,
269326 }
270327 data , _ := json .Marshal (response )
328+ r .writeMu .Lock ()
329+ defer r .writeMu .Unlock ()
271330 if err := conn .WriteMessage (websocket .TextMessage , data ); err != nil {
272331 r .logger .Printf ("failed to send pong: %v" , err )
273332 }
@@ -304,6 +363,8 @@ func (r *Runner) sendRequest(conn *websocket.Conn, method string, params any) er
304363 return fmt .Errorf ("marshal request: %w" , err )
305364 }
306365
366+ r .writeMu .Lock ()
367+ defer r .writeMu .Unlock ()
307368 if err := conn .SetWriteDeadline (time .Now ().Add (r .cfg .RequestTimeout )); err != nil {
308369 return fmt .Errorf ("set write deadline: %w" , err )
309370 }
0 commit comments