Fix bypass broker and stop command
- Wire BypassBroker into fs.Config before creating root tree
- Move defer bypassBroker.Close() outside block scope (was closing immediately)
- Rewrite /bypass/ FS nodes to match toolsrv's expected interface:
- bypass/pending: blocking read returning JSON {id,cmd,cwd,env}
- bypass/resolve: write JSON {id,approved,error}
- Add NextPending() method to broker with channel-based blocking
- Add JSON tags to Request struct for proper serialization
- Start bypass loop for restored sessions (was only for new sessions)
- Add context cancellation support to CallTool (closes fid on cancel)
This commit is contained in:
parent
9bb3cd76b8
commit
3a68b0be53
|
|
@ -28,6 +28,8 @@ type Broker struct {
|
|||
|
||||
rateMu sync.Mutex
|
||||
limiters map[string]*rate.Limiter // sessionID -> denial rate limiter
|
||||
|
||||
pendingCh chan *Request // channel for blocking NextPending
|
||||
}
|
||||
|
||||
// BrokerConfig configures the broker.
|
||||
|
|
@ -58,6 +60,7 @@ func NewBroker(cfg BrokerConfig) (*Broker, error) {
|
|||
pending: make(map[string]*Request),
|
||||
sessions: make(map[string]*Policy),
|
||||
limiters: make(map[string]*rate.Limiter),
|
||||
pendingCh: make(chan *Request, 16),
|
||||
}
|
||||
|
||||
cfg.Logf("elevate: broker ready")
|
||||
|
|
@ -80,6 +83,11 @@ func (b *Broker) Pending() []*Request {
|
|||
return out
|
||||
}
|
||||
|
||||
// NextPending blocks until a bypass request is available and returns it.
|
||||
func (b *Broker) NextPending() *Request {
|
||||
return <-b.pendingCh
|
||||
}
|
||||
|
||||
// PendingByID returns a specific pending request.
|
||||
func (b *Broker) PendingByID(id string) *Request {
|
||||
b.mu.RLock()
|
||||
|
|
@ -214,6 +222,9 @@ func (b *Broker) EvaluateRequest(sessionID, cmd, cwd string, env map[string]stri
|
|||
b.pending[req.ID] = req
|
||||
b.mu.Unlock()
|
||||
|
||||
// Send to pending channel for NextPending readers
|
||||
b.pendingCh <- req
|
||||
|
||||
// Notify user (D-Bus notification, etc.)
|
||||
b.notify(req)
|
||||
b.logf("elevate: pending id=%s cmd=%q session=%s", req.ID, cmd, sessionID)
|
||||
|
|
|
|||
|
|
@ -34,12 +34,12 @@ func (r Resolution) String() string {
|
|||
|
||||
// Request represents a pending bypass request.
|
||||
type Request struct {
|
||||
ID string
|
||||
Cmd string
|
||||
Cwd string
|
||||
Env map[string]string
|
||||
SessionID string
|
||||
CreatedAt time.Time
|
||||
ID string `json:"id"`
|
||||
Cmd string `json:"cmd"`
|
||||
Cwd string `json:"cwd"`
|
||||
Env map[string]string `json:"env,omitempty"`
|
||||
SessionID string `json:"-"` // not serialized
|
||||
CreatedAt time.Time `json:"-"` // not serialized
|
||||
|
||||
// Resolution channel — exactly one value sent when resolved.
|
||||
resolved chan Resolution
|
||||
|
|
|
|||
|
|
@ -169,65 +169,49 @@ func buildTreeSpec(cfg *Config) virtfs.FsNodeDecl {
|
|||
// bypass/
|
||||
virtfs.DirNode("bypass",
|
||||
virtfs.Doc("Bypass broker for sandboxed command approval"),
|
||||
virtfs.FileNode("policy", 0666,
|
||||
virtfs.Doc("Global bypass policy"),
|
||||
virtfs.Read(func() ([]byte, error) {
|
||||
virtfs.FileNode("pending", 0444,
|
||||
virtfs.Doc("Blocks until bypass request; returns JSON {id, cmd, cwd, env}"),
|
||||
virtfs.BlockOnce(func(_ context.Context, _ string) ([]byte, string, error) {
|
||||
if broker == nil {
|
||||
return nil, errNoBypass
|
||||
return nil, "", errNoBypass
|
||||
}
|
||||
p := broker.GlobalPolicy().Global()
|
||||
return p.Marshal()
|
||||
req := broker.NextPending()
|
||||
if req == nil {
|
||||
return nil, "", fmt.Errorf("no pending request")
|
||||
}
|
||||
data, err := json.Marshal(req)
|
||||
if err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
return append(data, '\n'), "", nil
|
||||
}),
|
||||
),
|
||||
virtfs.FileNode("resolve", 0222,
|
||||
virtfs.Doc("Resolve bypass request: write JSON {id, approved} or {id, error}"),
|
||||
virtfs.Write(func(data []byte) error {
|
||||
if broker == nil {
|
||||
return errNoBypass
|
||||
}
|
||||
var p bypass.Policy
|
||||
if err := bypass.ParsePolicy(data, &p); err != nil {
|
||||
return err
|
||||
var msg struct {
|
||||
ID string `json:"id"`
|
||||
Approved bool `json:"approved"`
|
||||
Error string `json:"error"`
|
||||
}
|
||||
return broker.GlobalPolicy().SetGlobal(p)
|
||||
if err := json.Unmarshal(data, &msg); err != nil {
|
||||
return fmt.Errorf("invalid JSON: %w", err)
|
||||
}
|
||||
var res bypass.Resolution
|
||||
if msg.Approved {
|
||||
res = bypass.ResolveApprove
|
||||
} else {
|
||||
res = bypass.ResolveDeny
|
||||
}
|
||||
if !broker.Resolve(msg.ID, res) {
|
||||
return fmt.Errorf("unknown request ID: %s", msg.ID)
|
||||
}
|
||||
return nil
|
||||
}),
|
||||
),
|
||||
virtfs.Each("{reqid}", func() ([]virtfs.FsNodeDecl, error) {
|
||||
if broker == nil {
|
||||
return nil, errNoBypass
|
||||
}
|
||||
pending := broker.Pending()
|
||||
var out []virtfs.FsNodeDecl
|
||||
for _, r := range pending {
|
||||
req := r
|
||||
out = append(out, virtfs.FsNodeDecl{
|
||||
Name: req.ID,
|
||||
Children: []virtfs.FsNodeDecl{
|
||||
virtfs.FileNode("request", 0666,
|
||||
virtfs.Read(func() ([]byte, error) {
|
||||
return []byte(req.Summary() + "\n"), nil
|
||||
}),
|
||||
virtfs.Write(func(data []byte) error {
|
||||
cmd := strings.TrimSpace(string(data))
|
||||
var res bypass.Resolution
|
||||
switch cmd {
|
||||
case "approve":
|
||||
res = bypass.ResolveApprove
|
||||
case "deny":
|
||||
res = bypass.ResolveDeny
|
||||
case "persist":
|
||||
res = bypass.ResolvePersist
|
||||
default:
|
||||
return fmt.Errorf("unknown resolution: %s (use approve/deny/persist)", cmd)
|
||||
}
|
||||
if !broker.Resolve(req.ID, res) {
|
||||
return fmt.Errorf("request already resolved or timed out")
|
||||
}
|
||||
return nil
|
||||
}),
|
||||
),
|
||||
},
|
||||
})
|
||||
}
|
||||
return out, nil
|
||||
}),
|
||||
),
|
||||
|
||||
// session/
|
||||
|
|
|
|||
|
|
@ -256,6 +256,11 @@ func restoreMultiAgentSession(ps *PersistedSession) (*RestoredSession, error) {
|
|||
sess.Uname = sess.Agents()[0].ID()
|
||||
Register(sess.Name(), sess)
|
||||
|
||||
// Start bypass approval loop for restored sessions (if broker is configured)
|
||||
if pkgBypassBroker != nil && !ps.Paused {
|
||||
sess.StartBypassLoop(pkgBypassBroker)
|
||||
}
|
||||
|
||||
toolCount := 0
|
||||
if conn := sess.ToolsConn(); conn != nil {
|
||||
if infos, err := conn.ListTools(); err == nil {
|
||||
|
|
|
|||
|
|
@ -60,30 +60,14 @@ func runServer(sockPath string) {
|
|||
agentsDirs := agent.AgentsDirs()
|
||||
sessionsDir := paths.DataDir() + "/sessions"
|
||||
|
||||
// Elevate broker (initialized after manager; closures capture the pointer).
|
||||
var bypassBroker *bypass.Broker
|
||||
|
||||
daemonCtx, daemonCancel := context.WithCancel(context.Background())
|
||||
defer daemonCancel()
|
||||
|
||||
modelCache := fs.NewModelCache()
|
||||
|
||||
rootTree := fs.NewRoot(fs.Config{
|
||||
Ctx: daemonCtx,
|
||||
AgentsDir: agentsDirs[0],
|
||||
SessionsDir: sessionsDir,
|
||||
Log: sink.NewLogger("9p"),
|
||||
Sink: sink,
|
||||
Yolo: *yolo,
|
||||
ModelCache: modelCache,
|
||||
})
|
||||
|
||||
var srv *Server
|
||||
// Server creation deferred until after bypass broker is ready (see below)
|
||||
|
||||
// Start bypass broker
|
||||
// Start bypass broker (must be before NewRoot so it can be passed in)
|
||||
policyPath := filepath.Join(paths.DataDir(), "bypass-policy.yaml")
|
||||
|
||||
var bypassBroker *bypass.Broker
|
||||
{
|
||||
notifyFn := func(req *bypass.Request) {
|
||||
notifyBypass(req)
|
||||
|
|
@ -96,18 +80,31 @@ func runServer(sockPath string) {
|
|||
SessionValid: func(id string) bool { return session.Lookup(id) != nil },
|
||||
})
|
||||
if err != nil {
|
||||
fmt.Fprintf(os.Stderr, "warning: %v\n", err)
|
||||
fmt.Fprintf(os.Stderr, "warning: bypass broker: %v\n", err)
|
||||
} else {
|
||||
defer bypassBroker.Close()
|
||||
// Desktop notifications for bypass prompts (uses D-Bus notifications API directly)
|
||||
if conn, err := dbus.SessionBus(); err == nil {
|
||||
initBypassNotifier(conn, bypassBroker)
|
||||
}
|
||||
}
|
||||
}
|
||||
if bypassBroker != nil {
|
||||
defer bypassBroker.Close()
|
||||
}
|
||||
|
||||
// Create 9P server (after broker is ready so it can be passed in)
|
||||
srv = New(Config{
|
||||
rootTree := fs.NewRoot(fs.Config{
|
||||
Ctx: daemonCtx,
|
||||
AgentsDir: agentsDirs[0],
|
||||
SessionsDir: sessionsDir,
|
||||
Log: sink.NewLogger("9p"),
|
||||
Sink: sink,
|
||||
Yolo: *yolo,
|
||||
ModelCache: modelCache,
|
||||
BypassBroker: bypassBroker,
|
||||
})
|
||||
|
||||
// Create 9P server
|
||||
srv := New(Config{
|
||||
Sink: sink,
|
||||
RootTree: rootTree,
|
||||
})
|
||||
|
|
|
|||
|
|
@ -142,6 +142,17 @@ func (c *Conn) CallTool(ctx context.Context, name string, args json.RawMessage)
|
|||
if err != nil {
|
||||
return nil, fmt.Errorf("open proc/new: %w", err)
|
||||
}
|
||||
|
||||
// Close fid on context cancellation to unblock the read
|
||||
done := make(chan struct{})
|
||||
defer close(done)
|
||||
go func() {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
fid.Close()
|
||||
case <-done:
|
||||
}
|
||||
}()
|
||||
defer fid.Close()
|
||||
|
||||
// Convert JSON args to key=value format
|
||||
|
|
@ -157,6 +168,9 @@ func (c *Conn) CallTool(ctx context.Context, name string, args json.RawMessage)
|
|||
}
|
||||
|
||||
if _, err := fid.Write([]byte(payload)); err != nil {
|
||||
if ctx.Err() != nil {
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
return nil, fmt.Errorf("write: %w", err)
|
||||
}
|
||||
|
||||
|
|
@ -174,6 +188,9 @@ func (c *Conn) CallTool(ctx context.Context, name string, args json.RawMessage)
|
|||
break
|
||||
}
|
||||
if err != nil {
|
||||
if ctx.Err() != nil {
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
return nil, fmt.Errorf("read: %w", err)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in New Issue