feat: implement swarm patterns (task auction, stigmergy, discovery)

Add task auction lifecycle with full permission enforcement:
- TaskStore (SQLite) for tasks and bids CRUD, expiry, channel cancellation
- SwarmService with PostTask, BidOnTask, AcceptBid, CompleteTask
- MCP tools: post_task, bid_task, accept_bid, complete_task, list_tasks
- ExpiryWorker background goroutine for deadline-based task cancellation
- Channel type enforcement (auction ops only on auction channels)
- Agent cannot bid on own task, only poster accepts bids, only assignee completes
- Wired into main.go with graceful shutdown

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
Algis Dumbris
2026-03-13 12:12:33 +02:00
co-authored by Claude Opus 4.6
parent 1cb0bb7a8f
commit 7909450473
10 changed files with 2013 additions and 2 deletions
+14 -2
View File
@@ -180,6 +180,10 @@ func runServe(cmd *cobra.Command, args []string) error {
channelStore := channels.NewSQLiteChannelStore(db.DB)
channelService := channels.NewService(channelStore, msgService, tracer)
// Create swarm service (task auction + stigmergy)
taskStore := channels.NewSQLiteTaskStore(db.DB)
swarmService := channels.NewSwarmService(taskStore, channelStore, tracer)
// Initialize auth subsystem
authSecret := make([]byte, 32)
if _, err := rand.Read(authSecret); err != nil {
@@ -220,10 +224,15 @@ func runServe(cmd *cobra.Command, args []string) error {
fmt.Printf("========================================\n\n")
}
// Create MCP server
mcpSrv := mcpserver.NewMCPServer(msgService, agentService, channelService)
// Create MCP server (with swarm tools)
mcpSrv := mcpserver.NewMCPServer(msgService, agentService, channelService, swarmService)
startTime := time.Now()
// Start task expiry worker
expiryWorker := channels.NewExpiryWorker(swarmService, 1*time.Minute)
expiryWorker.Start()
slog.Info("task expiry worker started")
// Set up chi router
r := chi.NewRouter()
@@ -284,6 +293,9 @@ func runServe(cmd *cobra.Command, args []string) error {
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 30*time.Second)
defer shutdownCancel()
// Stop expiry worker
expiryWorker.Stop()
// Stop retention cleaner
if retentionCleaner != nil {
retentionCleaner.Stop()
+66
View File
@@ -0,0 +1,66 @@
package channels
import (
"context"
"log/slog"
"sync"
"time"
)
// ExpiryWorker periodically checks for and cancels expired tasks.
type ExpiryWorker struct {
swarmService *SwarmService
interval time.Duration
logger *slog.Logger
done chan struct{}
wg sync.WaitGroup
}
// NewExpiryWorker creates a new expiry worker.
// The interval controls how often it checks for expired tasks (default: 1 minute).
func NewExpiryWorker(swarmService *SwarmService, interval time.Duration) *ExpiryWorker {
if interval <= 0 {
interval = 1 * time.Minute
}
return &ExpiryWorker{
swarmService: swarmService,
interval: interval,
logger: slog.Default().With("component", "expiry-worker"),
done: make(chan struct{}),
}
}
// Start begins the background expiry check loop.
func (w *ExpiryWorker) Start() {
w.wg.Add(1)
go func() {
defer w.wg.Done()
w.logger.Info("expiry worker started", "interval", w.interval.String())
ticker := time.NewTicker(w.interval)
defer ticker.Stop()
for {
select {
case <-ticker.C:
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
count, err := w.swarmService.ExpireTasks(ctx)
if err != nil {
w.logger.Error("expiry check failed", "error", err)
} else if count > 0 {
w.logger.Info("expired tasks processed", "count", count)
}
cancel()
case <-w.done:
w.logger.Info("expiry worker stopped")
return
}
}
}()
}
// Stop stops the expiry worker and waits for it to finish.
func (w *ExpiryWorker) Stop() {
close(w.done)
w.wg.Wait()
}
+126
View File
@@ -0,0 +1,126 @@
package channels
import (
"context"
"encoding/json"
"testing"
"time"
"github.com/smart-mcp-proxy/synapbus/internal/trace"
)
func TestExpiryWorker_ExpiresOverdueTasks(t *testing.T) {
db := newTestDB(t)
seedAgent(t, db, "poster-agent")
channelStore := NewSQLiteChannelStore(db)
taskStore := NewSQLiteTaskStore(db)
tracer := trace.NewTracer(db)
t.Cleanup(func() { tracer.Close() })
svc := NewSwarmService(taskStore, channelStore, tracer)
ctx := context.Background()
// Create auction channel
ch := &Channel{Name: "expiry-test", Type: TypeAuction, CreatedBy: "poster-agent"}
channelStore.CreateChannel(ctx, ch)
channelStore.AddMember(ctx, &Membership{ChannelID: ch.ID, AgentName: "poster-agent", Role: RoleOwner})
// Create a task with deadline in the past (bypass PostTask validation by using store directly)
pastDeadline := time.Now().Add(-1 * time.Hour)
taskStore.CreateTask(ctx, &Task{
ChannelID: ch.ID,
PostedBy: "poster-agent",
Title: "Expired Task",
Status: TaskStatusOpen,
Deadline: &pastDeadline,
Requirements: json.RawMessage(`{}`),
})
// Start worker with very short interval
worker := NewExpiryWorker(svc, 50*time.Millisecond)
worker.Start()
// Wait for at least one tick
time.Sleep(200 * time.Millisecond)
// Stop worker
worker.Stop()
// Verify the task was cancelled
tasks, err := taskStore.ListTasks(ctx, ch.ID, TaskStatusCancelled)
if err != nil {
t.Fatalf("ListTasks: %v", err)
}
if len(tasks) != 1 {
t.Errorf("cancelled tasks = %d, want 1", len(tasks))
}
}
func TestExpiryWorker_DoesNotExpireFutureTasks(t *testing.T) {
db := newTestDB(t)
seedAgent(t, db, "poster-agent")
channelStore := NewSQLiteChannelStore(db)
taskStore := NewSQLiteTaskStore(db)
tracer := trace.NewTracer(db)
t.Cleanup(func() { tracer.Close() })
svc := NewSwarmService(taskStore, channelStore, tracer)
ctx := context.Background()
ch := &Channel{Name: "no-expire-test", Type: TypeAuction, CreatedBy: "poster-agent"}
channelStore.CreateChannel(ctx, ch)
channelStore.AddMember(ctx, &Membership{ChannelID: ch.ID, AgentName: "poster-agent", Role: RoleOwner})
futureDeadline := time.Now().Add(1 * time.Hour)
taskStore.CreateTask(ctx, &Task{
ChannelID: ch.ID,
PostedBy: "poster-agent",
Title: "Future Task",
Status: TaskStatusOpen,
Deadline: &futureDeadline,
Requirements: json.RawMessage(`{}`),
})
worker := NewExpiryWorker(svc, 50*time.Millisecond)
worker.Start()
time.Sleep(200 * time.Millisecond)
worker.Stop()
// Task should still be open
tasks, _ := taskStore.ListTasks(ctx, ch.ID, TaskStatusOpen)
if len(tasks) != 1 {
t.Errorf("open tasks = %d, want 1", len(tasks))
}
}
func TestExpiryWorker_StartStop(t *testing.T) {
db := newTestDB(t)
channelStore := NewSQLiteChannelStore(db)
taskStore := NewSQLiteTaskStore(db)
tracer := trace.NewTracer(db)
t.Cleanup(func() { tracer.Close() })
svc := NewSwarmService(taskStore, channelStore, tracer)
worker := NewExpiryWorker(svc, 100*time.Millisecond)
worker.Start()
// Stop should not block/panic
worker.Stop()
}
func TestExpiryWorker_DefaultInterval(t *testing.T) {
db := newTestDB(t)
channelStore := NewSQLiteChannelStore(db)
taskStore := NewSQLiteTaskStore(db)
tracer := trace.NewTracer(db)
t.Cleanup(func() { tracer.Close() })
svc := NewSwarmService(taskStore, channelStore, tracer)
worker := NewExpiryWorker(svc, 0)
if worker.interval != 1*time.Minute {
t.Errorf("default interval = %v, want 1m", worker.interval)
}
}
+301
View File
@@ -0,0 +1,301 @@
package channels
import (
"context"
"encoding/json"
"fmt"
"log/slog"
"time"
"github.com/smart-mcp-proxy/synapbus/internal/trace"
)
// SwarmService handles task auction and stigmergy patterns.
type SwarmService struct {
taskStore TaskStore
channelStore ChannelStore
tracer *trace.Tracer
logger *slog.Logger
}
// NewSwarmService creates a new swarm service.
func NewSwarmService(taskStore TaskStore, channelStore ChannelStore, tracer *trace.Tracer) *SwarmService {
return &SwarmService{
taskStore: taskStore,
channelStore: channelStore,
tracer: tracer,
logger: slog.Default().With("component", "swarm"),
}
}
// PostTask creates a new task on an auction channel.
func (s *SwarmService) PostTask(ctx context.Context, channelID int64, agentName, title, description string, requirements json.RawMessage, deadline *time.Time) (*Task, error) {
// Verify channel exists and is auction type
ch, err := s.channelStore.GetChannel(ctx, channelID)
if err != nil {
return nil, err
}
if ch.Type != TypeAuction {
return nil, fmt.Errorf("post_task requires a channel of type 'auction', got '%s'", ch.Type)
}
// Verify agent is channel member
isMember, err := s.channelStore.IsMember(ctx, channelID, agentName)
if err != nil {
return nil, fmt.Errorf("check membership: %w", err)
}
if !isMember {
return nil, ErrNotChannelMember
}
// Validate deadline is in the future (if provided)
if deadline != nil && deadline.Before(time.Now()) {
return nil, fmt.Errorf("deadline must be in the future")
}
if requirements == nil || len(requirements) == 0 {
requirements = json.RawMessage("{}")
}
task := &Task{
ChannelID: channelID,
PostedBy: agentName,
Title: title,
Description: description,
Requirements: requirements,
Deadline: deadline,
Status: TaskStatusOpen,
}
if err := s.taskStore.CreateTask(ctx, task); err != nil {
return nil, fmt.Errorf("create task: %w", err)
}
s.logger.Info("task posted",
"task_id", task.ID,
"channel_id", channelID,
"posted_by", agentName,
"title", title,
)
if s.tracer != nil {
s.tracer.Record(ctx, agentName, "swarm.task_posted", map[string]any{
"task_id": task.ID,
"channel_id": channelID,
"channel_name": ch.Name,
"title": title,
})
}
return task, nil
}
// BidOnTask creates a bid on an open task.
func (s *SwarmService) BidOnTask(ctx context.Context, taskID int64, agentName string, capabilities json.RawMessage, timeEstimate, message string) (*Bid, error) {
// Get task
task, err := s.taskStore.GetTask(ctx, taskID)
if err != nil {
return nil, err
}
// Verify task is open
if task.Status != TaskStatusOpen {
return nil, fmt.Errorf("cannot bid on task with status '%s'; task must be 'open'", task.Status)
}
// Verify agent is not the poster
if task.PostedBy == agentName {
return nil, fmt.Errorf("cannot bid on your own task")
}
// Verify agent is member of task's channel
isMember, err := s.channelStore.IsMember(ctx, task.ChannelID, agentName)
if err != nil {
return nil, fmt.Errorf("check membership: %w", err)
}
if !isMember {
return nil, ErrNotChannelMember
}
if capabilities == nil || len(capabilities) == 0 {
capabilities = json.RawMessage("{}")
}
bid := &Bid{
TaskID: taskID,
AgentName: agentName,
Capabilities: capabilities,
TimeEstimate: timeEstimate,
Message: message,
}
if err := s.taskStore.CreateBid(ctx, bid); err != nil {
return nil, err
}
s.logger.Info("bid submitted",
"bid_id", bid.ID,
"task_id", taskID,
"agent", agentName,
)
if s.tracer != nil {
s.tracer.Record(ctx, agentName, "swarm.bid_submitted", map[string]any{
"bid_id": bid.ID,
"task_id": taskID,
})
}
return bid, nil
}
// AcceptBid accepts a bid and assigns the task to the bidding agent.
func (s *SwarmService) AcceptBid(ctx context.Context, taskID, bidID int64, agentName string) error {
// Get task
task, err := s.taskStore.GetTask(ctx, taskID)
if err != nil {
return err
}
// Verify caller is the task poster
if task.PostedBy != agentName {
return fmt.Errorf("only the task poster can accept bids")
}
// Verify task is open
if task.Status != TaskStatusOpen {
return fmt.Errorf("cannot accept bid on task with status '%s'; task must be 'open'", task.Status)
}
// Get the bid
bid, err := s.taskStore.GetBid(ctx, bidID)
if err != nil {
return err
}
// Verify bid belongs to this task
if bid.TaskID != taskID {
return fmt.Errorf("bid %d does not belong to task %d", bidID, taskID)
}
// Assign the task
if err := s.taskStore.UpdateTaskStatus(ctx, taskID, TaskStatusAssigned, bid.AgentName); err != nil {
return fmt.Errorf("assign task: %w", err)
}
// Accept the winning bid
if err := s.taskStore.UpdateBidStatus(ctx, bidID, BidStatusAccepted); err != nil {
return fmt.Errorf("accept bid: %w", err)
}
// Reject all other bids
allBids, err := s.taskStore.GetBids(ctx, taskID)
if err != nil {
return fmt.Errorf("get bids for rejection: %w", err)
}
for _, b := range allBids {
if b.ID != bidID && b.Status == BidStatusPending {
if err := s.taskStore.UpdateBidStatus(ctx, b.ID, BidStatusRejected); err != nil {
s.logger.Error("failed to reject bid", "bid_id", b.ID, "error", err)
}
}
}
s.logger.Info("bid accepted",
"task_id", taskID,
"bid_id", bidID,
"assigned_to", bid.AgentName,
"accepted_by", agentName,
)
if s.tracer != nil {
s.tracer.Record(ctx, agentName, "swarm.bid_accepted", map[string]any{
"task_id": taskID,
"bid_id": bidID,
"assigned_to": bid.AgentName,
})
}
return nil
}
// CompleteTask marks a task as completed.
func (s *SwarmService) CompleteTask(ctx context.Context, taskID int64, agentName string) error {
// Get task
task, err := s.taskStore.GetTask(ctx, taskID)
if err != nil {
return err
}
// Handle idempotent completion
if task.Status == TaskStatusCompleted && task.AssignedTo == agentName {
return nil
}
// Verify task is assigned
if task.Status != TaskStatusAssigned {
return fmt.Errorf("cannot complete task with status '%s'; task must be 'assigned'", task.Status)
}
// Verify caller is the assigned agent
if task.AssignedTo != agentName {
return fmt.Errorf("only the assigned agent can complete the task")
}
// Complete the task
if err := s.taskStore.UpdateTaskStatus(ctx, taskID, TaskStatusCompleted, ""); err != nil {
return fmt.Errorf("complete task: %w", err)
}
s.logger.Info("task completed",
"task_id", taskID,
"completed_by", agentName,
)
if s.tracer != nil {
s.tracer.Record(ctx, agentName, "swarm.task_completed", map[string]any{
"task_id": taskID,
})
}
return nil
}
// ListTasks returns tasks for a channel, optionally filtered by status.
func (s *SwarmService) ListTasks(ctx context.Context, channelID int64, status string) ([]*Task, error) {
return s.taskStore.ListTasks(ctx, channelID, status)
}
// GetTaskWithBids returns a task and all its bids.
func (s *SwarmService) GetTaskWithBids(ctx context.Context, taskID int64) (*Task, []*Bid, error) {
task, err := s.taskStore.GetTask(ctx, taskID)
if err != nil {
return nil, nil, err
}
bids, err := s.taskStore.GetBids(ctx, taskID)
if err != nil {
return nil, nil, err
}
return task, bids, nil
}
// ExpireTasks marks expired tasks as cancelled. Called by the expiry worker.
func (s *SwarmService) ExpireTasks(ctx context.Context) (int, error) {
count, err := s.taskStore.ExpireTasks(ctx)
if err != nil {
return 0, err
}
if count > 0 {
s.logger.Info("expired tasks cancelled", "count", count)
if s.tracer != nil {
s.tracer.Record(ctx, "system", "swarm.tasks_expired", map[string]any{
"count": count,
})
}
}
return count, nil
}
+466
View File
@@ -0,0 +1,466 @@
package channels
import (
"context"
"encoding/json"
"testing"
"time"
"github.com/smart-mcp-proxy/synapbus/internal/trace"
)
func newTestSwarmService(t *testing.T) (*SwarmService, *SQLiteChannelStore) {
t.Helper()
db := newTestDB(t)
seedAgent(t, db, "poster-agent")
seedAgent(t, db, "bidder-agent")
seedAgent(t, db, "bidder-agent-2")
seedAgent(t, db, "outsider-agent")
channelStore := NewSQLiteChannelStore(db)
taskStore := NewSQLiteTaskStore(db)
tracer := trace.NewTracer(db)
t.Cleanup(func() { tracer.Close() })
svc := NewSwarmService(taskStore, channelStore, tracer)
return svc, channelStore
}
func createTestAuctionChannel(t *testing.T, channelStore *SQLiteChannelStore) *Channel {
t.Helper()
ctx := context.Background()
ch := &Channel{Name: "test-auction", Type: TypeAuction, CreatedBy: "poster-agent"}
if err := channelStore.CreateChannel(ctx, ch); err != nil {
t.Fatalf("create auction channel: %v", err)
}
channelStore.AddMember(ctx, &Membership{ChannelID: ch.ID, AgentName: "poster-agent", Role: RoleOwner})
channelStore.AddMember(ctx, &Membership{ChannelID: ch.ID, AgentName: "bidder-agent", Role: RoleMember})
channelStore.AddMember(ctx, &Membership{ChannelID: ch.ID, AgentName: "bidder-agent-2", Role: RoleMember})
return ch
}
// --- PostTask tests ---
func TestSwarmService_PostTask(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
deadline := time.Now().Add(1 * time.Hour)
task, err := svc.PostTask(ctx, ch.ID, "poster-agent", "Test Task", "A test task",
json.RawMessage(`{"skill":"go"}`), &deadline)
if err != nil {
t.Fatalf("PostTask: %v", err)
}
if task.ID == 0 {
t.Error("task ID should not be 0")
}
if task.Status != TaskStatusOpen {
t.Errorf("status = %s, want open", task.Status)
}
if task.PostedBy != "poster-agent" {
t.Errorf("posted_by = %s, want poster-agent", task.PostedBy)
}
}
func TestSwarmService_PostTask_RequiresAuctionChannel(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ctx := context.Background()
// Create a standard channel
stdCh := &Channel{Name: "standard-ch", Type: TypeStandard, CreatedBy: "poster-agent"}
channelStore.CreateChannel(ctx, stdCh)
channelStore.AddMember(ctx, &Membership{ChannelID: stdCh.ID, AgentName: "poster-agent", Role: RoleOwner})
_, err := svc.PostTask(ctx, stdCh.ID, "poster-agent", "Task", "", nil, nil)
if err == nil {
t.Fatal("expected error for non-auction channel")
}
}
func TestSwarmService_PostTask_RequiresMembership(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
_, err := svc.PostTask(ctx, ch.ID, "outsider-agent", "Task", "", nil, nil)
if err == nil {
t.Fatal("expected error for non-member")
}
}
func TestSwarmService_PostTask_RejectsPastDeadline(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
past := time.Now().Add(-1 * time.Hour)
_, err := svc.PostTask(ctx, ch.ID, "poster-agent", "Past Deadline", "", nil, &past)
if err == nil {
t.Fatal("expected error for past deadline")
}
}
// --- BidOnTask tests ---
func TestSwarmService_BidOnTask(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
task, _ := svc.PostTask(ctx, ch.ID, "poster-agent", "Bid Test", "", nil, nil)
bid, err := svc.BidOnTask(ctx, task.ID, "bidder-agent",
json.RawMessage(`{"lang":"go"}`), "30m", "I can do this")
if err != nil {
t.Fatalf("BidOnTask: %v", err)
}
if bid.ID == 0 {
t.Error("bid ID should not be 0")
}
if bid.Status != BidStatusPending {
t.Errorf("status = %s, want pending", bid.Status)
}
}
func TestSwarmService_BidOnTask_CannotBidOnOwnTask(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
task, _ := svc.PostTask(ctx, ch.ID, "poster-agent", "Own Bid Test", "", nil, nil)
_, err := svc.BidOnTask(ctx, task.ID, "poster-agent", nil, "", "")
if err == nil {
t.Fatal("expected error when bidding on own task")
}
}
func TestSwarmService_BidOnTask_RequiresOpenTask(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
task, _ := svc.PostTask(ctx, ch.ID, "poster-agent", "Status Test", "", nil, nil)
// Bid and accept to move to assigned
bid, _ := svc.BidOnTask(ctx, task.ID, "bidder-agent", nil, "", "")
svc.AcceptBid(ctx, task.ID, bid.ID, "poster-agent")
// Try to bid on assigned task
_, err := svc.BidOnTask(ctx, task.ID, "bidder-agent-2", nil, "", "")
if err == nil {
t.Fatal("expected error when bidding on assigned task")
}
}
func TestSwarmService_BidOnTask_RequiresMembership(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
task, _ := svc.PostTask(ctx, ch.ID, "poster-agent", "Member Test", "", nil, nil)
_, err := svc.BidOnTask(ctx, task.ID, "outsider-agent", nil, "", "")
if err == nil {
t.Fatal("expected error for non-member bidder")
}
}
// --- AcceptBid tests ---
func TestSwarmService_AcceptBid(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
task, _ := svc.PostTask(ctx, ch.ID, "poster-agent", "Accept Test", "", nil, nil)
bid1, _ := svc.BidOnTask(ctx, task.ID, "bidder-agent", nil, "", "bid 1")
bid2, _ := svc.BidOnTask(ctx, task.ID, "bidder-agent-2", nil, "", "bid 2")
err := svc.AcceptBid(ctx, task.ID, bid1.ID, "poster-agent")
if err != nil {
t.Fatalf("AcceptBid: %v", err)
}
// Verify task is assigned
updatedTask, _, _ := svc.GetTaskWithBids(ctx, task.ID)
if updatedTask.Status != TaskStatusAssigned {
t.Errorf("task status = %s, want assigned", updatedTask.Status)
}
if updatedTask.AssignedTo != "bidder-agent" {
t.Errorf("assigned_to = %s, want bidder-agent", updatedTask.AssignedTo)
}
// Verify bid statuses
_, bids, _ := svc.GetTaskWithBids(ctx, task.ID)
for _, b := range bids {
if b.ID == bid1.ID && b.Status != BidStatusAccepted {
t.Errorf("winning bid status = %s, want accepted", b.Status)
}
if b.ID == bid2.ID && b.Status != BidStatusRejected {
t.Errorf("losing bid status = %s, want rejected", b.Status)
}
}
}
func TestSwarmService_AcceptBid_OnlyPosterCanAccept(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
task, _ := svc.PostTask(ctx, ch.ID, "poster-agent", "Auth Test", "", nil, nil)
bid, _ := svc.BidOnTask(ctx, task.ID, "bidder-agent", nil, "", "")
err := svc.AcceptBid(ctx, task.ID, bid.ID, "bidder-agent")
if err == nil {
t.Fatal("expected error when non-poster accepts bid")
}
}
func TestSwarmService_AcceptBid_CannotAcceptOnNonOpenTask(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
task, _ := svc.PostTask(ctx, ch.ID, "poster-agent", "Double Accept Test", "", nil, nil)
bid, _ := svc.BidOnTask(ctx, task.ID, "bidder-agent", nil, "", "")
// Accept once
svc.AcceptBid(ctx, task.ID, bid.ID, "poster-agent")
// Try to accept again
err := svc.AcceptBid(ctx, task.ID, bid.ID, "poster-agent")
if err == nil {
t.Fatal("expected error when accepting bid on non-open task")
}
}
// --- CompleteTask tests ---
func TestSwarmService_CompleteTask(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
task, _ := svc.PostTask(ctx, ch.ID, "poster-agent", "Complete Test", "", nil, nil)
bid, _ := svc.BidOnTask(ctx, task.ID, "bidder-agent", nil, "", "")
svc.AcceptBid(ctx, task.ID, bid.ID, "poster-agent")
err := svc.CompleteTask(ctx, task.ID, "bidder-agent")
if err != nil {
t.Fatalf("CompleteTask: %v", err)
}
updatedTask, _, _ := svc.GetTaskWithBids(ctx, task.ID)
if updatedTask.Status != TaskStatusCompleted {
t.Errorf("status = %s, want completed", updatedTask.Status)
}
}
func TestSwarmService_CompleteTask_OnlyAssignedAgentCanComplete(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
task, _ := svc.PostTask(ctx, ch.ID, "poster-agent", "Auth Complete Test", "", nil, nil)
bid, _ := svc.BidOnTask(ctx, task.ID, "bidder-agent", nil, "", "")
svc.AcceptBid(ctx, task.ID, bid.ID, "poster-agent")
// Non-assigned agent tries to complete
err := svc.CompleteTask(ctx, task.ID, "poster-agent")
if err == nil {
t.Fatal("expected error when non-assigned agent completes task")
}
err = svc.CompleteTask(ctx, task.ID, "bidder-agent-2")
if err == nil {
t.Fatal("expected error when non-assigned agent completes task")
}
}
func TestSwarmService_CompleteTask_IdempotentCompletion(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
task, _ := svc.PostTask(ctx, ch.ID, "poster-agent", "Idempotent Test", "", nil, nil)
bid, _ := svc.BidOnTask(ctx, task.ID, "bidder-agent", nil, "", "")
svc.AcceptBid(ctx, task.ID, bid.ID, "poster-agent")
svc.CompleteTask(ctx, task.ID, "bidder-agent")
// Complete again should be idempotent
err := svc.CompleteTask(ctx, task.ID, "bidder-agent")
if err != nil {
t.Fatalf("expected idempotent completion, got: %v", err)
}
}
func TestSwarmService_CompleteTask_CannotCompleteOpenTask(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
task, _ := svc.PostTask(ctx, ch.ID, "poster-agent", "Open Complete Test", "", nil, nil)
err := svc.CompleteTask(ctx, task.ID, "poster-agent")
if err == nil {
t.Fatal("expected error when completing open task")
}
}
// --- Full auction lifecycle test ---
func TestSwarmService_FullAuctionLifecycle(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
// 1. Post task
deadline := time.Now().Add(30 * time.Minute)
task, err := svc.PostTask(ctx, ch.ID, "poster-agent",
"Translate Document",
"Translate document X into French",
json.RawMessage(`{"language":"french","domain":"legal"}`),
&deadline,
)
if err != nil {
t.Fatalf("PostTask: %v", err)
}
// 2. Two agents bid
bid1, err := svc.BidOnTask(ctx, task.ID, "bidder-agent",
json.RawMessage(`{"languages":["french"]}`),
"10m", "Fast translation")
if err != nil {
t.Fatalf("BidOnTask 1: %v", err)
}
bid2, err := svc.BidOnTask(ctx, task.ID, "bidder-agent-2",
json.RawMessage(`{"languages":["french","german"]}`),
"20m", "Thorough translation")
if err != nil {
t.Fatalf("BidOnTask 2: %v", err)
}
// 3. Poster accepts bid1
if err := svc.AcceptBid(ctx, task.ID, bid1.ID, "poster-agent"); err != nil {
t.Fatalf("AcceptBid: %v", err)
}
// Verify state
task, bids, err := svc.GetTaskWithBids(ctx, task.ID)
if err != nil {
t.Fatalf("GetTaskWithBids: %v", err)
}
if task.Status != TaskStatusAssigned {
t.Errorf("task status = %s, want assigned", task.Status)
}
if task.AssignedTo != "bidder-agent" {
t.Errorf("assigned_to = %s, want bidder-agent", task.AssignedTo)
}
for _, b := range bids {
switch b.ID {
case bid1.ID:
if b.Status != BidStatusAccepted {
t.Errorf("bid1 status = %s, want accepted", b.Status)
}
case bid2.ID:
if b.Status != BidStatusRejected {
t.Errorf("bid2 status = %s, want rejected", b.Status)
}
}
}
// 4. Assigned agent completes
if err := svc.CompleteTask(ctx, task.ID, "bidder-agent"); err != nil {
t.Fatalf("CompleteTask: %v", err)
}
task, _, _ = svc.GetTaskWithBids(ctx, task.ID)
if task.Status != TaskStatusCompleted {
t.Errorf("final task status = %s, want completed", task.Status)
}
}
// --- ListTasks and GetTaskWithBids tests ---
func TestSwarmService_ListTasks(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
svc.PostTask(ctx, ch.ID, "poster-agent", "Task 1", "", nil, nil)
svc.PostTask(ctx, ch.ID, "poster-agent", "Task 2", "", nil, nil)
tasks, err := svc.ListTasks(ctx, ch.ID, "")
if err != nil {
t.Fatalf("ListTasks: %v", err)
}
if len(tasks) != 2 {
t.Errorf("got %d tasks, want 2", len(tasks))
}
}
func TestSwarmService_GetTaskWithBids(t *testing.T) {
svc, channelStore := newTestSwarmService(t)
ch := createTestAuctionChannel(t, channelStore)
ctx := context.Background()
task, _ := svc.PostTask(ctx, ch.ID, "poster-agent", "With Bids", "", nil, nil)
svc.BidOnTask(ctx, task.ID, "bidder-agent", nil, "", "bid 1")
svc.BidOnTask(ctx, task.ID, "bidder-agent-2", nil, "", "bid 2")
gotTask, gotBids, err := svc.GetTaskWithBids(ctx, task.ID)
if err != nil {
t.Fatalf("GetTaskWithBids: %v", err)
}
if gotTask.Title != "With Bids" {
t.Errorf("title = %s, want With Bids", gotTask.Title)
}
if len(gotBids) != 2 {
t.Errorf("got %d bids, want 2", len(gotBids))
}
}
// --- ExpireTasks test ---
func TestSwarmService_ExpireTasks(t *testing.T) {
db := newTestDB(t)
seedAgent(t, db, "poster-agent")
seedAgent(t, db, "bidder-agent")
seedAgent(t, db, "bidder-agent-2")
seedAgent(t, db, "outsider-agent")
channelStore := NewSQLiteChannelStore(db)
taskStore := NewSQLiteTaskStore(db)
tracer := trace.NewTracer(db)
t.Cleanup(func() { tracer.Close() })
svc := NewSwarmService(taskStore, channelStore, tracer)
ctx := context.Background()
ch := &Channel{Name: "expire-svc-test", Type: TypeAuction, CreatedBy: "poster-agent"}
channelStore.CreateChannel(ctx, ch)
channelStore.AddMember(ctx, &Membership{ChannelID: ch.ID, AgentName: "poster-agent", Role: RoleOwner})
// Create task with past deadline directly via store (bypassing PostTask validation)
pastDeadline := time.Now().Add(-1 * time.Hour)
taskStore.CreateTask(ctx, &Task{
ChannelID: ch.ID,
PostedBy: "poster-agent",
Title: "Expired",
Status: TaskStatusOpen,
Deadline: &pastDeadline,
Requirements: json.RawMessage(`{}`),
})
count, err := svc.ExpireTasks(ctx)
if err != nil {
t.Fatalf("ExpireTasks: %v", err)
}
if count != 1 {
t.Errorf("expired count = %d, want 1", count)
}
}
+333
View File
@@ -0,0 +1,333 @@
package channels
import (
"context"
"database/sql"
"encoding/json"
"fmt"
"log/slog"
"time"
)
// TaskStore defines the storage interface for task auction operations.
type TaskStore interface {
CreateTask(ctx context.Context, task *Task) error
GetTask(ctx context.Context, id int64) (*Task, error)
ListTasks(ctx context.Context, channelID int64, status string) ([]*Task, error)
UpdateTaskStatus(ctx context.Context, id int64, status, assignedTo string) error
CreateBid(ctx context.Context, bid *Bid) error
GetBids(ctx context.Context, taskID int64) ([]*Bid, error)
GetBid(ctx context.Context, bidID int64) (*Bid, error)
UpdateBidStatus(ctx context.Context, bidID int64, status string) error
ExpireTasks(ctx context.Context) (int, error)
CancelTasksByChannel(ctx context.Context, channelID int64) (int, error)
}
// SQLiteTaskStore implements TaskStore using SQLite.
type SQLiteTaskStore struct {
db *sql.DB
logger *slog.Logger
}
// NewSQLiteTaskStore creates a new SQLite-backed task store.
func NewSQLiteTaskStore(db *sql.DB) *SQLiteTaskStore {
return &SQLiteTaskStore{
db: db,
logger: slog.Default().With("component", "task-store"),
}
}
// sqliteTimeFormat is the format used for storing timestamps consistently in SQLite.
const sqliteTimeFormat = "2006-01-02 15:04:05"
// CreateTask inserts a new task.
func (s *SQLiteTaskStore) CreateTask(ctx context.Context, task *Task) error {
requirements := string(task.Requirements)
if requirements == "" {
requirements = "{}"
}
// Format deadline as SQLite-compatible timestamp string
var deadlineStr *string
if task.Deadline != nil {
s := task.Deadline.UTC().Format(sqliteTimeFormat)
deadlineStr = &s
}
result, err := s.db.ExecContext(ctx,
`INSERT INTO tasks (channel_id, posted_by, title, description, requirements, deadline, status, assigned_to, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`,
task.ChannelID, task.PostedBy, task.Title, task.Description, requirements,
deadlineStr, task.Status, task.AssignedTo,
)
if err != nil {
return fmt.Errorf("insert task: %w", err)
}
id, err := result.LastInsertId()
if err != nil {
return fmt.Errorf("get task id: %w", err)
}
task.ID = id
s.logger.Info("task created", "id", id, "title", task.Title, "channel_id", task.ChannelID)
return nil
}
// GetTask returns a task by ID.
func (s *SQLiteTaskStore) GetTask(ctx context.Context, id int64) (*Task, error) {
var task Task
var requirements string
var deadlineStr sql.NullString
var assignedTo sql.NullString
err := s.db.QueryRowContext(ctx,
`SELECT id, channel_id, posted_by, title, description, requirements, deadline, status, assigned_to, created_at, updated_at
FROM tasks WHERE id = ?`, id,
).Scan(&task.ID, &task.ChannelID, &task.PostedBy, &task.Title, &task.Description,
&requirements, &deadlineStr, &task.Status, &assignedTo, &task.CreatedAt, &task.UpdatedAt)
if err != nil {
if err == sql.ErrNoRows {
return nil, fmt.Errorf("task not found: %d", id)
}
return nil, fmt.Errorf("get task: %w", err)
}
task.Requirements = json.RawMessage(requirements)
if deadlineStr.Valid {
t, err := time.Parse(sqliteTimeFormat, deadlineStr.String)
if err == nil {
task.Deadline = &t
}
}
if assignedTo.Valid {
task.AssignedTo = assignedTo.String
}
return &task, nil
}
// ListTasks returns tasks for a channel, optionally filtered by status.
func (s *SQLiteTaskStore) ListTasks(ctx context.Context, channelID int64, status string) ([]*Task, error) {
var query string
var args []any
if status != "" {
query = `SELECT id, channel_id, posted_by, title, description, requirements, deadline, status, assigned_to, created_at, updated_at
FROM tasks WHERE channel_id = ? AND status = ? ORDER BY created_at DESC`
args = []any{channelID, status}
} else {
query = `SELECT id, channel_id, posted_by, title, description, requirements, deadline, status, assigned_to, created_at, updated_at
FROM tasks WHERE channel_id = ? ORDER BY created_at DESC`
args = []any{channelID}
}
rows, err := s.db.QueryContext(ctx, query, args...)
if err != nil {
return nil, fmt.Errorf("list tasks: %w", err)
}
defer rows.Close()
return scanTasks(rows)
}
// UpdateTaskStatus updates a task's status and optionally assigned_to.
func (s *SQLiteTaskStore) UpdateTaskStatus(ctx context.Context, id int64, status, assignedTo string) error {
var result sql.Result
var err error
if assignedTo != "" {
result, err = s.db.ExecContext(ctx,
`UPDATE tasks SET status = ?, assigned_to = ?, updated_at = CURRENT_TIMESTAMP WHERE id = ?`,
status, assignedTo, id)
} else {
result, err = s.db.ExecContext(ctx,
`UPDATE tasks SET status = ?, updated_at = CURRENT_TIMESTAMP WHERE id = ?`,
status, id)
}
if err != nil {
return fmt.Errorf("update task status: %w", err)
}
rowsAffected, _ := result.RowsAffected()
if rowsAffected == 0 {
return fmt.Errorf("task not found: %d", id)
}
return nil
}
// CreateBid inserts a new bid on a task.
func (s *SQLiteTaskStore) CreateBid(ctx context.Context, bid *Bid) error {
capabilities := string(bid.Capabilities)
if capabilities == "" {
capabilities = "{}"
}
result, err := s.db.ExecContext(ctx,
`INSERT INTO task_bids (task_id, agent_name, capabilities, time_estimate, message, status, created_at)
VALUES (?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP)`,
bid.TaskID, bid.AgentName, capabilities, bid.TimeEstimate, bid.Message, BidStatusPending,
)
if err != nil {
if isUniqueConstraintError(err) {
return fmt.Errorf("agent %s has already bid on task %d", bid.AgentName, bid.TaskID)
}
return fmt.Errorf("insert bid: %w", err)
}
id, err := result.LastInsertId()
if err != nil {
return fmt.Errorf("get bid id: %w", err)
}
bid.ID = id
bid.Status = BidStatusPending
s.logger.Info("bid created", "id", id, "task_id", bid.TaskID, "agent", bid.AgentName)
return nil
}
// GetBids returns all bids for a task.
func (s *SQLiteTaskStore) GetBids(ctx context.Context, taskID int64) ([]*Bid, error) {
rows, err := s.db.QueryContext(ctx,
`SELECT id, task_id, agent_name, capabilities, time_estimate, message, status, created_at
FROM task_bids WHERE task_id = ? ORDER BY created_at ASC`,
taskID,
)
if err != nil {
return nil, fmt.Errorf("get bids: %w", err)
}
defer rows.Close()
return scanBids(rows)
}
// GetBid returns a bid by ID.
func (s *SQLiteTaskStore) GetBid(ctx context.Context, bidID int64) (*Bid, error) {
var bid Bid
var capabilities string
var timeEstimate sql.NullString
err := s.db.QueryRowContext(ctx,
`SELECT id, task_id, agent_name, capabilities, time_estimate, message, status, created_at
FROM task_bids WHERE id = ?`, bidID,
).Scan(&bid.ID, &bid.TaskID, &bid.AgentName, &capabilities, &timeEstimate, &bid.Message, &bid.Status, &bid.CreatedAt)
if err != nil {
if err == sql.ErrNoRows {
return nil, fmt.Errorf("bid not found: %d", bidID)
}
return nil, fmt.Errorf("get bid: %w", err)
}
bid.Capabilities = json.RawMessage(capabilities)
if timeEstimate.Valid {
bid.TimeEstimate = timeEstimate.String
}
return &bid, nil
}
// UpdateBidStatus updates a bid's status.
func (s *SQLiteTaskStore) UpdateBidStatus(ctx context.Context, bidID int64, status string) error {
result, err := s.db.ExecContext(ctx,
`UPDATE task_bids SET status = ? WHERE id = ?`,
status, bidID,
)
if err != nil {
return fmt.Errorf("update bid status: %w", err)
}
rowsAffected, _ := result.RowsAffected()
if rowsAffected == 0 {
return fmt.Errorf("bid not found: %d", bidID)
}
return nil
}
// ExpireTasks marks all open tasks past their deadline as cancelled.
func (s *SQLiteTaskStore) ExpireTasks(ctx context.Context) (int, error) {
// Use a string-formatted timestamp for consistent SQLite comparison
now := time.Now().UTC().Format("2006-01-02 15:04:05")
result, err := s.db.ExecContext(ctx,
`UPDATE tasks SET status = ?, updated_at = CURRENT_TIMESTAMP
WHERE status = ? AND deadline IS NOT NULL AND deadline < ?`,
TaskStatusCancelled, TaskStatusOpen, now,
)
if err != nil {
return 0, fmt.Errorf("expire tasks: %w", err)
}
rowsAffected, _ := result.RowsAffected()
return int(rowsAffected), nil
}
// CancelTasksByChannel cancels all open tasks for a channel (used before channel deletion).
func (s *SQLiteTaskStore) CancelTasksByChannel(ctx context.Context, channelID int64) (int, error) {
result, err := s.db.ExecContext(ctx,
`UPDATE tasks SET status = ?, updated_at = CURRENT_TIMESTAMP
WHERE channel_id = ? AND status = ?`,
TaskStatusCancelled, channelID, TaskStatusOpen,
)
if err != nil {
return 0, fmt.Errorf("cancel tasks by channel: %w", err)
}
rowsAffected, _ := result.RowsAffected()
return int(rowsAffected), nil
}
// scanTasks scans multiple task rows.
func scanTasks(rows *sql.Rows) ([]*Task, error) {
var tasks []*Task
for rows.Next() {
var task Task
var requirements string
var deadlineStr sql.NullString
var assignedTo sql.NullString
err := rows.Scan(&task.ID, &task.ChannelID, &task.PostedBy, &task.Title, &task.Description,
&requirements, &deadlineStr, &task.Status, &assignedTo, &task.CreatedAt, &task.UpdatedAt)
if err != nil {
return nil, fmt.Errorf("scan task: %w", err)
}
task.Requirements = json.RawMessage(requirements)
if deadlineStr.Valid {
t, err := time.Parse(sqliteTimeFormat, deadlineStr.String)
if err == nil {
task.Deadline = &t
}
}
if assignedTo.Valid {
task.AssignedTo = assignedTo.String
}
tasks = append(tasks, &task)
}
if tasks == nil {
tasks = []*Task{}
}
return tasks, rows.Err()
}
// scanBids scans multiple bid rows.
func scanBids(rows *sql.Rows) ([]*Bid, error) {
var bids []*Bid
for rows.Next() {
var bid Bid
var capabilities string
var timeEstimate sql.NullString
err := rows.Scan(&bid.ID, &bid.TaskID, &bid.AgentName, &capabilities, &timeEstimate,
&bid.Message, &bid.Status, &bid.CreatedAt)
if err != nil {
return nil, fmt.Errorf("scan bid: %w", err)
}
bid.Capabilities = json.RawMessage(capabilities)
if timeEstimate.Valid {
bid.TimeEstimate = timeEstimate.String
}
bids = append(bids, &bid)
}
if bids == nil {
bids = []*Bid{}
}
return bids, rows.Err()
}
+349
View File
@@ -0,0 +1,349 @@
package channels
import (
"context"
"encoding/json"
"testing"
"time"
)
func newTestTaskStore(t *testing.T) (*SQLiteTaskStore, *SQLiteChannelStore) {
t.Helper()
db := newTestDB(t)
seedAgent(t, db, "poster-agent")
seedAgent(t, db, "bidder-agent")
seedAgent(t, db, "bidder-agent-2")
channelStore := NewSQLiteChannelStore(db)
taskStore := NewSQLiteTaskStore(db)
return taskStore, channelStore
}
func createAuctionChannel(t *testing.T, channelStore *SQLiteChannelStore) *Channel {
t.Helper()
ctx := context.Background()
ch := &Channel{
Name: "auction-ch",
Type: TypeAuction,
CreatedBy: "poster-agent",
}
if err := channelStore.CreateChannel(ctx, ch); err != nil {
t.Fatalf("create auction channel: %v", err)
}
channelStore.AddMember(ctx, &Membership{ChannelID: ch.ID, AgentName: "poster-agent", Role: RoleOwner})
channelStore.AddMember(ctx, &Membership{ChannelID: ch.ID, AgentName: "bidder-agent", Role: RoleMember})
channelStore.AddMember(ctx, &Membership{ChannelID: ch.ID, AgentName: "bidder-agent-2", Role: RoleMember})
return ch
}
func TestSQLiteTaskStore_CreateTask(t *testing.T) {
taskStore, channelStore := newTestTaskStore(t)
ch := createAuctionChannel(t, channelStore)
ctx := context.Background()
deadline := time.Now().Add(1 * time.Hour)
task := &Task{
ChannelID: ch.ID,
PostedBy: "poster-agent",
Title: "Test Task",
Description: "A test task",
Requirements: json.RawMessage(`{"skill":"go"}`),
Deadline: &deadline,
Status: TaskStatusOpen,
}
if err := taskStore.CreateTask(ctx, task); err != nil {
t.Fatalf("CreateTask: %v", err)
}
if task.ID == 0 {
t.Error("task ID should not be 0")
}
}
func TestSQLiteTaskStore_GetTask(t *testing.T) {
taskStore, channelStore := newTestTaskStore(t)
ch := createAuctionChannel(t, channelStore)
ctx := context.Background()
task := &Task{
ChannelID: ch.ID,
PostedBy: "poster-agent",
Title: "Get Test",
Description: "Description",
Requirements: json.RawMessage(`{}`),
Status: TaskStatusOpen,
}
taskStore.CreateTask(ctx, task)
got, err := taskStore.GetTask(ctx, task.ID)
if err != nil {
t.Fatalf("GetTask: %v", err)
}
if got.Title != "Get Test" {
t.Errorf("title = %s, want Get Test", got.Title)
}
if got.Status != TaskStatusOpen {
t.Errorf("status = %s, want open", got.Status)
}
}
func TestSQLiteTaskStore_GetTask_NotFound(t *testing.T) {
taskStore, _ := newTestTaskStore(t)
ctx := context.Background()
_, err := taskStore.GetTask(ctx, 99999)
if err == nil {
t.Error("expected error for non-existent task")
}
}
func TestSQLiteTaskStore_ListTasks(t *testing.T) {
taskStore, channelStore := newTestTaskStore(t)
ch := createAuctionChannel(t, channelStore)
ctx := context.Background()
// Create multiple tasks
taskStore.CreateTask(ctx, &Task{ChannelID: ch.ID, PostedBy: "poster-agent", Title: "Task 1", Status: TaskStatusOpen, Requirements: json.RawMessage(`{}`)})
taskStore.CreateTask(ctx, &Task{ChannelID: ch.ID, PostedBy: "poster-agent", Title: "Task 2", Status: TaskStatusOpen, Requirements: json.RawMessage(`{}`)})
t.Run("list all tasks", func(t *testing.T) {
tasks, err := taskStore.ListTasks(ctx, ch.ID, "")
if err != nil {
t.Fatalf("ListTasks: %v", err)
}
if len(tasks) != 2 {
t.Errorf("got %d tasks, want 2", len(tasks))
}
})
t.Run("filter by status", func(t *testing.T) {
tasks, err := taskStore.ListTasks(ctx, ch.ID, TaskStatusOpen)
if err != nil {
t.Fatalf("ListTasks: %v", err)
}
if len(tasks) != 2 {
t.Errorf("got %d tasks, want 2", len(tasks))
}
})
t.Run("filter by non-matching status", func(t *testing.T) {
tasks, err := taskStore.ListTasks(ctx, ch.ID, TaskStatusCompleted)
if err != nil {
t.Fatalf("ListTasks: %v", err)
}
if len(tasks) != 0 {
t.Errorf("got %d tasks, want 0", len(tasks))
}
})
}
func TestSQLiteTaskStore_UpdateTaskStatus(t *testing.T) {
taskStore, channelStore := newTestTaskStore(t)
ch := createAuctionChannel(t, channelStore)
ctx := context.Background()
task := &Task{ChannelID: ch.ID, PostedBy: "poster-agent", Title: "Update Test", Status: TaskStatusOpen, Requirements: json.RawMessage(`{}`)}
taskStore.CreateTask(ctx, task)
t.Run("update to assigned with assigned_to", func(t *testing.T) {
err := taskStore.UpdateTaskStatus(ctx, task.ID, TaskStatusAssigned, "bidder-agent")
if err != nil {
t.Fatalf("UpdateTaskStatus: %v", err)
}
got, _ := taskStore.GetTask(ctx, task.ID)
if got.Status != TaskStatusAssigned {
t.Errorf("status = %s, want assigned", got.Status)
}
if got.AssignedTo != "bidder-agent" {
t.Errorf("assigned_to = %s, want bidder-agent", got.AssignedTo)
}
})
t.Run("update non-existent task", func(t *testing.T) {
err := taskStore.UpdateTaskStatus(ctx, 99999, TaskStatusCompleted, "")
if err == nil {
t.Error("expected error for non-existent task")
}
})
}
func TestSQLiteTaskStore_CreateBid(t *testing.T) {
taskStore, channelStore := newTestTaskStore(t)
ch := createAuctionChannel(t, channelStore)
ctx := context.Background()
task := &Task{ChannelID: ch.ID, PostedBy: "poster-agent", Title: "Bid Test", Status: TaskStatusOpen, Requirements: json.RawMessage(`{}`)}
taskStore.CreateTask(ctx, task)
bid := &Bid{
TaskID: task.ID,
AgentName: "bidder-agent",
Capabilities: json.RawMessage(`{"lang":"go"}`),
TimeEstimate: "30m",
Message: "I can do this",
}
if err := taskStore.CreateBid(ctx, bid); err != nil {
t.Fatalf("CreateBid: %v", err)
}
if bid.ID == 0 {
t.Error("bid ID should not be 0")
}
if bid.Status != BidStatusPending {
t.Errorf("status = %s, want pending", bid.Status)
}
}
func TestSQLiteTaskStore_CreateBid_DuplicateRejected(t *testing.T) {
taskStore, channelStore := newTestTaskStore(t)
ch := createAuctionChannel(t, channelStore)
ctx := context.Background()
task := &Task{ChannelID: ch.ID, PostedBy: "poster-agent", Title: "Dup Bid Test", Status: TaskStatusOpen, Requirements: json.RawMessage(`{}`)}
taskStore.CreateTask(ctx, task)
bid1 := &Bid{TaskID: task.ID, AgentName: "bidder-agent", Capabilities: json.RawMessage(`{}`)}
taskStore.CreateBid(ctx, bid1)
bid2 := &Bid{TaskID: task.ID, AgentName: "bidder-agent", Capabilities: json.RawMessage(`{}`)}
err := taskStore.CreateBid(ctx, bid2)
if err == nil {
t.Error("expected error for duplicate bid from same agent")
}
}
func TestSQLiteTaskStore_GetBids(t *testing.T) {
taskStore, channelStore := newTestTaskStore(t)
ch := createAuctionChannel(t, channelStore)
ctx := context.Background()
task := &Task{ChannelID: ch.ID, PostedBy: "poster-agent", Title: "Bids Test", Status: TaskStatusOpen, Requirements: json.RawMessage(`{}`)}
taskStore.CreateTask(ctx, task)
taskStore.CreateBid(ctx, &Bid{TaskID: task.ID, AgentName: "bidder-agent", Capabilities: json.RawMessage(`{}`), Message: "bid 1"})
taskStore.CreateBid(ctx, &Bid{TaskID: task.ID, AgentName: "bidder-agent-2", Capabilities: json.RawMessage(`{}`), Message: "bid 2"})
bids, err := taskStore.GetBids(ctx, task.ID)
if err != nil {
t.Fatalf("GetBids: %v", err)
}
if len(bids) != 2 {
t.Errorf("got %d bids, want 2", len(bids))
}
}
func TestSQLiteTaskStore_GetBid(t *testing.T) {
taskStore, channelStore := newTestTaskStore(t)
ch := createAuctionChannel(t, channelStore)
ctx := context.Background()
task := &Task{ChannelID: ch.ID, PostedBy: "poster-agent", Title: "Get Bid Test", Status: TaskStatusOpen, Requirements: json.RawMessage(`{}`)}
taskStore.CreateTask(ctx, task)
bid := &Bid{TaskID: task.ID, AgentName: "bidder-agent", Capabilities: json.RawMessage(`{"x":1}`), Message: "my bid"}
taskStore.CreateBid(ctx, bid)
got, err := taskStore.GetBid(ctx, bid.ID)
if err != nil {
t.Fatalf("GetBid: %v", err)
}
if got.AgentName != "bidder-agent" {
t.Errorf("agent_name = %s, want bidder-agent", got.AgentName)
}
if got.Message != "my bid" {
t.Errorf("message = %s, want my bid", got.Message)
}
}
func TestSQLiteTaskStore_UpdateBidStatus(t *testing.T) {
taskStore, channelStore := newTestTaskStore(t)
ch := createAuctionChannel(t, channelStore)
ctx := context.Background()
task := &Task{ChannelID: ch.ID, PostedBy: "poster-agent", Title: "Bid Status Test", Status: TaskStatusOpen, Requirements: json.RawMessage(`{}`)}
taskStore.CreateTask(ctx, task)
bid := &Bid{TaskID: task.ID, AgentName: "bidder-agent", Capabilities: json.RawMessage(`{}`)}
taskStore.CreateBid(ctx, bid)
if err := taskStore.UpdateBidStatus(ctx, bid.ID, BidStatusAccepted); err != nil {
t.Fatalf("UpdateBidStatus: %v", err)
}
got, _ := taskStore.GetBid(ctx, bid.ID)
if got.Status != BidStatusAccepted {
t.Errorf("status = %s, want accepted", got.Status)
}
}
func TestSQLiteTaskStore_ExpireTasks(t *testing.T) {
taskStore, channelStore := newTestTaskStore(t)
ch := createAuctionChannel(t, channelStore)
ctx := context.Background()
// Create a task with deadline in the past
pastDeadline := time.Now().Add(-1 * time.Hour)
taskStore.CreateTask(ctx, &Task{
ChannelID: ch.ID, PostedBy: "poster-agent", Title: "Expired Task",
Status: TaskStatusOpen, Deadline: &pastDeadline, Requirements: json.RawMessage(`{}`),
})
// Create a task with deadline in the future
futureDeadline := time.Now().Add(1 * time.Hour)
taskStore.CreateTask(ctx, &Task{
ChannelID: ch.ID, PostedBy: "poster-agent", Title: "Future Task",
Status: TaskStatusOpen, Deadline: &futureDeadline, Requirements: json.RawMessage(`{}`),
})
// Create a task without deadline
taskStore.CreateTask(ctx, &Task{
ChannelID: ch.ID, PostedBy: "poster-agent", Title: "No Deadline Task",
Status: TaskStatusOpen, Requirements: json.RawMessage(`{}`),
})
count, err := taskStore.ExpireTasks(ctx)
if err != nil {
t.Fatalf("ExpireTasks: %v", err)
}
if count != 1 {
t.Errorf("expired count = %d, want 1", count)
}
// Verify the expired task is cancelled
tasks, _ := taskStore.ListTasks(ctx, ch.ID, TaskStatusCancelled)
if len(tasks) != 1 {
t.Errorf("cancelled tasks = %d, want 1", len(tasks))
}
if tasks[0].Title != "Expired Task" {
t.Errorf("cancelled task title = %s, want Expired Task", tasks[0].Title)
}
// Verify the other tasks are still open
openTasks, _ := taskStore.ListTasks(ctx, ch.ID, TaskStatusOpen)
if len(openTasks) != 2 {
t.Errorf("open tasks = %d, want 2", len(openTasks))
}
}
func TestSQLiteTaskStore_CancelTasksByChannel(t *testing.T) {
taskStore, channelStore := newTestTaskStore(t)
ch := createAuctionChannel(t, channelStore)
ctx := context.Background()
taskStore.CreateTask(ctx, &Task{ChannelID: ch.ID, PostedBy: "poster-agent", Title: "Task A", Status: TaskStatusOpen, Requirements: json.RawMessage(`{}`)})
taskStore.CreateTask(ctx, &Task{ChannelID: ch.ID, PostedBy: "poster-agent", Title: "Task B", Status: TaskStatusOpen, Requirements: json.RawMessage(`{}`)})
count, err := taskStore.CancelTasksByChannel(ctx, ch.ID)
if err != nil {
t.Fatalf("CancelTasksByChannel: %v", err)
}
if count != 2 {
t.Errorf("cancelled count = %d, want 2", count)
}
tasks, _ := taskStore.ListTasks(ctx, ch.ID, TaskStatusOpen)
if len(tasks) != 0 {
t.Errorf("open tasks = %d, want 0", len(tasks))
}
}
+66
View File
@@ -0,0 +1,66 @@
package channels
import (
"encoding/json"
"time"
)
// TaskStatus constants for task lifecycle.
const (
TaskStatusOpen = "open"
TaskStatusAssigned = "assigned"
TaskStatusCompleted = "completed"
TaskStatusCancelled = "cancelled"
)
// BidStatus constants for bid lifecycle.
const (
BidStatusPending = "pending"
BidStatusAccepted = "accepted"
BidStatusRejected = "rejected"
)
// Task represents a unit of work posted to an auction channel.
type Task struct {
ID int64 `json:"id"`
ChannelID int64 `json:"channel_id"`
PostedBy string `json:"posted_by"`
Title string `json:"title"`
Description string `json:"description"`
Requirements json.RawMessage `json:"requirements"`
Deadline *time.Time `json:"deadline,omitempty"`
Status string `json:"status"`
AssignedTo string `json:"assigned_to,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// Bid represents an agent's offer to complete a task.
type Bid struct {
ID int64 `json:"id"`
TaskID int64 `json:"task_id"`
AgentName string `json:"agent_name"`
Capabilities json.RawMessage `json:"capabilities"`
TimeEstimate string `json:"time_estimate"`
Message string `json:"message"`
Status string `json:"status"`
CreatedAt time.Time `json:"created_at"`
}
// ValidTaskStatus returns true if the given status is a valid task status.
func ValidTaskStatus(s string) bool {
switch s {
case TaskStatusOpen, TaskStatusAssigned, TaskStatusCompleted, TaskStatusCancelled:
return true
}
return false
}
// ValidBidStatus returns true if the given status is a valid bid status.
func ValidBidStatus(s string) bool {
switch s {
case BidStatusPending, BidStatusAccepted, BidStatusRejected:
return true
}
return false
}
+7
View File
@@ -26,6 +26,7 @@ func NewMCPServer(
msgService *messaging.MessagingService,
agentService *agents.AgentService,
channelService *channels.Service,
swarmService ...*channels.SwarmService,
) *MCPServer {
logger := slog.Default().With("component", "mcp-server")
@@ -46,6 +47,12 @@ func NewMCPServer(
channelRegistrar.RegisterAll(mcpSrv)
}
// Register swarm tools
if len(swarmService) > 0 && swarmService[0] != nil && channelService != nil {
swarmRegistrar := NewSwarmToolRegistrar(swarmService[0], channelService)
swarmRegistrar.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 {
+285
View File
@@ -0,0 +1,285 @@
package mcp
import (
"context"
"encoding/json"
"fmt"
"log/slog"
"time"
"github.com/mark3labs/mcp-go/mcp"
"github.com/mark3labs/mcp-go/server"
"github.com/smart-mcp-proxy/synapbus/internal/channels"
)
// SwarmToolRegistrar registers swarm-pattern MCP tools on the server.
type SwarmToolRegistrar struct {
swarmService *channels.SwarmService
channelService *channels.Service
logger *slog.Logger
}
// NewSwarmToolRegistrar creates a new swarm tool registrar.
func NewSwarmToolRegistrar(swarmService *channels.SwarmService, channelService *channels.Service) *SwarmToolRegistrar {
return &SwarmToolRegistrar{
swarmService: swarmService,
channelService: channelService,
logger: slog.Default().With("component", "mcp-swarm-tools"),
}
}
// RegisterAll registers all swarm tools on the MCP server.
func (str *SwarmToolRegistrar) RegisterAll(s *server.MCPServer) {
s.AddTool(str.postTaskTool(), str.handlePostTask)
s.AddTool(str.bidTaskTool(), str.handleBidTask)
s.AddTool(str.acceptBidTool(), str.handleAcceptBid)
s.AddTool(str.completeTaskTool(), str.handleCompleteTask)
s.AddTool(str.listTasksTool(), str.handleListTasks)
str.logger.Info("swarm MCP tools registered", "count", 5)
}
// --- Tool Definitions ---
func (str *SwarmToolRegistrar) postTaskTool() mcp.Tool {
return mcp.NewTool("post_task",
mcp.WithDescription("Post a task to an auction channel for agents to bid on"),
mcp.WithString("channel_name", mcp.Description("Name of the auction channel"), mcp.Required()),
mcp.WithString("title", mcp.Description("Task title"), mcp.Required()),
mcp.WithString("description", mcp.Description("Task description")),
mcp.WithString("requirements", mcp.Description("JSON object of task requirements")),
mcp.WithString("deadline", mcp.Description("Task deadline in ISO 8601 format (e.g. 2026-03-13T15:00:00Z)")),
)
}
func (str *SwarmToolRegistrar) bidTaskTool() mcp.Tool {
return mcp.NewTool("bid_task",
mcp.WithDescription("Submit a bid on an open task in an auction channel"),
mcp.WithNumber("task_id", mcp.Description("ID of the task to bid on"), mcp.Required()),
mcp.WithString("capabilities", mcp.Description("JSON object describing your relevant capabilities")),
mcp.WithString("time_estimate", mcp.Description("Estimated time to complete the task")),
mcp.WithString("message", mcp.Description("Message to the task poster explaining your bid")),
)
}
func (str *SwarmToolRegistrar) acceptBidTool() mcp.Tool {
return mcp.NewTool("accept_bid",
mcp.WithDescription("Accept a bid on a task you posted, assigning the task to the bidding agent"),
mcp.WithNumber("task_id", mcp.Description("ID of the task"), mcp.Required()),
mcp.WithNumber("bid_id", mcp.Description("ID of the bid to accept"), mcp.Required()),
)
}
func (str *SwarmToolRegistrar) completeTaskTool() mcp.Tool {
return mcp.NewTool("complete_task",
mcp.WithDescription("Mark a task as completed (only the assigned agent can do this)"),
mcp.WithNumber("task_id", mcp.Description("ID of the task to complete"), mcp.Required()),
)
}
func (str *SwarmToolRegistrar) listTasksTool() mcp.Tool {
return mcp.NewTool("list_tasks",
mcp.WithDescription("List tasks in an auction channel, optionally filtered by status"),
mcp.WithString("channel_name", mcp.Description("Name of the auction channel"), mcp.Required()),
mcp.WithString("status", mcp.Description("Filter by task status: open, assigned, completed, cancelled")),
)
}
// --- Tool Handlers ---
func (str *SwarmToolRegistrar) handlePostTask(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
agentName, ok := extractAgentName(ctx)
if !ok {
return mcp.NewToolResultError("authentication required"), nil
}
channelName := req.GetString("channel_name", "")
if channelName == "" {
return mcp.NewToolResultError("'channel_name' parameter is required"), nil
}
title := req.GetString("title", "")
if title == "" {
return mcp.NewToolResultError("'title' parameter is required"), nil
}
description := req.GetString("description", "")
requirementsStr := req.GetString("requirements", "{}")
deadlineStr := req.GetString("deadline", "")
// Parse requirements JSON
var requirements json.RawMessage
if requirementsStr != "" {
if !json.Valid([]byte(requirementsStr)) {
return mcp.NewToolResultError("requirements must be valid JSON"), nil
}
requirements = json.RawMessage(requirementsStr)
}
// Parse deadline
var deadline *time.Time
if deadlineStr != "" {
t, err := time.Parse(time.RFC3339, deadlineStr)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("deadline must be ISO 8601 format: %s", err)), nil
}
deadline = &t
}
// Resolve channel
ch, err := str.channelService.GetChannelByName(ctx, channelName)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("post_task failed: %s", err)), nil
}
task, err := str.swarmService.PostTask(ctx, ch.ID, agentName, title, description, requirements, deadline)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("post_task failed: %s", err)), nil
}
return resultJSON(map[string]any{
"task_id": task.ID,
"channel_id": task.ChannelID,
"title": task.Title,
"status": task.Status,
"posted_by": task.PostedBy,
"deadline": task.Deadline,
"created_at": task.CreatedAt,
})
}
func (str *SwarmToolRegistrar) handleBidTask(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
agentName, ok := extractAgentName(ctx)
if !ok {
return mcp.NewToolResultError("authentication required"), nil
}
taskID, err := req.RequireInt("task_id")
if err != nil {
return mcp.NewToolResultError("'task_id' parameter is required"), nil
}
capabilitiesStr := req.GetString("capabilities", "{}")
timeEstimate := req.GetString("time_estimate", "")
message := req.GetString("message", "")
var capabilities json.RawMessage
if capabilitiesStr != "" {
if !json.Valid([]byte(capabilitiesStr)) {
return mcp.NewToolResultError("capabilities must be valid JSON"), nil
}
capabilities = json.RawMessage(capabilitiesStr)
}
bid, err := str.swarmService.BidOnTask(ctx, int64(taskID), agentName, capabilities, timeEstimate, message)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("bid_task failed: %s", err)), nil
}
return resultJSON(map[string]any{
"bid_id": bid.ID,
"task_id": bid.TaskID,
"agent_name": bid.AgentName,
"time_estimate": bid.TimeEstimate,
"status": bid.Status,
})
}
func (str *SwarmToolRegistrar) handleAcceptBid(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
agentName, ok := extractAgentName(ctx)
if !ok {
return mcp.NewToolResultError("authentication required"), nil
}
taskID, err := req.RequireInt("task_id")
if err != nil {
return mcp.NewToolResultError("'task_id' parameter is required"), nil
}
bidID, err := req.RequireInt("bid_id")
if err != nil {
return mcp.NewToolResultError("'bid_id' parameter is required"), nil
}
if err := str.swarmService.AcceptBid(ctx, int64(taskID), int64(bidID), agentName); err != nil {
return mcp.NewToolResultError(fmt.Sprintf("accept_bid failed: %s", err)), nil
}
return resultJSON(map[string]any{
"task_id": taskID,
"bid_id": bidID,
"status": "accepted",
})
}
func (str *SwarmToolRegistrar) handleCompleteTask(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
agentName, ok := extractAgentName(ctx)
if !ok {
return mcp.NewToolResultError("authentication required"), nil
}
taskID, err := req.RequireInt("task_id")
if err != nil {
return mcp.NewToolResultError("'task_id' parameter is required"), nil
}
if err := str.swarmService.CompleteTask(ctx, int64(taskID), agentName); err != nil {
return mcp.NewToolResultError(fmt.Sprintf("complete_task failed: %s", err)), nil
}
return resultJSON(map[string]any{
"task_id": taskID,
"status": "completed",
})
}
func (str *SwarmToolRegistrar) handleListTasks(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
agentName, ok := extractAgentName(ctx)
if !ok {
return mcp.NewToolResultError("authentication required"), nil
}
_ = agentName // just verifying auth
channelName := req.GetString("channel_name", "")
if channelName == "" {
return mcp.NewToolResultError("'channel_name' parameter is required"), nil
}
statusFilter := req.GetString("status", "")
// Resolve channel
ch, err := str.channelService.GetChannelByName(ctx, channelName)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("list_tasks failed: %s", err)), nil
}
// Verify channel is auction type
if ch.Type != channels.TypeAuction {
return mcp.NewToolResultError(fmt.Sprintf("list_tasks requires a channel of type 'auction', got '%s'", ch.Type)), nil
}
tasks, err := str.swarmService.ListTasks(ctx, ch.ID, statusFilter)
if err != nil {
return mcp.NewToolResultError(fmt.Sprintf("list_tasks failed: %s", err)), nil
}
result := make([]map[string]any, len(tasks))
for i, task := range tasks {
result[i] = map[string]any{
"id": task.ID,
"title": task.Title,
"description": task.Description,
"status": task.Status,
"posted_by": task.PostedBy,
"assigned_to": task.AssignedTo,
"deadline": task.Deadline,
"created_at": task.CreatedAt,
}
}
return resultJSON(map[string]any{
"tasks": result,
"count": len(result),
})
}