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:
co-authored by
Claude Opus 4.6
parent
1cb0bb7a8f
commit
7909450473
+14
-2
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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),
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user