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:
co-authored by
Claude Opus 4.6
parent
2f55ce87c3
commit
e3dcfece48
@@ -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
|
||||
|
||||
@@ -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")
|
||||
)
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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
|
||||
@@ -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"`
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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,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 {
|
||||
|
||||
@@ -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);
|
||||
@@ -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);
|
||||
Reference in New Issue
Block a user