mirror of
https://github.com/tiennm99/goclaw.git
synced 2026-10-03 07:12:50 +00:00
notify channel of pairing approval (#1508)
* notify channel of pairing approval * purge cache when revoked
This commit is contained in:
1 parent
68a6d1c6c9
commit
c09a72a3bf
6 files changed
+172
-20
No files matched your search
@@ -240,7 +240,9 @@ func wireChannelEventSubscribers(
|
||||
}
|
||||
})
|
||||
|
||||
// Wire pairing revocation → force disconnect active WebSocket sessions.
|
||||
// Wire pairing revocation → force disconnect active WebSocket sessions and
|
||||
// clear the in-memory group approval cache so a revoked group re-enters the
|
||||
// pairing gate on its next message instead of the bot replying as usual.
|
||||
msgBus.Subscribe(bus.TopicPairingRevoked, func(event bus.Event) {
|
||||
if event.Name != bus.EventPairingRevoked {
|
||||
return
|
||||
@@ -250,6 +252,13 @@ func wireChannelEventSubscribers(
|
||||
return
|
||||
}
|
||||
go server.DisconnectByPairing(payload.SenderID, payload.Channel)
|
||||
// Group pairings use "group:<chatID>" as sender ID (telegram) or
|
||||
// "<chatID>" (other channels); only group entries carry an
|
||||
// approvedGroups cache entry worth clearing.
|
||||
if groupChatID, isGroup := strings.CutPrefix(payload.SenderID, "group:"); isGroup {
|
||||
slog.Debug("pairing revoked, clearing group approval cache", "channel", payload.Channel, "chat_id", groupChatID)
|
||||
channelMgr.ClearGroupApproval(payload.Channel, groupChatID)
|
||||
}
|
||||
})
|
||||
|
||||
// Cascade: when an agent becomes inactive, disable its linked channel instances.
|
||||
|
||||
@@ -171,6 +171,23 @@ func (m *Manager) GetChannel(name string) (Channel, bool) {
|
||||
return channel, ok
|
||||
}
|
||||
|
||||
// ClearGroupApproval removes a chat from a channel's in-memory pairing
|
||||
// approval cache (BaseChannel.approvedGroups). Used when a group pairing is
|
||||
// revoked so the bot re-enters the pairing gate on the next message instead of
|
||||
// continuing to reply as if it were still approved. Channels that don't embed
|
||||
// BaseChannel are ignored.
|
||||
func (m *Manager) ClearGroupApproval(channelName, chatID string) {
|
||||
m.mu.RLock()
|
||||
ch, ok := m.channels[channelName]
|
||||
m.mu.RUnlock()
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if c, ok := ch.(interface{ ClearGroupApproval(string) }); ok {
|
||||
c.ClearGroupApproval(chatID)
|
||||
}
|
||||
}
|
||||
|
||||
// GetStatus returns the running status of all channels.
|
||||
func (m *Manager) GetStatus() map[string]any {
|
||||
m.mu.RLock()
|
||||
|
||||
@@ -68,6 +68,13 @@ func (m *mockPairingStore) setPaired(senderID, channel string) {
|
||||
m.pairedDevices[senderID][channel] = true
|
||||
}
|
||||
|
||||
// setUnpaired removes a mock paired relationship.
|
||||
func (m *mockPairingStore) setUnpaired(senderID, channel string) {
|
||||
if m.pairedDevices[senderID] != nil {
|
||||
delete(m.pairedDevices[senderID], channel)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCheckDMPolicy_PolicyDisabled rejects all messages.
|
||||
func TestCheckDMPolicy_PolicyDisabled(t *testing.T) {
|
||||
bc := NewBaseChannel("test", nil, nil)
|
||||
@@ -566,6 +573,59 @@ func TestCheckGroupPolicy_PairingMarksGroupApproved(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestManagerClearGroupApproval verifies the Manager clears the per-channel
|
||||
// in-memory approval cache used to short-circuit the pairing gate after a
|
||||
// pairing is revoked.
|
||||
func TestManagerClearGroupApproval(t *testing.T) {
|
||||
bc := newApprovalChannel("telegram")
|
||||
ps := newMockPairingStore()
|
||||
bc.SetPairingService(ps)
|
||||
|
||||
mgr := NewManager(nil)
|
||||
mgr.RegisterChannel("telegram", bc)
|
||||
|
||||
chatID := "chat_group_1"
|
||||
ps.setPaired("group:"+chatID, "telegram")
|
||||
|
||||
ctx := context.Background()
|
||||
if got := bc.CheckGroupPolicy(ctx, "user1", chatID, "pairing"); got != PolicyAllow {
|
||||
t.Fatalf("paired group policy = %v; want PolicyAllow", got)
|
||||
}
|
||||
if !bc.IsGroupApproved(chatID) {
|
||||
t.Fatal("group should be cached as approved before revocation")
|
||||
}
|
||||
|
||||
// Revoke → cache must be cleared so the next message re-enters the pairing gate.
|
||||
mgr.ClearGroupApproval("telegram", chatID)
|
||||
if bc.IsGroupApproved(chatID) {
|
||||
t.Fatal("group approval cache should be cleared after revocation")
|
||||
}
|
||||
|
||||
// DB row is gone → next policy check must return PolicyNeedsPairing.
|
||||
ps.setUnpaired("group:"+chatID, "telegram")
|
||||
if got := bc.CheckGroupPolicy(ctx, "user1", chatID, "pairing"); got != PolicyNeedsPairing {
|
||||
t.Fatalf("post-revoke group policy = %v; want PolicyNeedsPairing", got)
|
||||
}
|
||||
|
||||
// Unknown channel names are a no-op (no panic).
|
||||
mgr.ClearGroupApproval("nonexistent", chatID)
|
||||
}
|
||||
|
||||
// approvalChannel wraps BaseChannel so it satisfies channels.Channel in tests.
|
||||
type approvalChannel struct {
|
||||
*BaseChannel
|
||||
}
|
||||
|
||||
func newApprovalChannel(name string) *approvalChannel {
|
||||
return &approvalChannel{BaseChannel: NewBaseChannel(name, nil, nil)}
|
||||
}
|
||||
|
||||
func (c *approvalChannel) Send(context.Context, bus.OutboundMessage) error { return nil }
|
||||
|
||||
func (c *approvalChannel) Start(context.Context) error { return nil }
|
||||
|
||||
func (c *approvalChannel) Stop(context.Context) error { return nil }
|
||||
|
||||
// TestCheckDMPolicy_AllPolicies_TableDriven comprehensive table test.
|
||||
func TestCheckDMPolicy_AllPolicies_TableDriven(t *testing.T) {
|
||||
tests := []struct {
|
||||
|
||||
@@ -4,11 +4,16 @@ import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
mcpgo "github.com/mark3labs/mcp-go/mcp"
|
||||
mcpserver "github.com/mark3labs/mcp-go/server"
|
||||
|
||||
"github.com/nextlevelbuilder/goclaw/internal/bus"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/channels"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/config"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/store"
|
||||
"github.com/nextlevelbuilder/goclaw/internal/systemmessages"
|
||||
)
|
||||
|
||||
// validMCPSenderIDRe mirrors internal/gateway/methods/pairing.go's
|
||||
@@ -23,48 +28,91 @@ func isValidMCPSenderID(id string) bool {
|
||||
return len(id) <= maxSenderIDLen && validMCPSenderIDRe.MatchString(id)
|
||||
}
|
||||
|
||||
// pairingCRUDDeps bundles the dependencies pairing CRUD tools need, including
|
||||
// the channel notification side effects on approve (mirrors the WS
|
||||
// PairingMethods onApprove callback in cmd/gateway_channels_setup.go).
|
||||
type pairingCRUDDeps struct {
|
||||
pairing store.PairingStore
|
||||
channelMgr *channels.Manager
|
||||
msgBus *bus.MessageBus
|
||||
cfg *config.Config
|
||||
}
|
||||
|
||||
// notifyPairingApproved mirrors the WS onApprove callback
|
||||
// (cmd/gateway_channels_setup.go) so MCP-triggered approvals also notify the
|
||||
// paired channel. Browser/internal channels are skipped — UI polls directly.
|
||||
// Logs the send result so MCP operators can verify delivery.
|
||||
func notifyPairingApproved(ctx context.Context, deps pairingCRUDDeps, paired *store.PairedDeviceData) {
|
||||
if deps.channelMgr == nil || deps.msgBus == nil || deps.cfg == nil || paired == nil {
|
||||
return
|
||||
}
|
||||
if channels.IsInternalChannel(paired.Channel) {
|
||||
slog.Debug("pairing approved for internal channel, skipping notification", "channel", paired.Channel)
|
||||
return
|
||||
}
|
||||
botName := deps.cfg.ResolveDisplayName("default")
|
||||
msg := systemmessages.NewResolver(deps.cfg).Render("", systemmessages.KeyPairingApproved, systemmessages.Vars{
|
||||
"app_name": botName,
|
||||
})
|
||||
// Group pairings need group_id metadata so channels (e.g. Zalo) route to group API.
|
||||
if strings.HasPrefix(paired.SenderID, "group:") {
|
||||
deps.msgBus.PublishOutbound(bus.OutboundMessage{
|
||||
Channel: paired.Channel,
|
||||
ChatID: paired.ChatID,
|
||||
Content: msg,
|
||||
Metadata: map[string]string{"group_id": paired.ChatID},
|
||||
})
|
||||
slog.Info("pairing approval notification published", "channel", paired.Channel, "chat_id", paired.ChatID)
|
||||
return
|
||||
}
|
||||
if err := deps.channelMgr.SendToChannel(ctx, paired.Channel, paired.ChatID, msg); err != nil {
|
||||
slog.Warn("failed to send pairing approval notification", "channel", paired.Channel, "chat_id", paired.ChatID, "error", err)
|
||||
return
|
||||
}
|
||||
slog.Info("pairing approval notification sent", "channel", paired.Channel, "chat_id", paired.ChatID)
|
||||
}
|
||||
|
||||
// registerPairingCRUDTools registers the goclaw_pairing_device_* and
|
||||
// goclaw_pairing_browser_status MCP tools backed by store.PairingStore.
|
||||
// Mirrors internal/gateway/methods/pairing.go minus the approve-callback
|
||||
// (channel notification) and event-broadcast side effects, which are WS/bus
|
||||
// concerns not applicable to this standalone MCP surface.
|
||||
func registerPairingCRUDTools(srv *mcpserver.MCPServer, pairing store.PairingStore) {
|
||||
// Mirrors internal/gateway/methods/pairing.go, including the approve-callback
|
||||
// (channel notification) and event-broadcast side effects.
|
||||
func registerPairingCRUDTools(srv *mcpserver.MCPServer, deps pairingCRUDDeps) {
|
||||
srv.AddTool(mcpgo.NewTool("goclaw_pairing_device_request",
|
||||
mcpgo.WithDescription("Request a device pairing code."),
|
||||
mcpgo.WithString("sender_id", mcpgo.Required(), mcpgo.Description("Sender identifier.")),
|
||||
mcpgo.WithString("channel", mcpgo.Required(), mcpgo.Description("Channel name.")),
|
||||
mcpgo.WithString("chat_id", mcpgo.Description("Chat ID.")),
|
||||
mcpgo.WithString("account_id", mcpgo.Description("Account ID; defaults to \"default\".")),
|
||||
), handlePairingDeviceRequest(pairing))
|
||||
), handlePairingDeviceRequest(deps.pairing))
|
||||
|
||||
srv.AddTool(mcpgo.NewTool("goclaw_pairing_device_approve",
|
||||
mcpgo.WithDescription("Approve a pending pairing code."),
|
||||
mcpgo.WithString("code", mcpgo.Required(), mcpgo.Description("Pairing code.")),
|
||||
mcpgo.WithString("approved_by", mcpgo.Description("Approver identifier; defaults to \"operator\".")),
|
||||
), handlePairingDeviceApprove(pairing))
|
||||
), handlePairingDeviceApprove(deps))
|
||||
|
||||
srv.AddTool(mcpgo.NewTool("goclaw_pairing_device_deny",
|
||||
mcpgo.WithDescription("Deny a pending pairing code."),
|
||||
mcpgo.WithString("code", mcpgo.Required(), mcpgo.Description("Pairing code.")),
|
||||
), handlePairingDeviceDeny(pairing))
|
||||
), handlePairingDeviceDeny(deps.pairing))
|
||||
|
||||
srv.AddTool(mcpgo.NewTool("goclaw_pairing_device_list",
|
||||
mcpgo.WithDescription("List pending and paired devices."),
|
||||
mcpgo.WithReadOnlyHintAnnotation(true),
|
||||
), handlePairingDeviceList(pairing))
|
||||
), handlePairingDeviceList(deps.pairing))
|
||||
|
||||
srv.AddTool(mcpgo.NewTool("goclaw_pairing_device_revoke",
|
||||
mcpgo.WithDescription("Revoke an approved device pairing."),
|
||||
mcpgo.WithString("sender_id", mcpgo.Required(), mcpgo.Description("Sender identifier.")),
|
||||
mcpgo.WithString("channel", mcpgo.Required(), mcpgo.Description("Channel name.")),
|
||||
mcpgo.WithDestructiveHintAnnotation(true),
|
||||
), handlePairingDeviceRevoke(pairing))
|
||||
), handlePairingDeviceRevoke(deps))
|
||||
|
||||
srv.AddTool(mcpgo.NewTool("goclaw_pairing_browser_status",
|
||||
mcpgo.WithDescription("Check the pairing status for a pending browser client."),
|
||||
mcpgo.WithString("sender_id", mcpgo.Required(), mcpgo.Description("Sender identifier.")),
|
||||
mcpgo.WithReadOnlyHintAnnotation(true),
|
||||
), handlePairingBrowserStatus(pairing))
|
||||
), handlePairingBrowserStatus(deps.pairing))
|
||||
}
|
||||
|
||||
func handlePairingDeviceRequest(pairing store.PairingStore) mcpserver.ToolHandlerFunc {
|
||||
@@ -90,17 +138,18 @@ func handlePairingDeviceRequest(pairing store.PairingStore) mcpserver.ToolHandle
|
||||
}
|
||||
}
|
||||
|
||||
func handlePairingDeviceApprove(pairing store.PairingStore) mcpserver.ToolHandlerFunc {
|
||||
func handlePairingDeviceApprove(deps pairingCRUDDeps) mcpserver.ToolHandlerFunc {
|
||||
return func(ctx context.Context, req mcpgo.CallToolRequest) (*mcpgo.CallToolResult, error) {
|
||||
code, err := req.RequireString("code")
|
||||
if err != nil {
|
||||
return toolError("pairing.approve", err)
|
||||
}
|
||||
approvedBy := req.GetString("approved_by", "operator")
|
||||
paired, err := pairing.ApprovePairing(ctx, code, approvedBy)
|
||||
paired, err := deps.pairing.ApprovePairing(ctx, code, approvedBy)
|
||||
if err != nil {
|
||||
return toolError("pairing.approve", err)
|
||||
}
|
||||
notifyPairingApproved(ctx, deps, paired)
|
||||
return jsonToolResult(map[string]any{"paired": paired})
|
||||
}
|
||||
}
|
||||
@@ -127,7 +176,7 @@ func handlePairingDeviceList(pairing store.PairingStore) mcpserver.ToolHandlerFu
|
||||
}
|
||||
}
|
||||
|
||||
func handlePairingDeviceRevoke(pairing store.PairingStore) mcpserver.ToolHandlerFunc {
|
||||
func handlePairingDeviceRevoke(deps pairingCRUDDeps) mcpserver.ToolHandlerFunc {
|
||||
return func(ctx context.Context, req mcpgo.CallToolRequest) (*mcpgo.CallToolResult, error) {
|
||||
senderID, err := req.RequireString("sender_id")
|
||||
if err != nil {
|
||||
@@ -141,9 +190,21 @@ func handlePairingDeviceRevoke(pairing store.PairingStore) mcpserver.ToolHandler
|
||||
slog.Warn("security.invalid_sender_id_format", "handler", "mcp.pairing.revoke")
|
||||
return mcpgo.NewToolResultError("pairing.revoke: invalid sender_id format"), nil
|
||||
}
|
||||
if err := pairing.RevokePairing(ctx, senderID, channel); err != nil {
|
||||
if err := deps.pairing.RevokePairing(ctx, senderID, channel); err != nil {
|
||||
return toolError("pairing.revoke", err)
|
||||
}
|
||||
// Broadcast revocation so the gateway force-disconnects active WebSocket
|
||||
// sessions and clears the in-memory group approval cache (parity with
|
||||
// internal/gateway/methods/pairing.go's handleRevoke).
|
||||
if deps.msgBus != nil {
|
||||
deps.msgBus.Broadcast(bus.Event{
|
||||
Name: bus.EventPairingRevoked,
|
||||
Payload: bus.PairingRevokedPayload{
|
||||
SenderID: senderID,
|
||||
Channel: channel,
|
||||
},
|
||||
})
|
||||
}
|
||||
return jsonToolResult(map[string]bool{"revoked": true})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -10,7 +10,7 @@ import (
|
||||
func TestPairingDeviceRequest_HappyAndInvalidSenderID(t *testing.T) {
|
||||
pairing := newFakePairingStore()
|
||||
srv := newTestMCPServer()
|
||||
registerPairingCRUDTools(srv, pairing)
|
||||
registerPairingCRUDTools(srv, pairingCRUDDeps{pairing: pairing})
|
||||
|
||||
result := callTool(t, srv, "goclaw_pairing_device_request", map[string]any{
|
||||
"sender_id": "user-123", "channel": "telegram",
|
||||
@@ -28,7 +28,7 @@ func TestPairingDeviceRequest_HappyAndInvalidSenderID(t *testing.T) {
|
||||
func TestPairingDeviceApproveDenyList(t *testing.T) {
|
||||
pairing := newFakePairingStore()
|
||||
srv := newTestMCPServer()
|
||||
registerPairingCRUDTools(srv, pairing)
|
||||
registerPairingCRUDTools(srv, pairingCRUDDeps{pairing: pairing})
|
||||
|
||||
req := callTool(t, srv, "goclaw_pairing_device_request", map[string]any{"sender_id": "user-1", "channel": "telegram"})
|
||||
require.False(t, toolIsError(req))
|
||||
@@ -47,7 +47,7 @@ func TestPairingDeviceApproveDenyList(t *testing.T) {
|
||||
func TestPairingDeviceRevoke(t *testing.T) {
|
||||
pairing := newFakePairingStore()
|
||||
srv := newTestMCPServer()
|
||||
registerPairingCRUDTools(srv, pairing)
|
||||
registerPairingCRUDTools(srv, pairingCRUDDeps{pairing: pairing})
|
||||
|
||||
// Not paired: revoke should fail.
|
||||
result := callTool(t, srv, "goclaw_pairing_device_revoke", map[string]any{"sender_id": "user-9", "channel": "telegram"})
|
||||
@@ -57,7 +57,7 @@ func TestPairingDeviceRevoke(t *testing.T) {
|
||||
func TestPairingBrowserStatus_Expired(t *testing.T) {
|
||||
pairing := newFakePairingStore()
|
||||
srv := newTestMCPServer()
|
||||
registerPairingCRUDTools(srv, pairing)
|
||||
registerPairingCRUDTools(srv, pairingCRUDDeps{pairing: pairing})
|
||||
|
||||
result := callTool(t, srv, "goclaw_pairing_browser_status", map[string]any{"sender_id": "user-1"})
|
||||
require.False(t, toolIsError(result))
|
||||
|
||||
@@ -195,7 +195,12 @@ func NewCRUDServer(deps CRUDDeps, version string) *mcpserver.StreamableHTTPServe
|
||||
registered += 8
|
||||
}
|
||||
if deps.Pairing != nil {
|
||||
registerPairingCRUDTools(srv, deps.Pairing)
|
||||
registerPairingCRUDTools(srv, pairingCRUDDeps{
|
||||
pairing: deps.Pairing,
|
||||
channelMgr: deps.ChannelManager,
|
||||
msgBus: deps.MessageBus,
|
||||
cfg: deps.Config,
|
||||
})
|
||||
registered += 6
|
||||
}
|
||||
if deps.ExecApproval != nil {
|
||||
|
||||
Reference in new issue
Block a user