diff --git a/cmd/synapbus/main.go b/cmd/synapbus/main.go index f16cccd..be7ac5b 100644 --- a/cmd/synapbus/main.go +++ b/cmd/synapbus/main.go @@ -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 diff --git a/internal/channels/errors.go b/internal/channels/errors.go new file mode 100644 index 0000000..8f7a261 --- /dev/null +++ b/internal/channels/errors.go @@ -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") +) diff --git a/internal/channels/service.go b/internal/channels/service.go new file mode 100644 index 0000000..a4b8057 --- /dev/null +++ b/internal/channels/service.go @@ -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) +} diff --git a/internal/channels/service_test.go b/internal/channels/service_test.go new file mode 100644 index 0000000..26e95b5 --- /dev/null +++ b/internal/channels/service_test.go @@ -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 diff --git a/internal/channels/store.go b/internal/channels/store.go new file mode 100644 index 0000000..f7b6c0c --- /dev/null +++ b/internal/channels/store.go @@ -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") +} diff --git a/internal/channels/store_test.go b/internal/channels/store_test.go new file mode 100644 index 0000000..fbcff9e --- /dev/null +++ b/internal/channels/store_test.go @@ -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 diff --git a/internal/channels/types.go b/internal/channels/types.go new file mode 100644 index 0000000..c386f9f --- /dev/null +++ b/internal/channels/types.go @@ -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"` +} diff --git a/internal/channels/validate.go b/internal/channels/validate.go new file mode 100644 index 0000000..301c842 --- /dev/null +++ b/internal/channels/validate.go @@ -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) +} diff --git a/internal/channels/validate_test.go b/internal/channels/validate_test.go new file mode 100644 index 0000000..6deeceb --- /dev/null +++ b/internal/channels/validate_test.go @@ -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) + } + }) + } +} diff --git a/internal/mcp/channel_tools.go b/internal/mcp/channel_tools.go new file mode 100644 index 0000000..b637bd9 --- /dev/null +++ b/internal/mcp/channel_tools.go @@ -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") +} diff --git a/internal/mcp/channel_tools_test.go b/internal/mcp/channel_tools_test.go new file mode 100644 index 0000000..7ac85fb --- /dev/null +++ b/internal/mcp/channel_tools_test.go @@ -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 diff --git a/internal/mcp/server.go b/internal/mcp/server.go index 0c8866c..231a3c7 100644 --- a/internal/mcp/server.go +++ b/internal/mcp/server.go @@ -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 { diff --git a/internal/storage/schema/002_channels.sql b/internal/storage/schema/002_channels.sql new file mode 100644 index 0000000..d9d2f0f --- /dev/null +++ b/internal/storage/schema/002_channels.sql @@ -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); diff --git a/schema/002_channels.sql b/schema/002_channels.sql new file mode 100644 index 0000000..d9d2f0f --- /dev/null +++ b/schema/002_channels.sql @@ -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);