feat: implement channels with membership and MCP tools

Add channel management feature with public/private channels,
membership (owner/member roles), invitations, broadcast messaging,
and full MCP tool exposure.

New files:
- internal/channels/ - types, store, service, validation, errors
- internal/mcp/channel_tools.go - 8 MCP tools for channel operations
- schema/002_channels.sql - migration for channel_invites table

MCP tools: create_channel, join_channel, leave_channel, list_channels,
invite_to_channel, kick_from_channel, send_channel_message, update_channel

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Algis Dumbris
2026-03-13 11:56:50 +02:00
co-authored by Claude Opus 4.6
parent 2f55ce87c3
commit e3dcfece48
14 changed files with 2915 additions and 1 deletions
+5 -1
View File
@@ -14,6 +14,7 @@ import (
"github.com/spf13/cobra"
"github.com/smart-mcp-proxy/synapbus/internal/agents"
"github.com/smart-mcp-proxy/synapbus/internal/channels"
mcpserver "github.com/smart-mcp-proxy/synapbus/internal/mcp"
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
"github.com/smart-mcp-proxy/synapbus/internal/storage"
@@ -90,8 +91,11 @@ func runServe(cmd *cobra.Command, args []string) error {
agentStore := agents.NewSQLiteAgentStore(db.DB)
agentService := agents.NewAgentService(agentStore, tracer)
channelStore := channels.NewSQLiteChannelStore(db.DB)
channelService := channels.NewService(channelStore, msgService, tracer)
// Create MCP server
mcpSrv := mcpserver.NewMCPServer(msgService, agentService)
mcpSrv := mcpserver.NewMCPServer(msgService, agentService, channelService)
startTime := time.Now()
// Set up chi router
+14
View File
@@ -0,0 +1,14 @@
package channels
import "errors"
// Sentinel errors for channel operations.
var (
ErrChannelNotFound = errors.New("channel not found")
ErrChannelNameConflict = errors.New("channel name already exists")
ErrNotChannelMember = errors.New("not a channel member")
ErrNotChannelOwner = errors.New("not the channel owner")
ErrOwnerCannotLeave = errors.New("channel owner cannot leave; transfer ownership or delete the channel first")
ErrNotInvited = errors.New("not invited to this private channel")
ErrInvalidChannelName = errors.New("invalid channel name")
)
+437
View File
@@ -0,0 +1,437 @@
package channels
import (
"context"
"fmt"
"log/slog"
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
"github.com/smart-mcp-proxy/synapbus/internal/trace"
)
// Service provides business logic for channel operations.
type Service struct {
store ChannelStore
msgService *messaging.MessagingService
tracer *trace.Tracer
logger *slog.Logger
}
// NewService creates a new channel service.
func NewService(store ChannelStore, msgService *messaging.MessagingService, tracer *trace.Tracer) *Service {
return &Service{
store: store,
msgService: msgService,
tracer: tracer,
logger: slog.Default().With("component", "channels"),
}
}
// CreateChannel creates a new channel and adds the creator as the owner.
func (s *Service) CreateChannel(ctx context.Context, req CreateChannelRequest) (*Channel, error) {
// Validate name
if err := ValidateChannelName(req.Name); err != nil {
return nil, err
}
name := NormalizeChannelName(req.Name)
// Set default type
chType := req.Type
if chType == "" {
chType = TypeStandard
}
ch := &Channel{
Name: name,
Description: req.Description,
Topic: req.Topic,
Type: chType,
IsPrivate: req.IsPrivate,
CreatedBy: req.CreatedBy,
}
if err := s.store.CreateChannel(ctx, ch); err != nil {
return nil, err
}
// Auto-add creator as owner
member := &Membership{
ChannelID: ch.ID,
AgentName: req.CreatedBy,
Role: RoleOwner,
}
if err := s.store.AddMember(ctx, member); err != nil {
return nil, fmt.Errorf("add creator as owner: %w", err)
}
s.logger.Info("channel created",
"id", ch.ID,
"name", ch.Name,
"type", ch.Type,
"is_private", ch.IsPrivate,
"created_by", req.CreatedBy,
)
if s.tracer != nil {
s.tracer.Record(ctx, req.CreatedBy, "channel.create", map[string]any{
"channel_id": ch.ID,
"channel_name": ch.Name,
"type": ch.Type,
"is_private": ch.IsPrivate,
})
}
return ch, nil
}
// JoinChannel adds an agent to a channel.
func (s *Service) JoinChannel(ctx context.Context, channelID int64, agentName string) error {
ch, err := s.store.GetChannel(ctx, channelID)
if err != nil {
return err
}
// Check if already a member (idempotent)
isMember, err := s.store.IsMember(ctx, channelID, agentName)
if err != nil {
return fmt.Errorf("check membership: %w", err)
}
if isMember {
return nil // idempotent
}
// If private, check for a pending invite
if ch.IsPrivate {
hasInvite, err := s.store.HasPendingInvite(ctx, channelID, agentName)
if err != nil {
return fmt.Errorf("check invite: %w", err)
}
if !hasInvite {
return ErrNotInvited
}
// Accept the invite
if err := s.store.AcceptInvite(ctx, channelID, agentName); err != nil {
return fmt.Errorf("accept invite: %w", err)
}
}
member := &Membership{
ChannelID: channelID,
AgentName: agentName,
Role: RoleMember,
}
if err := s.store.AddMember(ctx, member); err != nil {
return fmt.Errorf("add member: %w", err)
}
s.logger.Info("agent joined channel",
"channel_id", channelID,
"agent", agentName,
)
if s.tracer != nil {
s.tracer.Record(ctx, agentName, "channel.join", map[string]any{
"channel_id": channelID,
"channel_name": ch.Name,
})
}
return nil
}
// LeaveChannel removes an agent from a channel.
func (s *Service) LeaveChannel(ctx context.Context, channelID int64, agentName string) error {
ch, err := s.store.GetChannel(ctx, channelID)
if err != nil {
return err
}
// Check membership and role
member, err := s.store.GetMember(ctx, channelID, agentName)
if err != nil {
return err
}
if member.Role == RoleOwner {
return ErrOwnerCannotLeave
}
if err := s.store.RemoveMember(ctx, channelID, agentName); err != nil {
return err
}
s.logger.Info("agent left channel",
"channel_id", channelID,
"agent", agentName,
)
if s.tracer != nil {
s.tracer.Record(ctx, agentName, "channel.leave", map[string]any{
"channel_id": channelID,
"channel_name": ch.Name,
})
}
return nil
}
// InviteToChannel invites an agent to a private channel. Only the owner can invite.
func (s *Service) InviteToChannel(ctx context.Context, channelID int64, agentName, inviterAgent string) error {
ch, err := s.store.GetChannel(ctx, channelID)
if err != nil {
return err
}
// Verify inviter is the owner
inviterMember, err := s.store.GetMember(ctx, channelID, inviterAgent)
if err != nil {
return err
}
if inviterMember.Role != RoleOwner {
return ErrNotChannelOwner
}
// Check if already a member (idempotent)
isMember, err := s.store.IsMember(ctx, channelID, agentName)
if err != nil {
return fmt.Errorf("check membership: %w", err)
}
if isMember {
return nil // already a member, no-op
}
inv := &ChannelInvite{
ChannelID: channelID,
AgentName: agentName,
InvitedBy: inviterAgent,
}
if err := s.store.CreateInvite(ctx, inv); err != nil {
return fmt.Errorf("create invite: %w", err)
}
s.logger.Info("agent invited to channel",
"channel_id", channelID,
"agent", agentName,
"invited_by", inviterAgent,
)
if s.tracer != nil {
s.tracer.Record(ctx, inviterAgent, "channel.invite", map[string]any{
"channel_id": channelID,
"channel_name": ch.Name,
"invitee": agentName,
})
}
return nil
}
// KickFromChannel removes an agent from a channel. Only the owner can kick.
func (s *Service) KickFromChannel(ctx context.Context, channelID int64, agentName, kickerAgent string) error {
ch, err := s.store.GetChannel(ctx, channelID)
if err != nil {
return err
}
// Verify kicker is the owner
kickerMember, err := s.store.GetMember(ctx, channelID, kickerAgent)
if err != nil {
return err
}
if kickerMember.Role != RoleOwner {
return ErrNotChannelOwner
}
// Cannot kick yourself
if agentName == kickerAgent {
return fmt.Errorf("cannot kick yourself from the channel")
}
// Verify target is a member
if err := s.store.RemoveMember(ctx, channelID, agentName); err != nil {
return err
}
s.logger.Info("agent kicked from channel",
"channel_id", channelID,
"agent", agentName,
"kicked_by", kickerAgent,
)
if s.tracer != nil {
s.tracer.Record(ctx, kickerAgent, "channel.kick", map[string]any{
"channel_id": channelID,
"channel_name": ch.Name,
"kicked": agentName,
})
}
return nil
}
// ListChannels returns channels visible to the agent.
func (s *Service) ListChannels(ctx context.Context, agentName string) ([]*ChannelWithCount, error) {
channels, err := s.store.ListChannels(ctx, agentName)
if err != nil {
return nil, err
}
result := make([]*ChannelWithCount, len(channels))
for i, ch := range channels {
count, err := s.store.CountMembers(ctx, ch.ID)
if err != nil {
return nil, fmt.Errorf("count members for channel %d: %w", ch.ID, err)
}
result[i] = &ChannelWithCount{
Channel: *ch,
MemberCount: count,
}
}
if s.tracer != nil {
s.tracer.Record(ctx, agentName, "channel.list", map[string]any{
"count": len(result),
})
}
return result, nil
}
// GetChannel returns a channel by ID.
func (s *Service) GetChannel(ctx context.Context, id int64) (*Channel, error) {
return s.store.GetChannel(ctx, id)
}
// GetChannelByName returns a channel by name.
func (s *Service) GetChannelByName(ctx context.Context, name string) (*Channel, error) {
return s.store.GetChannelByName(ctx, NormalizeChannelName(name))
}
// UpdateChannel updates a channel's metadata. Only the owner can update.
func (s *Service) UpdateChannel(ctx context.Context, channelID int64, req UpdateChannelRequest, agentName string) (*Channel, error) {
ch, err := s.store.GetChannel(ctx, channelID)
if err != nil {
return nil, err
}
// Verify caller is the owner
member, err := s.store.GetMember(ctx, channelID, agentName)
if err != nil {
return nil, err
}
if member.Role != RoleOwner {
return nil, ErrNotChannelOwner
}
// Apply updates
if req.Description != nil {
ch.Description = *req.Description
}
if req.Topic != nil {
ch.Topic = *req.Topic
}
if err := s.store.UpdateChannel(ctx, ch); err != nil {
return nil, err
}
// Reload to get updated timestamps
ch, err = s.store.GetChannel(ctx, channelID)
if err != nil {
return nil, err
}
s.logger.Info("channel updated",
"channel_id", channelID,
"agent", agentName,
)
if s.tracer != nil {
s.tracer.Record(ctx, agentName, "channel.update", map[string]any{
"channel_id": channelID,
"channel_name": ch.Name,
})
}
return ch, nil
}
// BroadcastMessage sends a message to all channel members except the sender.
func (s *Service) BroadcastMessage(ctx context.Context, channelID int64, fromAgent, body string, priority int, metadata string) ([]*messaging.Message, error) {
ch, err := s.store.GetChannel(ctx, channelID)
if err != nil {
return nil, err
}
// Verify sender is a member
isMember, err := s.store.IsMember(ctx, channelID, fromAgent)
if err != nil {
return nil, fmt.Errorf("check membership: %w", err)
}
if !isMember {
return nil, ErrNotChannelMember
}
// Get all members
members, err := s.store.GetMembers(ctx, channelID)
if err != nil {
return nil, fmt.Errorf("get members: %w", err)
}
var messages []*messaging.Message
for _, m := range members {
if m.AgentName == fromAgent {
continue // skip sender
}
// Build metadata with channel info
channelMeta := fmt.Sprintf(`{"channel_id":%d,"channel_name":%q}`, channelID, ch.Name)
if metadata != "" {
// Merge user metadata with channel metadata
channelMeta = fmt.Sprintf(`{"channel_id":%d,"channel_name":%q,"user_metadata":%s}`, channelID, ch.Name, metadata)
}
opts := messaging.SendOptions{
Subject: fmt.Sprintf("channel:%s", ch.Name),
Priority: priority,
Metadata: channelMeta,
}
msg, err := s.msgService.SendMessage(ctx, fromAgent, m.AgentName, body, opts)
if err != nil {
s.logger.Error("failed to send channel message",
"channel_id", channelID,
"from", fromAgent,
"to", m.AgentName,
"error", err,
)
continue // best effort — don't fail the whole broadcast
}
messages = append(messages, msg)
}
s.logger.Info("channel message broadcast",
"channel_id", channelID,
"from", fromAgent,
"recipients", len(messages),
)
if s.tracer != nil {
s.tracer.Record(ctx, fromAgent, "channel.broadcast", map[string]any{
"channel_id": channelID,
"channel_name": ch.Name,
"recipients": len(messages),
})
}
if messages == nil {
messages = []*messaging.Message{}
}
return messages, nil
}
// GetMembers returns all members of a channel.
func (s *Service) GetMembers(ctx context.Context, channelID int64) ([]*Membership, error) {
return s.store.GetMembers(ctx, channelID)
}
+638
View File
@@ -0,0 +1,638 @@
package channels
import (
"context"
"database/sql"
"errors"
"testing"
"time"
_ "modernc.org/sqlite"
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
"github.com/smart-mcp-proxy/synapbus/internal/storage"
"github.com/smart-mcp-proxy/synapbus/internal/trace"
)
func newTestService(t *testing.T) (*Service, *sql.DB) {
t.Helper()
db := newTestDB(t)
// Seed test agents
seedAgent(t, db, "agent-a")
seedAgent(t, db, "agent-b")
seedAgent(t, db, "agent-c")
channelStore := NewSQLiteChannelStore(db)
msgStore := messaging.NewSQLiteMessageStore(db)
tracer := trace.NewTracer(db)
t.Cleanup(func() { tracer.Close() })
msgService := messaging.NewMessagingService(msgStore, tracer)
svc := NewService(channelStore, msgService, tracer)
return svc, db
}
// --- CreateChannel tests ---
func TestService_CreateChannel(t *testing.T) {
tests := []struct {
name string
req CreateChannelRequest
wantErr bool
errIs error
}{
{
name: "create public channel",
req: CreateChannelRequest{
Name: "alerts",
Type: TypeStandard,
CreatedBy: "agent-a",
},
},
{
name: "create private channel",
req: CreateChannelRequest{
Name: "secret",
Type: TypeStandard,
IsPrivate: true,
CreatedBy: "agent-a",
},
},
{
name: "default type is standard",
req: CreateChannelRequest{
Name: "default-type",
CreatedBy: "agent-a",
},
},
{
name: "with description and topic",
req: CreateChannelRequest{
Name: "research",
Description: "Research channel",
Topic: "Q1 findings",
Type: TypeStandard,
CreatedBy: "agent-a",
},
},
{
name: "invalid name (empty)",
req: CreateChannelRequest{
Name: "",
Type: TypeStandard,
CreatedBy: "agent-a",
},
wantErr: true,
errIs: ErrInvalidChannelName,
},
{
name: "invalid name (special chars)",
req: CreateChannelRequest{
Name: "ch@nnel!",
Type: TypeStandard,
CreatedBy: "agent-a",
},
wantErr: true,
errIs: ErrInvalidChannelName,
},
{
name: "invalid name (too long)",
req: CreateChannelRequest{
Name: "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaX",
Type: TypeStandard,
CreatedBy: "agent-a",
},
wantErr: true,
errIs: ErrInvalidChannelName,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
ch, err := svc.CreateChannel(ctx, tt.req)
if tt.wantErr {
if err == nil {
t.Fatal("expected error, got nil")
}
if tt.errIs != nil && !errors.Is(err, tt.errIs) {
t.Errorf("expected error %v, got %v", tt.errIs, err)
}
return
}
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
if ch.ID == 0 {
t.Error("channel ID should not be 0")
}
if ch.Name != NormalizeChannelName(tt.req.Name) {
t.Errorf("name = %s, want %s", ch.Name, NormalizeChannelName(tt.req.Name))
}
})
}
}
func TestService_CreateChannel_CreatorIsOwner(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
ch, err := svc.CreateChannel(ctx, CreateChannelRequest{
Name: "owned-channel",
Type: TypeStandard,
CreatedBy: "agent-a",
})
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
members, err := svc.GetMembers(ctx, ch.ID)
if err != nil {
t.Fatalf("GetMembers: %v", err)
}
if len(members) != 1 {
t.Fatalf("got %d members, want 1", len(members))
}
if members[0].AgentName != "agent-a" {
t.Errorf("member = %s, want agent-a", members[0].AgentName)
}
if members[0].Role != RoleOwner {
t.Errorf("role = %s, want owner", members[0].Role)
}
}
func TestService_CreateChannel_DuplicateName(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
svc.CreateChannel(ctx, CreateChannelRequest{Name: "alerts", Type: TypeStandard, CreatedBy: "agent-a"})
_, err := svc.CreateChannel(ctx, CreateChannelRequest{Name: "alerts", Type: TypeStandard, CreatedBy: "agent-b"})
if !errors.Is(err, ErrChannelNameConflict) {
t.Errorf("expected ErrChannelNameConflict, got %v", err)
}
}
func TestService_CreateChannel_NameNormalized(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
ch, err := svc.CreateChannel(ctx, CreateChannelRequest{Name: "MyChannel", Type: TypeStandard, CreatedBy: "agent-a"})
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
if ch.Name != "mychannel" {
t.Errorf("name = %s, want mychannel", ch.Name)
}
}
// --- JoinChannel tests ---
func TestService_JoinChannel(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
ch, _ := svc.CreateChannel(ctx, CreateChannelRequest{Name: "alerts", Type: TypeStandard, CreatedBy: "agent-a"})
t.Run("join public channel", func(t *testing.T) {
err := svc.JoinChannel(ctx, ch.ID, "agent-b")
if err != nil {
t.Fatalf("JoinChannel: %v", err)
}
is, _ := svc.store.IsMember(ctx, ch.ID, "agent-b")
if !is {
t.Error("agent-b should be a member after joining")
}
})
t.Run("join is idempotent", func(t *testing.T) {
err := svc.JoinChannel(ctx, ch.ID, "agent-b")
if err != nil {
t.Errorf("second JoinChannel should be idempotent, got: %v", err)
}
})
t.Run("join non-existent channel fails", func(t *testing.T) {
err := svc.JoinChannel(ctx, 99999, "agent-b")
if !errors.Is(err, ErrChannelNotFound) {
t.Errorf("expected ErrChannelNotFound, got %v", err)
}
})
}
func TestService_JoinChannel_PrivateRequiresInvite(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
ch, _ := svc.CreateChannel(ctx, CreateChannelRequest{
Name: "private-ch",
Type: TypeStandard,
IsPrivate: true,
CreatedBy: "agent-a",
})
t.Run("uninvited agent cannot join private channel", func(t *testing.T) {
err := svc.JoinChannel(ctx, ch.ID, "agent-c")
if !errors.Is(err, ErrNotInvited) {
t.Errorf("expected ErrNotInvited, got %v", err)
}
})
t.Run("invited agent can join private channel", func(t *testing.T) {
svc.InviteToChannel(ctx, ch.ID, "agent-b", "agent-a")
err := svc.JoinChannel(ctx, ch.ID, "agent-b")
if err != nil {
t.Fatalf("JoinChannel: %v", err)
}
is, _ := svc.store.IsMember(ctx, ch.ID, "agent-b")
if !is {
t.Error("agent-b should be a member after invite+join")
}
})
t.Run("invite status changes to accepted after join", func(t *testing.T) {
inv, err := svc.store.GetInvite(ctx, ch.ID, "agent-b")
if err != nil {
t.Fatalf("GetInvite: %v", err)
}
if inv.Status != InviteStatusAccepted {
t.Errorf("invite status = %s, want accepted", inv.Status)
}
})
}
// --- LeaveChannel tests ---
func TestService_LeaveChannel(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
ch, _ := svc.CreateChannel(ctx, CreateChannelRequest{Name: "alerts", Type: TypeStandard, CreatedBy: "agent-a"})
svc.JoinChannel(ctx, ch.ID, "agent-b")
t.Run("member can leave", func(t *testing.T) {
err := svc.LeaveChannel(ctx, ch.ID, "agent-b")
if err != nil {
t.Fatalf("LeaveChannel: %v", err)
}
is, _ := svc.store.IsMember(ctx, ch.ID, "agent-b")
if is {
t.Error("agent-b should not be a member after leaving")
}
})
t.Run("owner cannot leave", func(t *testing.T) {
err := svc.LeaveChannel(ctx, ch.ID, "agent-a")
if !errors.Is(err, ErrOwnerCannotLeave) {
t.Errorf("expected ErrOwnerCannotLeave, got %v", err)
}
})
t.Run("non-member cannot leave", func(t *testing.T) {
err := svc.LeaveChannel(ctx, ch.ID, "agent-c")
if !errors.Is(err, ErrNotChannelMember) {
t.Errorf("expected ErrNotChannelMember, got %v", err)
}
})
t.Run("can rejoin public channel after leaving", func(t *testing.T) {
err := svc.JoinChannel(ctx, ch.ID, "agent-b")
if err != nil {
t.Fatalf("JoinChannel after leave: %v", err)
}
})
}
// --- InviteToChannel tests ---
func TestService_InviteToChannel(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
ch, _ := svc.CreateChannel(ctx, CreateChannelRequest{
Name: "private-ch",
Type: TypeStandard,
IsPrivate: true,
CreatedBy: "agent-a",
})
t.Run("owner can invite", func(t *testing.T) {
err := svc.InviteToChannel(ctx, ch.ID, "agent-b", "agent-a")
if err != nil {
t.Fatalf("InviteToChannel: %v", err)
}
has, _ := svc.store.HasPendingInvite(ctx, ch.ID, "agent-b")
if !has {
t.Error("expected pending invite for agent-b")
}
})
t.Run("non-owner cannot invite", func(t *testing.T) {
// First join agent-b
svc.JoinChannel(ctx, ch.ID, "agent-b")
err := svc.InviteToChannel(ctx, ch.ID, "agent-c", "agent-b")
if !errors.Is(err, ErrNotChannelOwner) {
t.Errorf("expected ErrNotChannelOwner, got %v", err)
}
})
t.Run("inviting already-member is idempotent", func(t *testing.T) {
err := svc.InviteToChannel(ctx, ch.ID, "agent-b", "agent-a")
if err != nil {
t.Errorf("invite existing member should be idempotent, got: %v", err)
}
})
t.Run("inviting to non-existent channel fails", func(t *testing.T) {
err := svc.InviteToChannel(ctx, 99999, "agent-c", "agent-a")
if !errors.Is(err, ErrChannelNotFound) {
t.Errorf("expected ErrChannelNotFound, got %v", err)
}
})
}
// --- KickFromChannel tests ---
func TestService_KickFromChannel(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
ch, _ := svc.CreateChannel(ctx, CreateChannelRequest{Name: "team", Type: TypeStandard, CreatedBy: "agent-a"})
svc.JoinChannel(ctx, ch.ID, "agent-b")
t.Run("owner can kick a member", func(t *testing.T) {
err := svc.KickFromChannel(ctx, ch.ID, "agent-b", "agent-a")
if err != nil {
t.Fatalf("KickFromChannel: %v", err)
}
is, _ := svc.store.IsMember(ctx, ch.ID, "agent-b")
if is {
t.Error("agent-b should not be a member after kick")
}
})
t.Run("non-owner cannot kick", func(t *testing.T) {
svc.JoinChannel(ctx, ch.ID, "agent-b")
err := svc.KickFromChannel(ctx, ch.ID, "agent-b", "agent-b")
if !errors.Is(err, ErrNotChannelOwner) {
t.Errorf("expected ErrNotChannelOwner, got %v", err)
}
})
t.Run("owner cannot kick themselves", func(t *testing.T) {
err := svc.KickFromChannel(ctx, ch.ID, "agent-a", "agent-a")
if err == nil {
t.Error("expected error when owner kicks themselves")
}
})
t.Run("kicking non-member fails", func(t *testing.T) {
err := svc.KickFromChannel(ctx, ch.ID, "agent-c", "agent-a")
if !errors.Is(err, ErrNotChannelMember) {
t.Errorf("expected ErrNotChannelMember, got %v", err)
}
})
}
// --- ListChannels tests ---
func TestService_ListChannels(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
t.Run("empty list when no channels", func(t *testing.T) {
result, err := svc.ListChannels(ctx, "agent-a")
if err != nil {
t.Fatalf("ListChannels: %v", err)
}
if len(result) != 0 {
t.Errorf("got %d channels, want 0", len(result))
}
})
// Create channels
svc.CreateChannel(ctx, CreateChannelRequest{Name: "public-1", Type: TypeStandard, CreatedBy: "agent-a"})
svc.CreateChannel(ctx, CreateChannelRequest{Name: "public-2", Type: TypeStandard, CreatedBy: "agent-a"})
privCh, _ := svc.CreateChannel(ctx, CreateChannelRequest{Name: "private-1", Type: TypeStandard, IsPrivate: true, CreatedBy: "agent-a"})
t.Run("lists public channels for outsider", func(t *testing.T) {
result, err := svc.ListChannels(ctx, "agent-c")
if err != nil {
t.Fatalf("ListChannels: %v", err)
}
if len(result) != 2 {
t.Errorf("got %d channels, want 2 (public only)", len(result))
}
})
t.Run("includes member count", func(t *testing.T) {
result, err := svc.ListChannels(ctx, "agent-a")
if err != nil {
t.Fatalf("ListChannels: %v", err)
}
for _, ch := range result {
if ch.MemberCount < 1 {
t.Errorf("channel %s has member_count %d, want >= 1", ch.Name, ch.MemberCount)
}
}
})
t.Run("includes private channel for member", func(t *testing.T) {
result, err := svc.ListChannels(ctx, "agent-a")
if err != nil {
t.Fatalf("ListChannels: %v", err)
}
if len(result) != 3 {
t.Errorf("got %d channels, want 3 (2 public + 1 private owned)", len(result))
}
})
t.Run("includes private channel for invited agent", func(t *testing.T) {
svc.InviteToChannel(ctx, privCh.ID, "agent-b", "agent-a")
result, err := svc.ListChannels(ctx, "agent-b")
if err != nil {
t.Fatalf("ListChannels: %v", err)
}
if len(result) != 3 {
t.Errorf("got %d channels, want 3 (2 public + 1 invited private)", len(result))
}
})
}
// --- UpdateChannel tests ---
func TestService_UpdateChannel(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
ch, _ := svc.CreateChannel(ctx, CreateChannelRequest{
Name: "research",
Description: "Research topics",
Topic: "Q1 findings",
Type: TypeStandard,
CreatedBy: "agent-a",
})
t.Run("owner can update topic", func(t *testing.T) {
newTopic := "Q2 planning"
updated, err := svc.UpdateChannel(ctx, ch.ID, UpdateChannelRequest{Topic: &newTopic}, "agent-a")
if err != nil {
t.Fatalf("UpdateChannel: %v", err)
}
if updated.Topic != "Q2 planning" {
t.Errorf("topic = %s, want Q2 planning", updated.Topic)
}
})
t.Run("owner can update description", func(t *testing.T) {
newDesc := "Updated description"
updated, err := svc.UpdateChannel(ctx, ch.ID, UpdateChannelRequest{Description: &newDesc}, "agent-a")
if err != nil {
t.Fatalf("UpdateChannel: %v", err)
}
if updated.Description != "Updated description" {
t.Errorf("description = %s, want Updated description", updated.Description)
}
})
t.Run("non-owner cannot update", func(t *testing.T) {
svc.JoinChannel(ctx, ch.ID, "agent-b")
newTopic := "Unauthorized"
_, err := svc.UpdateChannel(ctx, ch.ID, UpdateChannelRequest{Topic: &newTopic}, "agent-b")
if !errors.Is(err, ErrNotChannelOwner) {
t.Errorf("expected ErrNotChannelOwner, got %v", err)
}
})
t.Run("update non-existent channel fails", func(t *testing.T) {
newTopic := "test"
_, err := svc.UpdateChannel(ctx, 99999, UpdateChannelRequest{Topic: &newTopic}, "agent-a")
if !errors.Is(err, ErrChannelNotFound) {
t.Errorf("expected ErrChannelNotFound, got %v", err)
}
})
}
// --- BroadcastMessage tests ---
func TestService_BroadcastMessage(t *testing.T) {
svc, _ := newTestService(t)
ctx := context.Background()
ch, _ := svc.CreateChannel(ctx, CreateChannelRequest{Name: "alerts", Type: TypeStandard, CreatedBy: "agent-a"})
t.Run("broadcast to zero other members", func(t *testing.T) {
msgs, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "hello", 5, "")
if err != nil {
t.Fatalf("BroadcastMessage: %v", err)
}
if len(msgs) != 0 {
t.Errorf("got %d messages, want 0 (no other members)", len(msgs))
}
})
t.Run("broadcast to one member", func(t *testing.T) {
svc.JoinChannel(ctx, ch.ID, "agent-b")
msgs, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "hello everyone", 5, "")
if err != nil {
t.Fatalf("BroadcastMessage: %v", err)
}
if len(msgs) != 1 {
t.Errorf("got %d messages, want 1", len(msgs))
}
if msgs[0].ToAgent != "agent-b" {
t.Errorf("to_agent = %s, want agent-b", msgs[0].ToAgent)
}
})
t.Run("broadcast to multiple members", func(t *testing.T) {
svc.JoinChannel(ctx, ch.ID, "agent-c")
msgs, err := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "broadcast test", 5, "")
if err != nil {
t.Fatalf("BroadcastMessage: %v", err)
}
if len(msgs) != 2 {
t.Errorf("got %d messages, want 2", len(msgs))
}
})
t.Run("sender does not receive own message", func(t *testing.T) {
msgs, _ := svc.BroadcastMessage(ctx, ch.ID, "agent-a", "no self-message", 5, "")
for _, m := range msgs {
if m.ToAgent == "agent-a" {
t.Error("sender should not receive their own broadcast message")
}
}
})
t.Run("non-member cannot broadcast", func(t *testing.T) {
// agent-c is a member but let's test someone who isn't
seedAgent(t, svc.store.(*SQLiteChannelStore).db, "outsider")
_, err := svc.BroadcastMessage(ctx, ch.ID, "outsider", "unauthorized", 5, "")
if !errors.Is(err, ErrNotChannelMember) {
t.Errorf("expected ErrNotChannelMember, got %v", err)
}
})
t.Run("broadcast delivers to inbox", func(t *testing.T) {
svc.BroadcastMessage(ctx, ch.ID, "agent-a", "inbox test", 5, "")
// Read agent-b's inbox
msgs, err := svc.msgService.ReadInbox(ctx, "agent-b", messaging.ReadOptions{IncludeRead: true})
if err != nil {
t.Fatalf("ReadInbox: %v", err)
}
found := false
for _, m := range msgs {
if m.Body == "inbox test" && m.FromAgent == "agent-a" {
found = true
// Channel info is in metadata
if string(m.Metadata) == "{}" {
t.Error("message metadata should contain channel info")
}
break
}
}
if !found {
t.Error("agent-b should have received the broadcast message in inbox")
}
})
}
// --- Trace recording tests ---
func TestService_TracesRecorded(t *testing.T) {
svc, db := newTestService(t)
ctx := context.Background()
// Create + join should generate traces
ch, _ := svc.CreateChannel(ctx, CreateChannelRequest{Name: "traced", Type: TypeStandard, CreatedBy: "agent-a"})
svc.JoinChannel(ctx, ch.ID, "agent-b")
svc.LeaveChannel(ctx, ch.ID, "agent-b")
// Wait for async trace writes
var count int
for i := 0; i < 20; i++ {
time.Sleep(50 * time.Millisecond)
db.QueryRow("SELECT COUNT(*) FROM traces WHERE action LIKE 'channel.%'").Scan(&count)
if count >= 3 {
break
}
}
if count < 3 {
t.Errorf("expected at least 3 channel traces, got %d", count)
}
}
// suppress unused import warning
var _ = storage.RunMigrations
+349
View File
@@ -0,0 +1,349 @@
package channels
import (
"context"
"database/sql"
"fmt"
"log/slog"
"strings"
)
// ChannelStore defines the storage interface for channel operations.
type ChannelStore interface {
CreateChannel(ctx context.Context, ch *Channel) error
GetChannel(ctx context.Context, id int64) (*Channel, error)
GetChannelByName(ctx context.Context, name string) (*Channel, error)
ListChannels(ctx context.Context, agentName string) ([]*Channel, error)
UpdateChannel(ctx context.Context, ch *Channel) error
DeleteChannel(ctx context.Context, id int64) error
AddMember(ctx context.Context, m *Membership) error
RemoveMember(ctx context.Context, channelID int64, agentName string) error
GetMember(ctx context.Context, channelID int64, agentName string) (*Membership, error)
GetMembers(ctx context.Context, channelID int64) ([]*Membership, error)
IsMember(ctx context.Context, channelID int64, agentName string) (bool, error)
CountMembers(ctx context.Context, channelID int64) (int, error)
CreateInvite(ctx context.Context, inv *ChannelInvite) error
GetInvite(ctx context.Context, channelID int64, agentName string) (*ChannelInvite, error)
HasPendingInvite(ctx context.Context, channelID int64, agentName string) (bool, error)
AcceptInvite(ctx context.Context, channelID int64, agentName string) error
}
// SQLiteChannelStore implements ChannelStore using SQLite.
type SQLiteChannelStore struct {
db *sql.DB
logger *slog.Logger
}
// NewSQLiteChannelStore creates a new SQLite-backed channel store.
func NewSQLiteChannelStore(db *sql.DB) *SQLiteChannelStore {
return &SQLiteChannelStore{
db: db,
logger: slog.Default().With("component", "channel-store"),
}
}
// CreateChannel inserts a new channel.
func (s *SQLiteChannelStore) CreateChannel(ctx context.Context, ch *Channel) error {
isPrivate := 0
if ch.IsPrivate {
isPrivate = 1
}
result, err := s.db.ExecContext(ctx,
`INSERT INTO channels (name, description, topic, type, is_private, created_by, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`,
ch.Name, ch.Description, ch.Topic, ch.Type, isPrivate, ch.CreatedBy,
)
if err != nil {
if isUniqueConstraintError(err) {
return ErrChannelNameConflict
}
return fmt.Errorf("insert channel: %w", err)
}
id, err := result.LastInsertId()
if err != nil {
return fmt.Errorf("get channel id: %w", err)
}
ch.ID = id
s.logger.Info("channel created", "id", id, "name", ch.Name)
return nil
}
// GetChannel returns a channel by ID.
func (s *SQLiteChannelStore) GetChannel(ctx context.Context, id int64) (*Channel, error) {
var ch Channel
var isPrivate int
err := s.db.QueryRowContext(ctx,
`SELECT id, name, description, topic, type, is_private, created_by, created_at, updated_at
FROM channels WHERE id = ?`, id,
).Scan(&ch.ID, &ch.Name, &ch.Description, &ch.Topic, &ch.Type, &isPrivate, &ch.CreatedBy, &ch.CreatedAt, &ch.UpdatedAt)
if err != nil {
if err == sql.ErrNoRows {
return nil, ErrChannelNotFound
}
return nil, fmt.Errorf("get channel: %w", err)
}
ch.IsPrivate = isPrivate != 0
return &ch, nil
}
// GetChannelByName returns a channel by name (case-insensitive).
func (s *SQLiteChannelStore) GetChannelByName(ctx context.Context, name string) (*Channel, error) {
var ch Channel
var isPrivate int
err := s.db.QueryRowContext(ctx,
`SELECT id, name, description, topic, type, is_private, created_by, created_at, updated_at
FROM channels WHERE LOWER(name) = LOWER(?)`, name,
).Scan(&ch.ID, &ch.Name, &ch.Description, &ch.Topic, &ch.Type, &isPrivate, &ch.CreatedBy, &ch.CreatedAt, &ch.UpdatedAt)
if err != nil {
if err == sql.ErrNoRows {
return nil, ErrChannelNotFound
}
return nil, fmt.Errorf("get channel by name: %w", err)
}
ch.IsPrivate = isPrivate != 0
return &ch, nil
}
// ListChannels returns all public channels plus private channels where the agent
// is a member or has a pending invite.
func (s *SQLiteChannelStore) ListChannels(ctx context.Context, agentName string) ([]*Channel, error) {
rows, err := s.db.QueryContext(ctx,
`SELECT DISTINCT c.id, c.name, c.description, c.topic, c.type, c.is_private, c.created_by, c.created_at, c.updated_at
FROM channels c
WHERE c.is_private = 0
OR EXISTS (SELECT 1 FROM channel_members cm WHERE cm.channel_id = c.id AND cm.agent_name = ?)
OR EXISTS (SELECT 1 FROM channel_invites ci WHERE ci.channel_id = c.id AND ci.agent_name = ? AND ci.status = 'pending')
ORDER BY c.name`,
agentName, agentName,
)
if err != nil {
return nil, fmt.Errorf("list channels: %w", err)
}
defer rows.Close()
var channels []*Channel
for rows.Next() {
var ch Channel
var isPrivate int
if err := rows.Scan(&ch.ID, &ch.Name, &ch.Description, &ch.Topic, &ch.Type, &isPrivate, &ch.CreatedBy, &ch.CreatedAt, &ch.UpdatedAt); err != nil {
return nil, fmt.Errorf("scan channel: %w", err)
}
ch.IsPrivate = isPrivate != 0
channels = append(channels, &ch)
}
if channels == nil {
channels = []*Channel{}
}
return channels, rows.Err()
}
// UpdateChannel updates a channel's mutable fields.
func (s *SQLiteChannelStore) UpdateChannel(ctx context.Context, ch *Channel) error {
result, err := s.db.ExecContext(ctx,
`UPDATE channels SET description = ?, topic = ?, updated_at = CURRENT_TIMESTAMP WHERE id = ?`,
ch.Description, ch.Topic, ch.ID,
)
if err != nil {
return fmt.Errorf("update channel: %w", err)
}
rowsAffected, _ := result.RowsAffected()
if rowsAffected == 0 {
return ErrChannelNotFound
}
s.logger.Info("channel updated", "id", ch.ID)
return nil
}
// DeleteChannel deletes a channel by ID.
func (s *SQLiteChannelStore) DeleteChannel(ctx context.Context, id int64) error {
result, err := s.db.ExecContext(ctx, `DELETE FROM channels WHERE id = ?`, id)
if err != nil {
return fmt.Errorf("delete channel: %w", err)
}
rowsAffected, _ := result.RowsAffected()
if rowsAffected == 0 {
return ErrChannelNotFound
}
s.logger.Info("channel deleted", "id", id)
return nil
}
// AddMember adds a member to a channel.
func (s *SQLiteChannelStore) AddMember(ctx context.Context, m *Membership) error {
result, err := s.db.ExecContext(ctx,
`INSERT INTO channel_members (channel_id, agent_name, role, joined_at)
VALUES (?, ?, ?, CURRENT_TIMESTAMP)`,
m.ChannelID, m.AgentName, m.Role,
)
if err != nil {
if isUniqueConstraintError(err) {
// Already a member — idempotent
return nil
}
return fmt.Errorf("add member: %w", err)
}
id, _ := result.LastInsertId()
m.ID = id
s.logger.Info("member added", "channel_id", m.ChannelID, "agent", m.AgentName, "role", m.Role)
return nil
}
// RemoveMember removes a member from a channel.
func (s *SQLiteChannelStore) RemoveMember(ctx context.Context, channelID int64, agentName string) error {
result, err := s.db.ExecContext(ctx,
`DELETE FROM channel_members WHERE channel_id = ? AND agent_name = ?`,
channelID, agentName,
)
if err != nil {
return fmt.Errorf("remove member: %w", err)
}
rowsAffected, _ := result.RowsAffected()
if rowsAffected == 0 {
return ErrNotChannelMember
}
s.logger.Info("member removed", "channel_id", channelID, "agent", agentName)
return nil
}
// GetMember returns a channel membership record.
func (s *SQLiteChannelStore) GetMember(ctx context.Context, channelID int64, agentName string) (*Membership, error) {
var m Membership
err := s.db.QueryRowContext(ctx,
`SELECT id, channel_id, agent_name, role, joined_at
FROM channel_members WHERE channel_id = ? AND agent_name = ?`,
channelID, agentName,
).Scan(&m.ID, &m.ChannelID, &m.AgentName, &m.Role, &m.JoinedAt)
if err != nil {
if err == sql.ErrNoRows {
return nil, ErrNotChannelMember
}
return nil, fmt.Errorf("get member: %w", err)
}
return &m, nil
}
// GetMembers returns all members of a channel.
func (s *SQLiteChannelStore) GetMembers(ctx context.Context, channelID int64) ([]*Membership, error) {
rows, err := s.db.QueryContext(ctx,
`SELECT id, channel_id, agent_name, role, joined_at
FROM channel_members WHERE channel_id = ? ORDER BY joined_at`,
channelID,
)
if err != nil {
return nil, fmt.Errorf("get members: %w", err)
}
defer rows.Close()
var members []*Membership
for rows.Next() {
var m Membership
if err := rows.Scan(&m.ID, &m.ChannelID, &m.AgentName, &m.Role, &m.JoinedAt); err != nil {
return nil, fmt.Errorf("scan member: %w", err)
}
members = append(members, &m)
}
if members == nil {
members = []*Membership{}
}
return members, rows.Err()
}
// IsMember checks if an agent is a member of a channel.
func (s *SQLiteChannelStore) IsMember(ctx context.Context, channelID int64, agentName string) (bool, error) {
var count int
err := s.db.QueryRowContext(ctx,
`SELECT COUNT(*) FROM channel_members WHERE channel_id = ? AND agent_name = ?`,
channelID, agentName,
).Scan(&count)
if err != nil {
return false, fmt.Errorf("check membership: %w", err)
}
return count > 0, nil
}
// CountMembers returns the number of members in a channel.
func (s *SQLiteChannelStore) CountMembers(ctx context.Context, channelID int64) (int, error) {
var count int
err := s.db.QueryRowContext(ctx,
`SELECT COUNT(*) FROM channel_members WHERE channel_id = ?`,
channelID,
).Scan(&count)
if err != nil {
return 0, fmt.Errorf("count members: %w", err)
}
return count, nil
}
// CreateInvite creates an invitation for an agent to join a private channel.
func (s *SQLiteChannelStore) CreateInvite(ctx context.Context, inv *ChannelInvite) error {
result, err := s.db.ExecContext(ctx,
`INSERT INTO channel_invites (channel_id, agent_name, invited_by, created_at, status)
VALUES (?, ?, ?, CURRENT_TIMESTAMP, 'pending')
ON CONFLICT(channel_id, agent_name) DO UPDATE SET
status = 'pending',
invited_by = excluded.invited_by`,
inv.ChannelID, inv.AgentName, inv.InvitedBy,
)
if err != nil {
return fmt.Errorf("create invite: %w", err)
}
id, _ := result.LastInsertId()
inv.ID = id
s.logger.Info("invite created", "channel_id", inv.ChannelID, "agent", inv.AgentName, "invited_by", inv.InvitedBy)
return nil
}
// GetInvite returns an invite by channel and agent.
func (s *SQLiteChannelStore) GetInvite(ctx context.Context, channelID int64, agentName string) (*ChannelInvite, error) {
var inv ChannelInvite
err := s.db.QueryRowContext(ctx,
`SELECT id, channel_id, agent_name, invited_by, created_at, status
FROM channel_invites WHERE channel_id = ? AND agent_name = ?`,
channelID, agentName,
).Scan(&inv.ID, &inv.ChannelID, &inv.AgentName, &inv.InvitedBy, &inv.CreatedAt, &inv.Status)
if err != nil {
if err == sql.ErrNoRows {
return nil, ErrNotInvited
}
return nil, fmt.Errorf("get invite: %w", err)
}
return &inv, nil
}
// HasPendingInvite checks if an agent has a pending invite to a channel.
func (s *SQLiteChannelStore) HasPendingInvite(ctx context.Context, channelID int64, agentName string) (bool, error) {
var count int
err := s.db.QueryRowContext(ctx,
`SELECT COUNT(*) FROM channel_invites WHERE channel_id = ? AND agent_name = ? AND status = 'pending'`,
channelID, agentName,
).Scan(&count)
if err != nil {
return false, fmt.Errorf("check invite: %w", err)
}
return count > 0, nil
}
// AcceptInvite marks a pending invite as accepted.
func (s *SQLiteChannelStore) AcceptInvite(ctx context.Context, channelID int64, agentName string) error {
result, err := s.db.ExecContext(ctx,
`UPDATE channel_invites SET status = 'accepted' WHERE channel_id = ? AND agent_name = ? AND status = 'pending'`,
channelID, agentName,
)
if err != nil {
return fmt.Errorf("accept invite: %w", err)
}
rowsAffected, _ := result.RowsAffected()
if rowsAffected == 0 {
return ErrNotInvited
}
s.logger.Info("invite accepted", "channel_id", channelID, "agent", agentName)
return nil
}
// isUniqueConstraintError checks if an error is a SQLite unique constraint violation.
func isUniqueConstraintError(err error) bool {
return strings.Contains(err.Error(), "UNIQUE constraint failed")
}
+471
View File
@@ -0,0 +1,471 @@
package channels
import (
"context"
"database/sql"
"fmt"
"testing"
_ "modernc.org/sqlite"
"github.com/smart-mcp-proxy/synapbus/internal/storage"
)
func newTestDB(t *testing.T) *sql.DB {
t.Helper()
dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name())
db, err := sql.Open("sqlite", dsn)
if err != nil {
t.Fatalf("open database: %v", err)
}
t.Cleanup(func() { db.Close() })
if _, err := db.Exec("PRAGMA foreign_keys=ON"); err != nil {
t.Fatalf("enable foreign keys: %v", err)
}
ctx := context.Background()
if err := storage.RunMigrations(ctx, db); err != nil {
t.Fatalf("run migrations: %v", err)
}
// Seed test user
db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (1, 'testowner', 'hash', 'Test Owner')`)
return db
}
func seedAgent(t *testing.T, db *sql.DB, name string) {
t.Helper()
db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (1, 'testowner', 'hash', 'Test Owner')`)
_, err := db.Exec(
`INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES (?, ?, 'ai', '{}', 1, 'testhash', 'active')`,
name, name,
)
if err != nil {
t.Fatalf("seed agent %s: %v", name, err)
}
}
func TestSQLiteChannelStore_CreateChannel(t *testing.T) {
tests := []struct {
name string
channel Channel
wantErr error
}{
{
name: "create public channel",
channel: Channel{
Name: "alerts",
Description: "System alerts",
Topic: "Current alerts",
Type: TypeStandard,
IsPrivate: false,
CreatedBy: "agent-a",
},
},
{
name: "create private channel",
channel: Channel{
Name: "core-team",
Description: "Core team only",
Type: TypeStandard,
IsPrivate: true,
CreatedBy: "agent-a",
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
db := newTestDB(t)
store := NewSQLiteChannelStore(db)
ctx := context.Background()
ch := tt.channel
err := store.CreateChannel(ctx, &ch)
if err != nil {
t.Fatalf("CreateChannel: %v", err)
}
if ch.ID == 0 {
t.Error("channel ID should not be 0")
}
// Verify retrieval
got, err := store.GetChannel(ctx, ch.ID)
if err != nil {
t.Fatalf("GetChannel: %v", err)
}
if got.Name != tt.channel.Name {
t.Errorf("name = %s, want %s", got.Name, tt.channel.Name)
}
if got.IsPrivate != tt.channel.IsPrivate {
t.Errorf("is_private = %v, want %v", got.IsPrivate, tt.channel.IsPrivate)
}
})
}
}
func TestSQLiteChannelStore_CreateChannel_DuplicateName(t *testing.T) {
db := newTestDB(t)
store := NewSQLiteChannelStore(db)
ctx := context.Background()
ch1 := &Channel{Name: "alerts", Type: TypeStandard, CreatedBy: "agent-a"}
if err := store.CreateChannel(ctx, ch1); err != nil {
t.Fatalf("CreateChannel 1: %v", err)
}
ch2 := &Channel{Name: "alerts", Type: TypeStandard, CreatedBy: "agent-b"}
err := store.CreateChannel(ctx, ch2)
if err != ErrChannelNameConflict {
t.Errorf("expected ErrChannelNameConflict, got %v", err)
}
}
func TestSQLiteChannelStore_GetChannelByName(t *testing.T) {
db := newTestDB(t)
store := NewSQLiteChannelStore(db)
ctx := context.Background()
ch := &Channel{Name: "alerts", Type: TypeStandard, CreatedBy: "agent-a"}
store.CreateChannel(ctx, ch)
tests := []struct {
name string
lookup string
wantErr bool
}{
{"exact match", "alerts", false},
{"case insensitive", "ALERTS", false},
{"not found", "nonexistent", true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := store.GetChannelByName(ctx, tt.lookup)
if tt.wantErr {
if err == nil {
t.Error("expected error, got nil")
}
return
}
if err != nil {
t.Fatalf("GetChannelByName: %v", err)
}
if got.Name != "alerts" {
t.Errorf("name = %s, want alerts", got.Name)
}
})
}
}
func TestSQLiteChannelStore_GetChannel_NotFound(t *testing.T) {
db := newTestDB(t)
store := NewSQLiteChannelStore(db)
ctx := context.Background()
_, err := store.GetChannel(ctx, 99999)
if err != ErrChannelNotFound {
t.Errorf("expected ErrChannelNotFound, got %v", err)
}
}
func TestSQLiteChannelStore_ListChannels(t *testing.T) {
db := newTestDB(t)
store := NewSQLiteChannelStore(db)
ctx := context.Background()
// Create public channels
pub1 := &Channel{Name: "alerts", Type: TypeStandard, IsPrivate: false, CreatedBy: "agent-a"}
pub2 := &Channel{Name: "general", Type: TypeStandard, IsPrivate: false, CreatedBy: "agent-a"}
store.CreateChannel(ctx, pub1)
store.CreateChannel(ctx, pub2)
// Create private channel
priv := &Channel{Name: "secret", Type: TypeStandard, IsPrivate: true, CreatedBy: "agent-a"}
store.CreateChannel(ctx, priv)
t.Run("lists public channels for uninvited agent", func(t *testing.T) {
channels, err := store.ListChannels(ctx, "agent-b")
if err != nil {
t.Fatalf("ListChannels: %v", err)
}
if len(channels) != 2 {
t.Errorf("got %d channels, want 2 (public only)", len(channels))
}
})
t.Run("includes private channel when agent is a member", func(t *testing.T) {
store.AddMember(ctx, &Membership{ChannelID: priv.ID, AgentName: "agent-b", Role: RoleMember})
channels, err := store.ListChannels(ctx, "agent-b")
if err != nil {
t.Fatalf("ListChannels: %v", err)
}
if len(channels) != 3 {
t.Errorf("got %d channels, want 3 (public + member of private)", len(channels))
}
})
t.Run("includes private channel when agent has pending invite", func(t *testing.T) {
priv2 := &Channel{Name: "invited-only", Type: TypeStandard, IsPrivate: true, CreatedBy: "agent-a"}
store.CreateChannel(ctx, priv2)
inv := &ChannelInvite{ChannelID: priv2.ID, AgentName: "agent-c", InvitedBy: "agent-a"}
store.CreateInvite(ctx, inv)
channels, err := store.ListChannels(ctx, "agent-c")
if err != nil {
t.Fatalf("ListChannels: %v", err)
}
if len(channels) != 3 {
t.Errorf("got %d channels, want 3 (2 public + 1 invited private)", len(channels))
}
})
t.Run("empty list when no channels", func(t *testing.T) {
db2 := newTestDB(t)
store2 := NewSQLiteChannelStore(db2)
channels, err := store2.ListChannels(ctx, "agent-x")
if err != nil {
t.Fatalf("ListChannels: %v", err)
}
if len(channels) != 0 {
t.Errorf("got %d channels, want 0", len(channels))
}
})
}
func TestSQLiteChannelStore_UpdateChannel(t *testing.T) {
db := newTestDB(t)
store := NewSQLiteChannelStore(db)
ctx := context.Background()
ch := &Channel{Name: "research", Topic: "Q1 findings", Type: TypeStandard, CreatedBy: "agent-a"}
store.CreateChannel(ctx, ch)
ch.Topic = "Q2 planning"
ch.Description = "Updated description"
if err := store.UpdateChannel(ctx, ch); err != nil {
t.Fatalf("UpdateChannel: %v", err)
}
got, err := store.GetChannel(ctx, ch.ID)
if err != nil {
t.Fatalf("GetChannel: %v", err)
}
if got.Topic != "Q2 planning" {
t.Errorf("topic = %s, want Q2 planning", got.Topic)
}
if got.Description != "Updated description" {
t.Errorf("description = %s, want Updated description", got.Description)
}
}
func TestSQLiteChannelStore_DeleteChannel(t *testing.T) {
db := newTestDB(t)
store := NewSQLiteChannelStore(db)
ctx := context.Background()
ch := &Channel{Name: "temp", Type: TypeStandard, CreatedBy: "agent-a"}
store.CreateChannel(ctx, ch)
if err := store.DeleteChannel(ctx, ch.ID); err != nil {
t.Fatalf("DeleteChannel: %v", err)
}
_, err := store.GetChannel(ctx, ch.ID)
if err != ErrChannelNotFound {
t.Errorf("expected ErrChannelNotFound after delete, got %v", err)
}
}
func TestSQLiteChannelStore_Membership(t *testing.T) {
db := newTestDB(t)
store := NewSQLiteChannelStore(db)
ctx := context.Background()
ch := &Channel{Name: "test-ch", Type: TypeStandard, CreatedBy: "agent-a"}
store.CreateChannel(ctx, ch)
t.Run("add member", func(t *testing.T) {
m := &Membership{ChannelID: ch.ID, AgentName: "agent-a", Role: RoleOwner}
if err := store.AddMember(ctx, m); err != nil {
t.Fatalf("AddMember: %v", err)
}
if m.ID == 0 {
t.Error("membership ID should not be 0")
}
})
t.Run("add duplicate member is idempotent", func(t *testing.T) {
m := &Membership{ChannelID: ch.ID, AgentName: "agent-a", Role: RoleMember}
err := store.AddMember(ctx, m)
if err != nil {
t.Fatalf("AddMember duplicate: %v", err)
}
})
t.Run("is member", func(t *testing.T) {
yes, err := store.IsMember(ctx, ch.ID, "agent-a")
if err != nil {
t.Fatalf("IsMember: %v", err)
}
if !yes {
t.Error("expected agent-a to be a member")
}
no, err := store.IsMember(ctx, ch.ID, "agent-x")
if err != nil {
t.Fatalf("IsMember: %v", err)
}
if no {
t.Error("expected agent-x to NOT be a member")
}
})
t.Run("get member", func(t *testing.T) {
m, err := store.GetMember(ctx, ch.ID, "agent-a")
if err != nil {
t.Fatalf("GetMember: %v", err)
}
if m.Role != RoleOwner {
t.Errorf("role = %s, want owner", m.Role)
}
})
t.Run("get member not found", func(t *testing.T) {
_, err := store.GetMember(ctx, ch.ID, "nonexistent")
if err != ErrNotChannelMember {
t.Errorf("expected ErrNotChannelMember, got %v", err)
}
})
t.Run("count members", func(t *testing.T) {
store.AddMember(ctx, &Membership{ChannelID: ch.ID, AgentName: "agent-b", Role: RoleMember})
count, err := store.CountMembers(ctx, ch.ID)
if err != nil {
t.Fatalf("CountMembers: %v", err)
}
if count != 2 {
t.Errorf("count = %d, want 2", count)
}
})
t.Run("get members", func(t *testing.T) {
members, err := store.GetMembers(ctx, ch.ID)
if err != nil {
t.Fatalf("GetMembers: %v", err)
}
if len(members) != 2 {
t.Errorf("got %d members, want 2", len(members))
}
})
t.Run("remove member", func(t *testing.T) {
if err := store.RemoveMember(ctx, ch.ID, "agent-b"); err != nil {
t.Fatalf("RemoveMember: %v", err)
}
is, _ := store.IsMember(ctx, ch.ID, "agent-b")
if is {
t.Error("agent-b should no longer be a member")
}
})
t.Run("remove non-member fails", func(t *testing.T) {
err := store.RemoveMember(ctx, ch.ID, "nonexistent")
if err != ErrNotChannelMember {
t.Errorf("expected ErrNotChannelMember, got %v", err)
}
})
}
func TestSQLiteChannelStore_Invites(t *testing.T) {
db := newTestDB(t)
store := NewSQLiteChannelStore(db)
ctx := context.Background()
ch := &Channel{Name: "private-ch", Type: TypeStandard, IsPrivate: true, CreatedBy: "agent-a"}
store.CreateChannel(ctx, ch)
t.Run("create invite", func(t *testing.T) {
inv := &ChannelInvite{ChannelID: ch.ID, AgentName: "agent-b", InvitedBy: "agent-a"}
if err := store.CreateInvite(ctx, inv); err != nil {
t.Fatalf("CreateInvite: %v", err)
}
})
t.Run("has pending invite", func(t *testing.T) {
has, err := store.HasPendingInvite(ctx, ch.ID, "agent-b")
if err != nil {
t.Fatalf("HasPendingInvite: %v", err)
}
if !has {
t.Error("expected pending invite for agent-b")
}
has, err = store.HasPendingInvite(ctx, ch.ID, "agent-c")
if err != nil {
t.Fatalf("HasPendingInvite: %v", err)
}
if has {
t.Error("expected no pending invite for agent-c")
}
})
t.Run("get invite", func(t *testing.T) {
inv, err := store.GetInvite(ctx, ch.ID, "agent-b")
if err != nil {
t.Fatalf("GetInvite: %v", err)
}
if inv.Status != InviteStatusPending {
t.Errorf("status = %s, want pending", inv.Status)
}
if inv.InvitedBy != "agent-a" {
t.Errorf("invited_by = %s, want agent-a", inv.InvitedBy)
}
})
t.Run("accept invite", func(t *testing.T) {
if err := store.AcceptInvite(ctx, ch.ID, "agent-b"); err != nil {
t.Fatalf("AcceptInvite: %v", err)
}
inv, err := store.GetInvite(ctx, ch.ID, "agent-b")
if err != nil {
t.Fatalf("GetInvite after accept: %v", err)
}
if inv.Status != InviteStatusAccepted {
t.Errorf("status = %s, want accepted", inv.Status)
}
// Should no longer be pending
has, _ := store.HasPendingInvite(ctx, ch.ID, "agent-b")
if has {
t.Error("should not have pending invite after acceptance")
}
})
t.Run("duplicate invite is idempotent (resets to pending)", func(t *testing.T) {
inv := &ChannelInvite{ChannelID: ch.ID, AgentName: "agent-b", InvitedBy: "agent-a"}
if err := store.CreateInvite(ctx, inv); err != nil {
t.Fatalf("CreateInvite duplicate: %v", err)
}
has, _ := store.HasPendingInvite(ctx, ch.ID, "agent-b")
if !has {
t.Error("re-invite should create a pending invite again")
}
})
t.Run("accept non-existent invite fails", func(t *testing.T) {
err := store.AcceptInvite(ctx, ch.ID, "nonexistent")
if err != ErrNotInvited {
t.Errorf("expected ErrNotInvited, got %v", err)
}
})
}
// suppress unused import warning
var _ = storage.RunMigrations
+91
View File
@@ -0,0 +1,91 @@
// Package channels provides channel management for SynapBus.
package channels
import "time"
// ChannelType constants.
const (
TypeStandard = "standard"
TypeBlackboard = "blackboard"
TypeAuction = "auction"
)
// MemberRole constants.
const (
RoleOwner = "owner"
RoleMember = "member"
)
// InviteStatus constants.
const (
InviteStatusPending = "pending"
InviteStatusAccepted = "accepted"
InviteStatusDeclined = "declined"
)
// Channel represents a named group communication space.
type Channel struct {
ID int64 `json:"id"`
Name string `json:"name"`
Description string `json:"description"`
Topic string `json:"topic"`
Type string `json:"type"`
IsPrivate bool `json:"is_private"`
CreatedBy string `json:"created_by"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// ChannelWithCount embeds Channel and adds a member count.
type ChannelWithCount struct {
Channel
MemberCount int `json:"member_count"`
}
// Membership represents the relationship between an agent and a channel.
type Membership struct {
ID int64 `json:"id"`
ChannelID int64 `json:"channel_id"`
AgentName string `json:"agent_name"`
Role string `json:"role"`
JoinedAt time.Time `json:"joined_at"`
}
// ChannelInvite represents a pending invitation for an agent to join a private channel.
type ChannelInvite struct {
ID int64 `json:"id"`
ChannelID int64 `json:"channel_id"`
AgentName string `json:"agent_name"`
InvitedBy string `json:"invited_by"`
CreatedAt time.Time `json:"created_at"`
Status string `json:"status"`
}
// CreateChannelRequest is the input for creating a channel.
type CreateChannelRequest struct {
Name string `json:"name"`
Description string `json:"description"`
Topic string `json:"topic"`
Type string `json:"type"`
IsPrivate bool `json:"is_private"`
CreatedBy string `json:"created_by"`
}
// UpdateChannelRequest is the input for updating a channel.
type UpdateChannelRequest struct {
Description *string `json:"description,omitempty"`
Topic *string `json:"topic,omitempty"`
}
// JoinChannelRequest is the input for joining a channel.
type JoinChannelRequest struct {
ChannelID int64 `json:"channel_id"`
AgentName string `json:"agent_name"`
}
// InviteRequest is the input for inviting an agent to a channel.
type InviteRequest struct {
ChannelID int64 `json:"channel_id"`
AgentName string `json:"agent_name"`
InviterAgent string `json:"inviter_agent"`
}
+30
View File
@@ -0,0 +1,30 @@
package channels
import (
"fmt"
"regexp"
"strings"
)
var channelNameRegex = regexp.MustCompile(`^[a-zA-Z0-9][a-zA-Z0-9_-]*$`)
// ValidateChannelName validates a channel name.
// Names must be alphanumeric plus hyphens and underscores, max 64 characters.
// Names are normalized to lowercase.
func ValidateChannelName(name string) error {
if name == "" {
return ErrInvalidChannelName
}
if len(name) > 64 {
return fmt.Errorf("%w: name exceeds 64 characters", ErrInvalidChannelName)
}
if !channelNameRegex.MatchString(name) {
return fmt.Errorf("%w: name must be alphanumeric with hyphens and underscores", ErrInvalidChannelName)
}
return nil
}
// NormalizeChannelName returns the lowercase form of a channel name.
func NormalizeChannelName(name string) string {
return strings.ToLower(name)
}
+69
View File
@@ -0,0 +1,69 @@
package channels
import (
"errors"
"testing"
)
func TestValidateChannelName(t *testing.T) {
tests := []struct {
name string
input string
wantErr bool
}{
{"valid simple", "alerts", false},
{"valid with hyphens", "my-channel", false},
{"valid with underscores", "my_channel", false},
{"valid mixed", "my-channel_123", false},
{"valid single char", "a", false},
{"valid numbers", "123", false},
{"empty", "", true},
{"spaces", "my channel", true},
{"special chars", "ch@nnel!", true},
{"starts with hyphen", "-channel", true},
{"starts with underscore", "_channel", true},
{"unicode", "ch\u00e4nnel", true},
{"dot", "my.channel", true},
{"too long (65 chars)", "aaaaaaaaaabbbbbbbbbbccccccccccddddddddddeeeeeeeeeeffffffffffggggg", true},
{"exactly 64 chars", "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa", false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
err := ValidateChannelName(tt.input)
if tt.wantErr {
if err == nil {
t.Error("expected error, got nil")
}
if !errors.Is(err, ErrInvalidChannelName) {
t.Errorf("expected ErrInvalidChannelName, got %v", err)
}
} else {
if err != nil {
t.Errorf("unexpected error: %v", err)
}
}
})
}
}
func TestNormalizeChannelName(t *testing.T) {
tests := []struct {
input string
want string
}{
{"alerts", "alerts"},
{"ALERTS", "alerts"},
{"MyChannel", "mychannel"},
{"Mixed-Case_123", "mixed-case_123"},
}
for _, tt := range tests {
t.Run(tt.input, func(t *testing.T) {
got := NormalizeChannelName(tt.input)
if got != tt.want {
t.Errorf("NormalizeChannelName(%q) = %q, want %q", tt.input, got, tt.want)
}
})
}
}
+373
View File
@@ -0,0 +1,373 @@
package mcp
import (
"context"
"fmt"
"github.com/mark3labs/mcp-go/mcp"
"github.com/mark3labs/mcp-go/server"
"github.com/smart-mcp-proxy/synapbus/internal/channels"
)
// ChannelToolRegistrar registers channel MCP tools on the server.
type ChannelToolRegistrar struct {
channelService *channels.Service
}
// NewChannelToolRegistrar creates a new channel tool registrar.
func NewChannelToolRegistrar(channelService *channels.Service) *ChannelToolRegistrar {
return &ChannelToolRegistrar{
channelService: channelService,
}
}
// RegisterAll registers all channel tools on the MCP server.
func (ctr *ChannelToolRegistrar) RegisterAll(s *server.MCPServer) {
s.AddTool(ctr.createChannelTool(), ctr.handleCreateChannel)
s.AddTool(ctr.joinChannelTool(), ctr.handleJoinChannel)
s.AddTool(ctr.leaveChannelTool(), ctr.handleLeaveChannel)
s.AddTool(ctr.listChannelsTool(), ctr.handleListChannels)
s.AddTool(ctr.inviteToChannelTool(), ctr.handleInviteToChannel)
s.AddTool(ctr.kickFromChannelTool(), ctr.handleKickFromChannel)
s.AddTool(ctr.sendChannelMessageTool(), ctr.handleSendChannelMessage)
s.AddTool(ctr.updateChannelTool(), ctr.handleUpdateChannel)
}
// --- Tool Definitions ---
func (ctr *ChannelToolRegistrar) createChannelTool() mcp.Tool {
return mcp.NewTool("create_channel",
mcp.WithDescription("Create a new channel for group communication"),
mcp.WithString("name", mcp.Description("Unique channel name (alphanumeric, hyphens, underscores, max 64 chars)"), mcp.Required()),
mcp.WithString("description", mcp.Description("Channel description")),
mcp.WithString("topic", mcp.Description("Current channel topic")),
mcp.WithString("type", mcp.Description("Channel type: 'standard', 'blackboard', or 'auction' (default 'standard')")),
mcp.WithBoolean("is_private", mcp.Description("Whether the channel is private (invite-only). Default false")),
)
}
func (ctr *ChannelToolRegistrar) joinChannelTool() mcp.Tool {
return mcp.NewTool("join_channel",
mcp.WithDescription("Join an existing channel"),
mcp.WithNumber("channel_id", mcp.Description("ID of the channel to join")),
mcp.WithString("channel_name", mcp.Description("Name of the channel to join (alternative to channel_id)")),
)
}
func (ctr *ChannelToolRegistrar) leaveChannelTool() mcp.Tool {
return mcp.NewTool("leave_channel",
mcp.WithDescription("Leave a channel you are a member of"),
mcp.WithNumber("channel_id", mcp.Description("ID of the channel to leave")),
mcp.WithString("channel_name", mcp.Description("Name of the channel to leave (alternative to channel_id)")),
)
}
func (ctr *ChannelToolRegistrar) listChannelsTool() mcp.Tool {
return mcp.NewTool("list_channels",
mcp.WithDescription("List all channels visible to the authenticated agent (all public channels plus private channels you are a member of or have been invited to)"),
)
}
func (ctr *ChannelToolRegistrar) inviteToChannelTool() mcp.Tool {
return mcp.NewTool("invite_to_channel",
mcp.WithDescription("Invite an agent to a channel (only the channel owner can invite to private channels)"),
mcp.WithNumber("channel_id", mcp.Description("ID of the channel")),
mcp.WithString("channel_name", mcp.Description("Name of the channel (alternative to channel_id)")),
mcp.WithString("agent_name", mcp.Description("Name of the agent to invite"), mcp.Required()),
)
}
func (ctr *ChannelToolRegistrar) kickFromChannelTool() mcp.Tool {
return mcp.NewTool("kick_from_channel",
mcp.WithDescription("Remove an agent from a channel (only the channel owner can kick)"),
mcp.WithNumber("channel_id", mcp.Description("ID of the channel")),
mcp.WithString("channel_name", mcp.Description("Name of the channel (alternative to channel_id)")),
mcp.WithString("agent_name", mcp.Description("Name of the agent to kick"), mcp.Required()),
)
}
func (ctr *ChannelToolRegistrar) sendChannelMessageTool() mcp.Tool {
return mcp.NewTool("send_channel_message",
mcp.WithDescription("Send a message to all members of a channel"),
mcp.WithNumber("channel_id", mcp.Description("ID of the channel")),
mcp.WithString("channel_name", mcp.Description("Name of the channel (alternative to channel_id)")),
mcp.WithString("body", mcp.Description("Message body text"), mcp.Required()),
mcp.WithNumber("priority", mcp.Description("Message priority (1-10, default 5)"), mcp.Min(1), mcp.Max(10)),
mcp.WithString("metadata", mcp.Description("JSON metadata object (optional)")),
)
}
func (ctr *ChannelToolRegistrar) updateChannelTool() mcp.Tool {
return mcp.NewTool("update_channel",
mcp.WithDescription("Update channel topic or description (only the channel owner can update)"),
mcp.WithNumber("channel_id", mcp.Description("ID of the channel")),
mcp.WithString("channel_name", mcp.Description("Name of the channel (alternative to channel_id)")),
mcp.WithString("topic", mcp.Description("New channel topic")),
mcp.WithString("description", mcp.Description("New channel description")),
)
}
// --- Tool Handlers ---
func (ctr *ChannelToolRegistrar) handleCreateChannel(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
agentName, ok := extractAgentName(ctx)
if !ok {
return mcp.NewToolResultError("authentication required"), nil
}
name := req.GetString("name", "")
if name == "" {
return mcp.NewToolResultError("'name' parameter is required"), nil
}
isPrivate := false
args := req.GetArguments()
if v, ok := args["is_private"]; ok {
if b, ok := v.(bool); ok {
isPrivate = b
}
}
createReq := channels.CreateChannelRequest{
Name: name,
Description: req.GetString("description", ""),
Topic: req.GetString("topic", ""),
Type: req.GetString("type", "standard"),
IsPrivate: isPrivate,
CreatedBy: agentName,
}
ch, err := ctr.channelService.CreateChannel(ctx, createReq)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("create_channel failed: %s", err)), nil
}
return resultJSON(map[string]any{
"channel_id": ch.ID,
"name": ch.Name,
"description": ch.Description,
"topic": ch.Topic,
"type": ch.Type,
"is_private": ch.IsPrivate,
"created_by": ch.CreatedBy,
})
}
func (ctr *ChannelToolRegistrar) handleJoinChannel(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
agentName, ok := extractAgentName(ctx)
if !ok {
return mcp.NewToolResultError("authentication required"), nil
}
channelID, err := ctr.resolveChannelID(ctx, req)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("join_channel failed: %s", err)), nil
}
if err := ctr.channelService.JoinChannel(ctx, channelID, agentName); err != nil {
return mcp.NewToolResultError(fmt.Sprintf("join_channel failed: %s", err)), nil
}
return resultJSON(map[string]any{
"channel_id": channelID,
"status": "joined",
})
}
func (ctr *ChannelToolRegistrar) handleLeaveChannel(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
agentName, ok := extractAgentName(ctx)
if !ok {
return mcp.NewToolResultError("authentication required"), nil
}
channelID, err := ctr.resolveChannelID(ctx, req)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("leave_channel failed: %s", err)), nil
}
if err := ctr.channelService.LeaveChannel(ctx, channelID, agentName); err != nil {
return mcp.NewToolResultError(fmt.Sprintf("leave_channel failed: %s", err)), nil
}
return resultJSON(map[string]any{
"channel_id": channelID,
"status": "left",
})
}
func (ctr *ChannelToolRegistrar) handleListChannels(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
agentName, ok := extractAgentName(ctx)
if !ok {
return mcp.NewToolResultError("authentication required"), nil
}
chList, err := ctr.channelService.ListChannels(ctx, agentName)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("list_channels failed: %s", err)), nil
}
result := make([]map[string]any, len(chList))
for i, ch := range chList {
result[i] = map[string]any{
"id": ch.ID,
"name": ch.Name,
"description": ch.Description,
"topic": ch.Topic,
"type": ch.Type,
"is_private": ch.IsPrivate,
"created_by": ch.CreatedBy,
"member_count": ch.MemberCount,
}
}
return resultJSON(map[string]any{
"channels": result,
"count": len(result),
})
}
func (ctr *ChannelToolRegistrar) handleInviteToChannel(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
agentName, ok := extractAgentName(ctx)
if !ok {
return mcp.NewToolResultError("authentication required"), nil
}
channelID, err := ctr.resolveChannelID(ctx, req)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("invite_to_channel failed: %s", err)), nil
}
targetAgent := req.GetString("agent_name", "")
if targetAgent == "" {
return mcp.NewToolResultError("'agent_name' parameter is required"), nil
}
if err := ctr.channelService.InviteToChannel(ctx, channelID, targetAgent, agentName); err != nil {
return mcp.NewToolResultError(fmt.Sprintf("invite_to_channel failed: %s", err)), nil
}
return resultJSON(map[string]any{
"channel_id": channelID,
"agent_name": targetAgent,
"status": "invited",
})
}
func (ctr *ChannelToolRegistrar) handleKickFromChannel(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
agentName, ok := extractAgentName(ctx)
if !ok {
return mcp.NewToolResultError("authentication required"), nil
}
channelID, err := ctr.resolveChannelID(ctx, req)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("kick_from_channel failed: %s", err)), nil
}
targetAgent := req.GetString("agent_name", "")
if targetAgent == "" {
return mcp.NewToolResultError("'agent_name' parameter is required"), nil
}
if err := ctr.channelService.KickFromChannel(ctx, channelID, targetAgent, agentName); err != nil {
return mcp.NewToolResultError(fmt.Sprintf("kick_from_channel failed: %s", err)), nil
}
return resultJSON(map[string]any{
"channel_id": channelID,
"agent_name": targetAgent,
"status": "kicked",
})
}
func (ctr *ChannelToolRegistrar) handleSendChannelMessage(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
agentName, ok := extractAgentName(ctx)
if !ok {
return mcp.NewToolResultError("authentication required"), nil
}
channelID, err := ctr.resolveChannelID(ctx, req)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("send_channel_message failed: %s", err)), nil
}
body := req.GetString("body", "")
if body == "" {
return mcp.NewToolResultError("'body' parameter is required"), nil
}
priority := req.GetInt("priority", 5)
metadata := req.GetString("metadata", "")
messages, err := ctr.channelService.BroadcastMessage(ctx, channelID, agentName, body, priority, metadata)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("send_channel_message failed: %s", err)), nil
}
msgIDs := make([]int64, len(messages))
for i, m := range messages {
msgIDs[i] = m.ID
}
return resultJSON(map[string]any{
"channel_id": channelID,
"recipients": len(messages),
"message_ids": msgIDs,
})
}
func (ctr *ChannelToolRegistrar) handleUpdateChannel(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
agentName, ok := extractAgentName(ctx)
if !ok {
return mcp.NewToolResultError("authentication required"), nil
}
channelID, err := ctr.resolveChannelID(ctx, req)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("update_channel failed: %s", err)), nil
}
updateReq := channels.UpdateChannelRequest{}
args := req.GetArguments()
if v, ok := args["topic"]; ok {
if s, ok := v.(string); ok {
updateReq.Topic = &s
}
}
if v, ok := args["description"]; ok {
if s, ok := v.(string); ok {
updateReq.Description = &s
}
}
ch, err := ctr.channelService.UpdateChannel(ctx, channelID, updateReq, agentName)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("update_channel failed: %s", err)), nil
}
return resultJSON(map[string]any{
"channel_id": ch.ID,
"name": ch.Name,
"description": ch.Description,
"topic": ch.Topic,
})
}
// resolveChannelID resolves a channel ID from either channel_id or channel_name parameter.
func (ctr *ChannelToolRegistrar) resolveChannelID(ctx context.Context, req mcp.CallToolRequest) (int64, error) {
if cid := req.GetInt("channel_id", 0); cid > 0 {
return int64(cid), nil
}
name := req.GetString("channel_name", "")
if name != "" {
ch, err := ctr.channelService.GetChannelByName(ctx, name)
if err != nil {
return 0, err
}
return ch.ID, nil
}
return 0, fmt.Errorf("either 'channel_id' or 'channel_name' is required")
}
+402
View File
@@ -0,0 +1,402 @@
package mcp
import (
"context"
"encoding/json"
"testing"
mcplib "github.com/mark3labs/mcp-go/mcp"
_ "modernc.org/sqlite"
"github.com/smart-mcp-proxy/synapbus/internal/channels"
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
"github.com/smart-mcp-proxy/synapbus/internal/storage"
"github.com/smart-mcp-proxy/synapbus/internal/trace"
)
func newTestChannelRegistrar(t *testing.T) (*ChannelToolRegistrar, *channels.Service) {
t.Helper()
db := newTestDB(t)
tracer := trace.NewTracer(db)
t.Cleanup(func() { tracer.Close() })
channelStore := channels.NewSQLiteChannelStore(db)
msgStore := messaging.NewSQLiteMessageStore(db)
msgService := messaging.NewMessagingService(msgStore, tracer)
channelService := channels.NewService(channelStore, msgService, tracer)
// Seed test agents
db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES ('agent-a', 'Agent A', 'ai', '{}', 1, 'hash', 'active')`)
db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES ('agent-b', 'Agent B', 'ai', '{}', 1, 'hash', 'active')`)
db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES ('agent-c', 'Agent C', 'ai', '{}', 1, 'hash', 'active')`)
registrar := NewChannelToolRegistrar(channelService)
return registrar, channelService
}
func parseResponse(t *testing.T, result *mcplib.CallToolResult) map[string]any {
t.Helper()
var resp map[string]any
text := result.Content[0].(mcplib.TextContent).Text
if err := json.Unmarshal([]byte(text), &resp); err != nil {
t.Fatalf("unmarshal response: %v", err)
}
return resp
}
func TestChannelToolHandler_CreateChannel(t *testing.T) {
ctr, _ := newTestChannelRegistrar(t)
authCtx := ContextWithAgentName(context.Background(), "agent-a")
t.Run("successful creation", func(t *testing.T) {
req := makeRequest(map[string]any{
"name": "test-channel",
"description": "A test channel",
"type": "standard",
})
result, err := ctr.handleCreateChannel(authCtx, req)
if err != nil {
t.Fatalf("handleCreateChannel: %v", err)
}
if result.IsError {
t.Fatalf("unexpected error: %v", result.Content)
}
resp := parseResponse(t, result)
if resp["name"] != "test-channel" {
t.Errorf("name = %v, want test-channel", resp["name"])
}
if resp["channel_id"] == nil || resp["channel_id"].(float64) == 0 {
t.Error("expected non-zero channel_id")
}
})
t.Run("create private channel", func(t *testing.T) {
req := makeRequest(map[string]any{
"name": "private-test",
"is_private": true,
})
result, err := ctr.handleCreateChannel(authCtx, req)
if err != nil {
t.Fatalf("handleCreateChannel: %v", err)
}
if result.IsError {
t.Fatalf("unexpected error: %v", result.Content)
}
resp := parseResponse(t, result)
if resp["is_private"] != true {
t.Errorf("is_private = %v, want true", resp["is_private"])
}
})
t.Run("missing name", func(t *testing.T) {
req := makeRequest(map[string]any{})
result, _ := ctr.handleCreateChannel(authCtx, req)
if !result.IsError {
t.Error("expected error for missing name")
}
})
t.Run("unauthenticated", func(t *testing.T) {
req := makeRequest(map[string]any{"name": "fail"})
result, _ := ctr.handleCreateChannel(context.Background(), req)
if !result.IsError {
t.Error("expected error for unauthenticated request")
}
})
t.Run("duplicate name", func(t *testing.T) {
req := makeRequest(map[string]any{"name": "test-channel"})
result, _ := ctr.handleCreateChannel(authCtx, req)
if !result.IsError {
t.Error("expected error for duplicate name")
}
})
}
func TestChannelToolHandler_JoinChannel(t *testing.T) {
ctr, svc := newTestChannelRegistrar(t)
ctx := context.Background()
ch, _ := svc.CreateChannel(ctx, channels.CreateChannelRequest{
Name: "join-test", Type: "standard", CreatedBy: "agent-a",
})
authCtx := ContextWithAgentName(ctx, "agent-b")
t.Run("join by channel_id", func(t *testing.T) {
req := makeRequest(map[string]any{
"channel_id": float64(ch.ID),
})
result, err := ctr.handleJoinChannel(authCtx, req)
if err != nil {
t.Fatalf("handleJoinChannel: %v", err)
}
if result.IsError {
t.Fatalf("unexpected error: %v", result.Content)
}
resp := parseResponse(t, result)
if resp["status"] != "joined" {
t.Errorf("status = %v, want joined", resp["status"])
}
})
t.Run("join by channel_name", func(t *testing.T) {
authCtxC := ContextWithAgentName(ctx, "agent-c")
req := makeRequest(map[string]any{
"channel_name": "join-test",
})
result, err := ctr.handleJoinChannel(authCtxC, req)
if err != nil {
t.Fatalf("handleJoinChannel: %v", err)
}
if result.IsError {
t.Fatalf("unexpected error: %v", result.Content)
}
})
t.Run("no channel identifier", func(t *testing.T) {
req := makeRequest(map[string]any{})
result, _ := ctr.handleJoinChannel(authCtx, req)
if !result.IsError {
t.Error("expected error when no channel identifier provided")
}
})
}
func TestChannelToolHandler_LeaveChannel(t *testing.T) {
ctr, svc := newTestChannelRegistrar(t)
ctx := context.Background()
ch, _ := svc.CreateChannel(ctx, channels.CreateChannelRequest{
Name: "leave-test", Type: "standard", CreatedBy: "agent-a",
})
svc.JoinChannel(ctx, ch.ID, "agent-b")
authCtx := ContextWithAgentName(ctx, "agent-b")
t.Run("successful leave", func(t *testing.T) {
req := makeRequest(map[string]any{
"channel_id": float64(ch.ID),
})
result, err := ctr.handleLeaveChannel(authCtx, req)
if err != nil {
t.Fatalf("handleLeaveChannel: %v", err)
}
if result.IsError {
t.Fatalf("unexpected error: %v", result.Content)
}
})
t.Run("owner cannot leave", func(t *testing.T) {
ownerCtx := ContextWithAgentName(ctx, "agent-a")
req := makeRequest(map[string]any{
"channel_id": float64(ch.ID),
})
result, _ := ctr.handleLeaveChannel(ownerCtx, req)
if !result.IsError {
t.Error("expected error for owner leaving")
}
})
}
func TestChannelToolHandler_ListChannels(t *testing.T) {
ctr, svc := newTestChannelRegistrar(t)
ctx := context.Background()
svc.CreateChannel(ctx, channels.CreateChannelRequest{Name: "pub-1", Type: "standard", CreatedBy: "agent-a"})
svc.CreateChannel(ctx, channels.CreateChannelRequest{Name: "pub-2", Type: "standard", CreatedBy: "agent-a"})
authCtx := ContextWithAgentName(ctx, "agent-b")
req := makeRequest(map[string]any{})
result, err := ctr.handleListChannels(authCtx, req)
if err != nil {
t.Fatalf("handleListChannels: %v", err)
}
if result.IsError {
t.Fatalf("unexpected error: %v", result.Content)
}
resp := parseResponse(t, result)
count := resp["count"].(float64)
if count != 2 {
t.Errorf("count = %v, want 2", count)
}
chList := resp["channels"].([]any)
ch0 := chList[0].(map[string]any)
if ch0["name"] == nil {
t.Error("expected name field in channel")
}
if ch0["member_count"] == nil {
t.Error("expected member_count field in channel")
}
}
func TestChannelToolHandler_InviteToChannel(t *testing.T) {
ctr, svc := newTestChannelRegistrar(t)
ctx := context.Background()
ch, _ := svc.CreateChannel(ctx, channels.CreateChannelRequest{
Name: "invite-test", Type: "standard", IsPrivate: true, CreatedBy: "agent-a",
})
ownerCtx := ContextWithAgentName(ctx, "agent-a")
t.Run("owner can invite", func(t *testing.T) {
req := makeRequest(map[string]any{
"channel_id": float64(ch.ID),
"agent_name": "agent-b",
})
result, err := ctr.handleInviteToChannel(ownerCtx, req)
if err != nil {
t.Fatalf("handleInviteToChannel: %v", err)
}
if result.IsError {
t.Fatalf("unexpected error: %v", result.Content)
}
})
t.Run("missing agent_name", func(t *testing.T) {
req := makeRequest(map[string]any{
"channel_id": float64(ch.ID),
})
result, _ := ctr.handleInviteToChannel(ownerCtx, req)
if !result.IsError {
t.Error("expected error for missing agent_name")
}
})
}
func TestChannelToolHandler_KickFromChannel(t *testing.T) {
ctr, svc := newTestChannelRegistrar(t)
ctx := context.Background()
ch, _ := svc.CreateChannel(ctx, channels.CreateChannelRequest{
Name: "kick-test", Type: "standard", CreatedBy: "agent-a",
})
svc.JoinChannel(ctx, ch.ID, "agent-b")
ownerCtx := ContextWithAgentName(ctx, "agent-a")
t.Run("owner can kick", func(t *testing.T) {
req := makeRequest(map[string]any{
"channel_id": float64(ch.ID),
"agent_name": "agent-b",
})
result, err := ctr.handleKickFromChannel(ownerCtx, req)
if err != nil {
t.Fatalf("handleKickFromChannel: %v", err)
}
if result.IsError {
t.Fatalf("unexpected error: %v", result.Content)
}
resp := parseResponse(t, result)
if resp["status"] != "kicked" {
t.Errorf("status = %v, want kicked", resp["status"])
}
})
}
func TestChannelToolHandler_SendChannelMessage(t *testing.T) {
ctr, svc := newTestChannelRegistrar(t)
ctx := context.Background()
ch, _ := svc.CreateChannel(ctx, channels.CreateChannelRequest{
Name: "msg-test", Type: "standard", CreatedBy: "agent-a",
})
svc.JoinChannel(ctx, ch.ID, "agent-b")
authCtx := ContextWithAgentName(ctx, "agent-a")
t.Run("send channel message", func(t *testing.T) {
req := makeRequest(map[string]any{
"channel_name": "msg-test",
"body": "Hello channel!",
})
result, err := ctr.handleSendChannelMessage(authCtx, req)
if err != nil {
t.Fatalf("handleSendChannelMessage: %v", err)
}
if result.IsError {
t.Fatalf("unexpected error: %v", result.Content)
}
resp := parseResponse(t, result)
recipients := resp["recipients"].(float64)
if recipients != 1 {
t.Errorf("recipients = %v, want 1", recipients)
}
})
t.Run("missing body", func(t *testing.T) {
req := makeRequest(map[string]any{
"channel_name": "msg-test",
})
result, _ := ctr.handleSendChannelMessage(authCtx, req)
if !result.IsError {
t.Error("expected error for missing body")
}
})
}
func TestChannelToolHandler_UpdateChannel(t *testing.T) {
ctr, svc := newTestChannelRegistrar(t)
ctx := context.Background()
ch, _ := svc.CreateChannel(ctx, channels.CreateChannelRequest{
Name: "update-test", Type: "standard", Topic: "Original", CreatedBy: "agent-a",
})
ownerCtx := ContextWithAgentName(ctx, "agent-a")
t.Run("update topic", func(t *testing.T) {
req := makeRequest(map[string]any{
"channel_id": float64(ch.ID),
"topic": "Updated topic",
})
result, err := ctr.handleUpdateChannel(ownerCtx, req)
if err != nil {
t.Fatalf("handleUpdateChannel: %v", err)
}
if result.IsError {
t.Fatalf("unexpected error: %v", result.Content)
}
resp := parseResponse(t, result)
if resp["topic"] != "Updated topic" {
t.Errorf("topic = %v, want 'Updated topic'", resp["topic"])
}
})
t.Run("non-owner cannot update", func(t *testing.T) {
svc.JoinChannel(ctx, ch.ID, "agent-b")
nonOwnerCtx := ContextWithAgentName(ctx, "agent-b")
req := makeRequest(map[string]any{
"channel_id": float64(ch.ID),
"topic": "Unauthorized",
})
result, _ := ctr.handleUpdateChannel(nonOwnerCtx, req)
if !result.IsError {
t.Error("expected error for non-owner update")
}
})
}
// suppress unused import
var _ = storage.RunMigrations
+8
View File
@@ -8,6 +8,7 @@ import (
"github.com/mark3labs/mcp-go/server"
"github.com/smart-mcp-proxy/synapbus/internal/agents"
"github.com/smart-mcp-proxy/synapbus/internal/channels"
"github.com/smart-mcp-proxy/synapbus/internal/messaging"
)
@@ -24,6 +25,7 @@ type MCPServer struct {
func NewMCPServer(
msgService *messaging.MessagingService,
agentService *agents.AgentService,
channelService *channels.Service,
) *MCPServer {
logger := slog.Default().With("component", "mcp-server")
@@ -38,6 +40,12 @@ func NewMCPServer(
registrar := NewToolRegistrar(msgService, agentService)
registrar.RegisterAll(mcpSrv)
// Register channel tools
if channelService != nil {
channelRegistrar := NewChannelToolRegistrar(channelService)
channelRegistrar.RegisterAll(mcpSrv)
}
// Create SSE transport with context func for auth propagation
sseServer := server.NewSSEServer(mcpSrv,
server.WithSSEContextFunc(func(ctx context.Context, r *http.Request) context.Context {
+14
View File
@@ -0,0 +1,14 @@
-- Channel invites for private channel membership gating
CREATE TABLE IF NOT EXISTS channel_invites (
id INTEGER PRIMARY KEY AUTOINCREMENT,
channel_id INTEGER NOT NULL REFERENCES channels(id) ON DELETE CASCADE,
agent_name TEXT NOT NULL,
invited_by TEXT NOT NULL,
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
status TEXT NOT NULL DEFAULT 'pending' CHECK (status IN ('pending', 'accepted', 'declined')),
UNIQUE(channel_id, agent_name)
);
CREATE INDEX IF NOT EXISTS idx_channel_invites_agent ON channel_invites(agent_name);
INSERT OR IGNORE INTO schema_migrations (version) VALUES (2);
+14
View File
@@ -0,0 +1,14 @@
-- Channel invites for private channel membership gating
CREATE TABLE IF NOT EXISTS channel_invites (
id INTEGER PRIMARY KEY AUTOINCREMENT,
channel_id INTEGER NOT NULL REFERENCES channels(id) ON DELETE CASCADE,
agent_name TEXT NOT NULL,
invited_by TEXT NOT NULL,
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
status TEXT NOT NULL DEFAULT 'pending' CHECK (status IN ('pending', 'accepted', 'declined')),
UNIQUE(channel_id, agent_name)
);
CREATE INDEX IF NOT EXISTS idx_channel_invites_agent ON channel_invites(agent_name);
INSERT OR IGNORE INTO schema_migrations (version) VALUES (2);